Skip to content

Softmax in C: three-pass and online two-pass

ModuleL9.2 · side · C · Pass 6 · 2 to 3 h
You buildc/src/kernels/softmax.c: tl_softmax_f32 (three passes over each row: max, exponentiate and sum, normalize) and tl_softmax_online_f32 (two passes: a running max with a rescaled running sum, then normalize)
Contractcourse/contracts/c/include/tinyllm/softmax.h · rules: c/ABI.md
Testscourse/tests/L9.2/: test_softmax.c (C, under ASan and UBSan, both kernels through one table) and shared file fixtures (what they check: section 4)
Needsrt.02 the error slot and the loader · M09.6 your tl_expf · M09.2 your Python softmax, the specification (or --ref-deps). Reading: S-M05 (loop invariants)
Used byL9.3 and L9.4 build their C-side attention oracle from tl_softmax_f32, and apply the online update of this chapter to tiles of keys
MilestoneMS-L9 (the C backend generates the same tokens as numpy)
Optional depthMilakov and Gimelshein, “Online normalizer calculation for softmax” (2018); Dao et al., “FlashAttention” (2022), section 3.1; Higham, Accuracy and Stability of Numerical Algorithms, chapter 4 (summation)
  • Subtract the row maximum first. softmax(x)=softmax(x−c)\mathrm{softmax}(x) = \mathrm{softmax}(x - c) for any constant cc, and with c=max⁡jxjc = \max_j x_j every exponent is at most 0, so rows of ±104\pm 10^4 are exact instead of ∞/∞\infty / \infty (large_logits_stay_finite).
  • The online version needs one read of xx for the statistics. It keeps a running maximum mm and a running sum ss, and rescales ss by emold−mnewe^{m_{\text{old}} - m_{\text{new}}} whenever the maximum grows; the invariant ”ss is the sum against the current mm” holds after every element (online_rescales_when_the_max_grows).
  • Masked entries are −∞-\infty and get exactly 0; a row with every entry masked gives zeros, not 0/00/0, and the online pass never evaluates −∞−(−∞)-\infty - (-\infty) (masked_entries_and_fully_masked_rows).
  • Each row is a distribution: non-negative, summing to 1 within about n⋅2−24n \cdot 2^{-24} for nn columns, and the two versions agree element by element (rows_sum_to_one_and_versions_agree).
  • Your C kernels agree with your Python softmax across row lengths from 1 to 4096 and scales from 0.01 to 10410^4 (test_matches_your_m09_2_softmax).
Terminal window
ol start L9.2 # stubs c/src/kernels/softmax.c into your repo
ol tests L9.2 # read the test catalog first
ol check L9.2 # exit code is the verdict
ol check L9.2 --ref-deps # only if rt.02, M09.6, or M09.2 is not passing yet
ol parity softmax.online # the golden parity suite against a float64 oracle
ol diff L9.2 # after passing: your code against the reference

Your Python sampler and every attention layer you wrote in Parts 5 to 7 call softmax from M09.2, in numpy. Part 9 moves the forward pass into C so that the Rust engine (the standalone Rust engine) can run it, and attention is where most of the softmax work is: one row of scores per query, per head, per layer, per step. A C softmax that overflows on a large score, turns a fully masked row into NaN, or reads its row three times when two would do, shows up later as a garbage token or as a slow decode. This module writes the row softmax in C twice. The three-pass version is the textbook. The online version is the one idea FlashAttention (L9.3) and paged attention (L9.4) are built on: you can normalize a row you are still reading, as long as you remember what maximum you normalized against.

SymbolMeaningType / shape
x∈Rnx \in \mathbb{R}^none row of scores (logits)float[cols]
nnthe row length, colsint64_t
y=softmax(x)y = \mathrm{softmax}(x)yj=exj/∑iexiy_j = e^{x_j} / \sum_i e^{x_i}float[cols]
mmthe row maximum, max⁡jxj\max_j x_j (or the running maximum)scalar
ssthe normalizer ∑jexj−m\sum_j e^{x_j - m} (or the running sum)scalar
mj,sjm_j, s_jrunning maximum and sum after reading x0,…,xjx_0, \dots, x_jscalars
uufloat32 unit roundoff, 2−24≈6×10−82^{-24} \approx 6 \times 10^{-8}

For any constant cc,

exj−c∑iexi−c=e−cexje−c∑iexi=exj∑iexi,\frac{e^{x_j - c}}{\sum_i e^{x_i - c}} = \frac{e^{-c} e^{x_j}}{e^{-c} \sum_i e^{x_i}} = \frac{e^{x_j}}{\sum_i e^{x_i}} ,

so subtracting a constant changes nothing mathematically (M09.2 proved this in Python). Numerically it changes everything: float32 overflows above about 3.4×1038=e88.73.4 \times 10^{38} = e^{88.7}, so e100e^{100} is ∞\infty and a row containing it computes ∞/∞=NaN\infty / \infty = \mathrm{NaN}. With c=m=max⁡jxjc = m = \max_j x_j, every exponent xj−m≤0x_j - m \le 0, every term is in [0,1][0, 1], the largest term is exactly e0=1e^0 = 1, and the sum is at least 1. Nothing overflows, and the division never divides by something tiny.

The direct algorithm reads the row three times:

  1. m=max⁡jxjm = \max_j x_j.
  2. ej=exj−me_j = e^{x_j - m} (stored into yjy_j), s=∑jejs = \sum_j e_j.
  3. yj=ej⋅(1/s)y_j = e_j \cdot (1 / s).

Computing 1/s1/s once and multiplying is cheaper than nn divisions and costs at most one extra rounding per element.

Can the maximum and the sum come from one read? Suppose that after reading x0,…,xjx_0, \dots, x_j we hold

mj=max⁡i≤jxi,sj=∑i≤jexi−mj.m_j = \max_{i \le j} x_i, \qquad s_j = \sum_{i \le j} e^{x_i - m_j} .

Read xj+1x_{j+1}. If it does not exceed mjm_j, the maximum stays and the new term joins the sum: mj+1=mjm_{j+1} = m_j, sj+1=sj+exj+1−mjs_{j+1} = s_j + e^{x_{j+1} - m_j}. If it does, every old term was computed against the wrong maximum. Since exi−mj+1=exi−mj emj−mj+1e^{x_i - m_{j+1}} = e^{x_i - m_j} \, e^{m_j - m_{j+1}}, the whole old sum rescales by one factor:

mj+1=xj+1,sj+1=sj emj−mj+1+e0=sj emj−mj+1+1.m_{j+1} = x_{j+1}, \qquad s_{j+1} = s_j \, e^{m_j - m_{j+1}} + e^{0} = s_j \, e^{m_j - m_{j+1}} + 1 .

Both branches preserve the two equations above, and they hold trivially for the empty prefix with m=−∞m = -\infty, s=0s = 0. By induction (S-M05), after the last element mm is the row maximum and ss the normalizer, so the second pass writes yj=exj−m/sy_j = e^{x_j - m} / s. That is two reads of xx and one write of yy, against three reads and two writes for the three-pass version: for a row that does not fit in cache, a third less memory traffic.

The same update merges two partial results: a prefix summarized by (ma,sa)(m_a, s_a) and a block summarized by (mb,sb)(m_b, s_b) combine to m=max⁡(ma,mb)m = \max(m_a, m_b), s=saema−m+sbemb−ms = s_a e^{m_a - m} + s_b e^{m_b - m}. FlashAttention applies exactly this to tiles of keys, with the weighted sum of values carried along and rescaled by the same factor.

Attention blocks a key by giving its score −∞-\infty, and e−∞=0e^{-\infty} = 0, so a masked entry contributes nothing and gets weight exactly 0. Two cases need code:

  • A fully masked row (a padded query, a window that excludes everything): m=−∞m = -\infty and s=∑0=0s = \sum 0 = 0, so the textbook gives 0/0=NaN0/0 = \mathrm{NaN}. The contract says zeros: such a query attends to nothing.
  • −∞-\infty in the online pass before any finite value: the update would compute exj−m=e−∞−(−∞)=eNaNe^{x_j - m} = e^{-\infty - (-\infty)} = e^{\mathrm{NaN}}. A −∞-\infty entry adds 0 whatever mm is, so the kernel simply skips it.

NaN is different: it means a bug upstream, and the contract makes it visible. No comparison with NaN is true, so NaN never becomes the maximum; it enters the sum, the sum becomes NaN, and the whole row becomes NaN. Rows are independent, so the next row is exact.

Each term is in [0,1][0, 1], the sum is accumulated in float32 left to right, and each of the nn additions rounds with relative error at most uu. The computed ss is within about nun u of the true sum (relative), and the outputs then sum to 1 within about nun u plus one rounding per element. The property test allows 2nu+10−72 n u + 10^{-7}. tl_expf (M09.6) itself is within 4 ulp, which the same bound absorbs. The three-pass and online versions round differently (the online sum is built from rescaled partial sums), so they agree to float32 tolerance, not bit for bit.

x=[1,2,3]x = [1, 2, 3], one row.

Three passes. Pass 1: m=3m = 3. Pass 2: e=[e−2,e−1,e0]=[0.1353353,0.3678794,1]e = [e^{-2}, e^{-1}, e^{0}] = [0.1353353, 0.3678794, 1], s=1.5032147s = 1.5032147. Pass 3: 1/s=0.66524101/s = 0.6652410, so

y=[0.1353353,0.3678794,1]×0.6652410=[0.0900306,0.2447285,0.6652410].y = [0.1353353, 0.3678794, 1] \times 0.6652410 = [0.0900306, 0.2447285, 0.6652410] .

Online. Start m=−∞m = -\infty, s=0s = 0.

jjxjx_jBranchss aftermm after
011>−∞1 > -\infty: s=0⋅e−∞+1s = 0 \cdot e^{-\infty} + 111
122>12 > 1: s=1⋅e−1+1s = 1 \cdot e^{-1} + 11.36787942
233>23 > 2: s=1.3678794⋅e−1+1=0.5032147+1s = 1.3678794 \cdot e^{-1} + 1 = 0.5032147 + 11.50321473

The same mm and ss as the three-pass version, and pass 2 writes the same yy. These are the numbers of the first test, hand_example, which runs both kernels, and of test_hand_example.

tinyllm/softmax.h
tl_status tl_softmax_f32(const float *x, float *y, int64_t rows, int64_t cols); /* 3-pass */
tl_status tl_softmax_online_f32(const float *x, float *y, int64_t rows, int64_t cols); /* 2-pass */
/* Row-major, contiguous; y may equal x (in place). A row of all -inf gives zeros;
NaN propagates to its row. TL_EINVAL for a negative dimension, or a NULL pointer
when rows * cols > 0. */

From Python, declare the two symbols on your loader and pass f32_ptr buffers, as for tl_matmul_f32.

TestKINDChecksWhy it matters downstream
hand_exampleunit, smokesection 3, both kernelsyou and the tests agree on the definition
online_rescales_when_the_max_growsunitan increasing row (rescale every step) and a decreasing one (never)the update FlashAttention applies to tiles
large_logits_stay_finiteboundary[104,104]→[0.5,0.5][10^4, 10^4] \to [0.5, 0.5] and [−104,0]→[0,1][-10^4, 0] \to [0, 1] exactlytrained logits and unscaled scores
masked_entries_and_fully_masked_rowsboundary−∞-\infty gets 0; an all-−∞-\infty row gives zeros; −∞-\infty first in the online passcausal and padding masks
nan_propagates_to_its_row_onlyboundarya NaN row is NaN, the next row exactbugs stay visible and local
in_placeunity == xnormalizing a score buffer in place
bad_arguments_are_einvalboundarynegative dims and NULL buffers give TL_EINVAL; empty work with NULL is fineerrors instead of crashes
rows_sum_to_one_and_versions_agreeproperty, differential200 random rows: non-negative, sum to 1 within 2nu2nu, zeros on masks, the two kernels agreethe definition, for every length up to 501
PitfallSymptomCaught by
multiplying by ss instead of 1/s1/srows sum to s2s^2, not 1hand_example (mutant s01)
a wrong normalizer in the online pass 2rows do not sum to 1hand_example (mutant s02)
not rescaling the running sum when the maximum growsan increasing row looks almost uniformonline_rescales_when_the_max_grows (mutant s03)
rescaling by emnew−molde^{m_{\text{new}} - m_{\text{old}}} (the sign flipped)the running sum explodesonline_rescales_when_the_max_grows (mutant s04)
not subtracting the maximume104=∞e^{10^4} = \infty, and ∞/∞\infty / \infty = NaNlarge_logits_stay_finite (mutants s05, s06)
no special case for a fully masked row0/00/0: NaN poisons the whole batch through the next matmulmasked_entries_and_fully_masked_rows (mutants s07, s09)
updating the online state with a −∞-\infty entry while m=−∞m = -\inftyeNaNe^{\mathrm{NaN}} in the running summasked_entries_and_fully_masked_rows (mutant s08)
zero-filling a row that is all NaNa broken forward pass looks like a masked rownan_propagates_to_its_row_only (mutant s10)
reading x[j] after writing y[j]wrong in place, right otherwisein_place (mutant s11)
skipping the NULL checka crash instead of TL_EINVALbad_arguments_are_einval (mutant s12)
computing x + r * cols when x is NULL and cols is 0UBSan: “applying zero offset to null pointer”; return before touching pointers when there is no workbad_arguments_are_einval
DirectionModuleHow it uses this
Backrt.02the error slot behind TL_EINVAL the Python tests use
BackM09.6tl_expf, your exponential: every exe^{x} in this file
BackM09.2the Python softmax that is the specification
ForwardL9.3FlashAttention: the online update over tiles of keys, with the output rescaled by the same factor; its C tests build the naive attention oracle from tl_softmax_f32
ForwardL9.4paged attention for decode: the online update over KV blocks; its C oracle uses tl_softmax_f32 too

If you skip this module, ol check L9.3 stops with BLOCKED ... needs L9.2: build it, or pass --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
online two-pass softmaxFlashAttention’s tile updatethe same rescaling applied to the output accumulator, so attention never stores the score matrixDao et al. (2022), algorithm 1
scalar loopsPyTorch softmax CPU kernelvectorized max, exp, and sum with SIMD, a vectorized polynomial expaten/src/ATen/native/cpu/SoftMaxKernel.cpp
one row per loopllama.cpp ggml_soft_maxfused scale and mask (ALiBi, causal) in the same passggml/src/ggml-cpu/ops.cpp
float32 sumcuDNN and Triton fused softmaxone block per row, a tree reduction for the max and the sumTriton tutorial “Fused Softmax”