Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

26. Speculative decoding

In this chapter

  • Why checking several tokens costs a memory-bound model about as much as generating one.
  • Draft-then-verify: greedy speculation, and rejection sampling that keeps the output distribution exactly the target model's.
  • The expected speedup as a function of acceptance rate, draft length and draft cost.
  • Rolling back the KV cache, and the zoo of draft methods: small models, early exits, n-gram lookup, Medusa, EAGLE and multi-token prediction.

You will build

accept_or_correct and speculative_generate in engine/speculative.py: lossless speculative decoding for any two models that share a tokenizer.

Time: 4-6 hours. GPU: recommended for speed measurements.

Verification is almost free

At batch 1, a decode step reads every weight to produce one token’s logits (Chapter 10). Feeding the model five tokens instead of one reads the same weights and returns five positions’ logits: the step is still memory-bound and takes nearly the same time. Chapter 24 used this by batching different users. Speculative decoding uses it along the sequence of one user.

The catch: to feed five tokens, you need to know them, and you only know the next token after computing it. Unless you guess. If a cheap draft model guesses the next $k$ tokens, the expensive target model can check all of them in one forward pass. Where its own predictions agree, those tokens are done. At the first disagreement, the target’s own prediction replaces the guess, so every round produces at least one token, and up to $k+1$.

Greedy speculation

With greedy decoding, the rule is simple. Feed the target [last, d1, d2, d3, d4]. Its logits at each position give its own choice for the next token:

target input:       last    d1     d2     d3     d4
target's argmax:     t1     t2     t3     t4     t5
draft proposed:      d1     d2     d3     d4

Accept $d_1$ if $t_1 = d_1$, then $d_2$ if $t_2 = d_2$, and so on. At the first mismatch, take the target’s token $t_i$ instead and stop. If all four match, $t_5$ is a free bonus token. Every committed token is exactly what greedy decoding of the target alone would have produced, so the output is identical; only the number of target calls changes.

Sampling: rejection sampling keeps the distribution exact

With temperature sampling, “agree with the argmax” is too strict and would also change the output distribution. Leviathan et al. (2023) and Chen et al. (2023) showed how to accept draft tokens so that every committed token is distributed exactly as if sampled from the target.

Let $q$ be the draft’s distribution at a position, $p$ the target’s, and $x \sim q$ the draft’s token:

  1. Accept $x$ with probability $\min!\big(1, p(x)/q(x)\big)$.
  2. Otherwise sample a replacement from the residual distribution $r(y) = \dfrac{\max(0,\ p(y) - q(y))}{\sum_z \max(0,\ p(z) - q(z))}$, and stop the round.

Why this produces exactly $p$: the probability of outputting $y$ through acceptance is $q(y) \cdot \min(1, p(y)/q(y)) = \min(q(y), p(y))$. The probability of rejecting is $1 - \sum_z \min(q(z), p(z)) = \sum_z \max(0, p(z) - q(z))$, which is exactly the residual’s normalizer. So rejection contributes $\max(0, p(y) - q(y))$. Adding the two:

$$ \min(p(y), q(y)) + \max(0,\ p(y) - q(y)) = p(y). \quad\checkmark $$

Worked example. Vocabulary of three tokens, $p = [0.6, 0.3, 0.1]$, $q = [0.2, 0.3, 0.5]$. The draft proposes token 2 half the time, but the target accepts it with probability $0.1/0.5 = 0.2$. Token 0 is accepted always ($0.6/0.2 > 1$). The residual is $\max(0, p - q) = [0.4, 0, 0]$: every rejection becomes token 0. Total probability of token 0: $0.2$ (proposed and accepted) $+ (0.5 \times 0.8)$ (token 2 rejected) $= 0.6$. The milestone test draws 20,000 samples and checks all three frequencies.

def accept_or_correct(p, q, token, generator=None):
    """One verification step for draft token `token` drawn from q.  (Your engine: Chapter 26)

    Accept with probability min(1, p[token] / q[token]). On rejection, draw from the
    residual max(p - q, 0) / sum(max(p - q, 0)). Returns (token, accepted).
    The two cases together produce exactly p: accepted mass min(p, q) plus residual mass
    max(p - q, 0) add up to p at every token.
    """
    ratio = (p[token] / q[token].clamp_min(1e-20)).clamp(max=1.0)
    if torch.rand((), generator=generator, device=p.device) < ratio:
        return token, True
    residual = (p - q).clamp_min(0)
    if residual.sum() <= 0:          # p == q numerically: nothing to correct
        residual = p
    return int(torch.multinomial(residual / residual.sum(), 1, generator=generator)), False
inline std::pair<size_t, bool> accept_or_correct(const std::vector<double>& p, const std::vector<double>& q, size_t x, Rng& rng) {
    if (rng.uniform() < std::min(1.0, p[x] / q[x])) return {x, true};
    std::vector<double> residual(p.size());
    double total = 0;
    for (size_t i = 0; i < p.size(); ++i) total += (residual[i] = std::max(0.0, p[i] - q[i]));
    double r = rng.uniform() * total;
    for (size_t i = 0; i < p.size(); ++i) { if (r < residual[i]) return {i, false}; r -= residual[i]; }
    return {x, false};
}
#![allow(unused)]
fn main() {
/// Draft token `x` came from q. Accept with probability min(1, p[x]/q[x]); otherwise draw
/// from the leftover distribution max(p - q, 0), renormalized. The result is distributed as p.
pub fn accept_or_correct(p: &[f64], q: &[f64], x: usize, rng: &mut Rng) -> (usize, bool) {
    if (rng.uniform() as f64) < (p[x] / q[x]).min(1.0) {
        return (x, true);
    }
    let residual: Vec<f64> = p.iter().zip(q).map(|(a, b)| (a - b).max(0.0)).collect();
    let total: f64 = residual.iter().sum();
    let mut r = rng.uniform() as f64 * total;
    for (i, mass) in residual.iter().enumerate() {
        if r < *mass {
            return (i, false);
        }
        r -= mass;
    }
    (residual.iter().rposition(|m| *m > 0.0).unwrap_or(x), false)
}
}

How much faster?

Suppose each draft token is accepted independently with probability $\alpha$ (a simplification, but a useful one). A round accepts $i$ tokens with probability $\alpha^i(1-\alpha)$ and always adds one target token, so the expected tokens per round is

$$ E[\text{tokens}] = 1 + \alpha + \alpha^2 + \dots + \alpha^k = \frac{1 - \alpha^{k+1}}{1 - \alpha}. $$

A round costs one target forward plus $k$ draft forwards. If a draft step costs a fraction $c$ of a target step, the speedup over plain decoding is

$$ \text{speedup} = \frac{1 - \alpha^{k+1}}{(1 - \alpha)(1 + k c)} . $$

With $\alpha = 0.8$, $k = 4$ and $c = 0.05$ (a draft about 20× smaller), that’s $3.36 / 1.2 = 2.8\times$. With $\alpha = 0.5$, it’s $1.94 / 1.2 = 1.6\times$. Larger $k$ helps only while $\alpha$ is high: the chance of reaching the $k$-th guess shrinks geometrically, while its cost doesn’t.

def expected_tokens_per_round(alpha, k):
    """With per-token acceptance probability alpha, a round yields (1 - alpha^(k+1)) / (1 - alpha) tokens."""
    if alpha >= 1:
        return k + 1
    return (1 - alpha ** (k + 1)) / (1 - alpha)

Explore: acceptance, draft length and speedup

Set the acceptance rate, the draft length and the draft's relative cost. See the expected tokens per round, the speedup curve over k, and a simulated run of rounds.

The algorithm, with cache rollback

@torch.inference_mode()
def speculative_generate(target, draft, prompt_ids, max_new_tokens, k=4, temperature=0.0,
                         stop_ids=(), generator=None):
    """Generate with a draft model. Both models must share the tokenizer.  (Your engine: Chapter 26)

    Invariant at the top of each round: both caches hold every committed token except the
    newest one, `last`, which has been chosen but not yet fed to either model.
    Returns (new_token_ids, stats).
    """
    device = next(target.parameters()).device
    ids = torch.as_tensor(prompt_ids, device=device).view(1, -1)
    capacity = ids.shape[1] + max_new_tokens + k + 1
    t_cache, d_cache = target.new_cache(1, capacity), draft.new_cache(1, capacity)
    target(ids[:, :-1], t_cache) if ids.shape[1] > 1 else None
    draft(ids[:, :-1], d_cache) if ids.shape[1] > 1 else None
    last = ids[:, -1:]
    output, rounds, accepted_total, proposed_total = [], 0, 0, 0
    while len(output) < max_new_tokens:
        rounds += 1
        base = t_cache.length
        # 1. Draft proposes k tokens, one cheap forward each.
        proposals, q_dists, token = [], [], last
        for _ in range(k):
            logits = draft(token, d_cache)[:, -1]
            if temperature == 0:
                token = logits.argmax(-1, keepdim=True)
            else:
                q = _probs(logits, temperature)[0]
                q_dists.append(q)
                token = torch.multinomial(q, 1, generator=generator)[None]
            proposals.append(token)
        # 2. Target scores last + all proposals in ONE forward: k+1 next-token distributions.
        block = torch.cat([last] + proposals, dim=1)
        target_logits = target(block, t_cache)[0]                       # [k+1, V]
        # 3. Accept a prefix of the proposals, then add one corrected or bonus token.
        new_tokens, all_accepted = [], True
        for i, proposal in enumerate(proposals):
            token_id = int(proposal)
            if temperature == 0:
                best = int(target_logits[i].argmax())
                ok, chosen = best == token_id, best
            else:
                chosen, ok = accept_or_correct(_probs(target_logits[i], temperature), q_dists[i], token_id, generator)
            new_tokens.append(chosen)
            if not ok:
                all_accepted = False
                break
        accepted = len(new_tokens) if all_accepted else len(new_tokens) - 1
        if all_accepted:                                                # all accepted: free bonus token
            bonus = target_logits[k]
            chosen = int(bonus.argmax()) if temperature == 0 else int(
                torch.multinomial(_probs(bonus, temperature), 1, generator=generator))
            new_tokens.append(chosen)
        accepted_total += accepted
        proposed_total += k
        # 4. Roll both caches back to "everything committed except the newest token".
        t_cache.truncate(base + accepted + 1)
        if accepted == k:
            draft(proposals[-1], d_cache)          # the draft never fed its own last proposal
        d_cache.truncate(base + accepted + 1)
        for token_id in new_tokens:
            output.append(token_id)
            if token_id in stop_ids or len(output) == max_new_tokens:
                return output, {"rounds": rounds, "acceptance_rate": accepted_total / proposed_total,
                                "tokens_per_target_call": len(output) / rounds}
        last = torch.tensor([[new_tokens[-1]]], device=device)
    return output, {"rounds": rounds, "acceptance_rate": accepted_total / max(proposed_total, 1),
                    "tokens_per_target_call": len(output) / max(rounds, 1)}

The bookkeeping that makes it correct is one invariant: at the top of each round, both caches hold every committed token except the newest one, last, which has been chosen but not yet fed to either model. Then:

  1. The draft feeds last and its own proposals one by one, writing $k$ positions to its cache.
  2. The target feeds last and all $k$ proposals at once, writing $k+1$ positions to its cache and returning $k+1$ distributions.
  3. Some prefix of the proposals is accepted, plus one corrected or bonus token.
  4. Both caches are truncated back to the base length plus the accepted tokens plus last, the new invariant. Rejected proposals’ keys and values are simply forgotten. That’s why Chapter 16’s cache has truncate: rollback costs nothing.

One edge case: if all $k$ proposals are accepted, the draft never fed its own last proposal, so it’s fed once more before truncating.

Run it

python run.py speculate --k 4

Both models are random here (a 6-layer tiny Qwen3 as target), so the numbers show mechanics, not real-world speed:

{"draft": "target itself (perfect draft)", "rounds": 7, "acceptance_rate": 1.0, "tokens_per_target_call": 4.57}
{"draft": "first 2 of 6 layers", "identical_to_target_greedy": true, "rounds": 14, "acceptance_rate": 0.339, "tokens_per_target_call": 2.29}

With a perfect draft, every round yields $k+1 = 5$ tokens (the last round is cut short at 32 tokens). The “early exit” draft, the target’s own first two layers followed by its final norm and head, agrees with the full model a third of the time, still halving the target calls. And the output is identical to the target’s greedy output, as it must be. With --k 2 the early-exit draft’s acceptance rose to 44% but tokens per call fell to 1.78; with --k 8 acceptance fell to 21%.

On real models, a well-matched pair does much better: drafts from the same family and training data agree on most “easy” tokens (punctuation, the rest of a word, boilerplate code). Typical acceptance rates are 0.6-0.8 for chat and higher for code.

Where drafts come from

drafthow it guessestrade-off
small model of the same family (Qwen3-0.6B for Qwen3-8B)its own forward passesneeds the same tokenizer; extra memory and its own KV cache
early exit / self-speculation (LayerSkip, 2024)the target’s first layers plus its headno extra model; works best if trained for early exit
n-gram / prompt lookupcopies the continuation of the last few tokens from earlier in the promptfree; excellent for editing, summarization, code with repetition
Medusa (Cai et al., 2024)extra heads on the target predict tokens $t+2, t+3, \dots$small trained heads; verify a tree of candidates
EAGLE 1-3 (Li et al., 2024-25)a one-layer model predicting the target’s next hidden stateamong the best acceptance rates; trained per target
multi-token prediction (MTP)heads trained with the model (DeepSeek-V3, Qwen3-Next, Qwen3.8-Flash-Next)free at inference if the checkpoint ships them

Tree verification checks several candidate continuations at once: the drafts form a tree, and one target forward pass evaluates all branches using an attention mask where each node sees only its ancestors. Your position-based attention (Chapter 5) supports this directly: give each node its depth as its position and an allowed mask from the tree.

When it doesn’t help

Speculation turns spare memory bandwidth into tokens. At large batch sizes, decode is already compute-bound (Chapter 24), so verifying $k+1$ tokens per request costs $k+1$ times as much, and rejected tokens are pure waste. Production engines reduce $k$ or turn speculation off as the batch grows. It also helps less when acceptance is low: high-temperature creative writing, or a draft trained on different data.

Build it

Engine milestone 26: speculative decoding. Implement accept_or_correct and speculative_generate in engine/speculative.py (expected_tokens_per_round is provided).

pytest tests/test_ch26_speculative.py
python run.py speculate --impl engine --k 4

The tests check that greedy speculation with an imperfect draft reproduces the target’s greedy output exactly, that a perfect draft accepts every proposal, that rejection sampling reproduces $p$ over 20,000 draws, and the expected-tokens formula.

Stretch exercises

  1. ★ Measure acceptance rate and wall-clock speedup for Qwen3-1.7B (or 8B) as target and Qwen3-0.6B as draft on a GPU, for $k = 2, 4, 6$. Compare with the formula. Where: experiments/ch26.py (create it), calling engine.speculative.speculative_generate.
  2. ★★ Implement prompt lookup decoding: find the most recent earlier occurrence of the last 3 tokens in the context and propose the $k$ tokens that followed. Measure acceptance on a “rewrite this paragraph” task. Where: add a prompt-lookup draft helper in engine/speculative.py.
  3. ★★ Make $k$ adaptive: track a running acceptance rate and choose the $k$ that maximizes the expected speedup each round. Where: the round loop in speculative_generate in engine/speculative.py.
  4. ★★★ Implement tree verification for greedy decoding: the draft proposes its top-2 tokens at each of 3 depths (a tree of 14 nodes), verify all branches in one target call with a tree attention mask, and commit the longest accepted path. Where: add tree drafting/verification in engine/speculative.py; pass tree masks through the target model to engine.attention.causal_attention.

Check your understanding

  1. Why does verifying five tokens take about as long as generating one, at batch 1?
  2. Show that the accept-or-resample rule outputs token $y$ with probability exactly $p(y)$.
  3. Why does every round produce at least one token?
  4. Why must both caches be truncated after a round, and to what length?
  5. Why can speculation make a large-batch server slower?

Going deeper

  • Leviathan, Kalman and Matias, Fast Inference from Transformers via Speculative Decoding (2023); Chen et al., Accelerating Large Language Model Decoding with Speculative Sampling (DeepMind, 2023).
  • GPU Mode L22 (Hacker’s guide to speculative decoding in vLLM, Cade Daniels).
  • PMPP §20.6 names batching and speculative decoding as the two main ways to raise the arithmetic intensity of the generation phase.
  • Cai et al., Medusa (2024); Li et al., EAGLE, EAGLE-2, EAGLE-3 (2024-25); Elhoushi et al., LayerSkip (2024); Gloeckle et al., Better & Faster Large Language Models via Multi-token Prediction (2024); DeepSeek-V3 technical report (MTP).