Rewrite speculative decoding so it runs and matches the target model - #18
Merged
Merged
Conversation
SpeculativeDecoder crashed on the first token (the decoder returns a (logits, kv_cache, cross_qk) tuple that was indexed as logits). Beyond that it re-encoded audio for both models on every draft/verify round with no kv cache, always built 80-band spectrograms (large-v3/turbo need 128), and used one tokenizer for models with different vocabularies. Now: - Audio is encoded once per model per window; both decoders keep a kv cache that is rolled back to the accepted tokens after each round. - Drafts are verified in one target pass; the output is exactly the target's greedy, timestamp-free decoding (tested token for token). - Each model gets a spectrogram with its own number of mel bands. - Draft and target must share a vocabulary; the default pair is now distil-large-v3 -> large-v3 (tiny -> large-v3 cannot work: their special-token ids differ). - SpeculativeDecoder.transcribe() does the windowing; the last window's end time is the end of the audio, not start + 30 s. - Stats report tokens per target pass instead of a "speed-up" figure; speedup_factor is kept as an alias. Verifying several tokens on top of a kv cache needs the matching rows of the causal mask; MultiHeadAttention sliced mask[:n, :n], which only works without a cache or for a single token. It now slices the last n rows over all keys, which is identical for the existing call patterns. Also: VADProcessor computes frame energies with a running sum instead of a Python loop (same regions as before, tested against the old code), and parallel_chunk_transcribe, which never ran chunks in parallel, is deprecated in favour of transcribe(batch_size=N). Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA
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.
What does this PR do?
Closes #8.
SpeculativeDecodercrashed on the first token: the decoder returns a(logits, kv_cache, cross_qk)tuple that was indexed as if it were the logits. Beyond that it:Now:
SpeculativeDecoder.transcribe()handles the 30 s windows. The last window's end time is the end of the audio, not start + 30 s.speedup_factoris kept as an alias. The docstring no longer promises 2-3x.Causal mask (
whisper.py): checking several tokens on top of a key/value cache needs the matching rows of the causal mask.MultiHeadAttentionslicedmask[:n, :n], which only works without a cache or for one token at a time. It now takes the lastnrows over all keys. For the existing call patterns (no cache, or one token at a time) the result is identical.Also:
VADProcessorcomputes frame energies from a running sum instead of a Python loop. It produces the same regions as before; a test compares it with the old implementation.parallel_chunk_transcribenever ran chunks in parallel; it is deprecated in favour oftranscribe(batch_size=N).How was this tested?
pytest): 55 passed, 1 skippedvayu audio.mp3)tests/test_speculative.py, all with random-weight models:target.decode(...)token for token, for 1, 3 and 4 drafted tokens.mainwith a broadcast error.Not verified: speed on real models. HuggingFace is unreachable from where this was built.
claude-reviewwill fail as on #12 (the repository'sCLAUDE_CODE_OAUTH_TOKENsecret, see this comment).🤖 Generated with Claude Code
https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA
Generated by Claude Code