LLM kernels 3: attention and online softmax
One attention head over 4,096 tokens takes 3 MiB of input and produces 1 MiB of output. The obvious way to compute it writes 64 MiB in between, and touches that buffer five times. Push the context to 131,072 tokens, which several shipping models accept, and the same buffer is 64 GiB for one head of one layer. No GPU has that, and those models run anyway.
So the matrix is not being stored. The kernel that skips it looks impossible at first, because softmax needs the largest element of a row before it can exponentiate anything, and you do not know the largest element until you have seen the row. This page derives the identity that gets around that, builds the kernel on top of it, and measures what the skipped traffic is actually worth in milliseconds, which is a different question with a less flattering answer.
What the matrix costs, and how often you touch it
One head is three matrices of shape N by d, the queries, keys and values.
Here d is 64 and N is the sequence length. Attention is
softmax(Q K^T / sqrt(d)) V: score every query against every key, turn
each row of scores into weights that sum to one, then mix the value vectors
with those weights.
The straightforward version does that in three kernels, and the middle one
forces the first to write its answer down. Q K^T produces an N by N
matrix of scores, S. The softmax reads S to find each row's maximum, reads
it again to write the exponentials, and the third kernel reads those back
to weight V. That is five passes over N^2 floats: 20 N^2 bytes, against
16 N d bytes for the queries, keys, values and output put together. At
N = 4096 those are 320 MiB and 4 MiB. The quadratic term is not a detail of
the implementation, it is the implementation.
Both versions do the same arithmetic, about 4 N^2 d floating-point operations. Divide that by the bytes and the naive path sits at 12.8 FLOP per byte, which day 49 measured to be under this T4's ridge point of 29.34, so it is memory bound and its arithmetic intensity does not depend on N at all. A kernel that never writes S has an intensity of N/4 on the same roofline, which climbs with the sequence length instead of sitting still.
Softmax needs the whole row, and does not
The reason people believe attention has to materialise S is a fact about
softmax, and the fact is true. Exponentiating raw scores overflows: a score
of 100 gives exp(100), which is inf in FP32, so every implementation
subtracts the row maximum first. Subtracting the maximum needs the maximum,
which needs the row.
The way out is not to avoid the maximum, it is to fix it up. Suppose you
have processed part of a row and you hold two numbers: m, the largest
score seen so far, and d, the sum of exp(x - m) over those same scores.
A new tile arrives whose largest score is bigger. Every term you have
already accumulated was scaled by the old maximum, and converting it to the
new one is one multiply, because the exponent is a sum:
// exp(x - mNew) == exp(x - mOld) * exp(mOld - mNew)
mNew = max(mOld, tileMax);
scale = exp(mOld - mNew); // one factor for the whole history
dNew = dOld * scale + sum over the tile of exp(x - mNew);
oNew = oOld * scale + (tile weights) . V_tile;
One factor rescales the running denominator and the running output
together, so the state a query row carries is m, d and its d-wide
output, and never a row of scores. Milakov and Gimelshein published this as
the online normalizer calculation
(https://arxiv.org/abs/1805.02867 , checked 2026-09-01); FlashAttention
applies it tile by tile with the value mix carried along
(https://arxiv.org/abs/2205.14135 , checked 2026-09-01). The widget above
steps a sixteen-element row in four tiles with the largest value arriving
last, which is the case intuition gets wrong.
Two paths, one head, one reference
Full program in
attention.cu, under
code/day97-attention/. It runs both paths at N of 256, 1024 and 4096 and
prints the bytes before it prints any time.
The baseline is not a straw man. Its two matmuls use day 16's 16 by 16 tile, with the key tile held transposed in shared memory so the global reads stay coalesced. The only thing the naive path is charged for on this page is the matrix it writes to global memory.
The reference is double precision on the host, holding one row of scores rather than the matrix. The two GPU paths sum in different orders, so an FP32 reference could not arbitrate between them. The gate is day 66's K-scaled rule with the depth set to N, since every output element is a sum over N keys, and the program prints the arithmetic beside it.
One key is deliberately loud. Key N-3 sits in the final tile of every sweep and scores far above the rest, so every query row's running maximum is still moving on the last tile. Without it the rescale could fire once and never again, leaving this day's branch untested.
In the online kernel each block owns 64 query rows, each warp owns eight of them, and each lane owns one key of the current 32-key tile and two of the 64 output dimensions. The state that replaces the score matrix is 32 floats per lane:
float acc[kRowsPerWarp][kDimsPerLane];
float runMax[kRowsPerWarp];
float runSum[kRowsPerWarp];
for (int r = 0; r < kRowsPerWarp; ++r) {
runMax[r] = -INFINITY;
runSum[r] = 0.0f;
for (int c = 0; c < kDimsPerLane; ++c) {
acc[r][c] = 0.0f;
}
}
And here is one tile's worth of the identity, for one query row. The two butterfly reductions are day 23's shuffles, which leave the tile maximum and the tile sum in every lane without a barrier or a scratch buffer:
float tileMax = score;
for (int off = kWarpSize / 2; off > 0; off >>= 1) {
tileMax =
fmaxf(tileMax, __shfl_xor_sync(0xffffffffu, tileMax, off));
}
const float newMax = fmaxf(runMax[r], tileMax);
// exp(-inf - newMax) is 0, so the first tile's rescale wipes the
// initialised state instead of special-casing it.
const float rescale = expf(runMax[r] - newMax);
const float weight = expf(score - newMax);
float tileSum = weight;
for (int off = kWarpSize / 2; off > 0; off >>= 1) {
tileSum += __shfl_xor_sync(0xffffffffu, tileSum, off);
}
runSum[r] = runSum[r] * rescale + tileSum;
runMax[r] = newMax;
for (int c = 0; c < kDimsPerLane; ++c) {
acc[r][c] *= rescale;
}
for (int j = 0; j < kKeyTile; ++j) {
const float w = __shfl_sync(0xffffffffu, weight, j);
for (int c = 0; c < kDimsPerLane; ++c) {
acc[r][c] += w * tileV[j][c * kWarpSize + lane];
}
}
The caveat is the query tile. Sixty-four rows per block means keys and values are re-read once per block, so the kernel's request count carries an N^2 term even though its allocation does not. That ask is printed next to the floor, and whether it reaches DRAM or stops at L2 is a cache question this program cannot answer. A production kernel raises the query tile and hands both matmuls to the tensor cores, which day 72 introduces and this page does not use.
Results
Re-verified on the same Tesla T4 with driver 580.173.02 and CUDA 13.0 (V13.0.88). Errors, resource rows and traffic arithmetic reproduced exactly; all cases passed and the process exited 0. At 4096, online improved from 6.7860 to 5.2910 ms while naive moved from 7.4483 to 8.3724 ms, widening the online win. Both transcripts are listed in front matter; this table is CUDA 13.0.
| n | path | max abs err | gate frac | time (ms) | naive/online |
|---|---|---|---|---|---|
| 256 | naive, 3 kernels | 5.941e-09 | 0.0614 | 0.0745 | 0.211 |
| 256 | online, 1 kernel | 5.678e-09 | 0.0614 | 0.3540 | 0.211 |
| 1024 | naive, 3 kernels | 3.023e-09 | 0.0553 | 0.8413 | 0.576 |
| 1024 | online, 1 kernel | 3.372e-09 | 0.0553 | 1.4604 | 0.576 |
| 4096 | naive, 3 kernels | 1.453e-09 | 0.0578 | 8.3724 | 1.582 |
| 4096 | online, 1 kernel | 1.310e-09 | 0.0550 | 5.2910 | 1.582 |
At n=4096 the arithmetic traffic ratio was 81.0x. The five predictions all held:
- Held. Every gate passed, while the online-versus-naive worst gate fractions remained nonzero: 0.0669, 0.0596 and 0.0614.
- Held. Online maximum error stayed within 1.12x of naive at every size and was smaller at 256 and 4096.
- Held.
attentionOnlineused 32,896 shared bytes, 61 registers, one block per SM and 25 percent occupancy. Each naive kernel reported four blocks and 100 percent. - Held. At 256, online took 0.3540 ms against naive's 0.0745 ms; naive/online was 0.211.
- Held. The ratio rose strictly, 0.211 to 0.576 to 1.582, crossing 1 only at 4096.
The byte win existed at every size; the time win appeared only at 4096. That is the measured distinction between deleting storage and accelerating the finite test case.
Run it yourself
Anything at sm_75 or above builds and runs this, including a free Colab session:
nvcc -std=c++17 -O3 -arch=sm_75 -o attention attention.cu
No profiler, no root: the occupancy table comes from cudaFuncGetAttributes
and cudaOccupancyMaxActiveBlocksPerMultiprocessor, neither of which needs
performance counters. No Compiler Explorer embed, because the
double-precision reference at 4,096 tokens is billions of multiply-adds on
one core and blows the 20 second execution cap by itself. Device memory
peaks near 140 MB, the naive path's two square buffers.
Exercise
Write the online kernel yourself. Keep the head width at 64 and the key tile at 32, and produce one kernel that takes Q, K and V and writes O with no N by N allocation anywhere.
Time: 60 to 90 minutes. Submit: your kernel, its worst gate fraction against the shipped naive path at n = 1024, and the size at which it starts beating that path in wall time on your card.
Check: your kernel must pass the two gates the shipped program prints,
against the double-precision host reference and against the naive path's
output, both with rtol = max(1e-5, 4*2^-23*sqrt(n)). A failure names the
query row and head dimension of the first offending element. Two mistakes
cause almost all of them: forgetting to rescale the accumulated output when
the maximum moves, which is right only when the largest score lands in the
first tile, and normalising per tile instead of dividing once at the end,
which gives rows that do not sum to one.
Hint 1
You cannot store a row of scores, so ask what a query row carries between tiles instead. Three things, one of them as wide as the head. What has to happen to all three when a later tile holds a bigger score than anything before it?
Hint 2
Write exp(x - mNew) in terms of exp(x - mOld). The exponent is a
difference, so the conversion is a single multiplicative factor that does
not depend on x. Which of your carried values does that factor apply to,
and which one does not need it because it is being replaced?
Solution
attentionOnline in attention.cu is the worked version. Per tile: reduce
the tile's scores to a maximum, take newMax = max(runMax, tileMax),
compute rescale = exp(runMax - newMax), multiply the running denominator
and every element of the running output by it, then add this tile's
contribution against newMax. Divide once at the end. Initialising the
maximum to negative infinity makes the first rescale exactly zero, so the
opening tile needs no special case.
Softmax's combining rule is associative: two partial results, each holding a maximum and a sum taken against it, merge by rescaling the smaller-maximum side. That is why attention tiles at all, and it is the property day 24 leaned on to turn a sum into a tree.
Pitfalls
Your kernel is right in testing and wrong in production. Test data whose largest score lands early never exercises the rescale, so a kernel that forgets it passes. Put the outlier in the last tile, the way this program does, and check that the failure appears.
You subtracted the maximum and still got inf. Then it was not the
maximum of everything you exponentiated: the usual cause is updating the
running sum before the running maximum, or weighting a tile against its own
maximum instead of the running one.
Your denominator drifts at long sequence lengths. The rescale factor is below one and applies to the whole history, so the running sum is a product of many such factors. In FP32 with a large spread of scores it can underflow to zero, and a denominator of zero produces NaN rather than a wrong number. Gate on finiteness, which this program does, rather than on error alone. Day 47 is where the exponential's own accuracy gets priced.
You benchmarked against a naive path nobody would ship. An untiled
Q K^T makes any comparison look good and measures the matmul rather than
the materialisation. This page gives the baseline day 16's tile for that
reason, and it still loses on bytes by the factor the arithmetic predicts.
You fused attention and expected the launch count to pay. Three kernels became one, and three launches are microseconds (day 48 priced an empty one at 2.668 us). What fusion bought here is the buffer, not the launches. Count the bytes before you attribute the result.
You padded every shared tile out of habit. The key tile here is padded by one float because lane j walks down key j, which without the pad puts all 32 lanes on one bank. The value tile is not, because lanes there read consecutive dimensions. Day 15 has the arithmetic that tells the two cases apart, and a bank conflict is worth ruling out before an occupancy change.
Go deeper
- Milakov and Gimelshein, "Online normalizer calculation for softmax", the one-pass algorithm this page derives: https://arxiv.org/abs/1805.02867 (checked 2026-09-01)
- Dao, Fu, Ermon, Rudra and Re, "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness", which applies it tile-wise and counts the HBM accesses: https://arxiv.org/abs/2205.14135 (checked 2026-09-01)
- CUDA Programming Guide, "Warp Shuffle Functions", for the two butterfly reductions in the kernel above: https://docs.nvidia.com/cuda/cuda-programming-guide/05-appendices/cpp-language-extensions.html (checked 2026-09-01)
cuda-samples,cpp/3_CUDA_Features/shfl_scan, warp-level reduction and scan without shared memory: https://github.com/NVIDIA/cuda-samples/tree/master/cpp/3_CUDA_Features/shfl_scan (checked 2026-09-01)- Programming Massively Parallel Processors, 4th edition, chapter 6, on thread coarsening and the memory cost of intermediate results: https://shop.elsevier.com/books/programming-massively-parallel-processors/hwu/978-0-323-91231-0
Next
Day 98 drops the precision this page kept in FP32, with INT8 on the T4's tensor cores and FP8 and 2:4 sparsity as stretch tiers. A tiny transformer forward pass is this day's stretch: combine the attention kernel with day 95's norms and day 96's elementwise chain, then check every stage before timing the whole block. Day 99 changes direction and teaches the backward shapes that day 100 needs. All three sit in module 10; time them with CUDA events and read occupancy from the runtime rather than guessing it.