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 z∈RK\vz\in\R^K and temperature T>0T>0, softmax is

pi=softmax⁡(z/T)i=exp⁡(zi/T)∑jexp⁡(zj/T).(12.1)p_i = \softmax(\vz/T)_i = \frac{\exp(z_i/T)}{\sum_j \exp(z_j/T)} .\tag{12.1}

Small TT sharpens the distribution; large TT moves it toward uniform. Adding the same constant to every logit changes neither the numerator ratios nor the probabilities:

softmax⁡(z+c1)=softmax⁡(z).(12.2)\softmax(\vz+c\one)=\softmax(\vz).\tag{12.2}

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: T<1T<1 widens them and makes sampling more greedy, while T>1T>1 compresses them and raises entropy. The argmax is unchanged for any positive TT, 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:

log⁡pi=zi−m−log⁡∑jexp⁡(zj−m),m=max⁡jzj.(12.3)\log p_i = z_i - m - \log\sum_j \exp(z_j-m), \qquad m=\max_j z_j .\tag{12.3}

With temperature, apply the formula to z/T\vz/T. This is the log-sum-exp pattern from Section B.7.

The log form is not an optional refinement. A wrong implementation computes exp⁡(zi)\exp(z_i) 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.

Listing 12.1 Stable softmax and log-softmax
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 T=1T=1, differentiate pi=ezi/Zp_i=e^{z_i}/Z, where Z=∑kezkZ=\sum_k e^{z_k}. If i=ji=j,

∂pi∂zi=eziZ−ezieziZ2=pi(1−pi).\frac{\partial p_i}{\partial z_i} =\frac{e^{z_i}Z-e^{z_i}e^{z_i}}{Z^2}=p_i(1-p_i).

If i≠ji\ne j,

∂pi∂zj=−eziezjZ2=−pipj.\frac{\partial p_i}{\partial z_j} =-\frac{e^{z_i}e^{z_j}}{Z^2}=-p_i p_j .

Both cases combine into the Jacobian

J=diag⁡(p)−pp⊤.(12.4)J = \diag(\vp)-\vp\vp^\T .\tag{12.4}

Rows sum to zero because a shift in all logits changes no probability. With temperature, the Jacobian gains a factor 1/T1/T. 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 diag⁡(p)−pp⊤\diag(\vp)-\vp\vp^\T is the covariance matrix of a one-hot draw from p\vp, so it is positive semidefinite and has the all-ones vector in its nullspace.

Listing 12.2 Softmax Jacobian
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 p−y\vp-\vy

Let y\vy be a target distribution over classes, usually one-hot, with ∑iyi=1\sum_i y_i=1. The cross-entropy loss on logits is

L(z,y)=−∑iyilog⁡pi=−∑iyizi+log⁡∑jezj.(12.5)L(\vz,\vy)=-\sum_i y_i \log p_i =-\sum_i y_i z_i+\log\sum_j e^{z_j}.\tag{12.5}

The second equality substitutes log-softmax and uses ∑iyi=1\sum_i y_i=1. Now the gradient is immediate:

∂L∂zj=−yj+ezj∑kezk=pj−yj.(12.6)\frac{\partial L}{\partial z_j} = -y_j + \frac{e^{z_j}}{\sum_k e^{z_k}} = p_j-y_j .\tag{12.6}

For a batch mean, divide by the batch size. For temperature TT, the derivative is (p−y)/T(\vp-\vy)/T. 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 00. 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.

Listing 12.3 Fused softmax-cross-entropy
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 aa, then

softmax⁡([0,a])2=ea1+ea=σ(a).(12.7)\softmax([0,a])_2=\frac{e^a}{1+e^a}=\sigma(a).\tag{12.7}

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 KK logits because the target can be any of KK classes, but it should still avoid computing probabilities and logs in separate unstable steps.

Label smoothing replaces a one-hot target by yϵ=(1−ϵ)y+ϵ1/K\vy^\epsilon=(1-\epsilon)\vy+\epsilon\one/K [szegedy2015rethinking]. The derivation above still works because the smoothed target sums to 1, so the gradient is p−yϵ\vp-\vy^\epsilon. 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]:

Lz=λ(log⁡Z)2,∂Lz∂zj=2λlog⁡Z pj.(12.8)L_z=\lambda(\log Z)^2, \qquad \frac{\partial L_z}{\partial z_j}=2\lambda\log Z\,p_j .\tag{12.8}

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 ∂log⁡Z/∂zj=pj\partial \log Z/\partial z_j=p_j, squaring the log-normalizer gives 2λlog⁡Z pj2\lambda\log Z\,p_j. 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.

Listing 12.4 Binary softmax helper
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].

Key equations
pi=ezi/T∑jezj/T,softmax⁡(z+c1)=softmax⁡(z)p_i=\frac{e^{z_i/T}}{\sum_j e^{z_j/T}}, \qquad \softmax(\vz+c\one)=\softmax(\vz)
log⁡pi=zi−m−log⁡∑jezj−m,m=max⁡jzj\log p_i=z_i-m-\log\sum_j e^{z_j-m}, \qquad m=\max_j z_j
Jsoftmax⁡=diag⁡(p)−pp⊤J_{\softmax}=\diag(\vp)-\vp\vp^\T
L=−∑iyilog⁡pi,∇zL=p−yL=-\sum_i y_i\log p_i, \qquad \nabla_{\vz}L=\vp-\vy
yϵ=(1−ϵ)y+ϵ1/K,∇zλ(log⁡Z)2=2λlog⁡Z p\vy^\epsilon=(1-\epsilon)\vy+\epsilon\one/K, \qquad \nabla_{\vz}\lambda(\log Z)^2=2\lambda\log Z\,\vp

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.

  1. Write pi=ezi/Zp_i=e^{z_i}/Z, then subtract max⁡z\max z before exponentiating.

  2. Differentiate pip_i for i=ji=j and i≠ji\ne j to get diag⁡(p)−pp⊤\diag(\vp)-\vp\vp^\T.

  3. Substitute log-softmax into cross-entropy and let ∑iyi=1\sum_i y_i=1 collapse the log-sum-exp.

  4. Point to the result: zˉ=p−y\bar{\vz}=\vp-\vy.

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

Exercise 12.1 ★ Shift and temperature

Show that adding cc to every logit leaves softmax unchanged. For logits (2,1,−1)(2,1,-1), compute how the largest probability changes when TT is 0.5, 1, and 2.

Exercise 12.2 ★★ The Jacobian

Derive ∂pi/∂zj\partial p_i/\partial z_j for the cases i=ji=j and i≠ji\ne j. Combine the result into diag⁡(p)−pp⊤\diag(\vp)-\vp\vp^\T, and explain why every row sums to zero.

Exercise 12.3 ★★ Cross-entropy gradient

Starting from log-softmax, derive ∇zL=p−y\nabla_{\vz}L=\vp-\vy for any target distribution y\vy that sums to 1. Then state the change when softmax uses temperature TT.

Exercise 12.4 ★★★ Fused implementation

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