最近在做一个检测模型部署,用PyTorch训练好的模型转ONNX,用onnxruntime推理发现输出和原模型差很多,不是小误差,是那种完全对不上的。我检查了输入预处理、归一化参数,都是对齐的,动态轴也设了。模型里有F.interpolate和自定义的ROIAlign,不知道是不是这些算子转换有问题。另外我试了opset 11和12,结果都差不多。有没有大佬遇到过类似情况?是不是我转的时候漏了什么参数,还是说某些层必须用onnx-script重写?求指点,部署卡在这好几天了。
PyTorch转ONNX后推理结果和原模型不一致,是哪里出了问题?
全部回复
共 19 条之前跑过类似的检测模型,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,这个在转换时容易出问题,固定输入尺寸试试说不定就好了。
大概率是ROIAlign在onnx里实现有精度差异,试着导出时把custom op拆成组合算子或转成onnx-script重写。
先试试把ROIAlign导出时加个onnx::Gather的静态shape,我之前就是这么解决的,onnxruntime对动态shape支持有点坑。
我之前也踩过类似的坑,最后发现是F.interpolate的mode和align_corners参数在ONNX里默认行为不一致导致的。PyTorch里你如果没显式设align_corners,有些版本默认False,但ONNX导出时可能会按True处理,这玩意儿对坐标映射影响特别大,尤其是上采样倍数大的时候,输出直接歪掉。ROIAlign的话,如果用的是torchvision的版本,导出ONNX时经常会把crop_and_resize拆成几个基础算子,但不同opset下这些算子的实现细节有差异,建议你先把自定义ROIAlign替换成ONNX官方支持的版本试试,或者干脆用onnxruntime的contrib op。另外你可以把原模型和ONNX模型的每一层输出都打印出来对比,用onnxruntime的IOBinding或者torch的hook,定位到底是从哪个节点开始分叉的,别只盯着最终输出。还有个小细节,你检查下输入张量的维度顺序,ONNX里NCHW是硬性要求,但如果你在PyTorch里用了NHWC或者中间有permute,导出的图可能带着额外的transpose,这会导致数值上看起来对不上,其实只是内存布局不同。如果实在找不到原因,建议用onnx-simplifier过一遍,它能把很多冗余的reshape和cast清掉,有时候问题就出在这些“隐形”操作上。我那次最后就是靠逐层对比发现是Resize的coordinate_transformation_mode设成了asymmetric,改成half_pixel就完全一致了。
先跑一下onnxruntime的onnx.checker和动态轴shape对不对,ROIAlign大概率要自己写个onnx算子替换。
我之前跑分割模型也撞到过这个坑,最后发现是F.interpolate的scale_factor和size在不同opset下的行为不一样,尤其是当输入尺寸是动态的时候,onnxruntime会默认用resize里的坐标变换模式,跟PyTorch默认的align_corners=False对不上。你那个自定义ROIAlign大概率是问题核心,PyTorch的autograd版本和onnx导出的静态算子实现细节差别很大,建议先用onnxruntime的CPUExecutionProvider跑一下,排除是CUDA provider的精度问题。另外可以试一下把模型里所有涉及shape计算的步骤都改成显式张量操作,有时候trace过程会把某些控制流拍平成固定值,导致动态轴根本没生效。你检查一下onnx输出的中间节点,对比一下原模型在ROIAlign前后的特征图差异,如果从那里就开始发散,那基本就是算子转换没跑了。还有个笨办法,把模型切成几段分别导出再比对,能快速定位是哪一层出的问题,我之前就是这么干的,比瞎猜效率高。最后如果实在不行,可以看看onnx-simplifier能不能把多余节点折叠掉,有时候是转换过程产生了精度丢失的冗余运算。
遇到过类似的,重点查一下自定义ROIAlign,这个大概率是转换时被拆成了多个基础算子,浮点精度和索引计算对不上就会完全跑偏。F.interpolate一般问题不大,但建议把mode和对齐方式显式写清楚。可以先试着把ROIAlign替换成torchvision官方版本再转一次对比下,能定位是不是它的问题。另外确认下onnxruntime是不是用了float16,有时候自动混合精度会悄悄改精度。
我之前也踩过类似的坑,最后发现是F.interpolate的mode默认值在转换时被改了,特别是align_corners这个参数,PyTorch和ONNX的默认行为不一致,建议你显式指定一下试试。另外自定义ROIAlign基本是必炸的,如果版本里没有匹配的算子,输出完全乱掉很正常,你可以先用onnxruntime的CPU执行模式跑一下看看是不是算子fallback的问题。还有个排查技巧,把模型简化一下,只保留到ROIAlign之前的层,对比中间输出,这样能快速定位是哪一段开始歪的。如果实在不行,建议直接用onnx-script把那几个自定义层重写,别硬刚转换器。
F.interpolate转ONNX大概率是坑,尤其mode是bilinear且align_corners没显式指定的时候,onnxruntime默认行为和PyTorch会对不上,输出偏移会逐层放大。自定义ROIAlign基本得用symbolic重写,torch.onnx默认那套导出十有八九是错的,建议先单独把这块抽出来验证。可以逐层对比ONNX和PyTorch的中间输出,定位到具体哪一层开始炸,比整体瞎猜快多了。opset影响没那么大,问题多半出在算子语义而不是版本。
F.interpolate在opset 11以下默认走upsample,坐标变换方式和PyTorch有细微差别,检测模型里这种偏差会被后续层放大。自定义ROIAlign基本得用symbolic重写,不然导出后可能直接变成一堆乱七八糟的算子组合。建议你先把模型拆开,逐层对比ONNX和PyTorch的输出,定位到第一个出问题的节点再针对性处理。
先单独跑下F.interpolate那段,八成是align_corners没对上,这坑我踩过。