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
26 changes: 25 additions & 1 deletion tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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():
Expand Down
84 changes: 84 additions & 0 deletions tests/test_word_timestamps.py
Original file line number Diff line number Diff line change
@@ -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,
)
23 changes: 22 additions & 1 deletion whisper_mlx/transcribe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
-------
Expand All @@ -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:
Expand Down Expand Up @@ -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"]
Expand Down
Loading