Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Describe Flash Attention at a high level—why is it faster than naive attention?


One query tile, never leaving SRAMLoad a Qtile into SRAMStream K andV tiles in,one by oneScore thetile, updaterunning max and sumRescale thepartial outputWrite onlythe finaloutput to HBMNaive attention writes a 134 MB score matrix per head at 8k tokens.
Same math, far fewer bytes moved: the T by T matrix only ever exists one tile at a time.

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

Text
S = Q Kᵀ / sqrt(d)      write T×T to HBMP = softmax(S)          read T×T, write T×TO = P V                 read T×T

At 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

  1. 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 S or P in HBM.
  2. Online softmax — keep a running max m and running sum l for each query row. When a new block arrives with a larger max, rescale the old sum and output by exp(m_old − m_new), then add the new block.
  3. 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.
Python
import numpy as nprng = np.random.default_rng(0)T, d, B = 8, 4, 2                      # 8 keys, processed 2 at a timeq = rng.normal(size=d)K, V = rng.normal(size=(T, d)), rng.normal(size=(T, d))s = K @ q / np.sqrt(d)                 # naive: whole row, then softmaxp = np.exp(s - s.max()); p /= p.sum()naive = p @ Vm, l, acc = -np.inf, 0.0, np.zeros(d)  # streaming: one block at a timefor start in range(0, T, B):    sb = K[start:start+B] @ q / np.sqrt(d)    m_new = max(m, sb.max())    scale = np.exp(m - m_new)          # shrink old sums to the new max    pb = np.exp(sb - m_new)    l = l * scale + pb.sum()    acc = acc * scale + pb @ V[start:start+B]    m = m_newprint(np.allclose(naive, acc / l))     # True

The 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.