使用torchsummary时触发Runtime Error的问题求助
问题根源
- 错误重载
__call__方法:PyTorch的nn.Module默认通过forward方法定义前向传播逻辑,直接重载__call__会绕过框架的默认机制。torchsummary在运行时会自动生成带batch维度的测试输入(默认batch_size=2),而你的__call__方法错误地将输入的第一个维度当成了图像的高度,导致view后输入形状变为(2,1,2,200),此时高度2小于卷积核尺寸3,触发报错。 - 输入格式不兼容:模型预期输入是
(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
相关产品推荐
相关产品推荐

