skip to content

Why does a numerically stable log-sum-exp subtract the row maximum before exponentiating?

level: middleimportance: must knowfreq 58%

answer

  1. the exponential has a hard ceiling
  2. make the largest exponent zero
  3. factor a constant out of the sum
  4. it becomes an additive term outside the log
  5. shifted sum sits between 1 and n

basics

~20 s

Subtracting the row maximum caps every exponent at zero, so the exponential cannot overflow. Adding it back outside the logarithm leaves the value provably identical: factoring it out of the sum turns it into an added constant.

solid answer

~50 s

The identity is `log(sum_j exp(x_j)) = m + log(sum_j exp(x_j - m))` with `m = max_j x_j`. It holds exactly: every term `exp(x_j)` equals `exp(m) * exp(x_j - m)`, so `exp(m)` factors out of the sum and the log turns it into an additive `m`. This is algebra, not an approximation. The point is what the shifted form does to the floating-point range: the largest shifted exponent is `exp(0) = 1`, so nothing can overflow, and the inner sum is bounded between `1` and `n`, making the inner log land in `[0, log n]`. A translation decoder emitting a logit of `+800` overflows the naive form immediately - the exponential's ceiling is around `709` in double precision and around `88` in single - while the shifted form returns a finite answer. Terms far below the maximum underflow to zero after the shift, and that is harmless.

code

python · 15 lines
python
import math

def log_sum_exp(xs):
    m = max(xs)
    return m + math.log(sum(math.exp(x - m) for x in xs))

logits = [800.0, 799.0, 795.0]

print(log_sum_exp(logits))               # ~800.3182, finite
print(log_sum_exp(logits) - logits[0])   # ~0.3182 = -log p for class 0

try:
    print(math.log(sum(math.exp(x) for x in logits)))
except OverflowError as err:
    print('naive:', err)                 # naive: math range error

go deeper

for a junior

Recall the shape of the fix: find the largest score in the row, subtract it before exponentiating, add it back outside the logarithm. Know it exists because the exponential overflows on large inputs.

for a middle

Derive the identity on the spot by factoring exp(m) out of the sum, and state the resulting range: shifted terms at most 1, the sum between 1 and n, the inner logarithm between 0 and log n. Explain why the maximum and not some other constant.

for a senior

Demonstrate that you know where this primitive hides - every log-probability, cross-entropy, mixture likelihood and beam score routes through it - and that a value outside the m to m plus log n band is a bug signature you can act on.

for a principal

Frame the tradeoff you actually decide: whether numerics like this belong in hand-written model code at all, or should be confined to a small set of reviewed, tested primitives that every team calls, so that a stability bug is fixed once rather than rediscovered per project.

## The quantity Log-sum-exp of a row of scores `x_1 ... x_n` is ``` LSE(x) = log(exp(x_1) + exp(x_2) + ... + exp(x_n)) ``` It is the normalizer of a softmax in log space: the log-probability of class `i` is exactly `x_i - LSE(x)`. Anything that computes a log-probability from raw scores computes an LSE somewhere. ## The identity Let `m = max_j x_j`. Then ``` exp(x_j) = exp(m) * exp(x_j - m) ``` for every `j`, purely because the exponential turns addition into multiplication. Factor the common `exp(m)` out of the sum: ``` sum_j exp(x_j) = exp(m) * sum_j exp(x_j - m) ``` Take the log of both sides, and the log of the product becomes a sum: ``` LSE(x) = m + log(sum_j exp(x_j - m)) ``` That is the shifted form. Note what has *not* happened: no term was dropped, no series was truncated, no error bound was invoked. In exact arithmetic the two expressions are the same number. The shift is a rearrangement chosen so that finite precision can carry it. ## Why the shift saves the computation **Overflow disappears by construction.** After the shift, every argument `x_j - m` is at most `0`, so every `exp(x_j - m)` is at most `1`. The sum of `n` such terms is at most `n`. Since the maximum entry contributes exactly `exp(0) = 1`, the sum is also at least `1`. So the inner sum always lies in `[1, n]` and the inner log always lies in `[0, log n]` - a range so small that no floating-point format has trouble with it, whatever the scale of the inputs. Compare the naive form. Floating-point exponentials have a hard ceiling: `exp(x)` overflows to infinity once `x` exceeds roughly `88` in single precision and roughly `709` in double. Those thresholds are not exotic. A decoder over a fifty-thousand-word vocabulary that has become very confident can put a logit at `+800`; `exp(800)` is infinity, the sum is infinity, `log(inf)` is infinity, and the loss and every gradient derived from it are ruined. The shifted form on the same row returns a perfectly ordinary number just above `800`. **Underflow becomes harmless.** After the shift, a term whose score is far below the maximum evaluates to something like `exp(-400)`, which rounds to exactly `0`. That is not a problem, and this is the part candidates most often get wrong. The sum is already at least `1` from the maximum's own term, so dropping a term that would have contributed less than the last representable bit changes the sum by less than a rounding error. Contrast this with the unshifted, unfused path, where an underflowing term becomes a probability of exactly zero and then `log(0)` produces negative infinity - there the underflow is fatal because the log is taken *of* the tiny number rather than of a sum that is safely order-one. ## Why the maximum specifically Algebraically any constant `c` works: `LSE(x) = c + log(sum_j exp(x_j - c))`. The maximum is chosen because it is the only choice that guarantees both properties at once. It makes the largest exponent exactly `0`, so overflow is impossible for any input, and it guarantees at least one term of size `1`, so the inner sum can never underflow to zero and make the outer log blow up in the other direction. Shifting by the mean, for instance, still overflows on a row with one enormous outlier. Shifting by the minimum makes every exponent non-negative and overflows even more readily. ## A useful bound that falls out Because the inner sum lies in `[1, n]`, ``` m <= LSE(x) <= m + log(n) ``` So log-sum-exp is a smooth upper bound on the maximum, never more than `log n` above it. With one dominant score it is essentially the maximum; with `n` equal scores it is exactly `m + log n`. That is why LSE is sometimes called a soft maximum - and it is a good sanity check when reading a computed value, since a result far outside that band means a bug. ## Where it is used The log-probability of class `i` is `x_i - LSE(x)`, which is the shifted form written out as `(x_i - m) - log(sum_j exp(x_j - m))`. Cross-entropy for the true class `c` is simply `LSE(x) - x_c`. Entropies, mixture likelihoods, KL terms and beam-search score accumulation all reduce to the same primitive. Getting the shift right once means every one of those is stable.

  • After the shift, terms far below the maximum underflow to exactly zero - does that corrupt the result?
    No. The maximum's own term contributes exactly `1`, so the sum is already order one; a term that underflows would have contributed less than the smallest representable fraction of that sum, so dropping it changes the answer by less than a rounding error. Underflow is fatal only when you take the log *of* the tiny number, which is what the unfused probability path does.
  • Could you shift by the mean instead of the maximum?
    Algebraically yes - any constant cancels the same way. But the mean does not bound the largest exponent. A row with one score at `+800` and the rest near zero has a mean far below `800`, so the shifted maximum is still enormous and the exponential still overflows. The maximum is the only shift that makes overflow impossible for every input.
  • How far can log-sum-exp be from the plain maximum of the row?
    Between `0` and `log(n)` above it, where `n` is the number of entries. The lower bound comes from the maximum's own term, the upper bound from all `n` shifted terms being at most `1`. One dominant score puts it essentially at the maximum; `n` tied scores put it exactly at `m + log n`. It is a smooth upper bound on the max.

saying these in an interview costs you the question

  • Calls the max shift an approximation that introduces small error
  • Shifts by a fixed constant or the mean and assumes overflow is impossible
  • Panics that small terms underflow to zero after the shift
  • Cannot state that the shifted sum lies between 1 and n
  • Thinks the shift is a speed optimization rather than a range fix

context