Concept
Written the obvious way, attention for one head over tokens is three steps, each a separate GPU kernel:
- , an matrix of scores, written to the GPU's main memory (HBM);
- read back, apply the softmax row by row, write (another matrix);
- read and , write the output .
At 16 bits, the two matrices for an 8,192-token sequence are 128 MiB each, per head, per layer, written and read again. The arithmetic is cheap next to that traffic.
FlashAttention never writes or . It loads a tile of queries and a tile of keys and values into the chip's fast on-chip memory (SRAM: the FlashAttention paper cites 192 KB per streaming multiprocessor on an A100, against 40 to 80 GB of HBM), computes that tile's scores, and folds them straight into the output. The obstacle is the softmax: it needs each row's maximum and sum, which seem to need the whole row first. The fix is an online softmax. Keep, per query row, a running maximum , a running sum and an unnormalised output. When a new tile raises the maximum, rescale what you have by and carry on. At the end, divide by . The result is exact, not an approximation: the widget below checks it against the textbook version.
FlashAttention-2 reorganised the loops and the work split across threads for better utilisation; later versions target newer GPUs. The tiling idea is the same.
Decode has a different problem. One new query per sequence, attending over a long cache. A kernel that parallelises over batch × heads has too little work to fill the GPU when the batch is small, while each query walks thousands of keys alone. Split-K (Flash-Decoding) also splits the keys: several workers each process a slice of the cache, producing a partial , and a small final step merges them with the same rescaling rule. Again exact.
Interactive
HBM traffic, and exact tiling
Top: elements moved between HBM and the chip for one attention head, standard (FlashAttention's Algorithm 0) against tiled (Algorithm 1), at 16-bit precision and the simulator's derated H100 bandwidth (80% of 3.35 TB/s). The default SRAM is the paper's 192 KB per A100 multiprocessor. The tiled count grows with d² and shrinks with SRAM: try d = 256 with little SRAM. Bottom: real computations in your browser.
48 queries × 48 keys, d = 16: tiled vs textbook, max |Δ| = 1.6e-15
One query over 4,096 keys, merged from 4 partials: max |Δ| = 1.5e-15
- #0: m=16.81 l=2.0
- #1: m=15.88 l=1.8
- #2: m=15.53 l=2.9
- #3: m=15.22 l=2.6
Maths
Online softmax for one query over key tiles : with scores ,
starting from , , ; the output is . Two partial states merge the same way, which is all split-K needs:
HBM traffic (FlashAttention, §3.2 and Theorem 2): standard attention moves elements (here ); FlashAttention with elements of SRAM moves . With key tiles of there are passes over the queries, each reading and and writing , plus the running statistics: . The widget plots both.
Code
// src/lib/inference/attentionKernels.ts — fold one tile into the state (excerpt)
const mNew = Math.max(st.m, mTile);
const alpha = st.m === -Infinity ? 0 : Math.exp(st.m - mNew); // rescale the past
const o = st.o.map((v) => v * alpha);
let l = st.l * alpha;
s.forEach((x, i) => {
const p = Math.exp(x - mNew);
l += p;
V[from + i]!.forEach((v, c) => (o[c]! += p * v));
});
The tests run tiled attention with tiles from 1 key to the whole sequence, and split-K with up to 1,500 workers over 1,000 keys, and require agreement with the textbook result to .