PidokuInfra

Sampling and Decoding Strategies

Basic Intermediate 1h 15m Difficulty 3/5 Topic 13 of 15

Prerequisites III.07, 04


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 → sample

The 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#

Go
// 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.027

Temperature 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:

Go
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.000

Note 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          multinomial

Order 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#

ParameterRangeEffect
temperature0 - 20 = greedy; <1 sharper; >1 more random
top_k1 - vocabkeep the k most likely
top_p0 - 1nucleus: smallest set covering p
min_p0 - 1keep tokens with p ≥ min_p × p_max — adaptive, robust
repetition_penalty1.0 - 1.5divide logits of already-seen tokens
presence_penalty-2 - 2subtract a constant from seen tokens
frequency_penalty-2 - 2subtract proportional to count
seedintRNG 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.

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 ms

Most 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 ms

Sampling 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=0 or temperature=-1 should 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#

  1. Explain temperature, top-k, and top-p. Why is top-p usually preferred over top-k?
  2. What is min-p and what problem does it solve that top-p doesn’t?
  3. Why must sampling happen on the GPU?
  4. How would you make top-p sampling fast for a 128k vocabulary?
  5. Why is greedy decoding not fully reproducible in a batched server?
  6. Why is beam search uncommon in LLM chat serving?
  7. How does constrained decoding work and what does it cost?
  8. 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.py and the sampler implementation
  • [REFERENCE] Outlines / XGrammar documentation for constrained decoding
  • Next: 14 — Streaming inference

↑↓ navigate↵ openesc close