PidokuInfra

Attention Computation in Practice

Basic Advanced 1h 30m Difficulty 4/5 Topic 06 of 12

Prerequisites III.07, III.08, 03


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#

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

Any 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 practice

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

Below 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 expensive

8. 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_attention picks among flash / mem-efficient / math backends based on dtype, head dim, mask type, and alignment. Falling back to math is a silent 10x regression. Use torch.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#

  1. Why can’t you materialize the attention score matrix at long context? Give numbers.
  2. Explain FlashAttention’s algorithm in terms of the memory hierarchy.
  3. Why do prefill and decode need different attention kernels?
  4. What is FlashDecoding and what problem does it solve?
  5. How does a PagedAttention kernel differ from a contiguous one, and what does the indirection cost?
  6. How should GQA be implemented in a kernel, and what’s the naive mistake?
  7. 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

↑↓ navigate↵ openesc close