Triton自定义算子在PyTorch中不自动生效,因torch.compile默认使用Inductor而非Triton,需显式封装为torch.autograd.Function并手动实现前向与反向;输入须为cuda设备、contiguous且dtype对齐,测时需同步GPU,profiling需确认kernel真实启动。

为什么 Triton 自定义算子在 PyTorch 里不直接生效
PyTorch 本身不内置 Triton 运行时,torch.compile 默认后端是 Inductor,不是 Triton。即使你写了 Triton kernel(@triton.jit 函数),它不会自动被 torch.compile 拾取或替换原生算子——必须显式调用、封装成 torch.autograd.Function 或注册为自定义算子才能参与训练流程。
常见错误现象:RuntimeError: Triton kernel not found 或 kernel 执行无加速效果,本质是没走 GPU kernel 调度路径,而是回退到了 CPU Python 循环或未编译的 PyTorch 原生实现。
- 确认已安装匹配版本:
pip install triton(注意:PyTorch 2.3+ 内置的torch._inductor.triton是私有模块,不建议直接用;应使用独立triton包) - Triton kernel 必须在 CUDA 设备上运行,且输入
Tensor的device为cuda,dtype需与 kernel 中tl.float16/tl.bfloat16显式对齐 - 避免在 kernel 内做 Python 控制流(如
if x > 0:),Triton 只支持tl.where或 warp-level 分支
如何把 Triton kernel 封装成可反向传播的 torch.Tensor 算子
不能只写 @triton.jit 函数,必须继承 torch.autograd.Function 并实现 forward 和 backward —— Triton 本身不提供自动微分,反向逻辑得你手写 kernel 或调用已有梯度公式。
示例场景:实现一个自定义的 silu(Sigmoid Linear Unit)激活函数,但用 Triton 手写前向+反向,绕过 PyTorch 原生算子调度开销:
class TritonSilu(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
output = torch.empty_like(x)
grid = lambda META: (triton.cdiv(x.numel(), META['BLOCK_SIZE']),)
_silu_kernel[grid](x, output, x.numel(), BLOCK_SIZE=1024)
ctx.save_for_backward(x)
return output
<pre class="brush:php;toolbar:false;">@staticmethod
def backward(ctx, grad_output):
x, = ctx.saved_tensors
grad_input = torch.empty_like(x)
grid = lambda META: (triton.cdiv(x.numel(), META['BLOCK_SIZE']),)
_silu_grad_kernel[grid](x, grad_output, grad_input, x.numel(), BLOCK_SIZE=1024)
return grad_input-
ctx.save_for_backward保存的是原始Tensor,不是 Triton kernel 输出结果;反向 kernel 输入必须和前向一致(例如这里需要x,因为silu'(x) = silu(x) + x * (1 - silu(x))) - 所有 kernel 调用前需确保输入
Tensor已 contiguous(x.contiguous()),否则data_ptr()可能出错 - 若 kernel 含 shared memory,
num_stages和num_warps参数会影响 bank conflict 和 occupancy,调试时先设num_warps=4保稳定
torch.compile 能否自动用 Triton 替换你的算子
不能。截至 PyTorch 2.4,torch.compile 的 Inductor 后端会生成自己的 Triton kernel(在 /tmp/torch_inductor_*/ 下可见),但它不会识别或复用用户写的 @triton.jit 函数。Inductor 生成的 kernel 是封闭的、不可干预的。
如果你的目标是“让 torch.compile 加速某段逻辑”,正确路径是:写标准 PyTorch 代码(用原生 torch ops),然后交给 torch.compile(model, backend="inductor");而不是手动插 Triton kernel。
- 想验证是否用了 Inductor 生成的 Triton:设置环境变量
TORCHINDUCTOR_DUMP_ASSETS=1,运行后检查/tmp/torch_inductor_*/下是否有*triton*文件 - 混合使用场景可行:用
torch.compile加速主干网络,再用自定义 Triton 算子替换其中某个瓶颈 op(如 custom attention),但二者要明确隔离,避免torch.compile尝试重写你的autograd.Function - Inductor 对
torch.autograd.Function默认跳过优化,所以你封装的 Triton 算子不会被二次编译 —— 这其实是优点,保证了控制权
容易被忽略的设备同步与 profiling 陷阱
Triton kernel 启动是异步的,但很多用户在测 latency 时直接用 time.time() 或没加 torch.cuda.synchronize(),导致测出来的时间远小于真实 GPU 执行耗时。
- 正确测时:在 kernel 调用前后加
torch.cuda.synchronize(),或用torch.cuda.Event计时 - kernel 报错信息极简(如
illegal memory access),实际常因指针越界(pid * BLOCK_SIZE + offset >= n_elements)或 dtype 不匹配(输入是torch.float32,kernel 却用tl.float16load) - 用
nsys profile --trace=cuda,nvtx python script.py查看 kernel 名称是否含triton前缀,确认是否真跑在 Triton 上;若显示void kernel或空名称,说明 kernel 未成功 launch
最麻烦的点往往不在 kernel 逻辑本身,而在 tensor layout、stride 处理、以及 autograd context 生命周期管理——这些细节不爆错,但会让梯度无声变零或数值异常。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











