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
相关产品推荐
相关产品推荐

