diff --git a/tests/conftest.py b/tests/conftest.py index 1643f31..12bed02 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -47,10 +47,16 @@ def to_float16(model: Whisper) -> Whisper: return model -def random_whisper(seed: int = 0, dims: ModelDimensions = TINY_EN_DIMS) -> Whisper: - """A float16 Whisper model with random weights, deterministic for a given seed.""" +def random_whisper( + seed: int = 0, dims: ModelDimensions = TINY_EN_DIMS, dtype: mx.Dtype = mx.float16 +) -> Whisper: + """A Whisper model with random weights, deterministic for a given seed.""" mx.random.seed(seed) - return to_float16(Whisper(dims, mx.float16)) + model = Whisper(dims, dtype) + if dtype == mx.float16: + return to_float16(model) + mx.eval(model.parameters()) + return model @pytest.fixture diff --git a/tests/test_speculative.py b/tests/test_speculative.py new file mode 100644 index 0000000..3548bb8 --- /dev/null +++ b/tests/test_speculative.py @@ -0,0 +1,169 @@ +"""Speculative decoding, VAD and the deprecated chunk helper.""" + +import importlib +from dataclasses import replace + +import mlx.core as mx +import numpy as np +import pytest + +from tests.conftest import TINY_EN_DIMS, random_whisper +from whisper_mlx.decoding import DecodingOptions +from whisper_mlx.speculative import ( + SpeculativeDecoder, + VADProcessor, + parallel_chunk_transcribe, +) + +speculative = importlib.import_module("whisper_mlx.speculative") + +MAX_TOKENS = 40 + + +@pytest.fixture(scope="module") +def target(): + return random_whisper(seed=1, dtype=mx.float32) + + +@pytest.fixture(scope="module") +def draft(): + return random_whisper(seed=2, dtype=mx.float32) + + +def mel(seed: int, n_mels: int = 80) -> mx.array: + mx.random.seed(seed) + return mx.random.normal((3000, n_mels)) + + +def greedy(model, window, language="en"): + options = DecodingOptions( + language=language, without_timestamps=True, sample_len=MAX_TOKENS, fp16=False + ) + return model.decode(window, options).tokens + + +def speculative_decoder(draft, target, **kwargs): + return SpeculativeDecoder( + draft_model=draft, + target_model=target, + dtype=mx.float32, + max_tokens=MAX_TOKENS, + **kwargs, + ) + + +@pytest.mark.parametrize("seed, num_draft_tokens", [(0, 1), (0, 4), (1, 3)]) +def test_output_matches_target_greedy_decoding(draft, target, seed, num_draft_tokens): + decoder = speculative_decoder(draft, target, num_draft_tokens=num_draft_tokens) + + result = decoder.decode_segment(mel(seed)) + + assert result["tokens"] == greedy(target, mel(seed)) + stats = decoder.get_stats() + assert stats["draft_tokens"] > 0 + assert stats["total_forward_passes"] > 0 + + +def test_identical_draft_is_always_accepted(target): + decoder = speculative_decoder(target, target, num_draft_tokens=4) + + result = decoder.decode_segment(mel(0)) + + stats = decoder.get_stats() + assert result["tokens"] == greedy(target, mel(0)) + assert stats["acceptance_rate"] == 1.0 + # every pass yields the drafts plus one extra token from the target + assert stats["tokens_per_target_pass"] > 4 + + +def test_models_with_different_vocabularies_are_rejected(target): + other = random_whisper(seed=3, dims=replace(TINY_EN_DIMS, n_vocab=51865)) + + with pytest.raises(ValueError, match="share a vocabulary"): + SpeculativeDecoder(draft_model=other, target_model=target) + + +def test_transcribe_gives_each_model_its_own_mel_bands(target): + wide_target = random_whisper(seed=4, dims=replace(TINY_EN_DIMS, n_mels=128)) + decoder = SpeculativeDecoder( + draft_model=random_whisper(seed=5), target_model=wide_target, max_tokens=8 + ) + audio = np.zeros(16000 * 35, dtype=np.float32) + + result = decoder.transcribe(audio) + + assert [(s["start"], s["end"]) for s in result["segments"]] == [ + (0.0, 30.0), + (30.0, 35.0), + ] + assert all(len(s["tokens"]) <= 8 for s in result["segments"]) + + +def test_multi_token_step_on_cache_matches_full_pass(target): + """The causal mask must line up when several tokens follow a kv cache.""" + features = target.encoder(mel(0)[None]) + tokens = mx.array([[50257, 50362, 1000, 2000, 3000, 4000, 5000, 6000]]) + + full, _, _ = target.decoder(tokens, features) + _, cache, _ = target.decoder(tokens[:, :3], features) + rest, _, _ = target.decoder(tokens[:, 3:], features, kv_cache=cache) + + assert np.allclose(np.array(full[:, 3:]), np.array(rest), atol=1e-4) + + +def reference_vad(audio, energy_threshold=0.01, min_speech_duration=0.5, sr=16000): + """The original frame-by-frame implementation, for comparison.""" + frame_size, hop_size = int(0.025 * sr), int(0.010 * sr) + num_frames = (len(audio) - frame_size) // hop_size + 1 + energy = np.zeros(num_frames) + for i in range(num_frames): + frame = audio[i * hop_size : i * hop_size + frame_size] + energy[i] = np.sqrt(np.mean(frame**2)) + energy = energy / (np.max(energy) + 1e-8) + is_speech = energy > energy_threshold + regions, in_speech, start = [], False, 0 + min_frames = int(min_speech_duration / 0.010) + for i, speech in enumerate(is_speech): + if speech and not in_speech: + start, in_speech = i, True + elif not speech and in_speech: + if i - start >= min_frames: + regions.append( + (start * hop_size, min(i * hop_size + frame_size, len(audio))) + ) + in_speech = False + if in_speech and len(is_speech) - start >= min_frames: + regions.append((start * hop_size, len(audio))) + return regions + + +@pytest.mark.parametrize("ends_in_speech", [False, True]) +def test_vad_matches_reference_implementation(ends_in_speech): + rng = np.random.default_rng(0) + audio = np.zeros(16000 * 6, dtype=np.float32) + for start, length in [(0.5, 1.0), (2.0, 0.3), (3.0, 1.2)]: + a, b = int(start * 16000), int((start + length) * 16000) + audio[a:b] = rng.standard_normal(b - a) * 0.3 + if ends_in_speech: + audio[-16000:] = rng.standard_normal(16000) * 0.3 + + vad = VADProcessor() + + assert vad.detect_speech_regions(audio) == reference_vad(audio) + assert len(vad.detect_speech_regions(audio)) == (3 if ends_in_speech else 2) + + +def test_vad_handles_audio_shorter_than_a_frame(): + assert VADProcessor().detect_speech_regions(np.zeros(100)) == [] + + +def test_parallel_chunk_transcribe_is_deprecated(monkeypatch): + transcribe_module = importlib.import_module("whisper_mlx.transcribe") + monkeypatch.setattr( + transcribe_module, + "transcribe", + lambda *a, **k: {"text": "", "segments": [], "language": "en"}, + ) + + with pytest.deprecated_call(): + parallel_chunk_transcribe(np.zeros(16000 * 5, dtype=np.float32)) diff --git a/whisper_mlx/speculative.py b/whisper_mlx/speculative.py index 3905b9a..f5199dc 100644 --- a/whisper_mlx/speculative.py +++ b/whisper_mlx/speculative.py @@ -1,220 +1,362 @@ # Copyright © 2024 Whisper MLX Contributors -# Novel: Speculative Decoding for Whisper -# -# Idea: Use a tiny/fast model to "draft" tokens, then verify with main model. -# The main model can verify multiple tokens in parallel (one forward pass), -# making it much faster than autoregressive decoding. +# Speculative decoding for Whisper """ Speculative Decoding for Whisper MLX -This implements speculative decoding where: -1. A small "draft" model (e.g., whisper-tiny) generates candidate tokens quickly -2. The main "target" model (e.g., whisper-large) verifies tokens in parallel -3. Accepted tokens are kept, rejected tokens trigger re-generation +A small "draft" model proposes a few tokens at a time, and the main "target" +model checks all of them in a single forward pass: -This can provide 2-3x speedup on top of batched decoding. +1. The draft model decodes `num_draft_tokens` tokens greedily. +2. The target model scores those tokens in one pass over its key/value cache. +3. Drafted tokens are kept up to the first one the target disagrees with; the + target's own choice is used there (or appended, if every draft matched). + +The output is exactly the target model's greedy (temperature 0, no timestamps) +transcript. Whether it is faster depends on how often the draft agrees with the +target and how cheap the draft is, so measure it on your own audio. The two +models must share a vocabulary, e.g. distil-large-v3 drafting for large-v3. Usage: from whisper_mlx.speculative import speculative_transcribe result = speculative_transcribe( "audio.mp3", - draft_model="tiny", + draft_model="distil-large-v3", target_model="large-v3", - batch_size=12, ) """ import time -from typing import Optional, Union - -from .utils import MODEL_REPOS +import warnings +from typing import Dict, List, Optional, Sequence, Tuple, Union import mlx.core as mx import numpy as np -from .audio import N_FRAMES, log_mel_spectrogram, pad_or_trim -from .constants import DEFAULT_DRAFT_TOKENS, MAX_TOKENS_PER_SEGMENT +from .audio import ( + HOP_LENGTH, + N_FRAMES, + N_SAMPLES, + SAMPLE_RATE, + load_audio, + log_mel_spectrogram, + pad_or_trim, +) +from .constants import DEFAULT_DRAFT_TOKENS +from .decoding import DecodingOptions, DecodingTask, LogitFilter from .load_models import load_model from .tokenizer import get_tokenizer +from .utils import MODEL_REPOS +from .whisper import Whisper + +KVCache = List[Tuple[Tuple[mx.array, mx.array], Tuple[mx.array, mx.array]]] + + +def _cache_length(cache: Optional[KVCache]) -> int: + return 0 if cache is None else cache[0][0][0].shape[1] + + +def _trim_cache(cache: KVCache, length: int) -> KVCache: + """Drop self-attention entries past `length`; cross-attention is unchanged.""" + return [((k[:, :length], v[:, :length]), cross_kv) for (k, v), cross_kv in cache] + + +def _forward( + model: Whisper, tokens: Sequence[int], features: mx.array, cache: Optional[KVCache] +) -> Tuple[mx.array, KVCache]: + """Feed `tokens` after the cached ones; returns one row of logits per token.""" + logits, cache, _ = model.decoder(mx.array([tokens]), features, kv_cache=cache) + return logits[0].astype(mx.float32), cache + + +def _greedy(logits: mx.array, context: List[int], filters: List[LogitFilter]) -> int: + """Pick the next token after `context` the way DecodingTask does at T=0.""" + logits = logits[None] + tokens = mx.array([context]) + for logit_filter in filters: + logits = logit_filter.apply(logits, tokens) + return mx.argmax(logits, axis=-1).item() class SpeculativeDecoder: """ - Speculative decoding: draft with fast model, verify with accurate model. + Speculative decoding: draft with a fast model, verify with an accurate one. - The key insight is that verification is parallelizable: - - Draft model generates K tokens autoregressively (slow but using tiny model) - - Target model verifies all K tokens in ONE forward pass (fast) - - If all tokens match, we get K tokens for ~1 forward pass cost + Both models keep a key/value cache, so each draft step costs one draft + decoder pass for one token, and each verification costs one target decoder + pass over all drafted tokens. Audio is encoded once per model per window. """ def __init__( self, - draft_model_path: str = "mlx-community/whisper-tiny-mlx", + draft_model_path: str = "mlx-community/distil-whisper-large-v3", target_model_path: str = "mlx-community/whisper-large-v3-mlx", num_draft_tokens: int = DEFAULT_DRAFT_TOKENS, dtype: mx.Dtype = mx.float16, + *, + draft_model: Optional[Whisper] = None, + target_model: Optional[Whisper] = None, + max_tokens: Optional[int] = None, ): """ Initialize speculative decoder. Args: - draft_model_path: Fast model for drafting (tiny recommended) - target_model_path: Accurate model for verification - num_draft_tokens: Number of tokens to draft before verification + draft_model_path: Fast model for drafting + target_model_path: Accurate model whose output is reproduced + num_draft_tokens: Number of tokens to draft before each verification dtype: Data type for computation + draft_model, target_model: Already-loaded models to use instead of + loading from the paths + max_tokens: Maximum tokens per 30-second window (default: the model's + limit, 224) """ - print(f"Loading draft model: {draft_model_path}") - self.draft_model = load_model(draft_model_path, dtype=dtype) - - print(f"Loading target model: {target_model_path}") - self.target_model = load_model(target_model_path, dtype=dtype) + if num_draft_tokens < 1: + raise ValueError(f"num_draft_tokens must be >= 1, got {num_draft_tokens}") + + if draft_model is None: + draft_model = load_model(draft_model_path, dtype=dtype) + if target_model is None: + target_model = load_model(target_model_path, dtype=dtype) + self.draft_model = draft_model + self.target_model = target_model + + draft_vocab = self.draft_model.dims.n_vocab + target_vocab = self.target_model.dims.n_vocab + if draft_vocab != target_vocab: + raise ValueError( + f"Draft and target models must share a vocabulary, but have " + f"{draft_vocab} and {target_vocab} tokens. Pair models from the same " + "family, e.g. distil-large-v3 with large-v3, or tiny with large-v2." + ) self.num_draft_tokens = num_draft_tokens self.dtype = dtype + self.max_tokens = max_tokens - # Get tokenizer (same for both models) - self.tokenizer = get_tokenizer(self.target_model.is_multilingual) + self.tokenizer = get_tokenizer( + self.target_model.is_multilingual, + num_languages=self.target_model.num_languages, + ) + self._tasks: Dict[Tuple[int, str], DecodingTask] = {} # Stats tracking self.stats = { "draft_tokens": 0, "accepted_tokens": 0, - "total_forward_passes": 0, + "generated_tokens": 0, + "draft_forward_passes": 0, + "total_forward_passes": 0, # target decoder passes } - def _draft_tokens(self, mel: mx.array, initial_tokens: list) -> list: - """Generate draft tokens using the fast model.""" - tokens = initial_tokens.copy() - - # Encode audio once - audio_features = self.draft_model.encoder(mel) - - for _ in range(self.num_draft_tokens): - # Decode next token - token_array = mx.array([tokens]) - logits = self.draft_model.decoder(token_array, audio_features) - next_token = mx.argmax(logits[0, -1]).item() - tokens.append(next_token) - - # Stop if end of text - if next_token == self.tokenizer.eot: - break - - self.stats["draft_tokens"] += len(tokens) - len(initial_tokens) - return tokens - - def _verify_tokens( + def _task(self, model: Whisper, language: str) -> DecodingTask: + """Greedy, timestamp-free decoding setup: prompt tokens and logit filters.""" + key = (id(model), language) + if key not in self._tasks: + options = DecodingOptions( + language=language, + without_timestamps=True, + sample_len=self.max_tokens, + fp16=self.dtype == mx.float16, + ) + self._tasks[key] = DecodingTask(model, options) + return self._tasks[key] + + def _encode(self, model: Whisper, mel: mx.array) -> mx.array: + if mel.ndim == 2: + mel = mel[None] + return model.encoder(mel.astype(self.dtype)) + + def decode_segment( self, mel: mx.array, - initial_tokens: list, - draft_tokens: list, - ) -> tuple[list, bool]: + language: str = "en", + *, + target_mel: Optional[mx.array] = None, + ) -> dict: """ - Verify draft tokens using target model. + Decode one 30-second window with speculative decoding. - This is the key: we verify ALL draft tokens in ONE forward pass! - """ - # Encode audio - audio_features = self.target_model.encoder(mel) - - # Get target model's predictions for all positions - token_array = mx.array([draft_tokens[:-1]]) # Input excludes last token - logits = self.target_model.decoder(token_array, audio_features) - - self.stats["total_forward_passes"] += 1 - - # Verify each drafted token - accepted_tokens = initial_tokens.copy() - all_accepted = True - - for i, draft_token in enumerate(draft_tokens[len(initial_tokens):]): - pos = len(initial_tokens) + i - 1 - if pos < 0: - pos = 0 - - # Get target model's prediction for this position - target_token = mx.argmax(logits[0, pos]).item() - - if target_token == draft_token: - accepted_tokens.append(draft_token) - self.stats["accepted_tokens"] += 1 - else: - # Rejection: use target's token and stop - accepted_tokens.append(target_token) - all_accepted = False - break - - return accepted_tokens, all_accepted - - def decode_segment(self, mel: mx.array, language: str = "en") -> dict: - """ - Decode a single audio segment using speculative decoding. + Args: + mel: Log-mel spectrogram of the window, shape (3000, n_mels), for the + draft model (and the target model, unless target_mel is given) + language: Language code; ignored by English-only models + target_mel: Spectrogram for the target model when it expects a + different number of mel bands than the draft (e.g. 128 vs 80) + + Returns: + dict with the sampled "tokens" (without end-of-text) and "text" """ - # Initial tokens - tokens = [ - self.tokenizer.sot, - self.tokenizer.special_tokens[f"<|{language}|>"], - self.tokenizer.special_tokens["<|transcribe|>"], - self.tokenizer.special_tokens["<|notimestamps|>"], - ] - - max_tokens = MAX_TOKENS_PER_SEGMENT - - while len(tokens) < max_tokens: - # Draft tokens - draft_tokens = self._draft_tokens(mel, tokens) - - # Verify with target model - tokens, all_accepted = self._verify_tokens(mel, tokens, draft_tokens) + if target_mel is None: + target_mel = mel + + draft_task = self._task(self.draft_model, language) + target_task = self._task(self.target_model, language) + tokenizer = target_task.tokenizer + eot = tokenizer.eot + sample_begin = target_task.sample_begin + max_length = sample_begin + target_task.sample_len + + draft_features = self._encode(self.draft_model, mel) + target_features = self._encode(self.target_model, target_mel) + draft_cache: Optional[KVCache] = None + target_cache: Optional[KVCache] = None + + tokens = list(target_task.initial_tokens) + while len(tokens) < max_length and tokens[-1] != eot: + # 1. Draft up to num_draft_tokens tokens, one cached pass each + context = list(tokens) + drafts: List[int] = [] + n_draft = min(self.num_draft_tokens, max_length - len(tokens)) + while len(drafts) < n_draft: + cached = _cache_length(draft_cache) + logits, draft_cache = _forward( + self.draft_model, context[cached:], draft_features, draft_cache + ) + self.stats["draft_forward_passes"] += 1 + token = _greedy(logits[-1], context, draft_task.logit_filters) + drafts.append(token) + context.append(token) + if token == eot: + break + + # 2. Score every drafted position with one target pass; + # row j predicts the token after context[cached + j] + cached = _cache_length(target_cache) + logits, target_cache = _forward( + self.target_model, context[cached:], target_features, target_cache + ) + self.stats["total_forward_passes"] += 1 + + # 3. Keep drafts while the target agrees, then take the target's token + accepted: List[int] = [] + for i, draft in enumerate(drafts + [None]): + position = len(tokens) + i + token = _greedy( + logits[position - 1 - cached], + context[:position], + target_task.logit_filters, + ) + accepted.append(token) + if token != draft or token == eot: + break + + matched = 0 + while matched < len(drafts) and accepted[matched] == drafts[matched]: + matched += 1 + self.stats["draft_tokens"] += len(drafts) + self.stats["accepted_tokens"] += matched + + # 4. Roll both caches back to the tokens that were kept + valid = len(tokens) + matched + tokens = (tokens + accepted)[:max_length] + target_cache = _trim_cache(target_cache, min(valid, len(tokens) - 1)) + draft_cache = _trim_cache(draft_cache, min(valid, len(tokens) - 1)) + + sampled = tokens[sample_begin:] + if eot in sampled: + sampled = sampled[: sampled.index(eot)] + self.stats["generated_tokens"] += len(sampled) - # Check for end - if tokens[-1] == self.tokenizer.eot: - break + return { + "tokens": sampled, + "text": tokenizer.decode(sampled).strip(), + } - # Decode tokens to text - text_tokens = [t for t in tokens if t < self.tokenizer.eot] - text = self.tokenizer.decode(text_tokens) + def transcribe( + self, + audio: Union[str, np.ndarray, mx.array], + language: str = "en", + verbose: bool = False, + ) -> dict: + """Transcribe audio window by window (30 s each) with speculative decoding.""" + if isinstance(audio, str): + audio = load_audio(audio) + elif not isinstance(audio, mx.array): + audio = mx.array(audio) + + draft_mels = self.draft_model.dims.n_mels + target_mels = self.target_model.dims.n_mels + # Pad 30 seconds of silence, as transcribe() does, for slicing + mels = { + n_mels: log_mel_spectrogram(audio, n_mels=n_mels, padding=N_SAMPLES) + for n_mels in {draft_mels, target_mels} + } + content_frames = mels[draft_mels].shape[0] - N_FRAMES + + segments = [] + start_time = time.time() + for seek in range(0, content_frames, N_FRAMES): + segment_size = min(N_FRAMES, content_frames - seek) + window = { + n_mels: pad_or_trim(mel[seek : seek + segment_size], N_FRAMES, axis=-2) + for n_mels, mel in mels.items() + } + result = self.decode_segment( + window[draft_mels], language, target_mel=window[target_mels] + ) + start = seek * HOP_LENGTH / SAMPLE_RATE + segments.append( + { + "start": start, + "end": (seek + segment_size) * HOP_LENGTH / SAMPLE_RATE, + "text": result["text"], + "tokens": result["tokens"], + } + ) + if verbose: + print(f"[{start:.1f}s] {result['text']}") return { - "tokens": tokens, - "text": text, + "text": " ".join(s["text"] for s in segments if s["text"]), + "segments": segments, + "stats": self.get_stats(), + "elapsed": time.time() - start_time, } def get_stats(self) -> dict: - """Get decoding statistics.""" - acceptance_rate = ( - self.stats["accepted_tokens"] / self.stats["draft_tokens"] - if self.stats["draft_tokens"] > 0 else 0 - ) + """ + Get decoding statistics. + + `tokens_per_target_pass` is the number of tokens produced per target + decoder pass (1.0 for ordinary greedy decoding); it bounds, but does not + measure, the wall-clock speed-up. + """ + draft_tokens = self.stats["draft_tokens"] + passes = max(1, self.stats["total_forward_passes"]) + tokens_per_pass = self.stats["generated_tokens"] / passes return { **self.stats, - "acceptance_rate": acceptance_rate, - "speedup_factor": self.stats["draft_tokens"] / max(1, self.stats["total_forward_passes"]), + "acceptance_rate": ( + self.stats["accepted_tokens"] / draft_tokens if draft_tokens else 0 + ), + "tokens_per_target_pass": tokens_per_pass, + "speedup_factor": tokens_per_pass, # kept for backward compatibility } def speculative_transcribe( audio: Union[str, np.ndarray], - draft_model: str = "tiny", + draft_model: str = "distil-large-v3", target_model: str = "large-v3", language: str = "en", verbose: bool = True, + num_draft_tokens: int = DEFAULT_DRAFT_TOKENS, ) -> dict: """ Transcribe audio using speculative decoding. Args: audio: Path to audio file or audio array - draft_model: Small model for drafting ("tiny", "base", "small") - target_model: Large model for verification ("large-v3", "turbo") + draft_model: Small model for drafting; must share the target's vocabulary + target_model: Model whose greedy output is reproduced language: Language code verbose: Print progress + num_draft_tokens: Tokens drafted per verification pass Returns: - dict with "text", "segments", and "stats" + dict with "text", "segments", "stats" and "elapsed" """ # Resolve model paths using centralized mapping draft_path = MODEL_REPOS.get(draft_model, draft_model) @@ -224,73 +366,28 @@ def speculative_transcribe( print(f"Speculative Decoding: {draft_model} → {target_model}") print("=" * 50) - # Initialize decoder decoder = SpeculativeDecoder( draft_model_path=draft_path, target_model_path=target_path, + num_draft_tokens=num_draft_tokens, ) - - # Load and process audio - from .audio import load_audio, log_mel_spectrogram, pad_or_trim, N_FRAMES, SAMPLE_RATE - - if isinstance(audio, str): - audio_array = load_audio(audio) - else: - audio_array = audio - - # Compute mel spectrogram - mel = log_mel_spectrogram(audio_array) - - # Process in segments - segments = [] - seek = 0 - total_frames = mel.shape[-2] - - start_time = time.time() - - while seek < total_frames: - # Get segment - segment_mel = mel[seek : seek + N_FRAMES] - segment_mel = pad_or_trim(segment_mel, N_FRAMES, axis=-2) - segment_mel = mx.expand_dims(segment_mel, axis=0) - - # Decode - result = decoder.decode_segment(segment_mel, language=language) - - # Add segment - time_offset = seek * 160 / SAMPLE_RATE # HOP_LENGTH = 160 - segments.append({ - "start": time_offset, - "end": time_offset + 30.0, - "text": result["text"], - }) - - if verbose: - print(f"[{time_offset:.1f}s] {result['text']}") - - seek += N_FRAMES - - elapsed = time.time() - start_time - stats = decoder.get_stats() + result = decoder.transcribe(audio, language=language, verbose=verbose) if verbose: + stats = result["stats"] print("=" * 50) - print(f"Time: {elapsed:.2f}s") + print(f"Time: {result['elapsed']:.2f}s") print(f"Acceptance rate: {stats['acceptance_rate']:.1%}") - print(f"Effective speedup: {stats['speedup_factor']:.2f}x") + print(f"Tokens per target pass: {stats['tokens_per_target_pass']:.2f}") - return { - "text": " ".join(s["text"] for s in segments), - "segments": segments, - "stats": stats, - "elapsed": elapsed, - } + return result # ============================================================ # IDEA 2: VAD-Guided Processing (Skip Silence) # ============================================================ + class VADProcessor: """ Voice Activity Detection guided processing. @@ -299,77 +396,72 @@ class VADProcessor: and only process those, skipping silence entirely. """ - def __init__(self, energy_threshold: float = 0.01, min_speech_duration: float = 0.5): + def __init__( + self, energy_threshold: float = 0.01, min_speech_duration: float = 0.5 + ): self.energy_threshold = energy_threshold self.min_speech_duration = min_speech_duration - def detect_speech_regions(self, audio: np.ndarray, sample_rate: int = 16000) -> list: + def detect_speech_regions( + self, audio: np.ndarray, sample_rate: int = 16000 + ) -> list: """ Simple energy-based VAD to detect speech regions. Returns list of (start_sample, end_sample) tuples. """ - # Convert to numpy if mlx array - if hasattr(audio, 'tolist') and not isinstance(audio, np.ndarray): - audio = np.array(audio.tolist()) + audio = np.asarray(audio, dtype=np.float64).ravel() # Frame-based energy calculation frame_size = int(0.025 * sample_rate) # 25ms frames - hop_size = int(0.010 * sample_rate) # 10ms hop + hop_size = int(0.010 * sample_rate) # 10ms hop + if len(audio) < frame_size: + return [] + # RMS energy per frame from a running sum of squares num_frames = (len(audio) - frame_size) // hop_size + 1 - energy = np.zeros(num_frames) - - for i in range(num_frames): - start = i * hop_size - frame = audio[start : start + frame_size] - energy[i] = np.sqrt(np.mean(frame ** 2)) + cumulative = np.concatenate([[0.0], np.cumsum(audio**2)]) + starts = np.arange(num_frames) * hop_size + energy = np.sqrt( + np.maximum(cumulative[starts + frame_size] - cumulative[starts], 0.0) + / frame_size + ) # Normalize energy energy = energy / (np.max(energy) + 1e-8) - # Find speech regions - is_speech = energy > self.energy_threshold - - # Merge close regions and filter short ones - regions = [] - in_speech = False - start = 0 + # Find runs of speech frames long enough to keep + is_speech = np.concatenate([[False], energy > self.energy_threshold, [False]]) + edges = np.flatnonzero(is_speech[1:] != is_speech[:-1]) + run_starts, run_ends = edges[::2], edges[1::2] # run_ends is exclusive min_frames = int(self.min_speech_duration / 0.010) - - for i, speech in enumerate(is_speech): - if speech and not in_speech: - start = i - in_speech = True - elif not speech and in_speech: - if i - start >= min_frames: - start_sample = start * hop_size - end_sample = min(i * hop_size + frame_size, len(audio)) - regions.append((start_sample, end_sample)) - in_speech = False - - # Handle case where audio ends during speech - if in_speech and len(is_speech) - start >= min_frames: - regions.append((start * hop_size, len(audio))) + regions = [] + for start, end in zip(run_starts, run_ends): + if end - start < min_frames: + continue + if end == num_frames: + # audio ends during speech + end_sample = len(audio) + else: + end_sample = min(end * hop_size + frame_size, len(audio)) + regions.append((int(start * hop_size), int(end_sample))) return regions def get_skip_ratio(self, audio: np.ndarray, sample_rate: int = 16000) -> float: """Calculate what percentage of audio can be skipped.""" - # Convert to numpy if mlx array - if hasattr(audio, 'tolist') and not isinstance(audio, np.ndarray): - audio = np.array(audio.tolist()) - + audio = np.asarray(audio) regions = self.detect_speech_regions(audio, sample_rate) speech_samples = sum(end - start for start, end in regions) return 1.0 - (speech_samples / len(audio)) # ============================================================ -# IDEA 3: Parallel Chunk Processing with Overlap Merging +# IDEA 3: Chunk Processing with Overlap Merging # ============================================================ + def parallel_chunk_transcribe( audio: Union[str, np.ndarray], model_path: str = "mlx-community/whisper-turbo", @@ -378,14 +470,17 @@ def parallel_chunk_transcribe( language: str = "en", ) -> dict: """ - Process ALL audio chunks in parallel, then merge with overlap handling. - - Unlike sequential processing, this transcribes all chunks simultaneously, - then uses overlap regions to stitch them together correctly. + Deprecated: transcribes overlapping chunks one after another, then merges them. - This maximizes GPU utilization for long audio files. + Despite the name, chunks are not processed in parallel. Use + `transcribe(audio, batch_size=N)`, which decodes N windows at once. """ - from .audio import load_audio, log_mel_spectrogram, SAMPLE_RATE, HOP_LENGTH + warnings.warn( + "parallel_chunk_transcribe processes chunks sequentially and is deprecated; " + "use transcribe(audio, batch_size=N) for batched decoding.", + DeprecationWarning, + stacklevel=2, + ) from .transcribe import transcribe # Load audio @@ -409,23 +504,21 @@ def parallel_chunk_transcribe( starts.append(pos / SAMPLE_RATE) pos += step_samples - print(f"Processing {len(chunks)} chunks in parallel...") - - # Transcribe all chunks with maximum batch size - # (In practice, you'd process all mel spectrograms as one big batch) results = [] for i, chunk in enumerate(chunks): result = transcribe( chunk, path_or_hf_repo=model_path, - batch_size=1, # Each chunk is one "batch" + batch_size=1, language=language, verbose=False, ) - results.append({ - "start": starts[i], - "result": result, - }) + results.append( + { + "start": starts[i], + "result": result, + } + ) # Merge overlapping segments merged_segments = merge_overlapping_segments(results, overlap_duration) @@ -440,7 +533,7 @@ def merge_overlapping_segments(results: list, overlap_duration: float) -> list: """ Merge transcription results from overlapping chunks. - Uses text similarity in overlap regions to find best merge points. + Where segments overlap, keeps the one with the longer text. """ if not results: return [] diff --git a/whisper_mlx/whisper.py b/whisper_mlx/whisper.py index 5c85195..9e1f797 100644 --- a/whisper_mlx/whisper.py +++ b/whisper_mlx/whisper.py @@ -79,7 +79,10 @@ def qkv_attention(self, q, k, v, mask=None): qk = q @ k if mask is not None: - qk = qk + mask[:n_ctx, :n_ctx] + # queries are the last n_ctx of the n_keys positions (earlier ones may + # come from the kv cache), so take the matching rows of the causal mask + n_keys = k.shape[-1] + qk = qk + mask[n_keys - n_ctx : n_keys, :n_keys] w = mx.softmax(qk, axis=-1, precise=True) out = (w @ v).transpose(0, 2, 1, 3)