PyTorch从零实现Vision Transformer矩阵维度不匹配报错解决
ViT实现中RuntimeError(维度不匹配)的解决方法
问题根源
报错RuntimeError: mat1 and mat2 shapes cannot be multiplied (30x50176 and 768x768)出现在Transformer编码器的线性层,核心原因是PatchEmbedding模块输出的张量维度错误:
- 原代码将图像分patch后错误保留了通道维度,导致输入线性层的张量形状为
(B, 3, 50176),而线性层输入特征数为768,维度完全不匹配。 - 同时存在MultiHeadAttention注意力计算错误、Transformer层缺少残差连接、未正确实现cls_token等问题。
具体修改步骤
1. 修复PatchEmbedding模块
调整张量维度转换逻辑,将每个patch展平为单向量,输出形状为(B, num_patches, embed_dim):
class PatchEmbedding(torch.nn.Module): def __init__(self, patch_size, in_channels, embed_dim): super().__init__() self.patch_size = patch_size self.embed_dim = embed_dim # 输入特征数是单个patch的总元素数:通道数*patch边长平方 self.projection = torch.nn.Linear(patch_size**2 * in_channels, embed_dim) def forward(self, x): B, C, H, W = x.shape # 将图像分割为patch:(B, C, H//p, p, W//p, p) x = x.reshape(B, C, H // self.patch_size, self.patch_size, W // self.patch_size, self.patch_size) # 转置维度,将patch的空间维度和通道维度合并:(B, H//p, W//p, C*p*p) x = x.permute(0, 2, 4, 1, 3, 5).flatten(3) # 调整为(B, num_patches, C*p*p),再投影到embed_dim num_patches = (H // self.patch_size) * (W // self.patch_size) x = x.reshape(B, num_patches, C * self.patch_size * self.patch_size) x = self.projection(x) return x
2. 正确实现cls_token
在VisionTransformer中初始化可学习的cls_token,并在patch嵌入后拼接至序列头部:
class VisionTransformer(torch.nn.Module): def __init__(self, patch_size, in_channels, embed_dim, num_heads, num_layers, num_classes): super().__init__() self.patch_embed = PatchEmbedding(patch_size, in_channels, embed_dim) self.encoder = TransformerEncoder(embed_dim, num_heads, num_layers) self.classifier = torch.nn.Linear(embed_dim, num_classes) # 初始化可学习的cls_token self.cls_token = torch.nn.Parameter(torch.randn(1, 1, embed_dim)) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) # 复制cls_token到batch的每个样本 cls_tokens = self.cls_token.expand(B, -1, -1) # 拼接cls_token到patch序列头部 x = torch.cat([cls_tokens, x], dim=1) x = self.encoder(x) # 提取cls_token的输出用于分类 x = x[:, 0] x = self.classifier(x) return x
3. 修复MultiHeadAttention的注意力计算
修正注意力权重的计算逻辑,确保维度匹配:
class MultiHeadAttention(torch.nn.Module): def __init__(self, embed_dim, num_heads): super().__init__() self.embed_dim = embed_dim self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.q_proj = torch.nn.Linear(embed_dim, embed_dim) self.k_proj = torch.nn.Linear(embed_dim, embed_dim) self.v_proj = torch.nn.Linear(embed_dim, embed_dim) self.out_proj = torch.nn.Linear(embed_dim, embed_dim) def forward(self, q, k, v): B, T, E = q.shape q = self.q_proj(q).reshape(B, T, self.num_heads, self.head_dim).transpose(1, 2) # (B, heads, T, head_dim) k = self.k_proj(k).reshape(B, T, self.num_heads, self.head_dim).transpose(1, 2) v = self.v_proj(v).reshape(B, T, self.num_heads, self.head_dim).transpose(1, 2) # 计算注意力权重:(B, heads, T, T) attn = torch.matmul(q, k.transpose(-2, -1)) / self.head_dim**0.5 attn = attn.softmax(dim=-1) # 加权求和得到输出 out = torch.matmul(attn, v) out = out.transpose(1, 2).reshape(B, T, E) out = self.out_proj(out) return out
4. 修复TransformerEncoder的层结构
添加残差连接,并调整为标准的Pre-LN结构:
class TransformerEncoder(torch.nn.Module): def __init__(self, embed_dim, num_heads, num_layers): super().__init__() self.layers = torch.nn.ModuleList([ torch.nn.Sequential( torch.nn.LayerNorm(embed_dim), MultiHeadAttention(embed_dim, num_heads), torch.nn.LayerNorm(embed_dim), torch.nn.Sequential( torch.nn.Linear(embed_dim, embed_dim * 4), torch.nn.GELU(), torch.nn.Linear(embed_dim * 4, embed_dim), torch.nn.Dropout(0.1) ) ) for _ in range(num_layers) ]) self.norm = torch.nn.LayerNorm(embed_dim) def forward(self, x): for layer in self.layers: # 注意力层残差连接 x = x + layer[0:2](x, x, x) # MLP层残差连接 x = x + layer[2:](x) x = self.norm(x) return x
完整修正后的代码
import torch import torchvision from torchvision import transforms from google.colab import drive drive.mount('/content/gdrive') import zipfile zip_ref = zipfile.ZipFile('/content/gdrive/MyDrive/dataset/data9k.zip', 'r') zip_ref.extractall("/content/dataset") zip_ref.close() import os from torchvision import datasets data_dir = 'dataset/datasets' train_dataset = datasets.ImageFolder(root=os.path.join(data_dir, 'train'), transform=transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]) ])) patch_size = 16 in_channels = 3 embed_dim = 768 num_heads = 8 num_layers = 12 num_classes = 4 epochs = 5 class PatchEmbedding(torch.nn.Module): def __init__(self, patch_size, in_channels, embed_dim): super().__init__() self.patch_size = patch_size self.embed_dim = embed_dim self.projection = torch.nn.Linear(patch_size**2 * in_channels, embed_dim) def forward(self, x): B, C, H, W = x.shape x = x.reshape(B, C, H // self.patch_size, self.patch_size, W // self.patch_size, self.patch_size) x = x.permute(0, 2, 4, 1, 3, 5).flatten(3) num_patches = (H // self.patch_size) * (W // self.patch_size) x = x.reshape(B, num_patches, C * self.patch_size * self.patch_size) x = self.projection(x) return x class MultiHeadAttention(torch.nn.Module): def __init__(self, embed_dim, num_heads): super().__init__() self.embed_dim = embed_dim self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.q_proj = torch.nn.Linear(embed_dim, embed_dim) self.k_proj = torch.nn.Linear(embed_dim, embed_dim) self.v_proj = torch.nn.Linear(embed_dim, embed_dim) self.out_proj = torch.nn.Linear(embed_dim, embed_dim) def forward(self, q, k, v): B, T, E = q.shape q = self.q_proj(q).reshape(B, T, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(k).reshape(B, T, self.num_heads, self.head_dim).transpose(1, 2) v = self.v_proj(v).reshape(B, T, self.num_heads, self.head_dim).transpose(1, 2) attn = torch.matmul(q, k.transpose(-2, -1)) / self.head_dim**0.5 attn = attn.softmax(dim=-1) out = torch.matmul(attn, v) out = out.transpose(1, 2).reshape(B, T, E) out = self.out_proj(out) return out class TransformerEncoder(torch.nn.Module): def __init__(self, embed_dim, num_heads, num_layers): super().__init__() self.layers = torch.nn.ModuleList([ torch.nn.Sequential( torch.nn.LayerNorm(embed_dim), MultiHeadAttention(embed_dim, num_heads), torch.nn.LayerNorm(embed_dim), torch.nn.Sequential( torch.nn.Linear(embed_dim, embed_dim * 4), torch.nn.GELU(), torch.nn.Linear(embed_dim * 4, embed_dim), torch.nn.Dropout(0.1) ) ) for _ in range(num_layers) ]) self.norm = torch.nn.LayerNorm(embed_dim) def forward(self, x): for layer in self.layers: x = x + layer[0:2](x, x, x) x = x + layer[2:](x) x = self.norm(x) return x class VisionTransformer(torch.nn.Module): def __init__(self, patch_size, in_channels, embed_dim, num_heads, num_layers, num_classes): super().__init__() self.patch_embed = PatchEmbedding(patch_size, in_channels, embed_dim) self.encoder = TransformerEncoder(embed_dim, num_heads, num_layers) self.classifier = torch.nn.Linear(embed_dim, num_classes) self.cls_token = torch.nn.Parameter(torch.randn(1, 1, embed_dim)) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat([cls_tokens, x], dim=1) x = self.encoder(x) x = x[:, 0] x = self.classifier(x) return x model = VisionTransformer(patch_size, in_channels, embed_dim, num_heads, num_layers, num_classes) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=0.001) criterion = torch.nn.CrossEntropyLoss() from torch.utils.data import DataLoader dataloader = DataLoader(train_dataset, batch_size=10, shuffle=True) for epoch in range(epochs): for batch_idx, (data, target) in enumerate(dataloader): data, target = data.to(device), target.to(device) output = model(data) loss = criterion(output, target) optimizer.zero_grad() loss.backward() optimizer.step() if batch_idx % 100 == 0: print('Epoch: {}, Batch: {}, Loss: {:.4f}'.format( epoch, batch_idx, loss.item()))
内容的提问来源于stack exchange,提问作者Fuji
相关产品推荐
相关产品推荐

