Team Ai
Modelpublic

mazesmazes/tiny-audio

sourceHugging Facemitupdated 1d agoView on Hugging Face
1likes177downloads
asr_processing.py381 linesDownload Raw Back to root
1"""Processor that turns raw audio (and optional text) into model inputs."""2 3from collections.abc import Mapping, Sequence4from typing import TYPE_CHECKING, Any, ClassVar, cast, overload5 6import numpy as np7import numpy.typing as npt8import torch9import transformers10from torch.nn.utils.rnn import pad_sequence11from transformers import (12    BatchFeature,13    PreTrainedTokenizerBase,14    ProcessorMixin,15    SequenceFeatureExtractor,16)17 18if TYPE_CHECKING:19    from .asr_config import (20        DEFAULT_ENCODER_CONV_LAYERS,21        ASRConfig,22        ConvLayerSpec,23        compute_encoder_output_length,24    )25    from .asr_types import AudioFeatureExtractor, AudioInput, PreparedChunk, Waveform26    from .projectors import MLPAudioProjector27else:28    try:29        from .asr_config import (30            DEFAULT_ENCODER_CONV_LAYERS,31            ASRConfig,32            ConvLayerSpec,33            compute_encoder_output_length,34        )35        from .asr_types import AudioInput, PreparedChunk36    except ImportError:  # flat layout on the Hub: sibling modules, no package37        from asr_config import (38            DEFAULT_ENCODER_CONV_LAYERS,39            ASRConfig,40            ConvLayerSpec,41            compute_encoder_output_length,42        )43        from asr_types import AudioInput, PreparedChunk44 45 46def collate_chunks(prepared: Sequence[PreparedChunk]) -> PreparedChunk:47    """Pad prepared chunks to the longest and stack them into one batch.48 49    The time axis is whichever feature axis matches the mask's length (Granite's50    features are `(1, T, D)`, Whisper's `(1, D, T)`); padded frames are zeros51    with a 0 in the mask, which the encoder honours (`encoder_attention_mask`).52    """53    longest = max(int(p["attention_mask"].shape[-1]) for p in prepared)54    features: list[torch.Tensor] = []55    masks: list[torch.Tensor] = []56    for p in prepared:57        feats, mask = p["input_features"], p["attention_mask"]58        length = int(mask.shape[-1])59        time_axis = 1 if feats.shape[1] == length else feats.dim() - 160        pad = longest - length61        # F.pad lists (left, right) pairs from the LAST axis backwards.62        spec = [0, 0] * (feats.dim() - 1 - time_axis) + [0, pad]63        features.append(torch.nn.functional.pad(feats, spec))64        masks.append(torch.nn.functional.pad(mask, (0, pad)))65    return {"input_features": torch.cat(features), "attention_mask": torch.cat(masks)}66 67 68# The instruction the model trained on (scripts/train_collator.py); the model69# and processor both default to it.70DEFAULT_TRANSCRIBE_PROMPT = "Transcribe the speech to text"71 72 73def render_audio_prompt(74    tokenizer: PreTrainedTokenizerBase,75    audio_token: str,76    num_audio_tokens: int,77    prompt: str | None,78    text: str | None = None,79) -> torch.Tensor:80    """Tokenize one chat prompt carrying exactly `num_audio_tokens` placeholders.81 82    The user turn is the placeholders, then `prompt` (if any); `text`, when83    given, is the assistant's reply, otherwise the generation prompt is added.84    """85    if num_audio_tokens > 0:86        user_content = audio_token * num_audio_tokens87        if prompt:88            user_content += " " + prompt89    else:90        user_content = prompt or ""91 92    messages = [{"role": "user", "content": user_content}]93    if text is not None:94        messages.append({"role": "assistant", "content": text})95 96    # With `tokenize=True, return_tensors="pt"` the ids come back as tensors.97    tokenized = cast(98        "torch.Tensor | Mapping[str, torch.Tensor]",99        tokenizer.apply_chat_template(100            messages,101            tokenize=True,102            add_generation_prompt=(text is None),103            return_tensors="pt",104            enable_thinking=False,  # Disable Qwen3 thinking mode for ASR105        ),106    )107 108    # apply_chat_template returns a bare tensor or a BatchEncoding/mapping.109    ids = tokenized if isinstance(tokenized, torch.Tensor) else tokenized["input_ids"]110    return (ids[0] if ids.dim() > 1 else ids).to(torch.long)111 112 113def left_pad_prompt_rows(114    rows: list[torch.Tensor], tokenizer: PreTrainedTokenizerBase115) -> tuple[torch.Tensor, torch.Tensor]:116    """Stack per-sample prompt rows into a left-padded batch: `(input_ids, attention_mask)`.117 118    Left, not right: these feed `generate`, so padding must not sit between119    the prompt and the first generated token. Pads with the tokenizer's pad120    token, falling back to eos, then 0. Pad positions never carry121    `audio_token_id`, so the model's masked_scatter is unaffected.122    """123    # transformers types special-token ids as any token value; a single id is an int.124    pad_id = cast("int | None", tokenizer.pad_token_id)125    if pad_id is None:126        pad_id = cast("int | None", tokenizer.eos_token_id) or 0127    input_ids = pad_sequence(rows, batch_first=True, padding_value=int(pad_id), padding_side="left")128    # Padded from ones rather than `input_ids != pad_id`: a real token may129    # equal `pad_id` when pad falls back to eos.130    attention_mask = pad_sequence(131        [torch.ones_like(row) for row in rows], batch_first=True, padding_side="left"132    )133    return input_ids, attention_mask134 135 136@overload137def prepend_lead_in[ScalarT: np.generic](138    audio: npt.NDArray[ScalarT], sampling_rate: int, seconds: float | None139) -> npt.NDArray[ScalarT]: ...140@overload141def prepend_lead_in[ScalarT: np.generic](142    audio: list[npt.NDArray[ScalarT]], sampling_rate: int, seconds: float | None143) -> list[npt.NDArray[ScalarT]]: ...144@overload145def prepend_lead_in(audio: AudioInput, sampling_rate: int, seconds: float | None) -> AudioInput: ...146def prepend_lead_in(audio: AudioInput, sampling_rate: int, seconds: float | None) -> AudioInput:147    """Prepend `seconds` of silence to a waveform (or each waveform in a list).148 149    Peoples ships fixed ~15s grid cuts rather than sentence-aligned segments,150    so a clip routinely opens mid-word and the model declines to emit the151    partial first token. Measured on 500 Peoples clips with a paired152    bootstrap: 20.51% -> 19.28% WER (delta -1.22, CI [-1.83, -0.64]) and153    utterances dropping a leading reference word fall 258/460 -> 170/460.154    CommonVoice, whose clips already start cleanly, is unaffected (+0.30,155    CI [-0.43, +1.17]).156 157    Inference only. Training feeds raw audio through the collator, so this is158    a test-time transform, and it recovers two thirds of the dropped onsets159    rather than all of them -- the remainder are clips whose first syllable160    was never recorded, which no amount of lead-in reconstructs.161    """162    if not seconds or seconds <= 0:163        return audio164 165    if isinstance(audio, (list, tuple)) and audio and not isinstance(audio[0], (int, float)):166        batch = cast("Sequence[Waveform]", audio)167        return [cast("Waveform", prepend_lead_in(a, sampling_rate, seconds)) for a in batch]168 169    waveform = cast("Waveform", audio)170    pad = round(sampling_rate * seconds)171    if pad <= 0:172        return waveform173    arr: npt.NDArray[Any] = np.asarray(waveform)174    padded: npt.NDArray[Any] = np.pad(arr, (pad, 0))175    return padded176 177 178# Audio is transcribed in chunks cut at the quietest point between179# these lengths. The model trained on clips of at most 19 s; 18 leaves room for180# the inference lead-in. Short clips are one chunk, so their text is unchanged.181CHUNK_MAX_S = 18.0182CHUNK_MIN_S = 8.0183 184 185def chunk_bounds(186    audio: npt.NDArray[np.float32],187    sample_rate: int,188    max_s: float = CHUNK_MAX_S,189    min_s: float = CHUNK_MIN_S,190) -> list[tuple[int, int]]:191    """Sample ranges of at most `max_s`, each cut at the quietest 100 ms frame after `min_s`."""192    frame = int(0.1 * sample_rate)193    bounds: list[tuple[int, int]] = []194    start, n = 0, len(audio)195    while n - start > max_s * sample_rate:196        lo = start + int(min_s * sample_rate)197        hi = start + int(max_s * sample_rate)198        cut = lo + int(np.argmin(_frame_rms(audio[lo:hi], frame))) * frame + frame // 2199        bounds.append((start, cut))200        start = cut201    bounds.append((start, n))202    return bounds203 204 205def _frame_rms(audio: npt.NDArray[np.float32], frame: int) -> npt.NDArray[np.float32]:206    """RMS of each whole `frame`-sample frame of `audio` (a trailing partial frame is dropped)."""207    k = len(audio) // frame208    return np.sqrt(np.mean(np.square(audio[: k * frame].reshape(k, frame)), axis=1))209 210 211# A chunk whose loudest 100 ms frame is this far below the recording's speech212# level (its 95th-percentile frame) holds no speech, only the room tone after213# the talker stopped. Decoded, such a tail comes back as a memorized sentence214# ("The film was directed by the director of the same name.", 0.7 WER on215# CommonVoice) or a stray "the"/"ok". On the cached eval clips over 18 s every216# noise-only chunk sat at -36 dB or below and every chunk with speech at217# -17.5 dB or above; -30 keeps the wider margin on the speech side.218QUIET_CHUNK_DB = -30.0219 220 221def audible_chunks(222    audio: npt.NDArray[np.float32], bounds: list[tuple[int, int]], sample_rate: int223) -> list[npt.NDArray[np.float32]]:224    """The chunks of `audio` at `bounds`, those quieter than `QUIET_CHUNK_DB` emptied.225 226    An empty chunk `is_silent`, so it transcribes as "" without the model. A227    one-chunk recording is never emptied: its loudest frame is its own level.228    """229    chunks = [audio[s:e] for s, e in bounds]230    if len(chunks) < 2:231        return chunks232    frame = int(0.1 * sample_rate)233    floor = np.percentile(_frame_rms(audio, frame), 95) * 10 ** (QUIET_CHUNK_DB / 20)234    return [235        chunk if len(chunk) >= frame and _frame_rms(chunk, frame).max() >= floor else chunk[:0]236        for chunk in chunks237    ]238 239 240# Below this RMS (-100 dBFS) a chunk is digital silence: exact zeros, as in241# edited or remixed recordings. Given one, the model answers with a memorized242# training sentence ("The film was directed by the same director who directed243# 'The Man with the Moustache'") -- ten such chunks cost 2.3 WER on one AMI244# meeting -- so it is skipped. Quiet real speech sits near -60 dBFS.245SILENCE_RMS = 1e-5246 247 248def is_silent(audio: npt.NDArray[np.float32]) -> bool:249    """True for digital silence (or an empty array): nothing for the model to hear."""250    return audio.size == 0 or float(np.sqrt(np.mean(np.square(audio)))) < SILENCE_RMS251 252 253class ASRProcessor(ProcessorMixin):254    """Processor for Whisper-based ASR models."""255 256    attributes: ClassVar[list[str]] = ["feature_extractor", "tokenizer"]257    feature_extractor: SequenceFeatureExtractor258    tokenizer: PreTrainedTokenizerBase259    feature_extractor_class = "AutoFeatureExtractor"260    tokenizer_class = "AutoTokenizer"261    # Fallback only. The real value comes from `ASRConfig.audio_token`, which262    # resolves to the decoder's native placeholder where it has one (Gemma 4's263    # pretrained "<|audio|>") and to "<audio>" otherwise. Hardcoding the264    # fallback here fails silently on a native-token decoder: "<audio>" was265    # never added to that vocab, so it tokenizes into ordinary subwords and266    # the prompt ends up with zero scatter positions for N audio embeddings.267    AUDIO_TOKEN = "<audio>"268    TRANSCRIBE_PROMPT = DEFAULT_TRANSCRIBE_PROMPT269 270    def __init__(271        self,272        feature_extractor: SequenceFeatureExtractor,273        tokenizer: PreTrainedTokenizerBase,274        projector: "MLPAudioProjector | None" = None,275        encoder_conv_layers: list[ConvLayerSpec] | None = None,276        audio_token: str | None = None,277        lead_in_seconds: float = 0.0,278    ):279        """Initialize the ASR processor.280 281        Args:282            feature_extractor: Audio feature extractor (WhisperFeatureExtractor)283            tokenizer: Text tokenizer for the language model284            projector: Audio projector module (for computing output lengths)285            encoder_conv_layers: Conv layer specs [(pad, kernel, stride), ...]286            audio_token: Placeholder token scattered with audio embeddings.287                Must match `ASRConfig.audio_token` / `ASRModel.audio_token`;288                defaults to AUDIO_TOKEN.289        """290        self.feature_extractor = feature_extractor291        self.tokenizer = tokenizer292        self.audio_token = audio_token or self.AUDIO_TOKEN293        self.audio_token_id = tokenizer.convert_tokens_to_ids(self.audio_token)294        self.projector = projector295        self.encoder_conv_layers = encoder_conv_layers or DEFAULT_ENCODER_CONV_LAYERS296        self.lead_in_seconds = float(lead_in_seconds)297 298    def _render_prompt(self, num_audio_tokens: int, text: str | None) -> torch.Tensor:299        """Tokenize one chat prompt carrying exactly `num_audio_tokens` placeholders."""300        return render_audio_prompt(301            self.tokenizer, self.audio_token, num_audio_tokens, self.TRANSCRIBE_PROMPT, text302        )303 304    def _stack_prompt_rows(self, rows: list[torch.Tensor]) -> tuple[torch.Tensor, torch.Tensor]:305        """Stack per-sample prompt rows into a batch (see `left_pad_prompt_rows`)."""306        return left_pad_prompt_rows(rows, self.tokenizer)307 308    def __call__(self, *args: Any, **kwargs: Any) -> BatchFeature:309        """Process audio and text inputs for inference; see `_process` for the arguments.310 311        `ProcessorMixin.__call__` takes `(images, text, videos, audio, ...)`; this312        processor takes audio first, so the arguments are forwarded unchanged to313        `_process`, which carries the real signature.314        """315        return BatchFeature(data=self._process(*args, **kwargs))316 317    def _process(318        self,319        audio: AudioInput | None = None,320        text: str | None = None,321        return_tensors: str = "pt",322        **kwargs: Any,323    ) -> dict[str, torch.Tensor]:324        """Process audio and text inputs for inference.325 326        Args:327            audio: Raw audio waveform(s). A batch gets one prompt per sample.328            text: Target transcription (optional, for training - but use DataCollator instead)329            return_tensors: Return format ("pt" for PyTorch)330 331        Returns:332            Dict with input_features, input_ids, attention_mask333        """334        result: dict[str, torch.Tensor] = {}335        token_counts = [0]336 337        # Process audio338        if audio is not None:339            sr = getattr(self.feature_extractor, "sampling_rate", 16000)340            padded_audio = prepend_lead_in(audio, sr, self.lead_in_seconds)341            extract = cast("AudioFeatureExtractor", self.feature_extractor)342            audio_inputs = extract(343                padded_audio,344                sampling_rate=sr,345                return_attention_mask=True,346                return_tensors=return_tensors,347                **kwargs,348            )349            result["input_features"] = audio_inputs["input_features"]350            result["audio_attention_mask"] = audio_inputs["attention_mask"]351 352            if self.projector is None:353                msg = (354                    "ASRProcessor needs a projector to size the audio prompt. Build it "355                    "with ASRModel.get_processor() instead of constructing it directly."356                )357                raise ValueError(msg)358 359            # One count per sample, from that sample's own mel length. Sizing a360            # single shared prompt from the batch max -- which this used to do --361            # returns batch-1 `input_ids` against batch-B `input_features`, and362            # gives every shorter row more `<audio>` placeholders than the363            # projector produced for it. `masked_scatter` then mis-scatters364            # silently. This is the same failure `_prepare_audio_inputs`365            # documents as fixed on the model side, and it only shows up on a366            # ragged batch, so batch-1 eval never sees it.367            mel_lengths = audio_inputs["attention_mask"].sum(dim=-1).reshape(-1).long()368            encoder_lengths = compute_encoder_output_length(mel_lengths, self.encoder_conv_layers)369            token_counts = self.projector.get_output_length(encoder_lengths).tolist()370 371        rows = [self._render_prompt(n, text) for n in token_counts]372        input_ids, attention_mask = self._stack_prompt_rows(rows)373        result["input_ids"] = input_ids374        result["attention_mask"] = attention_mask375 376        return result377 378 379ASRProcessor.register_for_auto_class()380transformers.AutoProcessor.register(ASRConfig, ASRProcessor)381