Skip to content

Rewrite speculative decoding so it runs and matches the target model - #18

Merged
CodeWithBehnam merged 1 commit into
mainfrom
claude/lucid-franklin-gsu7v4
Sep 30, 2026
Merged

CodeWithBehnam merged 1 commit into
mainfrom
claude/lucid-franklin-gsu7v4

Conversation

@CodeWithBehnam

Copy link
Copy Markdown
Owner

What does this PR do?

Closes #8.

SpeculativeDecoder crashed 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:

  • encoded the audio again for both models on every draft/verify round, with no key/value cache
  • always built 80-band spectrograms (large-v3 and turbo need 128)
  • used one tokenizer for models with different vocabularies

Now:

  • Audio is encoded once per model per window. Both decoders keep a key/value cache, rolled back to the accepted tokens after each round.
  • The drafted tokens are checked in one target pass. The output is exactly the target model's greedy, timestamp-free decoding; the tests check this 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. The old tiny → large-v3 default can't work, because the two use different special-token ids.
  • SpeculativeDecoder.transcribe() handles the 30 s windows. 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. 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. MultiHeadAttention sliced mask[:n, :n], which only works without a cache or for one token at a time. It now takes the last n rows over all keys. For the existing call patterns (no cache, or one token at a time) the result is identical.

Also:

  • VADProcessor computes 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_transcribe never ran chunks in parallel; it is deprecated in favour of transcribe(batch_size=N).

How was this tested?

  • Tested with audio file(s)
  • Ran existing tests (pytest): 55 passed, 1 skipped
  • Tested CLI (vayu audio.mp3)

tests/test_speculative.py, all with random-weight models:

  • Output equals target.decode(...) token for token, for 1, 3 and 4 drafted tokens.
  • An identical draft model is accepted 100% of the time.
  • A vocabulary mismatch is rejected.
  • 80-band draft with 128-band target, and the end times for a 35 s input.
  • A multi-token step on the cache matches a full pass; this fails on main with a broadcast error.
  • The VAD matches the old implementation.
  • The deprecation warning appears.

Not verified: speed on real models. HuggingFace is unreachable from where this was built.

claude-review will fail as on #12 (the repository's CLAUDE_CODE_OAUTH_TOKEN secret, see this comment).

🤖 Generated with Claude Code

https://claude.ai/code/session_01TMXYMqgLAykqRApmbRfpTA


Generated by Claude Code

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
@CodeWithBehnam
CodeWithBehnam merged commit f7e8126 into main Sep 30, 2026
3 of 4 checks passed
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.

speculative.py: speculative decoding crashes, and the helpers don't do what they claim

2 participants