llm-inference-explained
← /learn · 06

Attention kernels

FlashAttention's IO-aware tiling; split-K for long contexts.

Concept

Written the obvious way, attention for one head over NN tokens is three steps, each a separate GPU kernel:

  1. S=QK⊤/dS = QK^\top/\sqrt{d}, an N×NN \times N matrix of scores, written to the GPU's main memory (HBM);
  2. read SS back, apply the softmax row by row, write PP (another N×NN \times N matrix);
  3. read PP and VV, write the output O=PVO = PV.

At 16 bits, the two N×NN\times N 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 SS or PP. 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 mm, a running sum ll and an unnormalised output. When a new tile raises the maximum, rescale what you have by emold−mnewe^{m_{\text{old}} - m_{\text{new}}} and carry on. At the end, divide by ll. 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 (m,l,o)(m, l, o), 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.

Standard
136.3 MB
50.9 µs
Tiled
18.7 MB
7.0 µs
Reduction
7.3×
N×N matrix
33.6 MB
never written when tiled

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 qq over key tiles jj: with scores si=q⋅ki/ds_i = q\cdot k_i/\sqrt d,

m′=max⁡(m,max⁡i∈jsi),l′=l em−m′+∑i∈jesi−m′,o′=o em−m′+∑i∈jesi−m′ vi,m' = \max(m, \max_{i \in j} s_i),\quad l' = l\,e^{m-m'} + \sum_{i\in j} e^{s_i - m'},\quad o' = o\,e^{m-m'} + \sum_{i\in j} e^{s_i - m'}\,v_i,

starting from m=−∞m = -\infty, l=0l = 0, o=0o = 0; the output is o/lo/l. Two partial states merge the same way, which is all split-K needs:

m=max⁡(ma,mb),l=laema−m+lbemb−m,o=oaema−m+obemb−m.m = \max(m_a, m_b),\quad l = l_a e^{m_a - m} + l_b e^{m_b - m},\quad o = o_a e^{m_a - m} + o_b e^{m_b - m}.

HBM traffic (FlashAttention, §3.2 and Theorem 2): standard attention moves Θ(Nd+N2)\Theta(Nd + N^2) elements (here 4Nd+4N24Nd + 4N^2); FlashAttention with MM elements of SRAM moves Θ(N2d2/M)\Theta(N^2 d^2 / M). With key tiles of Bc=⌈M/4d⌉B_c = \lceil M/4d \rceil there are Tc=⌈N/Bc⌉T_c = \lceil N/B_c\rceil passes over the queries, each reading QQ and OO and writing OO, plus the running statistics: 2Nd+Tc(3Nd+4N)2Nd + T_c(3Nd + 4N). 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 10−1210^{-12}.