基于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")
关键修复说明
- 修复数据归一化:移除
divide_image_into_patches中的imgs/255,保留官方归一化参数。 - 重构MHA模块:采用标准ViT多头注意力结构,对多头输出做线性投影合并。
- 固定LayerNorm:将LayerNorm移至
__init__中,确保参数可被优化器跟踪。 - 动态获取batch_size:从输入图像中实时获取batch_size,避免硬编码维度错误。
- 合理参数初始化:使用
trunc_normal_和xavier_normal_控制初始参数方差。 - 修正分类头:最后一层仅保留线性层,输出原始logits供损失计算。
- 优化器调整:改用AdamW优化器,配合学习率调度器提升收敛效率。
- 简化模型结构:降低embedding维度和编码器块数量,适配MNIST小数据集。
内容的提问来源于stack exchange,提问作者Paras Mehta
相关产品推荐
相关产品推荐

