What does FlashAttention actually change?
Everything in the first four parts converges on one kernel. The roofline said naive attention is memory-bound at any sequence length. The hierarchy said a score matrix cannot fit in shared memory. The memory budget said the attention matrix is the quadratic term that made long context impossible. FlashAttention is the algorithm that answers all three at once, and it is the single best example I know of a program written for the memory hierarchy rather than for the arithmetic.
It is also the most over-explained and least understood kernel in the field. Most descriptions say “it fuses the attention operations” and stop. The interesting part is the two problems fusion creates, and how each version since has moved the work to whatever unit the hardware had left idle.
The example throughout is one head with $d = 128$, in BF16, on an H100 unless the version needs a B200.
The 30-second version
Standard attention computes the N by N score matrix, writes it to HBM, reads it back for the softmax, writes the probabilities, and reads them again to multiply by V. That is 8N² bytes of traffic for 4N²d FLOPs, an intensity of d/2, about 64, which is memory-bound on every GPU. FlashAttention gives each SM a block of queries and streams blocks of keys and values past it, computing each tile of scores in shared memory and discarding it, so HBM only sees Q, K, V and O once. Two problems follow. The softmax needs the whole row, so it keeps a running maximum, running sum and running output per row and rescales them when a new block raises the maximum: the online softmax. And the backward pass needs the probabilities, so instead of storing N² of them it stores one log-sum-exp per row and recomputes each tile from Q and K. The FLOPs go up by about a quarter in the backward pass; the bytes go down by a factor of N over a hundred, memory goes from quadratic to linear, and the kernel moves from the roofline's slope to its roof. Later versions kept the algorithm and chased the idle unit: FA2 parallelised over the sequence, FA3 overlapped softmax with matmul on Hopper, FA4 moved exponentials off the saturated special-function units on Blackwell.Why is the textbook algorithm memory-bound?
Attention for one head is two matmuls with a softmax between them:
\[S = QK^\top, \qquad P = \text{softmax}(S), \qquad O = PV\]Written as three kernels, $S$ is an $N \times N$ matrix that goes to HBM and comes back. Write $S$, read $S$ for the softmax, write $P$, read $P$ for the second matmul: four crossings of $N^2$ values, $8N^2$ bytes in BF16, plus $8Nd$ for $Q$, $K$, $V$ and $O$.
\[I_{\text{naive}} = \frac{4N^2 d}{8N^2 + 8Nd} \approx \frac{d}{2} = 64\]Independent of $N$, and a quarter of the way to the H100’s ridge of 295. At 8k tokens the score matrix is 134 MB per head, and on an H100 the round trips take about three times longer than the arithmetic. With 64 heads that is 8.6 GB of scores per layer, written and read twice, per forward pass.
In training it is worse, because $P$ is kept for the backward pass: five bytes per score across all heads, which is the $5as^2b$ term from Part 4. The arithmetic was never the problem. The algorithm was wrong for the memory hierarchy.
How does tiling remove the traffic?
Give each SM a block of $B_r$ queries and keep it in shared memory for the whole computation. Stream the keys and values past it in blocks of $B_c$. For each pair of blocks the SM computes a $B_r \times B_c$ tile of scores, uses it, and throws it away. The tile never exists in HBM. The output block for those queries accumulates in registers until every key block has passed, then is written once.
With $B_r = B_c = 128$ the working set is the query block, 32 KB, a key block and a value block, 32 KB each, and a 64 KB score tile in FP32. Under 227 KB, and it fits.
HBM now sees $Q$, $K$, $V$ and $O$ about once each: $8Nd$ bytes, intensity $N/2$, which at 8k tokens is 4096 and far right of the ridge. The keys and values are re-read once per query block, and in the FlashAttention paper’s accounting the HBM traffic is $O(N^2 d^2 / M)$ for shared-memory size $M$. In practice a head’s $K$ and $V$ are a few megabytes, so those re-reads are served by the L2 and HBM sees them roughly once. Same FLOPs, a hundred times fewer bytes.
Two problems remain, and they are the actual content of the paper.
How can you softmax a row you never see whole?
The softmax of row $i$ divides each $e^{s_{ij}}$ by the row’s sum, and for numerical safety subtracts the row’s maximum first. Both are over all $N$ scores, and the SM only ever holds $B_c$ of them.
The fix is the online softmax, from Milakov and Gimelshein in 2018. Carry three running quantities per row: the maximum so far $m$, the sum so far $\ell = \sum e^{s - m}$, and the unnormalised output so far $o = \sum e^{s - m} v$. When a block arrives with a new maximum $m’$:
\[\ell \leftarrow \ell \cdot e^{m - m'} + \sum_{\text{block}} e^{s - m'}, \qquad o \leftarrow o \cdot e^{m - m'} + \sum_{\text{block}} e^{s - m'} v\]The correction factor $e^{m - m’}$ rescales everything accumulated under the old maximum to the new one. At the end, $O = o / \ell$, and it equals the exact softmax to machine precision. The figure runs it on eight scores and checks.
Per row the state is three numbers. Per query block it is three small vectors that live in registers. That is the whole trick, and everything in FlashAttention that is not a matmul is this bookkeeping.
What does the backward pass do without P?
The backward pass needs the probabilities to form $dV = P^\top dO$ and $dS = P \odot (dP - \delta)$. The textbook stores them, $N^2$ per head, and that is the quadratic term that fills memory.
FlashAttention stores one number per row instead, the log-sum-exp:
\[L_i = m_i + \log \ell_i\]and recomputes each tile of $P$ when the backward pass needs it, $P_{ij} = e^{S_{ij} - L_i}$, from $Q$, $K$ and $L$, a tile at a time, on chip. That is one extra matmul on top of the five the backward pass already does, about a quarter more FLOPs, in exchange for the whole matrix.
Per head at 8k tokens the textbook keeps 336 MB; FlashAttention keeps $2Nd + 4N$ bytes, about 2 MB. Across 64 heads and 80 layers the difference is 1.7 TB against 11 GB, for one sequence. Recomputation is cheaper than remembering when what you would remember is quadratic and the tensor cores are idle anyway. This is Part 4’s activation formula losing its second term.
What did FlashAttention-2 fix?
The first version, in 2022, established tiling and recomputation and reached about half of the A100’s peak. The gap was parallelism, not arithmetic.
FlashAttention-2, in 2023, swapped the loop order. Each thread block owns a query block for the entire sequence, so the output accumulates in registers rather than round-tripping partial results through HBM. It split work between warps along the query axis, so warps stop synchronising through shared memory to combine their pieces. And it moved the non-matmul work out of the inner loop: the division by $\ell$ happens once at the end, and the rescaling is done as rarely as possible. The result was roughly double the throughput, 50 to 73% of the A100’s peak in the forward pass.
On Hopper the same code reached only 35%. The chip had changed underneath it.
What did FlashAttention-3 do with Hopper?
Hopper introduced two things the kernel had to be rewritten around: the tensor memory accelerator, which copies whole tiles from HBM to shared memory asynchronously, and warpgroup matmul instructions that run asynchronously too. A kernel written as a sequence of load, multiply, softmax, multiply leaves both idle for most of each step.
FlashAttention-3, in 2024, specialised the warps. Producer warps do nothing but issue TMA copies of the next key and value blocks. Consumer warpgroups issue the matmuls from shared memory and run the softmax in registers. And two consumer warpgroups ping-pong: while one is in its softmax on the vector units, the other is in its matmul on the tensor cores, so both units stay busy instead of alternating.
It also brought FP8, with a trick worth knowing. Quantising keys and values per block after a random Hadamard rotation spreads outliers across the block, and cut the error to 2.6 times below plain FP8. The numbers: 740 TFLOP/s in BF16, 75% of the H100’s peak, and about 1.2 PFLOP/s in FP8.
Why did Blackwell need FlashAttention-4?
Part 3 ended with the ratio: a Blackwell SM issues about 8192 tensor-core FLOPs per clock and 16 exponentials per clock, and the exponential units did not change from Hopper. Attention needs one exponential per score and $4d = 512$ matmul FLOPs per score. That is 16 scores per clock from the tensor cores and 16 from the MUFU, neck and neck on paper, and the vector pipeline also has to compute the max, apply the scaling, and convert precisions. The softmax lost.
FlashAttention-4, published in March 2026, rebalances in four moves.
Part of the exponentials run as software. Write $2^x = 2^n \cdot 2^f$ with $f \in [0, 1)$, and evaluate $2^f$ as a cubic polynomial, $1 + 0.6951 f + 0.2276 f^2 + 0.0771 f^3$, on the ordinary FMA units, which have slack. The MUFU handles the rest.
The rescaling becomes conditional. When a block’s new maximum is within a threshold of the old one, the correction of $\ell$ and $o$ is skipped and applied once at the end; most blocks do not move the maximum much, so most of the rescaling vanishes.
Accumulators move into the new tensor memory, 256 KB per SM wired to the tensor cores, freeing registers and shared memory. And two SMs execute one 256 by 256 matrix instruction together, each loading half of the key and value blocks, which halves shared-memory traffic per FLOP, the bottleneck of the backward pass.
The numbers: 1605 TFLOP/s in BF16 on a B200, 71% of peak, 1.1 to 1.3 times cuDNN 9.13 and 2.1 to 2.7 times Triton. PyTorch’s FlexAttention now uses it as a backend on Hopper and Blackwell, and an NVFP4 variant reaches about 2 PFLOP/s.
What is the pattern across four versions?
Each version found the unit that was idle and moved work onto it.
The first moved the score matrix off HBM and onto shared memory. The second moved the output out of HBM and the rescaling out of the inner loop, so the SMs stopped waiting on each other. The third moved the softmax underneath the matmul, so the tensor cores stopped waiting on the vector units. The fourth moved exponentials off the special-function units and accumulators off the registers, because Blackwell had made those the scarce resource.
The arithmetic has been the same since 2017. The algorithm has been rewritten four times to match what the hardware ran out of. That is the lesson of the whole series in one kernel, and it is why the next part, on the 2026 toolkit, reads as a list of things moved from one unit to another.
The 30-second version
Softmax needs the row max for stability and the row sum for normalisation, and both need the whole row. Keep m, the max so far, and accumulate the sum and the weighted output relative to it: ℓ = Σ e^(s−m), o = Σ e^(s−m) v. When a new block raises the max to m′, every term accumulated so far was computed relative to the old m and is too large by a factor of e^(m′−m), so multiply ℓ and o by e^(m−m′) before adding the new block's terms, which are computed relative to m′. The factor is 1 whenever the block does not raise the max, which is most blocks once a few have passed; FlashAttention-4 exploits that by skipping the rescale when the max barely moves. At the end O = o/ℓ, exactly the softmax-weighted sum.Rapid fire: can you do these from memory?
- Count the HBM bytes of textbook attention and derive its intensity of d/2.
- What is the working set of a FlashAttention block in shared memory at B = 128, d = 128?
- Write the online softmax update and explain the correction factor.
- What does FlashAttention store for the backward pass instead of P, and how does it recover P?
- Why is the FlashAttention backward pass about 2.5 times the forward FLOPs rather than 2?
- Name the three changes in FlashAttention-2 and what each fixed.
- What is ping-pong scheduling, and which two units does it keep busy?
- Why did the softmax become the bottleneck on Blackwell, and name two of FlashAttention-4's responses.
Part 6 is the catalogue: every technique in production in September 2026, placed against the bottleneck it attacks, from FP4 training to disaggregated serving, and the hardware that shipped this year.