模型导出 ONNX (Exporting to ONNX)


章节概述

ONNX(Open Neural Network Exchange)是从 Python 训练到 C++ 部署的桥梁。本章演示如何将 PyTorch 模型导出为 .onnx 文件,理解 ONNX 的内部结构(图、节点、张量),用 onnxruntime 在 Python 侧验证模型正确性,为下一章 C++ 部署做好准备。

核心理念:ONNX 是 AI 领域的”ELF 可执行文件”。就像 C 编译器将源代码编译为可在不同平台上运行的二进制文件,torch.onnx.export 将 Python 模型”编译”为可在不同推理引擎(ONNX Runtime、TensorRT、OpenVINO)上运行的 .onnx 文件。C++ 程序加载 .onnx → 执行推理 → 释放资源,就像加载 .so 动态库。理解 ONNX 的输入/输出形状是成功部署的关键。


第一节:ONNX 概述

1.1 什么是 ONNX

ONNX(Open Neural Network Exchange,开放神经网络交换格式)由 Facebook 和 Microsoft 于 2017 年推出,是一个开放的深度学习模型互操作标准。

graph LR
 subgraph Train["Python (训练/验证)"]
 PyTorch["PyTorch"]
 TensorFlow["TensorFlow"]
 Sklearn["sklearn (skl2onnx)"]
 Jax["JAX"]
 Mxnet["MXNet"]
 end

 PyTorch -- "export()" --> ONNX["ONNX Runtime<br/>Runtime (CPU/CUDA/DirectML)<br/>TensorRT (NVIDIA)<br/>OpenVINO (Intel)<br/>TVM (Apache)<br/>C API / C++ API"]
 TensorFlow --> ONNX
 Sklearn --> ONNX
 Jax --> ONNX
 Mxnet --> ONNX

1.2 ONNX vs TorchScript vs libtorch

方案格式优点缺点
ONNX.onnx跨框架、广泛支持部分算子不支持
TorchScript.ptPyTorch 原生、完整支持仅限 PyTorch 生态
libtorch.pt + C++直接用 C++ 调 PyTorch部署体积大 (>500MB)

推荐策略:优先 ONNX → 如果 ONNX 不支持某些算子 → 用 TorchScript/libtorch → 如果 GPU 推理 → 从 ONNX 转 TensorRT。


第二节:torch.onnx.export 实战

2.1 导出最简单的模型

python -c "
import torch
import torch.nn as nn
 
# 1. 定义一个简单模型
class SimpleModel(nn.Module):
 def __init__(self):
 super().__init__()
 self.linear = nn.Linear(4, 3)
 def forward(self, x):
 return self.linear(x)
 
model = SimpleModel()
model.eval()
 
# 2. 创建 dummy input — 形状必须和实际输入一致
dummy_input = torch.randn(1, 4) # batch=1, features=4
 
# 3. 导出 ONNX
torch.onnx.export(
 model,
 dummy_input,
 'simple_model.onnx',
 export_params=True, # 保存模型参数(权重)
 opset_version=17, # ONNX 算子集版本
 input_names=['input'], # 输入节点名称
 output_names=['output'], # 输出节点名称
 dynamic_axes={ # 动态轴:batch 维度可变
 'input': {0: 'batch'},
 'output': {0: 'batch'},
 },
)
print('Exported to simple_model.onnx')
print(f'File size: {__import__(\"os\").path.getsize(\"simple_model.onnx\")} bytes')
"

2.2 关键参数详解

参数含义注意事项
model要导出的模型必须先调用 model.eval()
args示例输入(tuple 或有多个参数)形状决定 ONNX 的参数尺寸!
f输出文件名后缀 .onnx
export_params是否保存权重一般为 True
opset_versionONNX 算子集版本11 是稳定版,17+支持更多算子
input_names输入节点名字列表C++ 侧用名字获取输入
output_names输出节点名字列表C++ 侧用名字获取输出
dynamic_axes可变的维度batch 维度经常是动态的

2.3 动态轴详解

# dynamic_axes 使 batch 维度可变
dynamic_axes = {
 'input': {0: 'batch_size'}, # 输入的第 0 维可变
 'output': {0: 'batch_size'}, # 输出的第 0 维可变
}
 
# 导出的模型可以接受任意 batch size:
# batch=1: (1, 4) 
# batch=32: (32, 4) 
# batch=None: 错误 — 没有标记为动态的维度必须匹配 dummy_input

C++ 部署关键:如果不设置 dynamic_axes,导出后 batch 维度被固定为 dummy_input 的大小。C++ 程序必须传入相同 batch 大小的输入。设置动态轴后,batch 大小可变,更灵活。

2.4 导出 CNN(图像模型)

python -c "
import torch, torch.nn as nn
 
class CNN(nn.Module):
 def __init__(self):
 super().__init__()
 self.conv = nn.Conv2d(3, 16, 3, padding=1)
 self.fc = nn.Linear(16, 10)
 def forward(self, x):
 x = self.conv(x) # (B,3,H,W) → (B,16,H,W)
 x = x.mean(dim=[2, 3]) # 全局平均池化 → (B,16)
 return self.fc(x) # → (B,10)
 
model = CNN().eval()
dummy_input = torch.randn(1, 3, 224, 224)
 
torch.onnx.export(
 model, dummy_input, 'cnn_model.onnx',
 input_names=['image'],
 output_names=['class_logits'],
 dynamic_axes={'image': {0: 'batch', 2: 'height', 3: 'width'}},
 opset_version=17,
)
print('CNN exported!')
# C++ 侧: 可以传入 (N, 3, H, W) 任意大小的图像
"

第三节:理解 ONNX 模型结构

3.1 ONNX 模型的内部表示

ONNX 模型是一个 protobuf 格式文件,内部结构如下:

graph TB
 MP["ModelProto"] --> G["graph"]
 G --> INPUT["input<br/>(ValueInfoProto)<br/>模型输入<br/>name, type, shape"]
 G --> OUTPUT["output<br/>(ValueInfoProto)<br/>模型输出<br/>name, type, shape"]
 G --> INIT["initializer<br/>(TensorProto)<br/>权重参数<br/>每个权重的名称、类型、原始数据"]
 G --> NODE["node<br/>(NodeProto × N)<br/>计算节点<br/>op_type, inputs[], outputs[], attributes"]
 MP --> OI["opset_import<br/>domain, version"]

每个 node 代表一个算子(如 Conv, Relu, Gemm, Softmax),initializer 存储权重。

3.2 用 Python 检查 ONNX 模型

python -c "
import onnx
 
model = onnx.load('simple_model.onnx')
 
# 输入信息
print('=== Inputs ===')
for inp in model.graph.input:
 print(f' Name: {inp.name}')
 shape = [d.dim_value if d.dim_value else 'dynamic' for d in inp.type.tensor_type.shape.dim]
 print(f' Shape: {shape}')
 
# 输出信息
print('=== Outputs ===')
for out in model.graph.output:
 print(f' Name: {out.name}')
 
# 算子列表
print('=== Nodes ===')
for node in model.graph.node:
 print(f' Op: {node.op_type}, Inputs: {list(node.input)}, Outputs: {list(node.output)}')
 
# 权重/参数
print(f'\\n=== Initializers ({len(model.graph.initializer)}) ===')
for init in model.graph.initializer:
 print(f' {init.name}: shape={list(init.dims)}')
 
# 验证模型合法性
onnx.checker.check_model(model)
print('\\n Model is valid!')
"

输出示例:

=== Inputs ===
 Name: input
 Shape: ['dynamic', 4]

=== Outputs ===
 Name: output

=== Nodes ===
 Op: Gemm, Inputs: ['input', 'linear.weight', 'linear.bias'], Outputs: ['output']

=== Initializers (2) ===
 linear.weight: shape=[3, 4]
 linear.bias: shape=[3]

 Model is valid!

注意:nn.Linear 在 PyTorch 中被转换为 ONNX 的 Gemm 算子(General Matrix Multiply:Y = αA×B + βC)。这是 ONNX 算子层面的”翻译”。

3.3 netron — 可视化 ONNX 模型

# 安装 netron
pip install netron
 
# 在浏览器中可视化
python -c "import netron; netron.start('simple_model.onnx')"
 
# 或者用命令行导出图片
# netron 是交互式工具,无法直接出图,但可以在浏览器中查看每个节点的结构

netron 让你像反汇编工具(objdump)一样检查模型内部结构。每个节点的输入/输出张量形状、权重值、算子属性都能看到。这对于调试 ONNX 导出问题非常有用。


第四节:用 onnxruntime 验证导出

4.1 Python 端验证

导出 ONNX 后最重要的一步是验证:Python 原始模型 和 ONNX Runtime 推理的结果必须一致。

import torch
import torch.nn as nn
import numpy as np
import onnxruntime as ort
 
# 1. 重新加载刚导出的模型
class SimpleModel(nn.Module):
 def __init__(self):
 super().__init__()
 self.linear = nn.Linear(4, 3)
 
model = SimpleModel()
model.eval()
 
# 2. 创建 ONNX Runtime session
session = ort.InferenceSession('simple_model.onnx')
 
# 3. 准备输入
x = torch.randn(2, 4) # batch=2
 
# 4. PyTorch 推理
with torch.no_grad():
 pytorch_out = model(x).numpy()
 
# 5. ONNX Runtime 推理
# 注意:ONNX Runtime 输入是 numpy 数组,键名必须匹配 input_names
onnx_out = session.run(
 None, # None = 获取所有输出
 {'input': x.numpy()}, # {'输入节点名': numpy 数组}
)[0] # run() 返回 list of outputs
 
# 6. 对比结果
diff = np.abs(pytorch_out - onnx_out).max()
print(f"PyTorch output: {pytorch_out[0]}")
print(f"ONNX RT output: {onnx_out[0]}")
print(f"Max difference: {diff:.10f}")
assert diff < 1e-5, f"Output mismatch! diff={diff}"
print(" PyTorch and ONNX Runtime outputs match!")

4.2 常见的 ONNX 导出问题

问题原因解决方案
输出不一致使用了不支持的 Python 控制流if/for 可能无法正确 trace
导出失败算子不支持(如自定义 CUDA kernel)torch.onnx.is_in_onnx_export() 提供 fallback
动态形状错误未设置 dynamic_axes为可变维度设置动态轴
eval() 忘调用Dropout/BatchNorm 行为不同始终在导出前调用 model.eval()

4.3 处理动态控制流

import torch
import torch.nn as nn
 
class DynamicModel(nn.Module):
 def __init__(self):
 super().__init__()
 self.linear_a = nn.Linear(10, 10)
 self.linear_b = nn.Linear(10, 10)
 
 def forward(self, x, use_branch_b=False):
 if use_branch_b:
 return self.linear_b(x)
 else:
 return self.linear_a(x)
 
model = DynamicModel().eval()
 
# 方案1:分别导出两个分支
# torch.onnx.export(model, (x, False), 'model_branch_a.onnx', ...)
# torch.onnx.export(model, (x, True), 'model_branch_b.onnx', ...)
 
# 方案2:用 torch.cond(PyTorch 2.0+)
def forward(self, x, use_branch_b):
 return torch.cond(use_branch_b, lambda: self.linear_b(x), lambda: self.linear_a(x))

PyTorch 的 ONNX 导出使用 Tracing(跟踪)模式:给定一个示例输入,实际执行一次前向传播,记录所有执行的操作,然后序列化为 ONNX。这类似于 C 代码的单路径执行记录——if/else 的分支、for 循环的迭代次数都被”固化”在导出的图中。


第五节:高级导出场景

5.1 导出带 BatchNorm 的模型

# 关键:先 eval(),再导出!
model.eval()
 
# 或者将 BatchNorm 折叠到前面的卷积中(减少推理计算)
torch.onnx.export(
 model, dummy_input, 'model.onnx',
 do_constant_folding=True, # 默认 True:常量折叠优化
 training=torch.onnx.TrainingMode.EVAL, # 确保是 EVAL 模式
)

5.2 导出 Transformer / BERT

from transformers import AutoModel, AutoTokenizer
 
model = AutoModel.from_pretrained('bert-base-uncased').eval()
tokenizer = AutoTokenizer.from_pretrained('bert-base-uncased')
 
# BERT 有两个输入: input_ids 和 attention_mask
text = "Hello, world!"
inputs = tokenizer(text, return_tensors='pt')
 
export_args = (inputs['input_ids'], inputs['attention_mask'])
 
torch.onnx.export(
 model,
 export_args,
 'bert.onnx',
 input_names=['input_ids', 'attention_mask'],
 output_names=['last_hidden', 'pooler_output'],
 dynamic_axes={
 'input_ids': {0: 'batch', 1: 'seq_length'},
 'attention_mask': {0: 'batch', 1: 'seq_length'},
 'last_hidden': {0: 'batch', 1: 'seq_length'},
 'pooler_output': {0: 'batch'},
 },
 opset_version=17,
)

5.3 简化 ONNX 模型(onnxsim)

pip install onnxsim
 
# 简化 ONNX 图:移除冗余节点、折叠常量
python -c "
import onnx
from onnxsim import simplify
model = onnx.load('bert.onnx')
simplified, check = simplify(model)
assert check, 'Simplification failed'
onnx.save(simplified, 'bert_simplified.onnx')
print('Simplified!')
"

onnxsim 对于大型模型非常实用,可以减少 10-30% 的模型大小和推理时间。



练习

以下题目用于验证本章所学内容:

题号题目链接涉及知识点
本章无对应力扣题请用动手练习题自检