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

PyTorch中DDPM-U-Net自注意力模块矩阵维度不匹配错误求助

问题排查与解决:PyTorch U-Net+DDPM 矩阵维度不匹配错误

错误信息

File "C:\Users\zzzz\miniconda3\envs\ddpm2\Lib\site-packages\torch\nn\modules\linear.py", line 114, in forward
    return F.linear(input, self.weight, self.bias)
RuntimeError: mat1 and mat2 shapes cannot be multiplied (35840x28 and 10x10)

问题根源

错误出在SelfAttention类的维度处理逻辑:

  • 输入到SelfAttention的是4D特征图张量:(batch_size, channels, height, width),比如第一个MyBlockWithAttention输出的(N,10,28,28)。
  • PyTorch的nn.Linear默认对张量最后一维做线性变换,但你期望Linear作用在channels维度(对应in_dim=10),实际输入最后一维是width=28,导致矩阵乘法维度不匹配:展平后mat1为(N*28*28, 28),mat2为(10,10),无法相乘。

修正方案

1. 修改SelfAttention类,适配特征图维度

将4D特征图转换为注意力机制需要的序列格式,计算后再还原回特征图格式:

class SelfAttention(nn.Module):
    def __init__(self, in_dim, out_dim):
        super(SelfAttention, self).__init__()
        self.query = nn.Linear(in_dim, out_dim)
        self.key = nn.Linear(in_dim, out_dim)
        self.value = nn.Linear(in_dim, out_dim)
        self.softmax = nn.Softmax(dim=-1)

    def forward(self, x):
        # 输入形状:(N, C, H, W)
        N, C, H, W = x.shape
        # 转换为序列格式:(N, H*W, C)
        x_flat = x.permute(0, 2, 3, 1).reshape(N, H*W, C)
        
        # 计算Q/K/V
        query = self.query(x_flat)  # (N, H*W, out_dim)
        key = self.key(x_flat)      # (N, H*W, out_dim)
        value = self.value(x_flat)  # (N, H*W, out_dim)
        
        # 计算注意力分数与权重
        scores = torch.matmul(query, key.transpose(-2, -1)) / torch.sqrt(key.size(-1))
        attention_weights = self.softmax(scores)
        
        # 计算输出并还原特征图格式
        output = torch.matmul(attention_weights, value)
        output = output.reshape(N, H, W, out_dim).permute(0, 3, 1, 2)  # (N, out_dim, H, W)
        
        return output

2. 修正MyBlockWithAttention中的LayerNorm逻辑

原代码中LayerNorm的参数错误,调整为适配通道维度的归一化:

class MyBlockWithAttention(nn.Module):
    def __init__(self, in_c, out_c, kernel_size=3, stride=1, padding=1, activation=None, normalize=True):
        super(MyBlockWithAttention, self).__init__()
        # 对通道维度做LayerNorm,指定归一化的特征数为in_c
        self.ln = nn.LayerNorm(in_c) if normalize else None
        self.conv1 = nn.Conv2d(in_c, out_c, kernel_size, stride, padding)
        self.attention = SelfAttention(out_c, out_c)
        self.conv2 = nn.Conv2d(out_c, out_c, kernel_size, stride, padding)
        self.activation = nn.SiLU() if activation is None else activation
        self.normalize = normalize

    def forward(self, x):
        if self.normalize:
            # 转置维度以适配LayerNorm对最后一维的归一化逻辑
            x = x.permute(0,2,3,1)
            x = self.ln(x)
            x = x.permute(0,3,1,2)
        out = self.conv1(x)
        out = self.attention(out)
        out = self.activation(out)
        out = self.conv2(out)
        out = self.activation(out)
        return out

3. 更新U-Net中MyBlockWithAttention的调用

去掉原代码中多余的shape参数,直接传入输入通道数和输出通道数:
例如原代码中的:

MyBlockWithAttention((1, 28, 28), 1, 10)

改为:

MyBlockWithAttention(1, 10)

所有MyBlockWithAttention的调用都需要做此修改(包括b1、b2、b3等模块)。

验证逻辑

修正后,注意力机制将特征图的每个空间位置视为序列元素,通道作为特征维度,符合自注意力计算逻辑;同时LayerNorm的维度处理与输入张量格式匹配,不会再出现维度不匹配错误。

内容的提问来源于stack exchange,提问作者Zahra Hosseini

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 01:55:01