Skip to main content

Speculative Decoding: Draft, Then Verify

Summary
A small draft model proposes tokens, the target model checks them all in one pass, and a rejection rule keeps the output distribution exactly the target's.

Autoregressive decoding produces one token per forward pass, strictly in sequence. Speculative decoding shortens that chain without changing what the model generates. A small draft model guesses the next few tokens, and the large target model checks all of them in one forward pass. Guesses that pass the check are kept; at the first rejection, the target supplies its own token and the round ends.

Leviathan, Kalman and Matias (2023) and Chen et al. (2023) published it independently, as speculative decoding and speculative sampling. A modified rejection-sampling rule makes the output follow the target model's distribution exactly.

The draft-then-verify loop

One round with draft length γ:

  1. Draft. The draft model runs γ ordinary decoding steps, proposing x1 … xγ and keeping its distribution q at each position.
  2. Verify. The target runs once over the prefix plus all γ drafts. A causal transformer predicts the next token at every input position, so this pass yields the target distribution p at each drafted position and one position beyond.
  3. Accept or reject, left to right. The first rejection ends the scan. Later drafts are discarded, because they were conditioned on a token that is no longer in the sequence.
  4. Emit one target token. After a rejection, the target samples a correction; if every draft survived, it samples a bonus token.

Every round commits at least one token, like a plain decoding step, and at most γ + 1.

Step through one run

Choose α, γ and the draft cost c, then step through the rounds; the cursor marks the same instant in every lane. With α = 0.4 and c = 0.5, speculation finishes last.

Why a verify pass costs about one decode step

At batch size one, a decode step multiplies one activation vector by every weight matrix. Each weight is read from high-bandwidth memory (HBM) and used for about two floating-point operations, so the step is limited by memory bandwidth while the compute units wait (data movement in transformers analyses this).

Scoring γ + 1 positions turns those matrix-vector products into thin matrix-matrix products. The weights are still read once, the extra arithmetic fills idle compute, and every position attends over the same cached keys and values. While γ + 1 stays small, the pass takes about as long as a single-token step, yet it can commit several tokens.

IO-aware kernels such as FlashAttention and the shared key/value heads of MQA and GQA (compared here) cut the bytes each step moves; speculation stacks on top of them.

The acceptance rule

Let q be the draft distribution, p the target distribution and x the drafted token. The target keeps x with probability

min(1,\ p(x)q(x))

A token the draft over-rated is kept only in proportion; any other is always kept. On rejection, the replacement comes from the normalised residual

r(x) \propto max\big(0,\ p(x) - q(x)\big)

Why the output still follows p

Token x arrives by two routes. The draft proposes it and it is kept, with probability min(p(x), q(x)). Or a draft is rejected and the residual picks x. Since p = min(p, q) + max(0, p − q) for every token, the rejection probability, 1 − Σ min(p, q), equals the residual's normaliser Σ max(0, p − q), so the second route contributes exactly max(0, p(x) − q(x)):

\begin{aligned} P\text{out}(x) &= min(p(x), q(x)) \ &+ max(0, p(x) - q(x)) \ &= p(x) \end{aligned}

A four-token example with acceptance 0.7, where the output is min(p, q) plus 0.3 times r:

Tokencatdogfoxowl
p (target)0.40.30.20.1
q (draft)0.20.50.10.2
min(1, p/q)10.610.5
min(p, q)0.20.30.10.1
r2/301/30
Output0.40.30.20.1

The rejected mass of 0.3 splits 2 : 1 between cat and fox and restores p exactly. The chance of keeping a draft is the sum of min(p, q), one minus the total variation distance between p and q.

Chen et al. state that the guarantee holds within hardware numerics. One source of drift is that the verify pass has a different batch shape from single-token decoding, so floating-point logits can differ slightly.

Greedy decoding

Under greedy decoding p is one-hot, and the rule becomes exact match: keep a draft if it equals the target's argmax, otherwise emit the argmax. The output matches plain greedy decoding, up to the same floating-point caveat, which can flip a near-tie.

Tokens per pass and speedup

If each draft is kept independently with probability α, the expected number of tokens one round commits is

𝔼[\text{tokens per pass}] = 1 - αγ+11 - α

It rises from 1 at α = 0 toward γ + 1 as α approaches 1. Leviathan et al. derive it under this independence simplification; real acceptance varies with the text. A round costs γc + 1, so the expected speedup is

\text{speedup} = 1 - αγ+1(1 - α)(γ c + 1)

For α = 0.8, γ = 4 and c = 0.1: 3.36 tokens per pass and about 2.4×. Growing the draft from γ to γ + 1 adds only α^(γ+1) expected tokens but another c of draft time, so long drafts stop paying off.

Variants

Later methods change where drafts come from, and several verify a tree of candidate continuations in one pass with a tree-shaped attention mask.

  • Medusa (Cai et al., 2024) adds decoding heads to the target's last hidden state, each predicting a token several positions ahead; their top candidates form the tree. Medusa-1 trains only the heads, Medusa-2 also fine-tunes the backbone. Its optional typical-acceptance rule keeps more candidates but no longer guarantees the target distribution.
  • EAGLE (Li et al., 2024) drafts features: a small head predicts the target's next second-to-top-layer feature from earlier features and the tokens advanced by one step, and the target's own output head turns features into tokens. It preserves the target distribution.
  • Lookahead decoding (Fu et al., 2024) needs no draft model. Jacobi iteration refines guesses for several future positions in parallel, the resulting n-grams go into a pool, and promising n-grams are verified in the same forward pass.
  • Self-speculative decoding drafts with the target itself: Zhang et al. (2023) skip intermediate layers while drafting, and LayerSkip (Elhoushi et al., 2024) trains the model to exit early and verifies with the remaining layers. No second model sits in memory.

When it does not help

  • Low acceptance. A round commits about one token but still pays γc for drafts, so decoding can get slower. Measure acceptance on the real workload.
  • A small target. A small target leaves little room for a much cheaper draft, and fixed per-round costs such as sampling and scheduling weigh more.
  • Compute-bound serving. In large batches the target is already limited by arithmetic, so verifying γ + 1 positions per sequence adds real work, and rejected drafts waste compute other requests could use. Throughput can fall even when latency improves.
  • Prefill. Only token-by-token decoding speeds up; a long prompt is already processed in parallel.

Production gotchas

Rolling back the KV cache

The verify pass writes keys and values for every drafted position, and the draft model caches its own proposals. After a rejection, both caches must be truncated to the accepted prefix; the correction or bonus token gets its entry as the first input of the next round. With a paged cache such as PagedAttention, rollback frees or overwrites the rejected blocks.

A shared vocabulary

The ratio p(x)/q(x) compares probabilities for the same token id, so draft and target need the same tokenizer and vocabulary, usually from one model family. Research methods handle mismatched vocabularies, but the standard algorithm assumes a shared one.

The same sampling settings

Temperature, top-k and top-p change the sampled distribution. The acceptance test and residual must use the processed p and q the samplers actually use, or the output drifts from plain sampling.

Batching

Sequences in a batch keep different numbers of drafts, so lengths turn ragged after each verify pass. Because verification work grows with batch size, some serving systems tune γ to the load or switch speculation off above a batch-size threshold.

Primary sources

If you found this explanation helpful, consider sharing it with others.

Mastodon