PidokuInfra

FlashAttention

Intermediate Advanced 2h Difficulty 4/5 Topic 10 of 14

Prerequisites III.07, III.08, IV.06, VI.03

★ 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 traffic

They 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#

Go
// 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:

Go
// 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-15

5. 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 budget

This 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 softmax

That 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 GB

Above 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.6x

8-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_attention chooses among flash/mem-efficient/math based on dtype, head_dim, mask type, and alignment. Falling back to math is a silent 10x regression at long context.
    Python
    with 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. Use is_causal=True or 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#

  1. What problem does FlashAttention solve? Give the memory arithmetic.
  2. Explain online softmax and why it makes tiling possible.
  3. Is FlashAttention an approximation? Justify your answer.
  4. What changed between FlashAttention-1, 2, and 3?
  5. What is FlashDecoding and why does decode need a different parallelization?
  6. How does head_dim affect which kernel you can use?
  7. How would you detect that your stack silently fell back to a non-flash backend?
  8. 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-attention and flashinfer repositories
  • Next: 11 — KV cache optimization

↑↓ navigate↵ openesc close