跳到主要内容

Triton 入门:写第一个 kernel

用 Triton 重写向量加法,理解 program、tl.arange 与 block 启动背后的模型。

Triton 让你用 Python 写 GPU kernel,编译器(MLIR)负责生成 接近手写 CUDA 的性能。代价是你需要理解它的抽象:block 级别的 program + 向量化的 tl 算子

向量加法

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(0)                    # 第几个 block
    offs = pid * BLOCK + tl.arange(0, BLOCK)  # 本 block 的元素下标
    mask = offs < n                           # 边界掩码
    x = tl.load(x_ptr + offs, mask=mask)
    y = tl.load(y_ptr + offs, mask=mask)
    tl.store(out_ptr + offs, x + y, mask=mask)

def add(x, y):
    out = torch.empty_like(x)
    n = x.numel()
    add_kernel[(triton.cdiv(n, 1024),)](x, y, out, n, BLOCK=1024)
    return out

对照 CUDA 的差异很明显:

概念CUDATriton
调度单位线程、warpblock(program)
索引threadIdx.xtl.arange
边界处理手写 ifmask=
内存手动管理tl.load / tl.store

三个入门建议

  1. tl.arange 开始——它替代了 CUDA 里最易错的索引计算;
  2. BLOCK: tl.constexpr 是编译期常量,调优时改它比改代码更常见;
  3. triton.testing.Benchmark 对照 PyTorch 的 torch.add, 感受内存带宽上限离你有多近。

Triton 已进入 PyTorch 2 的 torch.compile 后端,学会它等于 同时拿到一套可读性更好的 kernel 开发工具。

评论