PidokuInfra

Triton Kernels

Expert Advanced 1h 30m Difficulty 4/5 Topic 11 of 12

Prerequisites VI.08, VI.09, VII.07


1. What Triton is#

A Python-embedded language for writing GPU kernels, where you program at the block level rather than the thread level. The compiler handles the thread-level details.

CUDA                              TRITON
you manage threads                you manage BLOCKS of elements
you manage shared memory          the compiler manages it
you handle coalescing             the compiler handles it
you write ~200 lines              you write ~30 lines
you get 100% of achievable        you get 70-95% of achievable

The trade: ~5-10x less code, ~5-30% less performance, for kernels where the compiler’s choices are good. For elementwise, reduction, and fusion kernels — which is most of what inference needs beyond GEMM — that trade is excellent.


2. Why it matters for inference#

YOU WILL NOT BEAT cuBLAS/CUTLASS AT GEMM.

But inference needs many NON-GEMM kernels:
  fused add + RMSNorm
  SiLU × up (gated activation)
  RoPE application
  KV cache write with layout transformation
  sampling (temperature, top-k, top-p)
  quantize / dequantize
  custom attention variants

For all of these, Triton gets you 80-95% of hand-written CUDA
performance in a fraction of the time — and it's maintainable
by people who aren't CUDA experts.

This is why Triton is in production: vLLM, PyTorch Inductor, FlashAttention (the Triton version), and many others use it for exactly this class of kernel.


3. Your first Triton kernel#

Python
import triton
import triton.language as tl

@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
    pid = tl.program_id(axis=0)              # which block am I?
    offsets = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offsets < n                        # the bounds guard
    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 add(x, y):
    out = torch.empty_like(x)
    n = x.numel()
    grid = lambda meta: (triton.cdiv(n, meta['BLOCK']),)
    add_kernel[grid](x, y, out, n, BLOCK=1024)
    return out

Compare to the CUDA version (Section VI.08). Same structure — grid, offsets, bounds guard — but you operate on a vector of offsets rather than a single thread’s index, and there’s no explicit thread management.


4. The kernel that matters: fused add + RMSNorm#

The highest-value fusion in an LLM (Section VII.07):

Python
@triton.jit
def add_rmsnorm_kernel(
    x_ptr, residual_ptr, weight_ptr, out_ptr, new_residual_ptr,
    n_cols, eps,
    BLOCK: tl.constexpr,
):
    row = tl.program_id(0)
    x_row   = x_ptr + row * n_cols
    res_row = residual_ptr + row * n_cols
    out_row = out_ptr + row * n_cols
    nres_row = new_residual_ptr + row * n_cols

    cols = tl.arange(0, BLOCK)
    mask = cols < n_cols

    # load and add the residual — ONE pass over memory
    x = tl.load(x_row + cols, mask=mask, other=0.0).to(tl.float32)
    r = tl.load(res_row + cols, mask=mask, other=0.0).to(tl.float32)
    h = x + r

    # write the new residual (needed by the next layer)
    tl.store(nres_row + cols, h, mask=mask)

    # RMSNorm — the reduction happens in FP32
    var = tl.sum(h * h, axis=0) / n_cols
    rstd = 1.0 / tl.sqrt(var + eps)

    w = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
    tl.store(out_row + cols, (h * rstd * w).to(tl.bfloat16), mask=mask)


def add_rmsnorm(x, residual, weight, eps=1e-6):
    out = torch.empty_like(x)
    new_res = torch.empty_like(x)
    n_rows, n_cols = x.shape
    BLOCK = triton.next_power_of_2(n_cols)
    add_rmsnorm_kernel[(n_rows,)](
        x, residual, weight, out, new_res, n_cols, eps,
        BLOCK=BLOCK, num_warps=8)
    return out, new_res

Points to note:

  • One block per row (token). The whole row must fit in BLOCK.
  • tl.sum does the block-level reduction — the compiler generates the warp shuffles.
  • FP32 for the reduction (Section III.11), BF16 for storage.
  • Two outputs: the normalized value and the new residual (Section VII.07).

This kernel is ~35 lines and achieves 85-95% of a hand-tuned CUDA version. The CUDA version is ~150 lines with explicit warp reductions and vectorized loads.


5. Autotuning#

Python
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 32},
                      num_stages=4, num_warps=8),
        triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64,  'BLOCK_K': 32},
                      num_stages=4, num_warps=4),
        triton.Config({'BLOCK_M': 64,  'BLOCK_N': 64,  'BLOCK_K': 64},
                      num_stages=5, num_warps=4),
        # ... more configurations
    ],
    key=['M', 'N', 'K'],       # re-tune when these change
)
@triton.jit
def matmul_kernel(...):
    ...

Autotuning benchmarks each config on the actual hardware and caches the winner. This is how Triton kernels reach competitive performance without hand-tuning per architecture.

Cost: the first call for each new key combination runs all configs. For inference with varying shapes, cache the results (TRITON_CACHE_DIR) and warm up (Section VIII.08).


6. What Triton is good at, and what it isn’t#

EXCELLENT
  ✓ elementwise and fused elementwise chains
  ✓ reductions (norms, softmax)
  ✓ fusion of the above with each other
  ✓ custom attention variants (the FlashAttention tutorial is Triton)
  ✓ quantize/dequantize kernels
  ✓ sampling pipelines
  ✓ anything where you'd otherwise write 200 lines of CUDA

GOOD
  ~ GEMM (reaches 80-95% of cuBLAS; sometimes better on unusual shapes)
  ~ attention (the Triton FlashAttention is competitive)

POOR
  ✗ anything needing warp-level primitives you can't express
  ✗ very irregular memory patterns
  ✗ the last 5-15% of performance on GEMM
  ✗ Hopper-specific features (TMA, warp specialization) — support
    is improving but lags CUDA

Rule of thumb: if the kernel’s structure is “load a block, do arithmetic, reduce, store” — Triton is the right tool. If it’s “orchestrate an intricate pipeline of async copies and tensor core operations” — that’s CUTLASS territory.


7. Debugging and profiling Triton#

Shell
# See the generated code
TRITON_DEBUG=1 python your_script.py

# Dump intermediate representations
MLIR_ENABLE_DUMP=1 python your_script.py

# Interpret mode: run on CPU with Python semantics (SLOW, but debuggable)
TRITON_INTERPRET=1 python your_script.py
Python
# Print from inside a kernel (interpret mode, or with care)
tl.device_print("value: ", x)

# Static assertions
tl.static_assert(BLOCK % 32 == 0)

TRITON_INTERPRET=1 is the killer debugging feature: your kernel runs as ordinary Python with NumPy-like semantics, so you can print, use a debugger, and check intermediate values. Unusably slow, but it finds correctness bugs quickly.

Profile Triton kernels with ncu exactly as you would CUDA kernels (Section X.08) — they compile to ordinary PTX.


8. Triton in production#

WHO USES IT
  PyTorch Inductor    generates Triton for all its fused kernels
                      → every torch.compile fusion IS a Triton kernel
  vLLM                several kernels, and increasingly more
  FlashAttention      has a Triton implementation
  Unsloth, and many   fine-tuning and inference libraries
  others

WHAT THIS MEANS
  If you use torch.compile, you're already running Triton kernels.
  Reading the generated code (Section IV.09) teaches you Triton
  by example.

PRODUCTION CONSIDERATIONS
  □ compile time: first call per shape signature. CACHE IT.
      TRITON_CACHE_DIR=/persistent/path
  □ warm up all shapes you'll see (Section VIII.08)
  □ autotuning adds to first-call cost
  □ version compatibility: Triton, PyTorch, and CUDA must align
  □ AMD support exists (ROCm) — a portability advantage over CUDA

9. Production implications#

  • Use torch.compile first. It generates Triton for you, and covers most fusion opportunities.
  • Write Triton when you need a fusion the compiler won’t do — typically a custom architecture’s norm variant, an unusual activation, or a quantization scheme.
  • Cache the compilation. TRITON_CACHE_DIR on a persistent volume.
  • Warm up every shape you’ll encounter (Section VIII.08).
  • Profile with ncu normally. Triton kernels are ordinary kernels.
  • Triton is portable to AMD, which CUDA is not — relevant if your fleet is heterogeneous.
  • Read the generated code from torch.compile to learn the idioms.

10. Common mistakes#

Writing Triton when torch.compile would have fused it. Check first.

Trying to beat cuBLAS at GEMM. You probably won’t, and you don’t need to.

Not caching compilation. Minutes added to cold start.

Not warming up all shapes. The first request of each shape pays compile time.

BLOCK not a power of two. Triton requires it for many operations.

Forgetting the mask. Out-of-bounds access.

Reducing in low precision. Cast to FP32 for reductions.

Not using TRITON_INTERPRET for debugging. You’ll waste hours.


11. Hands-on exercise#

A. Vector add. Write, run, and verify the kernel from section 3. Benchmark against torch.add. What fraction of memory bandwidth do you achieve?

B. Fused RMSNorm. Implement the kernel from section 4. Verify correctness against PyTorch. Benchmark against (i) unfused PyTorch, (ii) torch.compiled PyTorch, (iii) vLLM’s CUDA kernel if available. Where do you land?

C. SiLU and multiply. Write the gated activation kernel: out = silu(gate) * up. This is the second-highest-value fusion. Benchmark it.

D. Autotune. Add @triton.autotune to your RMSNorm kernel with several num_warps and BLOCK configurations. Measure the improvement over a fixed configuration.

E. Read the generated code. Run a model with TORCH_COMPILE_DEBUG=1. Find the generated Triton kernels. Read three of them. What patterns do you see?

F. Debug with interpret mode. Deliberately introduce a bug (wrong mask, wrong offset). Find it with TRITON_INTERPRET=1.

G. Attention. Work through Triton’s FlashAttention tutorial. Benchmark against flash_attn. How close do you get?


12. Interview questions#

  1. What abstraction level does Triton work at, compared to CUDA?
  2. What class of kernels is Triton best for, and which should you leave to CUTLASS?
  3. Write a fused RMSNorm kernel structure from memory.
  4. What is Triton autotuning and what does it cost?
  5. Why does torch.compile matter to a Triton discussion?
  6. How would you debug a Triton kernel producing wrong results?
  7. What are the production considerations for Triton kernels?

13. Further reading#

  • [REFERENCE] Triton documentation and tutorials — https://triton-lang.org/ (work through the fused softmax and FlashAttention tutorials)
  • [REFERENCE] PyTorch Inductor’s generated Triton code — the best source of idioms
  • [ESTABLISHED] Tillet et al., “Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations” (2019)
  • [REFERENCE] vLLM’s Triton kernels
  • Next: 12 — Compilers: torch.compile and beyond

↑↓ navigate↵ openesc close