Why does a numerically stable log-sum-exp subtract the row maximum before exponentiating?
answer
- the exponential has a hard ceiling
- make the largest exponent zero
- factor a constant out of the sum
- it becomes an additive term outside the log
- shifted sum sits between 1 and n
basics
~20 sSubtracting 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 sThe 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 linesimport 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 errorgo deeper
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.
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.
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.
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