Prerequisite: GPU 与 Triton 入门 (in Chinese) covers the GPU memory hierarchy and Triton basics; 谁偷走了 5090 的算力和显存 (in Chinese) covers FLOPs, arithmetic intensity, the roofline model and peak memory, and measures on an RTX 5090 why the $N\times N$ score matrix dominates both the runtime and the activation memory of standard attention — the two problems this post starts from.
Self-attention is the one layer every Transformer pays for twice: once in FLOPs, and once — far more painfully — in memory traffic. FlashAttention is a family of exact-attention kernels built to fix the second problem, and each new version targets a different bottleneck that only becomes visible once the previous one is gone. This post walks through all four generations: what each one actually changed, why that change mattered on the hardware of its time, and what stays constant across all of them.
Here is the whole story at a glance; the rest of the post unpacks each column.
I’ll keep the math in inline notation throughout. The primary sources are linked at the end — I’d encourage reading at least the first paper directly rather than taking any secondhand summary (including this one) at face value.
Why attention is memory-bound, not compute-bound#
Standard self-attention computes $S = QK^\top/\sqrt{d}$, $P = \mathrm{softmax}(S)$, $O = PV$, for $Q, K, V \in \mathbb{R}^{N\times d}$. The FLOP count is $\Theta(N^2 d)$, and it stays $\Theta(N^2 d)$ in every version discussed here (the backward pass even adds a recomputed $QK^\top$, about $2N^2d$ FLOPs on top of the standard backward’s $\sim 8N^2d$). FlashAttention does not reduce the asymptotic amount of arithmetic attention requires. What it reduces is HBM traffic, and on a modern GPU that is almost always the thing you’re actually paying for.
The reason is the memory hierarchy. An A100 has 40–80GB of HBM (the “GPU memory” everyone quotes) with roughly 1.5–2.0TB/s of bandwidth, and 192KB of on-chip SRAM per streaming multiprocessor (SM) — about 20MB across its 108 SMs — with an estimated ~19TB/s of aggregate bandwidth: an order of magnitude faster, three orders of magnitude smaller. A naive implementation materializes $S$ and $P$ — both $N\times N$ — in HBM, writes them, reads them back for the softmax, writes again, reads again for the final matmul. For long sequences this round-tripping dominates wall-clock time. The elementwise passes over the $N\times N$ matrix (softmax, plus masking and dropout in training) do only $O(1)$ FLOPs per element moved, so they are purely bandwidth-bound; even the two matmuls are memory-limited at typical head dimensions, since producing an fp16 $N\times N$ matrix from $2N^2d$ FLOPs gives an arithmetic intensity of only about $d$ FLOPs/byte (64–128), below the A100’s roofline ridge point of roughly 150–200 FLOPs/byte. The whole pipeline sits left of the ridge: its speed is set by memory bandwidth, not by peak FLOPs.
The FlashAttention paper makes this concrete with a measurement of GPT-2’s attention layer. The surprise is which operations take the time:
FlashAttention (v1): make the N² matrix never exist in HBM#
Dao, Fu, Ermon, Rudra, and Ré’s original 2022 paper frames this as an IO-awareness problem and solves it with two classic systems ideas applied to softmax: tiling and recomputation. Here is the difference in one picture:
Tiling. Split $Q$, $K$, $V$ into blocks along $N$ small enough that a $Q$-block together with a $K$/$V$-block fits in SRAM. Compute each local score block $S_{ij}=Q_iK_j^\top/\sqrt d$ on-chip and never write the full $S$ or $P$ back to HBM:
The one wrinkle is that softmax is a row-wise global normalization — $P_{ij} = \exp(S_{ij})/\sum_k \exp(S_{ik})$ needs the sum over the entire row, which you don’t have yet when you’re only looking at block $j$. The fix is the “online softmax” recurrence: keep a running row-max $m$ and running row-sum $\ell$, and correct the running output every time the max changes. A tiny example with six scores in three blocks shows the whole trick:
In code, one query block at a time with a single normalization at the end:
for each query block i:
O_i ← 0, ℓ_i ← 0, m_i ← -∞ # kept in SRAM/registers only
for each key/value block j:
S_ij ← Q_i K_jᵗ / sqrt(d) # B×B, SRAM only, never hits HBM
m_new ← max(m_i, rowmax(S_ij))
P_ij ← exp(S_ij − m_new)
ℓ_i ← exp(m_i − m_new)·ℓ_i + rowsum(P_ij)
O_i ← exp(m_i − m_new)·O_i + P_ij·V_j
m_i ← m_new
O_i ← O_i / ℓ_i # write O_i (and the row statistics) to HBM onceStrictly, this is the loop order FlashAttention-2 adopted. The FA1 paper’s Algorithm 1 nests the loops the other way: the outer loop walks $K$/$V$ blocks and the inner loop walks $Q$ blocks, so every inner step reads $O_i$, $\ell_i$, $m_i$ from HBM, updates them — keeping $O_i$ normalized by $\mathrm{diag}(\ell_i)^{-1}$ at every step — and writes them back. The IO bound below holds for both orders; the reordering is one of the things FA2 changed.
Why is this exact rather than approximate? After processing blocks $1..j$, the accumulators satisfy $\ell_i=\sum_{k\le j}\mathrm{rowsum}(e^{S_{ik}-m_i})$ and $O_i=\sum_{k\le j}e^{S_{ik}-m_i}V_k$, with $m_i$ the running max. When the max grows, multiplying both by $e^{m_i-m_{\text{new}}}$ re-bases every earlier term, since $e^{S-m_i}\cdot e^{m_i-m_{\text{new}}}=e^{S-m_{\text{new}}}$. After the last block, $O_i/\ell_i=\mathrm{softmax}(S_{i,:})V$, because softmax is invariant to a per-row shift; the only difference from a one-shot implementation is floating-point summation order. Nothing of size $N\times N$ is ever written to HBM: the extra memory is $O(N)$ for the row statistics, instead of $O(N^2)$.
Recomputation for the backward pass. The backward pass of attention normally needs $P$ again. Rather than store the full $N\times N$ matrix for that (which is exactly the memory blow-up we just avoided), FlashAttention stores only $O$ and the small per-row statistics $(m, \ell)$, and recomputes the needed $S_{ij}$, $P_{ij}$ blocks on the fly inside SRAM during the backward pass. This trades a modest amount of extra matmul FLOPs for a large reduction in memory and memory traffic — a textbook example of recomputation beating storage when the storage is in the slow tier.
The paper’s complexity result: standard attention needs $\Theta(Nd+N^2)$ HBM accesses; FlashAttention needs $\Theta(N^2d^2M^{-1})$, where $M$ is the SRAM size. For typical $d$ (64–128) and $M$ (around 100KB), $d^2$ is many times smaller than $M$, so FlashAttention makes many times fewer HBM accesses; at $M=\Theta(d^2)$ the two bounds would coincide. The paper’s accompanying lower bound is deliberately weak: no exact algorithm can achieve $o(N^2d^2M^{-1})$ accesses for every $M$ in $[d, Nd]$. The paper’s own measurement shows the trade directly — more arithmetic, far less data movement, much less time:
In practice the attention op itself (forward + backward) became 2–4× faster than PyTorch’s standard implementation at sequence lengths 128–4K, which translated into end-to-end training speedups of 15% on BERT-large (seq. 512, vs. the MLPerf 1.1 record), 3× on GPT-2 (seq. 1K, vs. HuggingFace) and 2.4× on Long-Range Arena. Because memory now grows linearly instead of quadratically in $N$, the saving grows with sequence length — about 20× less memory at 4K tokens — which let people train with substantially longer context, while still computing the exact softmax, something the sparse- and low-rank-attention literature at the time could not claim.
FlashAttention-2: the kernel wasn’t even using the GPU well#
v1 solved the IO problem but left throughput on the table: on an A100 its forward pass reached only 30–50% of the GPU’s theoretical peak FLOPs/s (25–35% for the backward pass), well below the 80–90% that a well-tuned dense GEMM achieves. FlashAttention-2 (Dao, 2023) is a work-scheduling paper — same math, better engineering of how work is split across thread blocks and warps.
Three changes matter most:
- Fewer non-matmul FLOPs. Non-matmul work (exponentials, row max/sum, rescaling multiplies) runs on the FP32 ALUs and special-function units rather than the tensor cores. On A100 that is 19.5 TFLOPs/s of FP32 versus 312 TFLOPs/s of FP16 matmul, so each non-matmul FLOP costs roughly 16× more. v1 keeps the output normalized at every step, dividing by the running sum on top of the unavoidable $\exp(m_{\text{old}}-m_{\text{new}})$ correction. FA2 keeps an un-normalized accumulator, applies only the $\exp(m_{\text{old}}-m_{\text{new}})$ correction per block, divides by $\ell$ once at the end (the pseudocode above), and saves a single logsumexp $L=m+\log\ell$ per row for the backward pass instead of both $m$ and $\ell$.
- Swapped loops and parallelism over sequence length. v1 parallelized only over (batch × heads), which can be too few thread blocks to saturate the GPU’s SMs with today’s long sequences and small per-GPU batches. FA2 makes the query-block loop the outer one, so row blocks are independent: in the forward pass each thread block owns one query row block. In the backward pass each thread block owns one key/value column block, accumulating its $dK_j$, $dV_j$ locally and using atomic adds into a shared $dQ$ buffer.
- Better warp partitioning. Within a thread block, v1 split $K$/$V$ across warps while every warp could see all of $Q$ (a “split-K” scheme); each warp then had to write its partial result to shared memory, synchronize, and reduce across warps. FA2 flips this: it splits $Q$ across warps while every warp sees all of $K$/$V$ (“split-Q”), so each warp produces a complete output slice with no cross-warp reduction through shared memory. Warps still meet at block-wide barriers when they cooperatively load each $K$/$V$ tile, but they never exchange partial outputs.
Net effect: roughly 2× over v1, reaching up to 73% of theoretical peak FLOPs/s on A100 for the forward pass (about 230 TFLOPs/s) and up to 63% for the backward pass, which has an inherently less favorable data-reuse pattern. End to end, GPT-style training reached up to 225 TFLOPs/s per A100 (72% model FLOPs utilization). No new hardware feature was required — this is squarely a “use the GPU you already had better” paper, which is also why it’s the version most widely deployed across non-Hopper GPUs (and the AMD ROCm port).
FlashAttention-3: exploiting what Hopper added#
H100 introduced two primitives that neither v1 nor v2 was designed to use: the Tensor Memory Accelerator (TMA), a hardware unit that copies whole tiles between HBM and shared memory asynchronously, and warpgroup-level MMA (WGMMA), a wider, asynchronous tensor-core instruction issued by a group of four warps. FlashAttention-3 (Shah, Bikshandi, Zhang, Thakkar, Ramani, and Dao, 2024) is built around exploiting both, plus block quantization and incoherent processing for FP8.
- Warp specialization (producer/consumer). In FA2, every warp both issues its share of the tile copies (already asynchronous, via Ampere’s cp.async) and computes, with block-wide barriers between stages. FA3 instead assigns a producer warp to issue TMA loads into a multi-stage shared-memory ring tracked by hardware barriers, and gives most of the register file to consumer warpgroups that only run WGMMA on tiles already staged. The kernel becomes an explicit software pipeline.
- Overlapping GEMM and softmax. On a single tile, softmax depends on that block’s $QK^\top$, and the $PV$ GEMM depends on softmax, so they serialize. FA3 breaks the chain in two complementary ways. Inter-warpgroup ping-pong: the two consumer warpgroups each own a different query tile, and barriers force them to alternate, so warpgroup A’s GEMMs run while warpgroup B does its softmax, and vice versa; the paper reports this moving FP16 forward throughput from roughly 570 to 620 TFLOPs/s at head-dim 128, seqlen 8K on H100. Intra-warpgroup pipelining: within one warpgroup, the asynchronous WGMMAs for the next block’s $QK^\top$ and the current block’s $PV$ are issued so that softmax runs while they are in flight.
- FP8 with block quantization and incoherent processing. Quantizing $Q$ and $K$ directly to FP8 before the matmul is attractive for throughput but inaccurate, because a few outlier entries dominate the quantization range and crush the resolution for everything else. FA3 uses two complementary fixes. First, block quantization: one FP8 scale per tile instead of one per tensor, which keeps an outlier’s damage inside its own block and costs nothing extra in a tiled kernel. Second, incoherent processing: multiply $Q$ and $K$ by the same random orthogonal matrix $M$ before quantizing. Since $M$ is orthogonal, $(QM)(KM)^\top = QMM^\top K^\top = QK^\top$ — the attention scores are mathematically unchanged — but each entry of $QM$ is now a random mixture of many original entries, so no single outlier dominates any one coordinate. Implemented as a Hadamard transform with random sign flips, this costs $O(d\log d)$ per length-$d$ row instead of $O(d^2)$ for a dense rotation. On test inputs with simulated outlier features, the FP8 path has 2.6× lower RMSE than a baseline FP8 attention with per-tensor scaling.
Results: 1.5–2.0× over FA2 on H100, up to 840 TFLOPs/s in BF16 (about 85% of H100 SXM5’s 989 TFLOPs/s dense peak), and up to 1.3 PFLOPs/s in FP8. (The July 2024 arXiv v1 reported 740 TFLOPs/s and close to 1.2 PFLOPs/s; the numbers here are from the NeurIPS 2024 version cited below.)
FlashAttention-4: when the bottleneck itself shifts hardware generation#
FlashAttention-4 (Zadouri, Hoehnerbach, Shah, Liu, Thakkar, and Dao; Princeton, Meta, Colfax Research, NVIDIA, Georgia Tech, and Together AI) appeared on arXiv in March 2026. The paper targets Blackwell datacenter GPUs (B200/GB200); the open-source implementation (flash_attn/cute in the flash-attention repo, installed as flash-attn-4) also ships Hopper kernels. It is the newest generation here and the least battle-tested in production, so treat its numbers as the authors’ own measurements.
The paper’s motivating observation is asymmetric hardware scaling. Going from Hopper to Blackwell, dense BF16 tensor-core throughput doubled (8192 vs 4096 FLOPs per clock per SM; about 2.25 vs 1 PFLOPs/s per GPU), while shared-memory read bandwidth (128 B per clock per SM) and the MUFU exponential unit (16 ops per clock per SM) did not change at all. For a $128\times128$ tile at $d=128$, the two forward MMAs take $4\cdot128^3/4096=2048$ cycles on Hopper and the $128^2$ exponentials take $128^2/16=1024$, so FA3’s ping-pong could hide the exponentials behind a matmul twice as long. On B200 the MMAs take 1024 cycles and the exponentials still take 1024: overlap alone can no longer hide them. It is the same IO/compute-balance problem the original paper solved, one level further up the stack.
FA4’s answers:
- A pipeline rebuilt around fully asynchronous MMA and larger tiles. Blackwell’s MMA works on $128\times N$ tiles and writes its accumulator to tensor memory asynchronously, without going through registers. FA4 keeps FA3’s ping-pong, now between two 128-row $Q$ tiles per thread block, with one thread per row so the row max and row sum need no warp shuffles. Because $P$ reaches the $PV$ MMA through tensor memory rather than registers, the output rescale moves to a separate correction warpgroup, off the softmax critical path. A longest-processing-time-first tile scheduler improves load balance for causal and variable-length batches.
- Software-emulated exponential and conditional softmax rescaling. The emulation does not replace the hardware exp; it adds a second source of it. FA4 computes a tuned fraction (10–25%) of each row’s exponentials on the FMA pipes instead of MUFU: split $2^x = 2^{\lfloor x\rfloor}\cdot 2^{x-\lfloor x\rfloor}$ (Cody–Waite range reduction), evaluate $2^f$ for $f\in[0,1)$ with a degree-3 polynomial, and build $2^{\lfloor x\rfloor}$ by shifting the integer into the float’s exponent bits. Each emulated evaluation costs more than one hardware
ex2, but running both in parallel raises total exp throughput. Conditional rescaling attacks the other non-matmul cost, the $O \leftarrow e^{m_{\text{old}}-m_{\text{new}}}O$ correction. That rescale was never needed for correctness, only to keep numbers in range: $O$ and $\ell$ are accumulated against the same reference max, so $O/\ell$ is exact for any reference. FA4 therefore keeps the stale max unless a new block raises it by more than $\log_2 256 = 8$, which bounds the unnormalized entries of $P$ by a factor of 256 — harmless for FP32 accumulators.
- Tensor memory (TMEM) and 2-CTA MMA mode, mainly for the backward pass. Blackwell adds 256KB per SM of tensor memory that the tensor cores write accumulators into directly. In 2-CTA mode, two thread blocks (CTAs) on a pair of SMs in the same cluster execute one MMA with $M=256$; each keeps half of the A tile and the accumulator and stages only half of operand B, halving shared-memory traffic for B. The backward pass is shared-memory-bound on B200 (about 3328 cycles of shared-memory traffic against 2560 of MMA in the paper’s model); keeping more intermediates in TMEM and using 2-CTA mode brings it down to about 2688 cycles and reduces the atomic adds into $dQ$.
- Written in CuTe-DSL rather than raw CUDA C++ templates. CuTe-DSL is CUTLASS’s Python-embedded kernel language; the paper reports 20–30× faster compile times than template-heavy C++ with comparable expressiveness, and attention variants (ALiBi, sliding window, soft-capping) can be written as plain Python score-modification functions that get JIT-compiled into the kernel.
The paper reports up to 1613 TFLOPs/s in the BF16 forward pass on B200, about 71% of the GPU’s 2.25 PFLOPs/s dense peak. That is a lower fraction than FA3’s 85% on H100 — but against a peak that doubled while the exp unit and shared memory did not, holding ~70% took every change above, and it is about 1.9× FA3’s absolute throughput. Against baselines, FA4 is up to 1.3× faster than cuDNN 9.13 and up to 2.7× faster than Triton. The authors also worked with the cuDNN team to fold several of these techniques into cuDNN from 9.13/9.14 on, so the gap to the newest cuDNN (9.19) is much smaller.
What’s actually constant across all four#
It’s worth being explicit about what didn’t change, because that’s the part that’s easy to lose in a list of kernel tricks:
- All four compute the same mathematical function. $\mathrm{softmax}(QK^\top/\sqrt d)V$, up to floating-point rounding (and, for the FP8 path, quantization — which is itself bounded and measured, not hand-waved away).
- None of them reduce the asymptotic FLOP count. Compute stays $\Theta(N^2d)$ in every version; recomputation adds a little, and with a causal mask, skipping fully masked tiles roughly halves the work actually executed compared with computing and then masking the full matrix. The entire lineage is a sequence of answers to “how do we stop paying for data movement and non-matmul overhead that the FLOP count doesn’t actually require.”
- Each version targets whatever the previous version turned into the new bottleneck. v1 removes the $N^2$ HBM round-trip. v2 fixes the low occupancy and warp-communication overhead that v1’s scheduling left on the table. v3 exploits new async hardware (TMA/WGMMA) that v1/v2 predate, and separately attacks the precision/throughput trade-off with FP8. v4 responds to Blackwell’s asymmetric scaling: the exponentials are tiny in FLOP count but now take as many cycles as the matmuls in the forward pass, and shared-memory traffic exceeds the matmuls in the backward pass. This is the general pattern of hardware/software co-design: you don’t get to solve “attention is slow” once — you resolve whichever constraint is currently binding, and the next GPU generation hands you a new one.
Summary#
| Year | Target HW | Core idea | Headline result | |
|---|---|---|---|---|
| FlashAttention | 2022 | Ampere-era GPUs | Tiling + online softmax + recomputation → never materialize the N×N matrix in HBM | O(N) extra memory, exact; 2–4× faster attention op; 15% (BERT) to 3× (GPT-2) end-to-end |
| FlashAttention-2 | 2023 | Ampere/Ada/Hopper | Deferred normalization; swapped loops + parallelism over sequence length; split-Q warps | ~2× over v1, up to 73% of A100 peak (forward) |
| FlashAttention-3 | 2024 | Hopper (H100) | Warp-specialized async pipeline (TMA + WGMMA); GEMM/softmax overlap; FP8 with block quantization + incoherent processing | 1.5–2.0× over v2; up to 85% of H100 peak (840 TFLOPs/s); 1.3 PFLOPs/s FP8 |
| FlashAttention-4 | 2026 | Blackwell (B200/GB200); code also runs on Hopper | Async MMA + larger tiles; software-emulated exp + conditional rescaling; TMEM + 2-CTA MMA; CuTe-DSL | 71% of B200 peak (1613 TFLOPs/s); up to 1.3× over cuDNN 9.13, 2.7× over Triton |
References#
- Dao, Fu, Ermon, Rudra, Ré. “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”, NeurIPS 2022. Figures 1 and 2 of this paper are the basis for the time-breakdown, tiling and FLOPs-vs-traffic diagrams above.
- Dao. “FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning”, ICLR 2024; see also the Stanford Hazy Research write-up.
- Shah, Bikshandi, Zhang, Thakkar, Ramani, Dao. “FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision”, NeurIPS 2024; author’s blog post.
- Zadouri, Hoehnerbach, Shah, Liu, Thakkar, Dao. “FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling”, arXiv:2603.05451, March 2026 (the PDF also ships in the repo’s
assets/folder). - Dao-AILab/flash-attention — reference implementation for all four versions.
- Want the same kind of accounting for a whole Transformer — parameters, FLOPs, memory, training time? See my Transformer FLOPs calculator.