用 ONNX Runtime 把 PyTorch 模型变成跨平台极速推理引擎:导出、优化、量化完整实战
训练完一个 PyTorch 模型,下一步往往是 deployment 噩梦:生产环境没有 Python、没有 CUDA,甚至要跑在 Windows 工控机或浏览器里。直接
pip install torch把模型塞进服务端的做法,既重又慢,还绑死了技术栈。
一、为什么不能直接把 PyTorch 模型扔上生产
PyTorch 模型本质是「Python 代码 + 权重」,线上推理至少有三道坎:第一,必须带整套运行时(torch + 依赖动辄 1GB+);第二,默认 eager 模式逐算子走 Python 解释器,小 batch 下调度开销比计算还大;第三,跨语言(C++/C#/JS)调用极其别扭。ONNX(Open Neural Network Exchange)把模型算成一张与框架、语言无关的计算图,ONNX Runtime(ORT)再负责在这张图上做算子融合、常量折叠和执行------一次导出、到处快跑。
二、三步导出 ONNX:动态轴是关键
导出时最容易被忽略的是 dynamic_axes:不写死 batch 维度,否则线上只能固定条数推理。下面用一个最小文本分类器演示:
python
import torch, torch.nn as nn
class TextClf(nn.Module):
def __init__(self, dim=768, n_cls=3):
super().__init__()
self.net = nn.Sequential(
nn.Linear(dim, 256), nn.ReLU(),
nn.Linear(256, n_cls))
def forward(self, x):
return self.net(x)
model = TextClf().eval()
dummy = torch.randn(1, 768) # batch=1 的占位输入
torch.onnx.export(
model, dummy, "text_clf.onnx",
input_names=["input"], output_names=["logits"],
dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}},
opset_version=17,
)
print("导出完成:text_clf.onnx")
踩坑:opset 版本要和 ORT 版本匹配,过新会报
Unsupported operator;若模型含控制流(if/loop),需确认 torch 能 trace 展开。
三、用 ONNX Runtime 跑起来:一行 InferenceSession
推理端彻底摆脱 PyTorch,纯 numpy 即可,也能在 C++/C# 里用同一份 .onnx:
python
import onnxruntime as ort
import numpy as np
sess = ort.InferenceSession("text_clf.onnx",
providers=["CPUExecutionProvider"])
x = np.random.randn(32, 768).astype("float32") # 一次推 32 条
logits = sess.run(["logits"], {"input": x})[0]
print(logits.shape) # (32, 3)
providers 是性能开关:有 GPU 就填 ["CUDAExecutionProvider", "CPUExecutionProvider"],ORT 会自动回退。sess.run 的输入输出名必须和导出时的 input_names/output_names 对上。
四、再榨一层:图优化与动态量化
ORT 默认就做算子融合,但还能手动量化把 FP32 权重压成 INT8,体积和延迟一起降。
python
from onnxruntime.quantization import quantize_dynamic, QuantType
quantize_dynamic("text_clf.onnx", "text_clf.quant.onnx",
weight_type=QuantType.QInt8)
量化后权重文件通常缩到 1/4,CPU 上小模型推理常提速 1.5--3 倍。注意:动态量化只压权重不压激活,对矩阵乘占比高的模型收益最大;若精度掉得狠,改用配合校准集的静态量化。
五、上线前四道避坑清单
- 动态轴必须开:否则线上 batch 一变就报错或爆显存。
- Provider 顺序有讲究:GPU 机器把 CUDA 放前面,没卡会自动回 CPU,别写死。
- 量化先测精度:用一小撮真实样本比对量化前后输出分布,KL 散度超阈值就回退。
- 版本锁死 :
.onnx和 ORT 版本一起进仓库,避免换环境后算子不兼容。
总结
ONNX + ONNX Runtime 是把「研究模型」变成「生产组件」的最短路径:导出时盯住动态轴和 opset,推理端一行 InferenceSession 跨平台通用,量化再补一刀延迟。它不解决训练,但能让你的模型真正跑在用户机器上------轻、快、不挑语言。