一、先聊聊这个痛点
大家有没有遇到过这种情况:自己辛辛苦苦训好的模型,正准备部署到服务端,结果转ONNX的时候,“啪”一下报错了。错误信息里写着什么“Unsupported operator”,瞬间头大。尤其是当你用了PyTorch里一些很顺手、但ONNX不认识的API,比如torch.roll,又或者自己写了一个自定义层,那基本就是一场灾难。
我当年第一次遇到这种问题,第一反应是“ONNX怎么这么笨,明明PyTorch都跑得好好的”。后来才明白,ONNX不是笨,它是“约定优先”。它只认识一套大家共同遵守的算子集合,而不认识某个特定框架的“土办法”。这就像你带着方言去参加普通话比赛,评委听不懂不是评委的问题,是你没提前准备“翻译器”。
不过,问题总有解决办法。今天我们就从一个具体的场景出发,聊聊怎么用“自定义算子注册 + 梯度重写 + 兼容性声明”这三板斧,把一个不被支持的算子变得可导出、可移植,而且之后的模型完全不依赖PyTorch。
二、ONNX到底在做什么
要解决问题,先得知道ONNX的工作方式。ONNX说白了是一个“中间格式”。它包含了一套固定算子,比如卷积、全连接、激活函数、张量切片等等。PyTorch的模型,需要通过torch.onnx.export把网络里的每一个操作映射到ONNX的这套算子上。
这个映射过程,就叫“导出”。如果某个PyTorch操作能直接找到对应的ONNX算子,那就愉快地转换;如果找不到,就会报错,告诉你“我不认识这个操作”。
最要命的是,有些操作虽然没有现成的ONNX算子,但它能用几个已有的ONNX算子拼出来。比如torch.roll,把一个张量沿某个维度循环移动。你可以用Slice切两半,再用Concat拼接回来。这就是我们要做的事。
此时,有三件事必须要做:
- 告诉ONNX“如何用它的语言来翻译我们的自定义算子”,这就是自定义算子注册。
- 保证这个算子在训练时梯度正确,毕竟很多模型导出后还可能继续迁移学习,这就是梯度重写。
- 让导出的模型自带“说明”,让下游用户知道它依赖什么算子集、什么版本,这就是兼容性声明。
这三件事,缺一不可。
三、解决思路的三个零件
3.1 自定义算子注册
在PyTorch中,我们可以给torch.autograd.Function添加一个symbolic静态方法。这个方法就是“翻译官”,告诉ONNX导出器:“当遇到我这个算子时,你把它拆成哪些标准的ONNX操作”。有了这个方法,ONNX就能顺利导出。
如果你使用的是PyTorch自带的、但它不知道如何导出的算子(比如某些aten操作),也可以使用torch.onnx.register_custom_op_symbolic从外部注册。这两种方式本质是一样,都是把“自定义算子”和“ONNX表达式”绑定到一起。
3.2 梯度重写
很多朋友觉得“部署模型又不训练,梯度有什么好写的?”话是这么说,但PyTorch在导出的时候会构建一次反向传播图来检查一致性。如果你的自定义函数没有实现backward,导出过程可能都会报错,更别提以后你还可能拿着这个模型在别的框架里做微调。
所以,我们要在自定义Function中老老实实实现backward,把梯度传回去。这就叫梯度重写。
3.3 兼容性声明
ONNX模型是给人看的,也是给机器跑的。为了让对方明白这个模型需要什么环境,我们要在导出时指定opset_version(算子集版本),如果使用了自定义域名(比如com.example),还要在模型里记录这些信息。这样,不管是ONNX Runtime还是TensorRT,都能知道该怎么处理。
四、上手:把一个“循环移位”算子搬进ONNX
我们选一个非常有代表性的算子:torch.roll。它在信号处理、数据增强里很常用,而ONNX至今没有对应的原生算子。我们就拿它开刀。
4.1 定义自定义的PyTorch算子(含梯度)
先定义一个RollFunction,继承torch.autograd.Function。它就像我们家的“改造车间”,前向推理用torch.roll完成,反向传播就做一个反方向的移动。
# 技术栈:Python 3.8 / PyTorch 1.13
import torch
class RollFunction(torch.autograd.Function):
"""
自定义“循环移位”算子。
功能:把输入 x 在指定的维度 dim 上循环移动 shift 步。
例如:x=[1,2,3,4,5] shift=2 -> [4,5,1,2,3]
"""
@staticmethod
def forward(ctx, x, shift, dim):
"""
前向计算。ctx用来保存反向传播需要的参数。
shift 和 dim 是整数,不是张量。
"""
ctx.shift = shift
ctx.dim = dim
return torch.roll(x, shifts=shift, dims=dim)
@staticmethod
def backward(ctx, grad_output):
"""
反向传播。循环移动的梯度就是“再反向移动一次”。
例如,如果前向向右移了2,梯度就向左移2。
"""
grad_input = torch.roll(grad_output, shifts=-ctx.shift, dims=ctx.dim)
# shift 和 dim 是整数常量,没有梯度,所以返回 None
return grad_input, None, None
@staticmethod
def symbolic(g, x, shift, dim):
"""
注册到 ONNX:把 Roll 拆开成 Slice + Concat。
这样导出的模型里没有任何自定义节点,任何平台都能直接跑。
"""
# 先把需要用的常量包装成 ONNX Constant 节点
zero = g.op("Constant", value_t=torch.tensor([0], dtype=torch.int64))
neg_shift = g.op("Constant", value_t=torch.tensor([-shift], dtype=torch.int64))
big_num = g.op("Constant", value_t=torch.tensor([2**31 - 1], dtype=torch.int64))
axes = g.op("Constant", value_t=torch.tensor([dim], dtype=torch.int64))
# 第一部分:从位置 0 切到 倒数第 shift 个元素之前
part1 = g.op("Slice", x, zero, neg_shift, axes)
# 第二部分:从 倒数第 shift 个元素 一直切到末尾
part2 = g.op("Slice", x, neg_shift, big_num, axes)
# 把两部分顺序调换拼起来,得到循环移位结果
return g.op("Concat", part2, part1, axis_i=dim)
注意,这里的symbolic方法就是“自定义算子注册”的核心。它让ONNX在导出时调用这个函数,生成标准的Slice和Concat节点。这样一来,模型就不会指向任何PyTorch独有操作了。
4.2 构造一个使用了该算子的模型
为了演示,我们写一个特别简单的全连接网络,在输入后面接一个循环移位层。实际项目中,这个操作完全可以放在网络中间,比如做时间序列的平移特征。
# 技术栈:Python 3.8 / PyTorch 1.13
import torch.nn as nn
class ShiftNet(nn.Module):
"""
一个包含了自定义循环移位层的简单神经网络。
输入是 (B, 8),输出是 (B, 8)。
"""
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(8, 16) # 第一层全连接
self.fc2 = nn.Linear(16, 8) # 第二层全连接
def forward(self, x):
# 在维度1上循环移动2步,然后经过全连接层
x = RollFunction.apply(x, shift=2, dim=1)
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
这里要特别说明一下:RollFunction.apply的shift和dim最好写成固定值,因为在symbolic里我们需要它们作为常量。如果你非要用计算出来的值,也可以,但复杂度会高很多,我们建议先固定,性能往往也更好。
4.3 导出模型到ONNX
一切就绪,我们用torch.onnx.export把模型导出来。记得要把模型切到eval模式,不然中间会有Dropout之类的随机节点,导出结果不稳定。
# 技术栈:Python 3.8 / PyTorch 1.13 / onnx 1.13
model = ShiftNet()
model.eval() # 切换到推理模式
# 构造一个形状为 (1, 8) 的假数据,用于追踪计算图
dummy_input = torch.randn(1, 8)
# 导出 ONNX
torch.onnx.export(
model, # 要导出的模型
dummy_input, # 示例输入
"shift_net.onnx", # 输出文件名
input_names=["input"], # 输入节点名字
output_names=["output"], # 输出节点名字
opset_version=13, # 选择一个较新的算子集版本
do_constant_folding=True # 打开常量折叠,让图更干净
)
print("导出完成!")
导出之后,我们可以用onnx库来检查一下,确保模型格式合法。接着,再用ONNX Runtime跑一次,跟PyTorch的结果做对比,确认我们的“翻译”没有出错。
4.4 验证导出的模型
# 技术栈:Python 3.8 / onnx 1.13 / onnxruntime 1.15 / numpy
import onnx
import onnxruntime as ort
import numpy as np
# 1. 检查模型的合法性
onnx_model = onnx.load("shift_net.onnx")
onnx.checker.check_model(onnx_model) # 如果有格式问题,这里会抛异常
# 2. 用ONNX Runtime跑一遍
input_data = np.random.randn(1, 8).astype(np.float32) # 生成随机输入
sess = ort.InferenceSession("shift_net.onnx")
onnx_output = sess.run(None, {"input": input_data})[0] # 推理得到输出
# 3. 和PyTorch的输出做对比
with torch.no_grad():
pt_output = model(torch.from_numpy(input_data)).numpy()
max_err = np.max(np.abs(onnx_output - pt_output))
print(f"ONNX与PyTorch最大误差: {max_err:.6e}")
# 如果误差在1e-6级别,说明转换基本没问题
assert max_err < 1e-5, "误差太大,请检查转换逻辑"
到这里,一个不支持的算子就搞定啦。你可以在别的推理引擎里直接加载shift_net.onnx,再也不用每天守着PyTorch那台服务器了。
4.5 如果遇上的不是roll,怎么办
你可能已经发现,这套思路其实是通用的。以后不管遇到什么不支持的算子,都可以按这个顺序来:
- 先看看这个操作能不能拆成几个ONNX基础算子,比如Slice、Concat、Reshape、MatMul。
- 如果可以,就写一个
symbolic方法,告诉导出器怎么拆。 - 如果不可以,就得注册一个自定义节点,并准备一个对应的运行时实现。
不管哪条路,你都需要先把算子写成一个torch.autograd.Function,把前向和反向都实现好。这是整个流程的“地基”。
五、兼容性声明到底怎么写
有的同学会问:“你这个正好能用Slice和Concat拼出来。万一我遇到一个拼不出来的算子呢?”
这种情况也很常见。比如你要自己实现一个高效的非极大值抑制,或者一个特殊的稀疏注意力。ONNX基础算子确实没法表达,那就只能注册成一个“自定义节点”,给它一个独享的域名,比如com.example.roll。
做法是:在symbolic里返回一个自定义节点,类似这样:
# 技术栈:Python 3.8 / PyTorch 1.13
def my_special_symbolic(g, x, other_param):
# 创建一个自定义domain下的算子
return g.op("com.example.MySpecialOp", x, other_param,
domain="com.example", # 自定义命名空间
version=1) # 维护版本
但这样做,你在部署端也必须提供一个能解释com.example.MySpecialOp的算子实现,比如给ONNX Runtime写一个自定义Op插件。这就相当于又绑定了另一个框架,违背了我们“可移植”的初衷。
所以,我的建议是:优先尝试用基础算子组合,实在不行再用自定义节点。如果只能用自定义节点,那么一定要在模型里把自定义算子的domain、version和依赖库信息写清楚,用前检查好,避免别人拿到模型后一脸懵。
我们可以用onnx_model.metadata_props往模型里塞一段“说明书”:
# 技术栈:Python 3.8 / onnx 1.13
meta = onnx_model.metadata_props.add()
meta.key = "custom_op_info"
meta.value = "RollFunction has been expanded to Slice+Concat. No custom runtime required."
onnx.save(onnx_model, "shift_net_with_meta.onnx")
这段metadata虽然不是运行必需,但能提醒后续使用的人:“这个模型很干净,没有外部依赖。” 在团队协作或者开源时,这非常重要。
六、应用场景都在哪
说了这么多,这套方案实际用在哪里?我挑几个常见的:
1. 信号处理/时序模型。 比如语音识别里的趋势特征,经常需要把时间序列平移一下再拼接。torch.roll就很顺手,但ONNX不认。用我们的方法,导出后能轻松部署到手机端。
2. 数据增强。 做图像平移旋转的时候,有时候不想用已经封装好的torchvision.transforms,而是想自己控制边界。自写的循环填充就能用到roll。在训练时,PyTorch端用我们的RollFunction,导出后又能无缝迁移到ONNX Runtime。
3. 自定义激活函数。 很多人喜欢尝试新奇的激活函数,比如“带状态的分段函数”。这些函数往往由几个基础操作组合而成。如果不想把所有细节暴露在模型里,也可以注册成自定义算子。当然,还是那句话,能组合就不要新建节点。
4. 自定义采样/注意力模块。 一些最新论文里的算子,比如变形卷积、可变形注意力,ONNX通常不支持。我们同样可以用“先注册、再组合”的方式绕过去,只要最终能落在基础算子集合上,就皆大欢喜。
你可以把ONNX想象成一个国际机场,每个算子都是一架飞机。自己造的飞机没有航线许可,那我们就给它改装成已有型号,或者干脆申请一条新航线并在机场手册里写清楚。只要手续办全,照样能在全球飞行。
七、技术优缺点
7.1 优点
- 真正可移植:导出后的模型不带PyTorch依赖,可以在任何支持ONNX的平台上运行。
- 训练与推理一致:因为梯度重写正确,你在PyTorch里训练的模型,和部署到ONNX后的计算逻辑完全一致,误差极小。
- 运维成本低:不用为了一个算子搭一套自定义runtime,基础算子大家都有。
7.2 缺点
- 编写有门槛:你需要懂一点ONNX的符号层,会写
symbolic,还要熟悉基础算子的行为。 - 转换范围有限:不是所有逻辑都能用基础算子优雅地拼出来。碰上复杂逻辑,可能得写很多节点,性能反而下降。
- 调试较麻烦:ONNX的报错信息通常不友好,定位“哪里出问题”可能需要一点耐心。
八、注意事项
这部分很重要,能帮你少踩坑:
- 常量化参数:在
symbolic中,尽量把shift、dim写成Python整数或常量张量,不要搞成动态计算。动态计算会让你的转换函数复杂一百倍。 - 注意维度兼容:
Slice的axes、starts、ends必须严格匹配。尤其当dim为负数时,要提前转换成非负索引。 - 推导一下Shape:如果转换后的图里出现
Shape推断不出来的情况,会让后续优化失败。最好用torch.onnx.export后的静态输入,并打开do_constant_folding。 - 测梯度:不要只在推理模式下验证,建议用
torch.autograd.gradcheck跑一遍自定义算子的梯度,确保反向正确。 - 多版本opset测试:不同的opset对
Slice的行为略有差异,导出时选择一个版本后在目标运行环境上测试一下,别上来就选最新版。
九、文章总结
我们从一个“导出失败”的常见问题切入,看到了ONNX是一个“约定优先”的格式,也弄清楚了解决不支持的算子有三个要点:自定义算子注册、梯度重写、兼容性声明。然后我们拿torch.roll当例子,亲手实现了自定义Function,编写了symbolic方法,导出了模型,并用ONNX Runtime验证了结果。
整个过程听起来技术味很重,但实际操作起来并不复杂。只要你掌握了“基础算子组合”的思路,大部分不支持的算子都能轻松驯服。就算遇到组合不了的,你也能写出带自定义域的节点,并把它声明清楚,让所有合作伙伴都能理解。
其实,技术很多时候都是这样:别急着骂框架,先想想它为什么这样设计。ONNX的“不认识”不是针对你,只是需要你给它递上一个“翻译器”而已。希望这篇文章能让你再遇到类似问题时,不再发抖,而是淡定地打开PyCharm,写一个symbolic。
评论
围绕“PyTorch导出ONNX遇到不支持的算子时,通过自定义算子注册、梯度重写与兼容性声明构建完整的可移植替代方案,避免框架绑定”参与讨论