第一次看ViT论文的时候,不少人会忍不住倒吸一口凉气。明明是一张图片,怎么就能像处理句子一样,切成patch喂给Transformer?虽然论文里写得清清楚楚,但真到自己用PyTorch复现时,Patch Embedding怎么写、Position Embedding怎么加、Class Token放在哪,这些细节稍不留神就会翻车。这篇博文就是来填这个坑的。我会把Vision Transformer的PyTorch代码从里到外拆一遍,关键位置加上图解思路,配合完整可运行的代码,让不熟悉Transformer结构的人也能把这套架构吃透,并真正跑起来。
1. 核心设计拆解:ViT到底在模仿什么
1.1 先理解ViT解决的核心问题
ViT(Vision Transformer)这个思路,最颠覆的一点就是把图像彻底当成序列来处理。传统CNN靠卷积核逐层滑动,天然具备局部性和平移等变性,所以它能很快捕捉边缘、纹理这些局部特征。而ViT的做法是:把一张图切成固定大小的patch,每个patch展平后做一次线性映射,变成token,然后丢进标准的Transformer Encoder里面做全局自注意力建模。
这个转变带来了几个关键影响。一是感受野问题:CNN要看到全局信息,必须依赖深层堆叠,靠层数堆出足够大的感受野;而Transformer的每一层用自注意力直接建立任意两个patch之间的联系,第一层就已经是全局视野。二是归纳偏置问题:CNN把“相邻像素大概率相关”这种偏置内置到网络结构里,而ViT主动放弃了这个偏置,完全靠数据来学习。这也是为什么ViT需要更大的数据集或者更强的数据增强才能训练好,在ImageNet那种大规模数据集上效果才明显。
从代码实现角度看,ViT有一个特别友好的特点:结构上它跟NLP领域的BERT几乎一模一样,只是输入从词向量换成了图像patch的Embedding。所以你只要把Transformer Encoder那套吃透了,ViT的代码其实就剩三块新东西:Patch Embedding、Class Token、Position Embedding。这三块单独拎出来都不难,难的是理解它们各自解决的问题和拼接时的维度变化。
1.2 整体数据流向图解
很多教程贴出ViT结构图时会画得很复杂,但本质上数据流是这样的:
输入一张3通道、224x224的图片,先切成一堆16x16的patch,224/16=14,一共14x14=196个patch。每个patch展平后是16x16x3=768维的向量,经过一个线性层映射成D维(比如768维),这样我们就得到了196个token。然后在序列最前面拼一个可学习的Class Token,序列长度变成197。再给这197个token都加上一个可学习的位置编码,保持维度不变。接着丢进L层Transformer Encoder。最后取序列第一个token(也就是Class Token对应位置的输出)过一层分类头,得到类别概率。
这里有个特别容易混淆的点:最后一个输出到底取谁?Transformer Encoder输出的序列长度还是197,每个位置对应一个D维向量。分类只用第一个位置,也就是当初拼上的Class Token,经过N层编码之后的表示。后面代码里我会专门标出来这一行。
整个数据流对应到PyTorch的Tensor形状变化,就是下面这个流程:
- 输入:
(B, 3, 224, 224) - Patch Embedding后:
(B, 196, 768) - 拼接Class Token后:
(B, 197, 768) - 加Position Embedding后:
(B, 197, 768) - 经过Transformer Encoder后:
(B, 197, 768) - 取第一个token并过分类头后:
(B, num_classes)
后面所有代码都会围绕这个流程展开。
2. 环境配置与关键依赖准备
2.1 依赖安装与版本选择
ViT的实现并不复杂,核心依赖就是PyTorch和TorchVision。如果你没有现成的环境,用Anaconda建一个干净的环境是最省事的做法。
conda create -n vit_env python=3.9 -y conda activate vit_env pip install torch torchvision pip install matplotlib tqdm这里有几个版本细节需要注意。Python版本建议3.9或更高,PyTorch建议2.x版本,因为2.x里有一些对Transformer比较友好的改进。如果你有NVIDIA显卡,装CUDA版的PyTorch可以大幅加速训练;如果只有CPU也没关系,CIFAR-10这种小规模数据集用CPU跑几个epoch作为流程验证是完全可以接受的,就是慢一点。
还有个小提醒:如果图片预处理需要做RandomResizedCrop、RandomHorizontalFlip之类的基础增强,torchvision里全都有现成的,不需要额外装albumentations这些第三方库。训练ViT时数据增强很重要,但不必一开始就上太重的trick,先把主体流程跑通更重要。
2.2 数据准备:CIFAR-10举例
为了让大家能快速复现,我用CIFAR-10来做演示。这个数据集有6万张32x32的彩色图片,分10个类别,下载方便、单张图很小,训练起来压力不大。
import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms transform_train = transforms.Compose([ transforms.RandomCrop(32, padding=4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) transform_test = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) trainset = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform_train) testset = datasets.CIFAR10(root='./data', train=False, download=True, transform=transform_test) trainloader = DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2) testloader = DataLoader(testset, batch_size=256, shuffle=False, num_workers=2)这里要注意一个关键尺寸问题:CIFAR-10的图片是32x32,如果要切成16x16的patch,那整张图只有2x2=4个patch,信息量太少,Transformer几乎没法学出有效特征。所以对于CIFAR-10这种小图,一般把patch设为4x4,或者先用插值把图放大到64x64或224x224。我在后面的代码里采用patch_size=4的方式,这样能拿到8x8=64个token,配合较小的模型规模,在CIFAR-10上能跑出不错的效果。
3. ViT核心代码逐行拆解
3.1 Patch Embedding的两种实现方式与原理
Patch Embedding的目标是把图像转成token序列。最直观的做法:先用unfold把图像切成patch,然后对每个patch做线性变换。不过在实际工程中,更多是直接用nn.Conv2d一个卷积搞定,这也是很多开源实现的写法。
class PatchEmbed(nn.Module): def __init__(self, in_channels=3, patch_size=4, embed_dim=192): super().__init__() self.patch_size = patch_size # 用卷积实现patch切分 + 线性投影 self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): # x: (B, 3, H, W) B, C, H, W = x.shape assert H % self.patch_size == 0 and W % self.patch_size == 0, \ f"输入尺寸 {H}x{W} 不能被 patch_size {self.patch_size} 整除" # 卷积后: (B, embed_dim, H/patch_size, W/patch_size) x = self.proj(x) # 展平后两个维度: (B, embed_dim, num_patches) x = x.flatten(2) # 转置成序列格式: (B, num_patches, embed_dim) x = x.transpose(1, 2) return x用Conv2d实现的关键点在于:卷积核大小等于patch大小,步长也等于patch大小,这样卷积输出特征图的每个点就对应原图的一个patch,而且每个patch都经过了同一个卷积核的加权求和——这本质上就是一个线性映射。卷积输出的通道数设为embed_dim,相当于把每个patch展平后的向量映射到了D维空间。
为什么推荐用卷积而不是unfold?主要是因为Conv2d在GPU上的实现经过了深度优化,速度和显存利用效率都更好。而且卷积操作的表述非常简洁,一行代码就完成了切patch和投影两件事。不过直接用unfold其实也不难的,理解一下就行。
维度变化是最容易绕晕的地方,我把核心过程再捋一遍:输入(B, 3, 32, 32),patch_size=4,Conv2d输出(B, 192, 8, 8),flatten(2)后是(B, 192, 64),transpose(1,2)后是(B, 64, 192)。这64个token,每个都是192维。
3.2 Class Token和Position Embedding的细节实现
Patch Embedding之后,95%的人会在Class Token和Position Embedding这里犯迷糊。先看位置编码。Transformer本身没有顺序概念,而patch之间的相对位置对图像理解又至关重要,所以必须把位置信息揉进输入里。ViT用的是可学习的Position Embedding,也就是一个维度为(1, num_patches+1, embed_dim)的参数,训练时跟着网络一起更新。
Class Token的来历更有意思。因为Transformer Encoder输出的每个位置都对应一个token的表示,做分类时需要从这一堆token表示中汇聚出“整张图”的表示。最简单做法是对所有token做全局池化,但ViT论文里选择在序列最前面放一个可学习的Class Token,最后取这个token的输出当作整张图的特征。这个Class Token会在训练中学会“汇总”其他patch的信息。
class ViT(nn.Module): def __init__(self, image_size=32, patch_size=4, num_classes=10, embed_dim=192, depth=6, num_heads=3, mlp_ratio=4.0): super().__init__() self.patch_embed = PatchEmbed(in_channels=3, patch_size=patch_size, embed_dim=embed_dim) num_patches = (image_size // patch_size) ** 2 # Class Token: 一个可学习的向量 self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) # Position Embedding: 序列长度是 num_patches + 1,因为要算上 Class Token self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) self.pos_drop = nn.Dropout(p=0.1) self.blocks = nn.ModuleList([ TransformerBlock(embed_dim, num_heads, mlp_ratio) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) # (B, num_patches, embed_dim) # 把 Class Token 拼到序列前面 cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat((cls_tokens, x), dim=1) # (B, num_patches + 1, embed_dim) # 加位置编码 x = x + self.pos_embed x = self.pos_drop(x) for blk in self.blocks: x = blk(x) # 取 Class Token 对应的输出 x = self.norm(x) cls_out = x[:, 0] return self.head(cls_out)这里有几个容易踩坑的细节,我给你拆开说。
第一,cls_token的初始化我用了torch.zeros而不是随机初始化。实际训练中。两种方式差别不大,因为后续的Dropout和Transformer层会迅速打破对称性,所以zeros是完全可行的,这也是很多官方实现的做法。
第二,cls_token.expand(B, -1, -1)只是扩展维度,不会复制数据,所以额外显存开销可以忽略不计。torch.cat之后序列长度从64变成65,Position Embedding的维度必须对应改成65。如果你改了patch_size或者输入尺寸,这个数字很容易对不上,报错时会提醒你dimension mismatch,很多人第一次写会在这里卡住。
第三,我把位置编码初始化为0。源码里经常用截断正态分布来初始化,但在实践中,位置编码的信号在学习初期会被输入的embedding信号淹没,之后随着训练逐步调整。如果你的初始化方差过大,反而可能拖慢收敛速度。
3.3 自注意力机制的代码实现与计算过程
现在到了整个ViT最核心的模块——Multi-Head Self-Attention(MSA)。公式大家都见过:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V
关键问题是:Q、K、V在代码里怎么算?多头又是什么意思?
class Attention(nn.Module): def __init__(self, embed_dim, num_heads, qkv_bias=True): super().__init__() self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.scale = self.head_dim ** -0.5 # 一个全连接同时生成 Q、K、V self.qkv = nn.Linear(embed_dim, embed_dim * 3, bias=qkv_bias) self.proj = nn.Linear(embed_dim, embed_dim) def forward(self, x): B, N, C = x.shape # 生成 QKV,并拆成三个张量 qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] # 每个都是 (B, num_heads, N, head_dim) # 注意力分数: (B, num_heads, N, N) attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) # 加权求和 x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) return x这段代码我建议你对着张量形状一行行看。self.qkv(x)输出的维度是(B, N, C*3),然后通过reshape把最后一维拆成3份,对应Q、K、V。permute的作用是把维度顺序调整为(3, B, num_heads, N, head_dim),这样索引0、1、2分别就是Q、K、V。
为什么要除以sqrt(head_dim)?这是为了保证Q和K点积之后的结果方差维持在1左右,避免softmax在输入较大时梯度消失。你也许注意到这里用的是head_dim而不是整个embed_dim,因为每个头参与计算的是head_dim维的向量。
多头注意力的本质,是在同一个序列上并行运行多组不同的注意力。每个头有自己独立的Q、K、V投影,这意味着不同的头可以关注不同的位置关系——有的头可能偏向关注相邻patch,有的头可能偏向关注全局颜色分布,有的头则专门抓边界纹理。8个头就有8种不同的“视角”。
三维可视化一下这个过程:对第h个头,输入的token序列是(B, N, head_dim),Q和K做点积得到(B, N, N)的注意力矩阵,第i行第j列表示第i个token对第j个token的关注度。然后softmax归一化,确保一行加起来等于1。最后用这个权重矩阵去加权V,得到每个token的新表示。这个新表示里,第i个token的信息就是“整个序列所有token按注意力权重融合”的结果。
3.4 MLP与残差连接的细节处理
除了自注意力,Transformer Block里还有MLP、LayerNorm和残差连接。这里有一个ViT跟原始Transformer不一样的地方,值得专门讲一下:ViT用的是Pre-LN结构,也就是先做LayerNorm再进自注意力模块。
class TransformerBlock(nn.Module): def __init__(self, embed_dim, num_heads, mlp_ratio=4.0, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(embed_dim) self.attn = Attention(embed_dim, num_heads) self.norm2 = nn.LayerNorm(embed_dim) self.mlp = MLP(embed_dim, int(embed_dim * mlp_ratio), dropout) def forward(self, x): x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x残差连接的目的是让梯度可以顺畅地跨层传播。但Pre-LN和Post-LN有个重要区别:Post-LN(原始Transformer)把LayerNorm放在残差相加之后,训练深层Transformer时容易不稳定,需要 careful的warmup策略;Pre-LN把LayerNorm放在残差之前,梯度能更直接地反传,训练更稳定,对warmup的需求也更低。ViT的实现统一采用Pre-LN,好处是更稳,坏处是表征能力略有一点损失——不过在实践里这个损失基本可以忽略。
MLP部分相对简单,就是一个两层的全连接网络,中间跟一个GELU激活函数。为什么用GELU而不是ReLU?GELU在负区间不硬截断,而是平滑过渡,训练更稳定,这在Transformer类模型里已经成了事实标准。
class MLP(nn.Module): def __init__(self, in_features, hidden_features, dropout=0.1): super().__init__() self.fc1 = nn.Linear(in_features, hidden_features) self.act = nn.GELU() self.fc2 = nn.Linear(hidden_features, in_features) self.drop = nn.Dropout(dropout) def forward(self, x): x = self.fc1(x) x = self.act(x) x = self.drop(x) x = self.fc2(x) x = self.drop(x) return xMLP的hidden维度一般是embed_dim的4倍。这个比例来自Transformer论文的实验结论,4倍在计算量和表示能力之间取得了不错的平衡。再大效果提升不明显,训练成本反而涨得很厉害。
3.5 完整模型组装与参数量计算
把前面所有模块拼起来,就得到了我们完整的ViT模型。让我把每一层的维度变化和参数量算一下,这样你在训练时对模型规模能有个底。
def build_vit(image_size=32, patch_size=4, num_classes=10): model = ViT( image_size=image_size, patch_size=patch_size, num_classes=num_classes, embed_dim=192, depth=6, num_heads=3, mlp_ratio=4.0, ) return model model = build_vit() total_params = sum(p.numel() for p in model.parameters()) trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"Total params: {total_params / 1e6:.2f}M") print(f"Trainable params: {trainable_params / 1e6:.2f}M")在embed_dim=192、depth=6、patch_size=4、输入32x32的配置下,模型大概有几百万参数,跟一个小型CNN差不多。但要注意,参数量分布很不均匀:6层Transformer占了绝大部分,Patch Embedding和分类头的占比很小。
如果你的输入图片更大,比如224x224、patch_size=16,每个patch展平是768维,embed_dim通常也设为768或更大的值,参数量会直接涨到几千万甚至上亿。这就是为什么ViT模型动辄几百MB的原因——参数绝大部分都在Transformer Encoder里。
4. 训练细节与超参数选择实战
4.1 优化器、学习率与warmup策略
ViT训练跟CNN有个很大不同:它对优化器和学习率更敏感。如果用常规的SGD,训练速度会比较感人。ViT官方实现和社区复现普遍推荐AdamW优化器,并且配合warmup和cosine学习率衰减。
import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR epochs = 100 warmup_epochs = 5 optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.05) warmup_scheduler = LinearLR(optimizer, start_factor=0.1, end_factor=1.0, total_iters=warmup_epochs) cosine_scheduler = CosineAnnealingLR(optimizer, T_max=epochs - warmup_epochs, eta_min=1e-5) scheduler = SequentialLR(optimizer, schedulers=[warmup_scheduler, cosine_scheduler], milestones=[warmup_epochs])warmup的作用很关键。Transformer在训练初期如果直接用较大学习率,LayerNorm和注意力机制很容易产生震荡。先用小学习率跑几个epoch,让网络对数据分布有基本认识,再逐渐加大学习率进入正式训练阶段,这是ViT能稳定收敛的重要前置条件。
学习率本身的选择也值得斟酌。1e-3搭配batch_size=128在我的配置下表现不错。如果你的batch_size翻倍到256,学习率可以适当调到1.2e-3到1.5e-3,前提是你用AdamW或类似的自适应优化器——这种线性缩放经验在Transformer训练里比CNN更实用。weight_decay设置为0.05,这是ViT官方默认值,对控制过拟合有帮助。
4.2 训练循环完整代码
训练循环本身不复杂,关键是记得把模型切到train和eval两种模式,并且保证每个epoch都验证一下测试集准确率。下面我给出一个完整的训练代码框架。
def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss, correct, total = 0.0, 0, 0 for images, labels in loader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) _, preds = outputs.max(1) correct += preds.eq(labels).sum().item() total += images.size(0) return total_loss / total, 100.0 * correct / total @torch.no_grad() def evaluate(model, loader, criterion, device): model.eval() total_loss, correct, total = 0.0, 0, 0 for images, labels in loader: images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) total_loss += loss.item() * images.size(0) _, preds = outputs.max(1) correct += preds.eq(labels).sum().item() total += images.size(0) return total_loss / total, 100.0 * correct / total device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = build_vit().to(device) criterion = nn.CrossEntropyLoss() best_acc = 0.0 for epoch in range(epochs): train_loss, train_acc = train_one_epoch(model, trainloader, optimizer, criterion, device) test_loss, test_acc = evaluate(model, testloader, criterion, device) scheduler.step() if test_acc > best_acc: best_acc = test_acc torch.save(model.state_dict(), "best_vit_cifar10.pth") if (epoch + 1) % 10 == 0: print(f"Epoch [{epoch+1}/{epochs}] " f"Train Loss: {train_loss:.4f} Train Acc: {train_acc:.2f}% " f"Test Loss: {test_loss:.4f} Test Acc: {test_acc:.2f}% " f"Best Acc: {best_acc:.2f}%")这个训练循环里面的每个细节都有讲究。model.train()和model.eval()切换的是Dropout和LayerNorm的行为——Dropout在训练时随机丢弃、在评估时保持全量;LayerNorm在训练时用当前batch的统计数据,在评估时用运行时的统计均值。不切换的话,评估结果会忽高忽低,尤其是小的batch size时更明显。
zero_grad()在每个batch开始前清零梯度。这一步漏了的话,梯度会跨batch累积,loss直接崩掉。@torch.no_grad()告诉PyTorch不要构建计算图,推理时省大量显存和内存。
4.3 数据增强与正则化经验
在CIFAR-10这种相对小的数据集上训练ViT,数据增强和正则化是能否work的关键。ViT没有CNN那种先验偏置,如果不做增强,非常容易过拟合。我常用的一套增强组合如下:
transform_train = transforms.Compose([ transforms.RandomCrop(32, padding=4), transforms.RandomHorizontalFlip(), transforms.RandomApply([transforms.ColorJitter(0.4, 0.4, 0.4, 0.1)], p=0.8), transforms.RandomGrayscale(p=0.2), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ])这套增强的核心思想是让模型不能仅仅依赖颜色或简单的位置先验,而是必须学到更鲁棒的形状和纹理特征。RandomCrop和RandomHorizontalFlip是基础,ColorJitter可以扰动亮度、对比度、饱和度,RandomGrayscale则变相增强模型对颜色缺失的鲁棒性。
除了这些,Mixup和CutMix这类增强策略对ViT也有明显的正则化效果。它们本质上是把两张图的输入和标签都做线性插值,强迫模型学习更平滑的决策边界。如果你训练集不大、又追求更高精度,这两招值得一试,在timm库里都有现成实现可以直接调用。
5. 常见问题与踩坑实录
5.1 训练不收敛或收敛过慢的排查思路
我自己的经验是,跑ViT遇到的绝大多数问题都集中在以下几类,这里整理成一个排查清单:
| 问题现象 | 可能原因 | 排查/解决方案 |
|---|---|---|
| Loss不下降 | 学习率过大或过小 | 先用不同学习率做小规模实验,搭配warmup |
| 训练震荡严重 | 学习率太大 / 没做warmup | 降低初始学习率,增加warmup轮数 |
| 验证集准确率很低 | 模型太浅 / patch_size过大 | 适当增加depth,减小patch_size |
| 过拟合(训练好、验证差) | 数据集太小 / 增强不足 | 加数据增强,增大weight_decay和Dropout |
| 显存不足(OOM) | batch_size过大 / 序列过长 | 减小batch_size,减小输入分辨率或用梯度累积 |
| 准确率提升极慢 | 位置编码未生效 | 检查pos_embed是否正确加到每个token上 |
| 训练和测试时结果差异大 | Dropout/LayerNorm模式切换错误 | 确认train/eval模式是否正确调用 |
这里我想特别强调两个容易困扰新手的点。
第一个是关于Class Token的取法。有些实现会在最后对所有token做全局均值池化,跟取Class Token两个结果差别不算大,但如果你代码里取了x[:, 0](Class Token),却在Transformer之后忘了做LayerNorm,或者错把整个序列都过了一层nn.Linear,分类效果就会明显下滑。Classifier放在最后,输入维度必须是embed_dim不是num_patches + 1,这个维度对错了直接报错。
第二个是输入图片尺寸必须能被patch_size整除。比如输入是224x224,patch_size是16,刚好整除。但如果你临时换了数据集,比如图像是128x192,patch_size还是16,那就会出问题。代码里写了assert,但实际使用中建议根据数据集灵活调整patch_size,或者用插值先缩放图片。
5.2 推理时输出特征可视化的简单实践
虽然这次是代码解析,但理解模型学到的模式很能帮助调参。简单来说,可以取Transformer某个中间层(比如第一层)的注意力权重,然后用热力图形式可视化。
def visualize_attention(model, image, layer_idx=0, head_idx=0): model.eval() x = image.unsqueeze(0) # 拿到指定层 block = model.blocks[layer_idx] def hook_fn(module, input, output): global saved_attn # input[0] 是 norm1 之后的 x saved_attn = module.attn.attn.detach() # 注册 forward hook 来获取注意力矩阵 handle = block.attn.register_forward_hook(hook_fn) with torch.no_grad(): pred = model(x) handle.remove() # saved_attn 的维度是 (B, num_heads, N, N) attn_map = saved_attn[0, head_idx, 0, 1:].reshape(num_patches_side, num_patches_side) return attn_map当然这个实现依赖对Attention模块内部属性名的修改,如果你根据自己的代码结构调整,需要灵活适配一下,但思路是一样的:注册forward hook,在forward过程中把注意力矩阵捞出来,然后画热力图。
通过热力图你能直观看到:浅层attention往往关注局部结构(相邻patch相关性高),深层attention则可能学会跨距离的语义关联。如果发现浅层注意力看起来完全是均匀分布(每个位置都一样),那大概率是模型没有学到有效的位置信息,需要检查Position Embedding是否加了、学习率是否合适。
5.3 显存不足与训练速度优化
如果你在更大的分辨率或更大的模型上训练ViT,显存问题会非常突出。几个常用的减负手段分享给大家。
第一个是梯度累积。如果理想batch_size是128但显存只够放32,那就把batch_size设为32,每4个step做一次参数更新,效果基本等价。
accumulation_steps = 4 optimizer.zero_grad() for step, (images, labels) in enumerate(trainloader): images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) / accumulation_steps loss.backward() if (step + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()第二个是使用混合精度训练。PyTorch自带torch.cuda.amp,只需改动几行,训练速度通常能提升一倍不止,显存占用也明显下降。在支持Tensor Core的显卡上尤其明显。新手刚开始可能不想碰这块,但等你在更大的数据集上跑ViT时,这是绕不开的利器。
第三个是多卡并行。在单机多卡环境下,用torch.nn.DataParallel先做个最简单版本,也能有不小提升。不过要注意:PyTorch 2.x里DataParallel的调度开销比DistributedDataParallel大不少,如果卡数多或者模型大,建议直接上DDP。
6. 从零到一复现ViT的经验总结
我个人在实际操作中最大的体会是:ViT的代码门槛不在Transformer本身,而在于把图像转换成序列这个思维转变。只要把Patch Embedding、Class Token、Position Embedding这三块想明白,后面的Attention和MLP都是标准组件,照着NLP里成熟的实现搬过来就行。
最后再分享一个小技巧。在CIFAR-10上调试ViT时,建议先不急着上完整的数据增强和6层Transformer。先跑一个mini版本:depth=2、embed_dim=96、只做RandomCrop和HorizontalFlip,看看能不能在当前配置下过拟合训练集。如果连训练集都过拟合不了,说明是代码逻辑有问题;如果能过拟合但验证集很差,说明是数据增强和正则化不够。这种“先求过拟合,再谈泛化”的调试顺序,能帮你快速定位问题出在哪一层。
如果你后面打算在ImageNet这种大规模数据上使用ViT,我建议直接基于timm库做二次开发,里面的ViT实现经过充分验证、支持各种变体(DeiT、Swin Transformer等),比自己从头写稳妥得多。但在此之前,非常推荐先按这篇文章的思路手写一遍——只有当你能把每一步的维度变化和每个模块的输入输出都烂熟于心时,调参和改结构才能真正做到心里有数。