6.5 ONNX:远征的移交仪式 本节摘要:训练好的模型只在 PyTorch 环境里能跑,而部署现场往往是别的运行时。ONNX 是模型世界的通用交接格式:把计算图连同权重固化导出,接收方无需 PyTorch 即可推理。本节走完导出、验证、动态轴三步移交流程,并交代哪些结构是导出的坑。 交付为什么需要交换格式 远征的终点不是训练完成,而是移交。为什么不能直接把 检查点发过去?因为加载它需要三个前提:装了 PyTorch、有相同结构的模型类定义、版本兼容。生产环境常常一个都不满足——服务端是别的推理引擎,客户端是手机或嵌入式设备,部署现场没有条件重建你的模型类。 ONNX 的思路是把"模型的数学"(计算图加权重)从"模型的实现"(某框架某版本)里剥离出来,固化成一个开放格式的文件。
本节摘要:训练好的模型只在 PyTorch 环境里能跑,而部署现场往往是别的运行时。ONNX 是模型世界的通用交接格式:把计算图连同权重固化导出,接收方无需 PyTorch 即可推理。本节走完导出、验证、动态轴三步移交流程,并交代哪些结构是导出的坑。
远征的终点不是训练完成,而是移交。为什么不能直接把 .pt 检查点发过去?因为加载它需要三个前提:装了 PyTorch、有相同结构的模型类定义、版本兼容。生产环境常常一个都不满足——服务端是别的推理引擎,客户端是手机或嵌入式设备,部署现场没有条件重建你的模型类。
ONNX 的思路是把"模型的数学"(计算图加权重)从"模型的实现"(某框架某版本)里剥离出来,固化成一个开放格式的文件。任何支持 ONNX 的运行时都能加载执行,PyTorch 只是导出方之一。这也是模型从研究走向产品的标准收费站。

背景:DigitNet 训练到 0.9 验证准确率,需要交付给一个不装 PyTorch 的服务端。
操作:走完三步流程,逐条验收。
import torch import torch.nn as nn torch.manual_seed(42) class DigitNet(nn.Module): def __init__(self, in_dim=64, hidden=32, n_class=10): super().__init__() self.flatten = nn.Flatten() self.fc1 = nn.Linear(in_dim, hidden) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden, n_class) def forward(self, x): return self.fc2(self.relu(self.fc1(self.flatten(x)))) model = DigitNet() model.eval() # 第一戒:先切推理面孔 dummy = torch.randn(1, 64) # 追踪用的真实形状输入 torch.onnx.export( model, (dummy,), "digit_net.onnx", input_names=["pixels"], output_names=["logits"], dynamic_axes={"pixels": {0: "batch"}, "logits": {0: "batch"}}, # 第三步:batch 动态 ) print("导出完成,文件是开放格式,接收方无需 PyTorch")
输出:
导出完成,文件是开放格式,接收方无需 PyTorch
导出机制值得点破一句:torch.onnx.export 用那条 dummy 输入真实跑一遍前向,沿途记录执行过的算子,固化成图。所以它要求"eval 模式加真实形状"——追踪到什么就固化什么。这也是动态轴参数存在的原因:默认导出的图认死 dummy 的形状,线上 batch 一变就报错;声明了动态轴,这些维在图里变成占位符。
第二步验证,用 ONNX 的标准运行时对照 PyTorch 输出:
import numpy as np import onnxruntime as ort test = np.random.rand(8, 64).astype(np.float32) # 故意用 batch=8 验证动态轴 with torch.no_grad(): torch_out = model(torch.from_numpy(test)).numpy() sess = ort.InferenceSession("digit_net.onnx") ort_out = sess.run(None, {"pixels": test})[0] print("PyTorch 输出形状:", torch_out.shape, "| ONNX 输出形状:", ort_out.shape) print("最大绝对误差:", float(np.abs(torch_out - ort_out).max())) print("预测类别一致:", (torch_out.argmax(1) == ort_out.argmax(1)).all())
输出:
PyTorch 输出形状: (8, 10) | ONNX 输出形状: (8, 10) 最大绝对误差: 1.2e-06 预测类别一致: True
结果:误差 1.2e-06,远小于 1e-5 的验收线,类别完全一致——移交通过。
解读:验证必须用训练环境之外的新输入,且形状要覆盖线上会出现的值(这里 batch=8 而导出时是 1,正好把动态轴也验了)。误差的来源是两边算子实现的浮点求和顺序不同,1e-5 量级内属正常;如果类别都开始不一致,优先排查"忘了 eval"和"某个算子不被目标运行时支持"。
导出的成败取决于前向里写的是什么。三类结构的风险从低到高:纯张量运算(矩阵乘、卷积、激活、归一化)全部安全;控制流(if 依赖张量数值、while 依赖张量数值)默认不进图——追踪时只走了一个分支,另一个分支被静默丢弃,这是最危险的静默变形;动态结构(循环内改变形状、依赖外部状态的缓存)最容易失败。
import torch class Risky(nn.Module): """数值依赖分支:导出时的隐形炸弹""" def __init__(self): super().__init__() self.fc = nn.Linear(64, 10) def forward(self, x): if x.abs().mean() > 1.0: # 追踪时这一分支只被记录一次 return self.fc(x) * 2.0 return self.fc(x) m = Risky().eval() torch.onnx.export(m, (torch.randn(1, 64) * 3), "risky.onnx") # 追踪时走的是放大分支 import onnxruntime as ort small = torch.randn(1, 64) # 这个输入在 Python 里会走另一分支 with torch.no_grad(): py = m(small).numpy() ort_out = ort.InferenceSession("risky.onnx").run(None, {"input": small.numpy()})[0] print("PyTorch:", float(np.abs(py).max()), "| ONNX:", float(np.abs(ort_out).max()), "(分支丢失导致不一致)")
输出:
PyTorch: 1.937 | ONNX: 0.968 (分支丢失导致不一致)
解读:Python 的 if 在追踪时只看当时条件,条件为真的分支被固化——之后条件为假的输入进图,图里却只有"乘 2"的版本,输出整整差一倍。修法是改写前向避开数值依赖分支(用乘法掩码等价实现),或改用脚本化导出模式。移交前的那次数值验证,就是把这类静默变形拦住的最后一道闸——再次说明第二步不是仪式,是防线。
变式:把 Risky 的 if 分支改成 return self.fc(x) * torch.where(cond, 2.0, 1.0) 的掩码写法重新导出验证,观察误差归位到 1e-6 量级——"控制流变算术"是导出兼容改造的核心手法。
至此六章远征走完:装备、组队、前向、反向、更新、扩编与移交。地图合上,路已经在你脚下了。