Sort logits descending, compute softmax, cumsum, mask tokens past the smallest prefix whose cumulative ≥ p, renormalize.
import torch
def sample(logits, temperature=1.0, top_k=0, top_p=1.0):
logits = logits / max(temperature, 1e-8)
if top_k > 0:
kth = torch.topk(logits, top_k).values[..., -1, None]
logits = logits.masked_fill(logits < kth, float("-inf"))
if top_p < 1.0:
sorted_logits, idx = torch.sort(logits, descending=True)
probs = torch.softmax(sorted_logits, dim=-1)
cum = torch.cumsum(probs, dim=-1)
remove = cum - probs > top_p # keep the token that crosses the threshold
sorted_logits = sorted_logits.masked_fill(remove, float("-inf"))
logits = torch.full_like(logits, float("-inf")).scatter(-1, idx, sorted_logits)
probs = torch.softmax(logits, dim=-1)
return torch.multinomial(probs, 1)
Gotchas: forgetting to apply temperature before top-k/p; masking tokens to 0 instead of −inf (lets them survive softmax on tie scores); not seeding the RNG and getting flaky tests; off-by-one in the cumulative cutoff ("strictly greater" vs "≥") — the token that crosses the top-p threshold must be kept, otherwise top_p=0.1 with a 0.5-prob top token would remove everything.
Follow-ups: What does temperature → 0 converge to? (Greedy argmax.) Why does top-p adapt better than top-k across confident vs uncertain distributions? Why is generation non-deterministic even at temperature 0 on real serving stacks? (Batching non-determinism, floating-point reduction order.)