PyTorch转ONNX遇custom::deform_conv2d形状推理缺失警告求解决
解决PyTorch转ONNX时DeformConv2d的形状推理警告
你遇到的警告是因为ONNX无法自动推断自定义custom::deform_conv2d算子的输出形状,必须在自定义的symbolic函数中手动添加形状计算和设置逻辑。现有symbolic函数仅定义了算子调用,未处理形状推理,导致ONNX无法确认输出维度。
以下是修改后的完整代码,可消除该警告:
import torch import torchvision class Model(torch.nn.Module): def __init__(self): super().__init__() self.conv1 = torch.nn.Conv2d(3, 18, 3) self.conv2 = torchvision.ops.DeformConv2d(3, 3, 3) def forward(self, x): return self.conv2(x, self.conv1(x)) from torch.onnx import register_custom_op_symbolic from torch.onnx.symbolic_helper import parse_args @parse_args("v", "v", "v", "v", "v", "i", "i", "i", "i", "i", "i", "i", "i", "none") def symbolic(g, input, weight, offset, mask, bias, stride_h, stride_w, pad_h, pad_w, dil_h, dil_w, n_weight_grps, n_offset_grps, use_mask): # 获取输入张量的形状信息 input_shape = input.type().sizes() batch_size, _, in_h, in_w = input_shape # 从权重张量提取输出通道数和卷积核尺寸 weight_shape = weight.type().sizes() out_channels, _, kernel_h, kernel_w = weight_shape # 按卷积公式计算输出特征图的高宽 out_h = ((in_h + 2 * pad_h - dil_h * (kernel_h - 1) - 1) // stride_h) + 1 out_w = ((in_w + 2 * pad_w - dil_w * (kernel_w - 1) - 1) // stride_w) + 1 # 构造完整的输出形状 output_shape = (batch_size, out_channels, out_h, out_w) # 完整传递算子所需参数,处理可选参数的默认情况 mask_input = mask if use_mask else g.op("Constant", value_t=torch.zeros_like(offset)) bias_input = bias if bias is not None else g.op("Constant", value_t=torch.zeros(out_channels)) # 创建自定义算子节点 output = g.op( "custom::deform_conv2d", input, offset, weight, mask_input, bias_input ) # 显式设置输出张量的类型与形状,让ONNX完成形状推理 output.setType(input.type().with_sizes(output_shape)) return output register_custom_op_symbolic("torchvision::deform_conv2d", symbolic, 9) model = Model() input = torch.rand(1, 3, 10, 10) torch.onnx.export(model, input, 'dcn.onnx', opset_version=9)
关键修改说明
- 补充形状计算:依据卷积层的输入尺寸、padding、stride、dilation等参数,手动计算输出特征图的尺寸,确定完整输出形状。
- 完善算子参数:原symbolic函数仅传递了部分参数,修改后补充了weight、mask、bias等必要参数,保证算子逻辑完整。
- 显式设置输出类型:通过
setType方法将计算好的形状绑定到算子输出,让ONNX能够正确解析输出维度,消除警告。
运行上述代码后,ONNX转换时的形状推理警告会消失,生成的模型可被正确解析,后续转换为Paddle格式时也能避免形状相关错误。
内容的提问来源于stack exchange,提问作者阳铠行
相关产品推荐
相关产品推荐

