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

基于PyTorch的Vision Transformer模型训练异常排查及修正请求

ViT在MNIST上无法收敛的问题排查与修复

核心问题排查

模型始终维持10%准确率(随机猜测水平)、交叉熵2.31(对应均匀分布的logits),说明完全没有学习到有效特征,核心问题集中在以下几点:

  • 重复数据归一化:Transform中已执行Normalize((0.1307,), (0.3081,)),但divide_image_into_patches里又做imgs = imgs/255,导致数据特征分布混乱。
  • 多头注意力(MHA)实现错误:Head类的Q/K/V权重维度完全错误,且MHA未对多头输出做线性投影合并,不符合ViT标准结构。
  • LayerNorm动态创建:在forward中每次新建LayerNorm实例,导致参数无法被优化器跟踪,梯度传递中断。
  • 硬编码batch_size:使用初始化时的固定batch_size,若数据加载器实际batch_size不匹配,会引发维度错误。
  • 参数初始化不当:直接用torch.randn初始化参数,方差过大导致模型初始输出不稳定。
  • 分类头结构错误:最后一层使用GELU激活和Dropout,输出不是CrossEntropyLoss需要的原始logits,损失计算无效。
  • 优化器选择不合理:SGD对ViT优化效率低,初始学习率0.01过高且无学习率调度。

修正后的完整代码

# -*- coding: utf-8 -*-
import torch
from torch import nn
from torchvision import transforms
import torchvision.datasets as datasets
import torch.nn.functional as F
import torch.optim as optim
import numpy as np
import math

class Head(nn.Module):
    def __init__(self, embed_dim, head_dim):
        super().__init__()
        self.q = nn.Linear(embed_dim, head_dim)
        self.k = nn.Linear(embed_dim, head_dim)
        self.v = nn.Linear(embed_dim, head_dim)
        self.scale = math.sqrt(head_dim)

    def forward(self, x):
        q = self.q(x)
        k = self.k(x)
        v = self.v(x)
        
        # 计算注意力权重
        attn_weights = torch.matmul(q, k.transpose(-2, -1)) / self.scale
        attn_weights = F.softmax(attn_weights, dim=-1)
        
        output = torch.matmul(attn_weights, v)
        return output

class MHA(nn.Module):
    def __init__(self, embed_dim, num_heads):
        super().__init__()
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        assert embed_dim % num_heads == 0, "Embedding dimension must be divisible by number of heads"
        
        self.heads = nn.ModuleList([Head(embed_dim, self.head_dim) for _ in range(num_heads)])
        self.proj = nn.Linear(embed_dim, embed_dim)  # 多头输出合并投影

    def forward(self, x):
        # 并行计算所有head的输出
        head_outputs = [head(x) for head in self.heads]
        # 拼接所有head的输出,再做投影
        concat_output = torch.cat(head_outputs, dim=-1)
        output = self.proj(concat_output)
        return output

class EncoderBlock(nn.Module):
    def __init__(self, embed_dim, num_heads, hidden_dim, dropout=0.1):
        super().__init__()
        self.norm1 = nn.LayerNorm(embed_dim)
        self.mha = MHA(embed_dim, num_heads)
        self.norm2 = nn.LayerNorm(embed_dim)
        self.mlp = nn.Sequential(
            nn.Linear(embed_dim, hidden_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, embed_dim),
            nn.Dropout(dropout)
        )

    def forward(self, x):
        # 注意力残差连接
        x = x + self.mha(self.norm1(x))
        # MLP残差连接
        x = x + self.mlp(self.norm2(x))
        return x

class vit_model(nn.Module):
    def __init__(self, img_size, patch_size, embed_dim, num_heads, hidden_dim, n_classes, num_blocks=1):
        super().__init__()
        self.img_size = img_size
        self.patch_size = patch_size
        self.embed_dim = embed_dim
        self.num_patches = (img_size // patch_size) ** 2
        
        # Patch嵌入
        self.patch_embed = nn.Linear(patch_size * patch_size * 1, embed_dim)  # MNIST是单通道
        # 分类token
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))
        # 位置编码
        self.pos_embed = nn.Parameter(torch.randn(1, self.num_patches + 1, embed_dim))
        self.pos_drop = nn.Dropout(0.1)
        
        # 编码器块
        self.encoder_blocks = nn.Sequential(*[EncoderBlock(embed_dim, num_heads, hidden_dim) for _ in range(num_blocks)])
        # 分类头
        self.norm = nn.LayerNorm(embed_dim)
        self.head = nn.Linear(embed_dim, n_classes)
        
        # 参数初始化
        self._init_weights()

    def _init_weights(self):
        # 初始化patch嵌入和分类头
        nn.init.trunc_normal_(self.patch_embed.weight, std=0.02)
        nn.init.trunc_normal_(self.cls_token, std=0.02)
        nn.init.trunc_normal_(self.pos_embed, std=0.02)
        # 初始化MLP和线性层
        for m in self.modules():
            if isinstance(m, nn.Linear):
                nn.init.xavier_normal_(m.weight)
                if m.bias is not None:
                    nn.init.zeros_(m.bias)

    def divide_image_into_patches(self, imgs):
        batch_size, channels, height, width = imgs.shape
        # 分割patch: [B, C, H, W] -> [B, num_patches, C*patch_size*patch_size]
        patches = imgs.unfold(2, self.patch_size, self.patch_size).unfold(3, self.patch_size, self.patch_size)
        patches = patches.contiguous().view(batch_size, -1, channels * self.patch_size * self.patch_size)
        return patches

    def forward(self, images):
        batch_size = images.size(0)
        # 分割patch并嵌入
        patches = self.divide_image_into_patches(images)
        x = self.patch_embed(patches)
        
        # 添加分类token
        cls_tokens = self.cls_token.expand(batch_size, -1, -1)
        x = torch.cat((cls_tokens, x), dim=1)
        # 添加位置编码
        x = x + self.pos_embed
        x = self.pos_drop(x)
        
        # 编码器处理
        x = self.encoder_blocks(x)
        
        # 分类token输出
        x = self.norm(x[:, 0])
        x = self.head(x)
        return x

# 数据预处理
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.1307,), (0.3081,))  # MNIST官方归一化参数
])

# 加载数据集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
test_dataset = datasets.MNIST(root='./data', train=False, download=True, transform=transform)

# 数据加载器
batch_size = 128
train_loader = torch.utils.data.DataLoader(train_dataset, shuffle=True, batch_size=batch_size)
test_loader = torch.utils.data.DataLoader(test_dataset, shuffle=False, batch_size=batch_size)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# 初始化模型
model = vit_model(
    img_size=28,
    patch_size=4,
    embed_dim=128,  # 降低维度适配MNIST
    num_heads=4,
    hidden_dim=256,
    n_classes=10,
    num_blocks=2  # 减少编码器块数量,避免过拟合
).to(device)

# 损失函数与优化器
criterion = nn.CrossEntropyLoss()
optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-4)
# 学习率调度器
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.5)

# 训练循环
epoch_losses = []
epoch_accuracies = []
for epoch in range(20):
    epoch_loss = []
    epoch_acc = []
    model.train()
    for images, labels in train_loader:
        images = images.to(device)
        labels = labels.to(device)
        
        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        # 计算准确率
        with torch.no_grad():
            preds = torch.argmax(outputs, dim=-1)
            acc = torch.sum(preds == labels) / batch_size
        
        epoch_loss.append(loss.item())
        epoch_acc.append(acc.cpu().numpy())
    
    # 更新学习率
    scheduler.step()
    
    # 验证
    model.eval()
    val_loss = []
    val_acc = []
    with torch.no_grad():
        for images, labels in test_loader:
            images = images.to(device)
            labels = labels.to(device)
            outputs = model(images)
            loss = criterion(outputs, labels)
            preds = torch.argmax(outputs, dim=-1)
            acc = torch.sum(preds == labels) / batch_size
            val_loss.append(loss.item())
            val_acc.append(acc.cpu().numpy())
    
    avg_train_loss = np.mean(epoch_loss)
    avg_train_acc = np.mean(epoch_acc)
    avg_val_loss = np.mean(val_loss)
    avg_val_acc = np.mean(val_acc)
    
    epoch_losses.append(avg_train_loss)
    epoch_accuracies.append(avg_train_acc)
    
    print(f"Epoch {epoch+1}:")
    print(f"Train Loss: {avg_train_loss:.4f}, Train Acc: {avg_train_acc:.4f}")
    print(f"Val Loss: {avg_val_loss:.4f}, Val Acc: {avg_val_acc:.4f}\n")

关键修复说明

  1. 修复数据归一化:移除divide_image_into_patches中的imgs/255,保留官方归一化参数。
  2. 重构MHA模块:采用标准ViT多头注意力结构,对多头输出做线性投影合并。
  3. 固定LayerNorm:将LayerNorm移至__init__中,确保参数可被优化器跟踪。
  4. 动态获取batch_size:从输入图像中实时获取batch_size,避免硬编码维度错误。
  5. 合理参数初始化:使用trunc_normal_和xavier_normal_控制初始参数方差。
  6. 修正分类头:最后一层仅保留线性层,输出原始logits供损失计算。
  7. 优化器调整:改用AdamW优化器,配合学习率调度器提升收敛效率。
  8. 简化模型结构:降低embedding维度和编码器块数量,适配MNIST小数据集。

内容的提问来源于stack exchange,提问作者Paras Mehta

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 23:56:58