Add beam search decoding to whisper - #1454
Open
benediktgoodman wants to merge 3 commits into
Open
benediktgoodman wants to merge 3 commits into
benediktgoodman wants to merge 3 commits into
Conversation
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Implements beam search decoding for the Whisper example so it can be used with mlx-whisper. You already have the
beam_sizeandpatienceoptions in mlx_whisper/decode.py, butDecodingTaskraisesNotImplementedErrorwhenbeam_sizeis set. Hence I decided to follow OpenAIs torch implementation of the algorithm, but just using MLX methods instead.Decoding
BeamSearchDecoderas a separate object, ported from openai-whisper's torch implementation (dict-keyed sequence tracking for deduplication and scoring, per-beam top-beam_size+1candidate pool, patience stop rule), with unfinished beams drained and padded correctly atfinalize.Inference.rearrange_kv_cachegathers 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.ApplyTimestampRulesfixes 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
beam_sizeintranscribe(), where it now works through**decode_optionswith no other changes. Documented in docstring.Tests
Run from
whisper/withpython -m unittest discover tests -v...or using uv if you are a cool kid.tests/test_beam_search_parity.pydrives 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-whisperis not installed.tests/test_cross_kv_skip.pychecks the optimized KV rearrange against a full-gather reference, bit-identical on synthetic caches.