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

使用torchsummary时触发Runtime Error的问题求助

问题根源

  1. 错误重载__call__方法:PyTorch的nn.Module默认通过forward方法定义前向传播逻辑,直接重载__call__会绕过框架的默认机制。torchsummary在运行时会自动生成带batch维度的测试输入(默认batch_size=2),而你的__call__方法错误地将输入的第一个维度当成了图像的高度,导致view后输入形状变为(2,1,2,200),此时高度2小于卷积核尺寸3,触发报错。
  2. 输入格式不兼容:模型预期输入是(H,W)的cupy数组,但torchsummary传入的是(B,H,W)的torch张量,两者格式不匹配。

解决方案

方法1:规范实现forward方法(推荐)

修改模型类,用forward替代__call__,同时兼容torch张量和cupy数组的输入,正确处理batch维度:

import numpy as np
import torch
import torch.nn as nn
import cupy as cp
from torchviz import make_dot
from torchinfo import summary
from torchsummary import summary as summary_
    
def get_filter_torch(*args, **kwargs):
    
    class TraversabilityFilter(nn.Module):
        def __init__(self, w1, w2, w3, w_out, device="cuda", use_bias=False):
            super(TraversabilityFilter, self).__init__()
            self.conv1 = nn.Conv2d(1, 4, 3, dilation=1, padding=0, bias=use_bias)
            self.conv2 = nn.Conv2d(1, 4, 3, dilation=2, padding=0, bias=use_bias)
            self.conv3 = nn.Conv2d(1, 4, 3, dilation=3, padding=0, bias=use_bias)
            self.conv_out = nn.Conv2d(12, 1, 1, bias=use_bias)

            # Set weights.
            self.conv1.weight = nn.Parameter(torch.from_numpy(w1).float())
            self.conv2.weight = nn.Parameter(torch.from_numpy(w2).float())
            self.conv3.weight = nn.Parameter(torch.from_numpy(w3).float())
            self.conv_out.weight = nn.Parameter(torch.from_numpy(w_out).float())
            self.device = device

        def forward(self, x):
            # 处理输入:兼容cupy数组和torch张量,确保形状为(B,1,H,W)
            if isinstance(x, cp.ndarray):
                x = torch.as_tensor(x.astype(cp.float32), device=self.device)
                # 补充缺失的batch和channel维度
                if len(x.shape) == 2:
                    x = x.unsqueeze(0).unsqueeze(0)
                elif len(x.shape) ==3:
                    x = x.unsqueeze(1)
            elif isinstance(x, torch.Tensor):
                # 补充缺失的channel维度
                if len(x.shape) ==2:
                    x = x.unsqueeze(0).unsqueeze(0)
                elif len(x.shape)==3:
                    x = x.unsqueeze(1)
            
            with torch.no_grad():
                out1 = self.conv1(x)
                out2 = self.conv2(x)
                out3 = self.conv3(x)

                # 适配batch维度的裁剪
                out1 = out1[:, :, 2:-2, 2:-2]
                out2 = out2[:, :, 1:-1, 1:-1]
                out = torch.cat((out1, out2, out3), dim=1)
                out = self.conv_out(out.abs())
                out = torch.exp(-out)

            return out

    traversability_filter = TraversabilityFilter(*args, **kwargs).cuda().eval()
    return traversability_filter
   

# Define the weight values
w1 = np.random.randn(4, 1, 3, 3)
w2 = np.random.randn(4, 1, 3, 3)
w3 = np.random.randn(4, 1, 3, 3)
w_out = np.random.randn(1, 12, 1, 1)

model = get_filter_torch(w1, w2, w3, w_out)

cell_n = 200
x = cp.random.randn(cell_n, cell_n, dtype=cp.float32)
output = model(x)
print(model)

# 使用torchsummary时传入包含channel维度的输入尺寸
input_size=(1, cell_n, cell_n)
summary(model)
summary_(model, input_size)

方法2:直接传入测试张量给torchsummary(临时方案)

如果不想修改模型结构,手动创建符合模型预期的测试张量,直接传给torchsummary:

# 替换原来的summary_调用
test_input = torch.randn(200, 200).cuda()  # 模拟单个样本输入
summary_(model, input_data=test_input)

说明

  • 方法1遵循PyTorch规范,让模型兼容框架工具(torchsummary、torchinfo等),同时支持多格式输入,是长期维护的最优选择。
  • 方法2适合快速验证模型结构,但模型输入处理逻辑仍存在兼容性问题,不推荐长期使用。

内容的提问来源于stack exchange,提问作者YJ C

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 09:22:03