Era 6 · Generative models and systems · 2022

48 FlashAttention

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness · Dao, Fu, Ermon, Rudra & Ré · Stanford · NeurIPS
🟧 read selectively~2 horiginal ↗
The gist in 20 seconds. EXACT attention that minimizes reads and writes between the GPU's slow memory (HBM) and its fast memory (SRAM) by tiling. The bottleneck in attention is memory bandwidth, not FLOPs. Same result, several times faster, and memory O(N) instead of O(N²).

Context

Attention in the Transformer (#32) costs O(N²) in memory and time for a length of N — and the thing it actually runs into is not what people assumed.

The idea and the mechanism

The insight: the standard implementation materialises the full N×N score matrix in the GPU's slow memory (HBM), while the bottleneck on modern GPUs is not FLOPs but memory bandwidth (HBM↔SRAM traffic). The fix: IO-aware exact attention through tiling — compute in blocks, keep the intermediates in fast SRAM, never materialise the full matrix.

algorithms · systems Tiling and online softmax: why the memory is O(N)

GPU memory is a hierarchy: HBM is large but slow; SRAM is tiny but very fast. Naive attention writes an N×N matrix to HBM — that is O(N²) of traffic, and it is the traffic that stalls you, not the multiplications.

Tiling. Split Q, K, V into blocks, load them into SRAM and compute attention piece by piece. The problem: softmax needs to normalize over the whole row, and we only ever see a block. That is solved by online softmax — keep a running maximum m and a running sum ℓ, and rescale the accumulated result with every new block (numerically stable):

m ← max(m, mblock),   ℓ, O are rescaled with the new m

The full N×N matrix never sits in HBM — memory is O(N). The result is bit-for-bit the same (exact), but memory traffic drops sharply → a several-fold speed-up. The lesson: on modern hardware, optimize data movement, not arithmetic.

The backward pass avoids N×N too (the other half of the idea). During training, ordinary attention stores the N×N matrix from the forward pass in order to compute gradients. FlashAttention does NOT store it: in the backward pass it recomputes the score blocks it needs from Q, K, V and the saved normalizers (m, ℓ). This is rematerialization — the classic "extra FLOPs for memory" trade: a little more arithmetic in exchange for O(N) memory instead of O(N²). That is exactly what makes not only inference but also TRAINING on long contexts efficient.

PyTorch FlashAttention under the hood of sdpa
import torch.nn.functional as F
# PyTorch picks the FlashAttention kernel automatically:
out = F.scaled_dot_product_attention(Q, K, V)
# the result is identical to softmax(QKᵀ/√d)·V, but memory is O(N), not O(N²),
# because the full score matrix is never materialised in HBM
HBM (slow, large) naive: the N×N matrix SRAMfast blockspiece by piece,no N×N in HBM
Instead of writing the full N×N matrix to slow HBM — compute in blocks in fast SRAM. Same result, but memory traffic and the ceiling are O(N).
Analogy. Cooking a large order while running to a distant storeroom (HBM) for every ingredient is slow because of the running, not the cooking. FlashAttention keeps the ingredients it needs on the counter within reach (SRAM) and works through the order in small batches, without spreading the whole storeroom over the kitchen. The bottleneck was never the speed of your hands — it was the walk to the storeroom.

Why it matters

It made long contexts practical and sped up training and inference for every Transformer; it is the de facto standard — built into PyTorch, the default in vLLM, available in HF. The broader lesson: on modern hardware, optimize DATA MOVEMENT through the memory hierarchy, not just the arithmetic.

Connections

← speeds up32. Transformer

FlashAttention does not change the mathematics of attention (#32) — it computes the same thing faster and with less memory. The quadratic cost that limited how long a Transformer's context could be was largely lifted right here.

vLLM builds efficient serving on top of FlashAttention kernels. Two levels of inference optimization: FlashAttention inside a single attention operation, PagedAttention in the memory management across requests.

↔ contrast11. LSTM

A curious turn of the story: recurrence (LSTM) was dropped for the sake of parallelism, and what we got in return was the quadratic cost of attention. FlashAttention brings back the efficiency without bringing back recurrence — by optimizing memory rather than changing the architecture.

Questions worth asking

If FlashAttention is exact, why do approximate attention methods exist at all?

They attack something else: approximate methods (Linformer, Performer, sparse) cut the asymptotics from O(N²) down to linear in the number of operations, which matters for very long sequences. FlashAttention leaves O(N²) FLOPs in place but removes the quadratic memory and the traffic. For most practical lengths FlashAttention wins with no loss of accuracy; approximations are for the regime where even N² FLOPs are out of reach.

What does "memory-bound, not compute-bound" actually mean?

A GPU has a very high ratio of arithmetic to memory bandwidth: computing is cheaper than delivering the data. If an operation does little arithmetic per byte loaded (low arithmetic intensity), its speed is set by memory, not by the ALUs. Attention, with its N×N reads and writes, is exactly that case. Knowing whether a kernel is compute-bound or memory-bound is a basic GPU-optimization skill.

Online softmax sounds fiddly — why bother, can't you just take the softmax afterwards?

To take a softmax the usual way you need the whole row at once (to normalize and to subtract the maximum for stability) — and that row is precisely what we do not want to materialise. Online softmax maintains a running maximum and sum and correctly "re-weights" the result accumulated so far with every new block. It is the detail without which tiling would break the numerical stability of softmax.

What to read in the original

Read the essentials — the IO-aware framing and tiling/online softmax; the CUDA kernel details can be skipped unless you write kernels.