Speculative decoding
Be able to explain the draft model and the verification and measure the speed-up.
Prerequisites
Intuition
Decoding is memory-bandwidth-bound: for every token all the model weights are read. But verifying several tokens in one sweep costs almost as little as doing one — the same weights, more parallel work.
Speculative decoding exploits that:
- A small, fast draft model generates k tokens.
- The large model runs one forward pass over all k and checks them.
- The longest correct prefix is accepted; at the first deviation a token is sampled from the large model and the rest are discarded.
The decisive point: the acceptance criterion is constructed so that the output distribution becomes exactly the same as without speculation. It is not an approximation — you get the same quality, faster.
Formal
The acceptance rule (Leviathan et al. 2022): for a draft token with probability under the draft model and under the target model, accept with probability . On rejection, sample from the normalised residual distribution .
That guarantees that the resulting distribution is exactly — proved mathematically, not empirically.
The expected speed-up with an acceptance rate and a draft length :
With and : ≈ 2.8 tokens per round. Subtract the draft model's cost (typically 10–20 % of the target model's) → roughly 2–2.5× faster.
Where it works best: predictable text (code, structured output, repetitions) where the acceptance rate is high. Worst on creative text with high entropy.
Variants without a separate draft model: Medusa (extra heads on the same model guessing several tokens ahead), n-gram lookup (propose continuations that have already occurred in the context — nearly free and surprisingly effective with RAG and code editing).
Code
import numpy as np
def expected_gain(alpha, k, draft_cost=0.15):
"""Tokens per target-model forward pass, adjusted for the draft model's cost."""
accepted = (1 - alpha ** (k + 1)) / (1 - alpha)
cost = 1 + k * draft_cost
return accepted / cost
for alpha in (0.5, 0.7, 0.9):
row = [round(expected_gain(alpha, k), 2) for k in (2, 4, 8)]
print(f"alpha={alpha} k=2,4,8 → {row}")
# alpha=0.5 k=2,4,8 → [1.36, 1.22, 0.9]
# alpha=0.7 k=2,4,8 → [1.72, 1.79, 1.4]
# alpha=0.9 k=2,4,8 → [2.14, 2.6, 2.66]
def accept(p, q, x, rng):
"""p, q: probability vectors from the target and the draft model. x: the draft's token."""
if rng.random() < min(1.0, p[x] / max(q[x], 1e-12)):
return x, True
residual = np.maximum(p - q, 0)
return int(rng.choice(len(p), p=residual / residual.sum())), False
The table shows that a larger k is not always better: at a low acceptance rate too much work is thrown away and the gain goes negative.
Mastery means
- Explains the draft model and the verification
- Computes the speed-up from the acceptance rate
- Knows that the output distribution is unchanged
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Fast Inference from Transformers via Speculative Decoding — arXiv (open access; licence per article)
- arXiv — Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads — arXiv (open access; licence per article)