如何解决扩散模型训练中的张量维度不匹配RuntimeError
图像到图像扩散模型训练错误修复
问题背景
训练128×128尺寸的图像到图像扩散模型,batch size设为8,使用带注意力和残差块的UNet结构,训练时出现通道不匹配的RuntimeError:
RuntimeError: The size of tensor a (64) must match the size of tensor b (6) at non-singleton dimension 1
错误根源
- ConvBlock内部残差逻辑错误:
ConvBlock的forward方法中,直接将卷积后的输出(通道数为out_channels)与原始输入x(通道数为in_channels)相加。当ResidualBlock调用ConvBlock且in_channels≠out_channels时(比如第一个DownBlock中输入通道为6,输出为64),必然出现通道数不匹配。 - DownBlock返回值不匹配:
DownBlock的forward仅返回经过池化后的特征,但UNet的forward中写了x1, skip1 = self.down1(x),试图接收两个返回值,导致张量解包错误,进一步引发后续通道混乱。 - UNet中img_size未定义:
UNetWithAttention的forward最后使用了img_size变量,但未在类中定义或传入,会导致额外运行错误。
修复步骤
1. 移除ConvBlock内部的残差连接
ConvBlock的职责应为卷积+BN+激活+可选注意力,残差逻辑由外部的ResidualBlock处理,修改后的ConvBlock代码:
class ConvBlock(nn.Module): def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1, use_attention=False): super(ConvBlock, self).__init__() self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding) self.bn = nn.BatchNorm2d(out_channels) self.relu = nn.LeakyReLU(0.2, inplace=True) self.use_attention = use_attention self.attention = AttentionBlock(out_channels) if use_attention else None def forward(self, x): x = self.conv(x) x = self.bn(x) x = self.relu(x) if self.use_attention: x = self.attention(x) return x
2. 修正DownBlock的返回值
让DownBlock返回池化后的特征和残差块输出(作为skip连接),适配UNet的调用逻辑:
class DownBlock(nn.Module): def __init__(self, in_channels, out_channels, use_attention=False): super(DownBlock, self).__init__() self.residual_block = ResidualBlock(in_channels, out_channels, use_attention=use_attention) self.pool = nn.MaxPool2d(2) def forward(self, x): print(f"Input to DownBlock x shape: {x.shape}") skip = self.residual_block(x) # 残差块输出作为skip连接 x = self.pool(skip) # 对skip结果做池化得到下一层输入 print(f"Output from DownBlock x shape: {x.shape}, skip shape: {skip.shape}") return x, skip # 返回两个值
3. 定义UNet中的img_size变量
在UNetWithAttention初始化时添加img_size参数,避免未定义错误:
class UNetWithAttention(nn.Module): def __init__(self, in_channels, out_channels, img_size=128, base_channels=[64, 128, 256, 512], position_encoding_dim=128, timestep_dim=1, use_attention=True): super(UNetWithAttention, self).__init__() self.img_size = img_size # 定义图像尺寸 self.timestep_embed_proj = nn.Linear(position_encoding_dim, base_channels[3]) # 其余初始化代码保持不变... def forward(self, x, t=None): # 下采样、瓶颈、上采样代码保持不变... # 使用类内定义的img_size做上采样 x = F.interpolate(x, size=(self.img_size, self.img_size), mode='bilinear', align_corners=False) return x
同时初始化模型时传入img_size:
unet = UNetWithAttention(in_channels=6, out_channels=3, img_size=128, base_channels=[64, 128, 256, 512], position_encoding_dim=position_encoding_dim, timestep_dim=1, use_attention=True)
验证修复
修改后重新运行训练代码,通道不匹配的错误会被解决,同时DownBlock返回值和img_size的问题也得到修复,模型可正常前向传播。
内容的提问来源于stack exchange,提问作者teetee.py
相关产品推荐
相关产品推荐

