跳到正文

FlashAttention 1–4: How IO-Awareness Reshaped the Attention Kernel

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.

v1 · 2022A100 (Ampere)Bottleneck: the N×N score matrix makes round trips through slow HBM.Fix: tile it, fuse every step into one kernel, keep it on-chip.
v2 · 2023A100 (Ampere)Bottleneck: most of the GPU sits idle — too few thread blocks, too much non-matmul work.Fix: parallelize over the sequence; normalize once at the end.
v3 · 2024H100 (Hopper)Bottleneck: loading tiles and doing math still take turns.Fix: async producer/consumer pipeline; overlap softmax with matmul; FP8.
v4 · 2026B200 (Blackwell)Bottleneck: tensor cores doubled, but the exp unit and shared memory did not.Fix: compute some exps in software, skip most rescales, use tensor memory.

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
#

Intuition. Think of HBM as a big warehouse and on-chip SRAM as a small workbench beside the arithmetic units: the bench is about 10× faster but holds about 2000× less. Standard attention keeps carrying the entire N×N score matrix from the bench to the warehouse and back — once per step. FlashAttention never lets that matrix leave the bench.

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.

On-chip SRAM≈20 MB total (192 KB × 108 SMs) · ~19 TB/s aggregateHBM (device memory)40–80 GB · ~1.5–2.0 TB/snot drawn to scale
Two tiers of memory on an A100. HBM holds about 2000× more than all on-chip SRAM combined (40GB vs ~20MB), but SRAM has roughly 10× the bandwidth. Every FlashAttention version is, at its core, a different strategy for keeping the $N\times N$ intermediate out of the bottom tier.

The FlashAttention paper makes this concrete with a measurement of GPT-2’s attention layer. The surprise is which operations take the time:

≈ 13 ms: three elementwise passes, each reading and writing the whole N×N matrixPyTorch5 separate opsQKTmasksoftmaxdropout×V≈ 16.9 msFlashAttention1 fused kernelfused≈ 2.2 ms (7.6× faster)051015 msmatrix multiplies (almost all of the FLOPs)mask · softmax · dropout (memory-bound)
Where the time goes in standard attention (GPT-2 on an A100; approximate values read off Figure 1 of the FlashAttention paper). The two matrix multiplies hold nearly all of the FLOPs but take only about a quarter of the runtime. Mask, softmax and dropout do almost no arithmetic, yet each one streams the full $N\times N$ matrix out of HBM and back. Fusing everything into one kernel that never writes that matrix is the paper's 7.6× speedup.

FlashAttention (v1): make the N² matrix never exist in HBM
#

Intuition. Don't build the N×N score matrix at all. Work on one small tile at a time on the workbench, keep two running numbers per row — the largest score seen so far and the sum of exponentials — and fix up your partial answer whenever a larger score shows up. At the end you get exactly the same output.

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:

Standard attentionthree kernels, one per stepFlashAttentionone fused kernel, looping over tiles on-chipQKTQ, K → SsoftmaxS → P× VP, V → OHBMQ K VS (N×N)P (N×N)OQ, KSSPP, VOS and P (N×N each) go out to HBM and come backQKT→softmax→× VHBMQ K VOQ, K, VOonly N×d tensors ever touch HBMan N×N matrixan N×d matrix
Left: the standard implementation (Algorithm 0 in the FlashAttention paper) runs one kernel per step, so the $N\times N$ matrices $S$ and $P$ each make a round trip through HBM — and in training, mask and dropout add more round trips. Right: FlashAttention fuses all three steps into a single kernel that loops over tiles on-chip; the only HBM traffic is reading $Q$, $K$, $V$ and writing $O$.

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:

KTd × NQN × dVN × dON × dSij× V block j → O block iinner loop: every K/V block jouter loop over query blocks ionly this tile exists(on-chip, for one step)S = QKTis N × N — it is never stored whole
The same picture as Figure 1 of the FlashAttention paper, drawn in the loop order FlashAttention-2 later adopted. Query block $Q_i$ meets key block $K_j$ to form one tile $S_{ij}$ of the score matrix; that tile is multiplied by $V_j$, added into $O_i$, and thrown away before the next $j$. Only the finished $O_i$ (and its row statistics) is written to HBM. In FA2 each query block $i$ is handled by its own thread block.

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:

Block 1 · scores 2, 1first max: m = 2k1k2ℓ = 1 + 0.37 = 1.37bar height = e^(score − m)→Block 2 · scores 4, 3new max: m = 2 → 4shrunk × 0.14k1k2k3k4ℓ = 0.14 × 1.37 + 1.37 = 1.55old bars and ℓ shrink by e^(2 − 4)→Block 3 · scores 1, 0max unchanged: m = 4k1k2k3k4k5k6ℓ = 1.55 + 0.07 = 1.62nothing to rescaleCheck: softmax over all six scores at once has Σ e^(score − 4) = 1.62 — the same ℓ.this block's scoresearlier blocks, rescaledbefore rescaling
Online softmax on one row of six scores, two per block. Each bar is a score's unnormalized weight $e^{s-m}$ relative to the current running max $m$. When block 2 raises the max from 2 to 4, everything accumulated so far — the earlier weights, the running sum $\ell$, and (not drawn) the partial output $O$ — is multiplied by $e^{2-4}\approx 0.14$, which is exactly what the earlier terms would have been had we known the max was 4 all along. After the last block, $O/\ell$ is the exact softmax-weighted sum.

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 once

Strictly, 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:

standard attentionFlashAttentionGFLOPsarithmetic66.675.2 (+13%)HBM trafficreads + writes40.3 GB4.4 GB · 9.2× lessRuntimeforward + backward41.7 ms7.3 ms · 5.7× faster
Figure 2 (left) of the FlashAttention paper, redrawn: GPT-2 medium attention (sequence length 1024, head dimension 64, 16 heads, batch 64) on an A100, forward plus backward. Each row is scaled separately. FlashAttention does more arithmetic — the backward pass recomputes the score tiles — yet runs 5.7× faster, because its HBM traffic is 9× smaller. Runtime follows memory traffic, not FLOPs.

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
#

Intuition. v1 moves the right amount of data but leaves most of the GPU idle. v2 keeps the math and fixes the scheduling: hand out many more independent pieces of work so every SM stays busy, and spend fewer instructions on bookkeeping that isn't a matrix multiply.

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:

  1. 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$.
  2. 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.
v1: one thread block per (batch, head)batch 1 × 16 heads = 16 thread blocks16 of 108 SMs busy (15%)v2: also one per 128-query block16 heads × 64 query blocks = 1,024 thread blocks108 of 108 SMs busy (≈ 9.5 waves of blocks)
Why parallelizing over the sequence matters, for one illustrative case: a single 8K-token sequence with 16 heads on an A100, whose 108 SMs are drawn as squares. With one thread block per (batch, head), only 16 SMs get work and the rest idle; splitting each head's queries into 128-row blocks gives 1,024 independent thread blocks, enough to fill the GPU many times over.
  1. 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.
v1 — split-KQ — shared, visible to every warpK/V 0K/V 1K/V 2K/V 3warps 0–3: partial sum per K/V slicereduce in shared memwrite partial · barrier sync · addv2 — split-QK, V — shared, visible to every warpQ 0Q 1Q 2Q 3warps 0–3: complete output per Q-sliceO 0O 1O 2O 3independent outputs — no cross-warp reduction
Same four warps, opposite split. v1 splits $K$/$V$ and has to reconcile four partial sums through shared memory; v2 splits $Q$ so each warp's output slice is already complete when its matmuls finish.

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
#

Intuition. On Hopper, copying tiles and multiplying matrices are done by separate hardware engines that can run at the same time. v3 turns the kernel into an assembly line: one worker only fetches, others only compute, and two compute groups take turns so the tensor cores never sit waiting for softmax.

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.

  1. 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.
HBMK, V tilesproducer warpTMA loads onlyshared-memory ring bufferfullfullfillingemptyconsumer warpgroup 1WGMMA + softmaxconsumer warpgroup 2WGMMA + softmaxslot consumed → released back to the producer (hardware barrier)
Warp specialization as an assembly line. The producer warp does nothing but issue TMA copies into a ring of shared-memory slots; the consumer warpgroups do nothing but compute on slots that are already full. As long as the ring stays ahead, the tensor cores never wait on a load, and the producer gives most of its registers to the consumers.
  1. 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.
Warpgroup AWarpgroup BGEMMsoftmaxwaitGEMMsoftmaxwaitsoftmaxwaitGEMMsoftmaxwaitGEMMtime →A's softmax runs while B's GEMMs run — different functional unitsGEMM (tensor cores)softmax (exp/ALU units)
Two warpgroups, phases offset by one step. Whenever one warpgroup is doing softmax's scalar work, the other is busy on its tensor-core matmuls, so the "non-matmul tax" disappears into the matmul's shadow instead of serializing with it. Widths are schematic but proportioned for H100 at head dim 128, where the exponentials take about half as long as the GEMMs they hide behind; on B200 the two become equal (see the FA4 section).
  1. 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.
Before: one outlierAfter random signs + Hadamard rotation±8.0±3.1largest entry 8.0 sets the quantizer's scalelargest entry 3.1 · same length (8.03)
Incoherent processing on a toy 8-entry vector (computed exactly). A quantizer's scale is set by the largest entry in its block, so one outlier of 8.0 forces the seven small entries to share a tiny slice of the representable range. Random sign flips followed by a Hadamard transform spread the outlier's energy across every coordinate: the largest entry drops from 8.0 to 3.1, while the vector's length — and every dot product in $QK^\top$, since $Q$ and $K$ get the same rotation — is unchanged.

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
#

Intuition. Blackwell made matrix multiplies twice as fast but left the exponential unit and shared memory alone, so the softmax that used to hide behind the matmul now takes just as long as it. FA4 goes after that work directly: compute some exponentials on other units, skip rescales that don't change the answer, and keep more data in the new tensor memory.

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.

H100MMA (tensor cores)2048exp (MUFU)1024½ of MMA: hides in its shadowB200MMA (tensor cores)1024exp (MUFU)1024= MMA: nothing left to hideshared-mem reads7680512102415362048cycles per SM for one 128×128 tile, head dim 128 (fewer is faster)
A roofline estimate for one forward-pass tile ($M=N=d=128$), not a measurement. MMA time per tile halves from Hopper to Blackwell, from $4MNd/4096 = 2048$ to $1024$ cycles, but the $MN/16 = 1024$ cycles of exponentials do not: on H100 the exp work fits in half the MMA's shadow, on B200 it takes as long as the MMA itself. B200 bars follow Table 1 of the FA4 paper; the H100 bars apply the same formulas with Hopper's MMA rate (H100 shared-memory time is omitted because Hopper's 64-row MMA tiles change how often operands are re-read). In the backward pass the paper's estimate makes shared memory the binding resource instead: 3328 cycles against 2560 for MMA.

FA4’s answers:

  1. 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.
  2. 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.
running max of the scoresmax FA4 actually usesallowed gap: up to 256×0102030row max (log₂ units)rescalerescaleStandard online softmax rescales at all 9 increases; FA4 only twice.12345678910key/value block
Conditional rescaling on an illustrative row. The running max rises at nine of the ten blocks, and standard online softmax rescales $O$ and $\ell$ every time. FA4 keeps a stale reference max (thick line) and only moves it when the true max gets more than 8 above it in $\log_2$ units (the shaded band) — here twice. Between updates the unnormalized weights can reach $2^8 = 256$, which FP32 accumulators handle easily, and because $O$ and $\ell$ always share the same reference, $O/\ell$ is still exact.
  1. 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$.
  2. 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
#

100%50%0%FlashAttention (v1), A100 forward pass: 30–50% of peak (FlashAttention-2 paper)30–50%FA1A100FlashAttention-2, A100 forward pass: up to 73% of peak (230 of 312 TFLOPs/s)73%FA2A100FlashAttention-3, H100 forward pass, BF16: up to 85% of peak (840 of 989 TFLOPs/s)85%FA3H100FlashAttention-4, B200 forward pass, BF16: up to 71% of peak (1613 of 2250 TFLOPs/s)71%FA4B200
Peak forward-pass utilization reported for each generation's flagship kernel, as a share of its own GPU's dense FP16/BF16 tensor-core peak (A100 312, H100 989, B200 2250 TFLOPs/s). FA1's bar spans the 30–50% range the FlashAttention-2 paper measured for it. Not apples-to-apples across GPUs: FA4's 71% on B200 sits below FA3's 85% on H100 because Blackwell doubled the tensor cores without speeding up the exponential unit or shared memory, so the same fraction of peak is harder to reach.
YearTarget HWCore ideaHeadline result
FlashAttention2022Ampere-era GPUsTiling + online softmax + recomputation → never materialize the N×N matrix in HBMO(N) extra memory, exact; 2–4× faster attention op; 15% (BERT) to 3× (GPT-2) end-to-end
FlashAttention-22023Ampere/Ada/HopperDeferred normalization; swapped loops + parallelism over sequence length; split-Q warps~2× over v1, up to 73% of A100 peak (forward)
FlashAttention-32024Hopper (H100)Warp-specialized async pipeline (TMA + WGMMA); GEMM/softmax overlap; FP8 with block quantization + incoherent processing1.5–2.0× over v2; up to 85% of H100 peak (840 TFLOPs/s); 1.3 PFLOPs/s FP8
FlashAttention-42026Blackwell (B200/GB200); code also runs on HopperAsync MMA + larger tiles; software-emulated exp + conditional rescaling; TMEM + 2-CTA MMA; CuTe-DSL71% of B200 peak (1613 TFLOPs/s); up to 1.3× over cuDNN 9.13, 2.7× over Triton

References
#

GPU training optimization - 这篇文章属于一个选集。
→ FlashAttention 1–4: How IO-Awareness Reshaped the Attention Kernel (本文)

Noviorlu喵

多伦多大学计算机工程。热爱计算机图形学、渲染和 AI/ML。

留言

用 GitHub 账号登录后可以留言。