Skip to content

Add beam search decoding to whisper - #1454

Open
benediktgoodman wants to merge 3 commits into
ml-explore:mainfrom
benediktgoodman:main
Open

benediktgoodman wants to merge 3 commits into
ml-explore:mainfrom
benediktgoodman:main

Conversation

@benediktgoodman

Copy link
Copy Markdown

Implements beam search decoding for the Whisper example so it can be used with mlx-whisper. You already have the beam_size and patience options in mlx_whisper/decode.py, but DecodingTask raises NotImplementedError when beam_size is set. Hence I decided to follow OpenAIs torch implementation of the algorithm, but just using MLX methods instead.

Decoding

  • Adds BeamSearchDecoder as a separate object, ported from openai-whisper's torch implementation (dict-keyed sequence tracking for deduplication and scoring, per-beam top-beam_size+1 candidate pool, patience stop rule), with unfinished beams drained and padded correctly at finalize.
  • Inference.rearrange_kv_cache gathers only the self-attention cache. Cross-attention K/V rows are identical across beams of the same audio item, so gathering them copied identical rows (~1.2 GB per step on large models at beam 5). The result is bit-identical to a full gather of both caches.
  • ApplyTimestampRules fixes needed for beam search with timestamps: the boundary timestamp token is now included when masking past timestamps, and the timestamp-vs-text decision is computed on masked logits (as openai-whisper and transformers do). Previously it could mask every token after a closed timestamp pair.

Transcription

  • Adds beam_size in transcribe(), where it now works through**decode_options with no other changes. Documented in docstring.

Tests

Run from whisper/ with python -m unittest discover tests -v ...or using uv if you are a cool kid.

  • tests/test_beam_search_parity.py drives the MLX beam search decoder and openai-whisper's torch decoder on identical logits (beam sizes 1–3, multiple audio items, patience, early EOT termination, unfinished-beam drain) and asserts identical tokens, completion flags, cumulative log probabilities, KV-rearrange calls, and final candidates. Skips whenopenai-whisper is not installed.
  • tests/test_cross_kv_skip.py checks the optimized KV rearrange against a full-gather reference, bit-identical on synthetic caches.

--add added 3 commits September 25, 2026 23:33
tests/test_beam_search_parity.py drives the MLX beam search decoder
and openai-whisper's torch decoder on identical logits and asserts
identical tokens, completion flags, cumulative log probabilities,
KV-rearrange calls, and final candidates. Skips when openai-whisper
is not installed.

Also documents beam_size in transcribe(), where it now works through
**decode_options.
- tests/test_cross_kv_skip.py: the cross-attention KV skip in
  Inference.rearrange_kv_cache is bit-identical to a full gather
- ApplyTimestampRules: include the boundary timestamp when masking
  past timestamps, and decide the timestamp-vs-text comparison on the
  masked logits, as openai-whisper and transformers do
- drop stale file references from the rearrange_kv_cache docstring
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.

1 participant