Chapter 39
Decoding & Speculative Sampling
Greedy, beam, top-k, top-p, min-p, constrained decoding, and speculative sampling.
This chapter opens Part VII: turning a trained language model into a stream of tokens. At each step the model gives a next-token distribution; decoding either changes that distribution, searches over short continuations, or samples faster without changing it. The goal is to know exactly when we are changing the model and when we are only implementing the same distribution more efficiently.
39.1 Greedy and beam search
Let be the categorical distribution produced by the model. Greedy decoding emits at every step. It is cheap and deterministic, but it maximizes each local decision, not the whole sequence probability. A token that is second-best now can make a much better continuation later.
Beam search keeps prefixes. For a prefix , its score is the log-probability
Extending every beam by every token and keeping the top candidates is exact for and approximate otherwise. Logs turn products into sums and avoid underflow. The tiny implementation below deliberately omits length penalties and end tokens so the core search is visible.
Beam search is best understood as deterministic optimization, not as a sampler. With a fixed beam, it returns one of the high-scoring sequences under the model, and increasing the beam can change earlier choices because more prefixes survive long enough to reveal their continuations. Practical text decoders often add an end token, a length normalization, or diversity penalties, but those are extra objectives. The unmodified score in (39.1) is simply the model’s log-probability for the emitted sequence.
def tiny_beam_search(next_logits, prompt, steps, beam_size=2):
"""Keep the highest log-probability prefixes under a next-logits function."""
beams = [(tuple(prompt), 0.0)]
for _ in range(steps):
candidates = []
for prefix, score in beams:
probs = softmax(next_logits(prefix))
for token, prob in enumerate(probs):
candidates.append((prefix + (token,), score + np.log(prob)))
candidates.sort(key=lambda item: item[1], reverse=True)
beams = candidates[:beam_size]
return beams
39.2 Sampling filters are distribution transforms
Sampling starts from logits . Temperature samples from . If the largest probabilities sharpen; if they flatten. Greedy is the limit .
Top-k, nucleus top-p, and min-p are masks followed by renormalization. Let be the kept token set. The sampled distribution is
Top-k keeps the largest probabilities. Top-p sorts tokens and keeps the shortest prefix whose cumulative probability reaches a threshold. Min-p keeps tokens with , so the support grows and shrinks with confidence. All three are useful controls, but they are not neutral: after filtering, low-probability tokens outside have probability exactly zero.
The order of operations matters. Temperature acts on logits before softmax, changing the relative odds between every pair by . The filters then choose a support using the probabilities after temperature. Finally renormalization divides by the remaining mass, so a token’s probability can increase even though no logit changed. This is why the same top-p value can be conservative at low temperature and adventurous at high temperature.
def softmax(logits):
"""Stable softmax over the last axis."""
logits = np.asarray(logits, dtype=np.float64)
shifted = logits - np.max(logits, axis=-1, keepdims=True)
exp = np.exp(shifted)
return exp / exp.sum(axis=-1, keepdims=True)
def mask_logits(logits, allowed):
"""Set disallowed tokens to -inf before softmax or argmax."""
logits = np.asarray(logits, dtype=np.float64)
allowed = np.asarray(allowed, dtype=bool)
return np.where(allowed, logits, -np.inf)
def filtered_distribution(logits, temperature=1.0, top_k=None, top_p=None,
min_p=None, allowed=None):
"""Temperature, top-k, top-p, and min-p as probability transforms."""
z = np.asarray(logits, dtype=np.float64) / temperature
if allowed is not None:
z = mask_logits(z, allowed)
p = softmax(z)
keep = np.ones_like(p, dtype=bool)
if top_k is not None:
cutoff = np.partition(p, -top_k)[-top_k]
keep &= p >= cutoff
if top_p is not None:
order = np.argsort(-p)
cumulative = np.cumsum(p[order])
chosen = np.zeros_like(keep)
stop = int(np.searchsorted(cumulative, top_p, side="left"))
chosen[order[:stop + 1]] = True
keep &= chosen
if min_p is not None:
keep &= p >= min_p * np.max(p)
return np.where(keep, p, 0.0) / np.sum(np.where(keep, p, 0.0))
39.3 Constrained decoding is masked softmax
A constraint is also a mask. If a field must contain only digits, or a small grammar says that a comma must follow a number, set every illegal logit to before softmax. The resulting distribution is the original model distribution conditioned on the allowed set, because is exactly (39.2).
def tiny_json_number_mask(prefix, vocab):
"""Allowed tokens for a tiny grammar: '[' digit (',' digit)* ']'."""
if not prefix:
return np.array([token == "[" for token in vocab])
if prefix[-1] in {"[", ","}:
return np.array([token.isdigit() for token in vocab])
if prefix[-1].isdigit():
return np.array([token in {"]", ","} for token in vocab])
return np.zeros(len(vocab), dtype=bool)
def grammar_step(logits, prefix, vocab):
return softmax(mask_logits(logits, tiny_json_number_mask(prefix, vocab)))
This is the idea behind JSON-schema and grammar-constrained generation. The hard part is maintaining the allowed set quickly as the prefix changes; the math is only masking.
Masking also explains why constrained decoding can still produce fluent text. The model supplies preferences among allowed tokens, while the constraint supplies a hard support. If the allowed set is empty, the decoder has reached an invalid state; implementations either backtrack, force a repair token, or reject that partial sequence. For tiny grammars the mask is a few if statements. For real schemas it is usually a finite-state machine or parser state updated after each emitted token.
39.4 Speculative sampling
Speculative sampling uses a cheap draft distribution to propose a token and an expensive target distribution to verify it [leviathan2022fast] [chen2023accelerating]. Draw . Accept it with probability
If rejected, sample instead from the residual
This samples exactly from . The probability of outputting token through the accept path is . The rejection probability is , which equals . Therefore the reject path contributes , and the total is .
def speculative_sample_step(p, q, rng):
"""One exact speculative-sampling step for target p and draft q."""
p = np.asarray(p, dtype=np.float64)
q = np.asarray(q, dtype=np.float64)
draft = int(sample_categorical(q, rng))
if rng.random() < min(1.0, p[draft] / max(q[draft], 1e-300)):
return draft, True
residual = np.maximum(0.0, p - q)
return int(sample_categorical(residual / residual.sum(), rng)), False
def speculative_sample(p, q, n, rng):
"""Draw n tokens from p by proposing with q and correcting rejections."""
draws = np.empty(n, dtype=np.int64)
accepted = 0
for i in range(n):
draws[i], ok = speculative_sample_step(p, q, rng)
accepted += int(ok)
return draws, accepted / n
For a draft block of tokens, each position has acceptance probability . If verification stops at the first rejection, the expected number of accepted draft tokens is . Multi-token prediction heads and Medusa-style heads train the main model to draft several future tokens, replacing a separate small model with cheap heads attached to the same network [cai2024medusa] [gloeckle2024better].
The speedup comes from batching target-model work. The drafter proposes several tokens cheaply, then the verifier evaluates those positions in one forward pass using the proposed prefix. If many proposals are accepted, one expensive verification step advances several tokens. If is poor, most proposals are rejected and the method falls back toward ordinary sampling with extra overhead. The exactness proof above is local to one position; applying it sequentially preserves the target autoregressive distribution because every accepted or corrected token has the same conditional distribution the target model would have sampled at that prefix.
def acceptance_probability(p, q):
"""Probability that one proposed token is accepted: sum min(p_i, q_i)."""
return float(np.minimum(p, q).sum())
def expected_accepted_prefix(acceptance_probs):
"""Expected accepted draft tokens before the first rejection."""
expected = 0.0
prefix = 1.0
for alpha in acceptance_probs:
prefix *= alpha
expected += prefix
return expected
|
In practice
|
Production LLM APIs usually expose temperature, top-p, top-k, penalties, and stop constraints because they are simple transformations at the logits boundary. Speculative decoding is used when the verifier is much more expensive than the drafter and the acceptance rate is high; otherwise the extra draft work does not pay for itself [leviathan2022fast] [chen2023accelerating]. Multi-token heads are attractive because they reuse the target model’s hidden state and avoid serving a second model, but they still need the exact accept/reject correction to preserve the target distribution [cai2024medusa] [gloeckle2024better]. |
39.5 Teach it
One sentence: decoding is the policy that turns next-token probabilities into tokens, while speculative sampling accelerates exact sampling by correcting a cheap proposal. Analogy: beam search is keeping several promising chess lines; top-p is ignoring moves outside the plausible cluster; speculation is letting a junior player suggest moves that the expert accepts or fixes. Board steps: 1. write logits softmax ; 2. show masks and renormalization; 3. score two-token beams with log probabilities; 4. prove . Misconceptions: temperature and top-p do change the distribution; beam search is not sampling; speculative decoding is exact only with the residual correction. Check: if , what is the acceptance rate and residual mass?
39.6 Exercises
For probabilities , compute the support kept by top-k with , top-p with , and min-p with . Which filtered distribution is most peaked?
Fill in the proof that for normalized and , then use it to show the reject path contributes .
A three-token draft has acceptance probabilities . Compute the expected number of accepted draft tokens before the first rejection, and check it by simulation.
Implement a mask for the grammar [ d (, d)* ], where is one digit. Test the allowed next tokens after [], [3, and [3,.
References
-
[cai2024medusa] T. Cai et al. Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads. 2024. arXiv:2401.10774
-
[chen2023accelerating] C. Chen et al. Accelerating Large Language Model Decoding with Speculative Sampling. 2023. arXiv:2302.01318
-
[gloeckle2024better] F. Gloeckle et al. Better & Faster Large Language Models via Multi-token Prediction. 2024. arXiv:2404.19737
-
[leviathan2022fast] Y. Leviathan, M. Kalman, and Y. Matias. Fast Inference from Transformers via Speculative Decoding. 2022. arXiv:2211.17192