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
41 changes: 41 additions & 0 deletions tests/test_beam_search.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
"""Beam search is not implemented: say so instead of failing inconsistently."""

import pytest

from tests.conftest import StubModel
from whisper_mlx import LightningWhisperMLX
from whisper_mlx.cli import build_parser


@pytest.fixture
def stub(tokenizer, ts):
return StubModel(lambda o, w: [ts(0), *tokenizer.encode(" hello"), ts(2)])


@pytest.mark.parametrize("batch_size", [1, 4])
@pytest.mark.parametrize("option", [{"beam_size": 5}, {"patience": 1.0}])
def test_transcribe_rejects_beam_search_options(
stub, run_transcribe, batch_size, option
):
with pytest.raises(NotImplementedError, match="beam search"):
run_transcribe(stub, seconds=30, batch_size=batch_size, **option)


def test_unset_beam_search_options_are_fine(stub, run_transcribe):
result = run_transcribe(stub, seconds=30, beam_size=None, patience=None)

assert result["text"] == " hello"


def test_lightning_wrapper_rejects_beam_size():
whisper = LightningWhisperMLX(model="tiny")

with pytest.raises(NotImplementedError, match="beam search"):
whisper.transcribe("audio.mp3", beam_size=5)


def test_cli_has_no_patience_option(capsys):
with pytest.raises(SystemExit):
build_parser().parse_args(["audio.mp3", "--patience", "1.0"])

assert "--patience" in capsys.readouterr().err
6 changes: 0 additions & 6 deletions whisper_mlx/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,12 +118,6 @@ def build_parser():
default=5,
help="Number of candidates when sampling with non-zero temperature",
)
parser.add_argument(
"--patience",
type=float,
default=None,
help="Optional patience value to use in beam decoding, as in https://arxiv.org/abs/2204.05424, the default (1.0) is equivalent to conventional beam search",
)
parser.add_argument(
"--length-penalty",
type=float,
Expand Down
4 changes: 3 additions & 1 deletion whisper_mlx/lightning.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,9 @@ def transcribe(
- temperature, compression_ratio_threshold, logprob_threshold
- no_speech_threshold, condition_on_previous_text, initial_prompt
- prepend_punctuations, append_punctuations, clip_timestamps
- hallucination_silence_threshold, fp16, beam_size, patience, etc.
- hallucination_silence_threshold (batch_size=1 only), fp16
- suppress_tokens, best_of (batch_size=1 only), etc.
Beam search (beam_size, patience) is not implemented.

Returns
-------
Expand Down
9 changes: 8 additions & 1 deletion whisper_mlx/transcribe.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,9 @@ def transcribe(
to make it more likely to predict those word correctly.

decode_options: dict
Keyword arguments to construct `DecodingOptions` instances
Keyword arguments to construct `DecodingOptions` instances. Beam search
(`beam_size`, `patience`) is not implemented and raises NotImplementedError;
`best_of` only applies when batch_size is 1.

clip_timestamps: Union[str, List[float]]
Comma-separated list start,end,start,end,... timestamps (in seconds) of clips to process.
Expand All @@ -134,6 +136,11 @@ def transcribe(
f"batch_size={batch_size} may cause out-of-memory errors. "
"Consider using batch_size <= 64."
)
for option in ("beam_size", "patience"):
if decode_options.get(option) is not None:
raise NotImplementedError(
f"{option} is not supported: beam search is not implemented yet"
)
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(
Expand Down
Loading