Skip to content

whisper: add carry_initial_prompt to maintain context over sliding wi… - #1414

Open
Sahith59 wants to merge 2 commits into
ml-explore:mainfrom
Sahith59:feature/whisper-carry-prompt
Open

Sahith59 wants to merge 2 commits into
ml-explore:mainfrom
Sahith59:feature/whisper-carry-prompt

Conversation

@Sahith59

@Sahith59 Sahith59 commented Apr 8, 2026

Copy link
Copy Markdown

Description

This PR resolves the issue where context provided by the user's initial_prompt is lost/truncated as the cross-attention context window slides forward over long audio files.

Changes

  • Introduces the carry_initial_prompt keyword boolean to .transcribe().
  • Intercepts the prompt creation block inside transcribe.py. When enabled, it dynamically prepends the initial_prompt_tokens directly to the active prompt queue.
  • Calculates the window boundary natively via model.dims.n_text_ctx to correctly slice the previous_tokens so that the model evaluates without hitting dimension mismatch errors during Apple MLX evaluation.
  • Added native Apple MLX unit tests into test.py to test context-window bounds and propagation.
  • Default behavior rigorously maintains PyTorch backward compatibility (carry_initial_prompt=False).

Testing

Ran local conversion via setUpClass and verified native propagation logic inside the test suite. Formatted with black and passes pre-commit natively.

Resolves: #1410

@srmoorhead srmoorhead left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread whisper/test.py
Comment on lines +200 to +210
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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
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:]
)
)

@srmoorhead

Copy link
Copy Markdown

Follow-up after a full-length run: with carry_initial_prompt=True and a long prompt (192 tokens, near the 223-token limit), Whisper skipped stretches of speech. On a 3.5-hour recording it dropped 1,536 words across 44 gaps of 5 seconds or more, the longest 54 seconds. On 19 minutes of that audio: no prompt lost 10 words, a short names-only prompt carried lost 215, and the long prompt carried lost 515. It also wrote prompted names where they were not said. This looks like a property of carrying long prompts generally rather than of this PR's code, but users may want a note in the docstring that carrying a long prompt can cause skipped segments.

Tested and edited with Claude.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants