Triton Kernels
Writing GPU kernels in Python with Triton: block programs, fusion, tiled matmul and flash attention, autotuning, and one kernel for NVIDIA and AMD.
An interactive AI Infrastructure lesson: 25 steps, about 35 minutes, on a live simulation in your browser.
The inference team added a custom activation to their model: y = clamp(gelu(x + bias) * scale). One line of PyTorch. They run it in eager mode, where every operator is its own GPU kernel, on gpu0 of dgx-1, an H100.
The tensor is 8192 × 8192 bf16 values, 128 MiB. Eager PyTorch launches 4 kernels: add, GELU, multiply, clamp. Each reads the whole tensor from HBM and writes the whole tensor back, so 1.07 GB crosses the memory bus for 128 MiB of input. The four kernels take 372 µs.
What you will learn
Why Triton exists
- Four kernels for one line of Python: In eager mode every framework operator is a separate kernel that reads its input from HBM and writes its output back to HBM.
- Where do the 372 µs go?: A kernel below the ridge point is paid for in bytes, not FLOPs. To make it faster, move fewer bytes.
- Triton: kernels in Python, by the block: Triton is a block-level GPU language in Python: you describe what one program does to one tile, and the compiler handles threads, memory layout and tensor cores, for NVIDIA and AMD.
Programs, blocks and masks
- The first kernel: vector add: A Triton program is one instance of the kernel working on a whole block: program_id says which block, arange builds its offsets, a mask covers the ragged edge.
- Programs become blocks of warps: num_warps × 32 is the program's thread count; BLOCK_SIZE ÷ threads is the work per thread. Occupancy is a means to keep bytes in flight, not a goal.
- The tutorial's vector: 98,432 floats: Below a few hundred thousand elements a kernel is launch-bound: the fixed few microseconds of launch cost more than the data.
- Drill: the launch grid
Fusion and benchmarking
- Fuse the chain into one kernel: Fusing N memory-bound elementwise ops into one kernel divides the HBM traffic, and roughly the time, by N.
- Fusing a reduction: softmax: A reduction fuses when one program can hold what it reduces over: give each program a whole row and the row crosses HBM once in, once out.
- Measure it: triton.testing.do_bench: Benchmark a kernel with do_bench, then convert the time into GB/s or TFLOPS and compare with the GPU's peak: the percentage says whether more work is worth it.
Tiled matmul and autotuning
- A tiled matmul with tl.dot: In a tiled matmul each program owns one output tile and streams K through it; the tile size sets how often A and B are re-read, and tl.dot puts the inner product on tensor cores.
- Bigger tiles, fewer programs: For tensor-core kernels, data reuse per tile matters more than occupancy: bigger tiles at lower occupancy usually win, until they run out of shared memory or registers.
- Break it: stages until OutOfResources: Each pipeline stage keeps one more K-slice load in flight; stages pay off until the loads hide behind the maths, then only cost shared memory = num_stages × (BLOCK_M × BLOCK_K + BLOCK_K × BLOCK_N) × bytes. Over the per-block limit, Triton raises OutOfResources at launch.
- Drill: shared memory per program
- @triton.autotune picks the tile: @triton.autotune benchmarks every config the first time it sees a new key, prunes those that do not fit, and caches the fastest per key.
- Read the winner: Read an autotune result as a story: which resource each config was bound by. The winner is usually the largest tile that still fits, not the highest occupancy.
Flash attention
- Naive attention is memory-bound: Naive attention materialises a seq × seq matrix per head in HBM: memory traffic grows with seq², so long contexts are bound by HBM, not by the tensor cores.
- Flash attention: tile and never write S: Flash attention tiles Q, K and V, keeps a running max and sum per row, and never writes the seq × seq matrix: traffic grows with seq instead of seq², and attention becomes compute-bound.
One kernel, two vendors
- Break it: the CUDA kernel on an MI300X: Triton source is vendor-neutral: the same kernel compiles to PTX on NVIDIA and AMDGCN on AMD. CUDA C++ needs a HIP port.
- Break it: the H100 winner on AMD: Source is portable across vendors, tuning is not: tile sizes and stages that fit 227 KB of shared memory fail on 64 KB of LDS.
- Autotune per architecture: Compare vendors with the same kernel autotuned on each, on your shapes and software versions; datasheet peaks only bound the answer.
Triton in production
- Where Triton fits: Triton trades the last few percent of performance for Python, portability and speed of iteration; CUDA, CUTLASS and ThunderKittens buy that last few percent with more expertise per kernel.
- Caches, cold starts and version pins: Triton moves compilation into run time: platform teams own its cache, its warm-up and its version pin, or users see minutes of cold start.
Recap & playground
- Cheat sheet
- Playground