ONNX 模型格式的跨框架转换
bash pip install onnx onnxruntime onnxruntime-gpu # GPU版本 pip install torch torchvision pip install tensorflow # 或 tensorflow-cpu pip install tf2onnx pip install onnx2torch
### 从 PyTorch 导出 ONNX
PyTorch 提供了内置的 `torch.onnx.export()` 函数,可将模型的计算图导出为 ONNX 格式。
#### 1. 定义或加载 PyTorch 模型
这里我们使用一个简单的卷积网络示例,你也可以替换为自己的模型。
```python
import torch
import torch.nn as nn
import torch.nn.functional as F
class SimpleCNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 6, 5)
self.pool = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(6, 16, 5)
self.fc1 = nn.Linear(16 * 5 * 5, 120)
self.fc2 = nn.Linear(120, 84)
self.fc3 = nn.Linear(84, 10)
def forward(self, x):
x = self.pool(F.relu(self.conv1(x)))
x = self.pool(F.relu(self.conv2(x)))
x = x.view(-1, 16 * 5 * 5)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
x = self.fc3(x)
return x
model = SimpleCNN()
model.eval() # 一定要切换到评估模式
2. 创建虚拟输入并导出
torch.onnx.export() 需要一次示例输入来追踪模型的计算图。
batch_size = 1
dummy_input = torch.randn(batch_size, 3, 32, 32)
torch.onnx.export(
model, # 模型对象
dummy_input, # 示例输入
"simple_cnn.onnx", # 输出文件名
export_params=True, # 存储训练好的参数
opset_version=11, # ONNX 算子集版本,通常选 11 或 13
do_constant_folding=True, # 常量折叠优化
input_names=['input'], # 输入节点名称
output_names=['output'], # 输出节点名称
dynamic_axes={ # 动态轴配置(可选)
'input': {0: 'batch_size'},
'output': {0: 'batch_size'}
}
)
opset_version根据你的 ONNX Runtime 版本选择,一般 11 兼容性最好。- 导出后可使用 Netron 可视化
simple_cnn.onnx,确认模型结构正确。
从 TensorFlow/Keras 导出 ONNX
对于 TensorFlow 2.x 及 Keras 模型,推荐使用 tf2onnx 工具。
1. 准备 TensorFlow 模型
import tensorflow as tf
# 创建一个简单的 Keras 模型
model = tf.keras.models.Sequential([
tf.keras.layers.Conv2D(6, 5, activation='relu', input_shape=(32,32,3)),
tf.keras.layers.MaxPooling2D(2),
tf.keras.layers.Conv2D(16, 5, activation='relu'),
tf.keras.layers.MaxPooling2D(2),
tf.keras.layers.Flatten(),
tf.keras.layers.Dense(120, activation='relu'),
tf.keras.layers.Dense(84, activation='relu'),
tf.keras.layers.Dense(10)
])
# 或加载已训练的模型
# model = tf.keras.models.load_model('my_model.h5')
2. 使用 tf2onnx 转换
import tf2onnx
spec = (tf.TensorSpec((None, 32, 32, 3), tf.float32, name="input"),)
output_path = "tf_cnn.onnx"
model_proto, _ = tf2onnx.convert.from_keras(model, input_signature=spec, opset=13, output_path=output_path)
input_signature使用tf.TensorSpec定义输入张量的形状和类型,允许动态 batch。- 生成的文件同样可用 Netron 查看。
如果已有 SavedModel 格式,可直接使用命令行:
python -m tf2onnx.convert --saved-model ./saved_model --output model.onnx --opset 13
使用 ONNX Runtime 进行高性能推理
ONNX Runtime 是一个跨平台的推理引擎,支持 CPU、GPU 及多种硬件加速器,可直接加载 ONNX 模型进行预测。
加载模型并推理
import onnxruntime as ort
import numpy as np
# 选择执行提供器,GPU 可用 'CUDAExecutionProvider'
providers = ['CPUExecutionProvider']
session = ort.InferenceSession('simple_cnn.onnx', providers=providers)
# 获取输入输出名称
input_name = session.get_inputs()[0].name
output_name = session.get_outputs()[0].name
# 构造输入数据(与导出时的 dummy_input 形状一致)
input_data = np.random.randn(1, 3, 32, 32).astype(np.float32)
# 推理
pred_onx = session.run([output_name], {input_name: input_data})[0]
print(pred_onx.shape)
性能对比小贴士
- 使用
onnxruntime的 Graph Optimization 功能可进一步提升速度,通过session_options.graph_optimization_level设置。 - 对于动态 batch 的模型,输入数据 batch 维度可以灵活变化。
将 ONNX 转换为 TensorFlow 格式
当需要将 ONNX 模型集成到 TensorFlow 生态(例如使用 TensorFlow Serving、TensorFlow Lite)时,可以反向转换。
使用 onnx-tf 进行转换
首先安装 onnx-tf:
pip install onnx-tf
然后执行:
import onnx
from onnx_tf.backend import prepare
# 加载 ONNX 模型
onnx_model = onnx.load("simple_cnn.onnx")
# 准备 TensorFlow 表示
tf_rep = prepare(onnx_model)
# 导出为 SavedModel
tf_rep.export_graph("tf_from_onnx")
现在,你可以像加载普通 TensorFlow 模型一样使用它:
import tensorflow as tf
model = tf.saved_model.load("tf_from_onnx")
input_tensor = tf.constant(np.random.randn(1, 3, 32, 32).astype(np.float32))
output = model(input_tensor)
print(output)
注意事项
- 并非所有 ONNX 算子都能完美映射到 TensorFlow,复杂算子可能需要手动处理。
- 对于推理场景,转换后的 TensorFlow 模型通常功能正常,但不保证支持梯度计算(即无法直接训练)。
将 ONNX 转换为 PyTorch 格式
将 ONNX 模型转为 PyTorch 模型可以使用 onnx2torch 库,它能生成可微调的 torch.nn.Module。
安装与转换
pip install onnx2torch
import onnx
import onnx2torch
# 加载 ONNX 模型
onnx_model = onnx.load("simple_cnn.onnx")
# 转换为 PyTorch 模块
pytorch_model = onnx2torch.convert(onnx_model)
# 现在可以当作普通 PyTorch 模型使用
pytorch_model.eval()
dummy_input = torch.randn(1, 3, 32, 32)
output = pytorch_model(dummy_input)
print(output.shape)
优点与限制
- 转换后的模型支持反向传播,可用于微调(finetune)。
- 某些复杂的 ONNX 图(如循环、条件分支)可能无法完全转换,此时需要手动调整。
- 检查转换后的模型参数名称,可能与原始 PyTorch 模型不同。
常见问题与调试
1. 导出 ONNX 时版本不匹配
确保 opset_version 与目标推理环境兼容。例如,一些边缘设备可能只支持 opset 10。通常选择 11 或 12 能获得广泛支持。
2. 动态形状支持
导出时使用 dynamic_axes 可让 ONNX 模型接受可变 batch 尺寸或可变序列长度。转换回框架时,务必检查是否保留了动态维度。
3. 算子缺失或不支持
每个框架都有自己特有的算子,导出 ONNX 时可能遇到不支持的算子。常见解决方法:
- 降级 opset 版本。
- 使用
torch.onnx.export的custom_opsets参数。 - 在 TensorFlow 中使用
tf2onnx的--custom-ops添加自定义实现。
4. 模型精度验证
转换后可通过对比原框架与 ONNX Runtime 的输出误差来判断转换是否正确。一般允许 1e-5 以内的误差(对于 float32)。
# PyTorch 与 ONNX 结果比较
torch_output = model(dummy_input).detach().numpy()
onnx_output = session.run([output_name], {input_name: dummy_input.numpy()})[0]
assert np.allclose(torch_output, onnx_output, atol=1e-5)