You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为PyTorch转ONNX导出生成Dummy算子并注册符号函数

解决PyTorch导出ONNX自定义算子的问题

要让TorchScript将你的my_dummy_fn视为原子算子而非追踪内部逻辑,同时正确注册对应的ONNX符号函数,你需要通过注册自定义算子命名空间的方式实现,具体步骤如下:

完整可运行代码示例

import torch

# 1. 定义并注册自定义算子
torch.library.define("my_ops::my_dummy_fn", "(Tensor first, Tensor second) -> (Tensor, Tensor, Tensor)")

@torch.library.impl("my_ops::my_dummy_fn", "default")
def my_dummy_fn_impl(first: torch.Tensor, second: torch.Tensor):
    # Python端的兜底实现,导出ONNX时会被符号函数替换
    return torch.tensor(0), torch.tensor([1]), torch.tensor([2])

# 2. 定义ONNX符号生成函数
def my_symbolic_function(g, first, second):
    # 生成自定义ONNX节点,需指定输出数量与算子实现一致
    return g.op("my_custom_onnx", first, second, outputs=3)

# 3. 绑定符号函数到自定义算子
torch.onnx.register_custom_op_symbolic(
    symbolic_name="my_ops::my_dummy_fn",  # 对应自定义算子的完整命名空间路径
    symbolic_fn=my_symbolic_function,
    opset_version=12,
)

# 测试导出:封装模型并生成ONNX文件
class TestModel(torch.nn.Module):
    def forward(self, x, y):
        # 调用注册好的自定义算子
        return torch.ops.my_ops.my_dummy_fn(x, y)

model = TestModel()
dummy_input = (torch.randn(1), torch.randn(1))
torch.onnx.export(model, dummy_input, "custom_op_model.onnx", opset_version=12)

核心要点说明

  • torch.library.define:声明自定义算子的签名,告知TorchScript这是一个原子算子,不会追踪其内部实现代码。
  • 算子命名空间:my_ops::my_dummy_fn是自定义算子的唯一标识,注册符号函数时的symbolic_name必须完全匹配该路径。
  • 符号函数输出数量:通过outputs=3指定输出张量的个数,必须与算子实现的返回值数量一致,否则导出会报错。
  • 避免@torch.jit.script:该装饰器会解析函数体生成追踪代码,无法触发自定义符号函数的执行。

旧版本PyTorch兼容方案(1.12及以下)

如果使用未支持torch.library的PyTorch版本,可通过torch.ops手动注册:

import torch

# 初始化自定义算子命名空间
torch.ops.load_library("")  # 空字符串表示注册Python实现的算子
torch.ops.my_ops.my_dummy_fn = torch._ops.OpFunction(torch._ops._OpNamespace("my_ops"), "my_dummy_fn")

@torch.ops.my_ops.my_dummy_fn.impl
def my_dummy_fn_impl(first: torch.Tensor, second: torch.Tensor):
    return torch.tensor(0), torch.tensor([1]), torch.tensor([2])

# 后续符号注册与导出逻辑同前

内容的提问来源于stack exchange,提问作者GPhilo

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.26 04:37:27