最近在把一个训练好的图像分类模型(ResNet-50)从PyTorch导出到ONNX,再用ONNX Runtime做推理。结果发现精度掉得很明显,Top-1从原来的92.3%掉到了88%左右。我试了用torch.onnx.export,设置opset_version=11,也试了用onnx-simplifier简化模型,但效果不大。有没有老哥遇到过类似的问题?是ONNX对某些算子(比如BatchNorm、AdaptiveAvgPool)的转换有精度损失,还是我导出时没开正确的优化选项?另外,如果后续要部署到移动端,这个精度掉得能接受吗?还是说应该直接上TFLite?求指点。
PyTorch转ONNX后推理精度下降,是量化问题还是算子不支持?
全部回复
共 169 条这个精度掉4个点确实有点多,我怀疑不是算子转换的问题,ResNet-50这种经典结构在ONNX里早就成熟了。你试试检查下预处理是不是在导出时被并进图里了,比如Normalize的mean/std如果写错或者精度变成float16,影响比算子大得多。另外opset 11对某些层的融合确实不够好,可以试试opset 15以上,或者直接用torch.onnx.export的 dynamo模式,我上次用那个解决过类似问题。移动端的话,除非你模型本身很大,否则这精度损失对部署来说挺伤的,TFLite如果量化校准做得好反而可能更稳。
我之前也栽过这坑,ResNet-50转ONNX精度掉这么多大概率不是量化的问题,你先检查下导出时模型是不是误设成了train模式,BatchNorm和Dropout在eval和train下行为差挺多的。另外AdaptiveAvgPool在opset 11下有时会展开成动态shape的slice操作,数值上会有微小偏差,建议直接固定输入尺寸试试。移动端部署的话,4个点的精度损失确实肉疼,如果模型不大建议先用torch量化感知训练再导出,或者直接试TFLite的量化,通常比ONNX Runtime在ARM上稳。
先查下模型里有没有动态尺寸输入,AdaptiveAvgPool在opset11下容易出问题,固定输入尺寸试试。
老哥你这降幅不像量化,更像导出时某些层被重算精度丢了,建议对比下onnx和pytorch的逐层输出。
之前跑YOLOv5转ONNX也踩过类似的坑,精度掉1-2个点算常见,但你这掉了4个多点确实有点狠。我建议先别急着怀疑量化,你opset=11默认就是FP32导出,根本没量化,问题大概率出在算子实现差异上。AdaptiveAvgPool在静态shape下其实还好,真正容易出问题的是BatchNorm的folding——PyTorch推理时BN是跟卷积融合的,但ONNX导出有时会保留独立BN节点,导致数值计算顺序不同,精度就漂了。你先试试把模型设成eval模式再导出,然后检查一下导出的图里有没有多余的Reshape或Transpose,有时候onnx-simplifier反而会引入新的精度问题。另外一个排查思路是逐个替换可疑算子,比如把AdaptiveAvgPool改成固定kernel的AvgPool,看精度是否恢复。如果确认是BN的问题,可以手动把BN参数融合进卷积权重再导出,这步很关键。至于移动端部署,TFLite和ONNX Runtime的精度差异其实不大,关键看你量化方案——如果做INT8量化,掉2-3个点在分类任务上算正常,但你这FP32就掉4个点,说明不是量化的锅,得先把精度找回来再说。可以先用官方torchvision的ResNet50权重试一遍,排除是自己训练时某些层行为特殊导致的。
先检查下预处理和归一化参数是否一致,这步最容易出精度差;AdaptiveAvgPool在opset11下确实有精度问题,建议升到13试试。
我之前也踩过类似的坑,ResNet-50导出ONNX后掉点大概率不是量化的问题,因为你现在跑的是FP32,更像是算子映射差异导致的。AdaptiveAvgPool在ONNX里会展开成Gather+Div的组合,有些实现会在边界处理上跟PyTorch的ceil模式对不齐,这个精度损失在分类任务上可能不大,但确实会累积。你不如先跑个逐层对比,把PyTorch和ONNX Runtime的中间特征图导出来算下余弦相似度,定位到是哪个block开始漂移的,通常都是BN折叠的epsilon参数没对齐。另外opset_version=11对于ResNet来说够用了,但onnx-simplifier有时候会把一些融合搞出数值误差,建议原始导出先别简化,直接跑一遍看看。移动端部署的话,这个精度掉得确实有点多,如果你后续还要量化到INT8,那掉点会更严重,所以建议先解决FP32的转换精度问题再说。TFLite也不是银弹,如果ONNX这边搞不定,换框架大概率会遇到新的算子兼容问题,不如先排查清楚。你试试在export时把training参数设成False,并且手动把模型切到eval模式,有些人会漏掉这一步导致BN统计量被当成训练时的running_mean。
ResNet-50这个掉点幅度不太像纯算子转换的问题,更像是BatchNorm层在推理时被折叠进卷积后,权重精度没处理好。你可以先试试导出时把opset提到13以上,同时打开onnxruntime的graph optimization level,看看能不能拉回来一点。另外AdaptiveAvgPool在ONNX里会展开成动态形状的Mean操作,这个对输入尺寸敏感,如果模型不是固定输入尺寸导出,精度波动会很明显。移动端部署的话,这个精度损失确实有点大,建议先排查清楚再说,别急着换TFLite,毕竟你连ONNX的优化空间都没压榨完。
我之前也踩过这个坑,ResNet-50转ONNX掉点大概率不是算子不支持,而是导出时把训练模式和推理模式搞混了,BatchNorm的running_mean和running_var没冻结,你试试model.eval()之后再导出,精度应该能回来。另外AdaptiveAvgPool在opset 11下会展开成多个Slice+ReduceMean的组合,数值上可能有微小差异,但不会掉4个点这么多。移动端的话建议先量化感知训练再转TFLite,直接后量化掉点更狠,ONNX Runtime的FP16推理反而可能比FP32更稳一点。
你这精度掉4个点确实不正常,ResNet-50这种老模型ONNX转换很成熟了,大概率不是算子问题。我怀疑是导出时BN层被融合进了卷积,但某些版本的ONNX Runtime对融合后的数值处理有细微差异,你可以试试用onnxruntime的CUDA EP跑一下对比CPU结果。另外AdaptiveAvgPool在opset 11里实现是动态shape的,有时候会触发低精度路径,建议固定输入尺寸导出试试。移动端的话这精度肯定不能接受,但也不一定要TFLite,可以先用onnxruntime的量化工具做PTQ看看,能保住90%以上再考虑部署。
先试试把模型切成fp32单独跑onnx,排除量化影响,你这大概率是导出时某些层被重写导致数值漂移。
我最近也踩过类似的坑,ResNet50导出ONNX后精度掉到88%确实不太正常,这个幅度不像是单纯量化造成的。你试过把opset_version提到13或者更高吗?有些算子在低版本opset下会走fallback路径,比如AdaptiveAvgPool在旧版本可能会被展开成多个slice加mean,这中间浮点误差累积起来还挺可观的。另外你检查过导出前后BN层的running_mean和running_var有没有被正确冻结吗?有时候训练模式和eval模式没切换干净,BN层在导出时会带着训练时的统计数据走,那个误差比算子转换大多了。我记得onnx-simplifier虽然能简化图结构,但它不会管数值精度,如果原始图里某些算子本身就有精度问题,简化反而可能把误差放大。建议你先用onnxruntime的精度分析工具跑一下,看看每层输出和PyTorch的差异到底在哪一层开始变大的,不然瞎调参数效率太低了。至于移动端部署,如果精度要求这么苛刻,我建议要么试一下TensorRT或者CoreML的转换路径,要么干脆在ONNX Runtime里开启FP16试试看,有时候精度下降不是模型问题而是运行时的数值处理方式不同。TFLite也不一定就比ONNX好,量化感知训练才是关键,你现在这个精度掉法感觉更像是导出配置的问题,不是框架选型的问题。
这精度掉得有点狠,不太像纯量化损失,更像导出时某些层被替换成非等价的实现。建议先关掉onnx-simplifier试试,它有时候会乱合并节点,尤其对BatchNorm和残差结构。另外检查下模型的预处理(mean/std)在导出时有没有被固化进去,ONNX Runtime推理时如果输入没对齐,精度也会崩。移动端部署的话,这精度肯定不达标,TFLite量化调好的话一般能控制在1%以内,但ResNet-50转TFLite也有坑,不如先排查ONNX的问题。
你这情况我太熟了,之前我转YOLOv5的时候也掉过差不多的点,后来发现是模型里带了一个动态尺寸的AdaptiveAvgPool,ONNX导出后默认会展开成两个Resize节点,数值上跟PyTorch的均值计算有细微差别,尤其是特征图尺寸不是整数倍的时候误差会累积。你可以先试试把输入固定成单尺寸,或者干脆把AdaptiveAvgPool换成GlobalAvgPool,很多情况下能直接拉回精度。另外BatchNorm在训练和推理时是两套行为,PyTorch导出时如果你没把模型切到eval模式,BN的running_mean和running_var根本不会用上,这坑我踩过不止一次,八成你的问题也在这。至于opset_version,11其实够用,但如果你用了SiLU或者一些新激活函数,建议升到13以上,有些老版本算子会降级成近似实现。onnx-simplifier有时候会把图结构改得过度精简,反而破坏了数值流,我一般只在确认无精度问题后才跑它。移动端部署的话,这个精度掉4个点我觉得不太能接受,尤其是分类任务,用户感知很强。TFLite也不一定稳,量化感知训练才是关键,如果非要用ONNX Runtime,试试FP16或者INT8动态量化前先跑一遍per-channel校准,看看是不是量化误差主导。你要是方便的话,可以先把导出后的onnx用onnxruntime的python接口跑一遍,跟PyTorch的逐层输出对比,定位到具体是哪一层开始分叉,这样最省时间。
我之前也踩过类似的坑,ResNet-50转ONNX掉点大概率不是量化的问题,因为你现在还是FP32推理,根本没走量化流程。更可能出在导出时模型里的BatchNorm被fuse进Conv的方式跟PyTorch原版不一致,或者是AdaptiveAvgPool在ONNX里被展开成动态shape的Gather操作,导致数值上有一点点偏差,积少成多就掉了一个多点。你可以试试用onnxruntime的C++ API跑一下,对比Python端看结果是否一致,有时候Python的前后处理里带了imagenet的mean/std归一化,导出时没冻结进去,这个最容易忽略。另外opset_version=11对某些算子支持确实不友好,你可以试试升到13或者15,同时用torch.onnx.export里的dynamic_axes=False固定输入尺寸,有时候动态shape会触发不同的内核实现。至于移动端,88%的Top-1如果对业务影响不大倒也能用,但你要是追求和原模型一致,建议先排查是不是预处理差异,而不是急着换TFLite——TFLite转起来坑更多,尤其是量化感知训练没做的话,掉点会更狠。我之前有个分类模型排查到最后是torch的interpolate默认模式跟ONNX的resize对齐方式不一样,改一下align_corners就完全一致了,你可以先打印一下ONNX模型里每个算子的输出跟PyTorch逐层对比,定位到具体哪一层开始分叉。
大概率是导出时BN和AdaptiveAvgPool的转换问题,建议先关掉eval模式再试下,或者转成固定输入尺寸看看。
移动端这精度掉得有点狠,还是先排查算子吧,TFLite不一定更好调。
先查下模型有没有用faster模式导出,opset拉高到15试试,这精度差多半是算子映射问题。
先检查下预处理和后处理是不是对齐了,92掉88大概率不是算子问题,是数据流差异。
之前也踩过这坑,对比下onnx和pytorch每层输出,定位到具体层再调。
这种精度掉法大概率不是量化的问题,你opset 11导出默认就是FP32,跟量化没关系。ResNet-50里最容易出问题的其实是BatchNorm在训练和推理模式下的行为差异,以及AdaptiveAvgPool在ONNX里会被拆成多个算子导致数值累积误差,建议先导出后用onnxruntime的python API逐层对比中间输出。另外你试过把torch的model.eval()和torch.no_grad()都加上再导出吗?很多时候精度掉是因为BN层统计量没固定。移动端的话4个点的掉幅确实有点大,建议先排查清楚再考虑TFLite,不然换框架大概率还是同样的问题。
之前跑YOLOv5也踩过这个坑,92%掉到88%大概率不是量化的问题,你这还是FP32导出吧?先检查下预处理是不是对齐了,PyTorch和ONNX Runtime的输入归一化参数很容易不一致,另外AdaptiveAvgPool确实有坑,试试把输入尺寸固定死再导出。
移动端部署的话这个精度肯定没法接受,不过也别急着换TFLite,先试下opset 13+,或者用onnxruntime的graph optimization level调成all,有时候能救回来一点。实在不行再考虑TFLite,毕竟改框架成本也挺高的。
说实话你这个精度掉得有点多,正常ResNet-50转ONNX即使不动量化,top-1差距也就在0.1%以内,所以我第一反应不是算子转换问题,而是导出时模型本身有坑。你试试在export时把training参数显式设成False,再把model.eval()放到最前面,很多人的BatchNorm统计量就是在这时候被带进推理图的。另外AdaptiveAvgPool在opset 11里确实会转成动态shape的ReduceMean,ONNX Runtime跑起来可能跟PyTorch的静态shape行为不完全一致,但影响通常很小,不至于掉4个点。
我怀疑更可能的是你用onnx-simplifier的时候把某些fold BatchNorm的数值精度搞坏了,那个工具对fold后的常量浮点运算有时会引入误差,尤其是当模型里有大量小数值权重时。你可以先不开simplifier,直接跑一下原始ONNX的精度,如果恢复92%左右,那问题就锁定了。如果原始ONNX也掉,那再检查输入预处理,比如mean/std是否被重复归一化,或者导出时有没有把torchvision的transforms里的RandomResizedCrop的插值方式带进去。
至于移动端部署,88%的Top-1对实际产品肯定不够,但如果你本来就要做int8量化,那这点精度损失反而可能是量化误差和转换误差叠加的结果。我的建议是别急着上TFLite,先把ONNX这块搞干净,因为TFLite对ResNet-50的支持也不一定更友好,而且你从PyTorch转TFLite中间还得过ONNX或者直接转,反而多一层转换风险。不如先在ONNX Runtime里试一下fp16或int8量化,看能不能把精度拉回91%以上,如果量化后能稳住,那再考虑移动端。另外你可以跑一下onnxruntime的graph optimization level,默认是ALL,有时候改成ORT_ENABLE_EXTENDED反而会引入问题,手动关掉某些优化再测一下。