ONNX 导出与跨框架推理
本教程共 60 篇 · 第 51 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:认识 ONNX 格式,把 PyTorch 模型导出成 .onnx 文件,用 ONNX Runtime 跑通推理,并验证结果与 PyTorch 一致。
ONNX:模型界的通用格式
ONNX 全称开放神经网络交换格式(Open Neural Network eXchange),是一套描述机器学习模型的开放标准。2017 年由微软等公司推出,目标就一个:让模型在框架之间自由流动。
它解决什么问题?你在 PyTorch 里训好模型,但部署方可能用 C++ 服务、手机芯片、甚至浏览器。逐个框架去适配太痛苦,不如约定一个中间格式:大家都读写 ONNX。你可以把它理解成模型界的 PDF——写的人用 Word,看的人用各种阅读器。
ONNX 的历史也值得一提。2017 年微软和 Facebook 联手发布,当时 Facebook 的 Caffe2 是主要输出目标,所以老教程里满是「导出到 Caffe2、迁移到移动端」的内容。后来 Caffe2 并入 PyTorch 退役,ONNX 的定位反而更清晰了:它是模型交换的标准格式,谁家的推理引擎都能消费。今天你在网上看到的 ONNX 教程,绝大部分场景都是跨框架推理和硬件加速。
PyTorch 的导出入口在 torch.onnx 模块。ONNX 生态里最常用的运行时是微软的 ONNX Runtime,还有 NVIDIA 的 TensorRT(下一章的主角)。
安装依赖
导出器依赖三个包:onnx(标准库)、onnxscript(算子翻译脚本)、onnxruntime(运行时):
pip install --upgrade onnx onnxscript onnxruntime
装完在 Python 里逐个 import 并打印版本号,确认都能用再往下走。
导出一个小模型
从 PyTorch 2.5 起,导出器分两代。新的是 torch.onnx.export(..., dynamo=True),底层走上一章讲的 torch.export,官方推荐;旧的那套底层依赖 TorchScript,已弃用,别用。
导出流程三步:
- 定义模型并调
eval(); - 准备示例输入(元组形式);
- 调用
torch.onnx.export并保存。
import torch
import torch.nn as nn
import torch.nn.functional as F
class ImageClassifierModel(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 6, 5)
self.conv2 = nn.Conv2d(6, 16, 5)
self.fc1 = nn.Linear(16 * 5 * 5, 120)
self.fc2 = nn.Linear(120, 84)
self.fc3 = nn.Linear(84, 10)
def forward(self, x: torch.Tensor):
x = F.max_pool2d(F.relu(self.conv1(x)), (2, 2))
x = F.max_pool2d(F.relu(self.conv2(x)), 2)
x = torch.flatten(x, 1)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
return self.fc3(x)
torch_model = ImageClassifierModel().eval()
example_inputs = (torch.randn(1, 1, 32, 32),)
onnx_program = torch.onnx.export(torch_model, example_inputs, dynamo=True)
onnx_program.save("image_classifier_model.onnx")
导出结果是一个 ONNXProgram 对象,.save() 存成 .onnx 文件。用 onnx 标准库验证一下文件格式:
import onnx
onnx_model = onnx.load("image_classifier_model.onnx")
onnx.checker.check_model(onnx_model) # 没抛异常就是合法模型
Tip导出前一定先
model.eval()。训练模式下的 BatchNorm、Dropout 行为和推理不同,带着它们导出的图在线上会出错。
关于算子集版本(opset)多说一句。ONNX 的算子也在进化,每一版标准都有编号,叫 opset,比如 opset 18、opset 21。导出器按你指定的 opset 把 PyTorch 算子翻译成对应版本的 ONNX 算子。opset 太老,新算子翻不了;opset 太新,运行时可能不支持。默认值通常是稳妥选择,老教程里常见的 opset_version=13 之类参数,如今很少需要手动指定了。
旧教程里另一个常见参数是 dynamic_axes,用来声明动态维度,比如把 batch 维标成动态。新导出器改走 torch.export 的动态形状体系(dynamic_shapes,上一章讲过),语义更完整。看到旧代码里的 dynamic_axes,知道它想干嘛就行,新代码用 dynamic_shapes。
看图与跑图
想看看图长什么样,把 .onnx 文件拖进 Netron——浏览器打开 netron.app 就行,节点、连线、张量形状一目了然。排查导出问题时的神器。
跑图用 ONNX Runtime。注意它的输入输出是 NumPy 数组,不是张量:
import onnxruntime
ort_session = onnxruntime.InferenceSession(
"image_classifier_model.onnx", providers=["CPUExecutionProvider"]
)
# 张量转 NumPy,按输入名打包成字典
onnx_inputs = [t.numpy(force=True) for t in example_inputs]
ort_inputs = {inp.name: val for inp, val in zip(ort_session.get_inputs(), onnx_inputs)}
ort_output = ort_session.run(None, ort_inputs)[0]
# 和 PyTorch 逐元素对拍
torch_output = torch_model(*example_inputs)
torch.testing.assert_close(torch_output, torch.tensor(ort_output))
print("PyTorch 与 ONNX Runtime 输出一致")
assert_close 通过,说明导出无损,模型可以放心交给部署方。
这里有个细节值得记住:ONNX Runtime 的输入是字典,键是输入名,值是对应数组。输入名从哪来?ort_session.get_inputs() 里取,每个输入对象的 .name 属性。别自己编名字,跟图里的名字对不上会直接报错。
另外 providers 参数可以换成 ["CUDAExecutionProvider"] 让推理跑在 GPU 上,同样一份 .onnx 文件,换个运行时配置就能吃上不同硬件——这就是「通用格式」的实惠。
带分支的模型怎么导出
和 torch.export 一样,直接写 if-else 且条件依赖数据,导出会失败。需要把分支改写成 torch.cond(上一章讲过),ONNX 图里就有对应的 If 节点。依赖形状的分支则完全不受影响。
具体做法是把每个分支写成一个子函数,条件、两个分支、操作数分别传给 torch.cond。分支之间返回的张量形状、类型必须一致,否则导出的图无法确定输出形态。写的时候别嫌麻烦——这是静态图的代价,把不确定性挡在图外面。
不支持的算子怎么办
ONNX 的标准算子集(opset)覆盖不了所有 PyTorch 算子。遇到「No decompositions registered for …」这类报错,说明有算子翻译不了。可以用 onnxscript 自己补一个翻译,再通过 custom_translation_table 注册:
import onnxscript
from onnxscript import opset18 as op
def custom_aten_add(self, other, alpha: float = 1.0):
# self/other/alpha 必须与原算子的参数名一致
if alpha != 1.0:
alpha = op.CastLike(alpha, other)
other = op.Mul(other, alpha)
return op.Add(other, self)
# Model 是任意含该算子的模块,x、y 是它的示例输入
onnx_program = torch.onnx.export(
Model().eval(), (x, y),
dynamo=True,
custom_translation_table={torch.ops.aten.add.Tensor: custom_aten_add},
)
custom_translation_table 把 PyTorch 算子映射到你的 ONNX Script 实现。函数签名必须和原算子一致,参数要带类型标注。这个机制还有个大用处:把算子替换成运行时厂商的专有实现(比如 ONNX Runtime 的 com.microsoft 域加速算子),直接吃厂商优化。
大多数情况下你不需要走到这一步。官方导出器的覆盖范围一直在扩大,遇到报错先搜一下,看看是不是已知问题、有没有人贡献过实现。真要自己写,写完记得 onnx_program.optimize() 做一遍图优化,去掉冗余节点。
导出检查清单
把这一章的内容收拢成一份清单,导出前逐条过:
- 模型
eval(),退出训练模式; - 示例输入形状和实际部署时一致,batch 大小有变化就声明动态形状;
- 导出成功后
onnx.checker.check_model验证文件合法性; - 用 Netron 看一眼图结构,和模型结构对得上;
- ONNX Runtime 跑一遍,和 PyTorch 结果
assert_close对拍。
五条都过,这份 ONNX 文件才算真正交付。
小结
ONNX 的价值在「通用」两个字:导出一次,处处能跑。核心流程是 eval → export(dynamo=True) → save → 验证对拍 → 交给运行时。这套流程熟手十分钟就能走完,难点全在边界情况:控制流、动态形状、稀有算子。碰到一个解决一个,别慌,工具链已经比两年前成熟太多了。下一章走进推理加速的世界:TensorRT 怎么把这张图榨出更多性能,量化又是什么。