Softmax With Temperature
Problem
A language model scores every candidate next token with a raw number called a logit, and sampling turns those scores into probabilities with a softmax: divide each logit by the temperature T, exponentiate, and normalise so the results sum to 1. Write softmax(logits, temperature) returning the probabilities rounded to 3 decimals. It must not overflow on large logits such as 1000, and a temperature of 0 means greedy decoding, so all of the probability goes to the highest logit, the first one on a tie.
Examples
Input: logits = [2.0, 1.0, 0.1], temperature = 1.0
Output: [0.659, 0.242, 0.099]
Input: logits = [2.0, 1.0, 0.1], temperature = 0.5
Output: [0.864, 0.117, 0.019]
Why: a low temperature stretches the gaps, so the favourite takes more
Input: logits = [1000.0, 999.0], temperature = 1.0
Output: [0.731, 0.269]
Why: edge case, exp(1000) overflows, but only the differences between logits matter
Hints
0 / 3
Softmax only cares about differences between logits: adding the same constant to all of them leaves every probability unchanged.
Subtract the largest scaled logit before exponentiating. The biggest term becomes exp(0) = 1 and nothing can overflow.
Treat temperature 0 separately, because dividing by it is undefined: return 1.0 at the index of the maximum and 0.0 everywhere else. Otherwise scale, shift by the maximum, exponentiate, divide by the sum and round.
Solution
Dividing by the temperature before the softmax scales the gaps between logits: below 1 the gaps grow and the distribution sharpens toward the top token, above 1 they shrink and it flattens toward uniform. Subtracting the maximum scaled logit first is the standard stability trick, since it multiplies the numerator and the denominator by the same factor, leaving the answer unchanged while keeping every exponent at or below zero. Temperature 0 is the limit of that sharpening and cannot be computed by division, so it is handled as greedy decoding directly. Time and space are O(v) for v logits.
import math
def softmax(logits, temperature=1.0):
if temperature == 0: # greedy: all the mass on the top logit
best = logits.index(max(logits))
return [1.0 if i == best else 0.0 for i in range(len(logits))]
scaled = [x / temperature for x in logits]
top = max(scaled) # shift by the max so exp never overflows
exps = [math.exp(x - top) for x in scaled]
total = sum(exps)
return [round(e / total, 3) for e in exps]
print(softmax([2.0, 1.0, 0.1])) # -> [0.659, 0.242, 0.099]
print(softmax([2.0, 1.0, 0.1], 0.5)) # -> [0.864, 0.117, 0.019]
print(softmax([2.0, 1.0, 0.1], 5.0)) # -> [0.4, 0.327, 0.273]
print(softmax([1000.0, 999.0])) # -> [0.731, 0.269]
print(softmax([3.0, 5.0, 5.0], 0)) # -> [0.0, 1.0, 0.0]Stuck on the idea rather than the code? Temperature and Sampling covers it.