The one idea: attention is memory-bound, not compute-bound
The usual way to reason about a kernel is to count floating-point operations. For attention that instinct misleads you. On a modern GPU the arithmetic units are so fast that the bottleneck is almost never the multiplies — it is moving data between the large, slow off-chip memory and the tiny, fast on-chip memory.
FlashAttention’s founding observation (Dao et al., 2022) is that standard attention spends most of its wall-clock time reading and writing a large intermediate matrix, not computing it. Fix the memory traffic and the kernel speeds up even though the FLOP count is unchanged. This is what IO-aware means: design around the memory hierarchy, minimizing bytes moved to and from slow memory, not around the operation count. Everything else follows from taking that one goal seriously.
The memory hierarchy: HBM versus SRAM
A GPU has two kinds of memory that matter here. HBM (high-bandwidth memory) is the large pool — tens of gigabytes — where your tensors normally live. It is fast in absolute terms (an H100 moves roughly 3 TB/s) but it is the slow tier relative to the chip. SRAM is the on-chip scratchpad attached to each streaming multiprocessor — only about 100–200 KB per SM, but an order of magnitude faster and far lower latency.
The gap is the whole story. Reading from SRAM is cheap; reading from HBM is the expensive event you want to avoid. A kernel is memory-bound when it stalls waiting on HBM while the compute units sit idle. FlashAttention’s design question is blunt: how few times can we touch HBM and still get the exact right answer?
Why standard attention wastes memory bandwidth
Standard attention, for one head, takes Q, K, V of shape [N, d] and computes S = QK^T (shape [N, N]), then P = softmax(S) row-wise, then O = PV (shape [N, d]). The trouble is S: it is quadratic in sequence length and it is written to HBM, read back for the softmax, and read again for the PV multiply.
For N = 8192 and fp16, a single N-by-N matrix is about 128 MB — per head, per layer. Materializing it means O(N^2) HBM reads and writes, and that traffic, not the matmuls, dominates the runtime. It also costs O(N^2) memory just to exist, which is why naive attention runs out of memory long before the model runs out of useful context.
The fix: tile the computation
FlashAttention refuses to ever hold S whole. It splits Q into row blocks and K, V into column blocks — say 128 each — sized so a Q block and a K/V block fit together in SRAM. The kernel loops: load one Q block, stream the K and V blocks past it, and for each pair compute the small block of scores on-chip.
Because each score block is formed, used, and discarded inside SRAM, the full N-by-N matrix is never written to HBM at all. HBM sees only the inputs Q, K, V and the final output O — all of size O(N·d). The tiling is the mechanical half of the idea. The subtle half is how you take a correct softmax over a whole row when you only ever see one block of that row at a time.
The problem tiling creates: softmax needs the whole row
Softmax is not a per-element function — it normalizes across an entire row: softmax(x_i) = exp(x_i) / Σ_j exp(x_j). To divide by that denominator you seemingly need every score in the row at once, which is precisely what tiling denies you. Worse, in practice softmax is computed in a numerically stable form by subtracting the row max first: exp(x_i - m) / Σ_j exp(x_j - m) with m = max_j x_j, and you do not know that max until you have seen the whole row either.
So tiling alone breaks the softmax. FlashAttention resolves this with online softmax: a way to compute the row’s normalized result incrementally, updating a running answer block by block. The exact algebra is the sibling article’s territory; here we build the intuition for why it works.
Online softmax: the running-max, running-sum idea
The trick is to carry two small running statistics per Q row as you stream the K/V blocks: the maximum score seen so far, m, and the running sum of exponentials l = Σ exp(score - m). When a new block arrives with a larger maximum, the old m is stale — every exponential you accumulated was scaled by the wrong offset.
The repair is a single rescale. If the max rises from m to m', you multiply the accumulated sum and output by exp(m - m'), which corrects every earlier term at once, then add the new block. Because softmax is invariant to subtracting a constant from all its inputs, this yields exactly the same numbers as a one-shot softmax — no approximation, only a reordering of arithmetic.
Rescale and accumulate: the forward pass in one loop
Put the pieces together and the forward pass is a tidy nested loop. For each Q block, initialize a running max m = -∞, a running denominator l = 0, and an output accumulator O = 0. Then, for each K/V block:
S_ij = Q_i · K_j^T # small block, in SRAM
m_new = max(m, rowmax(S_ij))
P_ij = exp(S_ij - m_new) # rescaled probabilities
l = exp(m - m_new)*l + rowsum(P_ij)
O = exp(m - m_new)*O + P_ij · V_j
m = m_newAfter the last K/V block, divide once: O_i = O / l, and write that row block out to HBM. Every quantity in the loop is a small tile living in SRAM; the only HBM writes are the final outputs.
O(N) memory: the footprint that unlocks long context
Because the N-by-N scores are never stored, the memory FlashAttention needs beyond its inputs and output is just the per-row running statistics — a handful of numbers per query row, i.e. O(N), not O(N^2). That is the difference between a context length that fits and one that does not.
Concretely, the quadratic 128 MB score matrix for N = 8192 evaporates; the working set is a few tiles of a few tens of KB. Quadrupling the context still quadruples the compute — attention is inherently O(N^2) in FLOPs and FlashAttention does not change that — but the memory stays linear, which is why 128k- and million-token contexts became practical on hardware that could never hold their score matrices.
The backward pass: recompute instead of store
Training needs gradients, and gradients normally need the attention probabilities P that the forward pass produced. Storing P would drag O(N^2) memory right back in — defeating the point. FlashAttention’s answer is recomputation: the forward pass saves only the cheap O(N) softmax statistics (the per-row m and l), and the backward pass regenerates each P tile on the fly from Q, K, and those saved statistics.
This trades extra compute for far less memory traffic — gradient checkpointing specialized to attention. Because the kernel was memory-bound, redoing arithmetic on-chip is nearly free relative to the HBM round trips it avoids. The exact gradient formulas are derived in the math companion; the architectural point is that the backward pass, too, never materializes the full matrix.
A tiny worked example
Take one query row with four keys, split into two blocks. Scores are block A = [1, 3] and block B = [2, 5]. Process A first: running max m = 3, and l = exp(1-3) + exp(3-3) = 1.135.
Block B arrives with a larger value, 5, so m' = 5. Rescale the old sum by exp(3-5) = 0.135, then add block B’s terms exp(2-5) + exp(5-5) = 1.050, giving l = 0.135×1.135 + 1.050 = 1.203. The plain one-shot denominator over all four scores (max 5) is also 1.203 — identical. The streaming answer is exact, not approximate.