最近在把训练好的一个图像分类模型(ResNet50)转成ONNX,部署到移动端用。按照文档设了dynamic_axes={‘input’: {0: ‘batch_size’}},导出也成功了。但用onnxruntime推理时,只要batch_size不是1就报错,说“输入形状不匹配”?试了opset版本11和12都一样。是不是导出时少设了什么参数,比如输入尺寸的dynamic_axes必须和模型内部结构对应?还是说模型里有不支持动态batch的层(比如BatchNorm)?求有经验的大佬指点一下,我实在不想把batch固定为1。
PyTorch模型转ONNX时,动态轴设了但推理报错,有老哥踩过坑吗?
全部回复
共 165 条之前也遇到过类似问题,导出时设置dynamic_axes只是告诉ONNX这个维度是可变的,但实际推理时输入tensor的shape必须和模型内部算子的实际支持范围匹配。ResNet50本身没有BatchNorm的batch依赖问题,不过如果模型里用了reshape或者view这类算子,导出时可能没把batch维度传导给后续层,建议用onnx-simplifier过一遍,或者导完用netron看一眼中间节点的shape有没有变成-1。
另外onnxruntime对动态shape的优化有时会在特定opset下抽风,可以试试opset 13以上,或者显式指定输入输出名字再检查一下。我之前是改成了固定batch=1才通过,但你这需求的话,可以试试在导出时把input的shape写成[None,3,224,224],有些版本对None的处理更宽松。
可能是ONNX输入没写全,试试把input和output的dynamic_axes都显式设上,包括batch维。
我之前也卡过这问题,八成不是BatchNorm的锅,你查下onnxruntime的session options里有没有开优化,有时候图优化会把动态轴给折叠掉。另外建议用onnxruntime的shape inference跑一下,看看中间节点的shape是不是都带动态维度,如果中间层被推断成固定shape了,那大概率是导出时某些算子没映射全。还有个笨办法,直接把batch_size轴改成-1试试,有些版本对0的处理有bug。
我之前也卡过这个坑,多半不是BatchNorm的问题,而是onnxruntime的session配置里没开动态shape支持。你试试在推理前设置sess_options.add_free_dimension_override_by_name("batch_size", batch),或者直接用onnxruntime的IOBinding指定实际shape。另外检查下输入节点的dtype是不是和你feed的数据一致,有时候类型不匹配也会报成形状问题。如果还不行,可以把导出的onnx用onnxsimplifier优化一遍,某些冗余算子会导致动态维度传递失效。
大概率是onnxruntime的输入name没对上,导出时dynamic_axes的key得和实际输入张量名一致,检查下。
遇到过类似情况,多半不是BatchNorm的锅,ResNet50本身支持动态batch。你检查下onnxruntime的session配置,推理时输入数据维度要显式写清楚,比如用np.newaxis把shape补成(N,C,H,W),别只传个列表。另外试试用onnx-simplifier把模型过一遍,有时候导出时会有多余的reshape节点把动态轴写死了。我之前就是这么解决的,opset用13试试也行。
试试把input的shape写成[None,3,224,224]再导,光设dynamic_axes有时候不够。
遇到过类似的坑,十有八九是输入数据的维度写死了。你导出时虽然设了dynamic_axes,但onnxruntime的session输入可能还是按固定shape去读的,得在run之前把input的shape动态设一下,比如用session.get_inputs()拿name,再手动reshape成[batch, 3, 224, 224]。另外ResNet50里的BatchNorm在推理模式应该没问题,不是它的事,重点查下预处理或者dataloader是不是默认batch=1。
我之前也卡过一模一样的问题,后来发现是onnxruntime的session配置里没开动态shape支持,得设一下session_options的optimized_model_filepath或者直接关掉shape推断试试。另外BatchNorm本身不挡动态batch,但如果你用了torch.nn.DataParallel转出来的模型会有问题,建议先转成单卡再导出。还有个坑是resize或者全连接层之前如果用了flatten,维度写死也会报这个错,可以检查下onnx图里是不是有固定维度的reshape节点。
大概率是reshape或view里写死了batch维度,试试导出前把模型里所有reshape改成-1。
也可能是onnxruntime的session配置问题,检查下动态轴名字和实际输入tensor形状对不对得上。
我之前也踩过这个坑,导出的onnx里动态轴明明写了,但推理时batch一换就炸。后来发现问题是模型里有个reshape或者flatten操作把batch维写死了,ResNet50里global avg pooling之后那部分特别容易出这问题,你检查下onnx图里有没有shape输出是固定值的节点。另外你只设了input的dynamic_axes,但中间层如果用了torch的view或者tensor.size(0)这种,导出时计算图会把batch维当成常量,onnxruntime就会按固定shape去校验。建议你在导出时把模型的forward里所有用batch size的地方都改成用x.shape[0]动态获取,别用硬编码的变量。还有个笨办法,导出后用onnxsimplifier或者onnx-graphsurgeon把图里所有shape相关的常量改成动态,我上次就是这么救回来的。另外你试过opset 13或者14没,有些算子在新版本里动态shape支持更完善,老版本经常有bug。如果你用的是torch.onnx.export,建议把input_names和output_names的dynamic_axes都写上,只写input那边有时候输出shape也会被固定死。最后实在不行就干脆batch固定1,移动端推理很多场景单张图也够用,但你要是做视频流或者多帧聚合那还是得解决,可以贴下报错的具体节点名,我帮你看看是哪儿卡住的。
我之前也卡过这问题,后来发现是onnxruntime的session选项里没开动态shape支持,得设一下optimized_model_filepath或者用session_options.add_free_dimension_override_by_name手动指定维度。另外BatchNorm本身没问题,但如果你用了torchvision的ResNet,里面的AdaptiveAvgPool2d在导出时可能会把动态轴固化成具体尺寸,建议检查下导出后的图,或者试下opset13+。实在不行就先固定batch为1,跑通了再慢慢排查。
我之前也卡过这个坑,你导出时只设了input的dynamic_axes,但输出端也要对应设一下,不然ONNX的图里输出shape还是写死的,推理时batch变化就会报错。另外ResNet50里的BatchNorm在eval模式下转的话一般没问题,但如果你模型里还有flatten或者reshape层,最好检查下有没有把batch维写死成1。还有个笨办法,先用onnx-simplifier过一遍图,有时候能自动修正一些shape推断问题,挺管用的。
我之前也卡过这问题,多半不是BatchNorm的锅,ResNet50里BN在导出时是固定了统计量的,动态轴其实只对卷积和全连接这类层有效。你检查下onnxruntime的session设置,得显式指定providers和optimization level,有时候默认优化会把动态形状搞挂。另外试试用onnxruntime的symbolic shape inference跑一遍模型,有些算子的输出shape没被正确标记为动态。我之前是加了个reshape层把输入显式改成[None,3,224,224]才好的,你参考下。
大概率是模型里有个flatten或reshape把batch写死了,检查下导出前的forward里有没有硬编码维度。
我之前也卡这儿好久,后来发现是onnxruntime的session选项里没开动态shape支持,得设一下set_providers或者把graph优化级别调低,默认优化会把它固化掉。另外你检查下模型里有没有reshape或者view操作依赖固定batch,ResNet50理论上没问题,但预处理部分容易埋雷。实在不行试试把dynamic_axes的input和output都写上,只写一边有时候会漏掉输出维度的绑定。
你试试把输入名和输出名的dynamic_axes都配上,光设input有时不够,我之前就这样踩过坑。
遇到过类似情况,多半不是BatchNorm的问题,而是你导出时虽然设了动态轴,但模型里如果有reshape或者view操作,它会把batch维写死,onnxruntime就会按静态shape去校验。建议先onnx.shape_inference跑一下,看中间张量的batch维是不是还是-1,不是的话就得改模型把reshape的-1改成动态推导。另外检查下输入节点的dtype和实际喂的数据类型是否一致,移动端推理时经常是float16和float32不匹配报这个错。我之前还踩过坑是ONNX默认layout是NCHW,如果你预处理时用了NHWC但没转回来,也会报形状不匹配,这个最容易忽略。
我之前也遇到过一模一样的情况,折腾了两天才发现是onnxruntime的session配置问题。你导出时只设了input的dynamic_axes,但模型中间层的tensor shape其实也被ONNX固定下来了,特别是ResNet里那些reshape和flatten操作,它们对batch维度的传播有时会隐式写死。建议你导出后用netron看一眼整个图,重点检查第一个卷积和最后的全连接层之间有没有出现shape为固定值的节点。另外BatchNorm本身是支持动态batch的,这个倒不用担心,真正坑人的往往是onnxruntime的优化器,它默认会做图融合,可能把动态shape的元信息给丢掉了,你试试在推理时加上session.set_optimization_level(ORT_DISABLE_ALL)或者用onnxruntime的dynamic shape示例里的那种方式重新初始化session。还有个小技巧,导出时把opset调高到13或14,有些老版本的算子对动态shape支持不完整。我最后是把模型里所有reshape都换成了reshape_v2或者用expand_dims加squeeze绕过去,才彻底解决。你检查下是不是模型里用了view或者flatten这类对shape敏感的操作,如果是的话,改写成adaptive_avg_pool后再接全连接层就行。
试试把dynamic_axes里output的维度也加上,光设input有时候导出会漏。