PyTorch重复调用nn.Conv2d触发IndexError错误求助
错误原因分析
你遇到的IndexError: tuple index out of range本质是**self.reduction(x)返回的张量维度不足4维**,导致permute(0,2,3,1)时无法找到索引3。结合你的代码和现象,核心问题大概率出在以下几个环节:
1. 窗口尺寸的平方数断言失效
代码中通过shuf_size = int(shuf_size_2 ** 0.5)和win_size = int(win_size_2 ** 0.5)推导窗口尺寸,但如果shuf_size_2或win_size_2不是完全平方数,会导致实际窗口尺寸与预期不符,进而让window_reverse输出的张量维度错乱。
2. window_reverse函数实现异常
window_reverse负责将分块的窗口张量还原为4维特征图,如果函数内部的窗口数量计算、维度拼接逻辑错误,会导致输出张量不是预期的(B, C, H, W)4维格式,而是3维甚至更低维度,经过Conv2D后依然无法满足permute的维度要求。
3. self.input_resolution与实际输入不匹配
如果self.input_resolution设置的H、W与模型实际接收的输入图片分辨率不一致,会导致window_reverse在还原特征图时计算错误,输出异常维度的张量。
排查与修复步骤
步骤1:添加平方数校验
在计算shuf_size和win_size后立即添加断言,确保输入的窗口平方数合法:
shuf_size = int(shuf_size_2 ** 0.5) win_size = int(win_size_2 ** 0.5) # 断言检查是否为完全平方数 assert shuf_size ** 2 == shuf_size_2, f"shuf_size_2={shuf_size_2} must be a perfect square" assert win_size ** 2 == win_size_2, f"win_size_2={win_size_2} must be a perfect square"
步骤2:打印张量形状确认
在调用self.reduction前,强制打印张量形状,确认是否符合预期的4维格式:
# 检查msg_token的形状 msg_token = window_reverse( x[:, :, 0].unsqueeze(2), 1, shuf_size, H//win_size, W//win_size, nchw=True) print(f"msg_token before reduction: {msg_token.shape}") # 预期[128,64,8,8] msg_token = self.reduction(msg_token).permute(0, 2, 3, 1) # 检查x的形状 x = window_reverse(x[:, :, 1:], win_size, shuf_size, H, W, nchw=True) print(f"x before reduction: {x.shape}") # 预期[128,64,56,56] x = self.reduction(x).permute(0, 2, 3, 1)
如果打印出的形状不是4维,直接定位到window_reverse函数的实现问题。
步骤3:验证self.input_resolution的正确性
确认self.input_resolution的H、W与训练时输入图片的分辨率完全一致,比如输入图片是224x224,那self.input_resolution必须设置为(224,224),否则窗口还原时的尺寸计算会彻底错误。
步骤4:检查window_reverse的实现逻辑
确保window_reverse的输入输出维度符合预期:
- 输入格式:
(B*num_windows, window_size², C)(或对应nchw=True时的格式) - 输出格式:
(B, C, H, W)
如果函数内部存在维度顺序颠倒、窗口数量计算错误(比如num_windows != (H//window_size)*(W//window_size)),需要修正对应逻辑。
内容的提问来源于stack exchange,提问作者Yun

