= Decoding & 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. [#sec-search] == Greedy and beam search Let stem:[p_ heta(x_t \mid x_{> is simply the model's log-probability for the emitted sequence. .Tiny beam search over a next-logits function [source,python] ---- include::../../scratch/decoding.py[tag=beam] ---- [#sec-distribution-transforms] == Sampling filters are distribution transforms Sampling starts from logits stem:[\vz]. Temperature samples from stem:[\softmax(\vz / T)]. If stem:[T < 1] the largest probabilities sharpen; if stem:[T > 1] they flatten. Greedy is the limit stem:[T \to 0]. Top-k, nucleus top-p, and min-p are masks followed by renormalization. Let stem:[S] be the kept token set. The sampled distribution is [latexmath#eq-filter] ++++ \tilde{p}_i = \frac{p_i \mathbf{1}\{i \in S\}} {\sum_j p_j \mathbf{1}\{j \in S\}} . ++++ Top-k keeps the stem:[k] largest probabilities. Top-p sorts tokens and keeps the shortest prefix whose cumulative probability reaches a threshold. Min-p keeps tokens with stem:[p_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 stem:[S] have probability exactly zero. The order of operations matters. Temperature acts on logits before softmax, changing the relative odds between every pair by stem:[\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. .Temperature, top-k, top-p, min-p, and logit masks [source,python] ---- include::../../scratch/decoding.py[tag=logits] ---- [#sec-constrained] == 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 stem:[-\infty] before softmax. The resulting distribution is the original model distribution conditioned on the allowed set, because stem:[\softmax(\vz)_i / \sum_{j \in S} \softmax(\vz)_j] is exactly <>. .Digit-only and tiny-grammar constraints [source,python] ---- include::code/solutions.py[tag=tiny-grammar] ---- 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. [#sec-speculative] == Speculative sampling Speculative sampling uses a cheap draft distribution stem:[q] to propose a token and an expensive target distribution stem:[p] to verify it <> <>. Draw stem:[x \sim q]. Accept it with probability [latexmath#eq-accept] ++++ a(x) = \min\left(1, \frac{p(x)}{q(x)}\right) . ++++ If rejected, sample instead from the residual [latexmath#eq-residual] ++++ r(x) = \frac{\max(0, p(x) - q(x))}{\sum_y \max(0, p(y) - q(y))} . ++++ This samples exactly from stem:[p]. The probability of outputting token stem:[x] through the accept path is stem:[q(x)a(x) = \min(q(x), p(x))]. The rejection probability is stem:[1 - \sum_y \min(q(y), p(y))], which equals stem:[\sum_y \max(0, p(y)-q(y))]. Therefore the reject path contributes stem:[\max(0, p(x)-q(x))], and the total is stem:[p(x)]. .One-token speculative sampling [source,python] ---- include::../../scratch/decoding.py[tag=speculative] ---- For a draft block of stem:[K] tokens, each position has acceptance probability stem:[\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 stem:[\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 <> <>. 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 stem:[q] 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. .Expected accepted draft tokens [source,python] ---- include::../../scratch/decoding.py[tag=acceptance] ---- [NOTE,caption=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 <> <>. 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 <> <>. ==== [.key-equations#key-equations] .Key equations **** [latexmath] ++++ \tilde{x}_t = \argmax_i p_i ++++ [latexmath] ++++ s(y_{1:t}) = \sum_i \log p(y_i \mid y_{