14. Triton and kernel fusion
In this chapter
- Triton's programming model: you write the work of one block as vector code, and the compiler handles threads, shared memory and tensor cores.
- Kernel fusion: why combining operations into one kernel is the most effective optimization for memory-bound code.
- Fused softmax, fused residual-add + RMSNorm, and a tiled matmul with tensor cores, all in Python.
- How
torch.compilegenerates Triton for you, and when to write your own.
You will build
engine/kernels/triton_basics.py and triton_matmul.py: four kernels, tested on GPU, or on CPU with Triton's interpreter.
Time: 5-7 hours. GPU: optional (the interpreter runs everything, slowly).
Why another GPU language
CUDA C++ gives you full control, and you pay for it in detail: thread indices, shared-memory layouts, barriers, bank conflicts, vector loads, tensor-core fragment layouts. Triton (Tillet et al., 2019, now part of PyTorch) raises the level of abstraction by one step. You write the program for one block of data, using operations on whole tiles (tl.load a tile, tl.dot two tiles, tl.sum along an axis), and the compiler decides how threads map onto them, stages data through shared memory, coalesces loads and uses tensor cores. You still choose the tiling, which is the decision that matters most, but the bookkeeping disappears.
The trade-off: Triton reaches 80-100% of hand-written CUDA for most kernels an inference engine needs, in a fraction of the code. PyTorch’s compiler emits Triton, and most new kernels in vLLM and SGLang start as Triton. When it falls short (warp-specialized FlashAttention 3, the very latest hardware features), engines drop to CUDA, CUTLASS or newer tile languages.
The programming model
A Triton kernel is a Python function decorated with @triton.jit. It’s launched over a grid of program instances, and each instance asks which one it is with tl.program_id(axis):
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
pid = tl.program_id(0) # which block of BLOCK elements is mine
offsets = pid * BLOCK + tl.arange(0, BLOCK) # a vector of BLOCK indices
mask = offsets < n # the last block may run past the end
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
tl.store(out_ptr + offsets, x + y, mask=mask)
def vector_add(x, y, block=1024):
check_device(x, y)
if x.shape != y.shape or not x.is_contiguous() or not y.is_contiguous():
raise ValueError("vector_add needs equal-shape contiguous tensors")
out = torch.empty_like(x)
n = x.numel()
add_kernel[(triton.cdiv(n, block),)](x, y, out, n, BLOCK=block)
return out
Read it as “program pid handles elements [pid*BLOCK, (pid+1)*BLOCK)”:
tl.arange(0, BLOCK)is a vector ofBLOCKindices, not a loop.BLOCKmust be a compile-time constant (tl.constexpr) and a power of two.- Pointers support arithmetic:
x_ptr + offsetsis a vector of addresses. maskdisables lanes past the end. Every load and store of a partial block needs it, and masked loads take another=fill value.- The launch
add_kernel[grid](...)takes a grid tuple, computed withtriton.cdiv.
The first call compiles the kernel for the given constants and dtypes (seconds); later calls reuse the compiled binary. Time steady state, not the first call.
Tip
No GPU? Run Triton on the CPU. Set
TRITON_INTERPRET=1before importing triton and pass CPU tensors. The interpreter executes the kernel with NumPy, about 100-1,000x slower but numerically faithful, and you can put
Fusion: the optimization that matters most for memory-bound code
Consider softmax over the rows of a [4096, 32000] FP32 matrix (512 MB), written as separate PyTorch operations:
m = x.max(dim=-1, keepdim=True) # read x (512 MB) write m (small)
e = torch.exp(x - m) # read x, write x - m, read it back, write e (4 passes over 512 MB)
s = e.sum(dim=-1, keepdim=True) # read e
y = e / s # read e, write y
That’s around 8 full passes over memory for an operation whose minimum is 2 (read x once, write y once). Since softmax is memory-bound (Chapter 10), runtime is proportional to passes, so a fused kernel that keeps each row in registers is about 4x faster:
@triton.jit
def softmax_kernel(x_ptr, out_ptr, row_stride, width, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK)
mask = cols < width
x = tl.load(x_ptr + row * row_stride + cols, mask=mask, other=-float("inf")).to(tl.float32)
x = x - tl.max(x, axis=0) # stable: largest exponent is exp(0) = 1
num = tl.exp(x)
tl.store(out_ptr + row * row_stride + cols, num / tl.sum(num, axis=0), mask=mask)
def softmax(x):
"""Row softmax over the last dim of a contiguous 2-D tensor, one program per row.
Reads x once and writes once: the unfused version (max, sub, exp, sum, div) makes ~5 passes."""
check_device(x)
if x.ndim != 2 or not x.is_contiguous():
raise ValueError("softmax expects a contiguous matrix")
out = torch.empty_like(x)
block = triton.next_power_of_2(x.shape[1])
if block > 65536:
raise ValueError("Rows wider than 65536 need a multi-pass (online) kernel")
softmax_kernel[(x.shape[0],)](x, out, x.stride(0), x.shape[1], BLOCK=block)
return out
One program per row: load the row (padded with $-\infty$, the identity for max), subtract the max, exponentiate, sum, divide, and store. The intermediates never touch memory. Fusion is how almost every non-GEMM operation in a production engine is implemented: normalization, activations, rotary embeddings, residual additions, sampling.
Fused residual add + RMSNorm
Every transformer layer does x = x + branch_output followed by h = rmsnorm(x). Unfused, that’s a read of both inputs, a write of the sum, a read of the sum for the statistic, and another read to normalize, then a write. Fused, it’s two reads and two writes, all in one kernel:
@triton.jit
def add_rmsnorm_kernel(x_ptr, res_ptr, w_ptr, out_ptr, res_out_ptr, width, eps,
HAS_RESIDUAL: tl.constexpr, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK)
mask = cols < width
x = tl.load(x_ptr + row * width + cols, mask=mask, other=0.0).to(tl.float32)
if HAS_RESIDUAL: # fuse "x = x + residual" into the same pass
x += tl.load(res_ptr + row * width + cols, mask=mask, other=0.0).to(tl.float32)
tl.store(res_out_ptr + row * width + cols, x, mask=mask)
inv_rms = 1.0 / tl.sqrt(tl.sum(x * x, axis=0) / width + eps)
w = tl.load(w_ptr + cols, mask=mask, other=0.0).to(tl.float32)
tl.store(out_ptr + row * width + cols, x * inv_rms * w, mask=mask)
def rmsnorm(x, weight, eps=1e-6, residual=None):
"""RMSNorm(x [+ residual]) * weight over the last dim. With a residual, also returns x + residual,
so a transformer layer's 'add then normalize' costs one memory pass instead of three."""
check_device(x, weight)
shape = x.shape
x2 = x.reshape(-1, shape[-1]).contiguous()
out = torch.empty_like(x2)
res_out = torch.empty_like(x2) if residual is not None else out
res = residual.reshape(-1, shape[-1]).contiguous() if residual is not None else x2
add_rmsnorm_kernel[(x2.shape[0],)](x2, res, weight, out, res_out, shape[-1], eps,
HAS_RESIDUAL=residual is not None,
BLOCK=triton.next_power_of_2(shape[-1]))
if residual is not None:
return out.view(shape), res_out.view(shape)
return out.view(shape)
HAS_RESIDUAL is a constexpr, so Triton compiles two specialized versions and the if costs nothing at runtime. You’ll swap this kernel into the engine in Chapter 19.
Matmul in Triton
The tiled algorithm from Chapter 12, at tile level. Each program computes one [BLOCK_M, BLOCK_N] output tile by looping over K in BLOCK_K steps, with tl.dot doing the tile product, on tensor cores when the dtypes allow:
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr):
"""(Your engine: Chapter 14)"""
# Map the 1-D program id to an output tile. Walking GROUP_M tile-rows at a time keeps the
# B tiles those rows need hot in L2 cache ("grouped ordering", Triton tutorial 03).
pid = tl.program_id(0)
tiles_m, tiles_n = tl.cdiv(M, BLOCK_M), tl.cdiv(N, BLOCK_N)
group = pid // (GROUP_M * tiles_n)
first_m = group * GROUP_M
group_size = min(tiles_m - first_m, GROUP_M)
pid_m = first_m + (pid % group_size)
pid_n = (pid % (GROUP_M * tiles_n)) // group_size
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
rk = tl.arange(0, BLOCK_K)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) # accumulate in FP32 whatever the inputs
for k0 in range(0, K, BLOCK_K):
a = tl.load(a_ptr + rm[:, None] * stride_am + (k0 + rk)[None, :] * stride_ak,
mask=(rm[:, None] < M) & ((k0 + rk)[None, :] < K), other=0.0)
b = tl.load(b_ptr + (k0 + rk)[:, None] * stride_bk + rn[None, :] * stride_bn,
mask=((k0 + rk)[:, None] < K) & (rn[None, :] < N), other=0.0)
acc += tl.dot(a, b) # tensor cores when dtypes allow
c = acc.to(c_ptr.dtype.element_ty)
tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :] * stride_cn, c,
mask=(rm[:, None] < M) & (rn[None, :] < N))
def matmul(a, b, block_m=64, block_n=64, block_k=32, group_m=8):
check_device(a, b)
if a.ndim != 2 or b.ndim != 2 or a.shape[1] != b.shape[0]:
raise ValueError("matmul expects [M, K] @ [K, N]")
m, k = a.shape
n = b.shape[1]
c = torch.empty((m, n), device=a.device, dtype=a.dtype)
grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n),)
matmul_kernel[grid](a, b, c, m, n, k, a.stride(0), a.stride(1), b.stride(0), b.stride(1),
c.stride(0), c.stride(1), BLOCK_M=block_m, BLOCK_N=block_n,
BLOCK_K=block_k, GROUP_M=group_m)
return c
The grouped ordering at the top deserves a word. The GPU runs programs roughly in pid order. If consecutive programs walked along one row of output tiles, each would need a different column strip of B, and B’s tiles would be evicted from L2 before neighbouring rows reused them. Grouping GROUP_M tile-rows together makes nearby programs share B tiles while they’re still cached. This often gives a 10-20% speedup for free.
Autotuning
The best BLOCK_M, BLOCK_N, BLOCK_K, number of warps and number of pipeline stages depend on the GPU and the matrix shape. Triton can try several and remember the fastest per shape:
@triton.autotune(configs=[
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 32}, num_warps=8, num_stages=3),
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=4),
triton.Config({"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 64}, num_warps=4, num_stages=4),
], key=["M", "N", "K"])
@triton.jit
def matmul_kernel(...): ...
num_stages controls software pipelining: loading the next K-tile while computing on the current one, which hides memory latency behind arithmetic (PMPP §15.8).
Let the compiler write it: torch.compile
PyTorch’s compiler traces your model, fuses chains of elementwise and reduction operations, and generates Triton kernels. Wrapping a function in torch.compile often gives much of the benefit of hand fusion with no kernel code:
fast_norm = torch.compile(lambda x, r, w: rmsnorm_reference(x + r, w))
# TORCH_LOGS=output_code python script.py prints the Triton it generated
Reading the generated code is one of the best ways to learn Triton idioms. When does writing your own still pay? When the fusion crosses boundaries the compiler won’t cross (attention, quantized matmuls, paged memory access), when you need a specific algorithm (online softmax), or when you need control over memory layout.
Integrating kernels into an engine
Keep a readable PyTorch reference next to every custom kernel, and select between them in one place:
def rmsnorm(x, w, eps, backend="torch"):
if backend == "triton" and x.is_cuda and x.shape[-1] <= 65536:
return triton_basics.rmsnorm(x, w, eps)
return reference_rmsnorm(x, w, eps) # always-correct fallback
Then test at two levels. First the kernel against the reference, on awkward shapes (odd widths, one row, the largest supported width). Then the whole model’s logits with the kernel swapped in. A correct kernel can still be called with the wrong axis or a non-contiguous input. Record which backend produced every benchmark number.
| kernel | shape | reference ms | custom ms | max error | whole-engine effect |
|---|---|---|---|---|---|
| softmax | 32 × 32000 | ||||
| add+rmsnorm | 1 × 1024 (decode) | ||||
| add+rmsnorm | 4096 × 1024 (prefill) |
Fill a table like this with measurements, not expectations. Decode-sized inputs often benefit less than you’d hope: launch overhead, not memory traffic, dominates tiny kernels (Chapter 19).
Build it
Engine milestone 14: Triton kernels. Write the four @triton.jit kernels in engine/kernels/triton_basics.py (add_kernel, softmax_kernel, add_rmsnorm_kernel) and engine/kernels/triton_matmul.py (matmul_kernel). The Python wrappers that launch them are provided.
pytest tests/test_ch14_triton.py # on CPU this uses TRITON_INTERPRET=1 automatically
python run.py kernels --impl engine
On a GPU, benchmark your fused softmax and add+RMSNorm against the unfused PyTorch versions with measure_cuda, and report the achieved GB/s.
Stretch exercises
- ★ Add a
GELUepilogue to the matmul kernel (apply GELU toaccbefore storing) and test it againstF.gelu(a @ b). That’s a fused linear + activation. Where:matmul_kernelinengine/kernels/triton_matmul.py. - ★★ Add
@triton.autotuneto the matmul and plot TFLOP/s againsttorch.matmulfor square BF16 sizes 512-8,192. Where:matmul_kerneland its launcher inengine/kernels/triton_matmul.py. - ★★ Write a Triton RoPE kernel that rotates Q and K in place for given positions (Chapter 17 explains RoPE), and test it against the reference
apply_rope. Where: newengine/kernels/triton_rope.py; compare with the providedizh.qwen3.apply_rope. - ★★★ Write a softmax for rows too long for one block: two passes, or an online version that carries a running max and sum (preview of Chapter 15). Where:
softmax_kernelandsoftmaxinengine/kernels/triton_basics.py.
Check your understanding
- What does one Triton program instance correspond to, and how does it differ from one CUDA thread?
- Why does fusing a softmax make it several times faster, even though it does the same arithmetic?
- Why must masked loads in a max-reduction use
other=-inf? - Why should you keep a PyTorch reference path alongside every custom kernel?
Going deeper
- GPU Mode L14 (Umer Adil, A Practitioner’s Guide to Triton: notebook with debugging tips and the interpreter), L18 (fused kernels), L28 (Liger kernels: fused RMSNorm, RoPE and cross-entropy in production), L29 (Triton internals).
- The official Triton tutorials: vector add, fused softmax, matmul with grouped ordering and autotuning, layer norm.
- Tillet, Kung and Cox, Triton: an intermediate language and compiler for tiled neural network computations (2019).
- PMPP §15.8 (software pipelining) for what
num_stagesdoes underneath.