= Softmax & Cross-Entropy

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. xref:information-theory.adoc[] defines cross-entropy and perplexity as
information quantities; this chapter derives the logit-level math used by backprop.

[#sec-softmax-temperature]
== Softmax, temperature, and stability

For logits stem:[\vz\in\R^K] and temperature stem:[T>0], softmax is

[latexmath#eq-softmax]
++++
p_i = \softmax(\vz/T)_i =
\frac{\exp(z_i/T)}{\sum_j \exp(z_j/T)} .
++++

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

[latexmath#eq-shift-invariance]
++++
\softmax(\vz+c\one)=\softmax(\vz).
++++

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

[latexmath#eq-log-softmax]
++++
\log p_i = z_i - m - \log\sum_j \exp(z_j-m), \qquad m=\max_j z_j .
++++

With temperature, apply the formula to stem:[\vz/T]. This is the log-sum-exp pattern from
xref:numpy.adoc#sec-stability[].

The log form is not an optional refinement. A wrong implementation computes stem:[\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.

.Stable softmax and log-softmax
[source,python]
----
include::../../scratch/softmax_cross_entropy.py[tag=softmax]
----

[#sec-jacobian]
== The softmax Jacobian

Softmax couples the classes: increasing one logit decreases the other probabilities. For
stem:[T=1], differentiate stem:[p_i=e^{z_i}/Z], where stem:[Z=\sum_k e^{z_k}]. If stem:[i=j],

[latexmath]
++++
\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 stem:[i\ne j],

[latexmath]
++++
\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

[latexmath#eq-softmax-jacobian]
++++
J = \diag(\vp)-\vp\vp^\T .
++++

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

.Softmax Jacobian
[source,python]
----
include::../../scratch/softmax_cross_entropy.py[tag=jacobian]
----

[#sec-cross-entropy-gradient]
== Cross-entropy gives stem:[\vp-\vy]

Let stem:[\vy] be a target distribution over classes, usually one-hot, with stem:[\sum_i y_i=1].
The cross-entropy loss on logits is

[latexmath#eq-ce-logits]
++++
L(\vz,\vy)=-\sum_i y_i \log p_i
=-\sum_i y_i z_i+\log\sum_j e^{z_j}.
++++

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

[latexmath#eq-ce-gradient]
++++
\frac{\partial L}{\partial z_j}
= -y_j + \frac{e^{z_j}}{\sum_k e^{z_k}}
= p_j-y_j .
++++

For a batch mean, divide by the batch size. For temperature stem:[T], the derivative is
stem:[(\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 stem:[0]. 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.

.Fused softmax-cross-entropy
[source,python]
----
include::../../scratch/softmax_cross_entropy.py[tag=fused]
----

[#sec-variants]
== Binary, smoothed, and regularized forms

Sigmoid is the two-class softmax. If the negative-class logit is 0 and the positive-class
logit is stem:[a], then

[latexmath#eq-sigmoid-softmax]
++++
\softmax([0,a])_2=\frac{e^a}{1+e^a}=\sigma(a).
++++

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

Label smoothing replaces a one-hot target by
stem:[\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
stem:[\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>>:

[latexmath#eq-z-loss]
++++
L_z=\lambda(\log Z)^2, \qquad
\frac{\partial L_z}{\partial z_j}=2\lambda\log Z\,p_j .
++++

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
stem:[\partial \log Z/\partial z_j=p_j], squaring the log-normalizer gives
stem:[2\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.

.Binary softmax helper
[source,python]
----
include::../../scratch/softmax_cross_entropy.py[tag=extras]
----

[NOTE,caption=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#key-equations]
.Key equations
****
[latexmath]
++++
p_i=\frac{e^{z_i/T}}{\sum_j e^{z_j/T}}, \qquad
\softmax(\vz+c\one)=\softmax(\vz)
++++

[latexmath]
++++
\log p_i=z_i-m-\log\sum_j e^{z_j-m}, \qquad m=\max_j z_j
++++

[latexmath]
++++
J_{\softmax}=\diag(\vp)-\vp\vp^\T
++++

[latexmath]
++++
L=-\sum_i y_i\log p_i, \qquad
\nabla_{\vz}L=\vp-\vy
++++

[latexmath]
++++
\vy^\epsilon=(1-\epsilon)\vy+\epsilon\one/K, \qquad
\nabla_{\vz}\lambda(\log Z)^2=2\lambda\log Z\,\vp
++++
****

[.teach]
[#sec-teach]
== 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 stem:[p_i=e^{z_i}/Z], then subtract stem:[\max z] before exponentiating.
. Differentiate stem:[p_i] for stem:[i=j] and stem:[i\ne j] to get stem:[\diag(\vp)-\vp\vp^\T].
. Substitute log-softmax into cross-entropy and let stem:[\sum_i y_i=1] collapse the log-sum-exp.
. Point to the result: stem:[\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?

[#sec-exercises]
== Exercises

[#ex-softmax-cross-entropy-shift.exercise]
.★ Shift and temperature
====
Show that adding stem:[c] to every logit leaves softmax unchanged. For logits
stem:[(2,1,-1)], compute how the largest probability changes when stem:[T] is 0.5, 1, and 2.
====

[#ex-softmax-cross-entropy-jacobian.exercise]
.★★ The Jacobian
====
Derive stem:[\partial p_i/\partial z_j] for the cases stem:[i=j] and stem:[i\ne j]. Combine
the result into stem:[\diag(\vp)-\vp\vp^\T], and explain why every row sums to zero.
====

[#ex-softmax-cross-entropy-gradient.exercise]
.★★ Cross-entropy gradient
====
Starting from log-softmax, derive stem:[\nabla_{\vz}L=\vp-\vy] for any target distribution
stem:[\vy] that sums to 1. Then state the change when softmax uses temperature stem:[T].
====

[#ex-softmax-cross-entropy-fused.exercise]
.★★★ 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.
====

[bibliography]
[#sec-references]
== References

include::../../book/sources.adoc[tags=bridle1990;goodfellow2016;szegedy2015rethinking;vaswani2017attention;chowdhery2022palm]
