From 5083a82603af60fd51c8689ad395ddd97c4a1b2f Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 21:26:12 +0000 Subject: [PATCH] Stop timestamps going backwards during decoding ApplyTimestampRules collected the positions of timestamp tokens in the sampled sequence rather than their values, so the mask slice timestamp_begin: was always empty and the rule never applied: the model could emit a timestamp earlier than the previous one, or close a segment at the timestamp it opened with. Use the token values, as OpenAI's reference implementation does: mask every timestamp below the last one, and the last one itself unless the last token is a lone timestamp closing a segment (so "<|2.00|><|2.00|>" can still close one segment and open the next). Tests: unit checks on the mask, and a decode with a random-weight model asserting timestamps never decrease. Random test models are now cast to float16 like the published checkpoints. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA --- tests/conftest.py | 21 ++++++++--- tests/test_timestamp_rules.py | 68 +++++++++++++++++++++++++++++++++++ tests/test_word_timestamps.py | 25 ++++++------- whisper_mlx/decoding.py | 13 ++++--- 4 files changed, 102 insertions(+), 25 deletions(-) create mode 100644 tests/test_timestamp_rules.py diff --git a/tests/conftest.py b/tests/conftest.py index a06c1d0..1643f31 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -7,6 +7,7 @@ import mlx.core as mx import numpy as np import pytest +from mlx.utils import tree_map from whisper_mlx.decoding import DecodingOptions, DecodingResult from whisper_mlx.tokenizer import get_tokenizer @@ -32,14 +33,26 @@ ) -def random_whisper(seed: int = 0, dims: ModelDimensions = TINY_EN_DIMS) -> Whisper: - """A Whisper model with random weights, deterministic for a given seed.""" - mx.random.seed(seed) - model = Whisper(dims, mx.float16) +def to_float16(model: Whisper) -> Whisper: + """Cast weights to float16, as the published MLX checkpoints are.""" + model.update( + tree_map( + lambda p: ( + p.astype(mx.float16) if mx.issubdtype(p.dtype, mx.floating) else p + ), + model.parameters(), + ) + ) mx.eval(model.parameters()) 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.""" + mx.random.seed(seed) + return to_float16(Whisper(dims, mx.float16)) + + @pytest.fixture def tokenizer(): return get_tokenizer(False) # English-only (gpt2) tokenizer diff --git a/tests/test_timestamp_rules.py b/tests/test_timestamp_rules.py new file mode 100644 index 0000000..4bac9e4 --- /dev/null +++ b/tests/test_timestamp_rules.py @@ -0,0 +1,68 @@ +"""ApplyTimestampRules: timestamp tokens must come in order and in pairs.""" + +import mlx.core as mx +import numpy as np +import pytest + +from tests.conftest import random_whisper +from whisper_mlx.decoding import ApplyTimestampRules, DecodingOptions + +N_VOCAB = 51864 + + +@pytest.fixture +def allowed(tokenizer): + """Which tokens the rules leave open after `sampled` (sot is the only prompt).""" + rules = ApplyTimestampRules( + tokenizer, sample_begin=1, max_initial_timestamp_index=None + ) + + def run(sampled): + tokens = mx.array([[tokenizer.sot, *sampled]]) + return ~np.isneginf(np.array(rules.apply(mx.zeros((1, N_VOCAB)), tokens))[0]) + + return run + + +def test_timestamp_after_text_must_be_later(allowed, tokenizer, ts): + ok = allowed([ts(1.0), *tokenizer.encode(" hi")]) + + assert not ok[ts(0.2)] + assert not ok[ts(1.0)] # a segment can't end where it started + assert ok[ts(1.02)] + + +def test_single_closing_timestamp_may_repeat_to_open_next_segment( + allowed, tokenizer, ts +): + ok = allowed([ts(1.0), *tokenizer.encode(" hi"), ts(2.0)]) + + assert not ok[ts(1.98)] + assert ok[ts(2.0)] + assert ok[ts(2.5)] + + +def test_no_timestamp_straight_after_a_pair(allowed, tokenizer, ts): + ok = allowed([ts(1.0), *tokenizer.encode(" hi"), ts(2.0), ts(2.0)]) + + assert not ok[tokenizer.timestamp_begin :].any() + + +def test_decoded_timestamps_never_go_backwards(tokenizer): + model = random_whisper(seed=3) + mx.random.seed(0) + mel = mx.random.normal((2, 3000, 80)).astype(mx.float16) + + results = model.decode(mel, DecodingOptions(language="en", sample_len=96)) + + for result in results: + previous, text_since = None, False + for token in result.tokens: + if token >= tokenizer.timestamp_begin: + if previous is not None: + assert token >= previous # never earlier than the last one + if text_since: + assert token > previous # a segment has nonzero length + previous, text_since = token, False + elif token < tokenizer.eot: + text_since = True diff --git a/tests/test_word_timestamps.py b/tests/test_word_timestamps.py index aca1845..59478cd 100644 --- a/tests/test_word_timestamps.py +++ b/tests/test_word_timestamps.py @@ -5,7 +5,7 @@ import mlx.core as mx import pytest -from tests.conftest import TINY_EN_DIMS, StubModel +from tests.conftest import TINY_EN_DIMS, StubModel, to_float16 from whisper_mlx.whisper import Whisper @@ -14,7 +14,7 @@ class CannedWhisper(Whisper): def __init__(self, stub: StubModel): super().__init__(TINY_EN_DIMS, mx.float16) - mx.eval(self.parameters()) + to_float16(self) self.stub = stub def decode(self, mel, options): @@ -61,24 +61,21 @@ def test_batched_words_use_each_windows_offset(model, run_transcribe): assert seek / 100 <= start < seek / 100 + 30 -def test_batched_warns_that_hallucination_threshold_is_ignored(model, run_transcribe): +@pytest.fixture +def stub(tokenizer, ts): + return StubModel(lambda o, w: [ts(0), *tokenizer.encode(" hello"), ts(2)]) + + +def test_batched_warns_that_hallucination_threshold_is_ignored(stub, run_transcribe): with pytest.warns(UserWarning, match="hallucination_silence_threshold"): run_transcribe( - model, - seconds=30, - batch_size=2, - word_timestamps=True, - hallucination_silence_threshold=2.0, + stub, seconds=30, batch_size=2, hallucination_silence_threshold=2.0 ) -def test_sequential_does_not_warn_about_hallucination_threshold(model, run_transcribe): +def test_sequential_does_not_warn_about_hallucination_threshold(stub, run_transcribe): with warnings.catch_warnings(): warnings.simplefilter("error") run_transcribe( - model, - seconds=30, - batch_size=1, - word_timestamps=True, - hallucination_silence_threshold=2.0, + stub, seconds=30, batch_size=1, hallucination_silence_threshold=2.0 ) diff --git a/whisper_mlx/decoding.py b/whisper_mlx/decoding.py index 13b465e..fdf163a 100644 --- a/whisper_mlx/decoding.py +++ b/whisper_mlx/decoding.py @@ -369,16 +369,15 @@ def apply(self, logits: mx.array, tokens: mx.array) -> mx.array: else: # cannot be normal text tokens mask[k, : self.tokenizer.eot] = -np.inf - timestamps = [ - i for i, v in enumerate(seq) if v > self.tokenizer.timestamp_begin - ] + timestamps = [t for t in seq if t >= self.tokenizer.timestamp_begin] if len(timestamps) > 0: # timestamps shouldn't decrease; forbid timestamp tokens smaller than the last # also force each segment to have a nonzero length, to prevent infinite looping - last_timestamp = timestamps[-1] - if not last_timestamp or penultimate_was_timestamp: - last_timestamp += 1 - mask[k, self.tokenizer.timestamp_begin : last_timestamp] = -np.inf + if last_was_timestamp and not penultimate_was_timestamp: + timestamp_last = timestamps[-1] # may repeat to close the pair + else: + timestamp_last = timestamps[-1] + 1 + mask[k, self.tokenizer.timestamp_begin : timestamp_last] = -np.inf if len(tokens[0]) == self.sample_begin: # suppress generating non-timestamp tokens at the beginning