Chapter 12
Softmax & Cross-Entropy
Temperature, log-sum-exp stability, the softmax Jacobian, and why the gradient is p minus y.
Softmax turns arbitrary logits into a categorical distribution, and cross-entropy tells a classifier how surprising the correct class was. Together they are the final layer and loss of most language models: one vector of logits over the vocabulary, one stable fused loss, and one simple gradient. Chapter 7 defines cross-entropy and perplexity as information quantities; this chapter derives the logit-level math used by backprop.
12.1 Softmax, temperature, and stability
For logits and temperature , softmax is
Small sharpens the distribution; large moves it toward uniform. Adding the same constant to every logit changes neither the numerator ratios nor the probabilities:
Only differences between logits matter. A model can add 100 to every vocabulary score without changing its prediction, so the absolute logit level is not identifiable from cross-entropy alone. Temperature acts on those differences: widens them and makes sampling more greedy, while compresses them and raises entropy. The argmax is unchanged for any positive , but the probability assigned to non-argmax classes can change dramatically.
That identity is the numerical trick. Before exponentiating, subtract the maximum logit. The largest shifted value is 0, so no exponential overflows, and the common shift cancels. The same idea gives a stable log-softmax:
With temperature, apply the formula to . This is the log-sum-exp pattern from Section B.7.
The log form is not an optional refinement. A wrong implementation computes first, then divides, then takes a logarithm for the loss. With logits near 1000, the exponential already overflowed before the loss sees it. Stable log-softmax computes the normalizer in the shifted space and returns log-probabilities directly. The fused loss below therefore stores one set of log-probabilities for the forward pass and reuses the probabilities only for the backward pass.
def logsumexp(x, axis=-1, keepdims=False):
"""Stable log(sum(exp(x))) along an axis."""
x = np.asarray(x, dtype=np.float64)
shifted = x - np.max(x, axis=axis, keepdims=True)
result = np.max(x, axis=axis, keepdims=True) + np.log(
np.sum(np.exp(shifted), axis=axis, keepdims=True)
)
return result if keepdims else np.squeeze(result, axis=axis)
def log_softmax(logits, temperature=1.0, axis=-1):
z = np.asarray(logits, dtype=np.float64) / temperature
return z - logsumexp(z, axis=axis, keepdims=True)
def softmax(logits, temperature=1.0, axis=-1):
return np.exp(log_softmax(logits, temperature, axis))
12.2 The softmax Jacobian
Softmax couples the classes: increasing one logit decreases the other probabilities. For , differentiate , where . If ,
If ,
Both cases combine into the Jacobian
Rows sum to zero because a shift in all logits changes no probability. With temperature, the Jacobian gains a factor . The tests check vector-Jacobian products from this formula against finite differences.
The matrix also explains why classes compete. The diagonal entries are positive: increasing a logit increases its own probability. The off-diagonal entries are negative: the extra mass must come from the other classes because probabilities sum to one. The form is the covariance matrix of a one-hot draw from , so it is positive semidefinite and has the all-ones vector in its nullspace.
def softmax_jacobian(probabilities):
"""Jacobian of a single softmax vector p: diag(p) - p p^T."""
p = np.asarray(probabilities, dtype=np.float64)
return np.diag(p) - np.outer(p, p)
12.3 Cross-entropy gives
Let be a target distribution over classes, usually one-hot, with . The cross-entropy loss on logits is
The second equality substitutes log-softmax and uses . Now the gradient is immediate:
For a batch mean, divide by the batch size. For temperature , the derivative is . A fused implementation should never materialize unstable probabilities and then take logs; it should compute log-softmax once, return the scalar loss, and return the backward gradient. This is the tiny formula that makes the last layer of a language model easy to train, even when the vocabulary has tens of thousands of classes.
The gradient has a useful conservation law: its components sum to . The target class receives a negative gradient when its probability is below 1, which raises its logit under gradient descent. Every other class receives a positive gradient proportional to its predicted probability, which lowers overconfident wrong logits most. For soft targets, such as smoothed labels or a teacher distribution, the same formula moves the model toward the whole target distribution rather than toward a single class.
def smooth_one_hot(labels, num_classes, epsilon=0.0):
y = one_hot(np.asarray(labels, dtype=np.int64), num_classes)
return (1 - epsilon) * y + epsilon / num_classes
def softmax_cross_entropy(logits, labels, temperature=1.0,
label_smoothing=0.0, z_loss=0.0):
"""Return mean loss and gradient with respect to logits."""
logits = np.asarray(logits, dtype=np.float64)
targets = smooth_one_hot(labels, logits.shape[-1], label_smoothing)
log_p = log_softmax(logits, temperature)
p = np.exp(log_p)
batch = logits.shape[0]
loss = -np.sum(targets * log_p) / batch
grad = (p - targets) / (batch * temperature)
if z_loss:
log_z = logsumexp(logits, axis=-1)
loss = loss + z_loss * np.mean(log_z ** 2)
grad = grad + (2 * z_loss / batch) * log_z[:, None] * softmax(logits)
return float(loss), grad
12.4 Binary, smoothed, and regularized forms
Sigmoid is the two-class softmax. If the negative-class logit is 0 and the positive-class logit is , then
That is why binary logistic regression and two-class softmax regression have the same probability model, just with a redundant logit removed.
This equivalence is also a numerical hint. Binary classification can use a one-logit BCE-with-logits loss, because only the logit difference matters. Multiclass classification keeps all logits because the target can be any of classes, but it should still avoid computing probabilities and logs in separate unstable steps.
Label smoothing replaces a one-hot target by [szegedy2015rethinking]. The derivation above still works because the smoothed target sums to 1, so the gradient is . Smoothing prevents the target class from demanding probability 1 and assigns a small amount of loss pressure to every other class.
The cost is that a smoothed model is no longer trained to put all available probability on the observed label. That can improve calibration in classification, but it also changes the maximum-likelihood objective. For language modelling, use it deliberately: it lowers the penalty for plausible alternatives but also prevents the empirical next token from being the only target.
PaLM adds a z-loss to keep the log-normalizer small [chowdhery2022palm]:
It is not a replacement for cross-entropy; it is an extra penalty on the scale of the logits.
The derivative is another application of the chain rule. Since , squaring the log-normalizer gives . Unlike shift-invariant cross-entropy, z-loss depends on the absolute logit level. That is the point: it discourages the model from drifting to huge logits that leave probabilities unchanged but make optimization numerically harsher.
def sigmoid_as_two_class_softmax(logit):
"""The positive-class probability of softmax([0, logit])."""
pairs = np.stack([np.zeros_like(logit), logit], axis=-1)
return softmax(pairs)[..., 1]
|
In practice
|
Transformer language models train a linear vocabulary head with softmax cross-entropy [vaswani2017attention]. Production kernels usually fuse the shift, log-sum-exp, loss, and backward pass to avoid extra memory traffic and unstable intermediates. Temperature is used at sampling time to change entropy without retraining. Label smoothing is common in classification workloads, while PaLM reports using z-loss for training stability at scale [chowdhery2022palm]. |
12.5 Teach it
The one-sentence version. Softmax makes logits into probabilities; cross-entropy with a one-hot target sends back exactly "predicted minus correct."
An analogy. Logits are race scores. Softmax turns score gaps into win probabilities. The loss asks how much probability the winner received, and the gradient moves probability mass from overpredicted classes to the target.
At the board.
-
Write , then subtract before exponentiating.
-
Differentiate for and to get .
-
Substitute log-softmax into cross-entropy and let collapse the log-sum-exp.
-
Point to the result: .
Misconceptions to address.
-
"Softmax needs probabilities as input." It takes logits, not normalized values.
-
"Cross-entropy and softmax are separate in backprop." They should be fused for stability.
-
"Temperature changes the ranking." Positive temperature changes confidence, not argmax.
Check for understanding. Why can you subtract the maximum logit without changing the probabilities?
12.6 Exercises
Show that adding to every logit leaves softmax unchanged. For logits , compute how the largest probability changes when is 0.5, 1, and 2.
Derive for the cases and . Combine the result into , and explain why every row sums to zero.
Starting from log-softmax, derive for any target distribution that sums to 1. Then state the change when softmax uses temperature .
Use softmax_cross_entropy to compute a stable mean loss and gradient for a tiny batch.
Gradient-check the logits with label smoothing and with z-loss enabled.
References
-
[goodfellow2016] I. Goodfellow, Y. Bengio, and A. Courville. Deep Learning. MIT Press, 2016. https://www.deeplearningbook.org
-
[bridle1990] J. S. Bridle. Probabilistic interpretation of feedforward classification network outputs, with relationships to statistical pattern recognition. In Neurocomputing, Springer, 1990.
-
[chowdhery2022palm] A. Chowdhery et al. PaLM: Scaling Language Modeling with Pathways. 2022. arXiv:2204.02311
-
[szegedy2015rethinking] C. Szegedy et al. Rethinking the Inception Architecture for Computer Vision. 2015. arXiv:1512.00567
-
[vaswani2017attention] A. Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762