PyTorch U-Net测试时RuntimeError:输入与权重通道数不匹配
解决PyTorch U-Net通道不匹配的RuntimeError问题
错误信息
RuntimeError: Given transposed=1, weight of size [1024, 512, 2, 2], expected input[24, 512, 48, 48] to have 1024 channels, but got 512 channels instead
问题根源
错误出在U-Net的前向传播函数中:所有上采样步骤都错误复用了up1模块,没有按照U-Net的编码器-解码器结构,依次使用up1、up2、up3、up4模块。
- 小通道数(如2、4)时未触发错误,是因为通道数巧合匹配了模块的输入输出要求;
- 当通道数增大到64、128这类常规值后,模块预设的输入输出通道与实际输入不匹配,直接触发上述报错。
修复步骤
- 打开U-Net实现文件(
unet.py),找到类中的forward方法; - 替换所有重复的
up1调用,按照下采样的逆过程依次调用对应的上采样模块:- 最深层特征先用
up1上采样,再与对应下采样特征拼接; - 后续步骤分别使用
up2、up3、up4完成上采样与拼接操作。
- 最深层特征先用
代码示例对比
错误的前向传播片段
x = self.up1(x) x = torch.cat([x, self.down3], dim=1) x = self.up1(x) # 错误:应使用up2 x = torch.cat([x, self.down2], dim=1) x = self.up1(x) # 错误:应使用up3 x = torch.cat([x, self.down1], dim=1) x = self.up1(x) # 错误:应使用up4
修正后的前向传播片段
x = self.up1(x) x = torch.cat([x, self.down3], dim=1) x = self.up2(x) x = torch.cat([x, self.down2], dim=1) x = self.up3(x) x = torch.cat([x, self.down1], dim=1) x = self.up4(x)
内容的提问来源于stack exchange,提问作者김지한
相关产品推荐
相关产品推荐

