最近在把训练好的一个图像分类模型(ResNet50)转成ONNX,部署到移动端用。按照文档设了dynamic_axes={‘input’: {0: ‘batch_size’}},导出也成功了。但用onnxruntime推理时,只要batch_size不是1就报错,说“输入形状不匹配”?试了opset版本11和12都一样。是不是导出时少设了什么参数,比如输入尺寸的dynamic_axes必须和模型内部结构对应?还是说模型里有不支持动态batch的层(比如BatchNorm)?求有经验的大佬指点一下,我实在不想把batch固定为1。
PyTorch模型转ONNX时,动态轴设了但推理报错,有老哥踩过坑吗?
全部回复
共 165 条大概率是ResNet里的全局平均池化把动态轴吃掉了,试试把adaptive_avg_pool换成固定尺寸或者检查下输入张量有没有显式声明动态维度。
之前也踩过,光设dynamic_axes不够,得确认下模型里有没有reshape或全连接层把batch写死。
要不试试导出时把input的shape直接写成[None,3,224,224],onnxruntime一般能认。
我之前也卡过这问题,后来发现光设dynamic_axes还不够,onnxruntime那边跑的时候也得显式指定输入的实际shape,比如用ort的IOBinding或者直接传个带batch维度的numpy数组,不然它默认按静态图来。另外ResNet50里的BatchNorm本身是支持动态batch的,问题多半出在reshape或者view这类操作上,导出时没把那些层也标成动态。你可以先用onnx的shape inference检查下中间节点的输出维度,看是不是某个地方把batch给写死了。实在不行试试把模型里所有reshape都改成onnx::Reshape,有时候pytorch导出来会带些奇怪的常量。
大概率是ResNet里的BatchNorm在动态batch下跑不动,试试把opset拉到13以上,或者检查下输入shape有没有写全。
这问题我熟,之前转YOLOv5的时候也卡在这。动态轴设了不代表所有层都自动跟着变,尤其是ResNet里的BatchNorm,虽然它本身是per-channel的,但ONNX导出时如果某个分支的shape推断被固定了,实际跑起来就会拿静态shape去校验。你可以先试试用onnxruntime的symbolic shape inference,或者直接打印一下导出模型的输入输出shape,看是不是input的维度2和3也被写死了。另外,你只设了batch_size动态,但onnxruntime对动态shape的优化有时会要求整个图都支持动态,包括中间的reshape和pooling,ResNet50里刚好有GlobalAveragePooling,这个层对动态batch是友好的,但前面的卷积如果用了padding=‘same’这种,可能就有坑。我建议你把opset升到13+,然后把dynamic_axes里所有用到的维度都显式写出来,比如{0: ‘batch_size’, 2: ‘height’, 3: ‘width’},虽然你只要batch动态,但ONNX的shape推理有时会扩大范围。还不行的话,检查一下预处理阶段有没有固定尺寸的resize,或者干脆导出前把模型里的BatchNorm折叠进卷积,用torch的fuse_module试试,我那次就是这么解决的。
大概率是onnxruntime的输入shape没跟着动态轴走,试试显式设下input的shape再跑。
遇到过类似的,检查下预处理是不是把batch维度写死了,固定成1了。
我之前也卡过这问题,多半不是BatchNorm的锅,你查一下模型里有没有view或者reshape把batch维写死成1的操作,比如flatten那步。另外动态轴不止要设input,output那维也得对应设上,不然推理时输出形状固定了也会报不匹配。你可以用onnxruntime的shape inference脚本先跑一遍,看中间张量哪些维度被固定了,一目了然。实在不行就试下opset 13,某些算子对动态shape支持更友好。
动态轴其实只影响graph的输入输出声明,onnxruntime在推理时对中间tensor的形状推断是静态的,你ResNet50里如果有reshape或者flatten这种操作,batch维度很容易被写死,建议用onnx-simplifier过一遍看看图里有没有hardcode形状的节点。另外检查下有没有用torch.onnx.export时漏了把input的样例数据设成带batch维度的,有些层会直接根据样例输入的shape来固化逻辑。我之前转unet也碰到过类似问题,后来发现是export时忘了设do_constant_folding=False,某些折叠操作把动态维吞了。你要是方便的话可以把导出的onnx丢到netron里看下输入输出的shape是不是标了动态符号,如果显示的是固定数字那就是导出阶段的问题。
我之前也卡过这问题,后来发现多半是onnxruntime的session选项没开动态shape,得显式设置providers参数或者用ort的session_options里那个enable_all_available_optimizers,不然它默认按静态图优化了。另外你检查下ResNet50里的adaptive avg pool或者reshape操作,这类层对动态batch支持偶尔会有坑,可以试着把模型输入固定成[None,3,224,224]再导一次看看,有时候是shape推断的锅。
我之前也卡过这问题,多半不是BatchNorm的锅,那玩意儿本身是支持动态batch的。你试试导出时把模型的forward里所有reshape和view操作都检查一遍,尤其是用了-1的地方,很可能某个reshape把batch维度写死了。另外onnxruntime的session选项里有个optimization level,有时候默认优化会基于静态shape做图变换,你把它设成ORT_DISABLE_ALL再跑跑看。如果还不行,就写个几行脚本用onnxruntime的shape inference工具过一遍图,看哪一层输出的batch维变成具体数字了。
检查下输入数据是不是没把batch维传进去,onnxruntime有时得显式给足四维形状。
我之前也卡过这问题,多半不是BatchNorm的锅,你试试把dynamic_axes里input和output都配上,只设input有时候onnxruntime会拿不到完整的动态图信息。另外检查下预处理或者dataloader里有没有硬编码成固定shape,要是模型里有个Reshape或者view写死了维度,导出时不会报错但推理就炸了。我之前是加了个flatten然后手动reshape成固定尺寸,改成-1就好了。还有,如果移动端用,建议直接测下onnxruntime的CPU线程数,有时候报错是内存布局的问题,换下执行模式可能就过了。
我之前也卡过这个坑,大概率不是BatchNorm的问题,而是模型里有些reshape或者view操作把batch维度写死了。你检查下导出前的模型,看有没有用固定shape的flatten,或者全连接层前用了reshape(-1, 2048)这种,动态轴只对输入输出生效,中间层硬编码了尺寸就会炸。另外可以试试把opset升到13以上,有些老版本对动态shape支持有bug,不过最稳的办法还是用onnx-simplifier过一遍模型,能自动修正不少这类问题。
之前也踩过类似的坑,问题多半不在BatchNorm,而是ResNet里的adaptive avg pooling,PyTorch导出时对动态尺寸的算子支持有bug,尤其是opset 12以下。你试试把模型改成固定输入尺寸的avg pool,或者干脆升级opset到13+,同时用torch.onnx.export的dynamic_axes把input和output都标上,别只标input。另外,检查下onnxruntime是不是用的CPU版本,有些移动端推理引擎对动态shape支持很烂,建议先导出时用torch.jit.trace固定shape,再手动改onnx的batch维度,这样更可控。
我之前也卡过这问题,后来发现是onnxruntime的session配置里没开动态shape优化,光设dynamic_axes不够,还得用session.set_providers或者加个graph optimization level才行。另外你试试把输入张量用np.reshape成(batch, 3, 224, 224)再喂进去,有时候是数据维度没对齐。BatchNorm本身对动态batch没影响,倒是resize或全连接层容易出幺蛾子,你检查下导出前的模型有没有用torch.nn.functional.interpolate这类操作。
试试把dynamic_axes里input和output都配上,只配input的话某些节点会固化形状。另外检查下有没有reshape或view操作,那玩意儿最容易把动态batch搞崩。
我之前也栽过这个坑,问题大概率不在dynamic_axes本身,而是onnxruntime的session options里没开动态形状支持,得显式设一下session_options.add_free_dimension_override_by_name或者用onnxruntime.transformers优化图。另外你检查下导出时有没有把input的shape写成None而不是具体的[1,3,224,224],有时候PyTorch的fake tensor会把维度固定死。BatchNorm本身是支持动态batch的,不用太担心,倒是ResNet里的adaptive avgpool在ONNX里可能会被展开成固定shape,这个值得查一下。实在不行可以试试把onnxruntime升级到最新版,老版本对动态轴支持确实有bug。
我之前也卡过这问题,后来发现是ResNet里的AdaptiveAvgPool2d在ONNX里会固定输出尺寸,动态batch反而把特征图维度也搞乱了。你试试把模型的输入改成形状[batch, 3, 224, 224]然后导出前先用dummy input跑一次forward,再设dynamic_axes,有时候onnx的shape inference会漏掉一些中间节点。另外检查下onnxruntime是不是用的最新版,老版本对动态shape支持确实有bug,我换到1.16之后就正常了。如果还不行,可以试下用onnx-simplifier处理下模型,可能能去掉一些冗余的shape断言。
我之前也卡过这问题,多半不是BatchNorm的锅,ResNet50里BN在推理时是固定参数的,对动态batch没影响。你试试导出时把dynamic_axes同时加到output上,有些情况下ONNX只标了输入没标输出,推理时输出张量形状还是按固定batch算的,就会报不匹配。另外确认下onnxruntime是不是用的最新版本,老版本对动态轴支持有点bug,我升级到1.16之后就好了。要是还不行,可以把导出的模型用onnx.checker检查一下,有时候模型结构里会有隐式的reshape把batch维度写死,那种就得手动改图了。
我之前也碰到过一模一样的情况,折腾半天发现是onnxruntime的session配置里没开动态shape支持,光设dynamic_axes还不够,得在InferenceSession里传个providers参数或者设置execution_mode,你试试加上这个。另外ResNet50里的BatchNorm在ONNX里应该没问题,但如果你用了torch的F.interpolate或者某些view操作,导出的图可能把batch维写死了,建议用onnxsim简化一下模型再跑。我当时就是简化完就正常了,你可以先导出后看看输入节点的shape是不是带问号,不带的话八成是模型内部有层把维度固定了。