8. Generating text: decoding and sampling
In this chapter
- The decoding policy, a separate decision from the model: how logits become one chosen token.
- Greedy decoding, temperature, top-k, top-p (nucleus) and min-p, with the corner cases that trip up implementations.
- Repetition penalties, stop conditions and reproducible randomness.
- The generation loop, and the first hint of why it should not recompute the past.
You will build
engine/sampling.py: sample and generate_stream, the decoding layer of your engine.
Time: 3-4 hours. GPU: not needed.
The model proposes, the policy decides
The model’s job ends at logits: one score per vocabulary entry for the next position. Turning those scores into one token is a separate decoding policy, and the same model behaves very differently under different policies. Keeping them separate in code is good engineering: you can test the policy on fixed logits, and change it per request without touching the model.
Greedy decoding
Take the highest-scoring token: logits.argmax(-1). No softmax is needed, since the largest logit always has the largest probability. Greedy is deterministic, which makes it ideal for testing (two implementations of the same model must produce identical greedy text). It’s also the best choice for short factual answers and for code.
Its weakness shows on open-ended text. Always taking the locally most likely token tends to fall into loops (“the coree to doree of the coree to doree…”, as the early-stopped model in Chapter 7 did), because the model’s own repeated output makes the repetition ever more likely.
Temperature
To sample instead, convert logits to probabilities and draw from them. Temperature $\tau$ reshapes the distribution first:
$$ p_i = \operatorname{softmax}(z / \tau)_i . $$
For logits $[2, 1, 0]$:
| τ | probabilities |
|---|---|
| 0.25 | [0.982, 0.018, 0.000] |
| 0.5 | [0.867, 0.117, 0.016] |
| 1.0 | [0.665, 0.245, 0.090] |
| 2.0 | [0.507, 0.307, 0.186] |
| 10 | [0.367, 0.332, 0.301] |
Low temperature sharpens toward greedy, and high temperature flattens toward uniform. $\tau = 1$ is the model’s own distribution. $\tau = 0$ would divide by zero, so implementations treat it as a separate greedy branch.
Truncation: top-k, top-p and min-p
Even at $\tau = 1$, the long tail of a 150,000-token vocabulary holds real probability mass. Thousands of individually unlikely tokens add up, and drawing one of them occasionally derails a generation. Truncation removes the tail before sampling:
- Top-k keeps the $k$ highest-scoring tokens and renormalizes. Simple, but fixed: $k = 40$ is too many when the model is certain and too few when many continuations are reasonable.
- Top-p (nucleus sampling, Holtzman et al. 2020) keeps the smallest set of most-likely tokens whose total probability reaches $p$. For probabilities [0.6, 0.25, 0.1, 0.05] and $p = 0.8$, it keeps the first two tokens (0.6 + 0.25 = 0.85). The set adapts: it’s small when the model is confident and large when it isn’t. Note the boundary rule: keep the token that crosses $p$, not just the tokens before it. Otherwise $p = 0.5$ would keep nothing here.
- Min-p keeps tokens whose probability is at least $p_{\min}$ times the top token’s. For [0.5, 0.3, 0.15, 0.04, 0.01] and $p_{\min} = 0.1$, the threshold is 0.05, so it keeps the first three. It scales its cutoff with the model’s confidence.
The order of operations matters: temperature, then top-k, then top-p, then min-p, then renormalize and draw. Production engines apply them in this order (sometimes with options to reorder).
def sample(logits, temperature=0.0, top_k=None, top_p=None, min_p=None, generator=None):
"""Choose one token per row of logits [B, V]; returns [B, 1]. (Your engine: Chapter 8)
temperature == 0 is greedy argmax. Otherwise: divide by temperature, keep the top_k
scores, keep the smallest prefix of the sorted distribution whose mass reaches top_p
(always keeping the first token that crosses it), drop tokens whose probability is below
min_p times the top probability, then draw from what is left.
"""
if temperature < 0:
raise ValueError("temperature must be >= 0")
if temperature == 0:
return logits.argmax(dim=-1, keepdim=True)
scores = logits.float() / temperature
if top_k is not None:
if not 1 <= top_k <= scores.shape[-1]:
raise ValueError("top_k must be in [1, vocab]")
threshold = scores.topk(top_k, dim=-1).values[..., -1:]
scores = scores.masked_fill(scores < threshold, float("-inf"))
if top_p is not None:
if not 0 < top_p <= 1:
raise ValueError("top_p must be in (0, 1]")
ordered, order = scores.sort(dim=-1, descending=True)
probs = ordered.softmax(-1)
# Remove a token when the mass BEFORE it already reaches top_p.
remove = (probs.cumsum(-1) - probs) >= top_p
ordered = ordered.masked_fill(remove, float("-inf"))
scores = torch.full_like(scores, float("-inf")).scatter(-1, order, ordered)
if min_p is not None:
probs = scores.softmax(-1)
scores = scores.masked_fill(probs < min_p * probs.amax(-1, keepdim=True), float("-inf"))
return torch.multinomial(scores.softmax(-1), 1, generator=generator)
struct Rng { // xorshift64*: tiny, seeded, reproducible
uint64_t s;
explicit Rng(uint64_t seed) : s(seed * 0x9E3779B97F4A7C15ull | 1) {}
float uniform() { s ^= s >> 12; s ^= s << 25; s ^= s >> 27; return float((s * 0x2545F4914F6CDD1Dull) >> 40) / float(1ull << 24); }
};
inline size_t sample_top_p(std::vector<float> logits, float temperature, float top_p, Rng& rng) {
std::vector<size_t> order(logits.size());
std::iota(order.begin(), order.end(), 0);
if (temperature == 0) return size_t(std::max_element(logits.begin(), logits.end()) - logits.begin());
std::sort(order.begin(), order.end(), [&](size_t a, size_t b) { return logits[a] > logits[b]; });
std::vector<float> p(order.size());
for (size_t r = 0; r < order.size(); ++r) p[r] = logits[order[r]] / temperature;
softmax(p.data(), p.size());
float mass = 0;
size_t keep = 0;
while (keep < p.size() && mass < top_p) mass += p[keep++]; // keep the token that crosses top_p
float r = rng.uniform() * mass;
for (size_t i = 0; i < keep; ++i) { if (r < p[i]) return order[i]; r -= p[i]; }
return order[keep - 1];
}
#![allow(unused)]
fn main() {
pub struct Policy {
pub temperature: f32,
pub top_k: Option<usize>,
pub top_p: Option<f32>,
}
pub fn sample(logits: &[f32], policy: &Policy, rng: &mut Rng) -> usize {
let argmax = || logits.iter().enumerate().fold(0, |best, (i, &x)| if x > logits[best] { i } else { best });
if policy.temperature == 0.0 {
return argmax();
}
// Sort token ids by score, highest first, and keep only the candidates the policy allows.
let mut order: Vec<usize> = (0..logits.len()).collect();
order.sort_by(|&a, &b| logits[b].partial_cmp(&logits[a]).unwrap());
let keep = policy.top_k.unwrap_or(logits.len()).min(logits.len());
let mut probs: Vec<f32> = order[..keep].iter().map(|&i| logits[i] / policy.temperature).collect();
softmax(&mut probs);
if let Some(p) = policy.top_p {
let (mut mass, mut cut) = (0.0, probs.len());
for (rank, q) in probs.iter().enumerate() {
if mass >= p {
cut = rank; // the token that crossed p (at rank - 1) stays
break;
}
mass += q;
}
probs.truncate(cut);
let total: f32 = probs.iter().sum();
probs.iter_mut().for_each(|q| *q /= total);
}
let mut r = rng.uniform();
for (rank, q) in probs.iter().enumerate() {
if r < *q {
return order[rank];
}
r -= q;
}
order[probs.len() - 1]
}
}
The Python version handles a batch of rows at once and implements every filter with masking (-inf for removed tokens), so it runs on the GPU without sorting through Python lists. The C++ and Rust versions sort one row explicitly, which makes the top-p boundary rule easy to see.
Repetition penalties
The repetition penalty (from the CTRL paper) discourages tokens already present in the context: positive logits are divided by the penalty, and negative ones multiplied by it. Penalties of 1.05-1.2 reduce loops, but they also discourage legitimately repeated words like names and code identifiers. Related variants subtract a fixed amount per occurrence (the frequency penalty) or once per token seen (the presence penalty). They’re all heuristics layered on the model’s distribution, so use them sparingly.
def repetition_penalty(logits, previous_ids, penalty=1.1):
"""CTRL-style penalty: shrink positive logits and grow negative logits of tokens already seen."""
if penalty == 1.0:
return logits
seen = torch.zeros_like(logits, dtype=torch.bool).scatter(-1, previous_ids, True)
adjusted = torch.where(logits > 0, logits / penalty, logits * penalty)
return torch.where(seen, adjusted, logits)
Seeds and reproducibility
Pass an explicit torch.Generator seeded per request. Then the same seed, model, policy and prompt give the same text. In a server, every request needs its own generator: if requests share one, the tokens drawn for one user depend on which other users happened to be in the batch. Reproducibility also depends on the device and kernel versions, so record them (Chapter 18’s manifest does).
The generation loop
Generation repeats: run the model, take the last position’s logits, choose a token, append it, stop if it’s a stop token or the budget is spent:
@torch.inference_mode()
def generate_stream(model, ids, max_new_tokens, cached=True, eos_ids=(), **policy):
"""Yield one new token ID at a time for a single prompt ids [1, P]. (Your engine: Chapter 8)
Prefill the whole prompt once; its last logits choose the first new token. Each later step
feeds only the newest token (cached) or the whole sequence again (uncached).
"""
if ids.ndim != 2 or ids.shape[0] != 1 or ids.shape[1] == 0:
raise ValueError("generate expects one non-empty prompt of shape [1, P]")
if ids.shape[1] + max_new_tokens > model.context_limit:
raise ValueError("Prompt plus new tokens exceed the model's context")
model.eval()
if max_new_tokens == 0:
return
cache = model.new_cache(1, ids.shape[1] + max_new_tokens) if cached else None
sequence = ids
logits = model(ids, cache)
for step in range(max_new_tokens):
token = sample(logits[:, -1], **policy)
yield int(token) # .item(): one host sync per token (Chapter 19 removes it)
if int(token) in eos_ids or step + 1 == max_new_tokens:
return
sequence = torch.cat((sequence, token), dim=1)
logits = model(token, cache) if cached else model(sequence)
def generate(model, ids, max_new_tokens, cached=True, eos_ids=(), **policy):
"""The prompt followed by the generated tokens, as a [1, P+N] tensor."""
new = list(generate_stream(model, ids, max_new_tokens, cached, eos_ids, **policy))
return torch.cat((ids, torch.tensor([new], dtype=ids.dtype, device=ids.device)), dim=1)
A few rules, each of which fixes a bug that shows up in real systems:
- The prompt’s last logits choose the first new token. No extra forward pass is needed before the first sample.
- Stop conditions: a stop token (EOS, or
<|im_end|>for chat models), a maximum number of new tokens, or the model’s context limit. Decide whether the stop token is included in the output. Here it is; a chat UI would hide it. - Reject, don’t silently crop. If prompt plus budget exceed the context, raise an error. Silently dropping the start of the prompt changes what the model sees without anyone noticing.
- One host synchronization per token.
int(token)copies the token to the CPU, which waits for the GPU to finish. To stream text you have to do that eventually, but Chapter 19 shows how to avoid doing it every step.
The cached flag switches between feeding only the new token (with a KV cache, Chapter 16) and re-running the whole sequence. Both must give identical tokens, and your milestone test checks that. The cache is what makes generation affordable; for now, model.new_cache comes provided.
Watching policies on your trained model
Prompting the model you trained in Chapter 7 (python run.py generate) shows each policy’s character, even on a tiny overfit model:
--- greedy
I HAD always thought Jack Gisburn rather a cheap genies be sply oweagal note that Emperors of thereerly a
--- t=0.8 top_k=20
I HAD always thought Jack Gisburn rat Gisburn wife's biggartw through the hush, why Jvert is, and he said, my enough the mre up the
--- t=1.0 top_p=0.9
I HAD always thought Jack Gisburn rather tap get not feltt."
Beyond sampling
Two other ideas are worth knowing by name. Beam search keeps the $b$ most likely partial sequences instead of one. It’s useful for translation, but it produces bland, repetitive text from LLMs and is rarely used for chat. Constrained decoding masks logits so the output must follow a grammar (valid JSON, a regex, a function signature). It’s a masked_fill before sampling, the same as top-k, just with a smarter mask.
Build it
Engine milestone 8: decoding. In engine/sampling.py, implement sample (greedy, temperature, top-k, top-p, min-p) and generate_stream. generate and repetition_penalty are provided.
pytest tests/test_ch08_generation.py
python run.py generate --impl engine --checkpoint runs/model.pt
The tests check that top-k=1 equals greedy, that top-p keeps exactly the crossing token, that seeded sampling is reproducible, and that cached and uncached generation produce identical tokens.
Tip
Implement top-p on sorted probabilities:
remove = (cumsum - probs) >= top_pmarks every token whose preceding mass already reached $p$, so the crossing token survives. Then scatter the masked scores back to vocabulary order withscatter.
Stretch exercises
- ★ Add
stop_stringssupport togenerate: stop when the decoded text ends with any given string. Why is this harder than stopping on token IDs? Where:generate/generate_streaminengine/sampling.py; supply a tokenizer for decoded-text matching. - ★★ Implement presence and frequency penalties and compare them with the repetition penalty on your Chapter 7 model. Where: add penalty helpers in
engine/sampling.pyand call them fromgenerate_stream. - ★★ Implement beam search with width 4 and compare its output with greedy decoding. Which has the higher total log-probability? Which reads better? Where: add a beam-search helper in
engine/sampling.py. - ★★★ Implement constrained decoding that forces outputs to be a valid decimal number: precompute, for each vocabulary entry, whether it can continue a partial number, and mask the rest. Where: add a numeric-output mask helper in
engine/sampling.pyand apply it beforesample.
Check your understanding
- Why is temperature 0 implemented as a separate branch?
- For probabilities [0.4, 0.3, 0.2, 0.1] and top-p 0.5, which tokens remain?
- Why should every request in a server have its own random generator?
- Why does generation not need an extra forward pass before choosing the first new token?
Going deeper
- BALLM §5.3 (pp. 151-158): temperature scaling and top-k sampling.
- Holtzman et al., The Curious Case of Neural Text Degeneration (2020): why greedy and pure sampling fail, and nucleus sampling.
- Keskar et al., CTRL (2019) for the repetition penalty; Nguyen et al., Turning Up the Heat: Min-p Sampling (2024).
- The vLLM and llama.cpp sampler implementations, to see the full set of production options (logit bias, penalties, grammars).