PyTorch 最小模型转 ONNX 完整样例

ONNX简介

ONNX(Open Neural Network Exchange)即开放神经网络交换格式,是一种开源通用的深度学习模型标准格式。它统一定义了神经网络算子、计算图与数据存储规范,打破PyTorch、TensorFlow、Paddle等不同训练框架之间的模型壁垒,实现模型跨框架自由迁移。

开发者可将不同框架训练完成的模型导出为ONNX文件,再借助ONNX Runtime、TensorRT、NCNN等推理引擎完成端侧、服务器、嵌入式设备的高效部署,同时支持模型简化、算子优化与量化压缩,极大简化深度学习模型上线流程,是工业界AI模型部署最主流的中间格式。

1. 环境依赖

bash 复制代码
pip install torch onnx onnxsim onnxruntime onnxscript

2. 极简模型 + 导出 ONNX 代码

python 复制代码
import torch
import torch.nn as nn

# 1. 定义超简单单层网络
class MiniModel(nn.Module):
    def __init__(self):
        super().__init__()
        # 输入10维,输出2维
        self.fc = nn.Linear(10, 2)

    def forward(self, x):
        return self.fc(x)

# 2. 初始化模型
model = MiniModel()
model.eval()  # 推理模式

# 3. 构造虚拟输入 (batch=1, 输入维度10)
dummy_input = torch.randn(1, 10)

# 4. 导出 ONNX
torch.onnx.export(
    model,
    dummy_input,
    "mini_model.onnx",       # 导出文件名
    input_names=["input"],    # 输入节点名
    output_names=["output"],  # 输出节点名
    opset_version=17,         # ONNX算子版本
    do_constant_folding=True  # 常量折叠优化
)
print("✅ ONNX 导出完成:mini_model.onnx")

3. 验证 ONNX 是否合法

python 复制代码
import onnx

# 加载校验
onnx_model = onnx.load("mini_model.onnx")
onnx.checker.check_model(onnx_model)
print("✅ ONNX 模型格式合法无错误")

# 打印模型结构
print(onnx.helper.printable_graph(onnx_model.graph))

4. 简化 ONNX(去除冗余节点)

python 复制代码
import onnxsim

model_simplified, ok = onnxsim.simplify("mini_model.onnx")
assert ok, "模型简化失败"
onnx.save(model_simplified, "mini_model_simplified.onnx")
print("✅ ONNX 简化完成")

5. ONNX Runtime 推理测试

python 复制代码
import onnxruntime as ort
import numpy as np

# 加载模型
session = ort.InferenceSession("mini_model_simplified.onnx")
input_name = session.get_inputs()[0].name
output_name = session.get_outputs()[0].name

# 构造输入
inp = np.random.randn(1, 10).astype(np.float32)
res = session.run([output_name], {input_name: inp})
print("推理结果:", res[0])

关键参数说明

  1. opset_version:越高支持算子越多,部署兼容性优先选 13~17
  2. do_constant_folding:自动合并常量,减小模型体积
  3. 输入形状:(batch_size, feature_dim) 按需修改

扩展:带ReLU激活的常用小模型

python 复制代码
class MiniModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(10, 32),
            nn.ReLU(),
            nn.Linear(32, 2)
        )
    def forward(self, x):
        return self.net(x)

输出为

复制代码
✅ ONNX 模型格式合法无错误
/workspace/onnx/t2.py:9: DeprecationWarning: Deprecated since 1.19. Consider using onnx.printer.to_text() instead.
  print(onnx.helper.printable_graph(onnx_model.graph))
graph main_graph (
  %input[FLOAT, 1x10]
) initializers (
  %net.0.weight[FLOAT, 32x10]
  %net.0.bias[FLOAT, 32]
  %net.2.weight[FLOAT, 2x32]
  %net.2.bias[FLOAT, 2]
) {
  %linear = Gemm[alpha = 1, beta = 1, transA = 0, transB = 1](%input, %net.0.weight, %net.0.bias)
  %relu = Relu(%linear)
  %output = Gemm[alpha = 1, beta = 1, transA = 0, transB = 1](%relu, %net.2.weight, %net.2.bias)
  return %output
}

详解

相关推荐
彩讯股份3006345 小时前
彩讯股份与心洲科技签署战略合作协议,共建企业级模型后训练能力
人工智能·科技
Scott9999HH5 小时前
【IIoT流量实战】蒸汽管道阀门全关却仍有流量?用 Python 实现涡街信号 FFT 频谱分析与温压全补偿积算网关,深度拆解靠谱的涡街流量计厂家硬核技术标准
开发语言·python
迅易科技5 小时前
从场景验证到Agent上线:迅易 × WorkBuddy如何帮助企业建设AI能力?
人工智能·ai·腾讯云
PNP Robotics5 小时前
多伦多大学机器人峰会|物理AI与具身智能落地新趋势
人工智能·深度学习·机器学习·机器人
GIR1236 小时前
官方出品 | 多通道土壤呼吸测量系统市场现状与十五五规划深度报告:行业分析+趋势预测全收录
大数据·人工智能·机器学习
绿算技术6 小时前
绿算技术亮相第十八届HPC AI中国年会,擘画AI基础设施全栈协同新图景
人工智能
Litluecat6 小时前
2026年7月22日科技热点新闻
人工智能·科技·新闻·每日·速览
To_OC6 小时前
别再傻傻分不清:Workflow 和 Agent 到底不是一回事
人工智能·agent·workflow
AI云海6 小时前
python 列表、元组、集合和字典
开发语言·python
触底反弹6 小时前
🔥 2026 大模型选择指南:别再只看 Benchmark 了,这些维度才是关键!
人工智能·面试