为什么需要一条明确的自定义算子路径?
自定义算子真正难的不是把一个函数注册进命名空间,而是让它在 eager、自动微分、图捕获和后端执行之间保持同一份语义。仓库最新实现把这条路径拆成注册、分发、微分和原生下沉四个可检查的边界。
tensorplay.library 提供 custom_op、triton_op、wrap_triton,以及 register_kernel、register_fake 和 register_autograd。这样一个算子既可以先用 Python 实现验证数值,再替换成设备专用 kernel,而不必改写上层模型。
三个层面:注册、执行与微分
- 注册层:用命名空间和 schema 标识算子,提供默认实现与设备专用 kernel。
- 执行层:eager 直接调用 kernel;图捕获时生成一个可识别的 custom_op 节点。
- 微分层:通过 register_autograd 提供 setup_context 与 backward,不把梯度规则硬编码进计算核心。
@library.custom_op("demo::scaled", mutates_args=())
def scaled(x):
return tp.mul(x, 2.0)
@scaled.register_autograd
def backward(ctx, grad_out):
return (tp.mul(grad_out, 2.0),)Triton:eager 可以直通,捕获必须有契约
wrap_triton 负责让用户编写的 Triton JIT kernel 进入 TensorPlay 的算子边界。eager 路径可以直接发射 kernel;当编译器捕获到裸 launch 时,系统会明确报 GraphCaptureError,引导开发者使用 triton_op,而不是把一次不透明的启动伪装成普通算子。
triton_op 在捕获期表现为一个原生 custom_op 节点:kernel 内部的 add 不会泄漏成额外图节点,也不会用 Python GraphModule 解释器冒充编译结果。它既是执行边界,也是融合屏障,后续 lowering 可以据此做正确的优化取舍。
- eager:真实 Triton kernel 直接执行。
- capture:一个可定位的 opaque custom_op 节点。
- strict_native:没有原生执行产物时显式失败,而不是静默回退。
TVM:把点算子链下沉为 TIR
TensorPlay 的 TVM 后端以 backend="tvm" 注册,当前针对白名单中的 pointwise 算子,把多个节点组合成一个 TIR 计算并通过 DLPack 连接输入输出。它保留后端契约:不支持的 custom op 可以回到解释执行,但结果仍必须与 eager 路径一致。
这条路线的价值是把“编译后是什么”变成可以检查的后端选择:是 Stax 原生图、Triton kernel、TVM TIR,还是明确的 fallback。开发者不需要从一条模糊的加速开关猜测实际发生了什么。
验证:先看数值,再看图和执行器
自定义算子不能只凭一次输出正确就宣布完成。仓库测试覆盖默认 kernel、设备专用分发、fake/meta、autograd,以及 Triton 捕获时“单个 opaque 节点且函数体不执行”的边界。TVM 测试还检查 pointwise parity、shape 变化和训练区梯度。
- 对 eager 与编译结果做 allclose,并检查 requires_grad 与 backward。
- 检查图里是否只有预期的 custom_op 节点,避免内部算子泄漏。
- 检查 codegen 标记,确认结果确实来自 native、Triton 或 TVM 路径。
- 对不支持的后端能力验证 fallback 正确性,而不是把 fallback 当作性能结论。
