Skip to content
Road to Intelligence

Concept · Chapter 11: Inside Modern LLMs

FlashAttention

Should knowUnderstand12 minDifficulty

FlashAttention computes exact attention in tiles that stay in the GPU's small, fast on-chip memory, combining partial softmaxes as it goes, so the n × n score matrix is never written to slow memory: same result, much less memory traffic, memory linear in sequence length.

The problem

Standard attention writes the full score matrix to the GPU's main memory (HBM) and reads it back for the softmax and the multiplication by V. For long sequences that traffic, not the arithmetic, dominates the time.

The solution

Load blocks of Q, K and V into on-chip SRAM, compute the scores for one block, update a running maximum and sum so the softmax can be finished later, accumulate the output, and move on; recompute instead of storing scores for the backward pass.

The consequence

Attention became several times faster and could handle much longer sequences without approximation. It is a lesson that generalizes: on modern hardware, reorganizing work to reduce memory traffic can beat reducing arithmetic.

Memory, not arithmetic

A GPU has a large main memory (HBM) and a much smaller, faster memory on the chip. The FlashAttention paper's example: an A100 has 40–80 GB of HBM at 1.5–2.0 TB/s and 192 KB of on-chip SRAM in each of its 108 streaming multiprocessors, with bandwidth estimated around 19 TB/s Established.

Standard attention computes S=QK⊤S = QK^\top, writes SS to HBM, reads it back to compute the softmax PP, writes PP, reads it again to compute PVPV. With n×nn \times n matrices (the quadratic cost), those reads and writes dominate.

FlashAttention is an IO-aware exact attention algorithm: it uses tiling to reduce reads and writes between HBM and on-chip SRAM Established. It never stores SS or PP in HBM.

The trick that makes tiling possible

The softmax of a row needs the row's maximum and the sum of all its exponentials, which seem to require the whole row first. They don't: keep a running maximum mm and a running sum ℓ\ell, and rescale when a new block raises the maximum.

Tiny example. One row of scores arrives in two blocks, [1,3][1, 3] then [2,5][2, 5].

  • After block 1: m=3m = 3, ℓ=e1−3+e3−3=0.135+1=1.135\ell = e^{1-3} + e^{3-3} = 0.135 + 1 = 1.135.
  • Block 2 raises the maximum to 5. Rescale the old sum by e3−5=0.135e^{3-5} = 0.135 and add the new terms: ℓ=1.135×0.135+e2−5+e5−5=0.154+0.050+1=1.203\ell = 1.135 \times 0.135 + e^{2-5} + e^{5-5} = 0.154 + 0.050 + 1 = 1.203.
  • Computed in one go: e−4+e−2+e−3+e0=0.018+0.135+0.050+1=1.203e^{-4} + e^{-2} + e^{-3} + e^{0} = 0.018 + 0.135 + 0.050 + 1 = 1.203. Same answer.

The partial outputs (PP times VV for each block) are rescaled the same way, so the final output is exact. For the backward pass, the scores are recomputed from QQ, KK and VV instead of being stored, the same compute-for-memory trade as activation recomputation.

What it bought

FlashAttention trained BERT-large 15% faster than the MLPerf 1.1 speed record and GPT-2 3× faster at 1K tokens, and its memory use is linear rather than quadratic in sequence length Established. FlashAttention-2 changed how the work is divided across the GPU, roughly doubling speed again and reaching 50–73% of an A100's theoretical peak Established.

What to remember

  • Exact attention, not an approximation: the outputs are the same.
  • GPU memory hierarchy: big, slower HBM (A100: 40–80 GB at 1.5–2.0 TB/s) and small, fast on-chip SRAM (192 KB per SM).
  • Tiling plus an online softmax (running max and sum) avoids writing the n × n matrix.
  • Backward pass recomputes scores from Q, K, V instead of storing them.
  • FlashAttention: 3× faster GPT-2 training at 1K tokens; FlashAttention-2: about 2× more, 50–73% of A100 peak.

Key papers

Optional

Training Deep Nets with Sublinear Memory Cost

Tianqi Chen, Bing Xu et al. · 2016

Activation checkpointing: trade a little extra computation for a large cut in training memory. Every large model run uses some form of it.

~25 min readarXiv:1604.06174✓ verified 2026-10-04
Essential

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

Tri Dao, Daniel Y. Fu et al. · 2022

Showed that attention was slow because of memory traffic, not arithmetic, and fixed it exactly: same outputs, far fewer reads and writes, memory linear in sequence length.

How to read it: Figure 1 (the memory hierarchy and the tiling loop) and Algorithm 1 carry the idea; the IO-complexity proofs can wait.

~50 min readarXiv:2205.14135✓ verified 2026-10-05

Watch