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.

Let pheta(xt∣x<t)p_ heta(x_t \mid x_{<t}) be the categorical distribution produced by the model. Greedy decoding emits ildext=argmaxipiilde{x}_t = \\argmax_i p_i 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 BB prefixes. For a prefix y1:ty_{1:t}, its score is the log-probability

s(y1:t)=∑i=1tlog⁡p(yi∣y<i).(39.1)s(y_{1:t}) = \sum_{i=1}^{t} \log p(y_i \mid y_{<i}) .\tag{39.1}

Extending every beam by every token and keeping the top BB candidates is exact for B=VtB = V^t 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.

Listing 39.1 Tiny beam search over a next-logits function
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 z\vz. Temperature samples from softmax⁡(z/T)\softmax(\vz / T). If T<1T < 1 the largest probabilities sharpen; if T>1T > 1 they flatten. Greedy is the limit T→0T \to 0.

Top-k, nucleus top-p, and min-p are masks followed by renormalization. Let SS be the kept token set. The sampled distribution is

p~i=pi1{i∈S}∑jpj1{j∈S}.(39.2)\tilde{p}_i = \frac{p_i \mathbf{1}\{i \in S\}} {\sum_j p_j \mathbf{1}\{j \in S\}} .\tag{39.2}

Top-k keeps the kk largest probabilities. Top-p sorts tokens and keeps the shortest prefix whose cumulative probability reaches a threshold. Min-p keeps tokens with pi≥αmax⁡jpjp_i \ge \alpha \max_j p_j, so the support grows and shrinks with confidence. All three are useful controls, but they are not neutral: after filtering, low-probability tokens outside SS have probability exactly zero.

The order of operations matters. Temperature acts on logits before softmax, changing the relative odds between every pair by exp⁡((zi−zj)/T)\exp((z_i-z_j)/T). 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.

Listing 39.2 Temperature, top-k, top-p, min-p, and logit masks
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 −∞-\infty before softmax. The resulting distribution is the original model distribution conditioned on the allowed set, because softmax⁡(z)i/∑j∈Ssoftmax⁡(z)j\softmax(\vz)_i / \sum_{j \in S} \softmax(\vz)_j is exactly (39.2).

Listing 39.3 Digit-only and tiny-grammar constraints
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 qq to propose a token and an expensive target distribution pp to verify it [leviathan2022fast] [chen2023accelerating]. Draw x∼qx \sim q. Accept it with probability

a(x)=min⁡(1,p(x)q(x)).(39.3)a(x) = \min\left(1, \frac{p(x)}{q(x)}\right) .\tag{39.3}

If rejected, sample instead from the residual

r(x)=max⁡(0,p(x)−q(x))∑ymax⁡(0,p(y)−q(y)).(39.4)r(x) = \frac{\max(0, p(x) - q(x))}{\sum_y \max(0, p(y) - q(y))} .\tag{39.4}

This samples exactly from pp. The probability of outputting token xx through the accept path is q(x)a(x)=min⁡(q(x),p(x))q(x)a(x) = \min(q(x), p(x)). The rejection probability is 1−∑ymin⁡(q(y),p(y))1 - \sum_y \min(q(y), p(y)), which equals ∑ymax⁡(0,p(y)−q(y))\sum_y \max(0, p(y)-q(y)). Therefore the reject path contributes max⁡(0,p(x)−q(x))\max(0, p(x)-q(x)), and the total is p(x)p(x).

Listing 39.4 One-token speculative sampling
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 KK tokens, each position has acceptance probability αi=∑xmin⁡(pi(x),qi(x))\alpha_i = \sum_x \min(p_i(x), q_i(x)). If verification stops at the first rejection, the expected number of accepted draft tokens is ∑i=1K∏j=1iαj\sum_{i=1}^K \prod_{j=1}^i \alpha_j. 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 qq 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.

Listing 39.5 Expected accepted draft tokens
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].

Key equations
x~t=arg max⁡ipi\tilde{x}_t = \argmax_i p_i
s(y1:t)=∑ilog⁡p(yi∣y<i)s(y_{1:t}) = \sum_i \log p(y_i \mid y_{<i})
p~i=pi1{i∈S}∑jpj1{j∈S}\tilde{p}_i = \frac{p_i\mathbf{1}\{i\in S\}}{\sum_j p_j\mathbf{1}\{j\in S\}}
a(x)=min⁡(1,p(x)q(x))a(x) = \min\left(1, \frac{p(x)}{q(x)}\right)
r(x)∝max⁡(0,p(x)−q(x))r(x) \propto \max(0, p(x) - q(x))

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 →\to softmax pp; 2. show masks and renormalization; 3. score two-token beams with log probabilities; 4. prove min⁡(p,q)(p−q)=p\min(p,q)(p-q)_=p. 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 p=qp=q, what is the acceptance rate and residual mass?

39.6 Exercises

Exercise 39.1 ★ Filters

For probabilities (0.50,0.25,0.15,0.10)(0.50, 0.25, 0.15, 0.10), compute the support kept by top-k with k=2k=2, top-p with p=0.70p=0.70, and min-p with α=0.30\alpha=0.30. Which filtered distribution is most peaked?

Exercise 39.2 ★★ Speculative proof

Fill in the proof that ∑xmax⁡(0,p(x)−q(x))=1−∑xmin⁡(p(x),q(x))\sum_x \max(0,p(x)-q(x)) = 1 - \sum_x \min(p(x),q(x)) for normalized pp and qq, then use it to show the reject path contributes (p(x)−q(x))+(p(x)-q(x))_+.

Exercise 39.3 ★★ Expected accepted tokens

A three-token draft has acceptance probabilities 0.8,0.7,0.50.8, 0.7, 0.5. Compute the expected number of accepted draft tokens before the first rejection, and check it by simulation.

Exercise 39.4 ★★★ Constrained implementation

Implement a mask for the grammar [ d (, d)* ], where dd 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