mazesmazes/tiny-audio
1177
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 