概览
本教程说明如何为尚未支持的 PyTorch 算子创建 ONNX 实现,或用自己的实现替换现有实现。需要扩展导出器支持的场景有三种:
- 覆盖已有 PyTorch 算子的实现;
- 使用自定义 ONNX 算子;
- 支持自定义 PyTorch 算子。
完成教程后,将了解如何覆盖或增加 PyTorch 算子的 ONNX 支持、为特定运行时接入自定义 ONNX 算子,以及如何实现自定义 PyTorch 算子并将它转换成 ONNX。
前提条件
开始之前,需要满足以下条件:
torch >= 2.6;- 明确目标 PyTorch 算子;
- 已经学习 ONNX Script 教程;
- 准备好该算子的 ONNX Script 实现。
覆盖现有 PyTorch 算子的实现
ONNX 导出器团队努力支持所有 PyTorch 算子,但仍有算子可能尚未实现。本节展示向导出器提供这类算子的实现的方法。
实现尚未支持的算子,与用自定义版本替换现有算子的步骤相同。教程没有选择一个真正缺少实现的算子,而是用相同流程覆盖 torch.ops.aten.add.Tensor,模拟需要补充算子实现的情况。
模型因不支持某个算子而无法导出时,可能出现这样的错误:
No decompositions registered for [...]
这个例子中的目标是 torch.ops.aten.add.Tensor。它的类型是 torch._ops.OpOverload;注册自定义实现时,使用这个具体的算子重载作为目标。
import torch
import onnxscript
# Opset 18 is the standard supported version as of PyTorch 2.6
from onnxscript import opset18 as op
# Create a model that uses the operator torch.ops.aten.add.Tensor
class Model(torch.nn.Module):
def forward(self, input_x, input_y):
return torch.ops.aten.add.Tensor(input_x, input_y)
# NOTE: The function signature (including parameter names) must match the signature of the unsupported PyTorch operator.
# https://github.com/pytorch/pytorch/blob/main/aten/src/ATen/native/native_functions.yaml
# All attributes must be annotated with type hints.
def custom_aten_add(self, other, alpha: float = 1.0):
if alpha != 1.0:
alpha = op.CastLike(alpha, other)
other = op.Mul(other, alpha)
# To distinguish the custom implementation from the builtin one, we switch the order of the inputs
return op.Add(other, self)
x = torch.tensor([1.0])
y = torch.tensor([2.0])
# Then we provide the custom implementation to the ONNX exporter as a ``custom_translation_table``.
onnx_program = torch.onnx.export(
Model().eval(),
(x, y),
dynamo=True,
custom_translation_table={
torch.ops.aten.add.Tensor: custom_aten_add,
},
)
# Optimize the ONNX graph to remove redundant nodes
onnx_program.optimize()
官方示例导出日志:
[torch.onnx] Obtain model graph for `Model()` with `torch.export.export(..., strict=False)`...
[torch.onnx] Obtain model graph for `Model()` with `torch.export.export(..., strict=False)`... ✅
[torch.onnx] Run decompositions...
[torch.onnx] Run decompositions... ✅
[torch.onnx] Translate the graph into ONNX...
[torch.onnx] Translate the graph into ONNX... ✅
[torch.onnx] Optimize the ONNX graph...
[torch.onnx] Optimize the ONNX graph... ✅
现在检查模型,确认它使用了自定义实现:
print(onnx_program.model)
官方示例的模型输出:
<
ir_version=10,
opset_imports={'': 20},
producer_name='pytorch',
producer_version='2.14.0+cu130',
domain=None,
model_version=None,
>
graph(
name=main_graph,
inputs=(
%"input_x"<FLOAT,[1]>,
%"input_y"<FLOAT,[1]>
),
outputs=(
%"add"<FLOAT,[1]>
),
) {
0 | # node_add
%"add"<FLOAT,[1]> ⬅️ ::Add(%"input_y", %"input_x")
return %"add"<FLOAT,[1]>
}
加法节点中,input_y 出现在前面,input_x 出现在后面,说明采用了交换输入顺序的自定义实现。原文叙述提到的节点名在当前输出中显示为 node_add;节点名称可以随导出器版本变化。
直接对输入张量调用 torch.onnx.ONNXProgram,可以通过 ONNX Runtime 执行模型并核对结果:
result = onnx_program(x, y)[0]
torch.testing.assert_close(result, torch.tensor([3.0]))
使用自定义 ONNX 算子
在这一场景中,模型使用标准 PyTorch 算子,但运行时(例如微软的 ONNX Runtime)提供了专用内核,可以在导出时替换原有实现。
下面使用 ONNX Runtime 提供的 com.microsoft.Gelu。它与 ONNX 规范中的标准 Gelu 算子不是同一个算子。
class GeluModel(torch.nn.Module):
def forward(self, input_x):
return torch.ops.aten.gelu(input_x)
# Create a namespace for the custom operator using ONNX Script
# ``com.microsoft`` is an official ONNX Runtime namespace
microsoft_op = onnxscript.values.Opset(domain="com.microsoft", version=1)
# NOTE: The function signature (including parameter names) must match the signature of the unsupported PyTorch operator.
# https://github.com/pytorch/pytorch/blob/main/aten/src/ATen/native/native_functions.yaml
# NOTE: All attributes must be annotated with type hints.
# The function must be scripted using the ``@onnxscript.script()`` decorator when
# using operators from custom domains. This may be improved in future versions.
from onnxscript import FLOAT
@onnxscript.script(microsoft_op)
def custom_aten_gelu(self: FLOAT, approximate: str = "none") -> FLOAT:
return microsoft_op.Gelu(self)
onnx_program = torch.onnx.export(
GeluModel().eval(),
(x,),
dynamo=True,
custom_translation_table={
torch.ops.aten.gelu.default: custom_aten_gelu,
},
)
# Optimize the ONNX graph to remove redundant nodes
onnx_program.optimize()
官方示例导出日志:
[torch.onnx] Obtain model graph for `GeluModel()` with `torch.export.export(..., strict=False)`...
[torch.onnx] Obtain model graph for `GeluModel()` with `torch.export.export(..., strict=False)`... ✅
[torch.onnx] Run decompositions...
[torch.onnx] Run decompositions... ✅
[torch.onnx] Translate the graph into ONNX...
[torch.onnx] Translate the graph into ONNX... ✅
[torch.onnx] Optimize the ONNX graph...
[torch.onnx] Optimize the ONNX graph... ✅
检查模型,确认它使用 com.microsoft 命名空间中的 Gelu:
print(onnx_program.model)
官方示例的模型输出:
<
ir_version=10,
opset_imports={'com.microsoft': 1, '': 20},
producer_name='pytorch',
producer_version='2.14.0+cu130',
domain=None,
model_version=None,
>
graph(
name=main_graph,
inputs=(
%"input_x"<FLOAT,[1]>
),
outputs=(
%"gelu"<FLOAT,[1]>
),
) {
0 | # n0
%"gelu"<FLOAT,[1]> ⬅️ com.microsoft::Gelu(%"input_x")
return %"gelu"<FLOAT,[1]>
}
与前一个示例一样,通过 ONNX Runtime 执行模型并核对结果:
result = onnx_program(x)[0]
torch.testing.assert_close(result, torch.ops.aten.gelu(x))
该替换函数接受 approximate 参数,但示例实现没有根据它切换算法;不能直接把这里的验证推及 approximate="tanh" 等其他语义。目标运行时也必须支持 com.microsoft 域中的相应算子。
支持自定义 PyTorch 算子
这一场景的目标算子由用户实现并注册到 PyTorch。下面的自定义算子接收一个张量,先将输入与自身相加,再对结果取整,最终返回一个张量。
首先,用 torch.library.custom_op() 实现和注册算子。创建算子的详细步骤可参考 Creating new custom ops in Python。
# Define and use the operator in PyTorch
@torch.library.custom_op("mylibrary::add_and_round_op", mutates_args=())
def add_and_round_op(input: torch.Tensor) -> torch.Tensor:
return torch.round(input + input)
@add_and_round_op.register_fake
def _add_and_round_op_fake(tensor_x):
return torch.empty_like(tensor_x)
class AddAndRoundModel(torch.nn.Module):
def forward(self, input):
return add_and_round_op(input)
# Implement the custom operator in ONNX using ONNX Script
def onnx_add_and_round(input):
return op.Round(op.Add(input, input))
onnx_program = torch.onnx.export(
AddAndRoundModel().eval(),
(x,),
dynamo=True,
custom_translation_table={
torch.ops.mylibrary.add_and_round_op.default: onnx_add_and_round,
},
)
# Optimize the ONNX graph to remove redundant nodes
onnx_program.optimize()
print(onnx_program)
官方示例的导出日志与模型输出:
[torch.onnx] Obtain model graph for `AddAndRoundModel()` with `torch.export.export(..., strict=False)`...
[torch.onnx] Obtain model graph for `AddAndRoundModel()` with `torch.export.export(..., strict=False)`... ✅
[torch.onnx] Run decompositions...
[torch.onnx] Run decompositions... ✅
[torch.onnx] Translate the graph into ONNX...
[torch.onnx] Translate the graph into ONNX... ✅
[torch.onnx] Optimize the ONNX graph...
[torch.onnx] Optimize the ONNX graph... ✅
ONNXProgram(
model=
<
ir_version=10,
opset_imports={'': 20},
producer_name='pytorch',
producer_version='2.14.0+cu130',
domain=None,
model_version=None,
>
graph(
name=main_graph,
inputs=(
%"input"<FLOAT,[1]>
),
outputs=(
%"add_and_round_op"<FLOAT,[1]>
),
) {
0 | # node_Add_0
%"val_0"<FLOAT,[1]> ⬅️ ::Add(%"input", %"input")
1 | # node_add_and_round_op
%"add_and_round_op"<FLOAT,[1]> ⬅️ ::Round(%"val_0")
return %"add_and_round_op"<FLOAT,[1]>
}
,
exported_program=
ExportedProgram:
class GraphModule(torch.nn.Module):
def forward(self, input: "f32[1]"):
input_1 = input
# File: /var/lib/workspace/beginner_source/onnx/onnx_registry_tutorial.py:215 in forward, code: return add_and_round_op(input)
add_and_round_op: "f32[1]" = torch.ops.mylibrary.add_and_round_op.default(input_1); input_1 = None
return (add_and_round_op,)
Graph signature:
# inputs
input: USER_INPUT
# outputs
add_and_round_op: USER_OUTPUT
Range constraints: {}
)
这里用自定义翻译函数,将 torch.export.ExportedProgram 中的 torch.ops.mylibrary.add_and_round_op.default 翻译为 ONNX 的 Add 与 Round。
最后核对结果:
result = onnx_program(x)[0]
torch.testing.assert_close(result, add_and_round_op(x))
结语
本教程介绍了 custom_translation_table 选项,以及如何用 ONNX Script 为已有或尚未支持的 PyTorch 算子创建自定义实现。示例还使用 ONNX Runtime 执行模型并与 PyTorch 结果比较,形成了处理 ONNX 导出算子支持问题的完整流程。
继续阅读与代码下载
下面的教程涵盖基础例子到高级场景,顺序并非学习顺序。可以直接选择感兴趣的主题:
官方页面记录的脚本总运行时间为 3.331 秒,属于官方构建时的结果,并非本文在当前设备上的测量。
来源:PyTorch:Extending the ONNX Exporter Operator Support,作者 Ti-Tai Wang、Justin Chu。官方标记:创建于 2023-10-06,最后更新 2025-03-05,最后验证 2024-11-05;2026-10-03 所读页面属于 2.14.0+cu130 文档。中文版本翻译了正文,补充了节点名差异与 GELU 适用范围;全部 14 个代码与示例输出块保留原文。代码只完成静态语法核验,没有执行 ONNX 导出或运行时数值测试。
原项目版权、BSD 3-Clause 条件与免责声明
BSD 3-Clause License Copyright (c) 2017-2022, Pytorch contributors All rights reserved. Redistribution and use in source and binary forms, with or without modification, are permitted provided that the following conditions are met: * Redistributions of source code must retain the above copyright notice, this list of conditions and the following disclaimer. * Redistributions in binary form must reproduce the above copyright notice, this list of conditions and the following disclaimer in the documentation and/or other materials provided with the distribution. * Neither the name of the copyright holder nor the names of its contributors may be used to endorse or promote products derived from this software without specific prior written permission. THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.











暂无评论内容