Softmax
How a network turns raw scores into a probability distribution: the exp-over-sum formula, why exp, the subtract-the-max stability trick, and how temperature dials it from greedy argmax to coin-flip random.
What you'll learn
- The formula softmax(z)_i = exp(z_i) / sum_j exp(z_j) and why exp is the right choice
- The subtract-the-max trick — exact, not approximate — that every library uses
- Temperature — T below 1 sharpens toward argmax, T above 1 flattens toward uniform
- Softmax vs sigmoid — one competing distribution vs independent per-class scores
- Where it lives in practice — classifier output layers, attention, and LLM sampling
Before you start
The raw scores are called logits — the unnormalized outputs of the last linear layer. Softmax is the bridge from logits to a probability distribution: every output lands in the open interval between 0 and 1, and together they sum to exactly 1. One formula, three jobs across modern deep learning.
The formula
For a vector of logits z with K entries, the probability assigned to class i is:
Plug in the example: softmax([2.0, 1.0, 0.1]) gives [0.659, 0.242, 0.099],
which sums to 1. Class 0 had the biggest logit, so it gets the biggest
probability — softmax never reorders the classes.
Why exp, of all functions?
You could imagine normalizing logits some other way — divide each by the sum,
say. Softmax uses exp for four reasons that all matter at once:
That last one is the deepest. A hard argmax — “just pick the biggest” — is
flat almost everywhere, so its gradient is zero or undefined and a network
can’t learn through it. Softmax is a smooth, differentiable stand-in. Its
other name, softargmax, says it plainly: a soft, trainable version of
argmax.
Notice too that only the differences between logits matter — the ratio of two
probabilities, p_i / p_j, equals exp(z_i - z_j). Add the same constant to
every logit and nothing changes. Hold onto that fact.
See it across temperatures
The figure below is the whole lesson in one place: the same logits at three temperatures. The probability bars always sum to 100% — they sharpen as temperature drops and flatten as it rises.
The subtract-the-max trick
Here is the detail that separates a textbook formula from production code.
exp(1000) overflows to infinity in floating point, and your softmax returns
NaN.
The fix exploits the fact you just learned: softmax is unchanged when you add a constant to every logit. So before exponentiating, subtract the largest logit from all of them:
The biggest logit becomes 0, so its exponential is exp(0) = 1 — the largest
term can never overflow. Everything else is exp of a negative number, safely
between 0 and 1.
The result is bit-for-bit the mathematically correct answer, just computed without blowing up. This “safe softmax” is what every mainstream library does under the hood. Here it is in three lines of NumPy:
import numpy as np
def softmax(z, T=1.0):
z = np.asarray(z, dtype=float) / T
z = z - z.max() # the stability trick — exact, prevents overflow
e = np.exp(z)
return e / e.sum()
logits = [2.0, 1.0, 0.1]
p = softmax(logits)
print("probs:", np.round(p, 4))
print("sum: ", p.sum()) # exactly 1.0
# Big logits that would overflow a naive exp() — safe softmax handles them.
print("huge: ", np.round(softmax([1000.0, 1000.0, 1000.0]), 4))
probs: [0.659 0.2424 0.0986]
sum: 1.0
huge: [0.3333 0.3333 0.3333]
Temperature — one knob from greedy to random
Divide the logits by a temperature T before softmax: softmax(z / T). It is
a single dial over how peaked the distribution is.
| Temperature | Effect | Limit |
|---|---|---|
T = 1 | Plain softmax | — |
T < 1 | Sharpens — more peaked, more confident/greedy | T -> 0 becomes one-hot at the argmax |
T > 1 | Flattens — closer to uniform, more random | T -> infinity becomes uniform, 1/K each |
On [2.0, 1.0, 0.1]: at T = 0.5 you get [0.86, 0.12, 0.02] (sharper);
at T = 2 you get [0.50, 0.30, 0.19] (flatter).
This is the LLM temperature knob — T = 0 is greedy decoding (always the top
token), higher T gives more diverse, surprising output. The mental model to
keep: softmax is a temperature-controllable soft argmax.
Crank T down and it hardens into argmax; crank it up and it melts into a
uniform guess.
A common trap, worth saying out loud: higher temperature makes the model less confident, not more. People get this backwards constantly.
Softmax vs sigmoid — they are not interchangeable
This is the confusion that bites people in code review. A sigmoid squashes one logit into one independent probability. Stack N sigmoids and each class is scored on its own — the outputs need not sum to 1. That is multi-label: an image can be both “outdoor” and “sunset” at once.
Softmax couples all the classes into one distribution that competes and sums to
- That is multi-class, single-label: exactly one answer is right, so the classes fight over a fixed budget of probability.
Softmax is the multi-class generalization of logistic regression —
and in the two-class case it collapses back to a sigmoid of the logit difference:
softmax([z0, z1])_0 = sigmoid(z0 - z1). Pick by the question you are asking:
“which one?” is softmax; “which ones?” is sigmoid.
Where softmax shows up
Three places, same operation:
- Classifier output layer. Softmax over the final logits gives class probabilities, paired with cross-entropy. The softmax + cross-entropy combination has a beautifully clean gradient: predicted distribution minus the true one.
- Attention. Scaled dot-product attention
is
softmax(Q Kᵀ / sqrt(d_k)) V. Softmax runs row-wise over the similarity scores so each query’s attention weights are non-negative and sum to 1 — a weighted average over the values. The1/sqrt(d_k)scaling keeps the scores from growing so large that softmax saturates and its gradient vanishes. This is the Softmax block inside the transformer. - LLM sampling. The model emits a logit per vocabulary token; divide by temperature and softmax to get the next-token distribution you sample from, often after top-k or top-p truncation.
In one breath
- Softmax turns logits into a probability distribution: exponentiate each, divide by the sum, so every output is in (0,1) and they sum to 1.
- exp is the right choice because it’s positive, monotonic (ranking preserved), amplifies gaps into ratios (p_i/p_j = e^(z_i−z_j)), and is smooth — a differentiable “soft argmax” you can train through.
- Only differences between logits matter, which is why the subtract-the-max trick is exact, not approximate — it just stops exp from overflowing.
- Temperature is one knob: below 1 sharpens toward the top class (greedy), above 1 flattens toward uniform (random) — higher T means less confident.
- Softmax = one competing single-label distribution; stacked sigmoids = independent multi-label scores. It shows up in classifier heads, attention, and LLM next-token sampling.
Quick check
Quick check
Next
Softmax is the last step of the classifier; the loss it feeds is the engine of learning. Continue to loss functions to see why softmax and cross-entropy are an inseparable pair, or to sampling to see temperature drive an LLM’s word choices.
Practice this in an interview
All questionsSoftmax turns class logits into positive values that sum to one by exponentiating and normalizing them. It is used for mutually exclusive multiclass classification because the result forms a categorical distribution and pairs naturally with cross-entropy loss, although models are usually trained directly from logits for numerical stability.
We divide each query-key dot product by the square root of the key dimension because its standard deviation grows as the square root of the key dimension. Without this temperature adjustment, softmax saturates toward one-hot weights and sends very small gradients; the scaling keeps logits in a learnable range.
FlashAttention computes the same dense softmax attention as the standard formula, but tiles the work so the N by N score matrix stays out of HBM and is never materialized. Online softmax keeps the result exact in mathematical terms, reducing attention's memory footprint from quadratic in sequence length to linear while leaving its quadratic compute cost unchanged.
Under the usual zero-mean, unit-variance assumptions, a query-key dot product has variance dₖ and standard deviation √dₖ. Dividing by √dₖ keeps the logits at a stable scale before softmax, reducing dimension-dependent saturation and preserving useful gradients.