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 achievableThe 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#
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):
@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.sumdoes 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#
@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 CUDARule 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#
# 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
# 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 CUDA9. Production implications#
- Use
torch.compilefirst. 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_DIRon a persistent volume. - Warm up every shape you’ll encounter (Section VIII.08).
- Profile with
ncunormally. 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.compileto 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#
- What abstraction level does Triton work at, compared to CUDA?
- What class of kernels is Triton best for, and which should you leave to CUTLASS?
- Write a fused RMSNorm kernel structure from memory.
- What is Triton autotuning and what does it cost?
- Why does
torch.compilematter to a Triton discussion? - How would you debug a Triton kernel producing wrong results?
- 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