Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

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.compile generates 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 of BLOCK indices, not a loop. BLOCK must be a compile-time constant (tl.constexpr) and a power of two.
  • Pointers support arithmetic: x_ptr + offsets is a vector of addresses.
  • mask disables lanes past the end. Every load and store of a partial block needs it, and masked loads take an other= fill value.
  • The launch add_kernel[grid](...) takes a grid tuple, computed with triton.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=1 before 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 print and breakpoints inside it. The book’s tests switch it on automatically when no GPU is present. That’s how every Triton kernel in this book was validated during writing.

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.

kernelshapereference mscustom msmax errorwhole-engine effect
softmax32 × 32000
add+rmsnorm1 × 1024 (decode)
add+rmsnorm4096 × 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

  1. ★ Add a GELU epilogue to the matmul kernel (apply GELU to acc before storing) and test it against F.gelu(a @ b). That’s a fused linear + activation. Where: matmul_kernel in engine/kernels/triton_matmul.py.
  2. ★★ Add @triton.autotune to the matmul and plot TFLOP/s against torch.matmul for square BF16 sizes 512-8,192. Where: matmul_kernel and its launcher in engine/kernels/triton_matmul.py.
  3. ★★ 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: new engine/kernels/triton_rope.py; compare with the provided izh.qwen3.apply_rope.
  4. ★★★ 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_kernel and softmax in engine/kernels/triton_basics.py.

Check your understanding

  1. What does one Triton program instance correspond to, and how does it differ from one CUDA thread?
  2. Why does fusing a softmax make it several times faster, even though it does the same arithmetic?
  3. Why must masked loads in a max-reduction use other=-inf?
  4. 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_stages does underneath.