Skip to content

Speculative decoding: n-gram, prompt-lookup, and model drafts

ModuleL8.6 · build · Python · Pass 6 · 3 to 4 h
You buildpython/tinyllm/infer/spec.py: verify_draft, speculative_generate, and three drafts: NGramDraft, PromptLookupDraft, ModelDraft (the DraftModel protocol is in the contract)
Contractcourse/contracts/py/tinyllm/infer/spec.pyi
Testscourse/tests/L8.6/test_spec.py (and the toy models in _toy.py) (what they check: section 4)
NeedsM07.6 rejection sampling (speculative_step), L8.2 KV cache and generate (KVCache.truncate, cache_dims, Generation), L2.1 n-gram LM (NGramLM), L8.1 sampling (sample, sampling_distribution, token_logprobs, request_rng) (or --ref-deps). Reading: M11.1 entropy and KL
Used byL10.8 speculative decoding in the Rust engine, held to this module on shared draft and target logits (joins used_by when registered)
MilestoneMS-L8 (step spec-ngram-greedy: token-identical output, acceptance_rate reported)
Optional depthLeviathan, Kalman, and Matias, Fast Inference from Transformers via Speculative Decoding (ICML 2023); Chen et al., Accelerating Large Language Model Decoding with Speculative Sampling (2023); Saxena, Prompt Lookup Decoding (2023)
  • Decoding one token costs one pass over the weights; scoring k+1k + 1 tokens costs about the same, so a draft that guesses kk tokens and a target that checks them in one pass can emit several tokens per pass (test_stats_and_one_pass_per_round).
  • Greedy verification never changes the output: keep drafts while they equal the target’s argmax, then add the target’s own token (test_speculative_greedy_equals_greedy).
  • Sampled verification keeps the target’s distribution exactly: accept x∼qx \sim q with probability min⁡(1,p(x)/q(x))\min(1, p(x)/q(x)), otherwise draw from max⁡(0,p−q)\max(0, p - q) normalized; enumerated exactly, the output is pp as fractions (test_one_token_output_is_the_target_exactly, test_two_token_draft_is_exact_at_each_position).
  • Rejected drafts must leave no trace in the KV cache: truncate back to the kept length, then feed the replacement token first next round (test_rejected_drafts_are_rolled_back).
  • Cheap drafts are enough for repetitive text: copying from the prompt (test_prompt_lookup_hand_example) or an n-gram model (test_ngram_draft_follows_the_model) cost nothing to run.
Terminal window
ol start L8.6 # stubs spec.py into your repo, contract alongside
ol tests L8.6 # read the test catalog first
ol check L8.6 # exit code is the verdict
ol check L8.6 --ref-deps # only if you skipped M07.6, L8.2, L2.1, or L8.1
ol diff L8.6 # after passing: your code against the reference

You also write graded tests (rung R5) in python/tests/l8-6-spec/: they run against the reference with one planted bug at a time, and the share they catch is your grade.


Your generate (L8.2) produces one token per forward pass. Profile it on SmolLM2 and the pass is dominated by reading every weight from memory once; the arithmetic for one token is small. A pass over 5 tokens (with the KV cache, 5 new query rows) reads the same weights once and costs barely more. So decoding wastes most of what each pass could do. Speculative decoding spends that slack: something cheap guesses the next kk tokens, the target scores all of them in one pass, and every correct guess is a token you did not pay a pass for. The trap is correctness: a guess that is accepted when the target would not have produced it changes the output, and for sampling it changes the distribution. This module builds the verification rules that make the speedup free of either, plus three drafts, and Part 10’s Rust engine (L10.8) is checked against it.

SymbolMeaningType / shape
VVvocabulary sizeint
kkthe number of draft tokens per roundint ≥0\ge 0
x1,…,xmx_1, \dots, x_mthe draft tokens of one round, m≤km \le kids
pip_ithe target’s sampling distribution (L8.1) at draft position ii, given everything before xix_ifloat64 [V][V]
qiq_ithe draft’s distribution that xix_i was drawn from (one-hot for a deterministic draft)float64 [V][V]
uua uniform draw in [0,1)[0, 1) from the request’s PCG32 streamfloat
α=∑vmin⁡(p(v),q(v))\alpha = \sum_v \min(p(v), q(v))the acceptance probability of one draft token[0,1][0, 1]
TV(p,q)=12∑v∣p(v)−q(v)∣\mathrm{TV}(p, q) = \tfrac12 \sum_v \lvert p(v) - q(v) \rverttotal variation distance; α=1−TV\alpha = 1 - \mathrm{TV}[0,1][0, 1]

The target cache holds the positions that are kept. One round:

  1. the draft proposes x1..xmx_1..x_m after the context (prompt plus generated);
  2. the target runs one forward pass over the tokens not yet in its cache: the pending token (the previous round’s last token, or the whole prompt in the first round) followed by x1..xmx_1..x_m. Its last m+1m + 1 rows of logits are the target’s predictions after the context, after x1x_1, …, after xmx_m;
  3. verify_draft keeps x1..xnx_1..x_n for some n≤mn \le m and adds one more token yy of the target’s own (a correction, or a bonus when all mm were kept);
  4. the cache now holds mm draft positions, but only nn are kept: KVCache.truncate(committed + len(pending) + n) drops the rest. yy was never fed, so it becomes the next round’s pending token.

Every round emits n+1≥1n + 1 \ge 1 tokens for one target pass, so speculation can never be slower in passes than plain decoding (with k=0k = 0 it is plain decoding).

2.2 Exactness: why rejection sampling gives exactly pp

Section titled “2.2 Exactness: why rejection sampling gives exactly ppp”

Draw x∼qx \sim q. Keep it with probability min⁡(1,p(x)/q(x))\min(1, p(x)/q(x)) (M07.6’s rejection_accept: u<p(x)/q(x)u < p(x)/q(x)). If it is rejected, draw yy from the residual r=max⁡(0,p−q)/Zr = \max(0, p - q) / Z with Z=∑vmax⁡(0,p(v)−q(v))Z = \sum_v \max(0, p(v) - q(v)). For any token vv:

P(out=v)=q(v)min⁡ ⁣(1,p(v)q(v))⏟kept+(1−α)r(v)⏟replaced=min⁡(q(v),p(v))+max⁡(0,p(v)−q(v))=p(v),P(\text{out} = v) = \underbrace{q(v) \min\!\left(1, \tfrac{p(v)}{q(v)}\right)}_{\text{kept}} + \underbrace{\left(1 - \alpha\right) r(v)}_{\text{replaced}} = \min(q(v), p(v)) + \max(0, p(v) - q(v)) = p(v),

because 1−α=1−∑vmin⁡(p,q)=∑vmax⁡(0,p−q)=Z1 - \alpha = 1 - \sum_v \min(p, q) = \sum_v \max(0, p - q) = Z. Nothing in this argument needs qq to be good: a bad draft is rejected more often, never wrong. A deterministic draft (prompt lookup, greedy n-gram) proposes a fixed xx, which is a draw from the one-hot q=exq = e_x: then xx is kept with probability p(x)p(x) and the residual is pp with xx removed, renormalized.

At position i>1i > 1 the same step runs with pip_i, the target’s distribution after x1..xi−1x_1..x_{i-1} (which the target already computed in the same pass), and only if xi−1x_{i-1} was kept. Conditioned on the kept prefix, each position is a fresh exact step, so every emitted token follows the target’s distribution given the tokens before it. Penalties (repetition, presence, frequency) are part of pip_i: their history includes the drafts kept so far in this round.

At temperature 0 the target’s distribution is one-hot on its argmax gig_i (L8.1, ties to the lowest id). The rule becomes: keep xix_i while xi=gix_i = g_i; at the first mismatch emit gig_i and stop; after mm matches emit gm+1g_{m+1}. No uniform is drawn. The emitted tokens are exactly the target’s greedy tokens, which is the claim the milestone checks on a real model.

If each draft token is accepted independently with probability α\alpha, a round of kk drafts emits 1+α+α2+⋯+αk=1−αk+11−α1 + \alpha + \alpha^2 + \dots + \alpha^k = \frac{1 - \alpha^{k+1}}{1 - \alpha} tokens on average. With α=0.8\alpha = 0.8 and k=4k = 4 that is 3.36 tokens per target pass. The cost side is the draft: an n-gram lookup or a prompt lookup is free, a smaller model costs kk of its own passes per round. acceptance_rate (accepted drafts over proposed drafts) is the number that tells you whether a draft is worth it. Prompt lookup shines when the output repeats the input (code edits, summaries quoting their source); an n-gram model trained on similar text catches common continuations; a model draft generalizes best and costs the most.

One request generator (request_rng(seed), L8.1) serves the draft and the verifier in program order. Per verified draft position the verifier always draws two uniforms, u_accept then u_resample (M07.6’s speculative_step takes both, and drawing both keeps the count independent of the outcome); the bonus token is L8.1’s sample (one uniform). The Rust port (L10.8) draws in the same order, so on the same logits it emits the same ids.

Sampled. V=3V = 3, target p=[12,14,14]p = [\tfrac12, \tfrac14, \tfrac14], draft q=[14,12,14]q = [\tfrac14, \tfrac12, \tfrac14], draft token x=1x = 1.

StepComputationResult
acceptance probabilitymin⁡(1,p1/q1)=min⁡(1,1/41/2)\min(1, p_1/q_1) = \min(1, \tfrac{1/4}{1/2})12\tfrac12
uaccept=0.7u_{accept} = 0.70.7<0.50.7 < 0.5? norejected
residualmax⁡(0,p−q)=[14,0,0]\max(0, p - q) = [\tfrac14, 0, 0], Z=14Z = \tfrac14r=[1,0,0]r = [1, 0, 0]
uresample=0.4u_{resample} = 0.4inverse CDF of rrtoken 0; emitted [0], 0 accepted

With uaccept=0.3<0.5u_{accept} = 0.3 < 0.5 the draft is kept, uresampleu_{resample} is drawn and ignored, and the bonus token comes from the target’s next row with a third uniform. Averaged over the draft as well (x∼qx \sim q): token 1 comes out only when x=1x = 1 is kept, with probability q1⋅12=14=p1q_1 \cdot \tfrac12 = \tfrac14 = p_1; token 0 when x=0x = 0 (always kept, 14\tfrac14) or when x=1x = 1 is rejected (12⋅12\tfrac12 \cdot \tfrac12), total 12=p0\tfrac12 = p_0. That is section 2.2’s identity, which test_one_token_output_is_the_target_exactly checks by enumeration on a larger example.

Greedy. Target rows with argmax 2, 0, 1 ([[0, 1, 3], [2, 1, 0], [0, 5, 1]]) and draft [2, 1]: position 0, draft 2 = argmax 2, keep; position 1, draft 1 ≠\ne argmax 0, emit 0 and stop. Result [2, 0], one accepted, no uniform drawn. Draft [2, 0] matches twice and the bonus is row 2’s argmax, giving [2, 0, 1].

Prompt lookup. Context 1 2 3 9 1 2, k=3k = 3, max_ngram = 3: the last 3-gram 9 1 2 never occurred earlier; the last 2-gram 1 2 occurred at the start, followed by 3 9 1: that is the draft.

class DraftModel(Protocol):
def propose(self, ctx: Sequence[int], k: int, rng: Optional[UniformSource]) -> tuple[list[int], Optional[NDArray]]
class NGramDraft: def __init__(self, lm: NGramLM)
class PromptLookupDraft: def __init__(self, max_ngram: int = 3, min_ngram: int = 1)
class ModelDraft: def __init__(self, model, temperature: float = 1.0)
def verify_draft(target_logits, draft_ids, draft_probs, p, history, rng, prompt=()) -> tuple[list[int], int]
def speculative_generate(target, draft, tok, prompt, p, k=4, kv_dtype=np.float32, eos_ids=()) -> Generation

rng=None asks a draft for its greedy guess (and no distribution). speculative_generate reports drafted, accepted, acceptance_rate, and target_calls in Generation.stats, besides L8.2’s fields. The tests use _toy.BagLM, a model whose logits at a position depend only on the tokens up to it and are computed row by row in float64, so whether the target scores one token or five cannot change a bit; and it asserts on every call that the cache holds exactly the positions before the chunk.

TestKINDChecksWhy it matters downstream
test_hand_exampleunit, smokesection 3, sampled: [0] with 0 accepted for u=0.7,0.4u = 0.7, 0.4; [1, 2] for 0.3,0.9,0.10.3, 0.9, 0.1; three uniformsthe draw order L10.8 must follow
test_hand_example_greedyunit, smokesection 3, greedy: [2, 0], [2, 0, 1], [2]; no drawgreedy verification
test_one_token_output_is_the_target_exactlystatisticalV=5V = 5, all outcomes enumerated: the first token’s distribution equals pp as fractionsthe exactness theorem
test_two_token_draft_is_exact_at_each_positionstatisticalk=2k = 2: first token ∼p1\sim p_1; after a kept x1=y1x_1 = y_1, second ∼p(⋅∣y1)\sim p(\cdot \mid y_1)the right row and history per position
test_deterministic_draft_is_exactstatisticalone-hot drafts (every xx, including one with p(x)=0p(x) = 0) give exactly ppprompt lookup and n-gram drafts at temperature > 0
test_penalties_see_the_accepted_draftsunita repetition penalty rejects a second copy of a just-kept tokenpenalties match plain generate
test_verify_checks_shapesboundaryrows must be m+1m + 1, draft distributions m×Vm \times Vno off-by-one row
test_prompt_lookup_hand_exampleunit, smokesection 3; most recent occurrence; longest n-gram first; k=0k = 0the lookup rule L10.8 ports
test_ngram_draft_follows_the_modelunitgreedy ids are the argmax chain; sampled rows are the model’s distributionthe n-gram draft
test_model_draft_syncs_its_cacheunitthe draft’s own cache follows a changing context; proposals equal fresh greedy decodinga model draft after rejections
test_speculative_greedy_equals_greedydifferentialfour drafts: ids equal L8.2’s generate and a plain argmax loopthe MS-L8 claim
test_rejected_drafts_are_rolled_backpropertyevery target pass starts at the kept length with the right tokenthe cache rollback
test_stats_and_one_pass_per_roundunitself-draft: 40 tokens in 8 passes, acceptance 1; k=0k = 0: 40 passesthe speedup is real
test_sampled_speculation_matches_the_target_distributionstatisticalend to end, 400 seeds, chi-square p>10−3p > 10^{-3}sampling through the whole loop
test_eos_stop_and_budgetboundaryEOS ends without being emitted; a stop string cuts the text; never more than max_tokensgenerate’s contract kept
test_bad_argumentsboundaryk<0k < 0, empty prompt, prompt longer than the modelcaller errors early
test_logprobs_are_the_targetsunitlogprobs equal plain generate’sthe API reports the target’s numbers
PitfallSymptomCaught by
1. treating the draft as if it sampled from the target (q=pq = p), or ignoring draft_probsevery draft kept: the output follows the draft, not the targettest_one_token_output_is_the_target_exactly (mutants s01, s07)
2. a deterministic draft treated as uniform, or its one-hot at the wrong indexbiased output, or a crash on q(x)=0q(x) = 0test_deterministic_draft_is_exact (mutants s02, s09)
3. resampling from pp instead of the residual after a rejectiontokens the draft favors are over-representedtest_one_token_output_is_the_target_exactly (mutant s03)
4. drawing the uniforms in another orderPython and Rust streams divergetest_hand_example (mutant s04)
5. greedy: keeping the draft on a mismatchthe output changes with the drafttest_hand_example_greedy (mutant s05)
6. the bonus from the wrong rowthe token after a fully kept draft is wrongtest_hand_example_greedy (mutant s06)
7. penalties that do not see the drafts kept this roundsampled and greedy output differ from plain generate with penaltiestest_penalties_see_the_accepted_drafts (mutant s08)
8. prompt lookup: oldest occurrence, shortest n-gram first, or ignoring kkdifferent drafts than L10.8; longer drafts than askedtest_prompt_lookup_hand_example (mutants s10, s11, s12)
9. a model draft that does not truncate or re-feed its own cachestale keys: proposals drift from the draft model’s real greedy outputtest_model_draft_syncs_its_cache (mutants s15, s16)
10. no rollback, or the extra token counted as cached, or never fedthe next pass attends to rejected tokens; output divergestest_rejected_drafts_are_rolled_back (mutants s17, s18, s19)
11. EOS kept in ids, or a stop string not cutthe API returns text past the stoptest_eos_stop_and_budget (mutants s22, s23)
12. a round allowed k+1k + 1 tokens when fewer remainmore than max_tokens idstest_eos_stop_and_budget (mutant s24)
DirectionModuleHow it uses this
BackM07.6speculative_step is the sampled verification of one position
BackL8.2KVCache.truncate is the rollback; cache_dims sizes the cache; Generation is the result; generate is the reference the tests compare with
BackL2.1NGramDraft drafts from an NGramLM’s logprobs
BackL8.1sample (greedy and bonus draws), sampling_distribution (the pip_i), token_logprobs, request_rng
BackM11.1total variation and why α=1−TV(p,q)\alpha = 1 - \mathrm{TV}(p, q) (reading)
ForwardL10.8the Rust engine’s prompt-lookup and n-gram speculation, held to verify_draft on shared logits and to greedy equality

If you skip this module, ol check L10.8 stops with BLOCKED ... needs L8.6: build it, or pass --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
verify_draftvLLM’s rejection samplerbatched verification on the GPU, per-request draft lengths, the residual in log spacevLLM v1/sample/rejection_sampler.py
PromptLookupDraftvLLM’s n-gram proposer, HF prompt_lookup_num_tokensa KMP-style search over the whole context, minimum and maximum n-gram sizesvLLM v1/spec_decode/ngram_proposer.py
ModelDraftEAGLE, Medusaa draft head on the target’s own hidden states; a tree of candidates verified with one tree-masked attention passLi et al., EAGLE (2024); Cai et al., Medusa (2024)
one draft pathSGLang and vLLM tree speculationseveral candidate continuations per round, verified togetherSGLang srt/speculative/