From c0c9d6702f2cce0e7b920a845c70f18abe95ac00 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 21:23:21 +0000 Subject: [PATCH] Follow the temperature schedule in batched fallback Batched decoding always decoded at 0.0 and retried failing segments once at 1.0, whatever temperature the caller passed: temperature=0.0 still triggered a retry, the default schedule jumped straight to its noisiest step, and retried results were never re-checked. decode_batch_with_fallback now walks the same schedule as the one-window path, re-decoding only the segments that still fail at each step as one smaller batch. A single temperature means no fallback. The pass/fail check is shared by both paths. The test stub now numbers windows by their mel content so a re-decode maps back to the same window. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA --- tests/conftest.py | 33 +++++---- tests/test_temperature_fallback.py | 74 ++++++++++++++++++++ whisper_mlx/transcribe.py | 108 ++++++++++++----------------- 3 files changed, 139 insertions(+), 76 deletions(-) create mode 100644 tests/test_temperature_fallback.py diff --git a/tests/conftest.py b/tests/conftest.py index 7ba4120..a06c1d0 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -49,12 +49,15 @@ class StubModel: """Stands in for Whisper in transcribe(): decode() returns canned tokens. `respond(options, window)` returns the token list for one window, where - `window` is the index of the call's segment within the whole run. + `window` numbers the distinct mel segments in the order they are first seen, + so a re-decode of the same segment gets the same number. Use `noise=True` + in run_transcribe to make every window distinct. """ dims = SimpleNamespace(n_mels=80, n_audio_ctx=1500) is_multilingual = False num_languages = 99 # what an English-only checkpoint reports (n_vocab 51864) + no_speech_prob = 0.0 def __init__( self, @@ -65,39 +68,45 @@ def __init__( self.compression_ratio = compression_ratio self.calls: List[DecodingOptions] = [] self.batch_sizes: List[int] = [] - self._window = 0 + self._windows = {} def decode(self, mel, options: DecodingOptions): single = mel.ndim == 2 - n = 1 if single else mel.shape[0] + items = [mel] if single else [mel[i] for i in range(mel.shape[0])] self.calls.append(options) - self.batch_sizes.append(n) + self.batch_sizes.append(len(items)) results = [] - for _ in range(n): + for item in items: + key = np.array(item).tobytes() + window = self._windows.setdefault(key, len(self._windows)) results.append( DecodingResult( audio_features=None, language="en", - tokens=self.respond(options, self._window), + tokens=self.respond(options, window), avg_logprob=-0.1, - no_speech_prob=0.0, + no_speech_prob=self.no_speech_prob, temperature=options.temperature, - compression_ratio=self.compression_ratio(options, self._window), + compression_ratio=self.compression_ratio(options, window), ) ) - self._window += 1 return results[0] if single else results @pytest.fixture def run_transcribe(monkeypatch): - """Run transcribe() on `seconds` of silence against a StubModel.""" + """Run transcribe() on `seconds` of silence (or noise) against a stub model.""" - def run(model: StubModel, seconds: float, **kwargs): + def run(model: StubModel, seconds: float, noise: bool = False, **kwargs): monkeypatch.setattr( transcribe_module.ModelHolder, "get_model", lambda *a, **k: model ) - audio = np.zeros(int(SAMPLE_RATE * seconds), dtype=np.float32) + n_samples = int(SAMPLE_RATE * seconds) + if noise: + rng = np.random.default_rng(0) + audio = (0.1 * rng.standard_normal(n_samples)).astype(np.float32) + else: + audio = np.zeros(n_samples, dtype=np.float32) kwargs.setdefault("language", "en") return transcribe_module.transcribe(audio, **kwargs) diff --git a/tests/test_temperature_fallback.py b/tests/test_temperature_fallback.py new file mode 100644 index 0000000..1491a33 --- /dev/null +++ b/tests/test_temperature_fallback.py @@ -0,0 +1,74 @@ +"""Temperature fallback in transcribe(), sequential and batched.""" + +import pytest + +from tests.conftest import StubModel + +SCHEDULE = (0.0, 0.2, 0.4, 0.6, 0.8, 1.0) + + +@pytest.fixture +def tokens(tokenizer, ts): + return [ts(0), *tokenizer.encode(" hello"), ts(2)] + + +@pytest.mark.parametrize("batch_size", [1, 3]) +def test_single_temperature_means_no_fallback(tokens, run_transcribe, batch_size): + model = StubModel(lambda o, w: tokens, compression_ratio=lambda o, w: 3.0) + + run_transcribe( + model, seconds=90, noise=True, batch_size=batch_size, temperature=0.0 + ) + + assert {c.temperature for c in model.calls} == {0.0} + + +def test_batched_fallback_walks_schedule_for_failing_segments_only( + tokens, run_transcribe +): + # window 1 is too repetitive until t=0.4; the others pass straight away + def ratio(options, window): + return 3.0 if window == 1 and options.temperature < 0.4 else 1.0 + + model = StubModel(lambda o, w: tokens, compression_ratio=ratio) + + result = run_transcribe(model, seconds=90, noise=True, batch_size=3) + + assert [(c.temperature, n) for c, n in zip(model.calls, model.batch_sizes)] == [ + (0.0, 3), + (0.2, 1), + (0.4, 1), + ] + assert [s["temperature"] for s in result["segments"]] == [0.0, 0.4, 0.0] + + +def test_batched_fallback_keeps_last_result_when_schedule_runs_out( + tokens, run_transcribe +): + model = StubModel(lambda o, w: tokens, compression_ratio=lambda o, w: 3.0) + + result = run_transcribe(model, seconds=60, noise=True, batch_size=2) + + assert [c.temperature for c in model.calls] == list(SCHEDULE) + assert [s["temperature"] for s in result["segments"]] == [1.0, 1.0] + + +def test_silent_segments_do_not_fall_back(tokens, run_transcribe): + model = StubModel(lambda o, w: tokens, compression_ratio=lambda o, w: 3.0) + model.no_speech_prob = 0.9 + + run_transcribe(model, seconds=60, noise=True, batch_size=2, logprob_threshold=None) + + assert [c.temperature for c in model.calls] == [0.0] + + +def test_sequential_fallback_walks_schedule(tokens, run_transcribe): + def ratio(options, window): + return 3.0 if options.temperature < 0.4 else 1.0 + + model = StubModel(lambda o, w: tokens, compression_ratio=ratio) + + result = run_transcribe(model, seconds=30, noise=True, batch_size=1) + + assert [c.temperature for c in model.calls] == [0.0, 0.2, 0.4] + assert [s["temperature"] for s in result["segments"]] == [0.4] diff --git a/whisper_mlx/transcribe.py b/whisper_mlx/transcribe.py index a323854..8299ea2 100644 --- a/whisper_mlx/transcribe.py +++ b/whisper_mlx/transcribe.py @@ -212,11 +212,31 @@ def transcribe( if word_timestamps and task == "translate": warnings.warn("Word-level timestamps on translations may not be reliable.") + temperatures = ( + [temperature] if isinstance(temperature, (int, float)) else temperature + ) + + def needs_fallback(decode_result: DecodingResult) -> bool: + """Whether a result fails the quality checks and should be re-decoded.""" + if ( + no_speech_threshold is not None + and decode_result.no_speech_prob > no_speech_threshold + ): + return False # silence + if ( + compression_ratio_threshold is not None + and decode_result.compression_ratio > compression_ratio_threshold + ): + return True # too repetitive + if ( + logprob_threshold is not None + and decode_result.avg_logprob < logprob_threshold + ): + return True # average log probability is too low + return False + def decode_with_fallback(segment: mx.array) -> DecodingResult: """Decode a single segment with temperature fallback.""" - temperatures = ( - [temperature] if isinstance(temperature, (int, float)) else temperature - ) decode_result = None for t in temperatures: @@ -231,24 +251,7 @@ def decode_with_fallback(segment: mx.array) -> DecodingResult: options = DecodingOptions(**kwargs, temperature=t) decode_result = model.decode(segment, options) - - needs_fallback = False - if ( - compression_ratio_threshold is not None - and decode_result.compression_ratio > compression_ratio_threshold - ): - needs_fallback = True # too repetitive - if ( - logprob_threshold is not None - and decode_result.avg_logprob < logprob_threshold - ): - needs_fallback = True # average log probability is too low - if ( - no_speech_threshold is not None - and decode_result.no_speech_prob > no_speech_threshold - ): - needs_fallback = False # silence - if not needs_fallback: + if not needs_fallback(decode_result): break return decode_result @@ -256,55 +259,32 @@ def decode_with_fallback(segment: mx.array) -> DecodingResult: def decode_batch_with_fallback(segment_batch: mx.array) -> List[DecodingResult]: """Decode a batch of segments with per-segment temperature fallback. - Optimized: Collects all segments needing fallback and batches them together - instead of decoding each one individually (which destroys parallelism). + Walks the same temperature schedule as decode_with_fallback, but at each + step re-decodes only the segments that still fail, as one smaller batch. """ kwargs = {**decode_options} kwargs.pop("beam_size", None) kwargs.pop("patience", None) + # best_of samples several sequences per segment, which DecodingTask only + # supports for a single segment kwargs.pop("best_of", None) - # First pass: decode all segments at temperature 0 - options = DecodingOptions(**kwargs, temperature=0.0) - decode_results = model.decode(segment_batch, options) - - # Collect indices of segments needing fallback (instead of processing individually) - fallback_indices = [] - for i, decode_result in enumerate(decode_results): - needs_fallback = False - if ( - compression_ratio_threshold is not None - and decode_result.compression_ratio > compression_ratio_threshold - ): - needs_fallback = True - if ( - logprob_threshold is not None - and decode_result.avg_logprob < logprob_threshold - ): - needs_fallback = True - if ( - no_speech_threshold is not None - and decode_result.no_speech_prob > no_speech_threshold - ): - needs_fallback = False # Silence, no fallback needed - - if needs_fallback: - fallback_indices.append(i) - - # Batch all fallback segments together (instead of individual decoding) - if fallback_indices: - # Stack all segments needing fallback into one batch - fallback_segments = mx.stack([segment_batch[i] for i in fallback_indices], axis=0) - fallback_options = DecodingOptions(**kwargs, temperature=1.0) - fallback_results = model.decode(fallback_segments, fallback_options) - - # Ensure fallback_results is a list - if not isinstance(fallback_results, list): - fallback_results = [fallback_results] - - # Update decode_results with fallback results - for idx, fallback_result in zip(fallback_indices, fallback_results): - decode_results[idx] = fallback_result + n_segments = segment_batch.shape[0] + decode_results: List[Optional[DecodingResult]] = [None] * n_segments + pending = list(range(n_segments)) + + for t in temperatures: + batch = ( + segment_batch + if len(pending) == n_segments + else segment_batch[mx.array(pending)] + ) + results = model.decode(batch, DecodingOptions(**kwargs, temperature=t)) + for i, result in zip(pending, results): + decode_results[i] = result + pending = [i for i, result in zip(pending, results) if needs_fallback(result)] + if not pending: + break return decode_results