CS336 // FIELD MAP
← field map
LECTURE 05 · MAKE IT FASTTatsunori Hashimoto · 2025-04-15 · 74 min

GPUs

Stanford CS336 · Spring 2025 · lecture 05 of 17

Transcript: cleaned auto-captions with timestamps

TL;DR — A GPU is not a fast calculator; it is a very wide calculator sitting a long way from its data, and after twenty years of compute scaling roughly a thousand times faster than memory scaling, essentially every kernel you write is bottlenecked on moving bytes rather than on doing arithmetic. Tatsu builds the whole lecture around one plot — square-matmul throughput as a function of matrix size, which looks like unpredictable noise — and by the end every wiggle in it has a name: arithmetic intensity, tile alignment against DRAM burst boundaries, and wave quantization against the GPU's 108 streaming multiprocessors. The toolkit that follows is five moves, all the same move: use fewer bits, fuse operators, recompute instead of storing, coalesce your reads, and tile into shared memory. The one thing to remember: optimize data movement, not FLOPs. FlashAttention is the closing worked example — it is tiling plus an online softmax plus recomputation, three things from this lecture stacked, and nothing else.

Lectures 1–4 got a Transformer defined and costed on paper: tokenizer, PyTorch primitives, architecture choices, mixture-of-experts routing. All of that reasoning was in units of parameters and FLOPs. This lecture is where the course admits that FLOPs are the wrong currency. It opens the "make it fast" arc by replacing the paper cost model with the real one — a memory hierarchy, a SIMT execution model, and a roofline — so that Assignment 2 can ask you to write a Triton FlashAttention-2 kernel and have that be a reasonable thing to ask. Lecture 6 then hands you the tool (Triton); this lecture hands you the reasons.

Outline, with timestamps

One plot, held over your head for an hour

The lecture is structured as a whodunit. In the first two minutes Tatsu puts up a chart of achieved throughput for square matrix multiplies against matrix size, and it looks broken: not a smooth climb but four separate bands, each one rippling up and down, with cliffs where adding a single element to a dimension costs you a quarter of your performance (25:09). He promises that by the end you will find it boring — and that promise is the pedagogical spine: every hardware fact he introduces is one the plot eventually cashes out.

"So my goal today is to try to make CUDA and GPUs less magic."— Tatsunori Hashimoto, 00:37

The framing before that is the scaling argument the course keeps returning to (03:25): more compute reliably buys better models, so the interesting question is where compute comes from. Until roughly 2000 it came from Dennard scaling — smaller transistors, higher clocks, lower power per gate, all for free (04:35). That ended. Transistor counts kept climbing but single-thread performance flattened, which means every subsequent gain has had to be a parallel gain. Bill Dally's keynote chart of NVIDIA integer throughput from the K20 era to the H100 shows the replacement curve: something like a thousandfold in a decade (05:44). The slide's blunt version is that there is no LLM scaling without GPU scaling — which is why a language-modeling course spends a lecture on cache hierarchies.

The machine: wide, dumb, and far from its data

The CPU/GPU contrast is best held as a chip-area argument rather than a speed argument (06:18). A CPU spends most of its silicon on control: branch prediction, out-of-order machinery, deep caches — everything needed to make one instruction stream finish fast. A GPU spends almost all of it on arithmetic units with a thin ribbon of control logic driving them. The design goals follow: CPUs minimize latency per task, GPUs maximize total throughput, and a GPU will happily make every individual task slower if that finishes the batch sooner (07:28).

Concretely there are three nested units you must hold in your head (13:29). A streaming multiprocessor (SM) is the autonomous worker; an A100 has 108 of them enabled, out of 128 on the full GA100 die — Tatsu quotes both numbers at different points, and 108 is the one that matters for the punchline later. A block of threads is assigned to exactly one SM and gets that SM's shared memory. A warp is 32 consecutively numbered threads inside a block that issue the same instruction in the same cycle on different data (14:01). The warp is the unit that will keep reappearing, because it is simultaneously the unit of control (all 32 must agree on the instruction) and the unit of memory access (all 32 request at once).

The memory table is the other half, and it is a table about physical distance (10:45). Registers and shared memory / L1 live inside the SM: roughly 20 clock cycles. L2 is on-die but outside the SM. Global memory is DRAM chips physically beside the package, reached through the HBM connectors you can see at the edge of a die shot: roughly 200–300 cycles (12:24). SRAM is about 8x faster than DRAM and about 100x more expensive per byte, which is exactly why there is so little of it. And crucially, anything that has to travel between blocks has no choice but to round-trip through global memory. So the whole game is: get a small working set into shared memory, do as much arithmetic on it as possible, and leave.

The TPU aside (16:45) is worth two minutes of your attention precisely because it is boring: a TPU tensor core is an SM with the same skeleton — a scalar control unit, a vector unit, a big dedicated matmul unit (the MXU), fast on-core memory, slow HBM outside. Fewer, bigger cores instead of many small ones, and no warp abstraction. Everything in this lecture except the specific warp mechanics transfers.

Why matmuls are blessed and bytes are cursed

Two asymmetries drive every optimization that follows. First, matrix multiplication is a privileged operation. Researchers were bending graphics pipelines into matmul engines back in 2001 (20:37), but since the V100 introduced tensor cores, matmul FLOPs and non-matmul FLOPs have diverged by more than an order of magnitude on the same chip (21:14). That is an architecture constraint, not just a performance note: it retroactively justifies the design pressure in lectures 3 and 4 toward putting your parameters and your compute inside matmuls.

Second, and more important, compute and memory have scaled at wildly different rates (22:19). Over the generations the slide covers, host interconnect improved maybe an order of magnitude, global memory bandwidth roughly 100x (GDDR through HBM2E), and matmul throughput something like 100,000x. Those are log-scale gaps that only widen, because DRAM is genuinely hard to scale.

"Your bottlenecks are probably going to end up being memory, because the memory is not growing as fast."— Tatsunori Hashimoto, 23:28

The roofline model formalizes this (26:48). Plot achieved FLOP/s against arithmetic intensity — FLOPs performed per byte moved — and you get a diagonal segment where you are bandwidth-limited and a flat ceiling where you are compute-limited. Every trick in part 2 is a way of pushing a kernel rightward along that diagonal until it hits the roof. The left-hand rise of the mystery plot is nothing more exotic than small matmuls being memory-bound: below roughly 1536 there simply is not enough arithmetic per loaded byte to keep the tensor cores fed (58:06).

The five moves (and one hazard that isn't about memory)

The hazard first, because it is the exception. Under SIMT, a conditional inside a warp does not branch — it serializes (28:28). If half the threads take the if and half the else, the hardware runs the first path with the other threads asleep, then swaps. You pay for both branches. This is control divergence, and the practical rule is simply not to put data-dependent conditionals inside a warp's inner loop. Everything else on the list is memory.

Carry this away: arithmetic intensity is the number to compute before you optimize anything. FLOPs per byte moved, on the specific kernel, at the specific shape you actually run. If that number puts you on the diagonal of the roofline, no amount of clever math will help — you need fusion, tiling, or fewer bits. If it puts you on the roof, stop; you are done, and the remaining wins are in scheduling, not in the kernel.

Reading the mystery plot

Tiling is also where performance gets genuinely weird, and the two failure modes are what the plot has been hiding. The first is tile quantization. With a 128-wide tile, a 256-wide matrix is two clean tiles; a 257-wide matrix needs three, and the third is almost empty (53:35). Since each tile becomes a block assigned to an SM, that nearly-empty tile occupies a whole SM to do almost nothing. Tile size is therefore a three-way constraint: big enough to amortize, small enough to fit shared memory, and divisible into your actual dimensions.

The second is burst alignment (55:53). If a tile's rows begin exactly on DRAM burst boundaries, one read per row loads it. Add one element to the leading dimension and every row now straddles two burst sections, so every row costs two reads — you doubled global traffic by changing a dimension by one. This is where the divisibility bands in the plot come from: colour the points by the largest power of two dividing the matrix size and they sort into layers, with dimensions divisible by 32 at the top and prime dimensions at the bottom (59:49). Hence the standing advice: never pick a prime for a tensor dimension (60:22). The famous instance is Andrej Karpathy's nanoGPT result, quoted on the slide — padding the vocabulary from 50257 to 50304, the nearest multiple of 64, bought about a 25% end-to-end speedup for 47 dimensions of pure waste (57:34).

That still leaves the cliffs within a band, and the answer is the nicest arithmetic in the lecture: wave quantization (62:01). Take a 256×128 tile — a natural choice, since the matmul units themselves like operands around 128. At size 1792 you get 7 × 14 = 98 tiles. At 1793 you round up to 8 × 15 = 120. An A100 has 108 SMs. Ninety-eight tiles dispatch as a single wave with every SM busy; 120 tiles dispatch as 108 followed by a second wave of 12, during which 96 of 108 SMs sit idle waiting. You paid for two full waves and used a bit over half the second one. The lesson is not "avoid 1793" — it is that you want your tile count either comfortably below the SM count or many multiples above it, never barely over (62:35).

FlashAttention, disassembled

The last twelve minutes are the payoff: FlashAttention contains no new hardware idea at all, only three from this lecture composed correctly. Start with what it actually claims (64:44) — the paper's own framing is tiling plus recomputation to compute exact attention with sub-quadratic HBM accesses.

"So it's not subquadratic computation because you can't do that."— Tatsunori Hashimoto, 65:16

That distinction is the whole point. The FLOPs stay quadratic in sequence length; what becomes sub-quadratic is global memory traffic, which is the resource that was actually scarce. Attention is two matmuls with a softmax between them, and the matmuls tile exactly like any other matmul — Tatsu's observation is that Figure 1 of the FlashAttention paper is literally a tiled matmul diagram, blocks of K and Q copied into SRAM, multiplied, accumulated (66:24).

The softmax is the obstruction, because it is a global operation over each row: you cannot normalize until you have seen every score in the row, which seems to force materializing the full n×n matrix (67:29). The escape is the online softmax of Milakov and Gimelshein (2018) (68:01). Carry two running scalars per row — the max seen so far, m, and the running sum of exponentials, d. When a new tile arrives with a larger max, rescale the accumulated d by exp(m_old − m_new) and keep going. The correction telescopes, so the value you hold after the last tile is exactly the numerically-stable normalizer you would have computed in one pass. Softmax becomes streamable, and therefore tileable.

With that, the forward pass falls out (70:20): tile the QKᵀ product, fuse the exponential into the same kernel rather than round-tripping, maintain the running max and normalizer per tile, multiply by V tile-wise, and never write an n×n array to HBM. A student asks the sharp question — you cannot emit a normalized output until every tile has been seen — and the answer is that one pass suffices, because by its end the running normalizer is already sitting in shared memory. The backward pass, which the lecture skips, closes the loop with trick 3: recompute the n×n scores tile by tile from Q, K, V rather than storing them (72:10).

The closing recap is one sentence worth internalizing (73:16): when you are optimizing, do not ask how to reduce FLOPs — ask how to reduce movement between HBM and the SM. Every technique in this lecture, and most of what you will write in Assignment 2, is an answer to that question.

What you build with this

This lecture opens Assignment 2: Systems (repo, spring2025 tag · handout PDF · leaderboard), which builds a benchmarking and profiling harness, a Triton FlashAttention-2 kernel, distributed data-parallel training, and optimizer state sharding. The first half of the handout is this lecture made executable, in this order:

Lecture 6 (Kernels, Triton) supplies the programming model for that second half; this lecture supplies the reasons any of it is worth doing. The later parts of the assignment (DDP, bucketed gradient communication, optimizer state sharding) belong to lecture 7's parallelism material. Note that the repo's main branch has since moved to the Spring 2026 revision — the spring2025 tag above is the version this lecture was delivered against.

Supporting materials, verified

Exercises

  1. Reproduce the mystery plot code — On any CUDA device, time torch.matmul for square FP16 matrices at every size from 1024 to 2048. (1) Warm up, then time 50 iterations per size with torch.cuda.synchronize() around the loop. (2) Convert to achieved TFLOP/s using 2·N³ FLOPs. (3) Colour each point by the largest power of two dividing N. (4) Mark the size where your device's SM count times a plausible tile count changes wave. A good answer shows the divisibility bands separating cleanly, names at least one cliff, and computes the tile count on both sides of it against your GPU's SM count from torch.cuda.get_device_properties.
  2. Fusion, measured code — Write f(x) = sin(x)**2 + cos(x)**2 over a tensor large enough to exceed L2. (1) Time it eagerly. (2) Time torch.compile(f) after warm-up. (3) Compute the bytes each version must move, assuming eager materializes every intermediate. (4) Compare the measured speedup to the ratio of predicted bytes. A good answer explains any gap — kernel launch overhead, L2 hits on small tensors, or the compiler doing more than you modelled.
  3. Arithmetic intensity of your own attention — Without writing a kernel, work out bytes moved and FLOPs performed for one forward pass of naive scaled-dot-product attention at batch 8, one head, d = 64, and sequence lengths 1024 / 4096 / 16384, assuming the n×n score matrix is written to and read from HBM. Do the same assuming it never leaves SRAM. A good answer gives the two intensity curves, identifies where each crosses the FLOP/byte ratio of a real GPU (an A100 is roughly 312 TFLOP/s BF16 against about 2 TB/s of HBM bandwidth), and states which regime each sequence length lands in.
  4. The online softmax, by hand — Implement the running-max/running-normalizer recursion in twenty lines of NumPy over a stream of chunks, and verify it matches scipy.special.softmax on a vector containing values around 800 (where a naive exp overflows). A good answer states the invariant maintained after each chunk and shows why the exp(m_old − m_new) rescale keeps it exact rather than approximate.
Next: L06 Kernels, Triton · Back to the map.