Sample the Next Token
The same prompt gave two different answers, and someone proposed "turning the temperature down" without being able to say what that does. Write the sampler.
A model gives every vocabulary token a score, a logit; the sampler turns the scores into one choice. temperature and top_p are sampler settings.
next_token_probs(logits, temperature=1.0, top_p=1.0) returns one probability per token, in the same order as logits, summing to 1.
- Temperature. Divide every logit by
temperature, then take the softmax.temperature == 0means greedy: probability 1 for the highest logit and 0 for the rest. On a tie, the lowest index wins. - Stability. The exponential of 1000 overflows a float, so subtract the largest logit from all of them before exponentiating. Softmax only depends on the differences, so the result is the same.
- Top-p (nucleus). Rank the tokens from most to least probable, ties in index order. Keep the smallest leading group whose probabilities add up to at least
top_p(there is always at least one), set every other token to 0 and rescale the survivors to sum to 1.top_p == 1.0keeps everything.
sample(probs, u) picks a token index given a number u in [0, 1): walk the tokens in order with a running total of their probabilities and return the first index at which the total exceeds u. A token with probability 0 must never be returned.
Use math; no other imports are needed.