diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..fca99fa --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,86 @@ +"""Shared fixtures: a stub model for driving transcribe() without weights.""" + +import importlib +from types import SimpleNamespace +from typing import Callable, List + +import numpy as np +import pytest + +from whisper_mlx.decoding import DecodingOptions, DecodingResult +from whisper_mlx.tokenizer import get_tokenizer + +# transcribe() is re-exported by the package, so fetch the module itself +transcribe_module = importlib.import_module("whisper_mlx.transcribe") + +SAMPLE_RATE = 16000 + + +@pytest.fixture +def tokenizer(): + return get_tokenizer(False) # English-only (gpt2) tokenizer + + +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. + """ + + dims = SimpleNamespace(n_mels=80, n_audio_ctx=1500) + is_multilingual = False + num_languages = 99 # what an English-only checkpoint reports (n_vocab 51864) + + def __init__( + self, + respond: Callable[[DecodingOptions, int], List[int]], + compression_ratio: Callable[[DecodingOptions, int], float] = lambda o, w: 1.0, + ): + self.respond = respond + self.compression_ratio = compression_ratio + self.calls: List[DecodingOptions] = [] + self.batch_sizes: List[int] = [] + self._window = 0 + + def decode(self, mel, options: DecodingOptions): + single = mel.ndim == 2 + n = 1 if single else mel.shape[0] + self.calls.append(options) + self.batch_sizes.append(n) + results = [] + for _ in range(n): + results.append( + DecodingResult( + audio_features=None, + language="en", + tokens=self.respond(options, self._window), + avg_logprob=-0.1, + no_speech_prob=0.0, + temperature=options.temperature, + compression_ratio=self.compression_ratio(options, self._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.""" + + def run(model: StubModel, seconds: float, **kwargs): + monkeypatch.setattr( + transcribe_module.ModelHolder, "get_model", lambda *a, **k: model + ) + audio = np.zeros(int(SAMPLE_RATE * seconds), dtype=np.float32) + kwargs.setdefault("language", "en") + return transcribe_module.transcribe(audio, **kwargs) + + return run + + +@pytest.fixture +def ts(tokenizer): + """Token id of the timestamp `seconds` into a window, e.g. ts(2.0) -> <|2.00|>.""" + return lambda seconds: tokenizer.timestamp_begin + round(seconds / 0.02) diff --git a/tests/test_transcribe_batched.py b/tests/test_transcribe_batched.py new file mode 100644 index 0000000..1ded5cc --- /dev/null +++ b/tests/test_transcribe_batched.py @@ -0,0 +1,86 @@ +"""Batched decoding (batch_size > 1) in transcribe().""" + +from tests.conftest import StubModel + + +def test_keeps_text_after_last_timestamp_pair(tokenizer, ts, run_transcribe): + # Each window ends mid-sentence: "<|0.00|> hello <|2.00|><|2.00|> world" + tokens = [ + ts(0), + *tokenizer.encode(" hello"), + ts(2), + ts(2), + *tokenizer.encode(" world"), + ] + model = StubModel(lambda options, window: tokens) + + result = run_transcribe(model, seconds=65, batch_size=4) + + assert result["text"].count("world") == 3 + world = [s for s in result["segments"] if "world" in s["text"]] + # the unfinished segment runs from its timestamp to the end of its window + assert [(s["start"], s["end"]) for s in world] == [ + (2.0, 30.0), + (32.0, 60.0), + (62.0, 65.0), + ] + + +def test_segments_record_their_own_window_seek(tokenizer, ts, run_transcribe): + tokens = [ + ts(0), + *tokenizer.encode(" hello"), + ts(2), + ts(2), + *tokenizer.encode(" world"), + ts(3), + ] + model = StubModel(lambda options, window: tokens) + + result = run_transcribe(model, seconds=65, batch_size=4) + + assert [s["seek"] for s in result["segments"]] == [0, 0, 3000, 3000, 6000, 6000] + assert [s["id"] for s in result["segments"]] == list(range(6)) + + +def test_single_timestamp_ending_is_unchanged(tokenizer, ts, run_transcribe): + tokens = [ + ts(0), + *tokenizer.encode(" hello"), + ts(2), + ts(2), + *tokenizer.encode(" world"), + ts(3), + ] + model = StubModel(lambda options, window: tokens) + + result = run_transcribe(model, seconds=30, batch_size=2) + + assert [(s["start"], s["end"], s["text"]) for s in result["segments"]] == [ + (0.0, 2.0, " hello"), + (2.0, 3.0, " world"), + ] + + +def test_trailing_timestamp_without_text_adds_no_segment(tokenizer, ts, run_transcribe): + tokens = [ts(0), *tokenizer.encode(" hello"), ts(2), ts(2)] + model = StubModel(lambda options, window: tokens) + + result = run_transcribe(model, seconds=30, batch_size=2) + + assert [(s["start"], s["end"], s["text"]) for s in result["segments"]] == [ + (0.0, 2.0, " hello") + ] + + +def test_text_without_timestamp_pairs_spans_the_window(tokenizer, ts, run_transcribe): + tokens = [ts(0), *tokenizer.encode(" hello world")] + model = StubModel(lambda options, window: tokens) + + result = run_transcribe(model, seconds=45, batch_size=2) + + assert [(s["start"], s["end"]) for s in result["segments"]] == [ + (0.0, 30.0), + (30.0, 45.0), + ] + assert result["text"] == " hello world hello world" diff --git a/whisper_mlx/transcribe.py b/whisper_mlx/transcribe.py index bb9bd52..c876b68 100644 --- a/whisper_mlx/transcribe.py +++ b/whisper_mlx/transcribe.py @@ -319,7 +319,12 @@ def decode_batch_with_fallback(segment_batch: mx.array) -> List[DecodingResult]: initial_prompt_tokens = [] def new_segment( - *, start: float, end: float, tokens: mx.array, result: DecodingResult + *, + seek: int, + start: float, + end: float, + tokens: mx.array, + result: DecodingResult, ): tokens = tokens.tolist() text_tokens = [token for token in tokens if token < tokenizer.eot] @@ -419,6 +424,7 @@ def new_segment( ) current_segments.append( new_segment( + seek=segment_seek, start=time_offset + start_timestamp_pos * time_precision, end=time_offset + end_timestamp_pos * time_precision, tokens=mx.array(sliced_tokens), @@ -426,6 +432,24 @@ def new_segment( ) ) last_slice = current_slice + + # The one-window path re-decodes an unfinished last segment + # from its start timestamp. A batch has already moved on, so + # keep the text and end it at the window boundary instead. + remainder = tokens[last_slice:] + if np.any(remainder < tokenizer.eot): + start = time_offset + ( + remainder[0].item() - tokenizer.timestamp_begin + ) * time_precision + current_segments.append( + new_segment( + seek=segment_seek, + start=start, + end=max(start, time_offset + segment_duration), + tokens=mx.array(remainder), + result=result, + ) + ) else: duration = segment_duration timestamps = tokens[timestamp_tokens.nonzero()[0]] @@ -435,6 +459,7 @@ def new_segment( current_segments.append( new_segment( + seek=segment_seek, start=time_offset, end=time_offset + duration, tokens=mx.array(tokens), @@ -556,6 +581,7 @@ def next_words_segment(segments: List[dict]) -> Optional[dict]: ) current_segments.append( new_segment( + seek=seek, start=time_offset + start_timestamp_pos * time_precision, end=time_offset + end_timestamp_pos * time_precision, @@ -589,6 +615,7 @@ def next_words_segment(segments: List[dict]) -> Optional[dict]: current_segments.append( new_segment( + seek=seek, start=time_offset, end=time_offset + duration, tokens=mx.array(tokens),