★ The most important algorithmic optimization in modern transformer inference. Not because it’s the fastest — because without it, long context is impossible.
1. Problem → Why → Optimization#
PROBLEM Standard attention materializes an (S × S) score matrix in HBM.
At S=8192, B=8, h=32, FP16 that is 34 GB — for ONE layer.
WHY The textbook formulation computes softmax(QKᵀ/√d)V as three separate
steps, each writing its intermediate to memory.
OPTIMIZE Tile the computation and use online softmax so the score matrix
never leaves SRAM. Exact, not approximate.
TRADE-OFFS
✓ O(S) memory instead of O(S²)
✓ 2-10x faster (fewer HBM round trips)
✓ EXACT — bit-comparable to the naive version modulo floating point order
✗ requires a custom kernel per hardware generation
✗ constrained head dimensions and dtypes
✗ recomputation in the backward pass (irrelevant for inference)
WHEN TO USE Always.
WHEN NOT TO Never, for standard attention. (Only if your head_dim or dtype
isn't supported — in which case find a kernel that supports it.)2. Why it exists#
The insight, from Dao et al. (2022): attention is memory-bound, not compute-bound, and the standard implementation’s memory traffic is dominated by writing and reading a matrix you don’t actually need.
Standard attention HBM traffic (per head, per layer):
write S: S² × 2 bytes
read S: S² × 2
write P: S² × 2
read P: S² × 2
+ Q, K, V, O: 4 × S × d × 2
≈ 8S² + 8Sd bytes
FlashAttention:
read Q, K, V: 3 × S × d × 2, each block read O(S/block) times
write O: S × d × 2
≈ O(S²d/M) where M = SRAM size
→ for typical S, d, M: 5-20x less HBM trafficThey didn’t make attention compute less. They made it move fewer bytes. Which, given the roofline, is the same as making it faster.
3. Simple analogy#
Computing a weighted average of a million items with a small desk.
Naive: write all million weights on paper, spread across the floor, then normalize, then combine. Needs a warehouse.
FlashAttention: process items in batches of 100. Keep a running total and a running normalizer on your desk. When a new batch has a larger maximum, rescale what you have so far and continue. At the end, divide. Same answer, one desk.
The “rescale what you have so far” step is online softmax (Section III.07), and it’s the only non-obvious part.
4. Tiny example — the algorithm#
// FlashAttention for one head: exact attention with O(S) extra memory.
// Q, K, V are (S, d). The (S, S) score matrix is never built.
func flashAttention(Q, K, V [][]float64, blockQ, blockK int) [][]float64 {
S, d := len(Q), len(Q[0])
scale := 1 / math.Sqrt(float64(d))
O := make([][]float64, S)
for i0 := 0; i0 < S; i0 += blockQ { // a block of queries — stays in SRAM
i1 := min(i0+blockQ, S)
// running state per query in the block, in registers
m := make([]float64, i1-i0) // running max
l := make([]float64, i1-i0) // running sum of exp
acc := make([][]float64, i1-i0) // running weighted sum of V
for r := range m {
m[r], acc[r] = math.Inf(-1), make([]float64, d)
}
for j0 := 0; j0 < i1; j0 += blockK { // causal: keys up to the end of this query block
j1 := min(j0+blockK, S) // a block of K and V — loaded into SRAM
for r, i := 0, i0; i < i1; r, i = r+1, i+1 {
// scores for this (query, key block) tile — SRAM ONLY
hi := min(j1, i+1) // causal mask: query i sees keys 0..i
if hi <= j0 {
continue
}
s, mNew := make([]float64, hi-j0), m[r]
for j := j0; j < hi; j++ {
for c := 0; c < d; c++ {
s[j-j0] += Q[i][c] * K[j][c]
}
s[j-j0] *= scale
mNew = math.Max(mNew, s[j-j0])
}
// ---- online softmax update ----
alpha := math.Exp(m[r] - mNew) // rescale factor for what we already have
l[r] *= alpha
for c := range acc[r] {
acc[r][c] *= alpha
}
for j := j0; j < hi; j++ {
p := math.Exp(s[j-j0] - mNew)
l[r] += p
for c := 0; c < d; c++ {
acc[r][c] += p * V[j][c]
}
}
m[r] = mNew
}
}
for r, i := 0, i0; i < i1; r, i = r+1, i+1 { // normalize at the end
O[i] = acc[r]
for c := range O[i] {
O[i][c] /= l[r]
}
}
}
return O
}Read the three lines under “online softmax update” carefully. They are the entire
contribution. alpha rescales everything accumulated so far to the new maximum; then the new
block is added at the same scale. The result is bit-for-bit the same as computing the full
softmax (modulo floating-point ordering).
Verify:
// Check against the textbook version (IV.06's `naive`, which builds every row of scores):
Q, K, V := randn(2048, 64), randn(2048, 64), randn(2048, 64)
ref := naiveAttention(Q, K, V)
fmt.Println(maxAbsDiff(flashAttention(Q, K, V, 128, 128), ref)) // ~1e-155. Technical explanation#
The tiling and the SRAM budget#
Choose Bq, Bk so that this fits in shared memory:
Qi (Bq × d) + Kj (Bk × d) + Vj (Bk × d) + Sij (Bq × Bk)
For d=128, FP16, 228 KB shared memory (Hopper):
Bq=128, Bk=128:
128×128×2 × 3 (Q,K,V) + 128×128×2 (S) = 98 KB + 33 KB = 131 KB ✓
For d=256:
Bq=64, Bk=128 to stay within budgetThis is why head_dim affects which kernel you get and why unusual head dimensions may have no fast path.
FlashAttention-1 → 2 → 3#
FA-1 (2022)
The core idea: tiling + online softmax + recomputation in backward.
Parallelized over batch and heads only.
~2-4x over naive.
FA-2 (2023)
- Parallelize over the SEQUENCE dimension too (more blocks → better occupancy,
critical for long S with small batch)
- Reduce non-matmul FLOPs (fewer rescalings; defer the division)
- Better work partitioning between warps within a block
~2x over FA-1.
FA-3 (2024, Hopper-only)
- TMA for async bulk copies HBM→SRAM
- Warp specialization: producer warps load, consumer warps compute
- Overlap softmax (non-tensor-core) with GEMM (tensor-core) via pingpong scheduling
- FP8 support with incoherent processing for accuracy
~1.5-2x over FA-2 on Hopper; up to 75% of theoretical peak FLOPs.The decode variant: FlashDecoding#
Prefill has thousands of query rows → plenty of parallelism. Decode has one query per sequence:
Decode, batch 8, 32 heads, S=32768:
Parallelizing over (batch × heads) = 256 blocks on 132 SMs → 2 waves, poor.
And each block must serially process 32768/128 = 256 key blocks.FlashDecoding splits the key dimension:
Split the 256 key blocks across, say, 8 thread blocks.
Each computes a PARTIAL (unnormalized) output with its own (m, l).
A second kernel combines the partials using the same online-softmax rescaling.
→ 256 × 8 = 2048 blocks. GPU is full.
→ 2-4x faster decode at long context.Same mathematical trick (online softmax is associative), applied across thread blocks instead of within one.
Paged FlashAttention#
Combining with PagedAttention: the K/V blocks are gathered via a block table rather than read contiguously.
for blk in range(num_blocks_for_this_seq):
phys = block_table[seq][blk]
Kj = k_cache[phys] # ← the indirection
...Cost: one extra load per 16-token block, slightly worse coalescing at boundaries. 2-8% slower kernel, 2.5-4x more concurrency. Every production engine takes that trade.
6. Under the hood — the Hopper pipeline (FA-3)#
Producer warps (using TMA):
issue async copies of K, V tiles HBM → shared memory
signal a barrier when a tile arrives
Consumer warps:
wait on the barrier
WGMMA: Q × Kᵀ (tensor cores, from shared memory)
softmax on the result (non-tensor-core: exp, max, sum)
WGMMA: P × V (tensor cores)
rescale the accumulator
PINGPONG SCHEDULING:
while warpgroup A does softmax (non-tensor-core work),
warpgroup B does WGMMA (tensor-core work)
→ the tensor cores are never idle waiting for softmaxThat last point is the FA-3 insight: softmax’s exp/max/sum use the SFU and ALU, not tensor cores. Overlapping them with matmuls from another warpgroup keeps both pipelines busy.
7. Performance#
Prefill attention, A100, h=32, d=128, B=8, causal:
S Naive FA-2 Speedup Naive memory FA-2 memory
512 0.9 ms 0.4 ms 2.3x 0.5 GB 0.02 GB
2048 14.1 ms 4.1 ms 3.4x 8.6 GB 0.07 GB
8192 OOM 62 ms — 137 GB 0.27 GB
32768 OOM 980 ms — 2.2 TB 1.1 GB
65536 OOM 3.9 s — 8.8 TB 2.2 GBAbove S≈4096 the comparison is meaningless because the naive version cannot run. FlashAttention is what makes long context exist.
Decode with FlashDecoding, S=32768, B=1:
Standard FA-2 decode kernel: 3.2 ms
FlashDecoding (split-K): 0.9 ms 3.6x8-9. Production and mistakes#
Production:
- Always use a fused attention kernel. FlashAttention-2/3, FlashInfer, xFormers memory-efficient attention, or your engine’s built-in.
- Verify which backend is selected. PyTorch’s
scaled_dot_product_attentionchooses among flash/mem-efficient/math based on dtype, head_dim, mask type, and alignment. Falling back tomathis a silent 10x regression at long context.Pythonwith torch.nn.attention.sdpa_kernel([torch.nn.attention.SDPBackend.FLASH_ATTENTION]): out = F.scaled_dot_product_attention(q, k, v, is_causal=True) # raises if flash isn't usable — use this in tests to catch silent fallbacks - Check head_dim support. Common kernels support 64, 96, 128, 256. Unusual values may have no fast path.
- Use FA-3 on Hopper if your stack supports it — 1.5-2x over FA-2, and FP8 attention.
- Use FlashDecoding-style split-K for long-context decode.
- Watch for kernel regressions after upgrades. Attention kernel selection changes between versions.
Mistakes:
- Writing textbook attention in custom code. Fine at S=512, fatal at S=8192.
- Not noticing a backend fallback.
- Assuming FlashAttention is an approximation. It’s exact.
- Using a prefill kernel for decode. Wrong parallelization; 3-4x slower at long context.
- Materializing an explicit attention mask tensor
(B, h, S, S)— that defeats the purpose even if the kernel is flash. Useis_causal=Trueor a compact mask representation. - Padding sequences instead of using varlen. Wastes the kernel’s efficiency.
10. Hands-on exercise#
A. Implement it. Complete and verify the flash_attention function in section 4. Confirm it
matches the reference to floating-point precision. Then measure peak memory for both at
S = 512, 2048, 8192. Plot memory vs S for each; confirm O(S) vs O(S²).
B. Online softmax alone. Verify separately that your online softmax over blocks equals the full softmax exactly. This is the piece people get wrong.
C. Backend detection. For a matrix of (dtype, head_dim, mask type, alignment), determine which backend PyTorch’s SDPA selects. Build a compatibility table for your version. Which configurations silently fall back?
D. FlashDecoding. Implement decode attention two ways: parallelizing over queries only, and splitting over keys. At S=32768, B=1, measure both and explain the difference using SM occupancy.
E. Read the real kernel. Read the FlashAttention-2 CUDA source (csrc/flash_attn/).
Identify: the tiling loop, the online softmax update, the shared memory layout, and the causal
block-skipping logic.
11. Interview questions#
- What problem does FlashAttention solve? Give the memory arithmetic.
- Explain online softmax and why it makes tiling possible.
- Is FlashAttention an approximation? Justify your answer.
- What changed between FlashAttention-1, 2, and 3?
- What is FlashDecoding and why does decode need a different parallelization?
- How does head_dim affect which kernel you can use?
- How would you detect that your stack silently fell back to a non-flash backend?
- How does paged attention interact with FlashAttention, and what does the indirection cost?
12. Further reading#
- [ESTABLISHED] Dao, Fu, Ermon, Rudra, Ré, “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness” (NeurIPS 2022) — read this paper
- [ESTABLISHED] Dao, “FlashAttention-2” (2023)
- [ESTABLISHED] Shah et al., “FlashAttention-3” (2024)
- [ESTABLISHED] “Flash-Decoding for long-context inference” (Dao et al., blog post)
- [ESTABLISHED] Milakov & Gimelshein, “Online normalizer calculation for softmax” (2018)
- [REFERENCE]
flash-attentionandflashinferrepositories - Next: 11 — KV cache optimization