最近在做一个检测模型部署,用PyTorch训练好的模型转ONNX,用onnxruntime推理发现输出和原模型差很多,不是小误差,是那种完全对不上的。我检查了输入预处理、归一化参数,都是对齐的,动态轴也设了。模型里有F.interpolate和自定义的ROIAlign,不知道是不是这些算子转换有问题。另外我试了opset 11和12,结果都差不多。有没有大佬遇到过类似情况?是不是我转的时候漏了什么参数,还是说某些层必须用onnx-script重写?求指点,部署卡在这好几天了。
PyTorch转ONNX后推理结果和原模型不一致,是哪里出了问题?
全部回复
共 9 条之前跑过类似的检测模型,F.interpolate一般没事,但自定义ROIAlign大概率是罪魁祸首,onnxruntime对这类自定义算子的支持很看版本,建议先把这个替换成标准roi_align试试。另外你检查下转onnx前有没有设model.eval(),还有输入张量是不是带梯度的,这俩坑我踩过好几次,输出直接乱飘。如果替换后还不对,把onnx用onnxsimplifier简化一下,有时候是图优化把某些节点搞坏了。
之前跑分割模型也遇到过类似情况,最后发现是F.interpolate的坐标模式在转换时默认值变了,原模型里是align_corners=True,转出来没保留这个参数。你检查下这个,另外ROIAlign如果用的是torchvision实现,建议先用onnx-script或者把自定义算子注册成onnx::CustomOp试试,单纯靠torch.onnx.export有时候会把自定义逻辑折叠成一个奇怪的子图。还有个笨办法,把原模型里每个层输出都存下来,跟onnxruntime跑的中间结果逐层对比,很快能定位到是哪个节点开始崩的。
我之前也踩过这个坑,自定义ROIAlign大概率是罪魁祸首,PyTorch里很多自定义实现转ONNX时算子映射不全,容易直接垮掉。你可以先试着把模型里F.interpolate和ROIAlign单独抽出来转一下,看输出是否正常。另外检查下有没有用到torch.where或者mask这种动态shape的操作,这些在ONNX里经常出幺蛾子。不行的话就试试用onnx-script把ROIAlign重写一遍,或者干脆用onnxruntime的contrib op,我上次就是这么解决的。
我之前也踩过类似的坑,尤其是自定义ROIAlign,PyTorch里实现可能依赖了某些python控制流或者自定义autograd,ONNX导出时这些逻辑根本没法完整映射。F.interpolate本身没问题,但如果你用了align_corners或者mode参数不同版本默认值有差异,也会导致输出偏差,建议你把onnx模型用onnxruntime的graph优化关掉试试,有时候优化会改算子组合。另外你对比过中间层的输出吗?比如把原模型和onnx模型的某一层feature map导出来对齐一下,能快速定位是哪个算子开始漂移的。还有个小细节,torch转onnx时如果模型里有inplace操作,比如relu(inplace=True),某些版本下会导致计算图错误,可以先全局搜一下。如果实在不行,可以考虑把ROIAlign换成torchvision官方版本,或者用onnx-script重写那一块,我之前是重写了才过的。opset的话11和12差距不大,但如果模型里有比较新的算子,建议直接上13以上。你那个检测模型是两阶段的吗?如果是的话,后处理里的nms可能也受影响,建议把后处理逻辑放到onnx外面做,别让模型输出太多冗余框。
遇到过类似的,最后发现是ROIAlign的坐标映射问题,PyTorch里crop和resize的align_corners默认值和ONNX的算子实现不一致,这个坑特别隐蔽。另外F.interpolate建议先确认mode和align_corners是否在ONNX里有对应支持,不然很容易静默转换但结果错。你可以先把自定义ROIAlign替换成grid_sample试试,或者dump每层输出对比一下,看从哪一步开始分叉的,比瞎猜快。opset版本影响不大,别在这上面浪费时间。
我之前也踩过类似的坑,最后发现是F.interpolate的mode默认值在转ONNX时被固定成了nearest,跟PyTorch里默认的bilinear对不上,输出直接崩。你检查一下导出时有没有显式指定mode和align_corners,这两个参数很容易漏。ROIAlign的话,建议先用onnxruntime的推理日志把中间层输出打印出来,跟PyTorch逐层对比,定位是哪个节点开始发散。另外如果模型里有动态shape,试试固定输入尺寸导出,排除一下维度广播的问题。
试试把F.interpolate换成固定尺寸再转,ROIAlign建议用onnx-script重写,这俩最容易出问题。
我之前也是自定义层导致输出全乱,最后老老实实rewrite才搞定。
ROIAlign这块基本可以确定是转换重灾区,建议先单独导出这一层对比下输出,大概率是它的问题。
我之前也踩过类似的坑,最后发现是ROIAlign在转换时被拆成了几个基础算子,浮点精度和坐标对齐方式跟PyTorch原生实现有细微差别,累计起来结果就完全飘了。建议先单独把这两个模块抠出来测一下,或者试试用onnxruntime的CUDA执行提供程序跑,有时候CPU和GPU的算子实现也不一样。另外检查下模型里有没有动态shape相关的操作,比如reshape用了-1,这个在转换时容易出问题,固定输入尺寸试试说不定就好了。