ONNX 模型格式的跨框架转换

FreeGuideOnline 15阅读 2026-07-10

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.exportcustom_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)