ONNX Runtime输出动态轴问题:批量预测时输出形状不匹配
问题分析与解决
问题根源
你的模型forward函数里使用了无参数的self.output(x).squeeze(),这会引发形状匹配问题:
- 单张输入(形状
(1,3,255,255))时,self.output(x)输出形状为(1,1),squeeze()后变成标量(形状()) - 批量输入(形状
(2,3,255,255))时,self.output(x)输出形状为(2,1),squeeze()后变成(2,)
但导出ONNX时你用的是单张dummy_img,ONNX会将输出的静态形状固化为标量,即便配置了动态轴,也无法适配批量输入时的输出形状变化,最终触发形状不匹配报错。
解决方案
方案1:修改模型forward,指定squeeze维度
将return self.output(x).squeeze()改为return self.output(x).squeeze(dim=1),明确只压缩通道维度(假设self.output是输出维度为1的线性层,输出形状为(batch_size,1))。这样无论单张还是批量输入,输出形状都是(batch_size,),ONNX能正确识别动态batch轴。
修改后重新导出ONNX,推荐用批量dummy tensor(比如形状(2,3,255,255)),让ONNX更准确捕捉动态轴信息:
loaded_model.eval() # 构造批量dummy输入 dummy_img = torch.randn(2, 3, 255, 255) torch.onnx.export(loaded_model, dummy_img, "trained_model_3.onnx", export_params=True, do_constant_folding=True, verbose=False, dynamic_axes={'input' : {0 : 'batch_size'}, 'output' : {0 : 'batch_size'}}, input_names=input_names, output_names=output_names)
方案2:导出时强制使用批量dummy输入(不修改模型)
如果不想改动模型代码,导出ONNX时必须用批量dummy_img(比如batch_size=2),让ONNX将输出形状识别为(2,),确保动态轴配置生效。但这种方案存在风险:后续输入batch_size变化时,无参数squeeze()可能仍会引发形状异常,不如方案1稳妥。
验证
修改后重新导出ONNX,用ONNX Runtime加载批量输入(形状(2,3,255,255)),此时输出形状应为(2,),不会再触发形状不匹配错误。
内容的提问来源于stack exchange,提问作者Chema
相关产品推荐
相关产品推荐

