ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch转ONNX避坑指南:从算子兼容到推理优化

PyTorch转ONNX避坑指南:从算子兼容到推理优化 从 PyTorch 到 ONNX实际上是一条通往部署的必经之路但这条路坑比想象中多。尤其是当你手里的网络不再是教科书上的 ResNet18而是带着自定义算子、多分支、动态 shape 的复杂网络时torch.onnx.export那一行代码背后藏着无数个为什么报错。我最近刚把一个混合了 CNN 主干和 Transformer 编码器的模型完成 ONNX 转换并部署到 ONNX Runtime 上整个过程踩了十来个坑。这篇文章就把这些报错记录、排查思路和最终解决方案完整写出来希望能帮你少走几趟弯路。1. 转换前的真相ONNX 加速到底在加速什么先聊一个很多人误解的点ONNX 本身并不会让你的模型跑得更快它只是一个中间表示。真正提速的是 ONNX Runtime、TensorRT、OpenVINO、NCNN 这些推理引擎它们能拿到 ONNX 这样一份静态的计算图描述去做算子融合、内存复用、kernel 自动调优。而 PyTorch 在推理时要走 Python 调度、动态图解释这一层开销在工业级部署场景里是不可接受的。这次我转换的目标模型大概长这样一个 CSP 风格的 CNN 主干提取特征后面接了一个 4 层 Transformer Encoder 做全局建模最后输出三个尺度的检测头。模型里有nn.MultiheadAttention有动态的mask生成有torch.where、torch.cumsum、F.grid_sample这类稍不留神就出问题的算子。总参数量 46M输入是1x3x512x512。转换完用 ONNX Runtime 在 CPU 上测单帧从 PyTorch 的 860ms 降到了 540ms在 TensorRT FP16 下能跑到 23ms这就是转换的价值。在做任何转换之前强烈建议先确认一件事你到底要部署到哪里纯 CPU 服务器用 ONNX Runtime CPU 版就够了有 NVIDIA 显卡且追求极致性能直接 ONNX 转 TensorRT要上移动端或者嵌入式ONNX 转 NCNN/MNN 是常见路线。不同的目标决定了你要不要花精力处理动态维度、要不要做 int8 量化、要不要拆模型。别一上来就无脑转想清楚终点再出发。1.1 环境版本所有报错的第一个来源我见过太多人转换报错最后发现是版本不匹配。PyTorch、ONNX、ONNX Runtime 三者之间的兼容性说好听点是生态发展快说难听点就是互相跟不上。我这次用的组合是PyTorch 2.1.2 onnx 1.15.0 onnxruntime 1.17.1。这个组合在导出nn.MultiheadAttention时比较稳定。之前用 PyTorch 1.13 配 onnx 1.13 导同样的模型直接报Unsupported operator: aten::multi_head_attention_forward。所以第一件事先把版本对齐pip install torch2.1.2 onnx1.15.0 onnxruntime1.17.1另外一定要确认你的 onnxruntime 和推理时的环境一致。很多人转换和推理在两台机器上做版本不一致导致的结果就是我本地明明测试通过了部署机上怎么跑不起来。1.2 转换前先固化模型结构PyTorch 模型在eval()模式下有些层的行为会改变比如 Dropout、BatchNorm。导出之前必须先model.eval()这算是最基础的常识了。但还有一个容易被忽略的点把模型里所有不需要梯度计算的参数 freeze 住。model.eval() for param in model.parameters(): param.requires_grad False这一步不只是为了省内存更重要的是避免导出时计算图里混入梯度相关的节点。之前有人问过我为什么导出的 ONNX 里多了一堆奇怪的节点多半就是没有 freeze 参数或者没有 eval。2. 第一次报错维度推断失败的真正含义第一次运行torch.onnx.export时报错信息是这样的RuntimeError: Failed to export an ONNX attribute to, since its not constant, please try to make things这个报错其实挺误导人的它说某个属性不是常量但实际上问题出在torch.cumsum的返回值被当作后续操作的 shape 参数使用。ONNX 导出静态图时图的维度信息需要静态推断如果某个维度的值依赖运行时计算导出器就不知道该怎么处理。这类问题的本质是ONNX 是一个静态图格式它的 shape 信息在导出那一刻就固定了除非你指定动态轴。而 PyTorch 是动态图所有 shape 都是运行时算出来的。两者的世界观不一样。解决思路有两种思路一固定输入 shape避免动态计算检查模型里所有的 shape 相关操作比如x.size(-1)、x.shape[2]尽量改成直接传入常量或者用x.shape中确定的值。比如# 不好的写法 seq_len x.size(1) mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool() # 改进直接写死或者从外部传入 seq_len 64 # 已知的固定值 mask torch.triu(torch.ones(seq_len, seq_len), diagonal1).bool()思路二用动态轴如果输入尺寸确实不固定那就得在导出时声明dynamic_axes但动态轴会带来额外的性能损耗而且有些算子组合在动态 shape 下根本无法导出。能固定就固定不能固定再说。2.1 dynamic_axes 的正确打开方式如果你的模型确实需要支持多种输入尺寸比如检测模型要跑 640x640 和 320x320那必须在导出时设置动态轴torch.onnx.export( model, dummy_input, model.onnx, opset_version17, input_names[input], output_names[output], dynamic_axes{ input: {0: batch_size, 2: height, 3: width}, output: {0: batch_size, 2: height, 3: width}, } )但有一个坑我踩过不是所有层都能在动态 shape 下正常工作。比如nn.AdaptiveAvgPool2d在动态 shape 下是 OK 的但某些自定义的grid_sample配合动态 shape 就可能导出失败。此外动态 shape 下 ONNX Runtime 会做一些 shape 重推断性能会比静态 shape 慢一些。所以我的建议是能静态就静态动态只给 batch 维实在是业务需要再开放 H/W 维度。3. 第二波报错算子不兼容的连环拳固定了 shape 问题后新的报错又来了RuntimeError: Unsupported opset version 17 for op: ATen这是在 PyTorch 2.0 时代比较常见的问题。某些算子尤其是aten::*开头的在 torch 内部实现走的是 ATen 路径ONNX 导出器还没有对应的映射规则或者映射规则只支持到某个 opset 版本以下。我在这次转换中遇到的具体算子有四个3.1 nn.MultiheadAttention 的导出问题nn.MultiheadAttention在 PyTorch 2.x 里有原生导出支持但前提是你不能传入attn_mask时使用 bool 型 maskPyTorch 里 bool mask 表示不能看的位置而 ONNX 的attn_mask约定是 float 型 additive mask。这两者的语义不同导出时最容易翻车。报错信息经常是Unsupported operator: aten::_native_multi_head_attention我的解决方法是不用nn.MultiheadAttention模块而是手动实现 multi-head attention 的前向逻辑用torch.matmul、torch.softmax这些基础算子拼出来。这样虽然代码长了点但导出时每个算子都是 ONNX 认识的老朋友稳定性极高。class Attention(nn.Module): def __init__(self, dim, num_heads): super().__init__() self.num_heads num_heads self.head_dim dim // num_heads self.qkv nn.Linear(dim, dim * 3) self.proj nn.Linear(dim, dim) def forward(self, x, attn_maskNone): B, N, C x.shape qkv self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * (self.head_dim ** -0.5) if attn_mask is not None: attn attn attn_mask.unsqueeze(0).unsqueeze(0) attn attn.softmax(dim-1) x (attn v).transpose(1, 2).reshape(B, N, C) return self.proj(x)这样改完导出就顺畅多了。所以遇到复杂模块导出失败时先想能不能降级重写——用更基础的算子手动实现同样的逻辑。3.2 torch.where 的坑torch.where(condition, x, y)在大多数情况下能正常导出但如果condition是从tensor.shape推导出来的或者内部有非布尔张量参与就容易出问题。比如我有一段代码valid_mask points[..., 0] 0 output torch.where(valid_mask, values, torch.zeros_like(values))这个在 ONNX 里会被翻译成Where算子。但如果你在torch.where里用了x.shape相关判断比如torch.where(x.size(1) 10, y, z)这就是一个 Python 层的条件判断导出器会尝试把它变成一个If节点但If节点的处理在 ONNX 导出器中一直不太稳定。我的建议是把所有 Python 层的逻辑判断都放到模型外部模型内部只做张量运算。3.3 F.grid_sample 的版本兼容F.grid_sample在较老的 PyTorch 里导出时容易出问题尤其是配合align_cornersFalse时。新版 PyTorch2.x对这个算子的导出支持已经很完善了但如果你在用 1.x 版本建议升级。如果升级不了有一个绕行方案把grid_sample替换成多步插值组合。但这个方案比较复杂一般情况下不推荐。能升级就升级升级不了再考虑替换。3.4 控制流 if 语句的隐形问题如果你的模型 forward 里有if语句Python 层判断导出器会把判断执行的分支编译进静态图。这意味着导出时走if的哪条分支最终模型就只有那条分支。比如def forward(self, x): if self.training: return self._forward_train(x) else: return self._forward_infer(x)导出前你调用了model.eval()所以self.training是 False模型只会导出_forward_infer分支。这个逻辑是对的但很多人没意识到导出的 ONNX 模型已经固化了这条分支部署时不支持运行时切换。如果模型里存在基于输入数据的动态分支比如if x.sum() 0那 ONNX 导出器会尝试用If算子表示但If算子在很多推理引擎上支持度不高容易导致崩溃。遇到这种情况建议重构模型逻辑尽量消除运行时数据依赖的分支。4. 导出成功但推理结果不对检查这些隐藏雷区模型导出成功之后我以为万事大吉了结果用 ONNX Runtime 推理输出结果跟 PyTorch 比差距巨大。这类问题比报错更隐蔽因为整个过程没有任何异常提示。排查了三个多小时最终定位到三个雷区。4.1 BatchNorm 的统计量问题第一个问题出在 BatchNorm。PyTorch 模型在train模式下用的是 mini-batch 的均值和方差在eval模式下用的是 running_mean 和 running_var。如果在导出前忘了model.eval()会导致导出的 ONNX 里 BatchNorm 的均值和方差是错的推理结果自然不对。这个坑我在最开始就提到过但我发现很多人知道要eval()却不知道eval()要放在torch.onnx.export之前而且要确保作用在同一个模型实例上。更隐蔽的情况是模型有多个子模块导出前只对主模型调用了eval()但某个自定义子模块没有正确继承导致内部 BatchNorm 依然处于train模式。可以用一行代码排查for name, module in model.named_modules(): if isinstance(module, nn.BatchNorm2d): print(name, module.training)4.2 输入张量的内存格式差异PyTorch 和 ONNX Runtime 在输入张量的内存布局上可能有差异尤其是当模型内部使用了channels_last或者.permute()时。PyTorch 的 tensor 默认是contiguous的 NCHW 布局但某些操作会改变内存布局导出时如果没有正确标记ONNX Runtime 拿到的输入可能被错误解释。我遇到的情况是模型里有一段x x.permute(0, 2, 3, 1)再x x.contiguous()再x x.permute(0, 3, 1, 2)的代码这种来回 transpose 在 PyTorch 里没问题但导出成 ONNX 后某些引擎会优化掉中间的Transpose节点导致结果错乱。解决方法很粗暴在导出前用torch.onnx.export的check_traceTrue参数做一次一致性校验。但这个校验只检测 tensor 数值是否一致不保证内存布局。更稳妥的方法是在模型 forward 的入口和出口显式调用.contiguous()。4.3 动态 mask 的广播机制我之前提到模型里有动态 mask 的生成这个 mask 在 PyTorch 里是(N, L)形状但在 ONNX 里和 attention 的(B, H, N, N)张量做加法时广播规则可能和 PyTorch 不完全一致。这个问题的排查方法是导出前后各跑一遍用二分法逐个模块对比输出。具体做法是在模型 forward 里临时加几个 print 或者 hook记录中间 tensor导出前在 PyTorch 里记录一份用 ONNX Runtime 跑的时候再记录一份对比找到第一个不一致的模块。这个排查思路非常重要它帮我把问题精确定位到了 attention mask 的广播逻辑。最终修复方案是在生成 mask 之后显式地mask mask.unsqueeze(0).unsqueeze(0)将其扩展成(1, 1, N, N)这样导出的 ONNX 在广播时就和 PyTorch 保持一致了。5. 量化与进一步加速int8 量化的注意事项热搜里很多人也在问onnx 量化 int8这块我单独说一下。ONNX 的 int8 量化分为动态量化和静态量化两种。动态量化最简单但只对MatMul、Gemm这类算子有效静态量化需要校准数据集效果更好但步骤多。如果你的模型主要是 CNN静态量化能获得约 2-4 倍的 CPU 推理加速如果你的模型是 Transformer 结构效果可能没那么明显。量化常见的报错之一是RuntimeError: Quantization not supported for operator: LayerNormalization有些算子比如 LayerNorm、Softmax在 int8 下支持不佳或者需要特定版本的 ONNX Runtime 才支持。应对办法是对不支持量化的节点设置排除列表让它们保持 FP32 精度。from onnxruntime.quantization import quantize_static, QuantType from onnxruntime.quantization.shape_inference import quant_pre_process # 先做形状推断 quant_pre_process(model.onnx, model_preprocessed.onnx) # 静态量化 quantize_static( model_preprocessed.onnx, model_int8.onnx, calibration_data_readercalib_reader, quant_formatQuantFormat.QDQ, per_channelTrue, nodes_to_exclude[LayerNormalization_123, Softmax_456], )量化之后务必用同一份测试集对比 FP32 和 INT8 的精度差异。我之前做过一个分割模型的量化mIoU 从 0.82 掉到了 0.78虽然 4 个点看起来不多但在某些对精度敏感的业务场景里完全不可接受。所以量化前一定要先跑一遍精度评估量化后做对比。6. 完整导出流程我的最终方案经过前面的排查我终于跑通了一条相对稳定的导出流程。如果读者现在也面临 PyTorch 转 ONNX 的需求可以直接按下述流程来操作。6.1 导出前的模型梳理清单动手写导出代码之前先花半小时把模型的 forward 理一遍注意以下几点把所有 Python 层的if/else逻辑标注出来确认导出时走的是哪条分支把所有 shape 相关的运算如x.shape[1]、len(x)标注出来看看能不能改成常量把所有自定义的nn.Module检查一遍确认里面没有用到 ONNX 导出器不认识的算子确认输入张量的 dtype 和 shape固定住它们。这个过程非常像代码评审但评审对象是你的模型结构。很多时候报错只是表象深层原因在模型设计时就埋下了。6.2 一条完整的导出脚本模板这是我最终使用的导出脚本核心逻辑都在注释里了import torch import onnx import onnxruntime import numpy as np # 1. 加载模型并设置 eval model torch.load(model.pth, map_locationcpu) model.eval() # 2. 固定输入 shape dummy_input torch.randn(1, 3, 512, 512) # 3. 导出 torch.onnx.export( model, dummy_input, model.onnx, opset_version17, input_names[input], output_names[output], dynamic_axesNone, # 如果能固定就不要开动态 do_constant_foldingTrue, # 常量化折叠能省不少计算 verboseFalse, ) # 4. 检查 ONNX 模型 onnx_model onnx.load(model.onnx) onnx.checker.check_model(onnx_model) print(ONNX model check passed.) # 5. 用 onnxruntime 做一致性校验 sess onnxruntime.InferenceSession(model.onnx, providers[CPUExecutionProvider]) ort_outs sess.run(None, {input: dummy_input.numpy()}) with torch.no_grad(): torch_outs model(dummy_input) for i, (ort_out, torch_out) in enumerate(zip(ort_outs, torch_outs)): np.testing.assert_allclose(ort_out, torch_out.numpy(), rtol1e-3, atol1e-5) print(fOutput {i} matched. shape{ort_out.shape})do_constant_foldingTrue这个参数值得单独说一句。它会在导出时把一些只依赖常量的计算提前算好固化到 ONNX 图里。对于 BN 层、某些卷积层的融合非常有帮助。但注意如果你的模型里有动态 shape 相关的操作开启 constant folding 有时反而会引入不必要的常量节点。遇到这种问题可以试着关掉再对比。6.3 用 onnx-simplifier 做进一步优化官方导出完的 ONNX 图通常会有冗余节点我习惯再用onnx-simplifier优化一遍pip install onnx-simplifier python -m onnxsim model.onnx model_sim.onnxonnxsim会自动清理一些纯数学变换的冗余算子、融合部分 Transpose 和 Reshape并能对静态 shape 做进一步推断。简化后的模型通常会小 10%-30%推理速度也会快一些。需要注意的是onnxsim也不是万能的。我遇到过一次它把动态 shape 模型的 shape 相关节点优化掉导致运行时维度错误。所以运行完onnxsim之后一定要重新做一次数值一致性校验。7. 常见报错速查表把这次转换过程以及其他项目里遇到的常见报错整理成一张表方便大家快速定位问题报错信息关键词原因解决方案Unsupported opset算子需要更高/更低的 opset 版本检查算子支持的 opset 范围调整opset_versionFailed to export an ONNX attributeshape 或属性是动态计算的固定输入 shape或用常量替代动态属性not constant非 const 属性被用于图结构重写模型逻辑把动态计算放到模型外部Unsupported operator: aten::xxx该算子还没有 ONNX 映射升级 PyTorch/ONNX或手动重写该模块Exporting operator failed算子在当前 opset 下不匹配尝试切换 opset / 降级重写Python 层if分支错误导出时代码走的是未预期分支保证导出前eval()并手动检查分支推理结果和 PyTorch 不一致BatchNorm、内存格式、广播差异用前述的二分对比法逐模块定位onnxsim后模型出错shape 推断被错误优化不开动态轴或先做 simplify 再做动态性声明保留这张表的关键在于关键词。报错信息往往很长但真正有用的信息往往在最后几行。如果你在网上搜报错不要复制整段截取核心关键词去搜命中率会高很多。8. 关于 ONNX Runtime 的部署优化心得模型转换成功只是第一步真正上线前还有几个部署优化点值得关注。8.1 Provider 的选择ONNX Runtime 支持多个执行后端。CPU 场景下用默认的CPUExecutionProvider就行但如果是 Intel 平台可以考虑OpenVINOExecutionProvider需额外安装如果是 AMD有ROCmExecutionProviderNVIDIA GPU 上则用CUDAExecutionProvider或TensorrtExecutionProvider。不同 provider 对同一份 ONNX 模型的支持程度不一样尤其是一些新算子可能在默认 CPU 上能跑切到 TensorRT 后直接报不支持。所以部署前一定要先在小流量测试中验证。8.2 线程数与内存优化ONNX Runtime 的默认线程数往往不是最优的。我的实践经验是CPU 部署时intra_op_num_threads设置为物理核数的一半左右吞吐量更高设置OMP_NUM_THREADS环境变量也可能影响性能。具体还是要实测不同模型的最优配置不同。sess_options onnxruntime.SessionOptions() sess_options.intra_op_num_threads 4 sess_options.graph_optimization_level onnxruntime.GraphOptimizationLevel.ORT_ENABLE_ALL sess onnxruntime.InferenceSession(model.onnx, sess_options, providers[CPUExecutionProvider])graph_optimization_level也建议调到ORT_ENABLE_ALL这一项在不少模型上能带来 10%-30% 的加速。但同样要注意优化可能改变某些节点的行为用之前一定要过一遍数值校验。8.3 多路输入与动态 batch如果你要部署成 HTTP 服务一个很实际的问题是要不要支持动态 batch我的建议是不要贪心。动态 batch 带来的额外复杂度内存池管理、并发控制、超时处理往往比收益更大。与其做动态 batch不如直接固定 batch1然后用多进程/多线程横向扩展。具体做的时候我把服务的单次推理请求固定为 batch1用线程池承接并发请求实测 8 核 CPU 的机器可以稳定跑满 6 个并发任务P99 延迟没有明显劣化比动态 batch 的实现简单可靠得多。9. 更进一步的部署路线从 ONNX 到 TensorRT/NCNN最后再聊一下 ONNX 之后的方向。很多人在热搜词里搜yolo12 onnx转tensorrt其实就是在走这条更深层的部署优化路线。ONNX 是中间格式如果你最终目标是 TensorRT那么 ONNX 的导出精度直接影响 TensorRT 的转换效果。我的建议是导出 ONNX 时opset_version优先用 17 或 18这样 TensorRT 对算子的覆盖度最高。TensorRT 的转换工具是trtexec命令行格式大致是trtexec --onnxmodel.onnx --saveEnginemodel.engine --fp16 --workspace4096如果转换失败多半是某个算子 TensorRT 不支持。可以用onnx_graphsurgeon手动替换不支持节点或者反过来改回 PyTorch 模型设计时的算子选择尽量用 TensorRT 熟悉的基础算子。说到 Model Optimizer还有一个老工具NVIDIA 的Polygraphy可以用来对比 ONNX Runtime 和 TensorRT 的输出差异排查 TensorRT 推理结果不对的问题。这个工具我强烈推荐尤其是模型要落地到 TensorRT 上时它能帮你快速定位到是哪个层导致的结果不一致。10. 写在转模型之外的个人体会做模型转换这几年我最深的体会是转模型这件事不只是调 API的问题它逼着你去理解模型内部的算子构成、数据流向、数值精度。每一次为什么这个算子导不出来的追问都是在帮你梳理模型结构发现那些潜伏在代码里的性能问题和稳定性隐患。我踩过的坑里真正有价值的往往不是搜到一个 fix而是想明白为什么。想明白之后的每一段执行起来就非常顺固定 shape、重写注意力、检查广播、对比输出、量化校准每一步都有章可循。如果你现在也卡在某个 ONNX 转换的报错上建议先别急着到处复制别人代码不妨回到模型 forward 里把每一个可能产生动态行为的地方标出来再对照这篇文章的排查思路走一遍。问题的答案很多时候就藏在模型结构本身里。最后补一个实用的小技巧导出 ONNX 前可以用torch.jit.trace先做一次 trace 测试如果 trace 失败torch.onnx.export大概率也会失败。trace 报错信息往往更直白能帮你更快定位到是哪个算子出的问题。用这个方法我在不少项目里把排错时间缩短了一半以上。
返回列表