Why FlashAttention was a breakthrough
Same math, same exact outputs, same asymptotic compute — and yet it made attention several times faster and unlocked long context. The trick was noticing attention was a memory problem, not a compute problem.
On this page
- The picture version
- Why it exists
- Why it matters now
- The short answer
- How it works
- The naive kernel, and why it wastes the machine
- Attempt 1 — tile it, the way every good matmul kernel is tiled
- Why attempt 1 breaks: softmax needs the whole row
- The last leak: the tiles still leave the chip between steps
- What the numbers actually showed
- What it did not change
- The follow-ups, briefly
- Check yourself
- Famous related terms
- Going deeper
The picture version
Six pictures for a reader who has never written a line of GPU code. The prose below fills in the seams the pictures skip.
1 · The problem
The model fits comfortably. The job dies anyway.
2 · The mismatch
The maths didn’t change. Nothing was approximated. It just got much faster.
3 · Where it was going
A tiny fast shelf next to the workers, and a big slow warehouse across the yard.
4 · The obvious fix, and the wall it hits
Do it a tile at a time — except one step insists on seeing the whole row.
5 · The last leak
Do all that and the tiles still get carried to the warehouse between steps.
6 · Keep this card
None of the three ideas was new. Noticing which problem it was, was.
Why it exists
You paste a long document into a chat model and it answers in a couple of seconds. Anyone who has tried to run a transformer on their own GPU knows the other version of that experience: the model itself fits in memory comfortably, you feed it a long input, and the job dies with an out-of-memory error anyway. The thing that blew up wasn’t the model. It was a scratch value the model computes and throws away — and that scratch value is the running example for this whole post: the N×N table of attention scores, one row and one column per token, that a transformer builds every time it reads a sequence of length N. At N = 8,000 that’s 64 million numbers per attention head, per layer, and the standard implementation writes all of them to memory and reads them back.
Now the confusing part. If you read the FlashAttention paper expecting a clever new approximation that avoids building that table, you’ll be confused. There isn’t one. The output is the same softmax attention the original transformer paper described in 2017 — exact, up to floating-point reordering. The asymptotic FLOP count is the same (the backward pass actually does more arithmetic, because it recomputes intermediates instead of storing them — more on that below). The model isn’t changed. And yet it ran several times faster on real GPUs and made it practical to train and serve transformers at sequence lengths that previously ran out of memory.
That mismatch — same math, dramatically different wall-clock — is the entire point. Before FlashAttention, the prevailing intuition was that attention was compute-bound: it does an N×N matmul, and matmuls are what GPUs are for, so what’s left to optimize? The answer turned out to be everything. The bottleneck was never the multiplications. It was that N×N table being shuffled in and out of HBM three or four times per attention layer, with the GPU’s actual compute units sitting idle waiting for memory.
The breakthrough was noticing this and writing a kernel that respected the memory hierarchy. The technique — tiling plus an online softmax trick — is not novel in computer science; it’s the standard playbook for memory-bound numerical kernels. The novelty was applying that playbook to attention, demonstrating that attention had been a memory-bandwidth problem all along, and shipping a kernel everyone could use. The 2022 paper by Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, and Christopher Ré (arXiv:2205.14135) made that case, and it was unusually consequential — the FA-2/FA-3 lineage continues today, and the underlying kernel (or a derivative) is widely available through major training and inference stacks. (My read on adoption breadth, not a sourced market study.)
Why it matters now
The reason FlashAttention is still a load-bearing piece of infrastructure four years later, instead of a paper everyone moved on from, is that the bandwidth gap it exploits has only gotten wider:
- Long context is normal. Frontier models advertise 200k–1M token windows. Materializing an N×N scores matrix in HBM at N=1,000,000 is impossible — it’s a trillion entries per layer. FlashAttention (or a descendant of it) is the standard way current stacks keep attention’s activation memory linear in N rather than quadratic; without something equivalent, very long context wouldn’t fit on the device at all.
- Attention is more of the work than it used to be. Push N far enough for a given model and the quadratic attention term overtakes the linear-in-N feed-forward term. Where that crossover sits depends on the model’s width — a wide model pushes it out — but “long enough” is now a routine input length, not an exotic one, so optimizing attention went from a small win to a large share of the budget.
- Hardware kept getting more memory-skewed. Each new GPU generation adds compute throughput faster than it adds memory bandwidth. Algorithms that ignore the memory hierarchy fall further behind every generation, not closer. (See why VRAM is the bottleneck.)
- It set the template. The “treat it as an IO problem” reframe turned into a generation of follow-up work — FlashAttention-2 (2023), FlashAttention-3 (2024) targeting Hopper-class H100s, and a wave of fused kernels for everything else in the transformer stack.
The short answer
FlashAttention = exact softmax attention + tiling + online softmax + kernel fusion
Picture to keep: the N×N score table never exists all at once. Picture a small window sliding over it, one tile at a time, right next to the compute units — each tile computed, used, and thrown away before the next arrives, with two running numbers per row carried along to stitch the answer together.
It’s the same attention as before — same inputs, same exact outputs, same asymptotic FLOP count — restructured so the N×N intermediate matrix never gets written to slow GPU memory. The work is done in small blocks that fit in fast on-chip memory and are processed in a single fused kernel. The math is unchanged; only the data movement is. (The backward pass does add some recomputation FLOPs as the price of not storing the attention matrix.)
How it works
The honest way to tell this is as three attempts to get rid of that 64-million-number table, each one breaking in a way that forces the next fix. None of the fixes is novel computer science on its own; the sequence is what nobody had run to completion for attention.
The naive kernel, and why it wastes the machine
A GPU has two memory tiers that matter here. SRAM is on-chip, sits next to the compute units, and is fast — but tiny, on the order of tens of MB total across all streaming multiprocessors. HBM is off-chip DRAM, much slower per byte but much bigger (tens of GB).
The original FlashAttention paper cites an A100 example where SRAM bandwidth is roughly 19 TB/s versus 1.5 TB/s for HBM — about 13× the throughput in roughly 1/2000th the capacity. Standard attention writes the full N×N scores matrix out to HBM, reads it back to apply softmax, writes the result back, and reads it again to multiply by V. Those round-trips are the bottleneck. The matmul units finish early and wait.
This is the core empirical claim of the paper: an IO analysis plus benchmarks argue that standard attention is memory-bound, not compute-bound. Once you see that, the fix is structural: don’t materialize the N×N matrix in HBM at all.
flowchart LR
subgraph chip [On-chip — fast, tiny]
CU[Compute units] <--> SRAM["SRAM · ~19 TB/s · tens of MB"]
end
SRAM <--> HBM["HBM · ~1.5 TB/s · tens of GB"]
Standard attention parks the N×N scores matrix in HBM and crosses that slow link three or four times per layer. FlashAttention keeps every block left of the HBM boundary and only writes back the small final output.
Attempt 1 — tile it, the way every good matmul kernel is tiled
Tiling is the textbook technique for memory-bound numerical work — it’s how good BLAS matmul kernels have always worked. FlashAttention applies it to attention.
Conceptually:
for each block of K, V (loaded into SRAM):
for each block of Q (loaded into SRAM):
compute partial scores = Q_block · K_blockᵀ # in SRAM
compute partial softmax + partial output # in SRAM
update running output and softmax statistics # tiny
If that loop works, the N×N table is never assembled in HBM at all: each block is computed, used, and discarded inside SRAM, and the kernel writes back only the final output (shape N×d, same size as the input) plus a small O(N) per-row softmax statistic that the backward pass needs (the running max m and the running sum-of-exps l; some implementations store them combined as m + log l). Total HBM traffic drops from O(N²) to roughly O(N²·d²/M), where M is the SRAM size — a big concrete win because M is large enough that the d²/M factor is small.
The catch is hiding in the two innocuous middle lines of that pseudocode. Line 3 is fine — a matmul over a block is just a smaller matmul. Line 4 is not.
Why attempt 1 breaks: softmax needs the whole row
Tiling is only “trivial textbook stuff” if your operation is associative across blocks. Matmul is. Softmax isn’t, naively — softmax(x_i) = exp(x_i) / Σ exp(x_j) requires knowing the whole row to compute any entry, because of the denominator. That global dependency is what stops you from just tiling attention the obvious way.
The fix, often called the online softmax trick, predates FlashAttention — Milakov and Gimelshein described an online normalizer in 2018 (arXiv:1805.02867) — but FlashAttention is what plumbed it through to a full attention kernel. The idea: as you stream through blocks of a row, keep two running scalars per row — the running max m (for numerical stability) and the running sum-of-exps l. When a new block arrives with its own local max and local sum, you rescale the running quantities and the partial output, then fold in the new block. After the last block, the result equals what a one-shot softmax would have produced, up to floating-point reordering. No approximation, no drift.
This is the part of FlashAttention that took genuine engineering taste to land — getting the rescaling exactly right while staying numerically stable, doing it inside a single CUDA kernel, and making the backward pass also work. The forward pass keeps only O(N) softmax statistics; the backward pass uses recomputation (re-deriving the attention matrix on the fly during backprop instead of storing it) to keep memory linear too.
The last leak: the tiles still leave the chip between steps
Tiling and online softmax are algorithm-level fixes, and you could implement both and still leave much of the win on the table — because a framework runs each step as its own GPU kernel. The unfused baseline FlashAttention was measured against — attention written out as ordinary framework ops — launches the score matmul, then softmax, then the output matmul separately, and every kernel boundary is a place where results get written back to HBM and read in again. (Today’s frameworks ship fused attention by default.) Your carefully tiled blocks leave the chip between steps anyway.
The final ingredient closes that: fuse all of attention — score matmul, softmax, output matmul — into one CUDA kernel. FlashAttention does the whole thing inside one kernel, so blocks live in SRAM for the duration. This is the same fusion idea behind a lot of hand-tuned GPU work; FlashAttention’s contribution wasn’t inventing fusion, it was finding the version that fused all of attention without sacrificing exactness.
What the numbers actually showed
The reported speedups depend a lot on what you measure. The 2022 paper highlights:
- ~7.6× speedup on the attention computation itself on GPT-2.
- ~3× end-to-end speedup on GPT-2 training (sequence length 1k).
- ~15% end-to-end speedup on BERT-large training (sequence length 512), beating the then-current MLPerf record.
End-to-end gains are smaller than attention-only gains because attention isn’t 100% of the work, especially at short sequence lengths. The longer the sequence, the bigger the fraction of total time attention takes, and the bigger FlashAttention’s relative win — which is part of why long-context training became practical around the time this kernel and its successors became standard. (I’m not claiming sole causation; that’s an oversimplification — sequence-parallel training, ring attention, and architectural changes all contributed.)
What it did not change
It’s worth being precise about the limits, because the reframing only goes so far.
- Attention is still O(N²) in compute. The quadratic in N is unchanged. The asymptotic FLOP count is the same; the backward pass actually does more arithmetic because it recomputes the attention matrix instead of storing it. (The 2022 paper’s own benchmark shows a higher GFLOP count for FlashAttention than standard attention on one of its tests — and still wins, because it spends them on hot data instead of HBM round-trips.) FlashAttention changes the bandwidth, not the arithmetic.
- It is exact. No approximation, unlike sparse / linear / sliding-window attention. The output is identical to vanilla attention up to floating-point reordering.
- The bottleneck moved. After FlashAttention, attention is much closer to compute-bound. Follow-up work (FA-2, FA-3) focuses on extracting more compute utilization (parallelism, async pipelines, low precision) rather than further bandwidth reduction — my read is that the easy bandwidth wins were taken by the original kernel, though no head-to-head bandwidth attribution study has been published to settle it.
The follow-ups, briefly
FlashAttention-2 (2023) reorganized the parallelism (different work partitioning across thread blocks, fewer non-matmul operations) and roughly doubled throughput on Ampere-class GPUs. FlashAttention-3 (2024) targets Hopper (H100) specifically — exploiting TMA, warp-specialized async pipelines, and FP8 / BF16 with block-quantized FP8 — and reports lifting H100 attention utilization well above FA-2’s roughly 35%. Be careful quoting the headline number, because the versions differ: the NeurIPS 2024 paper reports up to ~840 TFLOPs/s in BF16 (~85% utilization) and up to ~1.3 PFLOPs/s in FP8, while the July 2024 arXiv preprint reported ~740 TFLOPs/s in FP16 (~75%) and close to 1.2 PFLOPs/s in FP8. My read, not consensus: each follow-up is a constant-factor win that gets harder to find as the easy bandwidth wins have already been taken.
You started with FlashAttention = exact softmax attention + tiling + online softmax + kernel fusion. Which of those three additions had never been done before? — none of them, individually. Tiling is decades old, online softmax was published in 2018, kernel fusion is standard practice. What FlashAttention added is + the observation that attention was memory-bound in the first place. That observation is the breakthrough: once you believe the N×N table’s trips to memory are the cost, all three fixes are the obvious ones. My read on why nobody had assembled them earlier — the paper doesn’t argue this, so take it as interpretation — is that the field was busy trying to speed up the multiplications.
Check yourself
Before you go — a colleague benchmarks FlashAttention against standard attention at sequence length 128 and finds it barely faster. Did they do something wrong?
Answer
Probably not — that’s the expected result. At N = 128 the score table is 16,384 entries, small enough that it costs little to move and may well stay in cache anyway. The saving FlashAttention offers scales with the size of the thing it refuses to materialize, so the win grows with N. This is the same reason the 2022 paper’s headline end-to-end numbers differ so much by workload: ~15% on BERT-large at length 512 versus ~3× on GPT-2 at length 1k. If your sequences are short, attention probably isn’t where your time is going, so an attention kernel has little to give you back — whatever it does win there comes from fusion overhead, not from the table it avoided.
And a harder one: FlashAttention’s backward pass does more floating-point arithmetic than the standard implementation, and it’s still faster. How can spending more FLOPs be a win?
Answer
Because FLOPs weren’t the scarce resource — bytes moved were. The backward pass needs the attention matrix, and it has two options: read it back from HBM (cheap in arithmetic, expensive in bandwidth) or recompute it from Q and K in fast on-chip memory (expensive in arithmetic, nearly free in bandwidth). On a chip where compute throughput far outruns memory bandwidth, the second trade wins. Generalize it: any time you find yourself on the memory-bound side of a machine, recomputing is a legitimate way to buy speed — this is the same bargain as gradient checkpointing.
Famous related terms
- Why attention is quadratic —
attention cost ∝ N² · d. The constraint FlashAttention works around, not the one it removes. Useful background for understanding what didn’t change. - Memory bandwidth —
memory bandwidth = the actual bottleneck for most LLM operations, not FLOPs. The general principle FlashAttention is one famous instance of. - Tiling / blocking —
tiling = split a big matrix op into small blocks that fit in fast cache + reuse each block while it's hot. The cornerstone trick for memory-bound numerical kernels since long before deep learning. - Online softmax —
online softmax = streaming softmax + a running max and a running sum that get rescaled as new blocks arrive. The piece that lets softmax be tiled exactly. Predates FlashAttention; Milakov & Gimelshein 2018. - Kernel fusion —
kernel fusion = collapse multiple GPU ops into one kernel + keep intermediates in registers/SRAM. The reason all of attention can stay on-chip in one pass. - Recomputation (gradient checkpointing) —
recomputation = trade FLOPs for memory by recomputing activations in the backward pass instead of storing them. How FlashAttention’s backward pass keeps memory linear in N. - KV cache —
KV cache = past tokens' K and V tensors stored across decode steps. Orthogonal to FlashAttention but related: KV cache addresses decode, FlashAttention addresses prefill and training.
Going deeper
- Dao, Fu, Ermon, Rudra, Ré, FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (NeurIPS 2022) — the primary source, and the place to check what the IO-bound claim is actually based on rather than taking my summary of it.
- Milakov & Gimelshein, Online normalizer calculation for softmax (2018) — answers “how can softmax possibly be computed one block at a time,” in four pages and without any GPU context to wade through.
- Rabbit hole: FlashAttention-3 (Shah et al., NeurIPS 2024) — answers “what’s left to optimize once the bandwidth problem is solved,” which turns out to be a completely different set of tricks (asynchrony, warp specialization, FP8).