Skip to content

Stable numerics: logsumexp, softmax, compensated sums

ModuleM09.2 · build · Python · Pass 2 · 2 to 3 h
You buildpython/tinyllm/num/stable.py: logsumexp, softmax, log_softmax, kahan_sum, pairwise_sum
Contractcourse/contracts/py/tinyllm/num/stable.pyi
Testscourse/tests/M09.2/ (what they check: section 4)
Needsno code dependency · reading: M09.1 IEEE 754 (why exp overflows float32 at 88.7), M02.1 Taylor series
Used byM11.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 10710^7 tokens, L9.2 the same math in C · later: L4.4, L5.1, L7.7
MilestoneMS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains)
Optional depthHigham, 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)
  • log⁡∑iexi=m+log⁡∑iexi−m\log\sum_i e^{x_i} = m + \log\sum_i e^{x_i - m} with m=max⁡ixim = \max_i x_i: after the shift the largest exponential is e0=1e^0 = 1, so nothing overflows, and the answer is finite for logits of any size (test_no_overflow, test_shift_invariance).
  • log_softmax is x - logsumexp(x), never log(softmax(x)): the second takes the log of an underflowed 0 and returns −∞-\infty where the answer is −104-10^4 (test_no_underflow_in_log_softmax).
  • A fully masked row (all −∞-\infty) is defined, not NaN: softmax gives zeros, log-sum-exp and log-softmax give −∞-\infty (test_fully_masked_row).
  • Plain summation of nn floats can be off by about nu∑∣xi∣n u \sum|x_i|; pairwise summation cuts that to ⌈log⁡2n⌉u∑∣xi∣\lceil\log_2 n\rceil u \sum|x_i| and Kahan’s compensation to about 2u∣S∣2u|S|, independent of nn (test_kahan_error_bound, test_pairwise_error_bound).
Terminal window
ol start M09.2 # stubs stable.py into your repo, contract alongside
ol tests M09.2 # read the test catalog first: rung R0, you write no tests here
ol check M09.2 # exit code is the verdict
ol diff M09.2 # after passing: your code against the reference

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 e−200=0e^{-200} = 0 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.

SymbolMeaningType / shape
x∈Rnx \in \mathbb{R}^none row of logits (scores)float[n]
m=max⁡ixim = \max_i x_ithe shift that prevents overflowscalar
LSE(x)=log⁡∑iexi\mathrm{LSE}(x) = \log\sum_i e^{x_i}log-sum-expscalar
softmax(x)i=exi/∑jexj\mathrm{softmax}(x)_i = e^{x_i} / \sum_j e^{x_j}probabilities from logitsfloat[n], sums to 1
logsoftmax(x)i=xi−LSE(x)\mathrm{logsoftmax}(x)_i = x_i - \mathrm{LSE}(x)log-probabilities from logitsfloat[n]
uuunit roundoff: 2−24≈6.0×10−82^{-24} \approx 6.0 \times 10^{-8} for float32, 2−53≈1.1×10−162^{-53} \approx 1.1 \times 10^{-16} for float64scalar
fl(a∘b)\mathrm{fl}(a \circ b)the floating-point result of an operation ∘\circscalar
S=∑i=1nxiS = \sum_{i=1}^n x_ithe exact sumscalar
S^\hat Sthe computed sumscalar
ccKahan’s compensation: the low-order part lost so farscalar

Where exp breaks. M09.1 showed that the largest float32 is about 3.4×1038=e88.723.4 \times 10^{38} = e^{88.72} and the largest float64 about e709.78e^{709.78}. So np.exp(89.0) in float32 is inf. At the other end, exe^x underflows to exactly 0 below about −103.9-103.9 in float32 (counting subnormals) and −745.1-745.1 in float64. Logits after training routinely exceed 100 in magnitude, and masks write −∞-\infty on purpose.

The shift identity. For any constant cc, ∑iexi=ec∑iexi−c\sum_i e^{x_i} = e^{c} \sum_i e^{x_i - c}, so

LSE(x)=c+log⁡∑iexi−c.\mathrm{LSE}(x) = c + \log \sum_i e^{x_i - c}.

Choose c=m=max⁡ixic = m = \max_i x_i. Every shifted exponent xi−mx_i - m is at most 0, so every term is at most 1, and the largest term is exactly e0=1e^0 = 1. The sum lies in [1,n][1, n]: it cannot overflow, it is never 0, and its log lies in [0,log⁡n][0, \log n]. Nothing is lost by the shift, because the identity is exact.

Softmax is shift invariant. Dividing by the sum cancels the factor ece^{c}:

softmax(x+c)i=exi+c∑jexj+c=ecexiec∑jexj=softmax(x)i.\mathrm{softmax}(x + c)_i = \frac{e^{x_i + c}}{\sum_j e^{x_j + c}} = \frac{e^{c} e^{x_i}}{e^{c} \sum_j e^{x_j}} = \mathrm{softmax}(x)_i.

So compute exi−m/∑jexj−me^{x_i - m} / \sum_j e^{x_j - m}. The same identity lets the C kernel in L9.2 process a row in one pass: when a new maximum m′m' appears, it rescales the running sum by em−m′e^{m - m'} instead of starting over.

Log-softmax from the identity, not from softmax. log⁡softmax(x)i=xi−LSE(x)\log \mathrm{softmax}(x)_i = x_i - \mathrm{LSE}(x). Computed this way it is a subtraction of two moderate numbers. Computed as log(softmax(x)), it first rounds exi−me^{x_i - m}, which underflows to 0 once xi−m<−104x_i - m < -104 in float32, and then takes log(0) = -inf. The cross-entropy loss is −logsoftmax(x)t-\mathrm{logsoftmax}(x)_t at the target tt, so this is the difference between a large finite loss (with a useful gradient) and an infinite one.

Masks. A masked entry is xi=−∞x_i = -\infty, so exi−m=0e^{x_i - m} = 0 and it gets probability exactly 0. If every entry is masked, m=−∞m = -\infty and xi−m=−∞−(−∞)x_i - m = -\infty - (-\infty) is NaN. The contract defines that row instead of propagating NaN: replace a non-finite mm by 0, so the shifted entries stay −∞-\infty, the sum is 0, softmax returns zeros (divide by 1 when the sum is 0), and log-sum-exp returns log⁡0=−∞\log 0 = -\infty. A padded position in a batch is exactly such a row.

Rounding in a sum. Every floating addition rounds: fl(a+b)=(a+b)(1+δ)\mathrm{fl}(a + b) = (a + b)(1 + \delta) with ∣δ∣≤u|\delta| \le u. Summing left to right, the kk-th partial sum carries the rounding errors of all earlier additions, and the first terms pass through n−1n - 1 additions. The standard bound is

∣S^−S∣≤(n−1) u∑i∣xi∣+O(u2).|\hat S - S| \le (n - 1)\, u \sum_i |x_i| + O(u^2).

For n=107n = 10^7 float32 values that is a relative error up to about 0.60.6: 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 ⌈log⁡2n⌉\lceil \log_2 n \rceil levels deep, and each element takes part in only one addition per level, so

∣S^−S∣≤⌈log⁡2n⌉ u∑i∣xi∣+O(u2).|\hat S - S| \le \lceil \log_2 n \rceil\, u \sum_i |x_i| + O(u^2).

For n=107n = 10^7 that is 24 instead of 10710^7. 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 cc holding the part the running sum ss could not absorb. For each xx:

y=x−c,t=s+y,c=(t−s)−y,s=t.y = x - c, \qquad t = s + y, \qquad c = (t - s) - y, \qquad s = t.

When ∣s∣≥∣y∣|s| \ge |y|, t−st - s is computed exactly, and it is the part of yy that actually made it into tt; subtracting yy leaves minus the part that was rounded away. Algebraically c=0c = 0, 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

∣S^−S∣≤2u∣S∣+O(nu2)∑i∣xi∣,|\hat S - S| \le 2u|S| + O(n u^2) \sum_i |x_i|,

which no longer grows with nn 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 xx (float32 stays float32) and return a Python float.

Log-sum-exp and softmax of x=[1,2,3]x = [1, 2, 3].

stepvalues
m=max⁡xm = \max x3
x−mx - m[−2,−1,0][-2, -1, 0]
ex−me^{x - m}[0.135335,0.367879,1][0.135335, 0.367879, 1]
sum1.5032151.503215
log⁡\log sum0.4076060.407606
LSE=m+log⁡\mathrm{LSE} = m + \log sum3.4076063.407606
softmax $= e^{x-m} / $ sum[0.090031,0.244728,0.665241][0.090031, 0.244728, 0.665241]
log-softmax =x−LSE= x - \mathrm{LSE}[−2.407606,−1.407606,−0.407606][-2.407606, -1.407606, -0.407606]

The softmax row sums to 1, and e−0.407606=0.665241e^{-0.407606} = 0.665241 confirms the last entry. Adding 1000 to every entry changes mm to 1003 and leaves every shifted value, and so the softmax, unchanged.

A masked row [0,−∞,0][0, -\infty, 0]: m=0m = 0, shifted [0,−∞,0][0, -\infty, 0], exponentials [1,0,1][1, 0, 1], softmax [0.5,0,0.5][0.5, 0, 0.5], log-softmax [−ln⁡2,−∞,−ln⁡2][-\ln 2, -\infty, -\ln 2]. A fully masked row [−∞,−∞,−∞][-\infty, -\infty, -\infty]: mm becomes 0, every exponential is 0, softmax is [0,0,0][0, 0, 0], and log-sum-exp is −∞-\infty.

Summing [224,1,1,−224][2^{24}, 1, 1, -2^{24}] in float32. At 224=167772162^{24} = 16777216 consecutive float32 values are 2 apart, so 224+1=167772172^{24} + 1 = 16777217 is exactly halfway between two floats and rounds to the even one, 1677721616777216. The exact sum is 2.

  • Left to right: 224+1→2242^{24} + 1 \to 2^{24}, +1→224+ 1 \to 2^{24}, −224→0- 2^{24} \to 0. Both ones are lost.
  • Pairwise: (224+1)+(1−224)=16777216+(−16777215)=1(2^{24} + 1) + (1 - 2^{24}) = 16777216 + (-16777215) = 1. The left pair loses its 1; 1−224=−167772151 - 2^{24} = -16777215 is representable, so the right pair keeps its 1. Better, not exact.
  • Kahan, one row per element:
xxy=x−cy = x - ct=s+yt = s + yc=(t−s)−yc = (t - s) - yss
2242^{24}2242^{24}2242^{24}02242^{24}
112242^{24} (rounded)(224−224)−1=−1(2^{24} - 2^{24}) - 1 = -12242^{24}
11−(−1)=21 - (-1) = 2224+22^{24} + 2 (exact)2−2=02 - 2 = 0224+22^{24} + 2
−224-2^{24}−224-2^{24}2(2−(224+2))+224=0(2 - (2^{24} + 2)) + 2^{24} = 02

The second row is the whole idea: the 1 that rounding threw away reappears as c=−1c = -1, and the next step adds it back. Kahan returns the exact 2.

In float64 the same thing happens one scale up: at 101610^{16} floats are 2 apart, so plain summation of [1016,1,1,−1016][10^{16}, 1, 1, -10^{16}] gives 0. Pairwise also gives 0 there, because 1−10161 - 10^{16} is a tie too and rounds back to −1016-10^{16}; 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.

python/tinyllm/num/stable.py
def logsumexp(x: ArrayLike, axis: int = -1, keepdims: bool = False) -> NDArray
def softmax(x: ArrayLike, axis: int = -1) -> NDArray
def log_softmax(x: ArrayLike, axis: int = -1) -> NDArray
def kahan_sum(x: ArrayLike) -> float
def pairwise_sum(x: ArrayLike) -> float

Float 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 nn is a power of two, so a vectorized version is fine. The contract has the full rules.

TestKINDChecksWhy it matters downstream
test_hand_exampleunitthe section 3 table for [1,2,3][1, 2, 3]you and the test agree on the definitions
test_hand_example_sumsunitfloat32 [224,1,1,−224][2^{24}, 1, 1, -2^{24}]: plain 0, Kahan 2, pairwise 1the three sums differ exactly as derived
test_matches_scipy_goldengoldenscipy’s three functions on seven shapes, axes, and scalesan independent implementation agrees
test_shift_invariancepropertyadding c∈[−1000,1000]c \in [-1000, 1000] changes nothing but LSE by ccthe online softmax of L9.2 rescales a running sum
test_no_overflowboundary[1000,1000][1000, 1000] gives [0.5,0.5][0.5, 0.5]; ±104\pm 10^4 stays finitelogits after training
test_no_underflow_in_log_softmaxboundary[0,−104][0, -10^4] gives [0,−104][0, -10^4], not −∞-\inftythe cross-entropy of L0.3
test_fully_masked_rowboundaryall −∞-\infty gives zeros and −∞-\infty, no NaN, no warningpadded rows in a batch
test_axis_and_keepdimsunitevery axis of a 3-D array, keepdims shapesattention over [B, H, T, T]
test_dtype_preservedunitfloat32 in, float32 out; integers become float64kernels compare in float32
test_softmax_rows_sum_to_onepropertyrows of scale 0.1 to 700 sum to 1the sampler’s CDF ends at 1
test_kahan_error_boundproperty65536 float32 values within 2u∣S∣2u\lvert S\rvert; plain summation is notL6.7 sums 10710^7 losses
test_pairwise_error_boundpropertywithin ⌈log⁡2n⌉u∑∣x∣\lceil\log_2 n\rceil u \sum \lvert x\rvert for n=65536n = 65536 and 5000150001the cheap accurate sum
test_sums_cover_every_elementboundarylengths 0, 1, odd, a 2-D array, integersno dropped or invented term
test_sums_work_in_the_input_dtypeunitthe float64 version of the hand examplecompensation, not a wider type, does the work
PitfallSymptomCaught by
1. exponentiating before subtracting the maxinf / inf = nan at a float32 logit of 89test_no_overflow (mutant s01)
2. log(softmax(x))−∞-\infty log-probability, infinite loss, NaN gradienttest_no_underflow_in_log_softmax (mutant s02)
3. subtracting a max of −∞-\infty, or dividing by a zero suma padded row of NaN that spreads through the batchtest_fully_masked_row (mutants s03, s04, s14)
4. computing the compensation and never using itKahan silently becomes plain summationtest_kahan_error_bound (mutants s06, s07)
5. a “pairwise” recursion that peels off one element at a timea depth-nn tree, the sequential error, and a recursion limittest_pairwise_error_bound (mutant s09)
6. forgetting to add mm back in log-sum-expoff by exactly the row maximumtest_hand_example (mutant s05)
7. dropping the middle element of an odd split, or indexing an empty arrayoff by a whole termtest_sums_cover_every_element (mutants s08, s10)
8. reducing the wrong axis, ignoring keepdimsattention normalized over the batchtest_axis_and_keepdims (mutants s11, s12)
9. upcasting float32 to float64the C kernel’s float32 reference no longer matchestest_dtype_preserved (mutant s13)

| 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. |

DirectionModuleHow it uses this
BackM09.1the overflow and underflow thresholds of exp, and the unit roundoff uu in every bound
BackM02.1Taylor’s view of exe^x near 0, behind why small shifted exponents are accurate
ForwardM11.1entropy_from_logits uses log_softmax
ForwardM08.3cross_entropy_vjp uses softmax
ForwardL0.2softmax and log-softmax become autograd ops
ForwardL0.3the fused cross-entropy is −logsoftmax-\mathrm{logsoftmax} at the target
ForwardL8.1temperature-scaled logits to sampling probabilities
ForwardL6.7perplexity over 10710^7 tokens accumulates with kahan_sum (through M11.2)
ForwardL9.2the 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.

Your pieceProduction equivalentWhat it addsWhere to look
softmaxPyTorch torch.softmaxfused vectorized kernels, autocast to float32 inside half-precision modelsaten/src/ATen/native/SoftMax.cpp
logsumexpscipy.special.logsumexpweights b (log of a weighted sum), sign handling for negative weightsscipy/special/_logsumexp.py
the shift identityFlashAttention’s online softmaxone pass over keys with a running max and a rescaled running sumDao et al., FlashAttention (2022), section 3.1
pairwise_sumnumpy np.add.reducepairwise with 8-way unrolled leaf blocks, chosen for speed and accuracynumpy/_core/src/umath/loops_utils.h.src
kahan_sumPython math.fsumShewchuk’s algorithm: the exactly rounded sum, at the cost of a list of partialsModules/mathmodule.c (math_fsum)