PyTorch实现U-Net遇维度不匹配错误,求助解决
问题描述
实现U-Net架构时运行报错,代码如下:
import torch import torch.nn as nn class UNet(nn.Module): def __init__(self, in_channels, out_channels): super(UNet, self).__init__() self.encoder1 = self.double_conv(in_channels, 64) self.encoder2 = self.down(64, 128) self.encoder3 = self.down(128, 256) self.encoder4 = self.down(256, 512) self.bottleneck = self.double_conv(512, 1024) self.decoder4 = self.up(1024, 512) self.decoder3 = self.up(512, 256) self.decoder2 = self.up(256, 128) self.decoder1 = self.up(128, 64) self.final_conv = nn.Conv2d(64, out_channels, kernel_size=1) # SAME convolution/padding def double_conv(self, in_channels, out_channels): # Convo Block return nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), nn.ReLU(inplace=True), nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), nn.ReLU(inplace=True), ) def down(self, in_channels, out_channels): return nn.Sequential( nn.MaxPool2d(kernel_size=2, stride=2), self.double_conv(in_channels, out_channels), ) def up(self, in_channels, out_channels): return nn.Sequential( nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2), self.double_conv(in_channels, out_channels), ) def forward(self, x): # Encoder enc1 = self.encoder1(x) # Output: [1, 64, 256, 256] print("enc1.shape",enc1.shape) enc2 = self.encoder2(enc1) # Output: [1, 128, 128, 128] print("enc2.shape",enc2.shape) enc3 = self.encoder3(enc2) # Output: [1, 256, 64, 64] print("enc3.shape",enc3.shape) enc4 = self.encoder4(enc3) # Output: [1, 512, 32, 32] print("enc4.shape",enc4.shape) bottleneck_output = self.bottleneck(enc4) # Output: [1, 1024, 32, 32] print("bottleneck_output",bottleneck_output.shape) # Decoder dec4 = self.decoder4(bottleneck_output) # Output: [1, 512, 64, 64] print(dec4.shape) dec4 = torch.cat((dec4, enc4), dim=1) # skip connect, Concatenate: [1, 1024, 64, 64] dec4 = self.double_conv(1024, 512)(dec4) # Corrected input channels to 1024 dec3 = self.decoder3(dec4) # Output: [1, 256, 128, 128] dec3 = torch.cat((dec3, enc3), dim=1) # Concatenate: [1, 512, 128, 128] dec3 = self.double_conv(512, 256)(dec3) # Corrected input channels to 512 dec2 = self.decoder2(dec3) # Output: [1, 128, 256, 256] dec2 = torch.cat((dec2, enc2), dim=1) # Concatenate: [1, 256, 256, 256] dec2 = self.double_conv(256, 128)(dec2) # Corrected input channels to 256 dec1 = self.decoder1(dec2) # Output: [1, 64, 512, 512] dec1 = torch.cat((dec1, enc1), dim=1) # Concatenate: [1, 128, 512, 512] dec1 = self.double_conv(128, 64)(dec1) # Corrected input channels to 128 return self.final_conv(dec1) # Output: [1, 1, 512, 512]
主函数执行代码:
unet = UNet(in_channels=1, out_channels=1) sample_input = torch.randn(1, 1, 256, 256) output = unet(sample_input)
各层输出形状:
enc1.shape torch.Size([1, 64, 256, 256]) enc2.shape torch.Size([1, 128, 128, 128]) enc3.shape torch.Size([1, 256, 64, 64]) enc4.shape torch.Size([1, 512, 32, 32]) bottleneck_output torch.Size([1, 1024, 32, 32])
报错信息:
---> 55 dec4 = self.decoder4(bottleneck_output) RuntimeError: Given groups=1, weight of size [512, 1024, 3, 3], expected input[1, 512, 64, 64] to have 1024 channels, but got 512 channels instead
瓶颈输出确实有1024通道,但decoder4无法正确识别,尝试过匹配维度无效,求解决思路。
问题分析
错误根源在up函数的实现逻辑:
ConvTranspose2d会把输入通道数从in_channels降到out_channels,比如处理1024通道的瓶颈输出时,输出是512通道、64x64的特征图。- 但紧接着调用的
double_conv传入的参数还是in_channels=1024和out_channels=512,导致这个卷积块的第一个卷积层期望输入是1024通道,而实际输入是512通道,维度完全不匹配。
解决方案
有两种修改方式可以解决问题:
方式一:修正up函数中double_conv的输入通道
让double_conv的输入通道数等于ConvTranspose2d的输出通道数(也就是out_channels):
def up(self, in_channels, out_channels): return nn.Sequential( nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2), self.double_conv(out_channels, out_channels), # 将第一个参数改为out_channels )
方式二:简化up函数(更贴合当前forward逻辑)
查看你的forward流程,在decoder部分已经手动做了特征拼接和后续的double_conv处理,原up函数里的double_conv属于重复操作,可直接简化为只保留转置卷积:
def up(self, in_channels, out_channels): return nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2)
两种方式都能解决通道不匹配的问题,第二种更符合你当前的解码流程设计。
内容的提问来源于stack exchange,提问作者greenbow
相关产品推荐
相关产品推荐

