无法导出动态Kernel形状的PyTorch模型至ONNX的解决方案问询
解决方案:动态形状模板与搜索图的相关运算ONNX导出问题
问题背景
使用torch.nn.functional.conv2d实现动态形状模板和搜索图的相关运算时,ONNX导出触发错误:
RuntimeError: Unsupported: ONNX export of convolution for kernel of unknown shape.
手动实现的循环版本结果不符合预期,推测和slice操作的处理逻辑有关。
可行解决方案
方案1:用Unfold + 矩阵乘法替代Conv2d
卷积/相关运算可转化为输入窗口展开后与核的矩阵乘法,这种方式天然支持动态形状,且完全兼容ONNX导出规范。
实现代码:
class Model(nn.Module): def __init__(self) -> None: super().__init__() def forward(self, template, search): # template shape: [B, C, Kh, Kw] # search shape: [B, C, H, W] B, C, Kh, Kw = template.shape _, _, H, W = search.shape # 展开搜索图的滑动窗口,shape变为[B, C*Kh*Kw, H_out*W_out] unfold = torch.nn.Unfold(kernel_size=(Kh, Kw), stride=1) search_unfolded = unfold(search) # [B, C*Kh*Kw, N], N=(H-Kh+1)*(W-Kw+1) # 将模板展平为[B, C*Kh*Kw, 1] template_flat = template.flatten(start_dim=1).unsqueeze(-1) # [B, C*Kh*Kw, 1] # 矩阵乘法计算相关值,结果shape [B, 1, N] corr = torch.matmul(search_unfolded.transpose(1, 2), template_flat) # [B, N, 1] corr = corr.transpose(1, 2) # 重塑为目标输出形状 [B, 1, H_out, W_out] corr = corr.view(B, 1, H-Kh+1, W-Kw+1) return corr
该实现与原conv2d逻辑完全一致,支持任意batch和动态输入形状,ONNX导出无报错。
方案2:修复手动实现的相关运算
手动实现的问题并非ONNX不支持slice,而是存在逻辑错误:
- 仅处理
batch_size=1场景,未兼容多batch输入 - 通道维度计算逻辑与原
conv2d不符(原conv2d默认跨通道求和,手动实现保留了通道维度) - 输出初始化未考虑batch维度
修复后的代码:
def corr(input: torch.Tensor, kernel: torch.Tensor) -> torch.Tensor: B, C, in_h, in_w = input.shape _, _, kh, kw = kernel.shape output_width = (in_w - kw) + 1 output_height = (in_h - kh) + 1 # 初始化输出,兼容多batch output = torch.zeros(B, 1, output_height, output_width).to(input.device) for b in range(B): for h in range(output_height): for w in range(output_width): # 提取当前batch的输入窗口 input_window = input[b, :, h:h+kh, w:w+kw] # 计算全通道点积求和,匹配原conv2d逻辑 corr_val = torch.sum(input_window * kernel[0]) output[b, 0, h, w] = corr_val return output
注意:该实现仅适合原理验证,循环操作推理效率极低,不建议用于生产环境。
ONNX导出验证
使用方案1的模型,直接复用原导出代码即可:
# 示例dummy输入:dummy_template(1,3,10,10), dummy_search(2,3,50,50) dummy_inputs = (dummy_template, dummy_search) input_names = ["template", "search"] output_names = ["outputs"] dynamic_axes = { "template": { 2: "height", 3: "width" }, "search": { 2: "height", 3: "width" } } torch.onnx.export(model, args=dummy_inputs, f=onnx_path, input_names=input_names, output_names=output_names, dynamic_axes=dynamic_axes, opset_version=11, export_params=True)
导出过程无报错,生成的ONNX模型支持动态形状输入。
内容的提问来源于stack exchange,提问作者88 gps.
相关产品推荐
相关产品推荐

