Triton is a Python-embedded language and compiler for writing GPU kernels at the level of blocks of data rather than individual threads. You decorate a Python function with @triton.jit, and inside it every value is either a scalar or a statically shaped block that lives in registers or shared memory. The compiler decides how threads, warps, vector loads and tensor-core instructions map onto that block. It is also the language PyTorch's torch.compile emits for most of the fused kernels it generates.
This article treats Triton as a language: values, shapes, types, constants, masks, tl.dot, loops and atomics. It then builds a tiled matrix multiply with a fused bias and GELU epilogue line by line. It closes with the failure modes that bite real users, how to test and ship kernels, and when not to write one. If you want the conceptual tour of the programming model, read Triton kernels in depth and come back.
Program instances, pointers and blocks
A Triton kernel runs as a grid of program instances. Each instance is a single logical program that operates on whole blocks; under the hood it is executed by num_warps warps (4 by default on NVIDIA, so 128 threads), but you never index a thread. The only identity you get is tl.program_id(axis) for up to three grid axes and tl.num_programs(axis) for the grid size.
Inside the kernel there are three kinds of values:
- Scalars, such as kernel arguments
Morstride_am, and the program id. Integer arguments are 32-bit unless their value needs 64 bits. - Pointers. A torch tensor passed as an argument arrives as a pointer to its element type,
*fp16for example. Pointer plus integer block gives a block of pointers, which is whattl.loadconsumes. - Blocks (Triton calls them tensors): N-dimensional arrays whose shape is fixed at compile time.
tl.arange(0, 128)produces a block of shape[128]; indexing with[:, None]and[None, :]adds axes so arithmetic broadcasts exactly as in NumPy.
Two rules follow from blocks being compile-time objects. First, every block dimension must be a power of two, so a 100-wide row is processed as a 128-wide block with a mask. Second, you cannot branch on a block value with Python if; per-element selection is tl.where(cond, x, y). You can branch on scalars and on compile-time constants, and the latter disappears entirely during compilation.
Compile-time constants and specialization
Arguments annotated tl.constexpr are baked into the compiled binary. Block sizes must be constexpr because they determine shapes. Anything else that changes code structure, such as a flag choosing an activation function, should also be constexpr so the dead branch is removed rather than evaluated at run time.
Every distinct combination of constexpr values is a separate compilation. Triton also specializes ordinary integer arguments: it records whether each one equals 1 and whether it is divisible by 16, and pointer arguments by their 16-byte alignment, and compiles a variant per pattern. Divisibility lets the compiler emit wide vectorized loads; a stride of 1 lets it prove contiguity. So an odd leading dimension can silently get a slower kernel. Compiled binaries are cached on disk (~/.triton/cache by default, overridable with TRITON_CACHE_DIR), so the cost is paid once per machine and shape pattern, not per call.
When the compiler cannot see a property you know, you can tell it. tl.multiple_of(x, 16) and tl.max_contiguous(x, 16) are hints on offset blocks that unlock vectorization, and tl.static_assert fails compilation early when a constexpr combination is invalid, for example a block smaller than 16 fed to tl.dot.
Loads, stores, masks and types
Memory access is explicit and block-shaped. tl.load(ptrs, mask=m, other=0.0) reads one element per pointer where m is true and substitutes other elsewhere; tl.store(ptrs, vals, mask=m) writes only the masked lanes. Masks are how Triton handles edges: you compute offsets for a full power-of-two block and mask out the part that falls past the end of the tensor.
The data-type rules mirror what a careful CUDA programmer would do by hand. Loads return the pointer's element type; arithmetic promotes in the usual way; .to(dtype) casts explicitly. Accumulate reductions in tl.float32 even when inputs are fp16 or bf16, because a 4,096-term sum in half precision loses most of its low bits.
Triton also provides tl.make_block_ptr and tl.advance, a structured way to describe a 2-D tile of a strided tensor (base, shape, strides, offsets, block shape, order) and step it along an axis. Recent releases also add tensor descriptors for Hopper-class TMA copies; that API has been renamed across versions, so check the docs for the release you pin.
Compute: reductions, tl.dot, loops and atomics
The compute primitives are deliberately few. Elementwise math works on blocks with operators and functions such as tl.exp, tl.sqrt and tl.maximum. Reductions take an axis: tl.sum(x, axis=1), tl.max, tl.argmax, and tl.reduce for a custom combine function. tl.cumsum and tl.associative_scan cover prefix operations.
The heart of most performance work is tl.dot(a, b, acc), a block matrix multiply of an [M, K] block by a [K, N] block, optionally accumulating into acc. The compiler lowers it to tensor-core instructions (MMA on NVIDIA, MFMA on AMD) when dtypes and shapes allow, which is why each dimension must be at least 16. For fp32 inputs the input_precision argument selects whether TF32 tensor cores may be used; leaving the default on an Ampere-or-later GPU usually means TF32, which is faster and less precise than IEEE fp32. See tensor cores in depth for what that trade costs.
Loops are ordinary Python for k in range(0, K, BLOCK_K) over scalars; the compiler turns them into a structured loop and, with num_stages greater than one, software-pipelines the loads so the next tile is in flight while the current one is multiplied. tl.static_range unrolls at compile time. Masked atomics such as tl.atomic_add combine partial results across program instances, as in split-K reductions.
Worked example: a fused matmul with bias and GELU
The canonical exercise is a matrix multiply C = act(A @ B + bias) with A of shape [M, K], B of shape [K, N], fp16 inputs and fp16 output. Each program instance computes one BLOCK_M x BLOCK_N tile of C, walking along K in steps of BLOCK_K.
import torch, triton, triton.language as tl
@triton.autotune(
configs=[
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 32, "GROUP_M": 8}, num_warps=8, num_stages=3),
triton.Config({"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 32, "GROUP_M": 8}, num_warps=4, num_stages=4),
],
key=["M", "N", "K"],
)
@triton.jit
def matmul_bias_gelu(a_ptr, b_ptr, bias_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, APPLY_GELU: tl.constexpr):
# 1. grouped ordering: neighbouring pids share rows of A for L2 reuse
pid = tl.program_id(0)
num_m = tl.cdiv(M, BLOCK_M)
num_n = tl.cdiv(N, BLOCK_N)
group = GROUP_M * num_n
first_m = (pid // group) * GROUP_M
size_m = min(num_m - first_m, GROUP_M)
pid_m = first_m + (pid % group) % size_m
pid_n = (pid % group) // size_m
# 2. tile offsets; int64 avoids overflow on very large tensors
offs_m = pid_m.to(tl.int64) * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n.to(tl.int64) * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn
# 3. K loop, fp32 accumulator
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, K, BLOCK_K):
k_mask = (k + offs_k) < K
a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & k_mask[None, :], other=0.0)
b = tl.load(b_ptrs, mask=k_mask[:, None] & (offs_n[None, :] < N), other=0.0)
acc = tl.dot(a, b, acc)
a_ptrs += BLOCK_K * stride_ak
b_ptrs += BLOCK_K * stride_bk
# 4. fused epilogue: bias + tanh-approximated GELU, then cast
bias = tl.load(bias_ptr + offs_n, mask=offs_n < N, other=0.0).to(tl.float32)
acc = acc + bias[None, :]
if APPLY_GELU:
inner = 0.7978845608 * (acc + 0.044715 * acc * acc * acc)
t = 2.0 * tl.sigmoid(2.0 * inner) - 1.0 # tanh via sigmoid
acc = 0.5 * acc * (1.0 + t)
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.float16), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))
def linear_gelu(a, w, bias):
M, K = a.shape
K2, N = w.shape
assert K == K2 and a.is_cuda and a.dtype == torch.float16
c = torch.empty((M, N), device=a.device, dtype=torch.float16)
grid = lambda meta: (triton.cdiv(M, meta["BLOCK_M"]) * triton.cdiv(N, meta["BLOCK_N"]),)
matmul_bias_gelu[grid](a, w, bias, c, M, N, K,
a.stride(0), a.stride(1), w.stride(0), w.stride(1),
c.stride(0), c.stride(1), APPLY_GELU=True)
return cWalk the four parts. The grouped ordering exists because the obvious row-major mapping of pid to tiles makes consecutive program instances sweep a whole row of output tiles, each needing a different column panel of B. Grouping GROUP_M rows together means instances running at the same time share both A row panels and B column panels, so more of each load is served from L2 instead of HBM.
The offsets are built by broadcasting: a column of row offsets plus a row of column offsets gives a 2-D block of pointers. Casting the program id to tl.int64 before multiplying is cheap insurance; without it, pid_m * BLOCK_M * stride_am is computed in 32 bits and wraps once the tensor passes about two billion elements.
The K loop keeps an fp32 accumulator in registers, masks the ragged last slice of K, and advances both pointer blocks. The epilogue is where Triton earns its keep: bias and activation are applied to the accumulator while it is still in registers, so the output touches memory once. Eager PyTorch would make three trips through HBM for the same output.
Testing and benchmarking
Test against a reference, on shapes that stress the masks, with tolerances that match the precision.
import pytest
@pytest.mark.parametrize("M,N,K", [(1, 1, 16), (127, 255, 33), (512, 512, 512), (4096, 1000, 4096)])
def test_linear_gelu(M, N, K):
torch.manual_seed(0)
a = torch.randn(M, K, device="cuda", dtype=torch.float16)
w = torch.randn(K, N, device="cuda", dtype=torch.float16) / K ** 0.5
bias = torch.randn(N, device="cuda", dtype=torch.float16)
ref = torch.nn.functional.gelu(a.float() @ w.float() + bias.float(), approximate="tanh")
out = linear_gelu(a, w, bias)
torch.testing.assert_close(out.float(), ref, atol=2e-2, rtol=2e-2)Odd sizes such as 127 and 33 exist to exercise every mask; a kernel that is only ever tested at multiples of 128 has untested edge code. Setting TRITON_INTERPRET=1 runs kernels on the CPU with NumPy semantics, so a debugger and print statements work; use small shapes. On the GPU, tl.device_print prints from inside a kernel and tl.static_print prints at compile time. Time with triton.testing.do_bench, which handles warm-up and repeats, and compare with torch.matmul; for a plain GEMM cuBLAS often wins.
Failure modes
The problems that recur in practice:
| Symptom | Cause | Fix |
|---|---|---|
| Wrong values only for large inputs | 32-bit offset overflow in pid * stride arithmetic | cast program ids or offsets to tl.int64 before multiplying |
| Garbage at tensor edges, illegal memory access | missing or wrong mask on a load or store | mask every load and store; test sizes that are not multiples of the block |
| Compile error mentioning shapes | block dimension not a power of two, or tl.dot block below 16 | round blocks up and mask; add tl.static_assert checks |
| Out of resources: shared memory | blocks times num_stages exceed shared memory | smaller blocks or fewer stages |
| First call takes seconds; latency spikes in serving | JIT compile plus autotuning on new shapes | warm up at startup on representative shapes; persist the cache directory |
| Slower than expected after a stride change | lost divisible-by-16 specialization, so no vector loads | pad leading dimensions to multiples of 16; add tl.multiple_of hints |
| Results differ slightly from PyTorch fp32 | TF32 used by tl.dot on fp32 inputs | set input_precision explicitly and document the choice |
Autotuning has its own trap: the key list decides when to re-tune. Omit an argument that matters and you keep a stale config; include a value that changes every call, such as sequence length, and you re-tune constantly. Bucket it first.
Shipping kernels in a PyTorch codebase
A kernel is not done when it is fast; it needs to fit into the framework around it. With PyTorch 2.6 and later, torch.library.triton_op and torch.library.wrap_triton register a Triton kernel as a custom operator that torch.compile can trace through. If the operation is used in training you also need a backward, either a second kernel registered through register_autograd or a composition of existing operators. Before hand-writing anything, check whether torch.compile already fuses it for you.
Pin the Triton version alongside PyTorch, since each PyTorch release is built against a specific Triton. Ship the autotune configs you tested and nothing else. Re-tune per GPU model; the best config on an A100 is often not the best on an H100.
Trade-offs
Triton sits between libraries and CUDA C++. Compared with calling cuBLAS or cuDNN, it lets you fuse arbitrary epilogues, prologues and odd shapes, which is where the wins are: attention variants, fused normalization, quantized matmuls with on-the-fly dequantization. See FlashAttention in depth for the most famous kernel of that kind. Compared with CUDA C++ or CUTLASS, Triton gives up direct control of thread layout and shared memory placement, which is why the fastest kernels on brand-new architectures usually appear in CUDA first. Use a library when the operation is standard, Triton when you need fusion, and CUDA when the last ten percent outweighs development time.
What to do next
- Install a Triton version matching your PyTorch and run the vector-add and softmax tutorials under
TRITON_INTERPRET=1and then on the GPU. - Port the matmul above, run the shape-parametrized test, and compare
do_benchnumbers withtorch.matmulplus separate bias and GELU. - Profile one fused operation in your own model that does several elementwise passes after a matmul; that is the most likely first win.
- Add int64 offsets, masks and
tl.static_assertguards to every kernel you own. - Register production kernels with
torch.library, warm them up at startup, and persistTRITON_CACHE_DIRbetween deployments. - Re-tune and re-benchmark on each new GPU generation and each Triton upgrade before rollout.