PyTorch模型部署一、模型转换为onnx
PyTorch模型转换为ONNX格式,从基础到实践一步步教学。

一、准备工作

  1. 安装必要的库

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/

下一步建议

  1. 先用一个简单的模型测试整个流程
  2. 验证推理结果的一致性
  3. 针对你的具体模型调整动态维度
  4. 使用ONNXRuntime测试部署效果
Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐