TensorPlay AI
博客

从自定义算子到原生图:Triton 与 TVM 如何接入 TensorPlay

让自定义 kernel 保持可执行、可捕获、可验证,而不是在编译阶段悄悄退回 Python 解释器。

为什么需要一条明确的自定义算子路径?

自定义算子真正难的不是把一个函数注册进命名空间,而是让它在 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 当作性能结论。
Ask DeepWiki