1. What is it?#
The model produces a probability distribution over the vocabulary. Sampling is choosing one token from it.
logits (128256 numbers) → temperature → filtering (top-k, top-p) → softmax → sampleThe choices here determine whether output is repetitive, incoherent, or good — and they have real performance and caching implications.
2. Why does it exist?#
Because always picking the most likely token (greedy decoding) produces degenerate text: repetitive, bland, and prone to loops. Human language is not the maximum-likelihood sequence; it has variety. Holtzman et al.’s “The Curious Case of Neural Text Degeneration” demonstrated this convincingly and gave us nucleus sampling.
For the inference engineer, sampling matters for three reasons: it’s on the critical path of every token, it affects cache keys, and its parameters are a client-facing API surface you must implement correctly.
3. Simple analogy#
Choosing what to say next in a conversation.
Greedy: always say the single most predictable thing. Result: “That’s interesting. That’s interesting. That’s interesting.”
Pure random: pick any word from the dictionary weighted by rough plausibility. Result: word salad.
Nucleus sampling: consider the handful of things that account for 90% of what a reasonable person might say next, and pick among them. Result: natural variety without nonsense.
Temperature is how adventurous you’re feeling.
4. Tiny example#
// sampling.go — what temperature, top-k and top-p do to a distribution.
package main
import (
"fmt"
"math"
"sort"
)
func softmax(logits []float64, temperature float64) []float64 {
out, sum := make([]float64, len(logits)), 0.0
for i, l := range logits {
out[i] = math.Exp(l / temperature)
sum += out[i]
}
for i := range out {
out[i] /= sum
}
return out
}
func show(name string, probs []float64) {
fmt.Printf("%-20s", name)
for _, p := range probs {
fmt.Printf(" %.3f", p)
}
fmt.Println()
}
// filter keeps only the tokens selected by `keep` (given tokens sorted by probability),
// then renormalizes.
func filter(probs []float64, keep func(rank int, cumBefore float64) bool) []float64 {
order := make([]int, len(probs))
for i := range order {
order[i] = i
}
sort.Slice(order, func(a, b int) bool { return probs[order[a]] > probs[order[b]] })
out, cum, sum := make([]float64, len(probs)), 0.0, 0.0
for rank, i := range order {
if keep(rank, cum) {
out[i] = probs[i]
sum += probs[i]
}
cum += probs[i]
}
for i := range out {
out[i] /= sum
}
return out
}
func main() {
logits := []float64{3.0, 2.5, 1.0, 0.5, -1.0, -2.0} // 6-token vocabulary
show("softmax (T=1)", softmax(logits, 1))
show("T=0.5 (sharper)", softmax(logits, 0.5))
show("T=2.0 (flatter)", softmax(logits, 2))
probs := softmax(logits, 1)
// top-k = 3: keep only the 3 most likely
show("top-k=3", filter(probs, func(rank int, _ float64) bool { return rank < 3 }))
// top-p = 0.9: keep the smallest set whose cumulative probability >= 0.9
show("top-p=0.9", filter(probs, func(_ int, cumBefore float64) bool { return cumBefore < 0.9 }))
}softmax (T=1) 0.503 0.305 0.068 0.041 0.009 0.003
T=0.5 (sharper) 0.716 0.263 0.013 0.005 0.000 0.000
T=2.0 (flatter) 0.336 0.262 0.124 0.096 0.045 0.027Temperature divides the logits before softmax. Low T amplifies differences (more deterministic); high T flattens them (more random). T→0 becomes argmax; T→∞ becomes uniform.
Now filtering:
probs := softmax(logits, 1)
// top-k = 3: keep only the 3 most likely
show("top-k=3", filter(probs, func(rank int, _ float64) bool { return rank < 3 }))
// top-p = 0.9: keep the smallest set whose cumulative probability >= 0.9
// (a token is kept if the tokens ranked above it have not yet reached 0.9)
show("top-p=0.9", filter(probs, func(_ int, cumBefore float64) bool { return cumBefore < 0.9 }))top-k=3 0.574 0.348 0.078 0.000 0.000 0.000
top-p=0.9 0.622 0.378 0.000 0.000 0.000 0.000Note the difference: top-k always keeps exactly 3; top-p keeps 2 here because the top 2 already cover 81%… and adding the third exceeds 0.9. Top-p adapts to the distribution’s shape, which is why it’s generally preferred.
5. Technical explanation#
The full sampling pipeline#
1. logits (B, V) from the LM head
2. logit bias add per-token biases (client-specified)
3. repetition/presence/frequency penalties
4. temperature logits /= T
5. top-k keep the k highest
6. top-p keep the smallest set with cumulative prob >= p
7. min-p keep tokens with prob >= min_p × max_prob
8. softmax
9. sample multinomialOrder matters. Applying temperature before top-p (standard) means T changes which tokens survive the nucleus. Applying it after would decouple them. Different implementations have historically differed here — a source of “why does this model behave differently on your platform?” reports.
The parameters#
| Parameter | Range | Effect |
|---|---|---|
temperature | 0 - 2 | 0 = greedy; <1 sharper; >1 more random |
top_k | 1 - vocab | keep the k most likely |
top_p | 0 - 1 | nucleus: smallest set covering p |
min_p | 0 - 1 | keep tokens with p ≥ min_p × p_max — adaptive, robust |
repetition_penalty | 1.0 - 1.5 | divide logits of already-seen tokens |
presence_penalty | -2 - 2 | subtract a constant from seen tokens |
frequency_penalty | -2 - 2 | subtract proportional to count |
seed | int | RNG seed for reproducibility |
min-p is underrated. Unlike top-p, it scales with the model’s confidence: when the model is certain (one token at 0.95), min_p=0.1 keeps only tokens above 0.095 — effectively just that one. When uncertain (top token 0.15), it keeps everything above 0.015 — a wide set. It behaves sensibly across confidence levels where fixed top-p does not.
Greedy (temperature 0)#
Special-case it. Do NOT implement as softmax with T=0 (division by zero).
token = argmax(logits)Greedy is deterministic given the same logits — but logits vary with batch composition (Section IV.12), so greedy is not fully reproducible in a batched server. This surprises people.
Beam search#
Keep the k highest-probability sequences, not tokens:
Step 1: expand each of k beams by all V tokens → k×V candidates
Step 2: keep the k with the highest cumulative log probability
Repeat.Used in translation and some structured tasks. Rarely used in modern LLM chat serving because:
- It multiplies KV cache by k.
- It produces bland, high-likelihood text (exactly what Holtzman showed is bad for open-ended generation).
- It complicates continuous batching (variable, branching sequences).
If you support it, PagedAttention’s copy-on-write sharing makes it much cheaper — the beams share their common prefix.
Structured / constrained decoding#
Forcing output to match a grammar (JSON schema, regex):
At each step, compute the set of tokens allowed by the grammar state,
mask all others to -inf, then sample normally.Implementations: Outlines, XGrammar, llguidance, LM Format Enforcer. The engineering challenge is computing the allowed-token mask fast enough — a naive implementation adds milliseconds per token. Modern ones precompute masks per grammar state and cache them.
Performance note: constrained decoding can be nearly free (a cached bitmask) or can dominate your ITL (recomputing a regex automaton per token). Check which you have.
6. Under the hood#
Sampling must run entirely on the GPU. A device→host copy of (B, 128256) logits plus a
synchronization costs ~0.5-2 ms per step — potentially 20-40% of your ITL.
The expensive part is top-p, which naively requires a full sort of 128k values per sequence:
Naive: sort (B, 128256) ~0.5 ms at B=64
Optimized: top-k first (k=1024 covers p=0.99 almost always), then sort those
~0.05 msMost engines do this: an approximate pre-filter followed by an exact computation on the survivors. The approximation is safe because tokens outside the top 1024 essentially never make the nucleus.
Per-request sampling parameters complicate batching: each sequence may have different T, top_p, etc. Engines handle this with vectorized per-row parameters rather than looping.
7. Performance implications#
Operation Cost at B=64, V=128k
Temperature scale negligible
Full sort for top-p ~0.5 ms
Top-k pre-filter + partial sort ~0.05 ms
Repetition penalty (gather/scatter) ~0.05 ms
Multinomial sample ~0.02 ms
D2H copy of logits (if done!) ~1.5 ms + sync ← avoid
Constrained decoding (naive) 0.5-5 ms ← can dominate
Constrained decoding (cached mask) ~0.02 msSampling should be < 5% of your ITL. If it’s more, you have a D2H copy or a naive sort.
8. Production implications#
- Keep everything on the GPU. Any per-step host round trip is a large ITL cost.
- Support per-request sampling parameters — clients expect it, and vectorizing it is straightforward.
- Include sampling parameters in any cache key. Two requests with the same prompt but different temperature are different requests.
- Special-case temperature 0.
- Validate parameter ranges at the gateway.
top_p=0ortemperature=-1should be a 400, not a NaN. - Document your parameter semantics. The order of operations differs between implementations and users will notice.
- Seed support is valuable for debugging and for customers who need reproducibility — but be honest that it’s only reproducible at fixed batch composition.
- Measure constrained-decoding overhead before promising JSON mode at low latency.
9. Common mistakes#
Sampling on the CPU. Adds milliseconds per token.
Implementing T=0 as division by zero.
Full 128k sort per step. Use a top-k pre-filter.
Forgetting sampling params in the cache key. Serves a temperature-0 response to a temperature-1 request.
Applying penalties after softmax. They operate on logits.
Assuming greedy is reproducible in a batched server. It isn’t (Section IV.12).
Naive constrained decoding. Can 10x your ITL.
Not validating client parameters. top_k=-5 should not reach the kernel.
10. Hands-on exercise#
A. Implement the pipeline. Write the full sampling pipeline from section 5 in PyTorch, vectorized over a batch with per-row parameters. Verify against a reference implementation.
B. Measure the cost. Benchmark each stage at B ∈ {1, 32, 128} with V=128256. Which dominates? Implement the top-k pre-filter optimization and re-measure.
C. Compare strategies. Generate 20 completions from the same prompt with: greedy, T=0.7 + top_p=0.9, T=1.0 pure, top_k=50, min_p=0.05. Compare diversity (distinct n-grams) and quality (by reading them). Which would you default to?
D. min-p vs top-p. Construct two distributions — one peaked, one flat — and show how top-p and min-p select different numbers of tokens in each. Explain why min-p is more robust.
E. Constrained decoding cost. Measure ITL with and without JSON-schema-constrained decoding using your engine. Is the implementation cached or naive?
11. Interview questions#
- Explain temperature, top-k, and top-p. Why is top-p usually preferred over top-k?
- What is min-p and what problem does it solve that top-p doesn’t?
- Why must sampling happen on the GPU?
- How would you make top-p sampling fast for a 128k vocabulary?
- Why is greedy decoding not fully reproducible in a batched server?
- Why is beam search uncommon in LLM chat serving?
- How does constrained decoding work and what does it cost?
- What must be included in a response cache key?
12. Further reading#
- [FUNDAMENTAL] Holtzman et al., “The Curious Case of Neural Text Degeneration” (2019) — nucleus sampling
- [ESTABLISHED] min-p sampling paper and discussion
- [REFERENCE] vLLM
sampling_params.pyand the sampler implementation - [REFERENCE] Outlines / XGrammar documentation for constrained decoding
- Next: 14 — Streaming inference