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
21 changes: 17 additions & 4 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
68 changes: 68 additions & 0 deletions tests/test_timestamp_rules.py
Original file line number Diff line number Diff line change
@@ -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
25 changes: 11 additions & 14 deletions tests/test_word_timestamps.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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):
Expand Down Expand Up @@ -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
)
13 changes: 6 additions & 7 deletions whisper_mlx/decoding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading