Triton is a Python-embedded language and compiler for writing GPU kernels. It exists because the two usual options are both uncomfortable. Calling library kernels (cuBLAS, cuDNN, the fused ops in PyTorch) is fast but only covers operations someone else anticipated. Writing CUDA C++ covers anything, but makes you responsible for every thread index, every coalesced load, every shared-memory buffer and every bank conflict. Triton sits between the two: you write the algorithm for one tile of data, and the compiler handles most of the per-thread details.
This article explains the model from first principles, works through a complete fused softmax kernel, shows how the knobs you set map onto the hardware, and covers the compile pipeline, autotuning, debugging and the ways Triton kernels fail in production. It assumes you know what a GPU kernel is. If not, start with the CUDA programming guide and come back.
The core idea: programs operate on blocks, not threads
In CUDA you write code for a single thread and launch thousands of them, grouped into thread blocks. Cooperation inside a block (sharing data through shared memory, synchronising, splitting loads so neighbouring threads touch neighbouring addresses) is your job. In Triton you write code for a single program instance, and the values it manipulates are whole blocks: vectors or small matrices whose shapes are compile-time constants. A statement such as x = tl.load(ptr + offs, mask=offs < n) loads an entire block of elements. The compiler decides which of the program's threads loads which element and how the load is vectorised.
You still choose the decomposition. You launch a grid of program instances, each reads its coordinates with tl.program_id(axis), computes which slice of the output it owns, and loops over the inputs it needs. That is the level where performance decisions live: tile shape, reuse, fusion and launch order. The compiler takes over below it: thread mapping, memory coalescing, staging through shared memory and emitting tensor-core instructions for tl.dot.
Three language rules follow. First, block shapes must be compile-time constants, which is why they are passed as tl.constexpr parameters. Second, tl.arange(0, BLOCK) requires a power-of-two length, so a row of 1,000 elements is processed with a 1,024-wide block and a mask. Third, every load and store that can run past the end of a tensor needs a mask; there are no automatic bounds checks.
Worked example: a fused row-wise softmax
Softmax over the last dimension of an M by N matrix is a good first kernel because the naive PyTorch version is memory-bound and wasteful. Computing max, subtracting, exponentiating, summing and dividing as separate ops reads and writes the matrix several times. A fused kernel reads each row once into registers, does all the arithmetic there, and writes once. When the row fits in one block, that is close to the minimum possible memory traffic.
import torch
import triton
import triton.language as tl
@triton.jit
def softmax_kernel(out_ptr, in_ptr, in_row_stride, out_row_stride, n_cols,
BLOCK: tl.constexpr):
row = tl.program_id(axis=0) # one program per row
offs = tl.arange(0, BLOCK) # BLOCK is a power of two >= n_cols
mask = offs < n_cols
x = tl.load(in_ptr + row * in_row_stride + offs, mask=mask, other=-float("inf"))
x = x.to(tl.float32) # do the maths in fp32
x = x - tl.max(x, axis=0) # numerically stable
num = tl.exp(x)
den = tl.sum(num, axis=0)
y = num / den
tl.store(out_ptr + row * out_row_stride + offs, y, mask=mask)
def softmax(x: torch.Tensor) -> torch.Tensor:
assert x.ndim == 2 and x.is_cuda
m, n = x.shape
out = torch.empty_like(x)
block = triton.next_power_of_2(n)
num_warps = 4 if block <= 2048 else (8 if block <= 4096 else 16)
softmax_kernel[(m,)](out, x, x.stride(0), out.stride(0), n,
BLOCK=block, num_warps=num_warps)
return out
x = torch.randn(4096, 1000, device="cuda", dtype=torch.float16)
torch.testing.assert_close(softmax(x), torch.softmax(x.float(), dim=-1).half(),
atol=1e-3, rtol=1e-3)The grid is one program per row, so 4,096 rows become 4,096 program instances scheduled across the streaming multiprocessors (SMs). Padding lanes load minus infinity through other=, so they vanish from the max and add exp(-inf) = 0 to the sum. The maths runs in fp32 because an fp16 sum of a thousand exponentials loses precision. The row stride is passed, not assumed, so sliced inputs work.
The kernel has a hard limit: the whole row must fit in one block of registers. For very long rows the block needs more registers than the program has, and the kernel spills to local memory or fails to compile. Longer rows need a different algorithm: loop over the row in chunks, keeping a running max and a running sum and rescaling the sum when the max changes (the online softmax trick that FlashAttention is built on). Extending a Triton kernel's range usually means changing the algorithm, not tuning a flag.
How your choices map onto the hardware
Triton hides threads, but the hardware still has them, and the knobs you set decide how they are used. Each program instance becomes one CUDA thread block (or the AMD equivalent). num_warps sets how many 32-thread warps that block has, so num_warps=4 means 128 threads share the work of one tile. More warps give more latency hiding, but fewer large blocks fit on an SM at once. The resident-block limit is covered in the occupancy guide; in Triton you influence it through block size, num_warps and num_stages rather than through launch bounds.
The block sizes set register and shared-memory pressure. A 128 by 128 fp32 accumulator holds 16,384 values; spread over 256 threads that is 64 registers per thread before any operands. Double the tile and you spill. For tl.dot, the operand tiles are staged through shared memory, and the compiler emits tensor-core instructions when the dtypes and shapes allow. Tensor cores need tile dimensions that are multiples of their fragment shapes, which in practice means BLOCK_K of at least 16 and fp16, bf16, fp8 or tf32 inputs.
num_stages controls software pipelining in loops over K: the compiler issues the loads for iterations i+1, i+2 and so on while computing iteration i, holding several copies of the operand tiles in shared memory. More stages hide more memory latency but multiply shared-memory use, and too many make the kernel fail to compile with an out-of-resources error or drop occupancy.
# From the Triton matmul tutorial: tile the output, loop over K, accumulate in fp32.
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
num_pid_in_group = GROUP_SIZE_M * num_pid_n # grouped ordering for L2 reuse
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0)
b = tl.load(b_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0)
accumulator = tl.dot(a, b, accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += BLOCK_SIZE_K * stride_bkThe grouped ordering is the one decision here that the compiler cannot make for you. With a naive row-major walk over output tiles, consecutive programs each need a different B column-panel, so B is streamed from DRAM repeatedly. Grouping GROUP_SIZE_M rows of tiles means the programs running together reuse a small set of panels that stay in L2. It is only a remapping of program_id.
The compile pipeline and specialization
When you first call a @triton.jit function, Triton parses its Python AST into Triton IR (TTIR), an MLIR dialect of block-level operations with no notion of threads. It then converts to TritonGPU IR (TTGIR), where each tensor gets a layout describing which thread and register hold which element. This is where coalescing, shared-memory staging, pipelining and tensor-core lowering happen. TTGIR is lowered to LLVM IR, then to PTX and a cubin for NVIDIA, or to AMDGCN for AMD GPUs. The binary is written to an on-disk cache (by default under the user's home directory, redirectable with TRITON_CACHE_DIR).
The compiled binary is specialized, not generic. The cache key includes argument dtypes, every constexpr value, the target GPU, the Triton version and some integer properties. By default integer arguments and pointers are checked for divisibility by 16, and integers equal to 1 are specialized, because knowing a stride is 1 or aligned enables vectorised loads. So a new shape can trigger a compile taking seconds, and a stride that stops being divisible by 16 produces a second, sometimes slower, binary. Warm the cache with representative shapes before serving traffic.
Autotuning without fooling yourself
No single tile configuration is best across shapes and GPUs, so Triton provides @triton.autotune: a list of triton.Config objects (constexpr values plus num_warps and num_stages) and a key naming the arguments whose change triggers a new search. On the first call with a new key, every config is compiled and benchmarked, and the winner is cached in memory for that key.
@triton.autotune(
configs=[
triton.Config({"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 64, "GROUP_SIZE_M": 8},
num_stages=3, num_warps=8),
triton.Config({"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8},
num_stages=4, num_warps=4),
],
key=["M", "N", "K"],
)
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K, ...): # strides and constexprs elided
...
grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),)
matmul_kernel[grid](a, b, c, M, N, K, *a.stride(), *b.stride(), *c.stride())Three pitfalls are common. First, keying on the exact M in an LLM server, where M is the token count of a batch and takes hundreds of values, turns every new batch size into a multi-second tuning pause. Bucket M before it reaches the key, or tune offline and hard-code a table. Second, autotuning runs the kernel several times, so a kernel that accumulates into its output or updates in place corrupts data during benchmarking; use the decorator's arguments to reset or restore those tensors between runs. Third, configs tuned on one GPU generation are not evidence for another; re-tune, and confirm with triton.testing.do_bench under realistic load.
Debugging and checking correctness
Setting the environment variable TRITON_INTERPRET=1 runs kernels in a NumPy-based interpreter on the CPU. It is slow but works with a normal Python debugger. Inside compiled kernels, tl.device_print prints runtime values, tl.static_print prints compile-time values such as shapes and constexprs, and tl.static_assert fails compilation when an assumption about block sizes is violated.
Always test against a PyTorch reference with tolerances chosen for the dtype. Test awkward shapes deliberately: sizes one either side of a block, a single row, non-contiguous inputs, and tensors with more than 2^31 elements. That last case catches the most insidious bug: offsets computed in int32 overflow silently. Cast the program id or row index to tl.int64 before multiplying by a stride when tensors can be that large.
Failure modes you will meet
- Missing or wrong masks. Out-of-bounds loads return whatever is in memory, or fault. Results look almost right on friendly shapes and fail on odd ones.
- Out of resources. Too much shared memory (large tiles times many stages) or too many registers; the fix is smaller tiles or fewer stages, not more warps.
- Register spills. The kernel compiles but runs several times slower; profile with Nsight Compute.
- Recompilation storms. A constexpr or autotune key that varies per call means compiling in the hot path.
- Precision surprises.
tl.doton fp32 inputs may use TF32 arithmetic by default; if you need IEEE fp32, request it explicitly for your Triton version. - Version drift. Experimental APIs change between releases. Pin Triton with PyTorch and re-benchmark on upgrade.
Running Triton kernels in production
Most teams meet Triton first through torch.compile, whose Inductor backend generates Triton kernels for fused pointwise, reduction and some matmul patterns. Hand-written Triton is worth it for a fusion the compiler does not find, a custom attention variant, or an op that fuses several memory-bound steps into one pass. Register it as a PyTorch custom operator so torch.compile and autograd treat it as an opaque op, with a separate backward.
Pin versions, warm the cache at startup, and keep a PyTorch fallback behind a flag so a bad kernel can be switched off without a deploy.
Trade-offs against CUDA and libraries
| Option | Strength | Cost |
|---|---|---|
| Library call (cuBLAS, cuDNN) | Fastest for standard shapes, no maintenance | Only the ops and fusions someone shipped |
| Triton | Custom fusions in Python, portable across NVIDIA and AMD, autotunable | Less control over warp-level scheduling; compile latency; fast-moving APIs |
| CUDA C++ / CUTLASS | Full control, access to the newest hardware features first | Much more code, per-architecture tuning, harder review |
Reach for Triton when a memory-bound chain of ops can be fused or a compute-bound op has a twist the libraries lack; drop to CUDA only when a profile shows Triton cannot express the schedule you need.
What to do next
- Install a Triton version matched to your PyTorch build and run the fused softmax above against
torch.softmax, including shapes that are not powers of two. - Run the same kernel with
TRITON_INTERPRET=1and step through it to see the block values. - Profile one memory-bound chain in your model and estimate the bytes saved by fusing it; write the fused kernel only if the saving is material.
- Add an autotune config list, key it on bucketed shapes, and use reset or restore arguments for any in-place kernel.
- Test at block boundaries, with non-contiguous inputs and with more than 2^31 elements; use int64 offsets where needed.
- Persist the Triton cache in your containers and warm it at startup.
- Keep a fallback path and a correctness canary, and re-benchmark on every Triton or GPU change.