An annotated companion · AI Primer

FlashAttention, annotated

About this page. This is a companion, not a copy. It follows the paper section by section, quotes at most a sentence or two per section (clearly marked and attributed), and explains everything in its own words. The paper is distributed under arXiv's standard licence, so its figures are redrawn from scratch and its tables are not reproduced; a few key numbers are restated in new tables with attribution. Every section heading links to the original.

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.

Compute unitsarithmetic happens here On-chip SRAM192 KB × 108 · about 19 TB/s the bottleneck HBM (GPU main memory)40 to 80 GB · 1.5 to 2.0 TB/s smaller and faster towards the top

Hover or tap a level of memory.

The memory hierarchy of an NVIDIA A100 GPU, redrawn with the numbers given in Dao et al. (2022), §2.1.

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.

Algorithm: Block size:

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
Measured on GPT-2 medium attention, forward plus backward (N = 1,024, d = 64, 16 heads, batch 64, A100). Restated from Dao et al. (2022), Figure 2
StandardFlashAttention
Arithmetic (GFLOPs)66.675.2
HBM reads and writes (GB)40.34.4
Runtime (ms)41.77.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.

Selected results, restated from Dao et al. (2022), Tables 1, 2 and 4 and §4.2
ExperimentBeforeWith FlashAttention
BERT-large to target accuracy, 8 × A10020.0 min (MLPerf 1.1 record)17.4 min (15% faster)
GPT-2 small training, 8 × A1009.5 days (Hugging Face)2.7 days, same perplexity 18.2
GPT-2 medium training21.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 chance61.4% accuracy
Path-256 (64K tokens, block-sparse)every earlier Transformer at chance63.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

DevelopmentWhat it does
FlashAttention-2 (Dao, 2023)Reorganizes the loops and work split to use the GPU better, roughly doubling speed again
Built into frameworksPyTorch's scaled_dot_product_attention can dispatch to a FlashAttention kernel, so most users get it without asking
Long-context modelsLinear memory in sequence length is one of the enablers of context windows of hundreds of thousands of tokens
IO-aware thinking elsewhereThe 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.