diff --git a/tests/test_beam_search.py b/tests/test_beam_search.py new file mode 100644 index 0000000..91f4631 --- /dev/null +++ b/tests/test_beam_search.py @@ -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 diff --git a/whisper_mlx/cli.py b/whisper_mlx/cli.py index c4e2c68..5248c68 100644 --- a/whisper_mlx/cli.py +++ b/whisper_mlx/cli.py @@ -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, diff --git a/whisper_mlx/lightning.py b/whisper_mlx/lightning.py index 9182dcf..d866773 100644 --- a/whisper_mlx/lightning.py +++ b/whisper_mlx/lightning.py @@ -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 ------- diff --git a/whisper_mlx/transcribe.py b/whisper_mlx/transcribe.py index 945471e..874430e 100644 --- a/whisper_mlx/transcribe.py +++ b/whisper_mlx/transcribe.py @@ -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. @@ -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(