Serving · Kernels
- result
- exact same as standard
- memory
- O(N) standard O(N²)
- A100 SRAM
- 192 KB × 108 SMs
- A100 HBM
- 40–80 GB
- shown
- toy: 8 tokens, d 4, tiles 4 × 4
All steps
Fast and slow memory
A GPU has a little very fast memory on the chip (SRAM, next to the arithmetic units) and a lot of slower memory beside it (HBM). Moving data between them often takes longer than the arithmetic, so fast kernels keep work in SRAM.
SRAM 19 TB/s · HBM 1.5 TB/sStandard attention
Ordinary attention runs as separate steps, each reading its inputs from HBM and writing its result back. Two of those results, the scores S and the weights P, are N × N: for long inputs they dominate both the traffic and the memory.
S, P: N × N written to HBM and read backAn online softmax
Softmax needs a row’s maximum and sum before any weight is final. An online softmax keeps a running maximum m and sum ℓ: when a new chunk brings a larger maximum, everything so far is rescaled by e^(mold − mnew). The result is exactly the ordinary softmax.
ℓ ← e^(mold − mnew) ℓ + Σ e^(s − mnew)Tile by tile
FlashAttention loads a block of queries and streams blocks of keys and values through SRAM. Each 4 × 4 tile of scores lives only on the chip; the running m, ℓ and output are rescaled as each tile arrives, and divided by ℓ at the end. Real arithmetic, checked against ordinary attention.
Oi = Σj e^(Sij − m) Vj / ℓ, tile by tileLess traffic, no N × N
The N × N matrices are never stored, so attention’s memory grows with N instead of N², and much less data crosses between HBM and the chip. The paper reports 2 to 4 times faster attention and the same exact result; the backward pass recomputes the tiles instead of storing them.
memory O(N) · traffic ~4× lower
Code
# standard: three kernels, S and P (N × N) go through HBM
S = Q @ K.T / sqrt(d); P = softmax(S); O = P @ V
# FlashAttention: one kernel; for each block of queries, stream K and V
for i in query_blocks: # Q_i stays in SRAM
m, l, acc = -inf, 0, 0
for j in key_blocks: # K_j, V_j loaded into SRAM
s = Q[i] @ K[j].T / sqrt(d) # one tile, never written out
m_new = max(m, s.max(-1))
p = exp(s - m_new)
l = exp(m - m_new) * l + p.sum(-1)
acc = exp(m - m_new) * acc + p @ V[j]
m = m_new
O[i] = acc / l # exactly softmax(S) @ V