All posts

Decoding flash attention

/AI/5 min read

Flash attention produces exactly the same numbers as ordinary attention. It is faster because it never writes the score matrix to memory at all.

Flash attention computes the same thing ordinary attention computes. Same inputs, same outputs, to the last bit. It is not an approximation.

What changes is where the intermediate work happens. Ordinary attention writes a large matrix out to GPU memory and reads it back; flash attention never writes it at all. That single difference is worth several times the speed.

What ordinary attention does

Attention scores every token against every other token. For a sequence of NN tokens that is an N×NN \times N matrix — one score per pair. Softmax turns each row into weights, and those weights multiply the value vectors.

The matrix is the problem. It is produced, stored, read back for the softmax, stored again, and read back once more for the final multiply. Each of those trips crosses the slowest link the GPU has.

Put a number on it. For 8,192 tokens with a head dimension of 128, in fp16:

entriessize
the N×NN \times N score matrix67,108,864128 MB
Q, K, V and the output together4,194,3048 MB

The scratch space is 16 times larger than all the real data combined. And it grows quadratically — double the sequence length and it quadruples.

Two kinds of GPU memory

The reason this matters is that a GPU has two very different places to put things.

HBM is the big pool, tens of gigabytes. It is what people mean by "GPU memory". It is also comparatively slow.

SRAM is on-chip, right next to the compute units. It is roughly an order of magnitude faster to read, and there is very little of it — on the order of a hundred kilobytes per streaming multiprocessor.

A 128 MB matrix has no chance of fitting in SRAM. So the standard implementation parks it in HBM and shuttles it back and forth. The arithmetic is not the bottleneck; the shuttling is.

The idea: never build the matrix

Flash attention breaks the computation into tiles small enough to live in SRAM while they are being worked on.

THE N×N SCORE MATRIX one tile · 32 KB fits in on-chip SRAM whole matrix · 128 MB never built at all 4,096 tiles cover it, one at a time
Only the bright tile is ever real. The rest is computed and discarded, one tile at a time.

A 128×128 tile of scores is 32 KB in fp16. That fits comfortably. So instead of computing all 67 million scores and storing them, the GPU loads a block of queries and a block of keys, computes just that tile of scores in SRAM, uses it immediately, and discards it.

The full matrix is never assembled anywhere. It exists only one tile at a time, in fast memory, and is gone before the next tile arrives.

The problem with doing softmax in pieces

There is an obstacle. Softmax is not computed per element — it needs a whole row.

softmax(s)i=esi−max⁡(s)∑jesj−max⁡(s)\text{softmax}(s)_i = \frac{e^{s_i - \max(s)}}{\sum_j e^{s_j - \max(s)}}

Both the maximum and the sum are over the entire row. But a tile only holds part of a row. You cannot normalise until you have seen everything, and the whole point is to avoid holding everything.

Online softmax

The fix is to keep a running maximum and a running sum, and correct them as new blocks arrive.

Process a block, note its maximum and its exponential sum. When the next block turns out to contain a larger maximum, the earlier numbers were computed against the wrong reference — so scale them by emold−mnewe^{m_{\text{old}} - m_{\text{new}}} to bring them onto the new one. The same correction applies to the accumulated output.

Take a row of eight scores, in two blocks of four:

blockscoresrunning maxrunning sum
11.4, 3.9, 0.7, 2.23.91.3055
24.6, 1.1, 3.1, 0.34.61.9152

The second block raises the maximum from 3.9 to 4.6, so everything carried from block 1 gets rescaled before the new numbers are folded in.

Comparing the streamed result against a plain softmax over the whole row:

output
plain softmax over the full row(1.765581, 3.078123)
streamed, block by block(1.765581, 3.078123)
difference0

Not close — identical. That is what makes this an optimisation rather than a trade-off.

The backward pass

Training needs the attention scores again to compute gradients, and flash attention threw them away.

Rather than store the matrix, it stores only the per-row statistics — the max and the sum, which are NN numbers rather than N2N^2 — and recomputes the tiles on the way back.

That trades extra arithmetic for less memory traffic. On modern GPUs, where compute has grown far faster than memory bandwidth, that trade is worth taking.

Later versions

Flash attention 2 reorganises the work: fewer rescaling operations, parallelism over query blocks as well as batch and heads, and a better split of work between threads.

Flash attention 3 targets newer hardware specifically, overlapping data movement with computation so the transfers happen while the arithmetic is still running.

Both are engineering on the same idea. The output stays identical.

The short version

  • Flash attention returns exactly the same numbers as standard attention.
  • Standard attention builds an N×NN \times N score matrix — 128 MB at 8,192 tokens, against 8 MB for all the actual data.
  • That matrix lives in slow HBM and gets shuttled back and forth; the memory traffic, not the arithmetic, is the bottleneck.
  • Flash attention computes it in tiles small enough for fast on-chip SRAM, and never assembles the whole thing.
  • Softmax is made block-friendly by carrying a running max and sum, rescaling when a bigger max appears.
  • The backward pass stores NN statistics instead of N2N^2 scores and recomputes what it needs.