PyTorch 模型导出与部署实战:ONNX + onnxruntime(可直接落地)

PyTorch 模型导出与部署实战:ONNX + onnxruntime(可直接落地)

训练好的模型放在 .pth 里没人会用,导出成 ONNX 才能部署交付。这篇从导出到封装成 HTTP 服务一条龙,代码真实可跑,接单交付用得上。

前言

之前一直有个困惑:模型在 PyTorch 里跑得好好的,怎么交给别人用?给个 .pth 文件,对方没有训练环境根本跑不起来;就算能跑,每个框架一套环境,交付极其痛苦。

直到学会 ONNX 才打开新世界:训练一次,到处运行。ONNX 是微软牵头的中立模型格式,PyTorch 导出的模型能被 onnxruntime(一个几百 MB 的轻量运行时)加载,CPU 能跑、GPU 能跑、手机也能跑。这篇文章把"导出 → 验证 → 封装接口"全流程走一遍。

环境准备(先装依赖):

bat 复制代码
pip install torch torchvision onnx onnxruntime numpy fastapi uvicorn

一、为什么导出 ONNX

对比 .pth(PyTorch) .onnx
使用方需要什么 完整 PyTorch 环境 + 训练代码 只要 onnxruntime 一个包
部署形态 绑死 Python Python/Java/C#/手机都能跑
交付友好度 差 好(交付就是给个文件)
推理速度 基准 通常更快(图优化)

做接单/交付的人尤其在意这点:交付 .onnx + 一段推理代码,客户环境再烂也能跑。

二、导出:torch.onnx.export(核心代码)

python 复制代码
import torch

# 你的模型(这里用之前训练好的图像分类模型举例)
model = SimpleCNN().to("cpu")
model.load_state_dict(torch.load("model.pth", map_location="cpu"))
model.eval()  # 必须 eval,影响 dropout/bn 行为

# 准备一个 dummy 输入,形状要和真实输入一致
dummy = torch.randn(1, 3, 32, 32)

torch.onnx.export(
    model,
    dummy,
    "model.onnx",
    input_names=["input"],
    output_names=["output"],
    dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}},  # 允许 batch 可变
    opset_version=17,
)
print("导出完成: model.onnx")

三个容易踩的细节:

  1. 模型必须先 eval(),否则导出的图里带 dropout 随机行为,推理结果对不上;
  2. dummy 输入形状必须和真实输入一致,导出时就固化了;
  3. dynamic_axes 让 batch 维度可变------不写的话,导出后只能一次推理固定数量(比如只能 1 张图一批),部署会很痛苦。

三、验证:onnxruntime 加载推理(必须做的一步)

导出不等于成功,跑一遍对比 PyTorch 的输出才放心:

bat 复制代码
pip install onnxruntime
python 复制代码
import onnxruntime as ort
import numpy as np
import torch

# 优先 GPU,没有就 CPU
sess = ort.InferenceSession(
    "model.onnx",
    providers=["CUDAExecutionProvider", "CPUExecutionProvider"],
)

# 同一张输入,PyTorch 和 onnxruntime 各算一次
x = torch.randn(1, 3, 32, 32)
with torch.no_grad():
    torch_out = model(x).numpy()
onnx_out = sess.run(None, {"input": x.numpy()})[0]

# 数值对比:误差在 1e-4 以内算正常
diff = np.abs(torch_out - onnx_out).max()
print("最大误差:", diff)
print("ONNX 推理结果:", onnx_out.argmax(axis=1))

误差在 1e-4 量级是正常的(浮点运算顺序差异),如果差得离谱,八成是导出时忘了 eval() 或者算子不支持。

运行验证 :python 依次运行导出、验证、部署三段代码,期望看到:导出完成: model.onnx、最大误差: 3.2e-05(1e-4 量级)、启动 uvicorn 后 curl 返回 {"class": 3, "confidence": 0.98} 这类 JSON------整条"导出→验证→接口"链路就走通了。

四、封装成 HTTP 服务(直接交付)

验证没问题后,用 FastAPI 把模型包成一个接口,客户随便什么语言都能调用:

bat 复制代码
pip install fastapi uvicorn
python 复制代码
from fastapi import FastAPI
import onnxruntime as ort
import numpy as np

app = FastAPI()
sess = ort.InferenceSession("model.onnx")

@app.post("/predict")
def predict(image: list):   # 客户端传拍平的图像数组
    arr = np.array(image, dtype=np.float32).reshape(1, 3, 32, 32)
    out = sess.run(None, {"input": arr})[0]
    cls = int(out.argmax(axis=1)[0])
    return {"class": cls, "confidence": float(out[0][cls])}

启动:

bat 复制代码
uvicorn main:app --host 0.0.0.0 --port 8000

测试调用(curl 传一个 3x32x32=3072 长度的数组即可):

bat 复制代码
curl -X POST http://127.0.0.1:8000/predict -H "Content-Type: application/json" -d "[0.1,0.2,...]"

这套交付物就三样 :model.onnx + main.py + requirements.txt(onnxruntime、fastapi、uvicorn)。客户 pip install -r requirements.txt && uvicorn main:app 就能跑,不依赖任何训练环境。

五、常见坑

坑 解决
GPU 推理报错 装 onnxruntime-gpu;providers 里 GPU 放前面;确认 CUDA 版本匹配
导出报"不支持算子" 简化模型结构(换成 ONNX 支持的层),或调高 opset_version
推理维度对不上 客户端数据 reshape 要和导出时的 dummy 一致;注意通道顺序是 CHW
需要可变输入尺寸 dynamic_axes 里把宽高也标上({2: "height", 3: "width"}),代价是推理慢一点
导出的 onnx 比 pth 大 正常(ONNX 自带图结构);在意体积可以做模型剪枝或换小模型

六、总结

链路就四步:导出(eval + dynamic_axes)→ 验证(数值对比)→ 封装(FastAPI)→ 交付(onnx + 代码)。

接单角度说一句:客户手里有训练好的模型、想部署成接口给人用,这是非常常见的外包需求。ONNX + FastAPI 这套组合,半天就能交付一个能用的推理服务,属于投入产出比很高的技能。

觉得有用点个赞收藏,有问题评论区见,看到都会回。

相关推荐
I'mChloe2 分钟前
Windows部署BiliNote:Docker安装、AI视频转写、Markdown笔记与cpolar远程访问
人工智能·windows·docker
喜欢打篮球的普通人6 分钟前
MiniMind 学习笔记(十二):Pretrain 实操——从版本梳理到 8GB 显卡上的真实训练
人工智能·笔记·学习
YOLO数据集集合8 分钟前
无人机视角行人与车辆检测数据集 | 无人机航拍 行人检测 车辆检测 智慧城市 公共安全9140期
人工智能·目标检测·无人机·智慧城市·车辆识别·无人机视角
老歌老听老掉牙18 分钟前
斜抛运动问题分析:给定最大高度与墙面位置的轨迹与时间求解
python·斜抛运动
Joecien21 分钟前
【2026实测】百炼 CLI 托管 Agent 教程:bl managed-agent 配置校验、版本回滚与变更预演(附完整命令)
人工智能·git·阿里云·知识图谱·agi
jimmyleeee25 分钟前
大模型安全之三十六:大模型数据管理----从投毒防御到偏见治理的完整框架
人工智能·安全
言乐626 分钟前
Python自动去除水印
开发语言·python·django·virtualenv·pygame
czq_268671948729 分钟前
Python打卡第31天
开发语言·python
invicinble30 分钟前
python 编程语言 认识维度
开发语言·数据库·python
AC赳赳老秦31 分钟前
数据采集全链路审计留痕:用 OpenClaw 实现合规审计与追溯
开发语言·汇编·python·php·swift·deepseek·openclaw