Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Describe Flash Attention at a high level—why is it faster than naive attention?
What you need to know
The GPU memory hierarchy
- HBM (high-bandwidth memory) is the GPU's main memory: 80 GB on an H100, about 3.35 TB/s.
- SRAM is small on-chip memory next to the compute units: a few hundred KB per streaming multiprocessor, tens of MB in total, but many times faster.
A kernel that keeps its data in SRAM and rarely touches HBM is fast. A kernel that writes big intermediates to HBM and reads them back is slow, even if it does few FLOPs.
What naive attention does
S = Q Kᵀ / sqrt(d) write T×T to HBMP = softmax(S) read T×T, write T×TO = P V read T×TAt T = 8192, one head's score matrix in fp16 is 8192 × 8192 × 2 = 134 MB. With 32 heads, one layer writes and reads about 4.3 GB of intermediate data — for one sequence — and training must keep P for the backward pass.
The two ideas
- Tiling — load a block of Q rows and a block of K/V rows into SRAM, compute their scores there, and never store the full
SorPin HBM. - Online softmax — keep a running max
mand running sumlfor each query row. When a new block arrives with a larger max, rescale the old sum and output byexp(m_old − m_new), then add the new block. - Recompute in backward — instead of storing
P, recompute score tiles from Q, K and V during the backward pass. Extra FLOPs are cheaper than the HBM traffic saved.
1import numpy as np2rng = np.random.default_rng(0)3T, d, B = 8, 4, 2 # 8 keys, processed 2 at a time4q = rng.normal(size=d)5K, V = rng.normal(size=(T, d)), rng.normal(size=(T, d))67s = K @ q / np.sqrt(d) # naive: whole row, then softmax8p = np.exp(s - s.max()); p /= p.sum()9naive = p @ V1011m, l, acc = -np.inf, 0.0, np.zeros(d) # streaming: one block at a time12for start in range(0, T, B):13 sb = K[start:start+B] @ q / np.sqrt(d)14 m_new = max(m, sb.max())15 scale = np.exp(m - m_new) # shrink old sums to the new max16 pb = np.exp(sb - m_new)17 l = l * scale + pb.sum()18 acc = acc * scale + pb @ V[start:start+B]19 m = m_new20print(np.allclose(naive, acc / l)) # TrueThe streaming loop only ever holds two scores at a time, yet gives the same answer as the full softmax. FlashAttention does this per tile on the GPU, for many rows at once.
Versions
- FlashAttention (Dao et al., 2022) — the paper reports about 3× faster GPT-2 training at sequence length 1K.
- FlashAttention-2 (2023) — better split of work across GPU thread blocks and warps.
- FlashAttention-3 (2024) — uses Hopper-specific asynchronous copies and FP8; the paper reports up to about 75% of H100 peak throughput in FP16.
You rarely call it directly: PyTorch's scaled_dot_product_attention picks a flash kernel when it can, and vLLM, TensorRT-LLM and SGLang use flash-style kernels for prefill and paged variants for decode.
A real-life example
A team fine-tunes a code-completion model on whole files at 16K context on A100s. With naive attention, one head's score matrix is 16384² × 2 bytes = 512 MB; 32 heads make 16 GB per layer, per sequence, before storing anything for backward. The run runs out of memory at batch size 1.
With FlashAttention, attention memory becomes proportional to T, and the same GPUs train at batch size 4. The final model's accuracy is the same as it would have been with naive attention, because the math is the same — only the memory traffic changed.
Follow-up questions to expect
- "Is FlashAttention an approximation?" — No. It is exact up to floating-point rounding; outputs and gradients match standard attention.
- "Does it reduce FLOPs?" — No, it does slightly more because of recomputation. It is faster because it moves far fewer bytes.
- "Does it make attention linear in T?" — Memory becomes linear; compute is still quadratic. At 128K tokens, attention FLOPs still dominate a layer.
- "How is decode different?" — With one query row, parallelism comes from splitting the KV sequence across thread blocks (Flash-Decoding) and combining partial results with the same max-and-sum trick.