Triton 自定义算子入门
本教程共 60 篇 · 第 49 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:认识 Triton 语言,学会写一个能用的 GPU 算子,并把它接进 torch.compile 与 PyTorch 生态。
什么时候需要自定义算子
PyTorch 自带两千多个算子,add、matmul、conv2d 应有尽有。绝大多数时候你不需要自己写。但总有几个场景例外:你在论文里看到一个新的激活函数,PyTorch 里没有;或者某个融合计算(比如「乘加一步做完」)用现有算子拼出来太慢。这时候就要自己动手。
自己写的基础运算,就叫自定义算子(Custom Operator)。PyTorch 给出三条路线,从易到难:
- 用 Triton 写 GPU 内核,配合 torch.compile 用——本章主角;
- 用 torch.library 把 Python 函数包装成正式算子;
- 用 C++/CUDA 扩展,走 cpp_extension——最底层,最灵活,也最麻烦。
Triton:用 Python 写 GPU 代码
Triton 是 OpenAI 开源的 GPU 编程语言。传统上给 NVIDIA GPU 写高性能代码要写 CUDA,指针、线程块、共享内存全要自己管,劝退了不少人。Triton 把这事拉回 Python 语法:你写一个带 @triton.jit 装饰器的函数,它帮你编译成 GPU 能跑的内核。装 PyTorch 时 Triton 会作为依赖自动装好,开箱即用。
GPU 的并行模型值得先花三十秒理解。想象一条流水线,把一大摞数据切成小块,每个工人领一块同时开工。GPU 上有成千上万个工人(线程),Triton 代码描述的是「一个工人怎么干它自己那块活」,至于开多少工人、怎么分工,由启动配置决定。你写的内核是单个工人的视角,不是全局视角——这是初学者最容易转不过来的弯。
NoteTriton 内核只能在 GPU 上运行。没有 NVIDIA 显卡的话,读代码理解思路就好,跑不起来是正常的。
先看官方教程的经典例子:向量加法。
import torch
from torch.utils._triton import has_triton
if not has_triton():
print("当前环境不支持 Triton,跳过示例")
else:
import triton
from triton import language as tl
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK_SIZE: "tl.constexpr"):
pid = tl.program_id(axis=0) # 我是第几块
start = pid * BLOCK_SIZE # 我这块从哪开始
offsets = start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n # 防越界
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)
@torch.compile(fullgraph=True)
def add_fn(x, y):
out = torch.zeros_like(x)
n = out.numel()
grid = lambda meta: (triton.cdiv(n, meta["BLOCK_SIZE"]),)
add_kernel[grid](x, y, out, n, BLOCK_SIZE=4)
return out
x = torch.randn(8, device="cuda")
y = torch.randn(8, device="cuda")
print(add_fn(x, y))
看几个关键点。GPU 的并行方式是「分块」:把数据切成 BLOCK_SIZE 大小的小块,每块交给一个程序实例去算。tl.program_id(0) 拿到当前块的编号,tl.arange(0, BLOCK_SIZE) 生成块内下标,两者一加就是全局下标。mask 处理末尾不足一块的情况,防止读写越界——向量长度不整除块大小时必须有它。
grid 是启动配置,告诉 GPU 总共开多少个块。triton.cdiv(n, BLOCK_SIZE) 算的是向上取整的块数。BLOCK_SIZE 标注为 tl.constexpr,意思是编译期常量,值一变内核就要重新编译。
外层再套一层 @torch.compile(fullgraph=True),这个自定义内核就嵌进了编译图里,和模型里其他算子一起被优化。fullgraph=True 要求整张图一次成型、不允许图中断,适合这种纯计算场景。
Tip写内核有个好习惯:先用小数据跑,拿
torch.testing.assert_close(out, x + y)和 PyTorch 原生结果对拍。内核 bug 无声无息,输错一个下标它也不会报错,只会给你错的结果。
自动调优:autotune
BLOCK_SIZE 取多大最快?不知道,得试。Triton 的 autotune 帮你把候选配置挨个跑一遍,挑最快的缓存下来:
@triton.autotune(
configs=[
triton.Config({"BLOCK_SIZE": 128}, num_warps=8),
triton.Config({"BLOCK_SIZE": 256}, num_warps=4),
],
key=[],
)
@triton.jit
def add_kernel_tuned(x_ptr, y_ptr, out_ptr, n, BLOCK_SIZE: "tl.constexpr"):
... # 内核体与上面相同
调用时不用再手动传 BLOCK_SIZE,autotune 自动接管。torch.compile 完整支持这种用法,两者可以叠加。
让内核成为「一等公民」:triton_op
直接写内核有个缺点:PyTorch 的很多子系统不认识它。CPU 上没有回退实现,反向传播不会自动求导,性能计数器也统计不到它的计算量。自己单打独斗可以,要给别人用就不行。
torch.library.triton_op 就是为补上这些而生。它把 Triton 内核包装成一个正式算子,同时 torch.compile 仍能追踪到内核内部做优化——这是它和 custom_op(对编译器是黑盒)最大的区别。先定义内核,再包装:
from torch.library import triton_op, wrap_triton
@triton.jit
def sin_kernel(x_ptr, out_ptr, n, BLOCK_SIZE: "tl.constexpr"):
pid = tl.program_id(axis=0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n
x = tl.load(x_ptr + offsets, mask=mask)
tl.store(out_ptr + offsets, tl.sin(x), mask=mask)
@triton_op("mylib::mysin", mutates_args={})
def mysin(x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
n = x.numel()
wrap_triton(sin_kernel)[(n,)](x, out, n, BLOCK_SIZE=4)
return out
包装之后就能像普通算子一样调用,还能通过 torch.ops.mylib.mysin 访问:
y = mysin(x) # 直接调用
z = torch.ops.mylib.mysin.default(x) # 走算子注册表
torch.testing.assert_close(y, x.sin())
再补三样东西,这个算子才算完整:
- CPU 回退:Triton 内核跑不了 CPU,用
mysin.register_kernel("cpu")注册一个普通 PyTorch 实现兜底。写清楚回退逻辑,代码才能在无 GPU 的机器上正常加载; - 反向传播:用
mysin.register_autograd(backward, setup_context=setup_context)注册求导公式,模型就能用这个算子训练。反向公式必须由 PyTorch 能理解的算子组成; - FLOP 统计:用
register_flop_formula(torch.ops.mylib.mysin)告诉性能计数器这算子算了多少浮点运算。
这三样不是写内核时就必须做,而是「交付给别人用」之前必须补。自己本地实验,跳过也能跑。
常见坑
我踩过和见人踩过的,列三个:
- 在 CPU 上调用内核:Triton 内核没有 CPU 实现,忘了做回退就会直接报错。所以正式交付的算子,都会像上面那样注册一个 CPU 回退;
- mask 写漏:向量长度不整除 BLOCK_SIZE 时,最后一块会越界。漏了 mask,有时不报错,只是读进来脏数据、写出去乱内存;
- heuristics 和 autotune 的顺序:
triton.heuristics必须放在triton.autotune之前使用,顺序反了会失效。
另外记住:torch.compile 对 autotune 的支持只覆盖 configs、keys、restore_value、reset_to_zero 这几个参数,别把复杂逻辑塞进其余参数里。
C++/CUDA 扩展:另一条路
当性能要求极致,或者要对接非 Python 环境时,可以写 C++/CUDA 扩展。流程是:用 torch.utils.cpp_extension 编译 .cpp/.cu 文件,TORCH_LIBRARY 宏注册算子,Python 侧用 torch.library.register_fake 补元信息,最后用 torch.library.opcheck 验证注册正确。PyTorch 2.10 起还提供稳定 ABI,编译一次能在多个 PyTorch 版本上跑。这条路线门槛高,入门阶段了解存在即可,细节留给官方手册。
Tip验证内核算得对不对,用
torch.testing.assert_close和 PyTorch 原生实现对比;验证注册规范不规范,跑一遍torch.library.opcheck。数值对、注册对,自定义算子才算真正交付。
什么时候不值得写
说句泼冷水的话:自定义算子是最容易被过度使用的优化手段。模型慢,八成不是因为缺算子,而是数据加载、设备搬运、训练循环写得不讲究。动手写内核之前,先问自己三个问题:现有算子真的拼不出来吗?这个算子真的是瓶颈吗?收益能覆盖维护成本吗?三个都答「是」,再动手。
Triton 的定位是「性能优化工具箱」:平时搁着,需要时拿出来用。会用、知道什么时候用,比会写更重要。
小结
自定义算子是「性能不够用」时的最后武器。优先用现有算子拼;拼不出来再上 Triton;要长期维护、给团队用,就包装成 triton_op;极致性能或跨语言场景,再考虑 C++/CUDA。下一章我们看模型怎么导出——把训练好的成果带出 Python,走向部署。