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
86 changes: 86 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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)
86 changes: 86 additions & 0 deletions tests/test_transcribe_batched.py
Original file line number Diff line number Diff line change
@@ -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"
29 changes: 28 additions & 1 deletion whisper_mlx/transcribe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -419,13 +424,32 @@ 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),
result=result,
)
)
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]]
Expand All @@ -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),
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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),
Expand Down
Loading