A visual, step-by-step guide to the algorithm that made long-context Transformers practical — exact attention, 2-4× faster, by refusing to waste trips between GPU memory levels.
Everyone was counting arithmetic operations. FlashAttention counted memory trips — and won.
On a GPU, attention is memory-bound, not compute-bound: the time is dominated by moving data between the GPU's slow large memory (HBM) and its fast tiny memory (SRAM), not by arithmetic. Approximate attention attacked the wrong bottleneck. FlashAttention keeps the math exact and removes the memory traffic instead.
Standard attention computes an N×N matrix nobody needs to keep — and pays for it in the slowest memory on the chip.
HBM is the walk-in pantry; SRAM is the countertop. Standard attention walks every ingredient to the countertop, then carries the half-finished dish back to the pantry, then returns to get it again — for every step of the recipe. FlashAttention rewrites the recipe so each countertop session completes a whole course: pantry trips: one in, one out.
Tiling, online softmax, and kernel fusion — each old individually, decisive together.
The physics the algorithm respects: big-and-slow HBM, small-and-fast SRAM.
| Level | Size | Bandwidth | Role |
|---|---|---|---|
| HBM (main GPU memory) | ~40–80 GB | ~1.5–2 TB/s | Stores weights, activations, KV cache |
| SRAM (on-chip) | ~20 MB | ~19 TB/s | Tiling workspace — ~10× the bandwidth |
| CPU RAM | 100s of GB | ~100s GB/s | Off-chip overflow (not in this story) |
An algorithm that keeps its working set in SRAM can be arithmetic-bound and fast; one that spills to HBM is memory-bound regardless of its FLOP count.
Sparse/linear attention reduced arithmetic — but their scattered memory patterns still strided through HBM. Many ran slower than the dense baseline they replaced while also losing quality. The paper's framing — "make attention IO-aware" — reframed the whole subfield: fix the data movement first, approximate only if you still need to.
Watch one output tile assemble itself as K/V blocks stream through SRAM, with the running softmax rescaling everything.
Softmax needs every score before it can normalize anything — normally forcing the full row into memory. The online version keeps a running max m and sum ℓ: when a new block arrives with a bigger max, earlier partial outputs are rescaled down (cheap, in SRAM). At the end: bit-for-bit the same math, zero global passes.
Once attention is tiled, masking whole blocks is one line: skip loading them. Block-sparse FlashAttention combines exact tiling with approximate sparsity — the paper's long-context results (16K–64K) come from this extension, achieving quality above dense baselines at lower cost.
The results table that ended the approximate-attention debate for training workloads.
| Workload | Speedup | Reference |
|---|---|---|
| BERT-large (512 seq) | 15% | vs MLPerf 1.1 training speed record |
| GPT-2 (1K seq) | 3× | end-to-end vs HuggingFace baseline |
| Long-range arena (1K–4K) | 2.4× | vs standard attention |
Exact outputs — speedups come purely from IO behavior, with zero change to model math.
FlashAttention stopped being a paper and became plumbing: the layer every long-context system stands on.
The deepest lesson for system design: before approximating, check whether your bottleneck is real.
FlashAttention's headline numbers are speedups; its real contribution is an epistemic correction. "Attention is quadratic" was a statement about arithmetic — the field acted on it as a statement about time. One careful IO analysis later, exact attention runs 2-4× faster and long context is plumbing. The lesson generalizes far past attention: measure the machine before modeling the math.
Check your understanding of the key concepts from the FlashAttention paper.
Everything you need to remember about this paper.