You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何解决扩散模型训练中的张量维度不匹配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

错误根源

  1. ConvBlock内部残差逻辑错误:ConvBlock的forward方法中,直接将卷积后的输出(通道数为out_channels)与原始输入x(通道数为in_channels)相加。当ResidualBlock调用ConvBlock且in_channels≠out_channels时(比如第一个DownBlock中输入通道为6,输出为64),必然出现通道数不匹配。
  2. DownBlock返回值不匹配:DownBlock的forward仅返回经过池化后的特征,但UNet的forward中写了x1, skip1 = self.down1(x),试图接收两个返回值,导致张量解包错误,进一步引发后续通道混乱。
  3. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.14 11:22:05