
Phase A · Step 3模型导出 ONNX做了这么多年模型部署我越来越觉得“导出”这件事才是整个落地链条里最容易翻车的一环。训练时 loss 掉得挺漂亮验证集准确率也说得过去结果到了导出 ONNX 这一步各种算子不支持、动态维度对不上、精度对不齐的问题全冒出来了。Phase A 的 Step 3 就是专治这种“训练一时爽部署火葬场”的尴尬。这个阶段的核心工作很明确把训练框架里跑得好好的模型转换成一个标准的、跨平台的中间表示——ONNX 模型文件让后续推理引擎能够直接加载和高效运行。这篇内容适合正在做模型部署、推理加速、或者准备把模型接到某个跨平台系统里的开发同学尤其是那些已经训练完模型、正准备往线上环境迁移的人。我会把整个导出流程拆开讲为什么要把模型转成 ONNX、导出前要确认哪些事情、主流框架导出时有哪些参数需要认真配置、以及导出之后怎么验证模型没有“变质”。这些都是实际项目里反复踩过的坑希望能帮你少走点弯路。1. 把模型变成通用的“交换格式”这一步到底在解什么题1.1 从训练代码到推理引擎中间隔着的不是一层皮很多人会觉得模型导出不就是把权重存下来吗其实远没这么简单。训练框架里描述模型用的是 Python 对象和动态计算图比如 PyTorch 里你定义一个nn.Module它内部包含了一堆参数张量和前向计算逻辑。推理引擎要跑这个模型需要的是静态的计算图知道有哪些节点、每个节点的输入输出张量形状、权重参数放在哪里、整个图的执行顺序是什么。这个信息差就是导出要解决的核心问题。我见过不少同学直接在部署环境里装一套 PyTorch然后加载训练好的权重做推理。这样做在原型验证阶段没毛病但一旦到了正式环境问题就来了推理服务每启动一次要等模型初始化CPU 内存占用居高不下多路并发时性能上不去。更麻烦的是训练框架升级一个新版本某些算子的行为可能就变了线上模型一夜之间“失灵”排查起来极其痛苦。ONNX 就是用来打破这种框架绑定的——它是一种计算图的规范描述把模型固化成纯结构的中间文件推理引擎只要解析这个文件就能执行不依赖任何训练框架。1.2 为什么不是直接用训练框架而是用 ONNXONNX 的全称是 Open Neural Network Exchange直译就是“开放的神经网络交换格式”。它本身不负责训练也不负责推理而是充当一个中立的数据交换标准。类比一下它就像是你把一份 Word 文档导出成 PDF——PDF 不是用来编辑的但任何设备、任何阅读器打开它版面都不会乱。模型导出成 ONNX 后你的模型就跟 PyTorch、TensorFlow 这些训练框架解耦了后面跑在哪个推理引擎上完全看部署需求。选 ONNX 而不直接选某个厂商的专有格式有几个实际原因。第一生态覆盖面广主流的推理引擎基本都支持加载 ONNX比如 ONNX Runtime、TensorRT、OpenVINO 等覆盖面广意味着你不必在选型阶段就被绑死。第二算子语义标准化ONNX 定义了一组标准算子集不同框架导出的模型都会映射到这组算子上虽然映射过程偶尔有坑但相比直接用框架私有格式还是规范得多。第三模型在导出后可以做图优化常量折叠、算子融合、精度校准这些操作都能在 ONNX 图层面进行为后续的性能调优打基础。1.3 ONNX 作为图的统一抽象节点、张量、初始器要搞清楚导出时哪些环节容易出问题先得理解 ONNX 模型文件的内部结构。一个 ONNX 文件本质上是一个 Protocol Buffers 序列化的图结构里面有三个最核心的元素节点Node代表一个算子操作比如 Conv、Relu、MatMul每个节点都有op_type算子类型、input列表、output列表。张量ValueInfo描述数据流每个张量有自己的名字、数据类型和形状信息。张量在节点之间传递构成整个数据流图。初始器Initializer存放模型的权重参数比如卷积核权重、BN 层的均值和方差。初始器会作为节点的输入之一参与计算。导出的过程本质上就是把训练框架里动态的计算逻辑逐步翻译成这种静态的、节点式的描述。翻译过程中最容易出问题的就是某个训练框架里的算子在 ONNX 算子集里找不到对应的映射或者映射过去之后语义有一点点不一样。这也是为什么导出的第一步不是写代码而是先搞清楚你用的算子和 ONNX 算子集的对应关系。2. 导出 ONNX 之前先过一遍硬件、环境和模型底座2.1 明确目标推理后端opset 才是第一个决策点导出 ONNX 不是一个按钮就能搞定的事动手之前先要回答一个问题这个模型导出后会跑在哪个推理引擎上不同的推理引擎对 ONNX 算子版本的支持程度不一样这个差异直接影响你选择的opset_version。opset_version算子集版本定义了导出时使用的 ONNX 算子版本。版本越高能用的新算子越多但目标推理引擎可能还没来得及支持。以我的经验当前大家在导出时默认选opset_version11到opset_version17之间具体选多少有个简单原则先看目标推理引擎支持的最高 opset 版本不要盲目追新。举个例子某次项目里目标环境是某个边缘设备推理引擎只支持 ONNX opset 11但我当时用 PyTorch 默认的 opset 17 导出模型文件在本地用 ONNX Runtime 跑得好好的一放到目标设备就报“unsupported operator”。后来把 opset 降回 11问题直接解决。所以在导出之前先去目标推理引擎的文档里确认支持的 opset 版本把这个参数定下来后面能省掉很多麻烦。2.2 把模型固定下来权重固化和输入输出定义导出之前还有一件事要做把模型的权重固化下来。训练好的模型权重一般存在 checkpoint 文件里导出时需要加载权重并把整个模型设置为推理模式。在 PyTorch 里对应model.eval()这一步不只是切换 dropout 和 BN 的行为更关键的是让模型的 forward 逻辑从训练状态切换到静态推理状态。我见过一个比较典型的错误模型没切eval()就直接导出结果 BN 层的统计量还是用 batch 内统计的导出来的 ONNX 模型在推理时表现和训练时完全不一致。所以导出前养成一个习惯先写三行代码model.eval() model model.to(device) # 构造一个虚拟输入走一次 forward确保模型能正常跑通 dummy_input torch.randn(batch_size, channels, height, width).to(device) with torch.no_grad(): output model(dummy_input)这个过程叫 warm-up一方面是确认模型 forward 没有报错另一方面是让模型内部的某些惰性初始化逻辑执行完。打包导出的时候模型的状态就锁定在“已就绪”的推理模式。输入输出的命名也很关键。ONNX 图里的每个张量都要有一个唯一的字符串名字导出时如果不显式指定框架会生成类似onnx::Conv_0这样的名字后续做推理引擎对接时非常不友好。实际项目中我一般会显式指定语义化的名称比如input、output这样后面的数据处理和结果解析代码可读性好很多。2.3 动态与静态 shape 的选择以什么维度为界导出 ONNX 时输入输出的张量形状要么固定要么标注为动态。固定形状的意思是输入张量的每一维都写死比如[1, 3, 224, 224]这样导出的图里所有中间节点的形状都能静态推断出来推理引擎可以做更多图优化性能通常更好。动态形状则允许某些维度是变量比如 batch 维度可以任意变化用batch这样的符号来表示。我的建议是除非你的业务场景确实需要不同 batch size 的输入比如服务端推理需要动态并发否则优先导出固定形状。原因很简单固定形状的模型在推理引擎里更容易做内存预分配和算子融合性能表现也更稳定。一些推理引擎对动态形状的支持还存在性能损耗对动态 input shape 会触发额外的 shape 推断跑起来比固定形状慢不少。如果确实需要动态维度操作时只把需要的维度标记为动态其他维度仍然固定别把所有维度都放开。比如目标检测类模型往往输入图片尺寸不固定就只把高和宽两个维度做成动态batch 仍然固定为 1这样平衡了灵活性和性能。导出后再验证一次各动态维度组合下的输出与原始模型是否一致这个验证不能省。3. PyTorch 和 TensorFlow 的导出实操常用配置逐个过一遍3.1 torch.onnx.export 必填参数讲解PyTorch 导出 ONNX 的核心 API 是torch.onnx.export这个函数看起来简单但参数配置讲究不少。基本调用形式是这样torch.onnx.export( model, # 要导出的模型 dummy_input, # 虚拟输入用于追踪计算图 model.onnx, # 输出文件路径 export_paramsTrue, # 是否导出权重参数 opset_version13, # ONNX 算子集版本 do_constant_foldingTrue, # 是否做常量折叠 input_names[input], # 输入节点名称 output_names[output], # 输出节点名称 dynamic_axes{ # 动态维度定义 input: {0: batch_size}, output: {0: batch_size} } )逐个说一下这些参数的实际作用。export_params必须为True否则导出的模型不包含权重只有一个空壳子。opset_version前面已经说了是第一个要确认的决策点。do_constant_folding建议开启它会把那些只依赖权重和常量的子图直接计算成常量比如 BN 层在推理模式下可以合并进前面的卷积层这会显著减少图的节点数和推理耗时。input_names和output_names前面提到了不只是为了好看而是在后续用 ONNX Runtime 或推理引擎加载模型时你可以通过名字直接获取输入输出的句柄不显式命名的话找接口的过程会非常别扭。dynamic_axes用字典结构描述哪些输入输出的哪些维度是动态的键是前面设置的名字值是一个字典里面维度索引: 维度语义名比如{0: batch_size}表示第 0 维是动态的 batch 维度。3.2 一个实例用 torch 导出并确认节点结构我拿一个模拟的图像分类模型来演示完整流程这个模型结构不复杂就是卷积加残差块的堆叠。假设模型文件已经训练好保存为一个 checkpoint。先加载模型和权重然后切推理模式import torch import torch.nn as nn # 以自定义的模拟网络为例 class SimpleNet(nn.Module): def __init__(self, num_classes10): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, kernel_size3, stride2, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue), nn.Conv2d(32, 64, kernel_size3, stride2, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), ) self.classifier nn.Linear(64 * 7 * 7, num_classes) def forward(self, x): x self.features(x) x x.view(x.size(0), -1) return self.classifier(x) model SimpleNet(num_classes10) state_dict torch.load(simulate_net.pt, map_locationcpu) model.load_state_dict(state_dict) model.eval()接着准备一个形状正确的虚拟输入。这里注意虚拟输入的形状必须符合模型期望的输入形状而且数值是什么不重要因为导出追踪的是计算流程不依赖具体数值。dummy_input torch.randn(1, 3, 28, 28) torch.onnx.export( model, dummy_input, simplenet.onnx, export_paramsTrue, opset_version13, do_constant_foldingTrue, input_names[input], output_names[output], dynamic_axesNone, )导出完成后可以用onnx库检查模型的基本信息import onnx model_onnx onnx.load(simplenet.onnx) onnx.checker.check_model(model_onnx) # 做一致性校验 print(fONNX 模型文件大小: {model_onnx.ByteSize()} bytes) print(f图中节点数: {len(model_onnx.graph.node)})这里onnx.checker.check_model会检查图结构是否合法节点输入输出是否引用正确、算子定义是否在算子集范围内等。这是导出后第一道验证关口如果这里都过不了后面推理引擎加载肯定也会出问题。3.3 TensorFlow SavedModel 转 ONNX 的常用流程TensorFlow 生态导 ONNX 稍微绕一点因为 TensorFlow 的图结构和 PyTorch 的动态图差别较大。如果是从 TensorFlow 出发我一般建议先用h5或SavedModel格式把模型保存好再用tf2onnx这个转换工具来处理。基本流程是pip install tf2onnx python -m tf2onnx.convert \ --saved-model ./saved_model_dir \ --output model.onnx \ --opset 13 \ --inputs input:0 \ --outputs output:0这里需要注意--inputs和--outputs指定的是 TensorFlow 图中张量的完整名称包括端口号后缀:0。如果不确定名字可以先加载 SavedModel打印出输入输出签名。tf2onnx在转换过程中会做算子映射大部分常见算子都能搞定但碰到一些自定义算子就会比较头疼。我的建议是如果项目里用了 TensorFlow 比较偏门或自定义的层提前检查一下tf2onnx支持的算子列表看看有没有覆盖没有的话要么改模型结构要么在转换后手动做子图替换。另外有个小坑TensorFlow 模型导出的 ONNX 里经常会出现一些多余的转换节点比如Identity节点、不必要的Transpose节点这些不影响正确性但会增加推理耗时。可以用onnxsimonnx-simplifier做一次图简化把冗余节点清掉。这个工具我用下来效果挺明显有时能减少百分之二三十的节点数。3.4 各框架导出时的对比与选择不同框架导出 ONNX 的体验和产物质量有差异。就以我常用的 PyTorch 和 TensorFlow 做对比PyTorch 的torch.onnx.export采用 tracing 机制也就是跟着虚拟输入实际跑一遍 forward把执行过的算子记录下来。这个机制的优点是实现简单、覆盖面广缺点是如果模型里有依赖于数据内容的分支逻辑比如某些循环次数由输入决定tracing 只会记录当前虚拟输入走的那条路径导出的图可能不完整。TensorFlow 这边的tf2onnx则是在图层面做映射因为是静态图所以不存在 tracing 路径不全的问题。但 TensorFlow 的图本身包含大量框架内部节点转换过程更复杂产物也更容易出现冗余。选择建议是如果项目从零开始优先用 PyTorch 训练因为导出链路更短如果是老项目本来就在 TensorFlow 上那就直接用tf2onnx转不用考虑迁移训练框架成本太高。至于 JAX、PaddlePaddle 之类的框架做法大同小异核心还是那几个参数和验证步骤。4. 导出不是结束验证才是离线上线最近的一步4.1 使用 ONNX Runtime 完成加载与输出对比模型导出后最基本的验证就是用 ONNX Runtime 加载跑一遍和原始模型的结果做对比。ONNX Runtime 是微软开源的推理引擎跨平台支持很好也是我最常用的验证工具。加载模型很简单import onnxruntime as ort import numpy as np sess ort.InferenceSession(simplenet.onnx, providers[CPUExecutionProvider]) input_name sess.get_inputs()[0].name output_name sess.get_outputs()[0].name # 构造一个输入 input_data np.random.randn(1, 3, 28, 28).astype(np.float32) output_ort sess.run([output_name], {input_name: input_data})这里有个细节输入张量在 ONNX Runtime 里的类型必须是 numpy 数组而且 dtype 要和模型的输入要求一致比如float32不一致的话会报类型错误。对比原始模型时用完全相同的输入喂给 PyTorch 模型和 ONNX Runtime然后比较输出向量。因为浮点数运算存在舍入误差两个结果不可能完全一致需要设置一个容差范围。一般来说数值精度控制在 1e-4 到 1e-5 范围内比较合理超过这个范围就要排查问题。4.2 数值一致性检查的正确打开方式数值一致性检查有几个容易忽略的细节。第一比较之前必须先做 softmax 或者 sigmoid把 logits 转成概率分布再比否则只看原始输出的绝对值没什么意义。第二要看最大绝对误差而不是平均误差因为平均误差会被大量微小误差稀释个别输出维度上的大偏差反而能说明问题。一个实用的比较脚本结构是import torch import numpy as np import onnxruntime as ort # 用原始框架生成参考输出 model.eval() with torch.no_grad(): ref_output model(torch.from_numpy(input_data)).numpy() # 用 onnxruntime 生成推理输出 sess ort.InferenceSession(simplenet.onnx) ort_output sess.run([output_name], {input_name: input_data})[0] # 计算最大绝对误差 max_diff np.max(np.abs(ref_output - ort_output)) mean_diff np.mean(np.abs(ref_output - ort_output)) print(fMax abs diff: {max_diff:.6f}) print(fMean abs diff: {mean_diff:.6f})如果 max diff 超过 1e-3说明导出过程至少在某一步做了数值类型转换或者算子实现有差异需要进一步定位。定位方法通常是二分法把模型前半段和后半段分别导出找出误差究竟是哪一层开始引入的。这个排查过程虽然耗时但能做一次之后你对模型内部结构和算子行为会有更清楚的认识。4.3 常见报错与排查记录速查表做过的导出项目多了之后我发现大多数报错类型都是可以预判和快速定位的。这里整理一个速查表方便你遇到问题时对照排查。报错现象可能原因排查方向提示Unsupported operator或Op not registered目标推理引擎支持的 opset 版本低于导出时的 opset降低opset_version或者换用兼容算子推理引擎加载时报 shape 不匹配动态维度配置有误或输入 shape 超出范围检查dynamic_axes定义确认实际输入 shape 与配置一致输出结果偏差大max diff 1e-3模型没有切到 eval 模式、输入 dtype 不一致、量化导致精度损失确认导出前model.eval()、检查输入输出 dtype、排除量化影响导出后节点数异常多框架生成了冗余的转换节点用 onnxsim 做图简化清除冗余节点自定义算子报错训练框架里的自定义算子不在 ONNX 算子集内在导出时注册自定义符号或者在 ONNX 图里用子图替换模型加载非常慢图过大、常量折叠没开启开启do_constant_folding适当做模型剪枝这张表是我整理的一个起点实际项目里会有更多变体但排查思路基本是一致的先分清楚是导出阶段出的问题还是推理引擎加载阶段出的问题再顺着图结构一步步排查。4.4 实际生产转化中的典型报错和应对挑一个实际遇到的案例讲。某次在做图像分割模型的 ONNX 导出时模型在 PyTorch 里正常导出后 ONNX Runtime 也能跑但输出结果始终和原始模型差出一大截。排查了一晚上最后发现是输入图片的归一化方式不一致训练时用的是均值[0.485, 0.456, 0.406]加标准差[0.229, 0.224, 0.225]的方式而用 ONNX Runtime 部署时预处理代码写错了顺序导致输入给模型的张量分布完全不对。这个报错形态不是直观的“跑不了”而是“跑出来的结果不对”。这类问题最隐蔽因为它源于导出流程之外的预处理链路恰恰说明验证阶段不能只看模型文件本身对不对要把整条推理链路串起来测试。后来我在验证脚本里把原始框架的预处理函数和 ONNX Runtime 的预处理逻辑做成同一个函数输入同一张测试图从数据入口到结果出口全链路对比才彻底解决这类问题。5. 导出后搭好模型“护照”走向多后端部署5.1 模型文件本身的信息检查与模型结构可视化模型导出成 ONNX 文件之后就像是一个人拿到了一份“护照”记录了他的身份信息。ONNX 文件里除了计算图还有模型的元信息生产者名称、使用的算子集版本、生成的工具版本等。这些信息在后续排错时非常有用。检查标准做法是import onnx model onnx.load(simplenet.onnx) print(fIR 版本: {model.ir_version}) print(f生产者: {model.producer_name} {model.producer_version}) print(f算子集: {[(op.domain, op.version) for op in model.opset_import]}) print(f输入: {[(inp.name, inp.type.tensor_type.elem_type, inp.type.tensor_type.shape) for inp in model.graph.input]}) print(f输出: {[(out.name, out.type.tensor_type.elem_type) for out in model.graph.output]})另一个推荐做的事情是用 Netron 打开模型文件直接把计算图可视化出来。Netron 是一个轻量的模型可视化工具支持 ONNX 格式打开后能看清每个节点的连接关系对于排查图结构异常很有帮助。我在导出后习惯先做一次onnx.checker再打开 Netron 扫一眼整体结构确认从输入到输出的路径符合预期。5.2 量化与剪枝对导出模型的影响模型导出 ONNX 之后很多团队会考虑做量化来减小体积、提升推理速度。量化指的是把 float32 的权重和激活值降低精度表示比如 int8。量化对 ONNX 模型的影响很大必须放在导出流程里统一考虑否则会导致精度显著下降。常见做法有两种一种是先导出 ONNX再做量化感知训练或训练后量化另一种是在原始训练框架里做量化感知训练再导出 ONNX。我个人的经验是后一种的效果通常会更好因为量化感知训练会把量化误差纳入训练过程模型权重会主动适应低精度的表示导出的 ONNX 模型在量化后精度损失更小。不过量化后再导出 ONNX有时会遇到量化算子比如 QLinearConv 这类带量化参数的算子和普通算子混在一起的情况。部分推理引擎对混合精度图的支持并不完善要提前查清楚目标平台支持哪些量化算子。如果目标平台不支持某些量化算子比较务实的替代方案是导出一个 float32 版本的 ONNX 模型再到推理引擎侧做离线量化和格式转换虽然中间多一道工序但兼容性最好。剪枝的影响主要在算子层面剪掉不重要的权重连接后某些层可能变成稀疏结构。ONNX 算子集对稀疏张量的支持还不算完善稀疏权重要么被转换为稠密表示要么需要特殊的稀疏算子替代。所以对剪枝后的模型做导出时不要预期它一定比剪枝前小——如果转成稠密表示文件体积可能不降反升。5.3 多后端部署意味着什么ONNX 模型真正的价值在于同一个.onnx文件可以在不同推理引擎之间切换而不用担心重复造模型。你把模型导出成 ONNX 之后基本就具备了一种“一次导出、多处部署”的能力。同一个文件既可以用 ONNX Runtime 跑普通 CPU 推理也可以转到 TensorRT 引擎做 GPU 优化还可以放到边缘设备的推理框架里做端侧部署。这带来的第一个好处是选型成本大幅下降。前期你可以先用 ONNX Runtime 快速验证精度和性能确认业务效果符合预期后再决定是否需要给某个特定的推理引擎做深度优化。第二个好处是跨平台能力。团队里有人用 CPU 环境调试有人用 GPU 环境跑批量推理有人要在嵌入式设备上做实时推理都可以从同一个 ONNX 文件出发各自做落地适配。但多后端部署也不是没有代价。不同推理引擎对 ONNX 算子支持差异很大同一个模型在一个引擎里能跑得好好的换一个引擎却可能报算子不支持。所以多后端部署的关键是导出的 ONNX 模型要尽量用“最通用”的算子子集避免使用太新或者太偏门的算子。有时候为了兼容老版本推理引擎不得不牺牲一点点图优化空间换取跨平台的稳定性这个权衡在项目规划时就要想清楚。6. 结尾关于导出这件事最后想分享几句心里话我在实际项目中多次被“说得容易做起来一堆坑”这句话教育模型导出这件事尤其如此。每次导出 ONNX 都会遇到新的边界情况可能是某个算子在新版本里的行为变了可能是目标推理引擎对动态 shape 的支持和预期不符也可能只是预处理代码里一个小小的dtype写错了。但反过来看每一次排查过程都是在加深对模型结构和推理链路的理解。给大家一个实用建议把导出 ONNX 的脚本固定成一个标准模板包含模型加载、eval 切换、虚拟输入构造、导出参数配置、ONNX Runtime 对比验证这几步每次新模型都走同一套流程。这样一旦出问题你很快能判断是模型的特殊结构导致的还是导出参数配置的问题而不是每次都要从零开始排查。另外一个忠告是不要在生产环境里临时改导出参数。opset 版本、dynamic_axes、输入输出名称这些一旦定了后面会嵌进部署系统里改一个模型文件名都可能引发联动问题。把导出流程当成软件工程的一部分来管理配置参数写成配置文件导出和验证做成固定脚本才能保证模型从训练到生产环境之间的通道稳定可靠。