You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

无法导出动态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,而是存在逻辑错误:

  1. 仅处理batch_size=1场景,未兼容多batch输入
  2. 通道维度计算逻辑与原conv2d不符(原conv2d默认跨通道求和,手动实现保留了通道维度)
  3. 输出初始化未考虑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.

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.13 11:33:10