Repository navigation
Conversation
srmoorhead
left a comment
There was a problem hiding this comment.
Tested and edited with Claude.
Tested this and it fixes the problem for us: with condition_on_previous_text=False, initial_prompt otherwise only reaches the first 30-second window.
- The prompt logic matches openai/whisper#2343.
- On a 5-minute clip with this branch (large-v3), 2 of 4 prompted terms came out as prompted with
carry_initial_prompt=True, against 0 of 4 without it, at the same speed.
test_carry_initial_prompt passes even with the new prompt-building block reverted, so it doesn't cover the change. The suggestion below records options.prompt on each model.decode call for audio longer than 30 seconds and checks that every window's prompt starts with the initial prompt tokens with carry on, and that later windows lose them with carry off. It fails with the block reverted and passes with it.
| def test_carry_initial_prompt(self): | ||
| result = mlx_whisper.transcribe( | ||
| TEST_AUDIO, | ||
| path_or_hf_repo=MLX_FP32_MODEL_PATH, | ||
| fp16=False, | ||
| initial_prompt="A test prompt.", | ||
| carry_initial_prompt=True, | ||
| ) | ||
| self.assertIn("text", result) | ||
| self.assertGreater(len(result["text"]), 0) | ||
|
|
There was a problem hiding this comment.
| def test_carry_initial_prompt(self): | |
| result = mlx_whisper.transcribe( | |
| TEST_AUDIO, | |
| path_or_hf_repo=MLX_FP32_MODEL_PATH, | |
| fp16=False, | |
| initial_prompt="A test prompt.", | |
| carry_initial_prompt=True, | |
| ) | |
| self.assertIn("text", result) | |
| self.assertGreater(len(result["text"]), 0) | |
| def test_carry_initial_prompt(self): | |
| from mlx_whisper.tokenizer import get_tokenizer | |
| from mlx_whisper.transcribe import ModelHolder | |
| # Tile the short test clip to > 30s so transcribe() must decode | |
| # at least two windows. | |
| data = audio.load_audio(TEST_AUDIO) | |
| long_audio = np.tile(data, 10) | |
| model = ModelHolder.get_model(MLX_FP32_MODEL_PATH, mx.float32) | |
| tokenizer = get_tokenizer( | |
| model.is_multilingual, | |
| num_languages=model.num_languages, | |
| language="en", | |
| task="transcribe", | |
| ) | |
| initial_prompt = "A test prompt." | |
| initial_prompt_tokens = tokenizer.encode(" " + initial_prompt.strip()) | |
| original_decode = model.decode | |
| recorded_prompts = [] | |
| def record_prompt(mel, options): | |
| recorded_prompts.append(list(options.prompt or [])) | |
| return original_decode(mel, options) | |
| def transcribe_and_record(carry_initial_prompt): | |
| recorded_prompts.clear() | |
| model.decode = record_prompt | |
| try: | |
| mlx_whisper.transcribe( | |
| long_audio, | |
| path_or_hf_repo=MLX_FP32_MODEL_PATH, | |
| fp16=False, | |
| language="en", | |
| initial_prompt=initial_prompt, | |
| condition_on_previous_text=False, | |
| carry_initial_prompt=carry_initial_prompt, | |
| ) | |
| finally: | |
| del model.decode | |
| return list(recorded_prompts) | |
| carried = transcribe_and_record(carry_initial_prompt=True) | |
| self.assertGreaterEqual(len(carried), 2) | |
| for prompt in carried: | |
| self.assertEqual( | |
| prompt[: len(initial_prompt_tokens)], initial_prompt_tokens | |
| ) | |
| dropped = transcribe_and_record(carry_initial_prompt=False) | |
| self.assertEqual( | |
| dropped[0][: len(initial_prompt_tokens)], initial_prompt_tokens | |
| ) | |
| self.assertTrue( | |
| any( | |
| prompt[: len(initial_prompt_tokens)] != initial_prompt_tokens | |
| for prompt in dropped[1:] | |
| ) | |
| ) |
|
Follow-up after a full-length run: with Tested and edited with Claude. |
Description
This PR resolves the issue where context provided by the user's
initial_promptis lost/truncated as the cross-attention context window slides forward over long audio files.Changes
carry_initial_promptkeyword boolean to.transcribe().transcribe.py. When enabled, it dynamically prepends theinitial_prompt_tokensdirectly to the active prompt queue.model.dims.n_text_ctxto correctly slice theprevious_tokensso that the model evaluates without hitting dimension mismatch errors during Apple MLX evaluation.test.pyto test context-window bounds and propagation.carry_initial_prompt=False).Testing
Ran local conversion via
setUpClassand verified native propagation logic inside the test suite. Formatted withblackand passespre-commitnatively.Resolves: #1410