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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 15:15:54