PyTorch Hourglass模型张量尺寸不匹配RuntimeError问题求助
解决Deep Image Prior 4通道图像修复的张量尺寸不匹配问题
核心原因
Deep Image Prior的U-Net结构依赖输入尺寸为2的整数次幂(比如64、128、256这类),如果你的4通道输入图像尺寸不符合这个要求,下采样(池化/卷积)再上采样后的特征图尺寸,会和跳跃连接的特征图出现1像素的偏差,导致torch.cat时维度不匹配(比如报错里的17 vs 18)。另外,修改通道数时如果没同步调整所有卷积层的输入输出通道,也可能引发这类问题。
具体解决步骤
1. 把输入图像调整为2的幂次尺寸
预处理阶段直接将图像resize到最近的2的幂次尺寸,比如原尺寸300x300就改成256x256或512x512:
import torchvision.transforms as transforms # 假设img是你的4通道PIL图像或张量 resize_transform = transforms.Resize((256, 256)) resized_img = resize_transform(img)
要是想保留原图像比例,可以先填充到2的幂次尺寸,修复完成后再裁剪回原大小:
import torch import math def pad_to_power_of_two(tensor): h, w = tensor.shape[1], tensor.shape[2] new_h = 2 ** int(math.ceil(math.log2(h))) new_w = 2 ** int(math.ceil(math.log2(w))) pad_h = new_h - h pad_w = new_w - w # 对H、W维度做对称填充(4通道张量维度为[C, H, W]) padded = torch.nn.functional.pad(tensor, (pad_w//2, pad_w - pad_w//2, pad_h//2, pad_h - pad_h//2)) return padded, (h, w) # 预处理时填充 padded_img, original_size = pad_to_power_of_two(input_tensor) # 修复后裁剪回原尺寸 restored_img = restored_img[:, :original_size[0], :original_size[1]]
2. 同步调整U-Net所有卷积层的通道数
确保你修改了所有卷积模块的输入输出通道,包括编码器下采样、解码器上采样以及跳跃连接对应的部分:
import torch.nn as nn class UNet(nn.Module): def __init__(self, in_channels=4, out_channels=4): super().__init__() # 编码器部分 self.encoder1 = self.conv_block(in_channels, 64) self.encoder2 = self.conv_block(64, 128) self.encoder3 = self.conv_block(128, 256) self.encoder4 = self.conv_block(256, 512) self.encoder5 = self.conv_block(512, 1024) # 解码器部分 self.decoder1 = self.upconv_block(1024, 512) self.decoder2 = self.conv_block(512+512, 256) # 注意跳跃连接的通道数相加 self.decoder3 = self.conv_block(256+256, 128) self.decoder4 = self.conv_block(128+128, 64) self.decoder5 = nn.Conv2d(64+64, out_channels, kernel_size=3, padding=1) def conv_block(self, in_ch, out_ch): return nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1), nn.ReLU(inplace=True) ) def upconv_block(self, in_ch, out_ch): return nn.Sequential( nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2), nn.ReLU(inplace=True) ) def forward(self, x): # 编码器前向 skip1 = self.encoder1(x) x = nn.MaxPool2d(2)(skip1) skip2 = self.encoder2(x) x = nn.MaxPool2d(2)(skip2) skip3 = self.encoder3(x) x = nn.MaxPool2d(2)(skip3) skip4 = self.encoder4(x) x = nn.MaxPool2d(2)(skip4) skip5 = self.encoder5(x) # 解码器前向 x = self.decoder1(skip5) x = torch.cat([x, skip4], dim=1) x = self.decoder2(x) x = torch.cat([x, skip3], dim=1) x = self.decoder3(x) x = torch.cat([x, skip2], dim=1) x = self.decoder4(x) x = torch.cat([x, skip1], dim=1) x = self.decoder5(x) return x
重点确认:每个torch.cat操作前,上采样后的特征图和对应跳跃连接的特征图,H、W维度完全一致。
3. 替换上采样方式(可选)
如果不想修改输入尺寸,可以把转置卷积换成双线性插值+卷积,这种方式更灵活,能减少尺寸偏差:
def upconv_block(self, in_ch, out_ch): return nn.Sequential( nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True), nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1), nn.ReLU(inplace=True) )
4. 打印张量尺寸定位问题
在模型的forward方法里,每次上采样后打印张量尺寸,和对应跳跃连接的张量对比,明确哪一层出现偏差:
def forward(self, x): skip1 = self.encoder1(x) print(f"skip1 size: {skip1.shape}") x = nn.MaxPool2d(2)(skip1) skip2 = self.encoder2(x) print(f"skip2 size: {skip2.shape}") # ... 其他层同理 x = self.decoder1(skip5) print(f"up5 size: {x.shape}, skip4 size: {skip4.shape}") # 对应报错的up_5和skip_5 x = torch.cat([x, skip4], dim=1) # ... 后续层
验证方法
修改后先单独测试模型前向传播,传入一个4通道、2的幂次尺寸的张量,确认无报错:
model = UNet(in_channels=4, out_channels=4) test_input = torch.randn(1, 4, 256, 256) # batch_size=1,4通道,256x256 output = model(test_input) print(f"Output shape: {output.shape}") # 应该输出(1,4,256,256)
内容的提问来源于stack exchange,提问作者Khubaib Khawar
相关产品推荐
相关产品推荐

