Week 13. GPUs and FlashAttention
Learn in the app: tutor, coding problems →
Core: SRAM and HBM, arithmetic intensity and the roofline, the critical batch, FlashAttention · Depth: Triton, measurements with intervals in bench.py · ≈ 11 h core / 22 h total
So far we have counted FLOPs. This week is about why FLOPs lie: most of a transformer's operations
are bound not by arithmetic but by reading and writing memory. FlashAttention is the main example here:
the same FLOPs (even more, with recomputation), yet several times faster, and O(S) memory instead of O(S²).
Without it, the long context of week 20 simply would not exist.
Step 1. The GPU model and the roofline
In plain terms. An H100 can do ~989 trillion operations per second but can read memory at only 3.35 TB/s. To keep the arithmetic units busy, you need ~300 operations for every byte you read. Adding two vectors in bf16 reads 4 bytes, writes 2 and does one operation: 1 operation per 6 bytes. The GPU is then less than 0.1% utilized: it is just waiting on memory. A large matmul reuses every number it reads hundreds of times and is bound by arithmetic instead.
- The GPU model: SMs, warps, registers, SRAM vs HBM, bandwidth. An SM (streaming multiprocessor; an H100 has 132) executes warps, groups of 32 threads running the same instruction. Registers and SRAM (shared memory, hundreds of KB per SM) are fast and tiny; HBM (main memory, 80 GB) is large and slower
- Arithmetic intensity, the roofline model. Why attention is memory-bound and matmuls are compute-bound.
Time ≥
max(FLOPs / P, bytes / W), intensityI = FLOPs / bytes, thresholdP/W ≈ 295FLOP/byte. Below the threshold an operation is memory-bound. In naive attention each element of theS×Smatrix gets only a few operations (mask, exp, division), andQKᵀhas an inner dimension of justH = 64…128 - Critical batch size. A pass over layers with
Nweights readsN·bbytes (bbytes per weight) and does2·N·BFLOPs, whereBis the number of tokens in that pass. Reading and computing take equal time atB* = P·b / (2W). For bf16 (b = 2) this is simplyP/W: about 295 for an H100, about 153 for an A100 (312 TFLOP/s, 2.04 TB/s). Below the threshold a pass costs one read of the weights, and extra tokens are almost free; above it, each token adds time. The threshold is counted in tokens, not sequences: prefill of one 2,000-token prompt is already compute-bound, decode for 64 users gives 64 tokens, and draft verification in speculative decoding givesB·(γ + 1). Quantization (storing weights or the cache in 8 or 4 bits instead of 16; the main treatment is in week 14) moves the threshold: int8 weights with bf16 arithmetic halve the bytes, andB* ≈ 148; if the arithmetic is int8 too (about 1979 TOPS on an H100), the threshold goes back to ≈ 295. In practice you often cannot reach the threshold: for Llama-3-8B on an H100 at 8k context, 59 sequences fit in memory; you run out of room for the cache before you run out of FLOPs (week 11) - In decode, batching does not save attention: each request has its own cache, and its intensity is about
N/KFLOP/byte at anyB(4 for Llama-3-8B, 1 with MHA). Hence the second point of GQA: the cache is not only smaller but also read more productively - Kernel fusion: why custom kernels exist at all. A chain of elementwise operations fused into one kernel
reads and writes the tensor once instead of
ktimes. Inscripts/profile_train.py(week 22) elementwise operations take 39% of the time
Step 2. FlashAttention
In plain terms. At S = 4096 with 32 heads, the score matrix in fp16 takes 1 GiB per layer per example.
Naive attention writes it to HBM, reads it for the softmax, writes the result, and reads it again to multiply by V.
FlashAttention takes a block of queries and a block of keys, computes their piece of the scores right in SRAM,
and immediately multiplies by a block of V. A running maximum and sum (online softmax from week 4) stitch the blocks together
correctly. Only the output goes to HBM; the full matrix never exists anywhere.
- FlashAttention: tiling into blocks + online softmax (week 4!) → we never materialize the
S×Smatrix. Backward with recomputation. FA2, FA3. For eachQblock we keepm,d,oand walk over theK,Vblocks; on a new maximum the olddandoare multiplied bye^{m_old − m_new}. HBM accesses:Θ(S²H²/M)versusΘ(SH + S²)for the naive version, whereMis the SRAM size


- Backward stores only the output and the logsumexp of each row, and recomputes the blocks of scores and probabilities
from
Q,K,V. The extra FLOPs are cheaper: arithmetic is idle, memory is not - The result is exact, not approximate: the math is the same, only the order of summation changes. With a causal mask, blocks entirely above the diagonal are skipped: another ~2× saving
- FA2 also parallelizes over sequence length and cuts non-matmul operations. FA3 uses Hopper's asynchrony (TMA, warp specialization) and FP8
- An introduction to Triton: what the programming model is, what a simple kernel looks like. One program
processes one block:
tl.program_id, offsets viatl.arange,tl.load/tl.storewith a mask on the tail,tl.dotfor a block matmul. The compiler handles threads and shared memory
Common mistakes
- Not rescaling the old
dandoon a new maximum:test_online_softmax_matches_direct_computation,test_online_softmax_block_size_does_not_matter - Starting with
m = 0or not handlingm = −inf: on a fully masked blocke^{−inf − (−inf)}gives NaN. The reference implementation checkstorch.isfinite(m); large logits are checked bytest_online_softmax_handles_large_logits - Timing the GPU without synchronization: kernels are asynchronous, so you only measure the enqueueing.
That is why the benchmark has
sync()and aWARMUP(the first call also compiles and allocates memory)
Code → scripts/attention_bench.py: naive attention (naive_attention from modules.py) versus
F.scaled_dot_product_attention in time and memory for S = 256…4096,
plus a table of the quadratic growth of the score matrix up to a 131k context. Functions: timed (median
after warmup), sync, score_matrix_bytes. FlashAttention in miniature is online_softmax_weighted_sum
in stability.py, with tests in tests/test_stability.py.
Do not focus on the speedup (~4×) but on two other numbers: the memory saving
is exactly S/H (the ratio of the S×S matrix to the S×H output; the script computes it rather than measuring it),
and the discrepancy between implementations is ~5e-7, meaning FlashAttention is exact, not approximate.
SDPA picks the kernel for the device itself (on CUDA it is FlashAttention, on CPU its own flash implementation,
which you can see in the week 22 profile): the numbers depend on the hardware, the picture does not. Optionally, write your own Triton kernel.
Code → nanolm/bench.py: time_fn (warmup is not part of the measurements, sync for CUDA, a replaceable clock),
BenchResult (mean, std, p50, p95, the half-width of a 95% interval from Student's t distribution; how to compare two variants with a statistical test is in week 18),
t_critical_95, pareto_frontier (the variants that cannot be improved on one axis without getting worse on another).
The task is in exercises_en/bench.py, check: NANOLM_IMPL=exercises_en pytest tests/test_bench.py -v.
The point here is discipline, not formulas: one number without warmup and spread is not an argument.
With 5 measurements the honest interval with t = 2.776 is 1.4 times wider than with the usual 1.96.
Math (Track D): D17: arithmetic intensity of a matmul and of decode on the roofline;
D18: multiplication order and linear attention: Q(KᵀV) instead of (QKᵀ)V.
Interview question of the week: "Why is FlashAttention faster if it does no fewer FLOPs?"
A 3-minute structure: the roofline and the intensity threshold → naive attention moves S×S through HBM →
tiling + online softmax remove this traffic → backward recomputes instead of reading → the result is exact →
the main consequence is not speed but O(S) memory: without it a 128k context does not fit.
Sources: Williams et al., Roofline (2009); Dao et al., FlashAttention (2022); Dao, FlashAttention-2 (2023); Shah et al., FlashAttention-3 (2024); Tillet et al., Triton (2019); Milakov & Gimelshein, Online softmax (2018).
Week outcomes
- I can compute the arithmetic intensity of a matmul and of an elementwise operation and place them on the roofline.
- I can derive the critical batch size
P·b / (2W)and say how quantization moves it. - I can explain FlashAttention through tiling and online softmax without materializing the
S×Smatrix. - I can run
attention_bench.pyand explain why the memory saving is exactlyS/H. - I can read a simple Triton kernel and say what each program does.
Self-check
- Why is attention memory-bound while a large matmul is compute-bound?
- What does the FlashAttention backward recompute, and why is recomputing cheaper than storing?
- What resource does kernel fusion save, and why does that matter more than FLOPs for elementwise operations?
- What is the critical batch of an H100 in bf16 and with int8 weights, and why is it counted in tokens?