Stable numerics: logsumexp, softmax, compensated sums
Overview
Section titled “Overview”| Module | M09.2 · build · Python · Pass 2 · 2 to 3 h |
| You build | python/tinyllm/num/stable.py: logsumexp, softmax, log_softmax, kahan_sum, pairwise_sum |
| Contract | course/contracts/py/tinyllm/num/stable.pyi |
| Tests | course/tests/M09.2/ (what they check: section 4) |
| Needs | no code dependency · reading: M09.1 IEEE 754 (why exp overflows float32 at 88.7), M02.1 Taylor series |
| Used by | M11.1 entropy from logits · M08.3 the cross-entropy VJP · L0.2 op library · L0.3 fused cross-entropy · later: L8.1 sampler, L6.7 perplexity over tokens, L9.2 the same math in C · later: L4.4, L5.1, L7.7 |
| Milestone | MS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains) |
| Optional depth | Higham, Accuracy and Stability of Numerical Algorithms (SIAM, 2nd ed.), ch. 1 and 4; Blanchard, Higham, and Higham, “Accurately computing the log-sum-exp and softmax functions” (IMA J. Numer. Anal., 2021) |
Key Takeaways
Section titled “Key Takeaways”- with : after the shift the largest exponential is , so nothing overflows, and the answer is finite for logits of any size (
test_no_overflow,test_shift_invariance). log_softmaxisx - logsumexp(x), neverlog(softmax(x)): the second takes the log of an underflowed 0 and returns where the answer is (test_no_underflow_in_log_softmax).- A fully masked row (all ) is defined, not NaN: softmax gives zeros, log-sum-exp and log-softmax give (
test_fully_masked_row). - Plain summation of floats can be off by about ; pairwise summation cuts that to and Kahan’s compensation to about , independent of (
test_kahan_error_bound,test_pairwise_error_bound).
How to work this chapter
Section titled “How to work this chapter”ol start M09.2 # stubs stable.py into your repo, contract alongsideol tests M09.2 # read the test catalog first: rung R0, you write no tests hereol check M09.2 # exit code is the verdictol diff M09.2 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Your tracer bigram (L0.0) already hides a numerically careful softmax: its nll subtracts each row’s maximum before exp. In this pass that one line becomes a dozen call sites. Your autograd op library (L0.2) needs softmax and log-softmax as ops, the fused cross-entropy (L0.3) needs log-softmax at the target, the sampler (L8.1) turns temperature-scaled logits into probabilities, and the model zoo (L6.7) sums ten million per-token losses into one perplexity. Written the obvious way, each of them fails the first time training succeeds: a confident model produces a logit above 88.7, np.exp overflows float32 to inf, and the loss becomes inf / inf = nan; a confidently wrong token gets probability in float32, and log(0) is an infinite loss with a NaN gradient. This module writes the stable versions once, proves them, and makes every later module import them.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| one row of logits (scores) | float[n] | |
| the shift that prevents overflow | scalar | |
| log-sum-exp | scalar | |
| probabilities from logits | float[n], sums to 1 | |
| log-probabilities from logits | float[n] | |
| unit roundoff: for float32, for float64 | scalar | |
| the floating-point result of an operation | scalar | |
| the exact sum | scalar | |
| the computed sum | scalar | |
| Kahan’s compensation: the low-order part lost so far | scalar |
Where exp breaks. M09.1 showed that the largest float32 is about and the largest float64 about . So np.exp(89.0) in float32 is inf. At the other end, underflows to exactly 0 below about in float32 (counting subnormals) and in float64. Logits after training routinely exceed 100 in magnitude, and masks write on purpose.
The shift identity. For any constant , , so
Choose . Every shifted exponent is at most 0, so every term is at most 1, and the largest term is exactly . The sum lies in : it cannot overflow, it is never 0, and its log lies in . Nothing is lost by the shift, because the identity is exact.
Softmax is shift invariant. Dividing by the sum cancels the factor :
So compute . The same identity lets the C kernel in L9.2 process a row in one pass: when a new maximum appears, it rescales the running sum by instead of starting over.
Log-softmax from the identity, not from softmax. . Computed this way it is a subtraction of two moderate numbers. Computed as log(softmax(x)), it first rounds , which underflows to 0 once in float32, and then takes log(0) = -inf. The cross-entropy loss is at the target , so this is the difference between a large finite loss (with a useful gradient) and an infinite one.
Masks. A masked entry is , so and it gets probability exactly 0. If every entry is masked, and is NaN. The contract defines that row instead of propagating NaN: replace a non-finite by 0, so the shifted entries stay , the sum is 0, softmax returns zeros (divide by 1 when the sum is 0), and log-sum-exp returns . A padded position in a batch is exactly such a row.
Rounding in a sum. Every floating addition rounds: with . Summing left to right, the -th partial sum carries the rounding errors of all earlier additions, and the first terms pass through additions. The standard bound is
For float32 values that is a relative error up to about : useless. In practice errors partly cancel, but when the addends have the same sign (losses, counts, probabilities) they accumulate.
Pairwise summation. Split the array in halves, sum each half the same way, add the two results. The recursion is levels deep, and each element takes part in only one addition per level, so
For that is 24 instead of . numpy’s np.sum does this internally (with blocks of 8 at the leaves), which is why the tests use np.cumsum, a strictly sequential sum, to show the problem.
Kahan’s compensated summation. Keep a second variable holding the part the running sum could not absorb. For each :
When , is computed exactly, and it is the part of that actually made it into ; subtracting leaves minus the part that was rounded away. Algebraically , which is the point: in floating point it is the rounding error, recovered exactly and fed back into the next addend. The error bound becomes
which no longer grows with to first order. It costs four operations per element and a sequential loop, so it is used where the count is huge and the budget is not: running totals of losses, token counts, and perplexity.
Work in the input’s dtype. The point of compensation is to get extra precision without a wider type. Both sums here keep every operation in the dtype of (float32 stays float32) and return a Python float.
3. Worked example by hand
Section titled “3. Worked example by hand”Log-sum-exp and softmax of .
| step | values |
|---|---|
| 3 | |
| sum | |
| sum | |
| sum | |
| softmax $= e^{x-m} / $ sum | |
| log-softmax |
The softmax row sums to 1, and confirms the last entry. Adding 1000 to every entry changes to 1003 and leaves every shifted value, and so the softmax, unchanged.
A masked row : , shifted , exponentials , softmax , log-softmax . A fully masked row : becomes 0, every exponential is 0, softmax is , and log-sum-exp is .
Summing in float32. At consecutive float32 values are 2 apart, so is exactly halfway between two floats and rounds to the even one, . The exact sum is 2.
- Left to right: , , . Both ones are lost.
- Pairwise: . The left pair loses its 1; is representable, so the right pair keeps its 1. Better, not exact.
- Kahan, one row per element:
| 0 | ||||
| 1 | 1 | (rounded) | ||
| 1 | (exact) | |||
| 2 | 2 |
The second row is the whole idea: the 1 that rounding threw away reappears as , and the next step adds it back. Kahan returns the exact 2.
In float64 the same thing happens one scale up: at floats are 2 apart, so plain summation of gives 0. Pairwise also gives 0 there, because is a tie too and rounds back to ; Kahan gives 2.
These numbers are the first cases in section 4: test_hand_example, test_hand_example_sums, and test_sums_work_in_the_input_dtype.
4. The interface
Section titled “4. The interface”def logsumexp(x: ArrayLike, axis: int = -1, keepdims: bool = False) -> NDArraydef softmax(x: ArrayLike, axis: int = -1) -> NDArraydef log_softmax(x: ArrayLike, axis: int = -1) -> NDArraydef kahan_sum(x: ArrayLike) -> floatdef pairwise_sum(x: ArrayLike) -> floatFloat inputs keep their dtype; integer inputs are computed in float64. axis and keepdims mean what they mean in numpy. The sums flatten their input in C order. pairwise_sum splits at h = n // 2 down to single elements; pairing neighbours bottom up builds the same tree when is a power of two, so a vectorized version is fine. The contract has the full rules.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | the section 3 table for | you and the test agree on the definitions |
test_hand_example_sums | unit | float32 : plain 0, Kahan 2, pairwise 1 | the three sums differ exactly as derived |
test_matches_scipy_golden | golden | scipy’s three functions on seven shapes, axes, and scales | an independent implementation agrees |
test_shift_invariance | property | adding changes nothing but LSE by | the online softmax of L9.2 rescales a running sum |
test_no_overflow | boundary | gives ; stays finite | logits after training |
test_no_underflow_in_log_softmax | boundary | gives , not | the cross-entropy of L0.3 |
test_fully_masked_row | boundary | all gives zeros and , no NaN, no warning | padded rows in a batch |
test_axis_and_keepdims | unit | every axis of a 3-D array, keepdims shapes | attention over [B, H, T, T] |
test_dtype_preserved | unit | float32 in, float32 out; integers become float64 | kernels compare in float32 |
test_softmax_rows_sum_to_one | property | rows of scale 0.1 to 700 sum to 1 | the sampler’s CDF ends at 1 |
test_kahan_error_bound | property | 65536 float32 values within ; plain summation is not | L6.7 sums losses |
test_pairwise_error_bound | property | within for and | the cheap accurate sum |
test_sums_cover_every_element | boundary | lengths 0, 1, odd, a 2-D array, integers | no dropped or invented term |
test_sums_work_in_the_input_dtype | unit | the float64 version of the hand example | compensation, not a wider type, does the work |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. exponentiating before subtracting the max | inf / inf = nan at a float32 logit of 89 | test_no_overflow (mutant s01) |
2. log(softmax(x)) | log-probability, infinite loss, NaN gradient | test_no_underflow_in_log_softmax (mutant s02) |
| 3. subtracting a max of , or dividing by a zero sum | a padded row of NaN that spreads through the batch | test_fully_masked_row (mutants s03, s04, s14) |
| 4. computing the compensation and never using it | Kahan silently becomes plain summation | test_kahan_error_bound (mutants s06, s07) |
| 5. a “pairwise” recursion that peels off one element at a time | a depth- tree, the sequential error, and a recursion limit | test_pairwise_error_bound (mutant s09) |
| 6. forgetting to add back in log-sum-exp | off by exactly the row maximum | test_hand_example (mutant s05) |
| 7. dropping the middle element of an odd split, or indexing an empty array | off by a whole term | test_sums_cover_every_element (mutants s08, s10) |
8. reducing the wrong axis, ignoring keepdims | attention normalized over the batch | test_axis_and_keepdims (mutants s11, s12) |
| 9. upcasting float32 to float64 | the C kernel’s float32 reference no longer matches | test_dtype_preserved (mutant s13) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Forward | L4.4 | Registered call site uses this module. |
| Forward | L5.1 | Registered call site uses this module. |
| Forward | L7.7 | Registered call site uses this module. |
| Direction | Module | How it uses this |
|---|---|---|
| Back | M09.1 | the overflow and underflow thresholds of exp, and the unit roundoff in every bound |
| Back | M02.1 | Taylor’s view of near 0, behind why small shifted exponents are accurate |
| Forward | M11.1 | entropy_from_logits uses log_softmax |
| Forward | M08.3 | cross_entropy_vjp uses softmax |
| Forward | L0.2 | softmax and log-softmax become autograd ops |
| Forward | L0.3 | the fused cross-entropy is at the target |
| Forward | L8.1 | temperature-scaled logits to sampling probabilities |
| Forward | L6.7 | perplexity over tokens accumulates with kahan_sum (through M11.2) |
| Forward | L9.2 | the same math in C, in one pass with a running maximum |
If you skip this module, ol check M11.1 stops with M11.1 needs M09.2: build it, or rerun with --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
softmax | PyTorch torch.softmax | fused vectorized kernels, autocast to float32 inside half-precision models | aten/src/ATen/native/SoftMax.cpp |
logsumexp | scipy.special.logsumexp | weights b (log of a weighted sum), sign handling for negative weights | scipy/special/_logsumexp.py |
| the shift identity | FlashAttention’s online softmax | one pass over keys with a running max and a rescaled running sum | Dao et al., FlashAttention (2022), section 3.1 |
pairwise_sum | numpy np.add.reduce | pairwise with 8-way unrolled leaf blocks, chosen for speed and accuracy | numpy/_core/src/umath/loops_utils.h.src |
kahan_sum | Python math.fsum | Shewchuk’s algorithm: the exactly rounded sum, at the cost of a list of partials | Modules/mathmodule.c (math_fsum) |