Skip to content
BytePatterns

Target Passes in Speculative Decoding

MediumAI & ML#speculative-decoding#draft-and-verify~25m

Problem

Speculative decoding speeds up generation with a small draft model. Both models here are greedy and deterministic, given as dicts from the last token to the next token; a token missing from a dict produces "<eos>". Each round, the draft proposes k tokens one after another, then one pass of the large target model checks them: it accepts proposed tokens while they match what the target itself would produce, and at the first mismatch it keeps its own token instead. If all k match, the same pass adds one bonus token. Generate exactly n tokens after the prompt, stopping mid-round if needed, and return the tokens and the number of target passes used.

Examples

Input:  prompt = ["the"], draft guesses "a" after "on" where the target says "the", k = 4, n = 8
Output: (['cat', 'sat', 'on', 'the', 'cat', 'sat', 'on', 'the'], 2)
Why:    3 guesses are accepted and the target fixes the 4th, so each pass yields 4 tokens
Input:  prompt = ["the"], draft = target, k = 3, n = 8
Output: (['cat', 'sat', 'on', 'the', 'cat', 'sat', 'on', 'the'], 2)
Why:    a perfect draft yields k + 1 tokens per pass: 3 accepted plus the bonus
Input:  prompt = ["the"], draft = {}, k = 3, n = 3
Output: (['cat', 'sat', 'on'], 3)
Why:    edge case, a draft that is always wrong still gives the target's exact text, one token per pass

Hints

0 / 3

Stuck on the idea rather than the code? Speculative Decoding covers it.