Skip to content
Snula
Curriculum
RU Open

Curriculum

Week 13. GPUs and FlashAttention

Phase 3. Inference, GPUs, scaling · week 13 of 24

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), intensity I = FLOPs / bytes, threshold P/W ≈ 295 FLOP/byte. Below the threshold an operation is memory-bound. In naive attention each element of the S×S matrix gets only a few operations (mask, exp, division), and QKᵀ has an inner dimension of just H = 64…128
  • Critical batch size. A pass over layers with N weights reads N·b bytes (b bytes per weight) and does 2·N·B FLOPs, where B is the number of tokens in that pass. Reading and computing take equal time at B* = P·b / (2W). For bf16 (b = 2) this is simply P/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 gives B·(γ + 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, and B* ≈ 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/K FLOP/byte at any B (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 k times. In scripts/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×S matrix. Backward with recomputation. FA2, FA3. For each Q block we keep m, d, o and walk over the K, V blocks; on a new maximum the old d and o are multiplied by e^{m_old − m_new}. HBM accesses: Θ(S²H²/M) versus Θ(SH + S²) for the naive version, where M is the SRAM size
FlashAttention: the S×S matrix never reaches HBMFlashAttention: the S×S matrix never reaches HBM
Diagram 19. Naive attention moves the S×S matrix through HBM four times; FlashAttention keeps blocks in SRAM, stitches them with online softmax (diagram 10), and writes only the output to HBM.
  • 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 via tl.arange, tl.load/tl.store with a mask on the tail, tl.dot for a block matmul. The compiler handles threads and shared memory

Common mistakes

  • Not rescaling the old d and o on a new maximum: test_online_softmax_matches_direct_computation, test_online_softmax_block_size_does_not_matter
  • Starting with m = 0 or not handling m = −inf: on a fully masked block e^{−inf − (−inf)} gives NaN. The reference implementation checks torch.isfinite(m); large logits are checked by test_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 a WARMUP (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×S matrix.
  • I can run attention_bench.py and explain why the memory saving is exactly S/H.
  • I can read a simple Triton kernel and say what each program does.

Self-check

  1. Why is attention memory-bound while a large matmul is compute-bound?
  2. What does the FlashAttention backward recompute, and why is recomputing cheaper than storing?
  3. What resource does kernel fusion save, and why does that matter more than FLOPs for elementwise operations?
  4. What is the critical batch of an H100 in bf16 and with int8 weights, and why is it counted in tokens?

In the app each week has skills to rate yourself on, questions with answer checking, Python coding problems and a tutor grounded in the course.

Learn in the app: tutor, coding problems
← PreviousWeek 12. Sampling strategies Next →Week 14. Scaling laws, precision, parallelism

Snula
Snula: LLMs from scratch

  • Home
  • Curriculum
  • App
  • Privacy
  • Terms

The course text is licensed under CC BY-NC-SA 4.0, nanolm code under Apache-2.0.