CS336 // FIELD MAP
← field map
LECTURE 06 · MAKE IT FASTTatsunori Hashimoto · 2025-04-17 · 80 min

Kernels, Triton

Stanford CS336 · Spring 2025 · lecture 6 of 17

Transcript: cleaned auto-captions with timestamps

TL;DR — Lecture 5 gave you the roofline; this one gives you the instruments and the hands. Half the hour is a discipline lesson — benchmark before you guess, and profile before you optimize — with the two traps that make GPU timings lie (no warmup, no torch.cuda.synchronize()). The other half writes the same GeLU five ways — naive PyTorch, PyTorch's fused kernel, hand-written CUDA, Triton, and torch.compile — and prices each one. The punchline is not "write CUDA": it is that fusion buys you an order of magnitude (8.1 ms → 1.1 ms), that the hand-written CUDA kernel is the slowest of the four fast options because it does one scalar element per thread, and that torch.compile, invoked with one line, gets you most of the way there.

The course's systems arc has a shape: lecture 5 tells you what the hardware can do, lecture 6 tells you how to find out what it is actually doing, and lecture 7 leaves one GPU behind. This is the lecture where the abstraction "PyTorch op" gets opened, twice — once down to the CUDA kernel that ATen dispatches to, and once further down to the PTX the Triton compiler emits. It exists because assignment 2 asks you to write a FlashAttention-2 kernel in Triton and to justify it with a profile, and neither of those is possible if a @ b is still an opaque box to you.

Outline, with timestamps

Two instruments, and what each one can tell you

Benchmarking answers "how long did the whole thing take". Profiling answers "which of the things inside it took the time". The lecture is emphatic that you need both and that you need them before you form a hypothesis, because the standard failure mode is spending an afternoon optimizing something that was never on the critical path (07:09).

"If there's one high-level thing to remember, it's if you want to write high performance code, you should remember to benchmark and profile your code."— Tatsunori Hashimoto, 07:09

The benchmark harness in the lecture is deliberately twelve lines rather than torch.utils.benchmark, so that the two things that matter stay visible. Warmup: the first call into a PyTorch op pays for JIT compilation, module loading and kernel-cache population, so you throw the first iteration away and measure steady state. Synchronize: the CPU dispatches CUDA kernels and returns immediately, so time.time() around a GPU call measures how long it took to enqueue the work, not to do it. Skip the sync and your 16384×16384 matmul appears to finish instantly. The harness runs one warmup, then three trials, syncing inside each trial, and averages — the averaging is there because thermal state and clock behaviour make single measurements noisy.

The profiler is torch.profiler with both CPU and CUDA activities on, sorted by cuda_time_total. What it buys you is the layer below the API: call a + b and the table shows aten::add (the C++ dispatch layer), then the actual kernel vectorized_elementwise_kernel<4, at::native::CUDAFunctor_add…>, then cudaLaunchKernel and cudaDeviceSynchronize as separate cost lines. The kernel name is itself diagnostic. A 2048×2048 fp32 matmul dispatches to cutlass_80_simt_sgemm_256x128_8x4_nn_align1 — CUTLASS is NVIDIA's templated linear-algebra library, simt means CUDA cores rather than tensor cores, and 256x128 is the tile size. Drop to 128×128 and PyTorch dispatches somewhere else entirely (an xmma GEMM). Same Python line, different machine code, different performance profile.

What the scaling curves actually said

Square fp32 matmul on an H100, wall clock, from the recorded trace:

dim102420484096819216384
ms0.870.843.1521.7162.9

Two regimes. Below 2048 the curve is flat: the work is too small to cover kernel launch and synchronization overhead, so you are measuring the harness, not the hardware. Above 4096 each doubling costs roughly 7–7.5×, converging on the 8× that n³ predicts. And the top of the curve is a sanity check worth doing yourself: 2·16384³ = 8.8 TFLOP in 162.9 ms is 54 TFLOP/s, against an H100 SXM fp32 non-tensor-core peak near 67 TFLOP/s. About 81% of roofline, which is exactly what you would expect from a simt sgemm kernel that is not using tensor cores at all.

The MLP sweep is the more interesting half, and the lecture skipped part of it live for time (17:37) — but the recorded trace has all four sweeps. Baseline is dim 256, 4 layers, batch 256, 2 steps: 6.29 ms. Scaling steps ×2…×5 gives 11.3 / 16.8 / 22.0 / 27.3 ms and scaling layers gives 9.7 / 13.3 / 16.9 / 20.4 ms — both clean straight lines, as everyone predicts. But scaling batch size gives 6.09 / 5.94 / 6.08 / 5.96 ms and scaling dimension gives 6.07 / 6.05 / 6.01 / 6.09 ms. Flat. Dead flat, through a 5× increase in FLOPs.

That is the whole lesson of the first half in one table. At these sizes the model is not compute-bound; it is bound by the fixed cost of launching a few hundred tiny kernels, and the GPU absorbs 5× more arithmetic for free. Steps and layers scale because they multiply the number of launches. Batch and dim scale the work inside each launch, which was never the constraint. If you had reasoned about this from FLOPs alone — the thing lecture 2 taught you to do — you would have predicted a 5× slowdown and been wrong by 5×.

The CPU is a whole step ahead of the GPU

The PyTorch profiler bottoms out at aggregate tables. To see ordering, the lecture switches to NVIDIA Nsight Systems, with the code annotated using NVTX ranges so each training step shows up as a labelled band (32:28). The timeline has two rows: CUDA hardware on top, CPU threads below. They do not line up, and that is the point. When the CPU thread is inside layer 1, the GPU is still finishing layer 1 from several layers back — at one point in the trace the CPU is at layer 9 while the GPU executes layer 1, and at the step level the CPU is a full forward-and-backward pass ahead.

This is not a bug, it is the execution model: the CPU enqueues kernels into a queue of bounded depth and only blocks when the queue fills. Which yields the observation that quietly justifies the entire course being written in Python:

"It doesn't matter that we're programming in Python and Python's not a very high performance language, right? Because the CPU is never the bottleneck, because the CPU can run ahead and sort of queue commands into the GPU."— Tatsunori Hashimoto, 43:20

It also explains a bug you will write. Add print(loss.item()) inside the training loop and the CPU can no longer run ahead: printing requires the value, the value requires the backward pass to finish, so a cudaStreamSynchronize appears in the timeline and the CPU sits idle until the GPU catches up (38:59). In the profiled example the GPU stayed near full utilization anyway, because one sync per step is cheap relative to the step — but the lecture is careful to say that logging heavily, or per-microbatch, turns this into a genuine CPU bottleneck. Anything that pulls a tensor to the host — .item(), .cpu(), a Python if on a loss value — is a barrier.

Fusion, priced

The mental model, borrowed from Horace He's "Making Deep Learning Go Brrrr", is a warehouse (DRAM) and a factory (SRAM/registers). Every elementwise op is a round trip: ship the data in, do a trivial amount of arithmetic, ship it back. Ten cheap ops in sequence pay ten shipping costs. Fusing them pays one.

The demonstration uses the tanh approximation to GeLU, 0.5·x·(1 + tanh(0.79788456·(x + 0.044715·x³))), written out as a chain of PyTorch ops on a 16384×16384 fp32 tensor, versus F.gelu(x, approximate="tanh"). Same numbers, checked. Very different cost: 8.1 ms versus 1.1 ms wall clock. In the profiler the naive version fires nine kernel launches — three multiply kernels, two adds, a tanh, and friends — for 7.669 ms of GPU time; the PyTorch version fires exactly one, GeluCUDAKernelImpl, for 701.6 µs.

Do the bandwidth arithmetic and the fused number stops looking like a speedup and starts looking like a ceiling. A 16384² fp32 tensor is 1.074 GB. One fused pass reads it and writes it: 2.147 GB in 701.6 µs is 3.06 TB/s, against roughly 3.35 TB/s of HBM3 bandwidth on an H100 SXM. That kernel is at about 91% of the memory roofline. There is no meaningful performance left in it — which is precisely lecture 5's claim that everything except matmul is memory-bound, now with a receipt.

Five ways to write GeLU, and what each one costs

The rest of the lecture writes the same function four more ways. In CUDA (48:18) it is about thirty lines: a wrapper that runs on the CPU, asserts the tensor is on the device and contiguous, allocates the output with empty_like rather than zeros_like (you are going to overwrite every element, so do not pay to zero it), computes num_blocks = cdiv(num_elements, 1024) and launches; plus a __global__ kernel that recovers its own index as blockIdx.x * blockDim.x + threadIdx.x, guards it with if (i < num_elements) because the last block overhangs the array, and writes one element. In Triton (64:10) the wrapper is nearly identical, but the kernel is written per block, not per thread: the index becomes a vector, offsets = pid*BLOCK_SIZE + tl.arange(0, BLOCK_SIZE), the bounds check becomes a mask passed to tl.load and tl.store, and the compiler handles memory coalescing, shared memory, and scheduling within an SM. Scheduling across SMs stays yours. And torch.compile(manual_gelu) is one line that generates a fused Triton kernel from the naive Python.

implementationwall clockGPU kernel timeeffective bandwidth
manual PyTorch chain8.10 ms7.669 ms (9 kernels)—
hand-written CUDA1.82 ms1.664 ms1.29 TB/s (~39%)
torch.compile1.47 ms——
PyTorch F.gelu1.11 ms701.6 µs3.06 TB/s (~91%)
Triton1.73 ms705.2 µs3.05 TB/s (~91%)

Wall clocks are the single five-way sweep in pytorch_compilation(); kernel times are from the corresponding profile runs; bandwidths are computed from 2 × 1.074 GB of traffic.

The lecture reads the wall-clock column and concludes that Triton bought ergonomics rather than speed (71:22). The kernel-time column tells a sharper story that is worth pulling out. The hand-written CUDA kernel is 2.4× slower on the GPU than either PyTorch's or Triton's. Not because C++ is slow — because it is naive in a specific way: one element per thread, scalar loads. The PTX walkthrough (68:02) shows what Triton did instead. Its generated assembly issues vectorized ld.global instructions pulling four floats into four registers at a time, and does two of them, so each thread processes eight elements — the transformation called thread coarsening. That is the whole 2.4×. Writing CUDA got you fusion, which was the big win; it did not get you the memory-access pattern that the compiler applies for free.

Read the wall-clock and kernel-time columns together and you also learn to distrust either one alone. Triton's kernel is a hair faster than PyTorch's, yet its end-to-end call is 0.6 ms slower, because Triton's Python-side launch path (grid computation, argument binding, JIT cache lookup) costs more CPU than ATen's C++ dispatch. At 268M elements that overhead is still visible. Which number is "true" depends on whether your kernel is called once or ten thousand times.

Softmax: the first kernel with a reduction in it

GeLU is elementwise, so every thread is independent and nothing has to be shared. Softmax is where it gets real: each output needs the max and the sum of its whole row. The design the lecture picks (75:16) is the simplest one that works — one block per row. Grid size is the number of rows; BLOCK_SIZE = triton.next_power_of_2(num_cols) so an entire row fits in one block's registers; you load the row with other=-inf for the padding, subtract the max, exponentiate, sum, divide, store. Because the whole reduction happens inside one block, the Triton code reads like NumPy. That constraint is also the design's limit: it assumes a row fits, which is why the assignment's FlashAttention kernel has to tile instead.

The read/write accounting in the lecture script is the part to internalize. The naive PyTorch softmax touches memory 5MN + M times for reads and 3MN + 2M for writes — eight passes over the matrix where two would do, so about a 4× speedup is available on paper. Measured, on a 16384² tensor: naive 3.258 ms of GPU time, Triton 705.5 µs. That is 4.6×, and the paper estimate predicted it. Note also that the Triton kernel (705.5 µs) beat PyTorch's own fused softmax kernel (1.137 ms) by 1.6×, and torch.compile (730.8 µs) landed between them. A hundred lines of Python beat the vendor library here — and both of them are within a few percent of the 641 µs that pure bandwidth would allow.

Carry this away: the order of operations is measure, fuse, then hand-tune, and usually stop after fusing. Fusion is where the order of magnitude lives (8.1 ms → 1.1 ms), and torch.compile gets most of it for one line of code. Hand-writing a kernel bought a 2.4× regression against the compiler here, because the compiler applies vectorized loads and thread coarsening that a naive kernel does not. Reach for Triton when you have a fused operation the compiler cannot discover — an attention variant, a custom norm, a new architecture that isn't hitting utilization — not because writing kernels feels productive.
"You shouldn't go home and say … I'm going to write CUDA kernels for every single part of my language model, you know, that's probably not a good use of your time. But if you're writing a new architecture with some complicated piece and you're not getting utilization, but you think you can, that's maybe the time to really bust out the Triton."— Tatsunori Hashimoto, 75:16

Where he hedges, and two things to watch

The lecturer hedges in three useful places. Asked why CPU time drops going from add to matmul in the profile, he says plainly he does not know (24:17). Asked whether block size matters for the GeLU kernel, he guesses it will not past roughly 1024 for a purely elementwise kernel but flags SM saturation and per-block work as the two things that would decide it (60:47). And on beating torch.compile: simple operator fusion and matmul kernel selection are settled in the compiler's favour, but FlashAttention-2 and especially -3 involve hardware-specific choices that a JIT would not have found on its own — "we know in hindsight that those are the right optimizations" (74:12). The lecture script adds a fourth hedge the video does not: timings are not reliably predictable because CUDA kernels, libraries and hardware are non-homogeneous.

Two things to watch. First, at 04:27 a group of 32 threads is called a wave; the standard term is a warp, and he uses it correctly minutes later in the Q&A. "Wave" properly means a scheduling round of thread blocks across SMs — hence wave quantization, the reason the lecture script's rule of thumb is to keep the number of thread blocks at four times the number of SMs or more, so the tail wave doesn't leave SMs idle. Second, throughout the benchmarking sections the spoken numbers are said as "8.1 seconds" and "1.1 seconds"; they are milliseconds everywhere.

What you build with this

This lecture feeds assignment 2 (systems), which opened two days earlier with lecture 5 — the Spring 2025 handout is v1.0.4. Concretely, the first half of the lecture is problem benchmarking_script (4 pts — warmup, timing, forward/backward on the five model sizes) and nsys_profile (5 pts — NVTX-annotated Nsight traces, exactly the workflow at 32:28), plus memory_profiling (4 pts). The second half is torch_compile (2 pts) and then the big one: flash_forward (15 pts) asks for FlashAttention-2's forward pass first in pure PyTorch as a debugging reference and then as a Triton kernel, flash_backward (5 pts) adds the backward, and flash_benchmarking (5 pts) sweeps sequence lengths from 128 to 65536 in bf16 and fp32 using triton.testing.do_bench. The A2 leaderboard times your FlashAttention-2 at batch 1, sequence 16384, d_model 1024, 16 heads, causal, bf16 — and Triton only, no CUDA. The parallelism half of the assignment belongs to lecture 7.

One thing the lecture did not cover but the executable script contains: triton_matmul_main() is written but not called from main(). It walks tiling, shared memory, and the grouped-vs-row-major block ordering that improves L2 reuse (90 block loads versus 54). If you are heading for the FlashAttention kernel, read it — tiling is the idea the softmax kernel deliberately avoided.

Supporting materials, verified

Exercises

  1. Break your own benchmark code — Time a 8192×8192 matmul three ways: (1) no warmup, no sync; (2) warmup, no sync; (3) warmup and torch.cuda.synchronize() inside the timed region. Then run all three under CUDA_LAUNCH_BLOCKING=1 and compare again. A good answer reports the three numbers, explains why version 1 can report sub-millisecond times for tens of milliseconds of work, and states which of the two omissions costs you more.
  2. Close the CUDA kernel's 2.4× gap code — Start from gelu.cu. (1) Reproduce the baseline: kernel time on a 16384² fp32 tensor, and convert it to effective GB/s. (2) Rewrite the kernel to process four elements per thread using a float4 load, adjusting the grid accordingly. (3) Measure again. (4) Try eight per thread. (5) Dump the PTX for both and confirm you now see vectorized ld.global.v4 instructions. A good answer reaches roughly 3 TB/s — PyTorch's and Triton's number — and says which change mattered, the vectorized load or the reduced launch count.
  3. Find the synchronization point code — Take any small training loop, annotate the steps with torch.cuda.nvtx.range_push/pop, and profile with nsys. Run it once clean and once with a print(loss.item()) inside the loop. A good answer shows the two timelines, identifies the cudaStreamSynchronize band, quantifies how far ahead the CPU ran in the clean version, and names one other line of ordinary training code that would create the same barrier.
  4. Predict before you measure — For the Triton softmax kernel on an M×N fp32 matrix, write down the minimum bytes that must move, convert to a time at 3.35 TB/s, and compare against the measured 705 µs for 16384². Then do the same accounting for the naive version using the 5MN + M reads / 3MN + 2M writes figures and check your predicted ratio against the measured 4.6×. A good answer states where the model is wrong, and why the naive version misses the prediction in the direction it does.
Next: L07 Parallelism 1 · Back to the map.