From 4eadab967fe8f43ec75841e29b246b1614f6ce42 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 30 Sep 2026 21:21:32 +0000 Subject: [PATCH] Add word-level timestamps to batched decoding add_word_timestamps was only called when batch_size=1, so word_timestamps=True had no effect with batching (the default for LightningWhisperMLX): segments had no words and the SRT/VTT word options did nothing. Each window in a batch is now aligned against its own mel segment, using that window's seek for the time offset. hallucination_silence_threshold relies on moving seek window by window, which a batch cannot do; it now warns when combined with batch_size > 1 instead of being dropped silently. Tests run a random-weight model so the alignment pass is real. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA --- tests/conftest.py | 26 ++++++++++- tests/test_word_timestamps.py | 84 +++++++++++++++++++++++++++++++++++ whisper_mlx/transcribe.py | 23 +++++++++- 3 files changed, 131 insertions(+), 2 deletions(-) create mode 100644 tests/test_word_timestamps.py diff --git a/tests/conftest.py b/tests/conftest.py index fca99fa..7ba4120 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,20 +1,44 @@ -"""Shared fixtures: a stub model for driving transcribe() without weights.""" +"""Shared fixtures: stub and random-weight models for tests without downloads.""" import importlib from types import SimpleNamespace from typing import Callable, List +import mlx.core as mx import numpy as np import pytest from whisper_mlx.decoding import DecodingOptions, DecodingResult from whisper_mlx.tokenizer import get_tokenizer +from whisper_mlx.whisper import ModelDimensions, Whisper # transcribe() is re-exported by the package, so fetch the module itself transcribe_module = importlib.import_module("whisper_mlx.transcribe") SAMPLE_RATE = 16000 +# English-only vocabulary with a deliberately small network: fast on CPU +TINY_EN_DIMS = ModelDimensions( + n_mels=80, + n_audio_ctx=1500, + n_audio_state=64, + n_audio_head=2, + n_audio_layer=2, + n_vocab=51864, + n_text_ctx=448, + n_text_state=64, + n_text_head=2, + n_text_layer=2, +) + + +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) + mx.eval(model.parameters()) + return model + @pytest.fixture def tokenizer(): diff --git a/tests/test_word_timestamps.py b/tests/test_word_timestamps.py new file mode 100644 index 0000000..aca1845 --- /dev/null +++ b/tests/test_word_timestamps.py @@ -0,0 +1,84 @@ +"""Word-level timestamps from transcribe(), sequential and batched.""" + +import warnings + +import mlx.core as mx +import pytest + +from tests.conftest import TINY_EN_DIMS, StubModel +from whisper_mlx.whisper import Whisper + + +class CannedWhisper(Whisper): + """Random-weight Whisper: decode() returns canned tokens, alignment runs for real.""" + + def __init__(self, stub: StubModel): + super().__init__(TINY_EN_DIMS, mx.float16) + mx.eval(self.parameters()) + self.stub = stub + + def decode(self, mel, options): + return self.stub.decode(mel, options) + + +@pytest.fixture +def model(tokenizer, ts): + tokens = [ + ts(0), + *tokenizer.encode(" hello there"), + ts(2), + ts(2), + *tokenizer.encode(" general"), + ts(3), + ] + mx.random.seed(0) + return CannedWhisper(StubModel(lambda options, window: tokens)) + + +@pytest.mark.parametrize("batch_size", [1, 4]) +def test_segments_get_words_within_their_window(model, run_transcribe, batch_size): + result = run_transcribe( + model, seconds=45, batch_size=batch_size, word_timestamps=True + ) + + segments = [s for s in result["segments"] if s["text"]] + assert segments + for segment in segments: + assert "".join(w["word"] for w in segment["words"]) == segment["text"] + window_start = segment["seek"] / 100 + for word in segment["words"]: + assert window_start <= word["start"] <= word["end"] <= window_start + 30 + + +def test_batched_words_use_each_windows_offset(model, run_transcribe): + result = run_transcribe(model, seconds=65, batch_size=4, word_timestamps=True) + + first_word_starts = sorted( + {s["seek"]: s["words"][0]["start"] for s in result["segments"]}.items() + ) + assert [seek for seek, _ in first_word_starts] == [0, 3000, 6000] + for seek, start in first_word_starts: + assert seek / 100 <= start < seek / 100 + 30 + + +def test_batched_warns_that_hallucination_threshold_is_ignored(model, 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, + ) + + +def test_sequential_does_not_warn_about_hallucination_threshold(model, 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, + ) diff --git a/whisper_mlx/transcribe.py b/whisper_mlx/transcribe.py index c876b68..a323854 100644 --- a/whisper_mlx/transcribe.py +++ b/whisper_mlx/transcribe.py @@ -118,7 +118,7 @@ def transcribe( hallucination_silence_threshold: Optional[float] When word_timestamps is True, skip silent periods longer than this threshold (in seconds) - when a possible hallucination is detected + when a possible hallucination is detected. Only applied when batch_size is 1. Returns ------- @@ -133,6 +133,12 @@ def transcribe( f"batch_size={batch_size} may cause out-of-memory errors. " "Consider using batch_size <= 64." ) + if batch_size > 1 and hallucination_silence_threshold is not None: + # skipping silence needs the window-by-window seeking of batch_size=1 + warnings.warn( + "hallucination_silence_threshold is only applied when batch_size=1; " + "ignoring it." + ) if isinstance(audio, str): if not audio: @@ -467,6 +473,21 @@ def new_segment( ) ) + if word_timestamps: + add_word_timestamps( + segments=current_segments, + model=model, + tokenizer=tokenizer, + mel=mel_segments[batch_idx], + num_frames=segment_size, + prepend_punctuations=prepend_punctuations, + append_punctuations=append_punctuations, + last_speech_timestamp=last_speech_timestamp, + ) + last_word_end = _get_end(current_segments) + if last_word_end is not None: + last_speech_timestamp = last_word_end + if verbose: for segment in current_segments: start, end, text = segment["start"], segment["end"], segment["text"]