1. What is it?#
How attention is actually computed on a GPU, as opposed to how it’s written in a paper. The difference is substantial: the textbook formula is unusable at production sequence lengths, and every real implementation restructures it.
2. Why does it exist as a separate topic?#
Because the naive implementation has a fatal flaw — it materializes an (S, S) matrix — and
because prefill and decode need structurally different kernels. Attention is the one operator
where “just call the library” isn’t enough knowledge; you need to know which of several kernels
you’re getting and why.
3. Simple analogy#
Computing a weighted average of a million items without writing down all the weights.
Naive: compute all million weights, store them, normalize, then combine. Needs a million-entry scratchpad.
Streaming: process items in chunks, maintaining a running total and a running normalizer, rescaling as you go. Needs a scratchpad of one chunk. Same answer.
That’s online softmax (Section III.07) applied to attention, and it’s FlashAttention.
4. Tiny example: three implementations#
// attention.go — naive attention vs tiled "online softmax" attention (the FlashAttention idea).
// One head, one sequence, to keep the indices readable.
package main
import (
"fmt"
"math"
"math/rand"
)
const S, D = 512, 64
type seq [S][D]float64
func dot(a, b *[D]float64) (s float64) {
for i := range a {
s += a[i] * b[i]
}
return s
}
// naive: for each query, MATERIALIZE the whole row of scores, softmax it, then mix values.
func naive(q, k, v *seq) *seq {
out, scale := new(seq), 1/math.Sqrt(D)
for i := 0; i < S; i++ {
scores := make([]float64, i+1) // row i of the (S, S) matrix; causal: only j <= i
m := math.Inf(-1)
for j := range scores {
scores[j] = dot(&q[i], &k[j]) * scale
m = math.Max(m, scores[j])
}
var sum float64
for j := range scores {
scores[j] = math.Exp(scores[j] - m)
sum += scores[j]
}
for j, p := range scores {
for c := 0; c < D; c++ {
out[i][c] += p / sum * v[j][c]
}
}
}
return out
}
// tiled: walk the keys in blocks, keeping only a running max, a running sum and a running
// output per query. The (S, S) score matrix never exists.
func tiled(q, k, v *seq, block int) *seq {
out, scale := new(seq), 1/math.Sqrt(D)
for i := 0; i < S; i++ {
m, l := math.Inf(-1), 0.0 // running max, running denominator
var acc [D]float64 // running (unnormalized) output
for j0 := 0; j0 <= i; j0 += block {
j1 := min(j0+block, i+1) // causal mask: stop at i
var s [256]float64
mNew := m
for j := j0; j < j1; j++ {
s[j-j0] = dot(&q[i], &k[j]) * scale // a small tile of scores
mNew = math.Max(mNew, s[j-j0])
}
alpha := math.Exp(m - mNew) // rescale what we accumulated under the old max
l *= alpha
for c := range acc {
acc[c] *= alpha
}
for j := j0; j < j1; j++ {
p := math.Exp(s[j-j0] - mNew)
l += p
for c := range acc {
acc[c] += p * v[j][c]
}
}
m = mNew
}
for c := range acc {
out[i][c] = acc[c] / l
}
}
return out
}
func main() {
q, k, v := new(seq), new(seq), new(seq)
for i := 0; i < S; i++ {
for c := 0; c < D; c++ {
q[i][c], k[i][c], v[i][c] = rand.NormFloat64(), rand.NormFloat64(), rand.NormFloat64()
}
}
a, b := naive(q, k, v), tiled(q, k, v, 128)
var diff float64
for i := range a {
for c := range a[i] {
diff = math.Max(diff, math.Abs(a[i][c]-b[i][c]))
}
}
fmt.Println("max diff:", diff) // ~1e-15 — same answer
}Run this. Then compare peak memory at S=8192: the naive version allocates
1·2·8192·8192·4 = 537 MB for the scores; the tiled version allocates ~1 MB per block. That
gap is why FlashAttention exists.
5. Technical explanation#
The four attention kernels you actually use#
Real engines dispatch among structurally different implementations:
1. PREFILL / CHUNKED PREFILL (many queries, many keys)
FlashAttention-2/3 forward. Tiles over both q and k.
Compute-bound. Uses tensor cores heavily.
Causal masking skips ~half the work.
2. DECODE / SINGLE-QUERY (1 query per sequence, many keys)
"FlashDecoding" / PagedAttention kernel.
Memory-bound: reads the whole KV cache.
Parallelizes over the KEY dimension (split-K style) because
there aren't enough queries to fill the GPU.
3. APPEND / MIXED (a few queries, many keys)
For speculative decoding (verify k tokens) and chunked prefill.
Varlen kernels with cu_seqlens.
4. CROSS ATTENTION (encoder-decoder, some VLMs)
Keys/values from a different source; no causal mask.The decode kernel deserves attention. With batch 8 and 32 heads you have 256 query rows — nowhere near enough to fill 132 SMs with meaningful work if you parallelize only over queries. FlashDecoding splits the key dimension across thread blocks, each computing a partial (un-normalized) result, then combines them with a second reduction pass. This is exactly split-K GEMM applied to attention, and it can give 2-4x on long-context decode.
PagedAttention’s kernel difference#
With a paged KV cache, keys and values are not contiguous:
Contiguous: k[b, h, 0:S, :] — one strided read
Paged: block_table[b] = [17, 3, 92, 45, ...]
read block 17 (16 tokens), block 3, block 92, ...The kernel takes the block table as an argument and gathers. Costs: an extra indirection per block, slightly worse coalescing at block boundaries. Benefit: zero fragmentation, sharing between sequences, instant free. The measured kernel overhead is a few percent; the memory win is 2-4x more concurrency. An easy trade.
GQA in the kernel#
Naive: k.repeat_interleave(h // h_kv, dim=1) ← materializes h copies. Wasteful!
Kernel: query head i reads KV head i // (h/h_kv) ← no copy, just index arithmeticAny implementation that does the repeat_interleave is reading (and writing) 8x more KV bytes
than necessary for a GQA-8 model. Fused kernels index directly. If you write custom attention,
this is the first thing to get right.
Masking variants#
Causal j <= i
Sliding window i - W < j <= i
Prefix-LM bidirectional over the prompt, causal over the generation
ALiBi add -m·(i-j) to the score instead of masking
Custom arbitrary (e.g. document boundaries in packed sequences)Fused kernels implement masks by skipping blocks entirely rather than computing and masking:
Causal, block (i_block, j_block):
if j_block_start > i_block_end: skip entirely — no computation
if j_block_end <= i_block_start: no masking needed — full block
else: apply element-wise mask (diagonal block)This is why causal attention costs ~half of full attention rather than the same.
6. Under the hood: FlashAttention’s memory hierarchy use#
HBM (3.35 TB/s) SRAM / shared memory (~19 TB/s)
───────────────── ────────────────────────────────
Q, K, V, O Q tile (Br × d)
K tile (Bc × d)
V tile (Bc × d)
S tile (Br × Bc) ← never leaves SRAM
running m, l ← in registers
Traffic: naive O(S² + S·d) HBM accesses
flash O(S² · d / M) where M = SRAM size
≈ 4-10x fewer HBM accesses in practiceThe trick: choose tile sizes so Br·d + 2·Bc·d + Br·Bc fits in the 228 KB of shared memory per
SM on Hopper. For d=128 and FP16, typical is Br=128, Bc=128.
FlashAttention-3 additionally uses Hopper’s TMA (async bulk copy) and warp specialization (some warps only load, some only compute) to overlap memory and math, and supports FP8.
7. Performance implications#
Measured, A100, causal attention, d_head=128, h=32, B=8:
S naive time naive memory flash time flash memory
512 0.9 ms 0.5 GB 0.4 ms 0.02 GB
2048 14 ms 8.6 GB 4.1 ms 0.07 GB
8192 OOM 137 GB 62 ms 0.27 GB
32768 OOM 2.2 TB 980 ms 1.1 GBBelow S=1024 the difference is modest. Above S=4096 the naive version simply cannot run. FlashAttention is not primarily a speed optimization — it’s a feasibility one.
For decode at long context:
S=32768, B=32, GQA-8, FP16:
KV bytes read per step = 2 × 32 layers × 8 × 128 × 2 × 32768 × 32 = 137 GB
→ 41 ms per step just to read KV
→ this is why long-context decode is expensive8. Production implications#
- Never ship naive attention. Use FlashAttention, FlashInfer, xFormers, or your engine’s kernel.
- Check which backend is selected. PyTorch’s
scaled_dot_product_attentionpicks among flash / mem-efficient / math backends based on dtype, head dim, mask type, and alignment. Falling back tomathis a silent 10x regression. Usetorch.backends.cuda.sdp_kernel(enable_math=False)in testing to catch it. - Head dim matters. Fast kernels support specific head dims (64, 128, 256). An unusual head dim may have no fast path.
- Long context needs FlashDecoding-style split-K or your decode will underutilize the GPU.
- Watch for the GQA repeat. Confirm your kernel indexes rather than materializes.
9. Common mistakes#
Materializing scores. Fatal above ~4k context.
Silent backend fallback. Check.
repeat_interleave for GQA. 8x wasted KV bandwidth.
Assuming one kernel works for both prefill and decode. Their shapes are opposite.
Applying the mask by adding a large negative number in FP16. -1e9 overflows to -inf
(fine), but -65504 is the FP16 limit — use -inf or a proper masked kernel.
Ignoring alignment. Some flash kernels require the head dim and pointers to be aligned; unaligned inputs fall back.
10. Hands-on exercise#
A. Implement tiled attention. Complete and verify the attention_tiled function above.
Measure peak memory vs the naive version at S ∈ {512, 2048, 8192}. Plot.
B. Which backend? For various (dtype, head_dim, mask) combinations, determine which backend
scaled_dot_product_attention selects. Build a compatibility table for your PyTorch version.
C. GQA correctness. Implement GQA attention two ways: with repeat_interleave and with
index arithmetic. Verify they agree. Measure the KV bytes read by each.
D. Decode parallelism. Implement decode attention parallelizing over queries only, then over keys (split-K). At S=32768, B=1, compare. Explain the difference in terms of SM occupancy.
E. Causal savings. Measure full vs causal attention time at S=4096. Is the ratio 2:1? Why or why not?
11. Interview questions#
- Why can’t you materialize the attention score matrix at long context? Give numbers.
- Explain FlashAttention’s algorithm in terms of the memory hierarchy.
- Why do prefill and decode need different attention kernels?
- What is FlashDecoding and what problem does it solve?
- How does a PagedAttention kernel differ from a contiguous one, and what does the indirection cost?
- How should GQA be implemented in a kernel, and what’s the naive mistake?
- How does a causal mask get implemented efficiently at the block level?
12. Further reading#
- [ESTABLISHED] Dao et al., “FlashAttention” (2022) and “FlashAttention-2” (2023)
- [EMERGING→ESTABLISHED] Shah et al., “FlashAttention-3” (2024)
- [ESTABLISHED] Kwon et al., “PagedAttention” (SOSP 2023) §4
- [REFERENCE] FlashInfer library — a good survey of attention kernel variants
- [REFERENCE] “Flash-Decoding for long-context inference” (Dao et al. blog post)
- Next: 07 — Tensor layouts and memory