最近在把训练好的一个图像分类模型(ResNet50)转成ONNX,部署到移动端用。按照文档设了dynamic_axes={‘input’: {0: ‘batch_size’}},导出也成功了。但用onnxruntime推理时,只要batch_size不是1就报错,说“输入形状不匹配”?试了opset版本11和12都一样。是不是导出时少设了什么参数,比如输入尺寸的dynamic_axes必须和模型内部结构对应?还是说模型里有不支持动态batch的层(比如BatchNorm)?求有经验的大佬指点一下,我实在不想把batch固定为1。
PyTorch模型转ONNX时,动态轴设了但推理报错,有老哥踩过坑吗?
全部回复
共 165 条我之前也遇到过一模一样的情况,onnxruntime对动态轴的支持其实挺挑的,尤其是ResNet这种带全局池化的结构,batch维度在中间层可能会被隐式固定住。你试试把dynamic_axes同时加到输出上,只设输入有时候会漏掉shape推断的依赖。另外检查下有没有用torch.onnx.export的dynamic_axes参数和input_names/output_names对应上,还有BatchNorm本身没问题,但如果你用了view或者reshape把batch维度硬编码了,那导出时得手动改改。实在不行可以先用onnx-simplifier跑一遍,有时候能帮你把冗余的shape约束清掉。
我遇到过类似的坑,问题大概率不在dynamic_axes本身,而是模型里有些层对shape是硬编码的。你检查下ResNet50里有没有用view或者reshape的地方,特别是global average pooling之后,有些实现会写死flatten后的维度,比如torch.flatten(x, 1)本身没问题,但如果你在forward里用了x.size(0)去构造别的tensor,那动态batch就会炸。另外,BatchNorm在推理模式下是没问题的,它只对channel维做归一化,batch维完全无关,所以基本可以排除。我建议你先把模型单独导出一次,用onnxruntime的shape inference工具看看中间节点的输出shape是不是都带动态维度,如果某些节点变成了固定值,那就得手动改模型代码。还有一个取巧的办法,就是导出时把batch维设成一个比较大的固定值比如8,然后推理时用pad到8再截断,虽然浪费点内存但能绕过这个问题。不过最好还是查下是不是torch.onnx.export里没设dynamic_axes的input_names和output_names的对应关系,我之前就是漏了output的dynamic_axes导致推理时输出维度对不上。你试试看能不能在导出后打印下onnx图的输入输出shape,如果输入是动态的但输出是静态的,那肯定是输出轴没设对。
我之前也遇到过一模一样的坑,后来发现不是dynamic_axes的问题,而是ResNet里自适应池化层在ONNX导出时把空间维度给固定了。你试试导出前把模型里的AdaptiveAvgPool2d换成固定kernel size的AvgPool2d,或者干脆在forward里用F.adaptive_avg_pool2d然后加个torch.onnx.export的input_names和output_names都配上,说不定能解决。另外检查下onnxruntime版本,太老的版本对动态shape支持确实有bug,升到1.16以上再试试看。
碰到过类似情况,当时也纠结了很久。你这个问题大概率不是BatchNorm的锅,ResNet50里BN在导出时是能跟着动态轴走的,更多可能是onnxruntime的session配置问题。你试试在创建InferenceSession的时候,显式设置providers参数,或者检查一下输入数据的维度是不是真的传成了[N, C, H, W]而不是反了。另外,动态轴设了不代表所有中间层的shape都能自动推导,有些算子比如Reshape或者Flatten会把batch维写死,建议用onnxruntime的shape inference工具跑一遍,看看中间张量的维度是不是还是固定值。我之前试过用onnxsim简化模型,有时候能解决这种隐藏的shape硬编码问题。还有个小细节,你导出时是不是用了torch.onnx.export的input_names和output_names?如果没给全,某些分支的输入可能没绑上动态轴。实在不行,先试着把opset调高到13或14,有些旧版本对动态shape支持确实有bug。最后如果还不行,干脆导出时把batch设成1,推理时用多个session或者手动循环,虽然笨但稳。
遇到过一模一样的情况,ResNet50转ONNX动态轴导出成功但推理必炸,后来排查发现不是BatchNorm的问题,是onnxruntime的C++ API里输入输出张量维度写死了。你导出时dynamic_axes只设了input的batch维度,但模型内部很多Reshape和View操作会自动把batch维度当成静态的,导出的图里这些节点对动态shape支持不友好,尤其是opset 11之后有些算子对动态维度有额外约束。
我当时的解决办法是把模型的输入输出都显式声明为["batch", 3, 224, 224],并且用torch.onnx.export时加上dynamic_axes同时指定input和output,关键是要对模型做一次trace后再导出,别直接用原始模型。另外你可以先用onnxruntime的Python接口试一下,如果Python下动态batch能用,基本就能确定是C++端shape没设置对。
还有个坑是onnxruntime的session options里有个execution_mode,默认是sequential,如果开parallel模式有些算子会要求固定batch。我最后是直接把导出时的opset降到10,动态batch就正常了,虽然牺牲一点算子兼容性但移动端够用。你试试先排除是不是推理代码的问题,用netron看看导出的图里有没有什么奇怪的Flatten或Gather节点。
如果实在不行,干脆转成TensorRT或者用OpenVINO,动态batch支持比ONNX Runtime成熟很多,特别是N卡上速度还快。别死磕这一个,我折腾了两天才发现是C++端shape初始化的问题,气死。
大概率是模型里有reshape或全连接层吃死了batch维度,导出前把模型改成自适应shape试试。
也可能是onnxruntime的session配置里没开动态shape,设个providers参数或者用onnxruntime.transformers优化下看看。
你这多半是模型里插了reshape或view,把batch维和通道维写死了,查下导出前的forward有没有动态shape操作。
遇到过类似情况,我当时也是ResNet系列,折腾了半天发现问题不在dynamic_axes本身,而在onnxruntime的session配置上。你导出时只设了input的dynamic_axes,但onnxruntime默认会用固定shape的优化策略,尤其是开启图优化后,某些融合算子会把batch维度写死,导致推理时shape不匹配。你可以试试在session初始化时加一个参数,比如sess_options.graph_optimization_level = ORT_ENABLE_BASIC,或者干脆把优化关掉看还报不报错,先定位是不是这个原因。
另外一个坑是,虽然你只设了input的dynamic_axes,但模型内部如果有Reshape或者Flatten操作,它们的输出shape可能没有跟着动态化。ONNX导出时,dynamic_axes是逐层传递的,但如果中间某些层用了固定尺寸的常量,比如view里的-1依赖具体batch数,导出器可能不会自动帮你去推断。建议你用onnxruntime的shape inference工具检查一下中间张量,或者用netron看看某个节点的输出shape是不是还是[1, 2048, 7, 7]这种固定值。
还有,BatchNorm本身是支持动态batch的,它只对每个样本独立做归一化,不依赖batch内统计量,所以问题大概率不在这。你可以试一下用onnxruntime的C++ API或者python API分别跑,有时候python端做了数据预处理会意外改变形状。另外,你确认一下输入数据是不是真的按[N,C,H,W]排的,我之前有个同事就是把H和W顺序搞反了,batch>1时直接越界报错,还以为是动态轴的问题。
最后实在不行,可以试一下改用opset 13以上,或者用torch.onnx.export时加上dynamic_axes的同时,也把output的dynamic_axes显式声明出来,有时候输出shape没动态化也会导致runtime内部校验失败。我那次后来是发现有个Gemm层因为输入维度被固定,导致输出特征图尺寸对不上,手动把那个层的权重reshape了一下才解决。
我之前也卡过这问题,后来发现是onnxruntime的session options里没开动态形状支持,得设一下optimized_model_filepath或者干脆用ort的transformers优化接口。另外你检查下模型里有没有reshape或者view操作写死了batch维度,ResNet50本身没这个问题,但如果你改了分类头就可能引入。还有个坑是导出时虽然设了dynamic_axes,但输入tensor的shape hint如果写死了,runtime还是会按静态处理。建议你把输入样例的shape改成[1,3,224,224]再导出,然后推理时用ort的IOBinding指定实际shape试试。
动态轴设了但推理报错,八成是模型里有些层对shape是硬编码的,比如reshape或者view用了固定batch的写法,ONNX导出时虽然标了动态,但实际图里已经写死了。你可以用onnxruntime的shape inference工具看下中间节点的输出维度,或者把模型导出后用onnx-simplifier过一遍,很多奇葩问题都是这玩意儿解决的。另外BatchNorm本身是支持动态batch的,基本可以排除。我之前遇到过类似情况,最后发现是数据预处理那边把输入张量固定成1了,检查下你的推理代码,别光盯着模型文件。
遇到过一模一样的情况,最后发现问题不在dynamic_axes本身,而是模型里有个自定义的reshape或者view操作把batch维写死了。ResNet50本身是支持动态batch的,你trace模型的时候如果用了固定shape的输入,有些算子会把shape当成常量固化下来,ONNX导出时虽然标了动态轴,但内部那些依赖shape的计算图逻辑还是按1来处理的。建议先用onnx.checker检查一下,再看下导出的onnx里shape信息是不是都变成动态了,尤其是GlobalAveragePooling后面接Flatten或者Reshape的地方,很容易踩坑。另外BatchNorm不会限制动态batch,它只是在channel维度上做归一化,batch维是自由的,所以别怀疑它。如果实在找不到具体是哪一层,可以试试用torch.onnx.export的input_names和output_names都加上,然后配合dynamic_axes把每个中间张量的batch维都声明成动态,有时候只设输入输出不够,中间某些节点的shape推断会出问题。最后实在不行,就换用onnx-simplifier处理一下,很多隐式的shape硬编码能被它清掉,我之前就是这么解决的。
我之前也撞到过这堵墙,搞了三天才发现问题不在dynamic_axes本身。你设的那个参数只告诉了ONNX图“这个维度可以变”,但模型内部的Reshape、Gather或者全连接层如果写死了形状,导出器不会自动帮你改。你试下用onnxruntime的session.get_inputs()打印一下实际输入要求,看是不是除了batch_size还有别的维度被固定了,比如某些中间tensor的shape是从输入硬编码推断的。另外BatchNorm在推理模式下是安全的,它只按通道缩放,不会限制batch,真正坑的往往是模型里手动reshape成固定shape的代码,或者opset版本对动态shape的支持差异——建议你导出时把dynamic_axes同时写到output上,然后跑一遍onnx.checker.check_model,再不行就用onnx-simplifier把图里那些显式shape操作替换成动态版本。还有个野路子,就是把输入改成NCHW的N用-1占位,然后推理前手动reshape成实际batch,有些老版本runtime会认这个。要是还不行,直接看onnxruntime的verbose日志,它会告诉你哪一层shape冲突。
试过把dynamic_axes里的batch_size同时加到output上吗?我之前也卡在这,光设input的dynamic没用,输出张量的形状也得跟着变,不然推理时内部图优化会把固定shape传下去。另外检查下模型里有没有reshape或者view操作,它们可能把batch维度写死了,ResNet50本身倒不至于被BatchNorm卡住。实在不行可以先转成静态batch=1跑通,再用onnx-simplifier看看图结构哪里被固化了。
大概率是模型里有个reshape或view把batch写死了,查下导出前的forward能不能过batch=2的trace。
试试把dynamic_axes里output也加上,只设input有时候onnxruntime会默认输出形状固定。
我之前也卡过这个坑,问题大概率不在dynamic_axes本身,而是ONNX导出时输入张量的实际shape写死了。你试试导出时把input的样例数据设成类似[1,3,224,224],同时检查下有没有用torch.onnx.export的input_names和output_names参数,并且确认onnxruntime的session选项里有没有开graph optimization,有时候优化会强行固定batch维度。
另外BatchNorm在推理模式下是没问题的,动态batch报错更可能是模型里其他reshape或者view操作把维度写死了。你可以先用onnx-simplifier跑一遍,再用onnxruntime的shape inference看看中间节点的输出维度是不是都带动态标记,我之前就是这么排查出来的。
我之前也踩过一模一样的坑,折腾了整整两天。你动态轴设的其实没问题,问题大概率出在onnxruntime的session配置上,推理的时候要显式地用ort.set_inputs_outputs_dynamic或者直接给输入数据传一个shape为(N, C, H, W)的tensor,光设dynamic_axes不够。另外ResNet50里的BatchNorm在ONNX里通常被融合进Conv了,一般不会成为动态batch的阻碍,但如果你用了adaptive_avg_pool或者某些自定义层,那确实可能卡住shape推断。我建议你导出后先用onnxruntime的onnxruntime.tools.check_onnx_model跑一下,或者用netron看看图里有没有显式shape的节点。还有一个坑是opset版本,我后来换到13才稳,特别是如果你用了torch的某些新算子,11和12可能没完全支持。实在不行就试下导出时把dynamic_axes的input和output都写上,output那边也得标batch维度,有时候只标input会导致IR版本不一致。最后如果还不行,干脆导出前把模型forward里所有reshape和view都改成用-1推断batch,有些框架内部写死了维度。
我之前也卡过这个坑,大概率不是BatchNorm的问题,你看下onnxruntime的版本,老版本对动态shape支持有点迷,换个1.16以上的试试。另外你导出时只设了input的dynamic_axes,但output的维度也得跟着设上,不然模型内部传递时维度信息会锁定成静态的,推理时batch一变就炸。还有个检查方法,用onnx.shape_inference跑一遍,看中间tensor的shape是不是都带动态符号,如果某层还是写死的数字,那就是那层之前没接住动态信息,得手动改下导出脚本里的input_names和output_names对应关系。
我之前也卡过这个坑,问题多半不在dynamic_axes本身,而是onnxruntime的session options里没开动态shape支持,要设置providers参数或者把graph optimization level调低试试。另外ResNet50里的BatchNorm在导出时如果用了training模式,有可能把running_mean这些当常量固化,导致维度写死,建议确认一下model.eval()之后再转。还有个更省事的办法,导出前直接用torch.onnx.export的input_names和output_names把batch维度也明确标出来,然后推理时用ort的IO Binding方式传入,能绕过不少检查。如果还不行,可以看看onnxruntime的日志,它会提示具体是哪一层不匹配。
大概率是模型里写了固定维度的reshape或view,查下ResNet的avgpool后面有没有展平操作,改成自适应维度就行。
试试把输入和输出的dynamic_axes都写上,光设input有时候onnxruntime不认。