"""Processor that turns raw audio (and optional text) into model inputs.""" from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar, cast, overload import numpy as np import numpy.typing as npt import torch import transformers from torch.nn.utils.rnn import pad_sequence from transformers import ( BatchFeature, PreTrainedTokenizerBase, ProcessorMixin, SequenceFeatureExtractor, ) if TYPE_CHECKING: from .asr_config import ( DEFAULT_ENCODER_CONV_LAYERS, ASRConfig, ConvLayerSpec, compute_encoder_output_length, ) from .asr_types import AudioFeatureExtractor, AudioInput, PreparedChunk, Waveform from .projectors import MLPAudioProjector else: try: from .asr_config import ( DEFAULT_ENCODER_CONV_LAYERS, ASRConfig, ConvLayerSpec, compute_encoder_output_length, ) from .asr_types import AudioInput, PreparedChunk except ImportError: # flat layout on the Hub: sibling modules, no package from asr_config import ( DEFAULT_ENCODER_CONV_LAYERS, ASRConfig, ConvLayerSpec, compute_encoder_output_length, ) from asr_types import AudioInput, PreparedChunk def collate_chunks(prepared: Sequence[PreparedChunk]) -> PreparedChunk: """Pad prepared chunks to the longest and stack them into one batch. The time axis is whichever feature axis matches the mask's length (Granite's features are `(1, T, D)`, Whisper's `(1, D, T)`); padded frames are zeros with a 0 in the mask, which the encoder honours (`encoder_attention_mask`). """ longest = max(int(p["attention_mask"].shape[-1]) for p in prepared) features: list[torch.Tensor] = [] masks: list[torch.Tensor] = [] for p in prepared: feats, mask = p["input_features"], p["attention_mask"] length = int(mask.shape[-1]) time_axis = 1 if feats.shape[1] == length else feats.dim() - 1 pad = longest - length # F.pad lists (left, right) pairs from the LAST axis backwards. spec = [0, 0] * (feats.dim() - 1 - time_axis) + [0, pad] features.append(torch.nn.functional.pad(feats, spec)) masks.append(torch.nn.functional.pad(mask, (0, pad))) return {"input_features": torch.cat(features), "attention_mask": torch.cat(masks)} # The instruction the model trained on (scripts/train_collator.py); the model # and processor both default to it. DEFAULT_TRANSCRIBE_PROMPT = "Transcribe the speech to text" def render_audio_prompt( tokenizer: PreTrainedTokenizerBase, audio_token: str, num_audio_tokens: int, prompt: str | None, text: str | None = None, ) -> torch.Tensor: """Tokenize one chat prompt carrying exactly `num_audio_tokens` placeholders. The user turn is the placeholders, then `prompt` (if any); `text`, when given, is the assistant's reply, otherwise the generation prompt is added. """ if num_audio_tokens > 0: user_content = audio_token * num_audio_tokens if prompt: user_content += " " + prompt else: user_content = prompt or "" messages = [{"role": "user", "content": user_content}] if text is not None: messages.append({"role": "assistant", "content": text}) # With `tokenize=True, return_tensors="pt"` the ids come back as tensors. tokenized = cast( "torch.Tensor | Mapping[str, torch.Tensor]", tokenizer.apply_chat_template( messages, tokenize=True, add_generation_prompt=(text is None), return_tensors="pt", enable_thinking=False, # Disable Qwen3 thinking mode for ASR ), ) # apply_chat_template returns a bare tensor or a BatchEncoding/mapping. ids = tokenized if isinstance(tokenized, torch.Tensor) else tokenized["input_ids"] return (ids[0] if ids.dim() > 1 else ids).to(torch.long) def left_pad_prompt_rows( rows: list[torch.Tensor], tokenizer: PreTrainedTokenizerBase ) -> tuple[torch.Tensor, torch.Tensor]: """Stack per-sample prompt rows into a left-padded batch: `(input_ids, attention_mask)`. Left, not right: these feed `generate`, so padding must not sit between the prompt and the first generated token. Pads with the tokenizer's pad token, falling back to eos, then 0. Pad positions never carry `audio_token_id`, so the model's masked_scatter is unaffected. """ # transformers types special-token ids as any token value; a single id is an int. pad_id = cast("int | None", tokenizer.pad_token_id) if pad_id is None: pad_id = cast("int | None", tokenizer.eos_token_id) or 0 input_ids = pad_sequence(rows, batch_first=True, padding_value=int(pad_id), padding_side="left") # Padded from ones rather than `input_ids != pad_id`: a real token may # equal `pad_id` when pad falls back to eos. attention_mask = pad_sequence( [torch.ones_like(row) for row in rows], batch_first=True, padding_side="left" ) return input_ids, attention_mask @overload def prepend_lead_in[ScalarT: np.generic]( audio: npt.NDArray[ScalarT], sampling_rate: int, seconds: float | None ) -> npt.NDArray[ScalarT]: ... @overload def prepend_lead_in[ScalarT: np.generic]( audio: list[npt.NDArray[ScalarT]], sampling_rate: int, seconds: float | None ) -> list[npt.NDArray[ScalarT]]: ... @overload def prepend_lead_in(audio: AudioInput, sampling_rate: int, seconds: float | None) -> AudioInput: ... def prepend_lead_in(audio: AudioInput, sampling_rate: int, seconds: float | None) -> AudioInput: """Prepend `seconds` of silence to a waveform (or each waveform in a list). Peoples ships fixed ~15s grid cuts rather than sentence-aligned segments, so a clip routinely opens mid-word and the model declines to emit the partial first token. Measured on 500 Peoples clips with a paired bootstrap: 20.51% -> 19.28% WER (delta -1.22, CI [-1.83, -0.64]) and utterances dropping a leading reference word fall 258/460 -> 170/460. CommonVoice, whose clips already start cleanly, is unaffected (+0.30, CI [-0.43, +1.17]). Inference only. Training feeds raw audio through the collator, so this is a test-time transform, and it recovers two thirds of the dropped onsets rather than all of them -- the remainder are clips whose first syllable was never recorded, which no amount of lead-in reconstructs. """ if not seconds or seconds <= 0: return audio if isinstance(audio, (list, tuple)) and audio and not isinstance(audio[0], (int, float)): batch = cast("Sequence[Waveform]", audio) return [cast("Waveform", prepend_lead_in(a, sampling_rate, seconds)) for a in batch] waveform = cast("Waveform", audio) pad = round(sampling_rate * seconds) if pad <= 0: return waveform arr: npt.NDArray[Any] = np.asarray(waveform) padded: npt.NDArray[Any] = np.pad(arr, (pad, 0)) return padded # Audio is transcribed in chunks cut at the quietest point between # these lengths. The model trained on clips of at most 19 s; 18 leaves room for # the inference lead-in. Short clips are one chunk, so their text is unchanged. CHUNK_MAX_S = 18.0 CHUNK_MIN_S = 8.0 def chunk_bounds( audio: npt.NDArray[np.float32], sample_rate: int, max_s: float = CHUNK_MAX_S, min_s: float = CHUNK_MIN_S, ) -> list[tuple[int, int]]: """Sample ranges of at most `max_s`, each cut at the quietest 100 ms frame after `min_s`.""" frame = int(0.1 * sample_rate) bounds: list[tuple[int, int]] = [] start, n = 0, len(audio) while n - start > max_s * sample_rate: lo = start + int(min_s * sample_rate) hi = start + int(max_s * sample_rate) cut = lo + int(np.argmin(_frame_rms(audio[lo:hi], frame))) * frame + frame // 2 bounds.append((start, cut)) start = cut bounds.append((start, n)) return bounds def _frame_rms(audio: npt.NDArray[np.float32], frame: int) -> npt.NDArray[np.float32]: """RMS of each whole `frame`-sample frame of `audio` (a trailing partial frame is dropped).""" k = len(audio) // frame return np.sqrt(np.mean(np.square(audio[: k * frame].reshape(k, frame)), axis=1)) # A chunk whose loudest 100 ms frame is this far below the recording's speech # level (its 95th-percentile frame) holds no speech, only the room tone after # the talker stopped. Decoded, such a tail comes back as a memorized sentence # ("The film was directed by the director of the same name.", 0.7 WER on # CommonVoice) or a stray "the"/"ok". On the cached eval clips over 18 s every # noise-only chunk sat at -36 dB or below and every chunk with speech at # -17.5 dB or above; -30 keeps the wider margin on the speech side. QUIET_CHUNK_DB = -30.0 def audible_chunks( audio: npt.NDArray[np.float32], bounds: list[tuple[int, int]], sample_rate: int ) -> list[npt.NDArray[np.float32]]: """The chunks of `audio` at `bounds`, those quieter than `QUIET_CHUNK_DB` emptied. An empty chunk `is_silent`, so it transcribes as "" without the model. A one-chunk recording is never emptied: its loudest frame is its own level. """ chunks = [audio[s:e] for s, e in bounds] if len(chunks) < 2: return chunks frame = int(0.1 * sample_rate) floor = np.percentile(_frame_rms(audio, frame), 95) * 10 ** (QUIET_CHUNK_DB / 20) return [ chunk if len(chunk) >= frame and _frame_rms(chunk, frame).max() >= floor else chunk[:0] for chunk in chunks ] # Below this RMS (-100 dBFS) a chunk is digital silence: exact zeros, as in # edited or remixed recordings. Given one, the model answers with a memorized # training sentence ("The film was directed by the same director who directed # 'The Man with the Moustache'") -- ten such chunks cost 2.3 WER on one AMI # meeting -- so it is skipped. Quiet real speech sits near -60 dBFS. SILENCE_RMS = 1e-5 def is_silent(audio: npt.NDArray[np.float32]) -> bool: """True for digital silence (or an empty array): nothing for the model to hear.""" return audio.size == 0 or float(np.sqrt(np.mean(np.square(audio)))) < SILENCE_RMS class ASRProcessor(ProcessorMixin): """Processor for Whisper-based ASR models.""" attributes: ClassVar[list[str]] = ["feature_extractor", "tokenizer"] feature_extractor: SequenceFeatureExtractor tokenizer: PreTrainedTokenizerBase feature_extractor_class = "AutoFeatureExtractor" tokenizer_class = "AutoTokenizer" # Fallback only. The real value comes from `ASRConfig.audio_token`, which # resolves to the decoder's native placeholder where it has one (Gemma 4's # pretrained "<|audio|>") and to "