如何为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
相关产品推荐
相关产品推荐

