最近在把一个训练好的图像分类模型(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 条你这掉法大概率不是量化问题,opset 11的BN和AdaptiveAvgPool在ONNX Runtime里都有对应的官方实现,精度损失一般不会这么夸张。建议先检查一下预处理环节,比如mean/std的归一化方式在导出时有没有被固化进模型,或者推理时输入数据的通道顺序是不是和训练时一致。另外可以试试把opset升到15以上,某些老版本算子映射确实有坑。移动端的话,如果这个精度差距对你业务影响大,还是别将就,直接上TFLite的量化感知训练流程会更稳。
这种精度掉法大概率不是算子不支持,ResNet-50那几个层ONNX支持都挺成熟了。你试试导出前把模型设成eval模式,然后核对下预处理和后处理是不是跟PyTorch里完全一致,有时候是输入归一化参数没对齐。另外opsert_version可以拉到13以上,顺便检查下onnxruntime的优化级别,默认开全量优化有时候会改图结构。移动端的话4个点确实有点肉疼,但量化后可能更离谱,建议先拿校准集做下动态量化对比再决定。
遇到过类似的,但你这个4个点的跌幅确实大了点,不太像单纯的量化问题。建议先检查下模型里有没有自定义的forward逻辑或者预处理没对齐,ONNX导出时这些很容易被“静默”掉,尤其是ResNet这种带残差的,batch norm折叠后数值分布可能变。另外AdaptiveAvgPool在opset 11下会转成动态shape的Gather+ReduceMean,某些runtime实现会有微小误差,可以试试固定输入尺寸导出看精度是否回升。如果后续要上移动端,建议直接对比下TFLite和ONNX的量化后精度,而不是只看浮点结果,毕竟真机部署通常要int8,那才是真正的考验。
我之前也栽在过这坑里,ResNet-50导出ONNX后精度掉4个点大概率不是算子问题,而是导出时模型默认进入了训练模式,BatchNorm的running_mean和running_var没被冻结。你试试在export前强制调model.eval(),再把opset提到13以上,有时能解决不少隐式转换的精度差。另外AdaptiveAvgPool在ONNX里会展开成动态shape的Gather,对输入尺寸敏感,最好用固定分辨率导出并加一个Resize层预处理,能稳住精度。移动端部署的话,如果硬件不支持INT8,这精度损失确实明显,建议先做PTQ或QAT量化校准,比直接换TFLite靠谱,TFLite对自定义算子支持更麻烦。
92.3掉到88这个幅度有点太大了,一般纯转换不会差这么多,八成是你导出时没切eval模式或者dtype对不上。建议先检查model.eval()和torch.no_grad(),再对比ONNX和PyTorch同一张输入的输出,看是从哪一层开始飘的。BatchNorm和AdaptiveAvgPool导出基本没啥精度问题,opset 11也够用,别急着甩锅给算子。移动端部署的话88其实也还行,但得先搞清楚这4个点到底丢哪儿了,不然换TFLite一样会踩坑。
BatchNorm在ONNX里一般不会掉这么多,先别急着怀疑算子。你导出时模型是不是还在train模式?忘了eval()的话BN统计量不对,精度直接崩。另外确认下预处理和后处理两边是不是完全一致,归一化参数差一点都能掉好几个点。如果这些都没问题,再对比下ONNX和PyTorch同一批输入的输出差异,定位到具体层。移动端88%其实也能用,但TFLite对ResNet-50支持更成熟,值得试试。
92.3掉到88确实有点狠,正常ONNX转换的精度损失通常在0.1到0.5个点以内,你这个幅度基本可以排除单纯的浮点误差了。我建议先别急着怀疑算子,第一步应该做逐层对比:把PyTorch和ONNX Runtime每一层的输出都dump出来,看是从哪一层开始偏差变大的。很多时候问题出在预处理上,比如PyTorch的Normalize和ONNX里你自己写的mean/std顺序不一致,或者Resize的插值方式不同,这些坑特别隐蔽。另一个常见原因是eval模式,导出前忘了model.eval()的话BatchNorm会走training分支,统计量直接用batch的,精度肯定崩。opset 11对ResNet这种结构支持已经很成熟了,AdaptiveAvgPool会被拆成AveragePool加Reshape,一般不会有大问题。至于移动端,88%的Top-1其实还能用,但前提是你要先搞清楚这4个点到底丢在哪,不然换TFLite也是一样的结果,因为根源可能在导出流程而不是框架本身。
先查下是不是没设eval模式,BN的running stats没同步过去,我踩过这坑。
先查一下ONNX和PyTorch的输出差多少,八成是某层算子对不上,别急着换TFLite。