从NLP到CV:手把手教你用PyTorch复现ViT(Vision Transformer)图像分类模型
从NLP到CV:手把手教你用PyTorch复现ViT(Vision Transformer)图像分类模型
在计算机视觉领域,卷积神经网络(CNN)长期占据主导地位,但Transformer架构的崛起正在改变这一格局。本文将带你从零开始实现Vision Transformer(ViT)模型,这是一个将自然语言处理中成功的Transformer架构直接应用于图像分类任务的创新方法。不同于传统CNN,ViT将图像分割为固定大小的块(patches),通过自注意力机制捕捉全局依赖关系,无需依赖局部感受野的渐进式扩展。
对于中高级开发者和学生而言,理解并实现ViT模型不仅能掌握前沿技术,还能深入体会自注意力机制在视觉任务中的应用。我们将使用PyTorch框架,从模型结构解析到完整训练流程,逐步构建一个可在CIFAR-10/100等常见数据集上运行的ViT实现。以下是本文的核心路线图:
- 模型架构拆解 :理解图像块嵌入、位置编码和Transformer编码器的协同工作
- PyTorch实现细节 :从张量操作到注意力矩阵计算的手写实现
- 训练技巧 :学习率调度、混合精度训练和梯度裁剪的实战应用
- 迁移学习 :如何加载预训练权重并适配自定义数据集
1. ViT模型架构深度解析
Vision Transformer的核心思想是将图像视为一系列局部块的集合,每个块经过线性投影后作为Transformer的输入token。这种处理方式与NLP中将句子分解为单词token的做法高度相似。下面我们拆解ViT-B/16(基础版本,块大小为16×16)的具体实现。
1.1 图像块嵌入与位置编码
传统CNN通过滑动窗口处理像素的局部邻域,而ViT首先将输入图像分割为N个非重叠的块。对于224×224的RGB图像和16×16的块大小,我们得到:
num_patches = (224 / 16) ** 2 = 196
每个块展开为16×16×3=768维向量,经过可训练的线性投影(全连接层)映射到模型维度D(通常为768)。这个过程可以用以下公式表示:
# 伪代码示意
class PatchEmbed(nn.Module):
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
super().__init__()
self.proj = nn.Conv2d(in_chans, embed_dim,
kernel_size=patch_size,
stride=patch_size)
def forward(self, x):
x = self.proj(x) # (B, 768, 14, 14)
x = x.flatten(2) # (B, 768, 196)
x = x.transpose(1, 2) # (B, 196, 768)
return x
位置编码是ViT理解图像空间结构的关键。与CNN不同,Transformer本身不具备位置感知能力,因此需要显式添加位置信息。ViT使用可学习的一维位置编码,其实现如下:
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
注意:实际实现中会包含一个额外的分类token(class token),因此位置编码的序列长度为num_patches + 1
1.2 Transformer编码器结构
ViT的编码器由L个相同的层堆叠而成(基础版本为12层),每层包含以下关键组件:
- 多头自注意力(MSA) :计算不同图像块之间的关系权重
- 前馈网络(MLP) :对每个位置的特征进行非线性变换
- 层归一化(LayerNorm) :稳定训练过程
- 残差连接 :缓解梯度消失问题
多头注意力的计算过程可以表示为:
Attention(Q, K, V) = softmax(QK^T/√d_k)V
其中Q、K、V分别是通过线性变换从输入得到的查询、键和值矩阵。PyTorch实现时可以使用优化过的 nn.MultiheadAttention 模块,但为了更好理解机制,我们展示一个简化实现:
class MultiHeadSelfAttention(nn.Module):
def __init__(self, dim, num_heads=8):
super().__init__()
self.num_heads = num_heads
self.scale = (dim // num_heads) ** -0.5
self.to_qkv = nn.Linear(dim, dim * 3)
self.proj = nn.Linear(dim, dim)
def forward(self, x):
B, N, C = x.shape
qkv = self.to_qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads)
q, k, v = qkv.unbind(2) # (B, H, N, D)
attn = (q @ k.transpose(-2, -1)) * self.scale
attn = attn.softmax(dim=-1)
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
return self.proj(x)
2. PyTorch完整实现指南
现在我们将上述组件整合为一个完整的ViT模型。以下实现支持灵活配置模型尺寸,适用于不同计算资源场景。
2.1 模型主体结构
class VisionTransformer(nn.Module):
def __init__(self, img_size=224, patch_size=16, in_chans=3,
num_classes=1000, embed_dim=768, depth=12,
num_heads=12, mlp_ratio=4.):
super().__init__()
self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim)
num_patches = self.patch_embed.num_patches
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
self.blocks = nn.ModuleList([
Block(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, N, D)
cls_tokens = self.cls_token.expand(B, -1, -1)
x = torch.cat((cls_tokens, x), dim=1) # (B, N+1, D)
x = x + self.pos_embed
for blk in self.blocks:
x = blk(x)
x = self.norm(x)
return self.head(x[:, 0]) # 仅使用分类token进行分类
2.2 关键超参数配置
下表展示了不同规模ViT模型的典型配置:
| 模型变体 | 层数 | 隐藏层维度 | MLP大小 | 头数 | 参数量 |
|---|---|---|---|---|---|
| ViT-Tiny | 12 | 192 | 768 | 3 | 5.7M |
| ViT-Small | 12 | 384 | 1536 | 6 | 22M |
| ViT-Base | 12 | 768 | 3072 | 12 | 86M |
| ViT-Large | 24 | 1024 | 4096 | 16 | 307M |
提示:对于CIFAR-10等小规模数据集,建议使用ViT-Tiny或ViT-Small以避免过拟合
3. 训练流程与优化技巧
ViT训练需要特别注意学习率调度和正则化策略。以下是我们推荐的训练配置:
3.1 数据预处理与增强
针对图像分类任务,标准的数据增强策略包括:
from torchvision import transforms
train_transform = transforms.Compose([
transforms.RandomResizedCrop(224),
transforms.RandomHorizontalFlip(),
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
3.2 优化器配置
AdamW优化器配合余弦退火学习率调度表现优异:
optimizer = torch.optim.AdamW(
model.parameters(),
lr=3e-4,
weight_decay=0.05
)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer,
T_max=epochs,
eta_min=1e-6
)
3.3 混合精度训练
使用AMP(自动混合精度)可以显著减少显存占用并加速训练:
scaler = torch.cuda.amp.GradScaler()
for inputs, labels in train_loader:
optimizer.zero_grad()
with torch.cuda.amp.autocast():
outputs = model(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
4. 迁移学习实践
ViT通常在��规模数据集(如ImageNet-21k)上预训练,然后迁移到下游任务。以下是微调的关键步骤:
- 加载预训练权重 :
model.load_state_dict(torch.load('vit_base_patch16_224.pth'))
- 替换分类头 :
model.head = nn.Linear(model.head.in_features, new_num_classes)
- 分层学习率 :
param_groups = [
{'params': model.patch_embed.parameters(), 'lr': base_lr*0.1},
{'params': model.pos_embed, 'lr': base_lr*0.5},
{'params': model.cls_token, 'lr': base_lr},
{'params': model.blocks.parameters(), 'lr': base_lr},
{'params': model.head.parameters(), 'lr': base_lr*5}
]
在实际项目中,我发现ViT对学习率非常敏感。在CIFAR-10上微调时,将基础学习率设为3e-5,配合早停策略(patience=5)通常能获得最佳结果。另外,添加CutMix或MixUp等增强技术可以进一步提升模型鲁棒性,特别是在小数据集场景下。
更多推荐


所有评论(0)