最近在把训练好的一个图像分类模型(ResNet50)转成ONNX,部署到移动端用。按照文档设了dynamic_axes={‘input’: {0: ‘batch_size’}},导出也成功了。但用onnxruntime推理时,只要batch_size不是1就报错,说“输入形状不匹配”?试了opset版本11和12都一样。是不是导出时少设了什么参数,比如输入尺寸的dynamic_axes必须和模型内部结构对应?还是说模型里有不支持动态batch的层(比如BatchNorm)?求有经验的大佬指点一下,我实在不想把batch固定为1。
PyTorch模型转ONNX时,动态轴设了但推理报错,有老哥踩过坑吗?
全部回复
共 165 条我遇到过类似情况,最后发现是模型里有个reshape操作把batch维度写死了,导出时动态轴没传进去。你试试用onnx-simplifier简化一下模型,或者用netron看看中间节点的shape,说不定能定位问题。另外BatchNorm在推理时一般不依赖batch大小,我感觉不是它的锅。
试过把input的batch维度直接设成-1或者None吗?有时候dynamic_axes的写法跟实际推理时的输入形状对不上。另外ResNet50里的BatchNorm层本身是支持动态batch的,问题可能出在全连接层或者reshape操作上——如果模型里用了某些硬编码尺寸的view或reshape,动态轴就会出问题。建议用onnx-simplifier优化一下模型,再检查下算子兼容性。
这个问题我去年也折腾过一阵,动态batch看着简单但坑确实不少。你opset版本没问题,不过光设dynamic_axes还不够,得确认模型里所有reshape和view操作是不是都用了动态维度,ResNet50本身没有显式的batch size依赖,但torch的某些算子比如flatten或adaptive_avg_pool2d在导出时可能会把batch写死。我建议你导出时加个verbose=True看看onnx图里有没有形状被硬编码成1的节点,另外试试用torch.onnx.export时显式指定一下input的size为(N,3,224,224),N设成-1或一个具体值,同时把dynamic_axes写全,比如input和output都要配。还有个常见问题是BatchNorm在eval模式下确实不影响动态batch,但如果你模型里混了某些自定义层或者用了F.interpolate这类函数,可能就得手动改onnx图了。实在不行可以先转成静态batch用onnx-simplifier优化一遍,再手动改动态轴,或者试试用onnxruntime的IOBinding接口直接绑定可变输入形状。
这个坑我确实也踩过,感觉问题大概率不在dynamic_axes本身,而是出在模型内部的某些操作上。ResNet50里如果有AdaptiveAvgPool或者reshape这类层,它们的实现可能对动态batch不友好,尤其是当某些算子在ONNX导出时被拆成了固定形状的中间节点。你可以先用onnx-simplifier把模型简化一下,看看是不是有些节点形状被写死了。另外BatchNorm其实对动态batch是支持的,它只是统计running mean和var,跟batch大小没关系,所以这个可以排除。我建议你导出的时候加上dynamic_axes的同时,也检查一下模型的输入输出是否真的都绑定了动态轴,有时候只设了input但output没设也会导致推理时形状冲突。还有一招比较笨但有效:用onnxruntime直接打印一下报错时的输入和模型期望的形状,对比一下具体是哪个维度对不上。如果还不行,可以试试把opset升到13或15,有些动态轴相关的bug在旧版本里修得不够彻底。
试试把input和output的dynamic_axes都显式设上,另外检查下模型里有没有reshape或flatten操作。
你这报错我熟,光设dynamic_axes还不够,onnxruntime那边也得显式声明输入的实际shape,比如用ort的IOBinding或者session.run的input_feed里传进去的tensor维度必须和动态轴匹配上。另外ResNet50里的BatchNorm在导出时通常是折叠进卷积的,一般不会卡动态batch,倒是检查下模型里有没有reshape或者view操作写死了batch维度,那才是真坑。我之前就遇到过某层reshape把batch硬编码成了1,导出不报错但一变batch就炸,你得用onnxruntime的profiler或者逐个节点排查。要是懒得查,直接导出时把opset拉到13以上,有些老版本bug会自动修复。
检查下onnxruntime的session输入,动态轴写对了但实际得传input_shape,用onnxsimplifier固化下试试。
我之前也栽在过这上面,大概率不是BatchNorm的问题,而是onnxruntime的session配置里没开动态shape的优化选项。你试试在InferenceSession创建时加上providers参数,并且用session.set_input_shapes手动指定一下动态维度,我之前这么搞定的。另外你导出的onnx里input的shape是不是还是固定的?用onnx.checker验证一下,有时候动态轴没真正写进graph里。实在不行就换个思路,导出时把batch维度设成None,然后推理时用ort的IOBinding接口传数据,别用run的输入列表。
试试把input和output的dynamic_axes都写上,只写input有时会漏掉中间张量的shape推断。
我之前也卡过这个坑,动态轴导出成功只是第一步,推理端用onnxruntime时输入张量必须显式指定成动态shape,比如用np.newaxis或者reshape一下,否则它默认按静态shape跑。另外检查下模型里有没有GlobalAveragePooling或者Flatten这类会隐式依赖batch维的层,ResNet50一般没这问题,但保险起见可以用onnxruntime的symbolic shape debug工具看下中间节点shape。还有个小技巧,试试把opset升到13或者15,有些老版本对动态shape支持有bug。实在不行可以先固定batch=1跑通流程,再逐个排查是哪个层限制了动态维度。
我之前也卡在这过,问题多半不在dynamic_axes本身,而是ONNX导出时输入数据没给足shape信息,试试在export里加上input_names和output_names,然后显式传一个batch=2的样例张量进去,让模型把维度关系锁死。另外BatchNorm在eval模式下是没问题的,但如果你转之前忘了切model.eval(),那跑动态batch确实会炸,先确认这个。还有个坑是onnxruntime的session选项里要开graph optimization,默认优化有时会误伤动态轴,可以设成ORT_ENABLE_EXTENDED试试。我上次是这么解决的,你排查下这几个点,应该能找到原因。
我之前也碰到过类似情况,后来发现是onnxruntime的session配置里没开动态shape,需要显式设置providers参数或者用IOBinding,光设dynamic_axes不够。另外ResNet50的BatchNorm在推理模式应该没问题,你检查下导出时是不是用了training=True?或者试下opset 13,有些算子对动态shape支持更友好。
我之前也卡过这问题,最后发现是onnxruntime的session选项里没开动态shape支持,得设session_options.add_free_dimension_override_by_name或者用onnxruntime.transformers优化一下图。另外你检查下ResNet50里的GlobalAveragePooling和Flatten,那俩在动态batch下有时会输出固定维度,导致后面全连接层对不上。还有个坑是导出时最好把input的shape写成[1,3,224,224]但dynamic_axes里同时给input和output都标上batch_size,只标输入不标输出的话推理时输出维度也是死的。你试试导出后先用onnx.checker和shape_inference跑一遍,看中间节点有没有把batch维度写死,我上次就是有个Reshape硬编码了1。
大概率是ResNet里adaptive_avg_pool把空间维固定了,试试把输入改成NCHW全动态,顺便检查下reshape有没有写死。
我之前也卡在这过,多半不是BatchNorm的问题,你试试导出的时候把input的shape写成[None,3,224,224]而不是[1,3,224,224],有些版本的torch.onnx.export对None的理解有偏差。另外检查下模型forward里有没有用x.size(0)或者view这种硬编码维度的操作,ResNet的global avg pooling之后有个flatten,如果写死batch=1就会炸。实在不行可以试下opset 13+,动态轴的处理逻辑更成熟些。
我之前也卡过这问题,多半不是BatchNorm的锅,是ONNX导出时把shape推死了。你试试导出前用torch.onnx.export的input_names和output_names参数,再把dynamic_axes里output的维度也对应写上,比如{‘output’: {0: ‘batch_size’}}。另外检查下模型forward里有没有用view或reshape固定了batch维度,ResNet的adaptive avg pool一般没事,但自定义的flatten可能要改改。我之前就是漏了output的dynamic_axes,改完就正常了。
大概率是模型里插入了固定shape的reshape或者flatten,转onnx时把batch维度写死了,建议用onnxsim简化下再导出试试。
我之前也卡过这问题,后来发现是模型里有个全连接层在作怪,导出时虽然设了动态轴,但全连接层的权重形状是固定的,batch维度没跟着变。你可以试试把input的dynamic_axes同时加到output上,或者检查下有没有用torch.nn.Flatten这种会隐式改变维度的层。另外BatchNorm本身没问题,我试过带BN的模型动态batch也能跑,多半还是导出参数没配全。实在不行就换opset 13+,新版本对动态形状支持更稳。
还有个思路,你导出前先打印下onnx图的输入输出shape,看看dynamic_axes到底生效没,有时候导出器会静默忽略掉某些节点的动态维度。我上次就是靠这招发现是reshape层把batch写死成1了,手动改下图结构就解决了。
大概率是ResNet里adaptive avgpool的flatten把batch维度写死了,检查下导出前模型的forward里有没有view成固定shape。
试试把dynamic_axes里output也加上,光设input有时候onnxruntime会漏掉维度推断。