FlashAttention, annotated
How to read this page
- Any dotted word explains itself on hover, focus or tap, and so does every symbol in every equation.
- The tiling animation in §3.1 is the core of the paper: press Play in both modes and watch the traffic counters.
This page assumes you know what attention computes; if not, start with the Attention Is All You Need companion or the attention lesson. Every idea climbs the ladder: everyday picture, tiny example, diagram, the math, why it matters today.
Abstract · original
“We argue that a missing principle is making attention algorithms IO-aware […] accounting for reads and writes between levels of GPU memory.”Dao et al. (2022), Abstract
Everyday picture
A chef works at a small counter; the ingredients live in a big pantry down a flight of stairs. However fast the chef chops, if every step of the recipe means a trip down and up the stairs, the meal is slow. Standard attention is that chef: it computes quickly, but it writes a huge intermediate table to slow memory and reads it back several times. FlashAttention reorganizes the recipe so each ingredient is carried up once, everything possible is done at the counter, and only the finished dish goes back down.
What the paper claims
- An exact attention algorithm: the same result as standard attention, not an approximation.
- Much less memory traffic between the GPU's large, slow memory and its small, fast memory, and a proof that no exact algorithm can do asymptotically better across all fast-memory sizes.
- Real speedups: 15% faster BERT-large training than the MLPerf 1.1 record, 3× faster GPT-2 training, 2.4× on long-sequence benchmarks.
- Memory that grows only linearly with sequence length, enabling much longer contexts and the first Transformers to beat chance on the Path-X (16K tokens) and Path-256 (64K tokens) tasks.
Why it matters today
FlashAttention and its successors are now the default way attention runs on GPUs. Much of the jump from 2,048-token to million-token context windows rests on this idea.
1 Introduction · original
Everyday picture
Attention's cost grows with the square of the sequence length (see the O(n²) discussion in the attention lesson). Many earlier papers tried to reduce the arithmetic with approximations: skip some word pairs, compress others. The authors point out that these tricks often did not make training faster in practice, because arithmetic was not the bottleneck. Moving data was. It is like trying to speed up the chef by teaching faster chopping when the real problem is the stairs.
Tiny example
Count the reads and writes for a 1,024-token sequence and a head width of 64. The inputs Q, K and V hold 3 × 1,024 × 64 ≈ 197,000 numbers. The score table S holds 1,024 × 1,024 ≈ 1,049,000 numbers: over five times more than all the inputs put together, and standard attention writes it to slow memory and reads it back, then does the same again for the softmax table P. Most of the traffic is intermediate results nobody asked for.
Why it matters
The lesson generalizes far beyond attention: on modern hardware, the number of memory transfers often predicts speed better than the number of arithmetic operations. The inference lesson uses the same idea to explain why generating text is memory-bound.
2 Background · original
2.1 GPU memory and speed · original
Everyday picture
A GPU has two kinds of memory. HBM is the big pantry: tens of gigabytes, but relatively slow to reach. SRAM is the counter: tiny, but right next to the arithmetic units and about ten times faster. Every calculation has to bring its data up to SRAM first.
Hover or tap a level of memory.
Reading it: the higher a box, the faster and smaller it is. HBM at the bottom holds gigabytes; SRAM in the middle holds only 192 KB per processor (about 20 MB across all 108 of them) but moves data roughly ten times faster. Arithmetic can only use data that has come up to the top. The arrow between SRAM and HBM is labelled the bottleneck because modern GPUs can compute far faster than HBM can feed them, so the question is how many times data crosses that arrow.
Two kinds of operation
An operation is compute-bound if its time is spent on arithmetic (a big matrix multiply) and memory-bound if its time is spent waiting for data (softmax, masking, dropout, layer norm: a few operations per number read). The measure that decides which is arithmetic intensity, operations per byte moved. The standard fix for memory-bound work is kernel fusion: load once, do several steps, write once.
Why it matters today
Compute has grown faster than memory bandwidth for years, so this imbalance keeps getting worse. Thinking in terms of data movement is now a core skill for anyone optimizing models.
2.2 Standard attention · original
Everyday picture
The textbook recipe in three trips: compute every score and store the whole table downstairs; fetch the table, turn each row into shares with softmax, store the new table downstairs; fetch it again with the values and blend. Each trip moves an N × N table.
In words: “score every query against every key, turn each row of scores into shares, then use the shares to blend the values.” (The usual 1/√d scaling is folded into Q here, as in the paper.)
With the numbers: for GPT-2's N = 1,024 and d = 64, S and P each hold 1,048,576 numbers per head. In 16-bit precision that is 2 MB per head per layer, per sequence; at N = 16,384 it is 512 MB per head. Both tables are written to HBM and read back.
In Python:
N, d = 1024, 64
# numbers in S (and again in P), per head
N * N # → 1048576
# MB at 2 bytes per number
N * N * 2 // 2 ** 20 # → 2
N = 16_384
N * N * 2 // 2 ** 20 # → 512
Why it matters
The math is fine; the problem is where S and P live. Standard implementations need O(N²) memory just to hold them, which is what capped context lengths, and the repeated round trips are what made attention slow.
3 FlashAttention: the algorithm · original
3.1 Softmax in pieces · original
Everyday picture
The obstacle to working in small tiles is softmax: each share is a score divided by the total of the whole row, and you don't know the total until you have seen every score. The trick is to keep a running tally. Imagine counting votes from several ballot boxes: count one box, note the leader and the tally; when the next box arrives, rescale the old tally if there's a new leader and add the new votes. You never need all the ballots on the table at once.
In words: “subtract the block's largest score before exponentiating (so nothing overflows), add up the results, and divide.”
Now the key step: two blocks' summaries can be merged without revisiting their scores.
In words: “the overall maximum is the bigger of the two block maxima; the overall total is each block's total, shrunk by how far its own maximum sits below the overall maximum, added together.”
With the numbers: scores x = (1, 3 | 2, 0), split into two blocks. Block 1: m = 3, f = (e−2, e0) = (0.135, 1), ℓ = 1.135. Block 2: m = 2, f = (1, 0.135), ℓ = 1.135. Merge: m = 3, ℓ = e0 × 1.135 + e−1 × 1.135 = 1.135 + 0.418 = 1.553. Directly: e1−3 + e0 + e2−3 + e0−3 = 0.135 + 1 + 0.368 + 0.050 = 1.553. Same answer, one block at a time.
In Python:
import math
# one block's (m, ℓ)
def summary(x):
# m(x)
m = max(x)
# f(x)
f = [math.exp(x_i - m) for x_i in x]
# ℓ(x) = Σ_i f(x)_i
return m, sum(f)
m1, l1 = summary([1, 3])
m2, l2 = summary([2, 0])
m1, round(l1, 3), m2, round(l2, 3) # → (3, 1.135, 2, 1.135)
# merged m(x)
m = max(m1, m2)
# merged ℓ(x)
l = math.exp(m1 - m) * l1 + math.exp(m2 - m) * l2
m, round(l, 3) # → (3, 1.553)
# all four scores at once
round(summary([1, 3, 2, 0])[1], 3) # → 1.553
The partial outputs are merged the same way: FlashAttention keeps each query row's running output O, its running maximum m and running total ℓ, and rescales O whenever a new block raises the maximum. After the last block, O is exactly softmax(QKᵀ)V.
Why it matters
This “online softmax” was known before; the paper's contribution is building a whole fast GPU kernel around it. The same stable-softmax trick (subtract the max) appears in the attention lesson's softmax function.
3.1 Tiling, animated
“The main idea is that we split the inputs Q, K, V into blocks, load them from slow HBM to fast SRAM, then compute the attention output with respect to those blocks.”Dao et al. (2022), §3.1
Everyday picture
Carry one tray of keys and values up to the counter, then bring up the queries a tray at a time and do all the work for that pair of trays before sending anything back. The giant score table never exists anywhere: each tile of it lives briefly on the counter and is thrown away.
Reading it: the square is the N × N score matrix for N = 1,024 and head width d = 64, cut into tiles of the chosen block size; rows are blocks of queries, columns are blocks of keys and values. In standard mode, every tile is written to HBM as S (orange), read back and rewritten as P (purple), then read again to produce the output (green): three full passes over the whole matrix. In FlashAttention mode the outer loop walks across the columns (one block of keys and values loaded once), and the inner loop walks down the rows; the red tile is the one being computed in SRAM, and once done it is never stored. The bars count the numbers moved to or from HBM so far, using the accounting of the paper's Algorithms 0 and 1. Bigger blocks mean fewer passes over the queries, so less traffic, but a block must fit in SRAM, which caps how big it can be.
3.1 Recomputation and kernel fusion
Everyday picture
Training needs a backward pass, which normally reuses the stored S and P tables. FlashAttention stores neither. Instead it keeps only the output and the two small running statistics (m and ℓ) per row, and simply recomputes each tile of S and P during the backward pass. That means more arithmetic, but it is cheaper to redo the chopping at the counter than to fetch pre-chopped ingredients from the pantry.
And because the whole recipe (matrix multiply, mask, softmax, dropout, matrix multiply) happens tile by tile on-chip, it all fits in one GPU program: one fused kernel that reads the inputs once and writes the output once.
Why it matters
This is a clear case of spending more arithmetic to save memory traffic, and winning. In the paper's measurement (next section) FlashAttention does more floating-point operations than standard attention and still runs about 5.7 times faster.
3.2 How much traffic is saved · original
Everyday picture
How many stair trips does each recipe need? Standard attention's count grows with the size of the table, N². FlashAttention's count also grows with N², but divided by how much the counter can hold, so a bigger counter means proportionally fewer trips.
In words: “standard attention moves the inputs plus the full N × N table; FlashAttention moves the N × N work shrunk by the factor d² / M, which is small because a head's width squared is much smaller than the fast memory.”
With the numbers: with d = 64 and M around 100,000 (the paper's “around 100KB”), d2 / M = 4,096 / 100,000 ≈ 0.041, so FlashAttention needs roughly 24 times fewer HBM accesses, ignoring constant factors. The paper also proves that no exact attention algorithm can do asymptotically better for every SRAM size.
In Python:
d, M = 64, 100_000
# the shrink factor d²/M
d ** 2 / M # → 0.04096
# N² (standard) over N² d² / M (FlashAttention)
round(M / d ** 2) # → 24
| Standard | FlashAttention | |
|---|---|---|
| Arithmetic (GFLOPs) | 66.6 | 75.2 |
| HBM reads and writes (GB) | 40.3 | 4.4 |
| Runtime (ms) | 41.7 | 7.3 |
The measured table is the whole argument in three rows: 13% more arithmetic, 9 times less memory traffic, 5.7 times faster. The theory's “24 times” ignores constants; the measured 9 times is what the real kernel achieved.
Why it matters today
Counting memory traffic, not just arithmetic, is now the standard way to reason about kernel speed. The same accounting explains why the KV cache and weight reads dominate text generation; see the inference lesson's roofline discussion.
3.3 Block-sparse FlashAttention · original
Everyday picture: if you know in advance that some trays don't matter (say, words very far apart), skip those tiles entirely. The paper extends the kernel to take a fixed pattern of tiles to compute and skip. Traffic and runtime shrink in proportion to the fraction skipped. Unlike FlashAttention itself this is an approximation, and it is what reached 64K-token sequences in §4.
4 Experiments · original
Everyday picture
Faster is only interesting if the model stays exactly as good. Because FlashAttention is exact, the models reach the same quality; the question is how much sooner, and what the freed-up memory makes possible.
| Experiment | Before | With FlashAttention |
|---|---|---|
| BERT-large to target accuracy, 8 × A100 | 20.0 min (MLPerf 1.1 record) | 17.4 min (15% faster) |
| GPT-2 small training, 8 × A100 | 9.5 days (Hugging Face) | 2.7 days, same perplexity 18.2 |
| GPT-2 medium training | 21.0 days (Hugging Face) | 6.9 days, same perplexity 14.3 |
| GPT-2 small with 4× the context (4K) | 1K context, perplexity 18.2, 4.7 days (Megatron-LM) | perplexity 17.5, 3.6 days |
| Path-X (16K tokens) | every earlier Transformer at chance | 61.4% accuracy |
| Path-256 (64K tokens, block-sparse) | every earlier Transformer at chance | 63.1% accuracy |
The highlighted row shows the second benefit: with memory now growing linearly, GPT-2 could be trained on 4 times longer texts, still faster than the previous 1K-context training, and it got better (lower perplexity) because it could see more context. The memory footprint measured up to 20 times smaller than exact-attention baselines.
Why it matters today
Longer context turned out to be one of the most valuable capabilities of language models, and this is one of the papers that made it affordable.
5 Limitations and future directions · original
The authors name three: every new attention variant needs a hand-written low-level CUDA kernel, which is hard and not portable across GPU generations; the IO-aware idea should extend beyond attention to other layers; and a single GPU is the limit of the analysis, while multi-GPU attention adds another level of the memory hierarchy. All three became active research areas.
What happened next
| Development | What it does |
|---|---|
| FlashAttention-2 (Dao, 2023) | Reorganizes the loops and work split to use the GPU better, roughly doubling speed again |
| Built into frameworks | PyTorch's scaled_dot_product_attention can dispatch to a FlashAttention kernel, so most users get it without asking |
| Long-context models | Linear memory in sequence length is one of the enablers of context windows of hundreds of thousands of tokens |
| IO-aware thinking elsewhere | The same data-movement accounting drives work on serving, such as PagedAttention, and on fused kernels generally |
The authors' reference implementation is open source at github.com/Dao-AILab/flash-attention.
Glossary
Every term with hover guidance on this page, in one place.