Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 21 additions & 12 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)

Expand Down
74 changes: 74 additions & 0 deletions tests/test_temperature_fallback.py
Original file line number Diff line number Diff line change
@@ -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]
108 changes: 44 additions & 64 deletions whisper_mlx/transcribe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -231,80 +251,40 @@ 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

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

Expand Down
Loading