Concept · Chapter 11: Inside Modern LLMs
FlashAttention
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.
You should understand first
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 , writes to HBM, reads it back to compute the softmax , writes , reads it again to compute . With 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 or 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 and a running sum , and rescale when a new block raises the maximum.
Tiny example. One row of scores arrives in two blocks, then .
- After block 1: , .
- Block 2 raises the maximum to 5. Rescale the old sum by and add the new terms: .
- Computed in one go: . Same answer.
The partial outputs ( times for each block) are rescaled the same way, so the final output is exact. For the backward pass, the scores are recomputed from , and 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
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.
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.
FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
Tri Dao · 2023
A lesson in how much speed hides in work partitioning: the same exact algorithm, about twice as fast.
How to read it: Read after FlashAttention; Section 3 lists the three changes.
Watch
Stanford Online
Stanford CS336 I Language Modeling from Scratch | Spring 2025 | Lecture 5: GPUs
The hardware background for this chapter: why memory movement, not arithmetic, so often sets the speed.