先把一个让我印象很深的画面放在这里:训练脚本里loss曲线终于贴着地板走,模型在验证集上也刷出了漂亮数字,我以为把.pth文件复制到生产机器上就万事大吉。真到了部署那天才发现,线上推理程序是C++写的,加速引擎要的是TensorRT的engine,边缘盒子只认自家转换工具,没人愿意直接吃PyTorch权重。最后绕了一大圈,我用的过渡格式就是ONNX。这篇内容就围绕一个非常具体的动作展开:如何把训练好的PyTorch模型转成ONNX,转出来之后又怎么验证、怎么避坑。无论是刚接触模型落地的同学,还是已经在折腾TensorRT、OpenVINO、瑞芯微这类工具链的朋友,把ONNX这一关打通,后续路会好走很多。
1. 为什么模型训练完还要先变成ONNX
1.1 .pth不是拿来即用的部署格式
PyTorch在训练和快速实验阶段确实无敌好用,模型结构就是Python代码,想怎么改怎么改,想在哪一行打印就打印,梯度也是自动求的。但这套便利性在部署阶段反而成了负担。线上服务通常希望模型是一个独立的、可被高性能引擎直接解析的“产物”,而不是一个需要拉起Python解释器、导入PyTorch、再把整个模型类重新实例化的依赖。
一个典型的矛盾是:你的服务端或边缘设备可能根本没有Python环境,也可能是C++为主的推理框架,或者只能跑特定厂商的加速SDK。这时候你把训练好的.pth交给对方,对方第一句话往往就是“这是什么结构?层定义文件呢?预处理逻辑呢?”因为.pth本质上只是参数,不包含完整网络结构定义,结构还需要代码来重建。虽然可以把整个PyTorch模型torch.save成一个包含结构的文件,但推理时依旧绕不开PyTorch框架本身,体积大、启动慢、对运行时版本敏感,部署确实不方便。
1.2 ONNX是“模型界的普通话”
ONNX的全称是Open Neural Network Exchange,相当于一类开放的模型表示格式。它把网络结构、权重、输入输出定义都统一在一个.onnx文件里,不绑定任何特定训练框架。PyTorch训练完的模型可以导出成ONNX,TensorFlow训练完也有工具能导成ONNX,其他框架多数也认这个格式。对模型来说,ONNX就像“普通话”,让不同框架、不同推理引擎之间能直接对话。
而且ONNX文件里的图结构是显式的,包含了每个算子的类型、输入输出张量名、权重值等。你可以用工具像看电路图一样查看它,也可以做算子层面的优化、剪枝、量化。这种中间表示天然适合做部署链路的枢纽。
1.3 从ONNX继续流向不同加速后端
ONNX本身通常不是最后跑推理的那个引擎,更多像是一个通用的“传输格式”。真正跑起来的时候,根据你的硬件和场景,后面还会接不同的执行后端:
- ONNX Runtime:微软开源的跨平台推理引擎,CPU、GPU、手机端都能跑,也是目前验证ONNX文件最方便的runtime。
- TensorRT:NVIDIA GPU上的高性能推理引擎,极擅长把模型优化成针对特定GPU的推理计划。
- OpenVINO:Intel CPU、集显、VPU上优化推理的主流选择。
- RKNN、NNRT等边缘NPU工具链:瑞芯微、地平线等芯片平台,通常也是先把模型转成ONNX,再导入自家SDK做量化和格式转换。
所以在很多实际项目里,路径通常是“PyTorch权重 → ONNX → 各平台格式”。理解了这条链路,就知道ONNX这一个节点有多重要。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境准备:把基础工具链搭好
2.1 安装PyTorch
如果只是想走通转换流程,不需要特别大的GPU算力,CPU版本的PyTorch就足够了。安装命令也比较简单:
bash复制pip install torch torchvision
如果你需要GPU训练并想从GPU环境里直接导出模型,就需要根据自己机器的CUDA版本去PyTorch官网选对应的安装命令。这里我的建议是:导出ONNX时尽量切到CPU上操作,不要图方便直接在GPU显存里做trace。原因后面细说,但这一步先记住,能省很多莫名其妙的兼容性问题。
如果你用的是Anaconda,也可以用conda来建环境:
bash复制conda create -n deploy python=3.9
conda activate deploy
pip install torch torchvision
2.2 安装onnx和onnxruntime
接下来装两个和ONNX强相关的Python库:
bash复制pip install onnx onnxruntime
这里容易混淆的是两个包的分工:onnx负责解析、检查、修改ONNX模型文件,比如加载.onnx、跑格式校验、查看节点信息;onnxruntime才负责真正把ONNX模型跑起来做推理。如果你只装了onnxruntime就跑去调onnx.load,会直接报ModuleNotFoundError,反过来也一样。
如果后面要在GPU上用ONNX Runtime加速,可以安装onnxruntime-gpu,并且需要装对应CUDA和cuDNN版本。如果只是做模型转换和一致性验证,CPU版本的onnxruntime完全够用,轻量又稳定。这里我习惯用CPU版本先把流程跑通,验证模型没问题之后再考虑GPU推理,避免一开始就陷入环境依赖泥潭。
2.3 快速验证环境是否正常
装完先别急着写转换代码,用一行命令确认三个库的版本都能正常加载:
bash复制python -c "import torch, onnx, onnxruntime; print(torch.__version__); print(onnx.__version__); print(onnxruntime.__version__)"
如果这三行版本号都正常打印出来,说明基础环境没问题。常见问题无非是Python版本太老或太新导致某个包没装好,或者系统里存在多个Python环境导致pip装的库和当前python不是同一个。用which python和which pip先对齐环境,是排查这类问题最直接的方法。
3. torch.onnx.export核心参数拆解
3.1 一段能跑的最小转换代码
PyTorch转ONNX的核心就是torch.onnx.export,它做的事情可以粗分成三步:把模型转成TorchScript格式,记录计算图,再映射成ONNX的算子集合。最小可用的代码长这样:
python复制import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.conv = nn.Conv2d(3, 16, kernel_size=3, padding=1)
self.relu = nn.ReLU()
self.pool = nn.AdaptiveAvgPool2d((1, 1))
self.fc = nn.Linear(16, 10)
def forward(self, x):
x = self.conv(x)
x = self.relu(x)
x = self.pool(x)
x = torch.flatten(x, 1)
return self.fc(x)
model = SimpleModel()
model.eval()
dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
model,
dummy_input,
"simple_model.onnx",
input_names=["input"],
output_names=["output"],
opset_version=14,
do_constant_folding=True,
)
这段代码跑完,当前目录下就会多出一个simple_model.onnx。需要注意:dummy_input并不是一个可有可无的摆设,它的shape、dtype、乃至送进去的是不是CPU张量,都会直接影响导出结果。因为PyTorch导出ONNX默认走的是trace(追踪)机制,它要把一个假的输入真正跑一遍前向,记录执行过的算子和张量流转关系,最终生成计算图。
3.2 导出前必须调用的model.eval()
很多刚开始接触ONNX的同学最容易翻车的一点,就是忘了在导出前调用model.eval()。PyTorch的模型默认是训练模式,如果模型里含有BatchNorm、Dropout这类层,不同模式下的前向逻辑是完全不同的。
BatchNorm在训练模式下会用当前batch的统计数据做归一化,同时更新running_mean和running_var;Dropout在训练模式下会随机丢弃一部分神经元,导致每次前向结果不一样。你在这种状态下做trace,导出计算图本质上是拿一个“不稳定”的推理逻辑去作图,后面在ONNX Runtime里跑出来的结果自然很可能和模型正常推理结果对不上。
所以无论你加载的是PyTorch官方权重还是自己训练的checkpoint,只要准备导出了,就统一执行:
python复制model.eval()
model.to("cpu")
有人可能觉得把模型搬到CPU多此一举。我的经验是:如果模型里有某些算子在GPU上的实现导出的算子和CPU上不一样,或者目标推理环境根本只有CPU,后续就很容易出现“导出成功但Runtime加载报错”或者“精度对不上”的问题。老老实实CPU导出,能避免至少一半的部署兼容性麻烦。
3.3 用input_names/output_names给输入输出起名
input_names和output_names这两个参数看起来不起眼,但非常重要。它们直接决定ONNX模型中输入节点和输出节点的命名,而你在后续用ONNX Runtime或者TensorRT时,就是靠这些名字去喂数据和取结果的。
命名建议和使用模型时的习惯对齐。比如说分类模型的输入叫input,输出叫output;目标检测模型的输出可能叫boxes、scores、labels。一个模型中如果有多个输入或多个输出,可以用列表依次写清楚:
python复制torch.onnx.export(
model,
(input_image, input_meta),
"model.onnx",
input_names=["image", "meta"],
output_names=["pred"],
...
)
后面用onnxruntime推理时,就可以用这些名字构造输入字典:
python复制ort_out = sess.run(
["pred"],
{"image": img_array, "meta": meta_array}
)
如果不设置名字,PyTorch会生成一串形如input.1、output.1这种没规律的名。虽然也能跑通,但一旦模型结构改动或者要接后端转换,很容易把自己坑到。尤其是一些支持“按名字绑定输入”的框架,名字稳定非常重要。
3.4 dynamic_axes怎么设置才不容易出错
模型导出时默认会把dummy_input的shape当作固定shape写死。也就是说,你用(1, 3, 224, 224)导出的模型,推理时输入shape也必须严格是(1, 3, 224, 224),想换成(4, 3, 224, 224)甚至(1, 3, 320, 320)都可能直接报输入维度不匹配。
现实场景中批量大小经常要变,比如线上服务可能要支持动态batch,或者目标检测模型需要处理不同尺寸的输入。这时候需要设置dynamic_axes,它本质是在告诉ONNX:“这个维度的长度不固定,推理时由实际输入决定”。
python复制torch.onnx.export(
model,
dummy_input,
"simple_model.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={
"input": {0: "batch"},
"output": {0: "batch"}
},
opset_version=14,
)
上面代码表示输入张量的第0维是可变维度,名字可以自己起,示意含义是batch。如果不只batch需要动态,把一个维度的数字继续往下加即可。比如有些语义分割模型的输入输出H、W也需要动态:
python复制dynamic_axes={
"input": {0: "batch", 2: "height", 3: "width"},
"output": {0: "batch", 2: "height", 3: "width"},
}
这里必须提醒一下:不要看到dynamic_axes能设动态维度,就把所有维度都标成动态。模型内部有些层对具体维度大小是有要求的,比如全连接层的输入维数必须和权重矩阵匹配,它天然只能动态batch这一维。某些reshape或slice操作如果标了过宽的动态范围,导出的模型可能在推理时报shape推理错误。我的做法是:先只让必要的维度动态,例如分类网络只让batch动态;如果检测或分割模型确实需要多尺度输入,再逐步把H、W加进去,并且每加一个维度都用不同尺寸输入实测一遍。
3.5 opset_version选多少,背后的兼容性逻辑
opset_version是很多初学者会忽略但坑最多的参数。ONNX每代版本会维护一套算子集,每一代opset可能会新增算子、修改已有算子的属性或输入输出。PyTorch导出时选择opset_version,实际上在决定“允许用哪个版本的ONNX算子来描述你的模型”。
选太低,某些PyTorch算子找不到对应的ONNX实现,会直接导出失败;选太高,后端的ONNX Runtime版本如果太老,又可能不认识新算子集,推理加载时报类似“Unsupported opset version”的错误。所以opset版本并不是越高越好,而是要看目标runtime的“接受范围”。
我的实践建议是:在没有特殊算子需求的情况下,先选择opset_version=14,这个版本对绝大多数常规CNN、Transformer模型都支持得不错,主流的onnxruntime和TensorRT版本也兼容。如果导出时报某个算子无法映射,再根据报错提示逐步调高opset版本,比如15、16、17甚至更高,直到能导出为止。
如果你要部署到瑞芯微这种边缘NPU平台,尤其要提前查一下目标SDK工具链支持的opset上限。很多厂商工具链对opset版本要求比较严格,有的只支持到11或者12,这种情况你就需要把模型调整到该工具链约束范围内,而不是一味追求新版。
4. 完整实操:从PyTorch模型到ONNX并验证
4.1 准备一个用于示例的模型并加载权重
下面用一个简单但结构完整的CNN分类模型走一遍全流程。你只需要把下面的模型类替换成自己项目的模型即可。
python复制import torch
import torch.nn as nn
class SimpleCnn(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.features = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding=1),
nn.ReLU(inplace=True),
nn.MaxPool2d(2),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.ReLU(inplace=True),
nn.AdaptiveAvgPool2d((1, 1)),
)
self.classifier = nn.Linear(64, num_classes)
def forward(self, x):
x = self.features(x)
x = torch.flatten(x, 1)
return self.classifier(x)
真实项目里从.pth恢复权重,核心就一行:
python复制model = SimpleCnn()
state_dict = torch.load("your_model.pth", map_location="cpu")
model.load_state_dict(state_dict)
model.eval()
如果你暂时没有训练好的权重,只想快速体验转换,也可以直接实例化模型导出。转换是否成功、前后推理结果是否一致,本来就和权重具体长什么样没有关系。它只验证一件事:同一种输入经过PyTorch计算出来的结果,和一个经过ONNX Runtime计算出来的结果,是不是能对齐。
4.2 导出ONNX并做静态检查
把上一节的SimpleCnn导出成ONNX文件:
python复制model = SimpleCnn(num_classes=1000)
model.eval()
dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
model,
dummy_input,
"simple_cnn.onnx",
export_params=True,
opset_version=14,
do_constant_folding=True,
input_names=["input"],
output_names=["output"],
dynamic_axes={
"input": {0: "batch"},
"output": {0: "batch"},
},
)
导出完成后不要急着拿去部署,先做一个快速检查:
python复制import onnx
onnx_model = onnx.load("simple_cnn.onnx")
onnx.checker.check_model(onnx_model)
print("onnx check passed")
onnx.checker.check_model会检查模型结构是否合法,包括节点的输入输出、张量维度、算子定义等。但它只能证明这个ONNX文件“看起来没问题”,不能证明里面的数值逻辑和你原来的PyTorch模型一致。真正的一致性要靠推理对比。
4.3 用onnxruntime推理,对比和PyTorch的输出差异
这一步是整条链路里最值得投入时间的环节。很多同学导出ONNX后直接丢到服务端,结果推理结果一团糟,又找不到原因,就是因为省掉了数值一致性验证。
用ONNX Runtime把导出的模型读回来,用同一个随机输入,分别跑PyTorch和ONNX Runtime:
python复制import onnxruntime as ort
import numpy as np
import torch
input_data = torch.randn(4, 3, 224, 224)
# PyTorch结果
with torch.no_grad():
torch_out = model(input_data).numpy()
# ONNX Runtime结果
sess = ort.InferenceSession("simple_cnn.onnx", providers=["CPUExecutionProvider"])
input_name = sess.get_inputs()[0].name
output_name = sess.get_outputs()[0].name
ort_out = sess.run(
[output_name],
{input_name: input_data.numpy()},
)[0]
print("torch output shape:", torch_out.shape)
print("ort output shape:", ort_out.shape)
diff = np.abs(torch_out - ort_out)
print("max abs diff:", diff.max())
当输出shape都是(4, 1000),max abs diff在1e-5量级甚至更小时,基本可以认为转换成功。这种极小误差来自浮点数累加顺序不同,是正常现象。
这里有一个需要注意的细节:sess.run接受的是numpy数组,不是torch Tensor。所以喂数据前记得做.numpy()转换,输入dtype也要保持一致,通常都是float32。如果dtype不匹配,Runtime会报“input type mismatch”。
如果发现输出的运行结果和PyTorch差异很大,第一时间不要怀疑工具,先从模型本身找原因。最常见的问题是忘记model.eval(),导致BatchNorm统计方式和Runtime不一致;其次是模型里存在动态控制流,而导出过程只记录了某条固定路径,导致其他分支的逻辑根本没进ONNX图里。
4.4 用Netron可视化查看导出的图结构
检查数值没有问题之后,我还有一个习惯:把.onnx文件拖进Netron里看一眼结构。Netron是一个开源的神经网络可视化工具,支持ONNX、TorchScript、TensorFlow等格式,浏览器打开网页版就能用,也可以本地安装。
可视化能看到的东西很直观:输入节点叫什么、输出节点叫什么、权重都存在哪些节点上、整张图的算子链路是否和你预期一致。对于一些大模型,不需要一个一个节点看,只确认输入输出和几个关键分支没走错就够。这个动作成本很低,却经常能帮你提前发现“结构里混进了一个不该有的算子”之类的问题。
5. 转换过程中最容易踩的坑
5.1 数值对不上,先检查是不是漏了model.eval()
推理结果不一致是所有人遇到的第一类问题。排查顺序我一般固定为:
- 导出前有没有调用
model.eval()? - 输入给模型的预处理是否和训练时完全一致?
- 模型里是否存在依赖Python控制流的动态分支?
- 对比时是否保证了同一输入、同一dtype、同一个随机状态?
其中第一点出现的概率最高。训练模式下的BatchNorm会使用当前batch统计量,而部署推理普遍期望使用训练阶段累计的running_mean/running_var。忘记model.eval()等于拿了一个实时统计的逻辑去作图,误差自然大。
第二个点经常出现在目标检测模型里。很多YOLO训练代码会做归一化、letterbox、颜色通道调整,如果没有把输入数据处理好,ONNX Runtime里的推理结果异常,不一定是谁的问题。
5.2 算子不支持或导出失败的兜底思路
torch.onnx.export遇到一些PyTorch自定义算子或非常新的实现时,偶尔会抛“Unsupported operator”这类错误。比如某些模型内部用了复杂的切片、花式索引,或者依赖了不常见的第三方算子,导出器可能真的很难处理。
我的兜底思路通常分几步:
- 尝试把
opset_version调高,看看是不是旧版本算子集缺少对应映射。 - 如果还是失败,把出错的局部模块从模型里“摘”出来,单独导出,确认是哪一层造成的。
- 对实在无法导出的算子,在部署阶段把它移到ONNX图外,用numpy或纯Python实现后处理。比如NMS这类后处理,把模型输出原始预测再在Runtime之外做非极大值抑制,是很多目标检测模型的实际部署方案。
- 必要时手动改写模型前向,把花哨操作改写成更常见的卷积、池化、矩阵运算组合。
经验谈:大多数导出失败的根因不是PyTorch太弱,而是业务代码里在forward里塞了太多“部署时才需要处理”的逻辑,例如可视化、绘图、调试打印。导出前在模型类边上保留一份“干净的推理版本”,会省很多事。
5.3 动态batch失效与reshape相关报错
动态batch失效的场景通常有两种表现:导出时明明配了dynamic_axes,运行时换batch却依然报错;或者导出的模型在Netron里看到输入still是静态维度。
第一种情况一般是dynamic_axes里的维度没写全,或者输入的某个中间tensor经历了把shape写死的操作。比如某段代码用x.view(-1, 64),view本身能感知前导维度,但如果你写的是x[:1]这种硬编码第一个维度size的切片,动态性就可能被破坏。排查这种问题时,我会去源码里搜view、reshape、slice、flatten,逐个判断它们是否适配动态shape推理。
第二种常见原因是后端加载模型时指定了静态输入,比如某些SDK只接受固定shape的模型,那么ONNX层面即便动态也白搭。所以要分清问题究竟出在导出的模型里,还是出在加载模型的那层runtime上。
5.4 模型出来又大又慢的初步优化建议
转出来的ONNX文件如果特别大,第一反应不一定是量化,先看一下有没有导出不必要的东西。比如模型里包含优化器状态、训练日志、缓存等,通常会在导出前用load_state_dict重新加载一次权重,而不是直接把整个checkpoint塞进去。
对已经变小的ONNX还觉得慢,可以尝试开启do_constant_folding,把图中固定的常量计算折叠掉,减少运行时计算。也可以看一下模型里有没有冗余的层或分支,在导出前的“干净版本”里去提醒自己“哪些层在推理时根本不需要”。
关于INT8量化,ONNX Runtime本身提供了相对成熟的量化工具,可以做动态量化和静态量化。但量化这块涉及精度明显下降、校准集收集、算子融合等一系列问题,一次讲不清楚,更适合单独开一篇。如果你打算把导出的ONNX部署到边缘NPU,通常量化是逃不开的一步,但先把不带量化的ONNX流程跑通验证正确,再考虑量化,会稳妥得多。
最后说点我自己的体会。这个转换流程我做了很多次,从最初的ResNet到后来的各种检测分割模型,踩坑踩多了以后发现,大部分问题其实都出现在最开始那几步:模型没置为eval、输入样例和真实部署输入不一致、输入输出名字不固定、opset版本选得随意。这些看似不起眼的细节,放到生产环境里一个比一个致命。现在我每做一个模型,都会先把“PyTorch输出和ONNX Runtime输出能不能对得上”作为最基本的一道门槛,过了这道门槛才敢继续往下走。你要是正在为部署发愁,我的建议就是先把这套最小闭环跑通,它虽然不能保证解决所有平台问题,但能让你在后面面对各种加速SDK时,少浪费很多排查时间。
