PyTorch模型部署一、模型转换为onnx
·
PyTorch模型部署一、模型转换为onnx
PyTorch模型转换为ONNX格式,从基础到实践一步步教学。
一、准备工作
- 安装必要的库
pip install onnx
pip install onnxruntime
二、基础转换示例
示例1:最简单的模型转换
import torch
import torch.nn as nn
import onnx
import onnxruntime as ort
import numpy as np
# 1. 定义一个简单的模型
class SimpleModel(nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.linear = nn.Linear(10, 5)
self.relu = nn.ReLU()
def forward(self, x):
return self.relu(self.linear(x))
# 2. 实例化并设置为评估模式
model = SimpleModel()
model.eval()
# 3. 创建示例输入
dummy_input = torch.randn(1, 10) # batch_size=1, input_size=10
# 4. 转换为ONNX
torch.onnx.export(
model, # 要转换的模型
dummy_input, # 示例输入
"simple_model.onnx", # 输出文件名
export_params=True, # 导出模型参数
#opset_version=13, # ONNX算子集版本
opset_version=18,
# 问题一: 你设置了 opset_version=13(旧版本)
# PyTorch内部实际用了更新版本的算子(版本18)
# ONNX尝试从版本18降到版本13,但Relu算子不支持这个降级
# 这是版本兼容性问题
do_constant_folding=True, # 是否进行常量折叠优化
input_names=['input'], # 输入节点名称
output_names=['output'], # 输出节点名称
# dynamic_axes={ # 动态维度设置(可选)
## 问题二:旧api,'dynamic_axes' is not recommended when dynamo=TruePyTorch >= 2.0Dynamo引擎dynamic_shapes更强大,支持更多控制
# 'input': {0: 'batch_size'},
# 'output': {0: 'batch_size'}
# }
dynamic_shapes={ # 使用新的API
'x': {0: torch.export.Dim("batch_size")},# 修正:dynamic_shapes 的key必须是 forward 的参数名 'x'
}
# 问题一:确实是opset版本不匹配,新PyTorch要用opset 18+
# 问题二:确实是API变化,新PyTorch要用dynamic_shapes
, verbose=True#verbose=True 会打印出报错前最后处理的层或算子信息
)
print("模型已保存为 simple_model.onnx")
三、实际项目中的完整流程
步骤1:加载并准备你的PyTorch模型
## 假设你有一个训练好的模型
def load_your_model():
# 这里替换为你的模型加载代码
model = YourModelClass() # 你的模型类
model.load_state_dict(torch.load('your_model.pth'))
model.eval() # 重要:必须设置为评估模式
return model
model = load_your_model()
步骤2:创建合适的示例输入
# 根据你的模型输入要求创建dummy input
# 示例1:图像分类模型
batch_size = 1
channels = 3
height = 224
width = 224
dummy_input = torch.randn(batch_size, channels, height, width)
# 示例2:NLP模型
seq_length = 50
batch_size = 1
dummy_input = torch.randint(0, 10000, (batch_size, seq_length))
步骤3:详细配置导出参数
#详细配置的导出函数
def export_to_onnx(model, dummy_input, output_path="model.onnx"):
"""
将PyTorch模型导出为ONNX格式
参数:
model: PyTorch模型
dummy_input: 示例输入
output_path: 输出文件路径
"""
torch.onnx.export(
model=model,
args=dummy_input,
f=output_path,
export_params=True,
opset_version=13, # 常用版本:11, 12, 13, 14
do_constant_folding=True,
input_names=['input'],
output_names=['output'],
dynamic_axes={
'input': {0: 'batch_size'}, # 第0维是batch,设为动态
'output': {0: 'batch_size'}
},
verbose=False
)
print(f"模型已导出到: {output_path}")
# 执行导出
export_to_onnx(model, dummy_input, "my_model.onnx")
四、验证转换结果
验证1:检查ONNX模型
def check_onnx_model(model_path):
"""检查导出的ONNX模型"""
# 加载ONNX模型
onnx_model = onnx.load(model_path)
# 检查模型结构
print(f"模型IR版本: {onnx_model.ir_version}")
print(f"生产者名称: {onnx_model.producer_name}")
print(f"生产者版本: {onnx_model.producer_version}")
# 验证模型格式
try:
onnx.checker.check_model(onnx_model)
print("✓ ONNX模型格式正确")
except onnx.checker.ValidationError as e:
print(f"✗ 模型格式错误: {e}")
return False
return True
check_onnx_model("my_model.onnx")
验证2:推理结果对比
def verify_inference(pytorch_model, onnx_path, dummy_input):
"""对比PyTorch和ONNX的推理结果"""
# PyTorch推理
with torch.no_grad():
pytorch_output = pytorch_model(dummy_input).numpy()
# ONNX推理
ort_session = ort.InferenceSession(onnx_path)
ort_inputs = {ort_session.get_inputs()[0].name: dummy_input.numpy()}
ort_output = ort_session.run(None, ort_inputs)[0]
# 对比结果
print(f"PyTorch输出形状: {pytorch_output.shape}")
print(f"ONNX输出形状: {ort_output.shape}")
# 计算差异
diff = np.abs(pytorch_output - ort_output).max()
print(f"最大绝对误差: {diff}")
if diff < 1e-5:
print("✓ 推理结果一致")
else:
print("⚠ 推理结果存在差异")
return diff
verify_inference(model, "my_model.onnx", dummy_input)
五、常见问题解决方案
问题1:动态维度处理
如果你的模型需要支持不同的输入形状:
# 支持动态batch和序列长度
dynamic_axes={
'input': {
0: 'batch_size',
1: 'sequence_length' # 对于NLP模型
},
'output': {
0: 'batch_size',
1: 'sequence_length'
}
}
问题2:多输入模型
# 假设模型有多个输入
class MultiInputModel(nn.Module):
def forward(self, x1, x2):
return x1 + x2
model = MultiInputModel()
dummy_input1 = torch.randn(1, 10)
dummy_input2 = torch.randn(1, 10)
torch.onnx.export(
model,
(dummy_input1, dummy_input2), # 多个输入作为元组
"multi_input.onnx",
input_names=['input1', 'input2'],
output_names=['output']
)
问题3:模型简化
from onnxsim import simplify
def simplify_onnx_model(input_path, output_path):
"""简化ONNX模型"""
# 加载原始模型
model = onnx.load(input_path)
# 简化模型
model_simp, check = simplify(model)
if check:
onnx.save(model_simp, output_path)
print(f"简化模型已保存: {output_path}")
return True
else:
print("简化失败")
return False
# 使用
simplify_onnx_model("my_model.onnx", "my_model_simplified.onnx")
六、完整工作流程示例
import torch
import torch.nn as nn
import onnx
import onnxruntime as ort
import numpy as np
def load_your_model():
# 这里替换为你的模型加载代码
model = YourModelClass() # 你的模型类
model.load_state_dict(torch.load('your_model.pth'))
model.eval() # 重要:必须设置为评估模式
return model
class YourCNNModel(nn.Module):
"""示例CNN模型"""
def __init__(self):
super(YourCNNModel, self).__init__()
self.conv1 = nn.Conv2d(3, 16, 3, padding=1)
self.pool = nn.MaxPool2d(2)
self.conv2 = nn.Conv2d(16, 32, 3, padding=1)
self.fc = nn.Linear(32 * 56 * 56, 10) # 假设输入224x224
def forward(self, x):
x = torch.relu(self.conv1(x))
x = self.pool(x) # 112x112
x = torch.relu(self.conv2(x))
x = self.pool(x) # 56x56
x = x.view(x.size(0), -1)
return self.fc(x)
def main():
# 1. 加载模型
model = load_your_model()
model.eval()
# 2. 创建示例输入
dummy_input = torch.randn(1, 3, 224, 224)
# 3. 导出ONNX
print("正在导出ONNX模型...")
torch.onnx.export(
model=model,
args=dummy_input,
f="final_model.onnx",
export_params=True,
opset_version=13,
do_constant_folding=True,
input_names=['input'],
output_names=['output'],
dynamic_axes={'input': {0: 'batch_size'},
'output': {0: 'batch_size'}},
verbose=False
)
# 4. 验证模型
print("\n验证模型...")
onnx_model = onnx.load("final_model.onnx")
onnx.checker.check_model(onnx_model)
# 5. 对比推理结果
print("\n对比推理结果...")
with torch.no_grad():
torch_out = model(dummy_input).numpy()
ort_session = ort.InferenceSession("final_model.onnx")
ort_inputs = {ort_session.get_inputs()[0].name: dummy_input.numpy()}
ort_out = ort_session.run(None, ort_inputs)[0]
diff = np.abs(torch_out - ort_out).max()
print(f"PyTorch输出: {torch_out.shape}, 值范围: [{torch_out.min():.4f}, {torch_out.max():.4f}]")
print(f"ONNX输出: {ort_out.shape}, 值范围: [{ort_out.min():.4f}, {ort_out.max():.4f}]")
print(f"最大绝对误差: {diff:.6f}")
if diff < 1e-5:
print("✓ 转换成功!模型可以正常使用。")
else:
print("⚠ 有微小误差,但可能在可接受范围内")
if __name__ == "__main__":
main()
七、最佳实践建议
1. 版本控制:
• ONNXRuntime最新版本
2. 模型优化:
安装优化工具
pip install onnxruntime-tools
pip install onnxoptimizer
3. 常见算子支持:
• 大部分PyTorch算子都支持
• 注意:自定义算子需要单独实现
• 某些动态控制流可能不支持
4. 调试技巧:
导出时添加verbose查看详细信息
torch.onnx.export(…, verbose=True)
使用Netron可视化ONNX模型
访问:https://netron.app/
下一步建议
- 先用一个简单的模型测试整个流程
- 验证推理结果的一致性
- 针对你的具体模型调整动态维度
- 使用ONNXRuntime测试部署效果
更多推荐
所有评论(0)