Softmax and FlashAttention
Online softmax keeps a running maximum and rescales; FlashAttention uses it to never write the score matrix to HBM.
Loading the animation…
Concept
Attention computes (an matrix of scores for a sequence of tokens), takes a softmax along each row to get , and returns . Done as three separate kernels, standard attention writes to HBM, reads it back for the softmax, writes , and reads it again for the product: four passes over numbers. At and head dimension 64 that is 136 MB for one head, almost all of it the score matrix.
FlashAttention (Dao et al., 2022) never stores . It takes a block of queries, streams the keys and values past in blocks of , and for each pair computes the score tile on chip, folds it into a running output, and discards it. The animation is FlashAttention-2's loop order (Dao, 2023): one row of tiles per query block. What goes to and from HBM is and once, and and once per query block, between 2.1 MB (if every re-read of and hits L2, the "K and V from L2" bar) and 34.6 MB (if none does) at the same size.
The obstacle is the softmax: each output row needs over the whole row before any of it can be normalised, and FlashAttention sees the row a tile at a time. The trick that makes it possible is online softmax.
Loading the animation…
Concept
A safe softmax subtracts the row's maximum before exponentiating, so nothing overflows: . That seems to need two passes (find , then sum), and a third to normalise. Milakov and Gimelshein (2018) showed that one pass is enough: keep a running maximum and a running sum of . When a new block brings a larger maximum , every term already in was computed against the old and is too large by , the same factor for all of them, so multiplying by fixes them all at once. Step through the animation and watch the bars below the scores: they all shrink together when the maximum rises.
FlashAttention keeps exactly these two numbers per query row, and rescales its partial output row by the same factor, which is why it can stream and through without ever seeing a whole row of scores.
Maths
Online softmax is exact. Let and be the maximum and the sum after blocks. Suppose . After the next block, with ,
so by induction ends as with the true maximum, and is the ordinary softmax, in exact arithmetic. In floating point the two differ only by the rounding of the different order of additions: the animation's last frame shows the largest difference, of order .
HBM traffic. With bytes per value, standard attention moves : read and , write and read , write and read , read , write . FlashAttention-2 reads once (), writes once () and reads and () once for each of the query blocks:
The term survives, but divided by and halved, and with L2 reuse it disappears: the floor is , every value read or written once. This is the paper's result with for on-chip memory .
Code
The online softmax step, cut from src/lib/gpu/model.ts:
const mNew = m > bmax ? m : bmax;
const scale = Math.exp(m - mNew);
let s = 0.0;
for (const v of blk) s += Math.exp(v - mNew);
const lPrev = l;
l = l * scale + s;
And the bytes FlashAttention-2 moves per tile (no L2 reuse in moved, every re-read an L2 hit in once):
if (j === 0) {
moved += eb * br * d;
once += eb * br * d;
}
moved += 2 * eb * bc * d;
if (i === 0) once += 2 * eb * bc * d;
One head, forward pass (model, BF16, blocks of 128)
| Sequence × head dim | Standard attention | FlashAttention, no L2 reuse | Every value once (floor) |
|---|---|---|---|
| 4096 × 64 | 136 MB | 34.6 MB | 2.1 MB |
| 4096 × 128 | 138 MB | 69.2 MB | 4.19 MB |
| 8192 × 128 | 545 MB | 273 MB | 8.39 MB |