返回专辑
·Johan·5 分钟阅读

onnxruntime assert_allclose 验收

PyTorch 能跑不代表 ONNX 能导出——dynamic axes、unsupported op 和 train/eval 差异要在 target runtime 上数值对齐验证。

onnxruntime assert_allclose 验收

1. PyTorch val 99%,ORT 推理全错类

分类模型 PyTorch eval() val acc 99%,导 ONNX 用 onnxruntime 跑同一 batch,top-1 掉到随机——常见是 export 时忘了 model.eval()(Dropout/BN 统计错),或 dynamic control flow trace 走了 training 分支。训练能跑 ≠ 部署图正确;必须在 target runtime 上做数值对齐。部署事故复盘里「模型文件有了」但无 ORT 对齐记录,几乎无法区分 export bug 与预处理 bug。

2. 导出契约

python
model.eval()
dummy = torch.randn(1, 3, 224, 224, device=device)
torch.onnx.export(
    model, dummy, "model.onnx",
    input_names=["input"], output_names=["logits"],
    dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}},
    opset_version=17,
    do_constant_folding=True,
)

dynamic_axes 声明 batch 维;固定 batch 部署可省略。opset 17+ 覆盖更多算子。custom autograd.Function 可能导出 ORT 无 kernel 的 op——export 后 onnx.checker.check_model。export PR 应附 checker 通过 log 与 ORT smoke 结果,而不是只附 .onnx 文件。

3. ORT 数值对齐

python
import onnxruntime as ort
import numpy as np

sess = ort.InferenceSession(
    "model.onnx", providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
ort_out = sess.run(None, {"input": x.cpu().numpy().astype(np.float32)})[0]

with torch.no_grad():
    pt_out = model(x).cpu().numpy()

np.testing.assert_allclose(ort_out, pt_out, rtol=1e-3, atol=1e-4)

logits L∞ diff 应 <1e-3(fp32);分类 argmax 一致是更强验收。部署用哪个 EP 就在哪个 EP 上验——CI 里 CPU EP 通过不代表 CUDA EP 数值相同。固定 32 张 val 图做 regression golden,防后续改 op 静默 drift。

4. trace 陷阱

torch.onnx.export 是 trace:data-dependent branch 被固化成 export 时走的那条。真正 dynamic control flow 考虑 torch.exportdo_constant_folding=True 减图大小,极少数 op 折叠后与 eager 有差——仍要 assert_allclose。export 前用 representative dummy(含边界 batch size)扫一遍。if self.training 分支必须在 export 前 eval 掉。

5. 预处理链一并验收

部署常是 resize → normalize → NCHW → model。ORT 输入若只验 random tensor,现场仍可能错——用 与生产相同的 preprocessing 跑 fixed image,对比 PT eager 与 ORT end-to-end。swapRB、mean/std、input range [0,1] vs [0,255] 是高频坑。预处理应与训练 script 共用模块,而不是 C++ 里重写一份。Mobile 端 int8 量化前必须先 fp32 ORT 对齐,再谈 calibration。

6. 二次转换:TensorRT / OpenVINO

PyTorch → ONNX → ORT 对齐后,再转 TensorRT 或 OpenVINO IR,每层转换后复验同一 input。TRT fp16/int8 校准引入 drift——业务 metric 与数值 atol 一起看。dynamic batch engine 要测 min/opt/max shape 各一组。转换链 artifact 版本与 ORT 基线 assert 结果一并归档,便于回滚。

7. 失败模式

现象原因
导出报错unsupported op / opset 低
ORT 与 PT 差大train 模式 / BN 统计
batch>1 失败未声明 dynamic_axes
TRT 与 ORT 不一致fp16 或 plugin

8. 验收

固定 seed + batch:PT vs ORT assert_allclose;多 batch size 各验一次。部署 target 上 latency 与数值双过关。val best 权重 export,不用 test 调导出选项。CI 加 export + ORT smoke,防 merge 后 silent break。

9. 案例:BN 在 train 模式导出

某 ResNet 部署 top-1 随机,查 export notebook 未调用 model.eval(),BN running mean 用错 batch 统计。加 eval + assert_allclose 后 ORT 与 PT 一致。团队规定:export 脚本第一行 assert not model.training,CI 失败则 block merge。

10. 与量化、剪枝的先后

INT8 量化前必须 fp32 ORT 对齐;量化后再用 calibration set 验 metric,不能跳过 ORT 基线。剪枝改 graph 后重新 export + ORT assert。部署链每一环 artifact 版本:fp32 onnx → ort-golden → trt-engine,回滚能指认是哪一环 drift。

11. 自定义算子与 fallback

aten:: 自定义 op 若 ORT 无 kernel,应改模型结构或注册 ORT custom op,而不是假设「PyTorch 能跑就行」。export 失败时先查 unsupported op 列表,再决定改 forward 还是升 opset。团队维护「可导出算子白名单」,新层 merge 前过 export smoke。Mobile 部署若用 NNAPI/CoreML 再转一层,每一环都保留 ORT fp32 golden 输入输出对,回滚能指认是哪一环引入 argmax 翻转。export 脚本与 train 脚本共用同一 eval() 与 preprocess 函数,禁止 copy-paste 两份 normalize。dynamic batch 服务要在 max batch 上 export 并验 ORT,再测 batch=1 latency。Control flow 含 Python if on tensor value 的模型,优先 refactor 为静态 graph 再 export,而不是指望 trace 运气。Serving 与 train 的 opset 版本不一致时,在 CI 用 deploy 侧 opset 做 assert,避免训练环境 export 过新、边缘设备 ORT 过旧。Export PR 模板:eval 截图、checker log、assert_allclose 数值、deploy EP 名称四件套齐全才 merge。

12. 案例:BN running mean 导出

MobileNet deploy acc 崩:export notebook 未 eval,ORT 与 PT 差 0.3 L∞。加 assert_allclose 后 block merge;export 脚本首行 assert not model.training。Serving 预热 batch 应用 train 统计已固定的 BN,勿在 serving 再跑 train 模式 forward。Edge ORT 版本写在 deploy manifest,export CI 用同版本 runner 跑 assert。Opset 升级 PR 应重跑全量 golden vectors,防算子语义微变。Segmentation 导出注意 output 是 logits 还是 mask,deploy 侧 argmax 与 train 一致。Dynamic axes 仅 batch 不够时,height/width 也要声明,否则 resize 服务 crash。Custom op 导出失败应报 issue 到训练侧改 arch,而不是 production 硬补 Python 后处理。INT8 量化 PR 在 ORT fp32 golden 通过后再提交,量化 regression 单独 job。Triton/TensorRT plugin 自定义 op 在 ORT 后再验 golden,plugin 版本纳入 deploy manifest。Segmentation 导出 output 名与 deploy 后处理约定一致,避免 logits/mask 混用。OpenVINO 与 ORT 数值差有时来自 NHWC layout,转换后 end-to-end 验 argmax 一致。Export 失败时保留 torch 侧 golden pickle,便于 ORT 侧 bisect 哪一层 drift。Mobile 部署 checklist:eval 模式、opset、input size、EP 四件与训练侧签字一致。Regression golden 集版本与 model checkpoint 同 tag,回滚成对进行。

ONNX 是部署契约——ORT 对齐是 export PR 的 merge 门槛。

← 全部文章

johan's blog