训练好的PyTorch模型要真正上线,很多人的第一反应是直接把 .pt 文件拷过去,或者在服务器上装一套 PyTorch 环境跑推理。你猜怎么着?等到环境不一致、推理速度上不去、目标硬件根本不支持 PyTorch 的时候,才想起来中间缺了一道工序——把模型转成 ONNX。这篇内容就是 PyTorch 模型部署系列的第一篇,聚焦在“模型转换为 ONNX”这个环节,讲清楚 ONNX 在部署链路里到底扮演什么角色、torch.onnx.export 的每个参数该怎么填、转完之后怎么验证、以及我实际踩过的那些坑。
这篇文章适合三类人:刚训完模型准备做服务端部署的算法工程师、要把模型搬到边缘设备或嵌入式平台的开发人员、以及正在学习模型部署、想知道“转换这一步到底发生了什么”的学生。我不会只贴一段能跑的代码就完事,而是把转换前后的准备、参数选择、验证方法、报错排查逻辑都串起来讲,方便你拿着这篇内容直接照着做。
在开始之前先给一个总览结论:ONNX 转换不是一个“一键导出”动作,它是一次静态化重构。你交给 torch.onnx.export 的模型,会被 PyTorch 用 JIT 追踪(trace)的方式跑一遍,把动态的 Python 执行过程变成静态的计算图。这一步里,任何你藏在 if、for 循环里的逻辑,都可能是之后报错的起点。
1. 为什么要转 ONNX:部署链路的角色定位
1.1 ONNX 不是推理引擎,是模型中间格式
很多人第一次接触 ONNX 时会有个误解:觉得 ONNX 是一种推理框架,能像 PyTorch 一样把模型跑起来。其实 ONNX(Open Neural Network Exchange)是一种计算图的中间表示格式,它描述的是一张有向无环图,图中的每个节点代表一个算子,边代表张量的流动。它既不像 PyTorch 那样负责动态建图,也不像 ONNX Runtime 那样负责真正在硬件上执行算子,它只负责“描述”模型。
打个比方:PyTorch 训练好的模型像一篇用中文写的手稿,ONNX 是把这篇手稿翻译成国际通用的世界语版本。翻译完之后,中文原稿还是中文原稿,世界语版本谁都能读,但真正要把它朗读出来,还得靠不同国家的播音员——这些播音员就是 ONNX Runtime、TensorRT、OpenVINO、RKNN 这些推理引擎或硬件编译器。
理解这一点很重要,因为它决定了你在转换阶段的目标:你交付的是一张能被各种推理后端正确解析的计算图,而不是一个能自己跑的二进制程序。所以判断转换是否成功,标准不是“ONNX 文件生成了”,而是“这张图在目标后端上能否正确、高效地执行”。
1.2 为什么不能只用 TorchScript 一条路走到黑
PyTorch 官方其实提供了自己的静态图方案 TorchScript,torch.jit.script 或 torch.jit.trace 也能把模型序列化。那为什么部署链路上,大家还是热衷转 ONNX?答案在于生态互通性。
TorchScript 是 PyTorch 自家的格式,如果你整套部署栈都基于 PyTorch,比如服务器端用 LibTorch(C++ 版本的 PyTorch)加载模型,那 TorchScript 完全够用。但如果你的目标环境是 NVIDIA 的 TensorRT、Intel 的 OpenVINO、ARM 的 NPU 工具链、或者用 C# 在 Windows 桌面端接入模型,这些后端几乎不会直接支持 TorchScript,它们统一认 ONNX。
我用过一个实际案例来说明:之前做一个图像分类项目,训练阶段在 PyTorch 里完成,交付的时候对方技术栈是 C# 桌面应用。如果只给 .pt 文件,对方基本无法处理;给了 ONNX 文件之后,对方直接用 ONNX Runtime 的 C# API 就接上了,前后不到半天。这就是 ONNX 的价值——它让模型脱离了训练框架和硬件平台的绑定。
1.3 到底什么场景才真正需要 ONNX
不需要所有项目都强行转 ONNX,我自己判断是否要转的标准有三个:
| 场景 | 是否需要 ONNX | 原因 |
|---|---|---|
| 服务端用 Python + PyTorch 直接推理 | 不一定 | 如果环境完全可控,直接加载 .pt 也可以 |
| 服务端要接入 TensorRT 加速 | 需要 | TensorRT 官方工具链原生读写 ONNX |
| 边缘设备/NPU 部署(瑞芯微、地平线等) | 需要 | 厂商工具链基本只认 ONNX 等中间格式 |
| 跨语言调用(C#/Java/Go) | 需要 | ONNX Runtime 提供了完整的跨语言 API |
| 模型量化与格式标准化 | 推荐 | ONNX 有标准 QDQ 量化描述,方便做 PTQ/QAT |
如果你只是在自己的服务器上用 PyTorch 跑推理、不换后端、不跨语言、不碰边缘硬件,那不一定非转 ONNX。但如果模型要往外走一步,ONNX 几乎是绕不开的必经之路。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 转换前的准备工作:版本、模型和输入一个都别漏
2.1 环境版本对齐:pytorch、onnx、opset 的一致性问题
转换 ONNX 不是把 torch.onnx.export 调出来就行,环境的版本匹配会直接影响导出结果。PyTorch 不同版本对 ONNX 的支持能力差别很大,我建议在开始之前先统一这几部分的版本:
bash复制pip list | grep -E "torch|onnx|onnxruntime"
版本对齐的核心是理解 opset(ONNX Operator Set) 的概念。ONNX 规范里每个算子都有自己的版本,opset_version 就是声明这张 ONNX 图“使用哪个版本的算子集合”。比如 Resize 算子,在 opset 10 和 opset 19 里的定义和行为差异很大,如果你用 opset 10 导出,后续现代推理引擎可能需要做不少兼容处理。
我目前的个人经验是:PyTorch 1.12 以上建议配合 opset 12 或 13;PyTorch 2.x 可以尝试 opset 17 甚至 18。如果你的目标推理环境比较老,比如某个嵌入式的 ONNX Runtime 版本停留在 1.10 以下,那么 opset 11 是更稳妥的下限。版本选太高,有可能转换成功,但部署端解析不了算子;选太低,一些复杂算子(如部分动态 shape 下的 Resize、Einsum)又无法完整表达。
2.2 模型准备:eval 模式、去掉梯度、用完整模型
这一步很多人会栽跟头。我见过不少同事转换前忘记调 model.eval(),结果导出的 ONNX 模型在推理时结果诡异——尤其是模型里含 BatchNorm 或 Dropout 的时候。原因很简单:model.train() 模式下,BatchNorm 会使用当前 batch 的统计量更新 running mean/running variance,而推理时需要冻结这些统计量。转换 ONNX 的过程本质上是在做一次推理采样,如果不切 eval 模式,BN 的归一化统计就可能错乱。
另外,建议在转换前把模型包装成完整的 nn.Module,而不是只把 model.state_dict() 拿出来。torch.onnx.export 接受的是模型实例,会真正执行一次 forward,所以模型结构必须完整且已经加载了训练好的权重。
还有一个细节:给模型加一个 .eval() 之后,记得用 torch.no_grad() 包裹整个导出过程。虽然 torch.onnx.export 内部会处理梯度,但养成这个习惯能避免某些自定义算子里的梯度逻辑干扰追踪过程。
python复制model = MyModel()
model.load_state_dict(torch.load("best.pth", map_location="cpu"))
model.eval()
with torch.no_grad():
torch.onnx.export(model, dummy_input, "model.onnx", ...)
2.3 输入张量:静态图的“形状锚点”
torch.onnx.export 需要一个 dummy_input 作为示例输入。很多人随手传一个 torch.randn(1, 3, 224, 224) 就完事了,但在传之前,最好想清楚一个问题:这个输入张量的形状会成为计算图中很多张量形状的锚点。
在 JIT trace 过程中,PyTorch 会把实际运行时的张量形状信息记录进图里。如果你的模型内部有 view、reshape、flatten 这类对形状敏感的操作,trace 出来的 ONNX 图可能把这些操作固化成“针对某个具体形状”的版本。也就是说,dummy_input 是什么 shape,后续推理时 ONNX 图就能稳定处理的 shape 范围,和你这个示例输入的 shape 有强相关。
实际操作中,我会准备和真实推理场景一致的 dummy input。如果线上推理的图片是 640x640 的,就不要用 224x224 去导出;如果 batch 会变化,就设置好 dynamic_axes(后面详细讲),或者先固定一个值。这一条看似不起眼,但能避免掉大量“输出 shape 对不上”的问题。
2.4 前后处理逻辑必须在导出之前剥离
ONNX 只描述模型内部的张量计算,不负责 Python 层面的图像解码、归一化、NMS 后处理这些逻辑。举个例子:如果代码里 forward 函数内部先做了 cv2.resize,或者调用了某个 Python 的 for 循环对预测结果做过滤,这些逻辑不会被正确转换。它们要么直接报错,要么被忽略,要么以错误的方式固化下来。
解决思路是:让模型保持纯张量计算。图像大小调整、归一化、通道变换等全部挪到模型外部,在调用 ONNX Runtime 推理之前用宿主语言(Python/C++/C#)处理;后处理如 NMS、阈值过滤也放到推理之后。这样既符合 ONNX 的图描述能力,也让推理后端可以专心做算子优化。
如果一些预处理必须在图内完成,比如 Resize 或者 Normalize,可以使用 PyTorch 的 torch.nn.functional 或 torchvision.transforms 里能被 JIT trace 的算子来实现,但要非常小心,因为它们很容易引入额外的算子版本兼容问题。
3. torch.onnx.export 核心参数逐项解读
3.1 导出 API 的最小可用示例
先给一个最精简的示例,后续再逐个解释参数:
python复制import torch
class MyModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv = torch.nn.Conv2d(3, 16, 3, padding=1)
self.fc = torch.nn.Linear(16 * 64 * 64, 10)
def forward(self, x):
x = self.conv(x)
x = torch.flatten(x, 1)
x = self.fc(x)
return x
model = MyModel().eval()
dummy_input = torch.randn(1, 3, 64, 64)
torch.onnx.export(
model,
dummy_input,
"model.onnx",
export_params=True,
opset_version=13,
do_constant_folding=True,
input_names=["input"],
output_names=["output"],
dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}}
)
这段代码跑完,model.onnx 就生成了。但“生成了”不等于“能用了”,先别急着走,把下面几个参数吃透。
3.2 opset_version:版本选不对,转换成功也白搭
opset_version 是 torch.onnx.export 里最重要的参数之一,也最容易让人困惑。它表示导出时使用的 ONNX 算子集版本。每个 opset 版本会新增算子或修改已有算子的行为,比如:
- opset 10:新增
Resize的现代版本 - opset 11:
Resize支持更多坐标系模式,同时新增一系列算子 - opset 13:进一步统一了
Resize等算子的属性,减少歧义 - opset 17:新增了
LayerNormalization等算子
我建议按这个原则选择:目标推理引擎支持的最高 opset 版本,与 PyTorch 支持能力取交集。如果目标环境是 ONNX Runtime 1.14 以上,opset 可以选 17 或 18;如果是老的嵌入式环境,保守一点选 13 或 11。
有一个经典坑:在 PyTorch 里用 F.interpolate 做上采样,导出时如果用 opset 10,生成的可能是旧版 Upsample 节点,行为和新版 Resize 有细微差别;如果用 opset 13+,导出的是 Resize 节点,坐标系模式可以明确指定。不同版本下,同一个 PyTorch 算子的输出可能完全不同。
3.3 do_constant_folding:什么时候开,什么时候关
do_constant_folding=True 表示在导出时执行常量折叠优化,把一些只依赖常量输入的计算提前算完,固化到图里,减少推理时的计算量。典型例子是 BatchNorm 的 scale/shift 和前一层的卷积权重融合,或者某些固定 mask 的乘加操作。
绝大多数场景下这个开关应该保持 True。但我遇到过一种情况:模型里某个算子在 trace 阶段因为常量折叠被错误地简化了,导致输出 shape 变化。如果你在验证阶段发现 ONNX Runtime 的输出和 PyTorch 差异大,而且排查算子也看不出问题,可以临时把 do_constant_folding=False 重新导出对比一次。这相当于排查嫌疑人的手段,不是推荐部署时关掉。
3.4 input_names 和 output_names:命名不是给人看的,是给推理引擎看的
input_names 和 output_names 用来给 ONNX 图的输入输出张量起名字。很多人觉得“随便起不起都行”,但如果后续要用 ONNX Runtime 的 C++/C# API,或者要在 TensorRT 里按名字绑定输入输出,这些名字就直接决定你写代码的清晰度。
更重要的是,这个命名和 dynamic_axes 配置强相关。如果你给输入取名叫 input,那么在 dynamic_axes 里就必须用同名配置:
python复制dynamic_axes={
"input": {0: "batch", 2: "height", 3: "width"},
"output": {0: "batch"}
}
如果不给名字,默认的输入输出名是 input.1、output.1 这种带编号的,后续写绑定代码时很容易写错。我个人的习惯是统一用语义化名字:images、input_ids、logits、boxes 这类,方便在 Netron 里看图和写部署代码时一眼识别。
还有一个容易被忽略的参数是 export_params。默认是 True,表示把模型的权重参数固化进 ONNX 文件。如果你只想导出图结构、不想要权重,可以设成 False。但这种情况极少,正常部署时让它保持 True。
4. dynamic_axes 动态轴:设置得当是加速器,设置不当是绊脚石
4.1 dynamic_axes 的三种配置写法
dynamic_axes 用于声明模型哪些维度是动态的,在推理时可以改变。常见的动态维度有两种:batch 维度(一次推理几张图)和空间分辨率维度(输入图片长宽不固定)。配置方式有三种等价写法:
python复制# 写法一:字典形式,明确指定维度
dynamic_axes={
"input": {0: "batch", 2: "height", 3: "width"},
"output": {0: "batch"}
}
# 写法二:字符串形式,表示所有维度都可变
dynamic_axes={
"input": "batch",
"output": "batch"
}
# 写法三:整数索引形式(老版本 API)
dynamic_axes={"input": [0, 2, 3]}
我推荐第一种写法,因为它最精确。动态维度声明得越具体,推理引擎越容易做形状推断和内存规划。字符串形式表示“所有维度都可能变”,反而会给一些推理引擎增加不必要的负担。
4.2 动态维度的隐形成本
说句实在话,动态 shape 不是免费的。每次推理时,如果输入 shape 与上一次不同,推理引擎需要重新做一次形状推断、内存分配甚至图优化,这会明显增加延迟。实测下来,在一些边缘设备上,动态 batch 模式下的推理耗时可能比固定 batch 高出 20% 到 50%。
另一个隐患是后续量化。如果你计划做 int8 量化(比如转成 ONNX 的 QDQ 格式,再走 ONNX Runtime 量化或者瑞芯微 RKNN 工具链),很多量化工具要求输入 shape 固定。动态 shape 的模型在量化时,要么直接报错,要么部分算子无法量化,导致最终性能不达标。我之前在 RKNN 工具链上转一个目标检测模型,就是因为输入分辨率被设成了动态,RKNN 工具链直接拒绝转换,改成固定 640x640 之后才顺利通过。
4.3 固定 shape 优先:一个来自边缘设备的教训
接到一个树莓派 5 上部署自训练 YOLOv5 的项目时,最初我把输入设成动态宽高,觉得这样灵活。结果在 ONNX Runtime 的 CPU 推理时,每次传入不同分辨率的图片,推理时间波动特别大。后来用固定 640x640 输入重新导出,图优化效果明显提升,平均推理耗时下降了约 30%。
当然,如果你的业务场景确实需要动态分辨率,不能为了性能强行固化成正方形。但至少做一次评估:你的线上输入尺寸是不是真的会频繁变化?如果基本固定,或者只有少数几个档位,可以导出多个静态 shape 的 ONNX 文件,在推理代码里按实际输入尺寸选择对应的 session。这样既满足了灵活性,又保住了性能。
5. 转换后的验证:数值对比和结构检查一个都不能少
5.1 结构校验:onnx.checker.check_model
转换完成后的第一件事,不是拿去推理,而是先做结构校验。ONNX 官方提供了 onnx.checker:
python复制import onnx
model = onnx.load("model.onnx")
onnx.checker.check_model(model)
print("check passed")
这个检查会验证 ONNX 图的结构合法性:节点连接是否正确、输入输出是否匹配、算子属性是否合法等。如果这一步就报错,说明模型在导出过程中已经产生了结构性问题,不用继续往下走了。
但要注意,check_model 只检查结构,不检查数值。我自己经历过一次:check_model 完全通过,但 ONNX Runtime 跑出来的结果和 PyTorch 差了十万八千里。所以结构校验只是第一步,数值验证才是关键。
5.2 数值验证:用 ONNX Runtime 与 PyTorch 输出对比
数值验证的核心逻辑很简单:用同一个输入分别跑 PyTorch 模型和 ONNX Runtime 加载的模型,对比输出张量的差异。下面这个脚本我几乎每个模型都要跑一遍:
python复制import numpy as np
import onnxruntime as ort
import torch
# PyTorch side
model.eval()
with torch.no_grad():
torch_out = model(dummy_input)
# ONNX Runtime side
ort_session = ort.InferenceSession(
"model.onnx",
providers=["CPUExecutionProvider"]
)
ort_out = ort_session.run(
None,
{"input": dummy_input.numpy()}
)
# Compare
for i, (t_out, o_out) in enumerate(zip(torch_out, ort_out)):
t_arr = t_out.detach().numpy()
o_arr = np.asarray(o_out)
max_diff = np.abs(t_arr - o_arr).max()
mean_diff = np.abs(t_arr - o_arr).mean()
print(f"output[{i}] shape={o_arr.shape} max_diff={max_diff:.8f} mean_diff={mean_diff:.8f}")
判断标准上,我一般遵循这个经验:max_abs_diff 在 1e-4 到 1e-5 量级属于正常范围,说明是浮点精度导致的微小误差;如果达到 1e-2 甚至更大,基本可以断定图里某个算子被错误地转换或优化了,需要进一步排查。
5.3 多组输入的边界测试
只测一组固定输入远远不够,尤其当模型带动态 shape 或者输入分布变化大时。我建议至少准备三组输入做对比:
- 标准输入:最常规的 shape 和数值范围,确认基本正确性。
- 边界输入:比如接近纯黑/纯白的图像、全零输入、或者特别大的数值,检查是否存在数值溢出或除零问题。
- 不同 shape 输入:如果你配置了 dynamic_axes,用不同的 batch size、不同分辨率分别测试,确认动态轴真的生效,且所有 shape 下输出都一致。
这个习惯帮我抓出来过一个隐藏很深的问题:某个模型在 224x224 输入下完全正常,但换成 1080p 输入时,ONNX Runtime 端出现了奇怪的边缘伪影,最后定位到是 F.interpolate 的 align_corners 参数在导出时和 ONNX Resize 的坐标模式不一致导致的。这个坑你只测标准输入永远发现不了。
5.4 误判案例:数值偏差多大才算异常
有一次我在验证一个语义分割模型时,看到 mean_diff 是 1e-6,心想“完美,没问题”。但在后续接 TensorRT 时,TensorRT 的输出和 PyTorch 的最大误差达到了 0.05,而且集中在分割边缘区域。刚开始我还以为是 TensorRT 的精度问题,后来仔细排查发现,问题在 ONNX 图里的 Resize 节点上——PyTorch 默认的 align_corners=False 和 ONNX Resize 默认的 half_pixel 坐标系虽然在许多情况下等价,但在某些 opset 版本和具体尺寸下,会存在像素对齐差异。
所以现在我的验证策略升级了:不只看整体误差,还会按输出张量的语义做分区域检查。分割模型就看边缘区域和主体区域的误差,检测模型就按检测框内的特征图单独对比。数值验证的目的是发现潜在风险,不能只满足于“大体对上”。
6. 转换报错的常见根因与排查思路
6.1 追踪失败类:动态控制流和 Python 原生逻辑
PyTorch 的 torch.onnx.export 底层依赖 TorchScript 的 JIT tracing,它会执行一次模型 forward,然后把执行过的算子记录成图。这意味着,forward 函数里的 Python 原生逻辑并不是被“翻译”到 ONNX 里的,而是在 trace 时被直接执行,只有涉及张量运算的算子会被记录。
所以最常见的报错场景是:
code复制RuntimeError: Could not run 'aten::native_batch_norm' with arguments from the 'CPU' backend...
或者 trace 过程中抛出 Python 异常。原因通常是 forward 里有 if x.sum() > 0: 这类依赖张量数值的 Python 控制流。trace 阶段 x.sum() 会被执行,得到一个确定的 bool 值,模型只会走其中一个分支,另一个分支的逻辑永远不会出现在 ONNX 图里。
解决方案很直接:把动态控制流从 forward 里移出去,或者在导出前把模型做成一个只包含确定计算路径的版本。如果实在绕不开,可以考虑用 torch.onnx.is_onnx_supported() 或检查算子支持列表,但更推荐的做法是调整模型设计,让它的 forward 对输入 shape 和数值不敏感。
6.2 算子不支持类:aten::xxx 没有 ONNX 映射怎么办
这是第二大类的报错,特征比较明显:
code复制RuntimeError: Exporting the operator 'aten::einsum' to ONNX opset version 11 is not supported.
PyTorch 算子到 ONNX 的映射不是一一对应的,PyTorch 有几千个算子的变体,而 ONNX 算子集相对少很多。像 torch.einsum 在 opset 12 之后才有对应的 Einsum 节点,某些老的自定义算子则根本没有映射。
碰到这种情况,我的排查思路是三步:
- 升级 opset_version,有些算子在新 opset 里才有映射,比如
einsum在 12+、LayerNorm在 17+。 - 用等价算子替换,把
einsum拆成permute、matmul、reshape的组合;把某些自定义激活函数换成F.gelu、F.relu等标准算子。 - 给 PyTorch 注册自定义 ONNX 符号函数,用
torch.onnx.register_custom_op_symbolic告诉 PyTorch“这个算子怎么映射到 ONNX”。这个方法功能强,但成本高,需要同时了解 PyTorch 算子的底层实现和 ONNX 算子的属性定义,非必要不推荐。
另外要留意:torchvision 里的 nms、roi_align、deform_conv2d 等算子在部分版本有 ONNX 支持,但依赖 torchvision 的版本。用之前先查一下你手里的 torchvision 版本是否支持导出这些算子。
6.3 推理引擎兼容类:opset、Resize、动态 shape 的连锁反应
这一类最隐蔽,因为转换阶段完全不报错,问题出现在部署阶段。典型的两个表现:
表现一:ONNX Runtime 加载时报错 No Op registered for Resize with domain_version=...
原因几乎都是 opset_version 和 ONNX Runtime 版本不匹配。比如你用了 opset 18 导出,但部署端的 ONNX Runtime 是 1.10 版本,它只支持到 opset 15。解决方案就是在导出时降低 opset,或升级部署端的 ONNX Runtime。
表现二:推理结果 shape 对不上
这个多见于带动态 shape 的模型。view、reshape 这类算子在动态 shape 下,如果导出时把某个维度固化了,推理时输入尺寸一变,后面的形状就错位了。排查方法是先固定 shape 重新导出,看问题是否消失;如果消失,就逐层检查动态轴的传导路径,看是哪个节点把动态信息丢了。
6.4 排查工具链:Netron 和 onnx-simplifier 配合使用
排查 ONNX 问题,最好用的两个工具是 Netron 和 onnx-simplifier。
Netron 是一个可视化工具,直接在浏览器里打开 ONNX 文件,就能看到完整的计算图结构、每个节点的属性和 shape 信息。排查 shape 问题时,我通常先在 Netron 里沿着 graph 看一遍,确认每个节点的输入输出 shape 是否符合预期,很快就能定位到形状断掉的节点。
onnx-simplifier 的作用是简化图结构:去掉冗余节点、常量折叠、合并算子序列等。有些 PyTorch 导出的图会有大量无用的 Identity 节点或多余的 transpose 对,虽然不影响正确性,但会影响推理性能,有时还会干扰后续 TensorRT/RKNN 的转换。用法很简单:
bash复制pip install onnx-simplifier
python -m onnxsim model.onnx model_sim.onnx
但要注意:onnx-simplifier 在某些场景下会把动态维度错误地固化成常量,所以简化之后一定要重新做一遍数值验证。我遇到过一次,简化前后的输出完全一致,但输入 shape 一变就崩了,最后发现是 simplifier 把动态 shape 推断成了常量 shape。
7. 从 ONNX 继续走:量化、TensorRT 与边缘 NPU 的衔接
7.1 最省事的部署路径:直接交给 ONNX Runtime
如果模型只需要在 CPU 或 GPU 上做常规服务,ONNX Runtime 是最省事的落点。转换完后直接用 Python 或 C++ 加载就可以跑:
python复制import onnxruntime as ort
session = ort.InferenceSession(
"model.onnx",
providers=["CUDAExecutionProvider", "CPUExecutionProvider"]
)
ONNX Runtime 还支持对图做 level 优化,默认 ORT_ENABLE_ALL 会让推理引擎自动做算子融合、内存规划和并行执行。很多 PyTorch 模型在 ONNX Runtime 上跑,即使不做任何手工优化,性能也比直接用 PyTorch 推理高出一截——因为 ONNX Runtime 的图优化是专门针对静态图设计的,省掉了动态图的调度开销。
7.2 TensorRT 转换时的 ONNX 准备
要用 TensorRT 加速,常规路线是把 ONNX 文件丢给 trtexec 或者 onnx-tensorrt 解析器生成 engine。TensorRT 对 ONNX 图的要求更苛刻,几个常见问题我在实践中反复碰到:
- 动态 shape 需要额外配置 optimization profile:TensorRT 不是简单地“支持动态”,你得告诉它 batch/height/width 的最小值、最优值、最大值,它才能为每个 shape 区间做规划。
- 部分 ONNX 算子 TensorRT 不支持:比如
Einsum在 TensorRT 8.x 早期版本支持有限,我一般会在导出前用算子替换来规避。 - 解析器版本和 ONNX opset 要匹配:onnx-tensorrt 解析器对 opset 的支持有上限,用太新的 opset 导出 TensorRT 可能直接报 parse error。
所以如果你确定后续要走 TensorRT,建议导出 ONNX 时就用相对保守的 opset(比如 13 或 17),并且在 PyTorch 侧尽量避免用过于新颖的算子。
7.3 边缘 NPU 与量化 int8:瑞芯微和树莓派的常见路线
边缘设备的部署路线里,转换 ONNX 只是万里长征第一步。以瑞芯微 RK3588/RK3568 为例,官方工具链 RKNN-Toolkit2 支持直接读取 ONNX 模型,但通常要求先做以下几件事:
- 将模型输入固定为静态 shape,尤其是 batch 维度。
- 用代表性数据集做 int8 PTQ(训练后量化),数据集要覆盖真实业务的数据分布,不能随便拿几张图凑数,否则量化后掉点严重。
- 逐层查看量化效果,必要时对特定层使用不量化策略或混合精度。
树莓派 5 上部署自己训练的 YOLOv5 模型,大部分人会选择 ONNX Runtime 的 CPU 推理,或者走 NCNN、RKNN 这类第三方推理栈。无论哪种,ONNX 导出环节的质量都直接影响后续转换的成功率。我见过太多人抱怨 RKNN 工具链报各种奇怪的错误,最后发现问题出在 ONNX 导出时使用了动态 shape,或者 Resize 节点的 coordinate_transformation_mode 与工具链预期不一致。
7.4 int8 量化的 ONNX 表达:QDQ 格式
int8 量化在 ONNX 里通常以 QDQ(QuantizeLinear/DequantizeLinear)格式表达:图上每隔一段就插入一对量化和反量化算子,让推理引擎知道这段计算需要在 int8 下执行。如果你用 PyTorch 做 QAT(量化感知训练),PyTorch 2.x 已经支持导出带 QDQ 节点的 ONNX 模型;如果用 PTQ,则可以用 ONNX Runtime 的 quantization 工具或厂商的工具链完成。
这里有一条经验:量化一定要在确定了部署目标后再做。不同部署目标的量化实现不同,TensorRT 的 QDQ 解释和 RKNN 的 QDQ 解释不完全一致,你在 PyTorch 里导出的 QDQ 图未必能被目标工具链原样接受。更稳妥的做法是,让 ONNX 保持 FP32 精度作为“母版”,到了具体部署阶段,再基于这个母版做对应的 ptq/量化转换。
最后的实际操作心得
转 ONNX 看起来是一个小步骤,但它的质量直接决定后续 TensorRT、RKNN、ONNX Runtime 等所有部署环节的顺畅程度。我踩过这么多次坑之后,现在的标准动作是先写一个可复用的验证脚本,每次转换完模型都固定跑一遍 PyTorch 与 ONNX Runtime 的数值对比,再按需检查动态 shape 和算子兼容性。
还有个小技巧分享给做部署的朋友:转换后的 ONNX 文件建议用 git 管理,同时保留转换脚本和当时的 PyTorch 权重。模型迭代时,新旧 ONNX 之间做 diff 会方便很多。如果你在部署链路上卡在某个奇怪报错上,多半不是目标推理引擎的问题,而是 ONNX 导出环节留下了隐患。回到图里,用 Netron 一步步看,比反复试错高效得多。
