GPUs
Transcript: cleaned auto-captions with timestamps
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
- 00:04 — Goals, and why hardware decides the ceiling: make CUDA less magic; Dennard scaling ended, so all remaining headroom is parallel.
- 06:18 — CPU vs GPU: one is built to finish a task quickly, the other to finish all tasks quickly; the chip area goes to different places.
- 08:33 — Anatomy: SMs, SPs, and how far away the memory is — 20 clock cycles to shared memory, 200–300 to global.
- 13:29 — Execution and memory model: blocks land on SMs, threads execute in warps of 32, anything crossing blocks goes through global memory.
- 16:45 — TPU aside: same skeleton (small control, big matmul unit, fast local memory), no warps; and why SIMT scales.
- 20:37 — Matmuls are blessed: tensor cores opened a 10x-plus gap over non-matmul FLOPs, while memory bandwidth barely moved.
- 25:09 — The mystery plot, and the roofline model that explains its left half.
- 27:52 — Control divergence (the one non-memory hazard), then trick 1: low and mixed precision.
- 33:02 — Trick 2: operator fusion — the factory-and-warehouse picture, and what torch.compile does for free.
- 37:33 — Trick 3: recomputation — throw activations away and recompute them; 8 memory accesses become 5.
- 41:23 — Trick 4: DRAM burst mode, and why a warp's access pattern is worth 4x.
- 46:55 — Trick 5: tiling — a factor-T reduction in global reads — and the two ways it goes wrong.
- 58:06 — The mystery solved: divisibility bands, and wave quantization at 1792 → 1793.
- 64:44 — FlashAttention as the worked example: tiling + online softmax + recomputation; then the recap.
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.
- 1. Low precision — Fewer bits per element is directly fewer bytes to move. Tatsu's worked example is elementwise ReLU: in FP32 you read x and write the result, so 8 bytes moved per FLOP; in FP16 the same operation costs 4 bytes per FLOP (31:17). You just doubled your effective bandwidth for free. The catch is that precision is not uniform across a network — the standard mixed-precision matmul takes 16-bit inputs but accumulates partial sums in FP32 inside the tensor core, and operations needing dynamic range (exponentials, normalizers) want BF16 or FP32 (31:51).
- 2. Operator fusion — Horace He's factory picture: memory is the warehouse, the SM is the factory, and the conveyor between them is the bottleneck (33:37). A naive chain of pointwise ops ships each intermediate back to the warehouse and fetches it again. Writing sin(x)**2 + cos(x)**2 in PyTorch launches five kernels and four needless round trips; fused, it is one kernel and the data never leaves. This class of fusion is mechanical enough that torch.compile finds it automatically, which is Tatsu's actual recommendation — use it everywhere (36:27).
- 3. Recomputation — Three stacked sigmoids, counted honestly: the forward pass reads x and writes S1, S2, out; the backward pass reads S1, S2 and the incoming gradient and writes dx. Eight global accesses, and essentially zero arithmetic intensity because there is no matmul anywhere. Drop the stored activations and recompute them inside the SM during backward and the count falls to five — 5/8 of the traffic for identical results (40:17). The mechanism is the same as gradient checkpointing, but the motivation is inverted: not "I ran out of memory" but "I had idle ALUs and no bandwidth, so I spent the cheap resource."
- 4. Coalescing — DRAM does not serve single bytes. Moving a row to the sense amplifier is the slow step, so once it is there you get the whole burst section — ask for element 0 and you are handed 0, 1, 2, 3 (42:31). Since a warp's 32 threads issue their loads together, whether they land inside one burst section or scatter across 32 of them changes your effective bandwidth by up to 4x on the slide's example (44:10). For a row-major matrix, threads striding along a row are not coalesced; threads whose lane index walks the contiguous dimension are. Tatsu notes he had to stare at this diagram before believing the direction (45:16) — it is a common place to get it backwards.
- 5. Tiling — The big one. A naive matmul reads every input element N times from global memory, once per output it contributes to (51:23). Cut both operands into T×T tiles that fit in shared memory, and the loop restructures into phases: load one M tile and one N tile, accumulate every partial sum they support, move on. Each input is now read N/T times from global and T times from shared — a factor of T reduction in the traffic that actually costs you (52:27). And because you control the load order within a tile, you can make those loads coalesced as a bonus.
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:
- §1.1 profiling and benchmarking — an end-to-end timing script over five model sizes (small through 2.7B), an Nsight Systems compute profile, and a memory profile. This is the roofline discipline: measure where the time and the bytes actually go before optimizing anything.
- §1.1.5 mixed precision — the mixed_precision_accumulation problem makes you watch FP16 accumulation lose to FP32 accumulation numerically, then benchmarking_mixed_precision has you run the model under torch.autocast and reason about which layers (notably layer norm) get kept in higher precision. That is trick 1, with the caveats.
- §1.2–1.3 attention benchmarking and torch.compile — sweep head dimension against sequence length until naive attention OOMs, account for where the memory went, then compare against the compiled version. That is trick 2, and the point where the n² score matrix becomes visibly the problem.
- §1.3.1 onwards, the Triton kernels — a warm-up weighted-sum kernel to learn tiles and block pointers, then flash_forward (15 points), flash_backward, and flash_benchmarking. This is tricks 3, 4 and 5 together, and it is exactly the disassembly in the last section of this lecture, rebuilt.
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
- 2025 Lecture 5 — GPUs (slides) — Tatsunori Hashimoto (2025) · The 51 slides this page is checked against; the wave-quantization arithmetic and the ReLU intensity numbers are on them verbatim.
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré (2022) · The paper the whole third act unpacks; Figure 1 is the tiled-matmul diagram Tatsu points at.
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning — Tri Dao (2023) · The version Assignment 2 asks you to implement in Triton; better work partitioning across warps and thread blocks.
- Online normalizer calculation for softmax — Maxim Milakov, Natalia Gimelshein (2018) · The telescoping running-max trick that makes softmax streamable, and therefore makes FlashAttention possible. (The slide spells the first author "Mikailov".)
- Making Deep Learning Go Brrrr From First Principles — Horace He (2022) · Source of the factory/warehouse mental model and the compute-bound vs memory-bound vs overhead-bound taxonomy.
- What Shapes Do Matrix Multiplications Like? — Horace He (2023) · The mystery plot itself, plus the divisibility-band colouring and the wave-quantization walkthrough.
- Min-cut optimal recomputation with AOTAutograd — Horace He, PyTorch dev-discuss (2022) · Where the stacked-sigmoid 8-accesses-to-5 example comes from, and how the compiler chooses what to recompute automatically.
- Matrix Multiplication Background User's Guide — NVIDIA · The vendor's own writeup of tile quantization and wave quantization, with the same plots in more detail.
- AI and Memory Wall — Amir Gholami, Zhewei Yao, Sehoon Kim, Coleman Hooper, Michael W. Mahoney, Kurt Keutzer (2024) · The compute-outruns-bandwidth chart on the slides, with the numbers behind it.
- Scaling Laws for Neural Language Models — Jared Kaplan et al. (2020) · The "compute buys capability" premise the lecture opens on, and the reason hardware efficiency is a modeling concern.
- Fast Matrix Multiplies Using Graphics Hardware — E. Scott Larsen, David McAllister, SC '01 (2001) · The pre-CUDA paper Tatsu cites for researchers hacking texture buffers into matmul engines. (The ACM page blocks automated fetches but resolves in a browser.)
- Mixed Precision Training — Paulius Micikevicius et al. (2018) · Field map extra. The canonical treatment of FP16 storage with FP32 accumulation and loss scaling — the details behind the "not every layer goes low" caveat.
- How to Scale Your Model — Jacob Austin et al., Google DeepMind (2025) · Field map extra. The "nice TPU book" Tatsu credits; its GPU chapter re-derives these rooflines for NVIDIA parts.
- How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance — Simon Boehm (2022) · Field map extra. Ten successive kernels from naive to near-cuBLAS, each step being one idea from this lecture; the best way to feel coalescing and tiling rather than just hear about them.
- GPU MODE lectures — the GPU MODE (formerly CUDA MODE) reading group · Field map extra. The community resource Tatsu credits at the top; recordings and code for kernel-level topics this lecture only samples.
Exercises
- 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.
- 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.
- 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.
- 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.