模型导出 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 | .pt | PyTorch 原生、完整支持 | 仅限 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_version | ONNX 算子集版本 | 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_inputC++ 部署关键:如果不设置
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% 的模型大小和推理时间。
练习
以下题目用于验证本章所学内容:
| 题号 | 题目 | 链接 | 涉及知识点 |
|---|---|---|---|
| — | 本章无对应力扣题 | — | 请用动手练习题自检 |