最近在把训练好的一个图像分类模型(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里所有输入输出都配上,光设input不够,输出shape也得跟着变。
- 之前遇到过,BatchNorm没毛病,可能是模型里reshape或view把batch维度写死了。
老哥这问题我太熟了,之前做检测模型转onnx也卡了好几天。你dynamic_axes只设了input的batch维度,但onnxruntime推理时如果输入tensor的shape是[2,3,224,224],而模型内部有些节点比如reshape或者flatten是写死batch=1的,那就会报不匹配。建议你导出后用netron看一下整个图,重点找有没有shape是硬编码的节点,尤其是ResNet最后的全局池化后面接全连接那部分,有时候pytorch的adaptive_avg_pool会带出固定的batch信息。另外BatchNorm本身是支持动态batch的,这个不用担心,问题大概率出在reshape或view操作上,pytorch转onnx时这些操作有时会生成固定的shape张量。你可以试试输入一个batch=2的dummy数据导出,如果导出时shape是[2,3,224,224],那动态轴就会跟着变,但如果你导出时用的dummy是[1,3,224,224],那动态轴可能只对batch=1生效。还有个土办法,就是导出前把model.eval()加上,然后手动把模型里所有view/reshape的-1参数检查一遍,或者干脆用torch.onnx.export的input_names和output_names都加上dynamic_axes,包括输出层。要是实在搞不定,可以试下用onnx-simplifier处理一下,有时候能自动修复这些batch写死的问题。
动轴设了但batch不是1就报错,这情况太典型了,我当初也被卡过一阵。你只设了input的dynamic_axes,但ONNX的batch维度是全局的,模型中间那些reshape、flatten或者全连接层如果默认写死了第一维是1,导出时根本不会自动帮你改。ResNet50本身BatchNorm和全局池化对动态batch是没问题的,问题多半出在你那个分类头或者预处理里,比如view或者permute写死了1。建议你导出后用onnx.checker验证一下,再拿onnxruntime的session.get_inputs()和get_outputs()看一眼实际形状,顺便用onnx.shape_inference跑一遍,排查哪个节点把batch维给固定了。另外你试试把dynamic_axes同时加到output上,有时候输出层没设动态轴也会引发这种隐性问题。实在不行就转成onnx后用onnx-simplifier过一遍,它能帮你把很多冗余的固定形状节点清理掉,我之前就是这么解决掉的。
试试把dynamic_axes里所有维度的名字都设全,另外检查下模型里有没有reshape或view把batch写死了。
我之前也卡在这过,问题多半不在dynamic_axes本身,而是模型里有些层(比如reshape或者全连接前的flatten)对batch维度写死了。你导出后用netron看看计算图,找找有没有把batch_size当成固定值的节点,尤其注意ResNet50最后的avgpool和fc之间。另外onnxruntime有个优化选项,建议把graph优化级别调低试试,有时候是优化器自作主张把动态轴给折叠了。实在不行就先固定batch导出,推理时用onnxruntime的IOBinding手动改shape,比调dynamic_axes省心。
这个问题我当初也折腾过一阵,最后发现多半不是BatchNorm的锅,而是onnxruntime的session配置里没开动态shape的优化。你导出时dynamic_axes确实写了,但推理端如果没设置providers参数或者没指定session_options,有些版本默认会按静态图执行,导致batch维度被锁死。
另外一个坑是ResNet里的GlobalAveragePooling,它本身对batch是友好的,但如果你在forward里用了view或者reshape并且硬编码了维度,比如x.view(x.size(0), -1),导出时ONNX会把那个0当成常量,这样动态轴就失效了。建议先在导出前用torch.onnx.export加上input_names和output_names,然后检查一下onnx.checker,再用onnxruntime的get_providers确认用的不是CPU的旧版实现。
我最后是换成了opset 13,并且在推理时手动把输入dict的shape改成(batch, 3, 224, 224),同时用ort.SessionOptions里的add_free_dimension_override_by_name来指定动态维度,才解决问题。你可以先试试在导出时加一句dynamic_axes={'input': {0: 'batch'}, 'output': {0: 'batch'}},然后把opset_version提到13,如果还报错,大概率是模型内部某个自定义层或F.interpolate的坐标生成方式不支持,那就得检查onnx图里有没有Resize节点带scales常量了。
我之前也卡在这过,折腾半天发现大概率不是BatchNorm的问题,那玩意儿在ONNX里本来就是静态的,跟batch维度无关。你只设了input的dynamic_axes,但ONNX导出时模型内部那些中间tensor的shape信息其实是被固定下来的,尤其ResNet里那个flatten和全连接层,它们对batch维度的传播要求很严格,你可以试下把output的dynamic_axes也加上,比如{'output': {0: 'batch_size'}},有时候onnxruntime会拿输出shape去校验输入。另外检查下你输入数据是不是NHWC和NCHW搞混了,移动端经常有这种暗坑,报错信息里说的形状不匹配可能不是batch,而是channel排布问题。还有个骚操作,用onnx-simplifier跑一遍,它能帮你把那些冗余的shape操作给折叠掉,我之前有个模型就是这么治好的。如果还不行,就打印一下onnx graph里每个节点的输出shape,定位到底是哪一层开始变的,大概率是某个reshape或者gather节点写死了维度。
遇到过一模一样的坑,折腾了两天才发现是onnxruntime的session配置问题,光设dynamic_axes不够,还得在session options里把execution_mode设成ORT_PARALLEL或者给输入绑定具体的shape。另外ResNet50里的BatchNorm在eval模式转出来应该没问题,重点检查下模型里有没有reshape或者flatten操作,那些层对动态shape特别敏感,建议先用onnx-simplifier过一遍再试。
之前也踩过类似的坑,问题多半不在dynamic_axes本身,而是模型内部有reshape或者view操作把batch维写死了。ResNet50的全局池化后面那个flatten层经常搞事,建议先netron看看onnx图里batch维是不是被固定成1了。另外onnxruntime的session选项里要记得设优化级别,有时候默认优化会把动态轴给折叠掉。实在不行就导出时把input的shape写成None试试,或者用onnx-simplifier过一遍再测。
大概率是onnxruntime的输入name没对上,你导出时用torch.onnx.export的input_names重新指定下试试。
之前也卡这过,试试把reshape和flatten那几个层也加进dynamic_axes,光设输入没用。
我上次是模型里有个view写死了batch,改成-1就通了,你检查下有没有类似操作。
大概率是转出时还有个隐藏的shape没绑动态轴,试试导出前用torch.onnx.export的input_names/output_names全配上。
我之前也栽在过这个坑里,你设dynamic_axes其实没毛病,问题多半出在导出后的模型内部。ResNet50里BatchNorm在推理时是折叠成卷积的,按理说不影响动态轴,但你这个报错更像是onnxruntime那边对输入形状的约束没放开——试试在session配置里把execution_mode设成ORT_PARALLEL,或者干脆检查下导出时有没有把模型的input shape写死,有时候torch.onnx.export的input_names里如果带了固定维度,动态轴会被忽略。另外,你可以在导出后直接用onnxruntime的shape inference跑一下,看看中间节点的维度是不是也跟着动态了,如果中间层还是固定shape,那问题就出在模型里某些reshape或flatten操作把batch维写死了。我之前遇到类似情况是模型里有个自定义op,转的时候没注册,导致动态轴没传到后面去,你检查下有没有这种隐藏雷点。要不你先把导出时的fixed_batch_size参数显式设为-1试试?我那次就是这么解决的,虽然不知道原理,但确实管用。
我之前也卡过这个问题,导出成功不代表动态轴真的对所有层生效。你检查下模型里有没有reshape或者view操作,这类层经常把动态batch写死,尤其是ResNet里的adaptive avg pool后面,建议把输入shape写成(1,3,224,224)导出,再用onnx-simplifier处理一遍。另外onnxruntime的session选项里有个optimization level,默认开优化有时会破坏动态轴,试试设成ORT_DISABLE_ALL,我上次就是关掉优化才跑通的。如果还不行,把导出的onnx用netron打开看下input的shape是不是真的带动态维度,有时候只是导出时没打印出来而已。
遇到过类似的坑,问题大概率不是BatchNorm,而是你导出时只设了input的dynamic_axes,但模型中间某些reshape或全连接层对batch维度是写死的,比如view成固定形状。建议用onnx.helper打印一下图结构,重点看ResNet最后的全局池化后面有没有把batch硬编码进shape的节点,有的话得手动改onnx图或者用onnxsim简化一下。另外onnxruntime的session选项里可以试试开启动态shape优化,有时候是执行provider没选对,比如CPU上用默认的CPUExecutionProvider对动态batch支持不好,换CoreML或TensorRT试试可能就好了。还有个土办法,导出时把batch维设成symbolic的字符串,比如‘batch’,别用0,然后推理前用onnxruntime的IOBinding重新绑定输入输出,能绕开一些形状检查的bug。
动态轴设了但推理报错,八成是导出时没把输入的真实shape写全,比如只标了batch维度,但onnxruntime内部还是按固定shape去优化了,你试试导出前把input的size设成[None,3,224,224]再转一次。另外ResNet50里的BatchNorm在eval模式下是没问题的,但如果你转的时候模型还在train模式,那些running_mean/var会被当成动态计算,换个思路,转之前务必model.eval()一下,我上次就是栽在这上面。还有个小坑,如果用了torch.onnx.export里的dynamic_axes,最好把output也一起标上,比如{‘output’: {0: ‘batch_size’}},不然某些runtime会拿固定输出shape去反推输入。实在不行你可以先试下固定batch=4导出,然后推理时用onnxruntime的IOBinding手动改shape,绕开这个校验,但那样移动端部署就不太方便了。
之前也栽在过这上面,光设dynamic_axes不够,得检查下模型里有没有reshape或者view依赖固定shape,ResNet50的池化层一般没问题,但如果你改过前向逻辑就难说。另外你推理时输入得是np数组吧,onnxruntime要求输入形状完全匹配动态轴的默认值,试试在session里显式指定输入名称的shape,或者把dynamic_axes的input也加上3和2两个维度。还有个坑是opset版本太高或太低对动态轴支持不同,11和12都不行的话,试试13。
我之前也卡在这过,问题多半不在dynamic_axes本身,而是模型里reshape或者view这类操作把维度写死了,导出时虽然没报错,但实际推理走的还是固定shape。你可以先试下用onnx-simplifier把图精简一遍,再检查下有没有hard-coded的shape常量,另外确认下输入节点名字跟dynamic_axes里的key对得上,有时候就差这一步白折腾半天。
遇到过,动态轴设了但推理报错,大概率是onnxruntime的session配置里没开dynamic shapes支持,光靠导出参数不够。你试下在InferenceSession里加个providers参数,或者用onnxruntime的dynamic_axes图优化选项,具体我记不太清API名了。另外检查下模型里有没有reshape或者flatten这类硬编码维度的层,ResNet50理论上没这问题,但保险起见用onnx工具看下中间节点的shape是不是都带batch维。我之前是卡在全连接层前的自适应池化上,改成全局平均池化就好了。
我之前也遇到过一模一样的情况,最后发现是onnxruntime的session配置问题,默认会用静态输入shape去优化图,得在session里显式设置providers和optimized_model_filepath,或者用onnxruntime.transformers优化时关掉某些融合选项。另外你试试把dynamic_axes里input和output都写上,只写input有时候导出器会漏掉内部tensor的shape推断。还有,BatchNorm本身是支持动态batch的,问题大概率出在ResNet最后的全局池化或者全连接层对维度写死了,建议导出后用netron看一眼每个节点的输出shape。要是还不行,可以试试固定batch导出然后推理时用reshape,虽然丑但能用。