skip to content

Why is chaining a softmax into a separate log less stable than one fused log-softmax?

level: seniorimportance: should knowfreq 44%

answer

  1. the intermediate is where precision dies
  2. tail rounds to exactly zero
  3. one infinite row poisons the batch mean
  4. confident-correct loses its digits too
  5. a subtraction of two order-one terms

basics

~20 s

The chain materializes a probability first, and that intermediate destroys information: a tail value rounds to exactly zero, so its logarithm is negative infinity. Fusing computes the log-probability as a subtraction that never forms the probability at all.

solid answer

~50 s

Fused log-softmax computes `log p_i = (x_i - m) - log(sum_j exp(x_j - m))` with `m` the row maximum - a difference of two order-one quantities, exact to full relative precision even when the answer is `-400`. The chained version computes `p_i` first, and that intermediate is where the damage happens. If `x_i` sits far enough below the maximum, `p_i` rounds to exactly `0`, and `log(0)` is negative infinity: one row makes the batch mean infinite and the next backward pass produces NaN parameters. At the other end, a confident correct prediction gives `p_i` so close to `1` that the true loss - say `1e-9` - is not representable, so the chain reports exactly `0` and the gradient signal vanishes. Patching it by clamping the probability away from zero removes the infinity but caps the loss and flattens the gradient precisely on the examples the model gets most wrong.

go deeper

for a junior

Know the rule even if you cannot yet derive it: hand raw scores to the loss, do not compute probabilities and then take their logarithm yourself. An infinite or NaN loss is often that mistake.

for a middle

Explain both ends of the precision failure - a tail probability rounding to exactly zero so its logarithm is negative infinity, and a near-one probability rounding to one so its loss reads as zero - and write the fused expression that avoids forming the probability.

for a senior

Show diagnostic judgment: recognise an infinite batch loss as an unfused or unclamped log-likelihood, explain why an epsilon clamp is a masking patch rather than a fix, and give the stable binary form straight from the logit.

for a principal

Own the contract: your serving and storage interfaces should carry logits, with probabilities derived downstream, so that calibration, ensembling and threshold changes stay possible and no team can accidentally route a loss through a lossy intermediate.

## The two paths Given logits `x_1 ... x_n` and true class `c`, negative log-likelihood is `-log p_c`. There are two ways to get there. **Chained:** ``` p_i = exp(x_i) / sum_j exp(x_j) # a softmax layer loss = -log(p_c) # a separate log, then a gather ``` **Fused:** ``` m = max_j x_j log_p_i = (x_i - m) - log(sum_j exp(x_j - m)) loss = -log_p_c ``` In exact arithmetic they agree. In floating point they do not, and the disagreement is not a rounding curiosity - it is the difference between a run that trains and a run that dies. ## Failure one: the tail rounds to zero A probability is stored in the same finite format as everything else. When a class score sits far below the row maximum, its normalized probability is not merely small, it falls off the bottom of the representable range and becomes *exactly* `0`. In single precision that happens once the gap to the maximum is more than roughly `104`; the value that should have been `1e-46` is stored as zero. Then `log(0)` is negative infinity, so the per-example loss is `+inf`. Averaging over the batch propagates that to every element - a mean containing one infinity is infinity - and the backward pass turns infinities into NaN gradients and NaN parameters. From that step onward every output is NaN and the run is unrecoverable. The fused form never computes that probability. It computes `log p_i` directly as `(x_i - m) - log(sum_j exp(x_j - m))`. With a gap of `-400`, the first bracket is `-400` and the second is a small number in `[0, log n]`, so the result is about `-400`: an entirely ordinary double or float, carrying full precision, exactly the large loss the optimizer needs in order to punish that prediction. ## Failure two: the confident end loses its digits The opposite extreme is subtler and quieter. Suppose the model is right and confident, so the true probability is `1 - 1e-9`. In single precision that value is not representable at all; it rounds to `1.0`. The chain then computes `-log(1.0) = 0`, reporting a loss of exactly zero when it should be `1e-9`. The gradient computed from that is zero too, so an example that still had a little signal left contributes nothing. This is catastrophic cancellation: subtracting two nearly equal quantities and keeping only the noise in the low bits. The fused form avoids it because `log p_c` is `(x_c - m)` minus a small positive number, and both operands are order one, so the tiny negative result retains its significant digits. ## Failure three: the softmax itself may overflow If the softmax in the chain does not internally subtract the row maximum, a large logit makes `exp(x_i)` infinite, the normalizer infinite, and the ratio NaN before the log is even reached. Even a well-implemented softmax that does shift internally cannot help the log that comes after it - the information has already been squeezed into the `[0, 1]` interval. ## Why clamping is not the fix The usual patch is to clamp the probability into `[eps, 1 - eps]` before taking the log. It does stop the infinity, and that is the whole of its virtue. What it also does: - **It caps the loss.** However wrong the model is, the reported loss saturates at `-log(eps)`. Two rows, one badly wrong and one catastrophically wrong, look identical. - **It flattens the gradient where you need it most.** Inside the clamp the output is constant, so the derivative through it is zero. The examples the model handles worst stop contributing any learning signal - the exact inverse of what you want. - **It silently changes the objective.** You are no longer minimizing the negative log-likelihood, and the training curve you compare against a baseline is measuring a different quantity. - **It hides a real bug.** An infinite loss often means exploding logits or a mislabelled target. Clamping suppresses the alarm and lets the underlying problem persist. Fusing is strictly better: it is exact, it costs nothing extra, and it leaves the objective alone. ## The binary case The same argument applies to a single sigmoid score. Naively, `p = sigmoid(x)` then `-[y log p + (1 - y) log(1 - p)]`. At `x = +50` the sigmoid rounds to exactly `1.0`, so `log(1 - p)` is `log(0)` - infinite loss for any negative-labelled example. Rewritten straight from the logit, with label `y` in `{0, 1}`: ``` loss = max(x, 0) - x * y + log(1 + exp(-abs(x))) ``` Check it. At `x = +50, y = 0`: `50 - 0 + log(1 + exp(-50))`, which is about `50` - a large finite loss, correct for a confident wrong prediction. At `x = +50, y = 1`: `50 - 50 + ~0`, about `0`. At `x = -50, y = 1`: `0 + 50 + ~0`, about `50`. The `abs(x)` inside the exponential guarantees the argument is never positive, so nothing overflows, and the `max(x, 0)` term carries the linear behaviour that the logarithm can no longer represent. ## The operational rule Define losses on logits. If you need probabilities for reporting, calibration or a downstream consumer, compute them on a separate path from the same logits - just never route the loss through them. And treat the appearance of an infinite loss as a signal to investigate, not a signal to clamp.

  • Write binary cross-entropy directly from a logit x and label y so it stays finite at x = plus or minus 50.
    `loss = max(x, 0) - x * y + log(1 + exp(-abs(x)))`. The absolute value keeps the exponential's argument non-positive so it can never overflow, and `max(x, 0)` supplies the linear growth for large magnitudes. At `x = 50, y = 0` it gives about `50`; at `x = 50, y = 1` about `0`; at `x = -50, y = 1` about `50`. The naive sigmoid-then-log path gives infinity in two of those three cases.
  • Is clamping the probability to [1e-7, 1 - 1e-7] before the log a sufficient fix?
    No. It removes the infinity but caps the loss at `-log(1e-7)` and makes the derivative through the clamp zero, so the worst-predicted examples stop producing any gradient. It also changes the objective you are optimizing and masks the underlying cause - usually exploding logits or a bad label - which is the thing actually worth finding.
  • Your serving stack needs calibrated probabilities. Does using a fused loss stop you from producing them?
    Not at all. Compute the loss from logits and produce probabilities on a separate path from those same logits for reporting or thresholding. The rule is one-directional: probabilities may be derived from logits, but the loss must never be derived from probabilities. Persist the logits too, since any class below the underflow floor is unrecoverable from a stored probability.

saying these in an interview costs you the question

  • Assumes floating point is precise enough that both paths agree
  • Adds an epsilon inside the log and calls the numerics fixed
  • Treats a single infinite row as harmless because it is one example
  • Believes only very small probabilities lose precision, not near-one ones
  • Assumes single precision has enough range that log of zero cannot occur

context