PyTorch实战:从零构建Vision Transformer图像分类模型
第一次看到Vision Transformer(ViT)将自然语言处理领域的Transformer成功迁移到计算机视觉任务时,我意识到这不仅仅是技术上的突破,更是一种思维方式的革新。传统的卷积神经网络(CNN)在图像处理领域统治多年后,ViT用完全不同的方式证明了自注意力机制在视觉任务中的强大潜力。本文将带你从第一行代码开始,完整实现一个可运行的ViT模型,并深入探讨每个关键组件的设计原理和实现技巧。
1. 环境准备与数据预处理
在开始构建ViT之前,我们需要确保开发环境配置正确。建议使用Python 3.8+和PyTorch 1.10+版本,这些版本对Transformer相关操作有更好的支持。
pip install torch torchvision einops matplotlib
ViT模型对输入图像有特定要求——通常需要将图像调整为224×224像素并分割成16×16的小块。以下是一个完整的图像预处理流程:
from torchvision import transforms
from PIL import Image
# 定义图像预处理流水线
transform = transforms.Compose([
transforms.Resize((224, 224)), # 调整图像尺寸
transforms.ToTensor(), # 转换为张量
transforms.Normalize( # 标准化
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
# 加载并预处理图像
img = Image.open("example.jpg")
x = transform(img).unsqueeze(0) # 添加batch维度
print(x.shape) # 输出: torch.Size([1, 3, 224, 224])
提示:在实际项目中,建议使用torchvision.datasets.ImageFolder或自定义Dataset类来批量处理数据,这里简化处理仅展示单张图像的预处理过程。
2. 核心组件实现
2.1 Patch Embedding层
ViT的核心创新之一是将图像视为一系列patch的序列。我们需要实现一个PatchEmbedding模块,将图像分割并转换为嵌入向量:
import torch
from torch import nn
from einops import rearrange
class PatchEmbedding(nn.Module):
def __init__(self, in_channels=3, patch_size=16, emb_size=768, img_size=224):
super().__init__()
self.patch_size = patch_size
self.projection = nn.Sequential(
# 使用卷积层替代线性层提升效率
nn.Conv2d(in_channels, emb_size, kernel_size=patch_size, stride=patch_size),
# 重排维度: [batch, emb_size, height/patch, width/patch] -> [batch, num_patches, emb_size]
nn.Flatten(2)
)
self.cls_token = nn.Parameter(torch.randn(1, 1, emb_size))
num_patches = (img_size // patch_size) ** 2
self.positions = nn.Parameter(torch.randn(num_patches + 1, emb_size))
def forward(self, x):
batch_size = x.shape[0]
x = self.projection(x) # [batch, emb_size, h, w] -> [batch, emb_size, num_patches]
x = x.transpose(1, 2) # [batch, num_patches, emb_size]
# 添加CLS token
cls_tokens = self.cls_token.expand(batch_size, -1, -1)
x = torch.cat([cls_tokens, x], dim=1)
# 添加位置编码
x += self.positions
return x
关键点解析:
- 卷积投影:使用卷积核大小和步长等于patch_size的卷积层,比原始论文中的线性投影更高效
- CLS Token:类似BERT的[CLS]标记,用于最终的分类任务
- 位置编码:可学习的位置编码,为模型提供空间信息
2.2 多头注意力机制
自注意力是Transformer的核心,让我们实现一个高效的多头注意力模块:
class MultiHeadAttention(nn.Module):
def __init__(self, emb_size=768, num_heads=8, dropout=0.1):
super().__init__()
self.emb_size = emb_size
self.num_heads = num_heads
self.head_dim = emb_size // num_heads
# 合并QKV投影以提升效率
self.qkv = nn.Linear(emb_size, emb_size * 3)
self.attn_drop = nn.Dropout(dropout)
self.proj = nn.Linear(emb_size, emb_size)
self.proj_drop = nn.Dropout(dropout)
def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape
# 生成QKV并分割多头
qkv = self.qkv(x).reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4) # [3, batch, heads, seq_len, head_dim]
q, k, v = qkv[0], qkv[1], qkv[2]
# 计算注意力分数
attn_scores = (q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5)
# 应用注意力掩码(可选)
if mask is not None:
attn_scores = attn_scores.masked_fill(mask == 0, float('-inf'))
attn_probs = self.attn_drop(torch.softmax(attn_scores, dim=-1))
# 注意力加权求和
output = (attn_probs @ v).transpose(1, 2).reshape(batch_size, seq_len, self.emb_size)
output = self.proj(output)
output = self.proj_drop(output)
return output
注意:实际应用中,可以使用PyTorch内置的nn.MultiheadAttention,但自定义实现有助于理解底层原理。
2.3 Transformer编码器块
将多头注意力和前馈网络组合成完整的Transformer编码器块:
class FeedForward(nn.Module):
def __init__(self, emb_size=768, expansion=4, dropout=0.1):
super().__init__()
self.net = nn.Sequential(
nn.Linear(emb_size, expansion * emb_size),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(expansion * emb_size, emb_size),
nn.Dropout(dropout)
)
def forward(self, x):
return self.net(x)
class TransformerBlock(nn.Module):
def __init__(self, emb_size=768, num_heads=8, dropout=0.1, expansion=4):
super().__init__()
self.norm1 = nn.LayerNorm(emb_size)
self.attn = MultiHeadAttention(emb_size, num_heads, dropout)
self.norm2 = nn.LayerNorm(emb_size)
self.ffn = FeedForward(emb_size, expansion, dropout)
def forward(self, x):
# 残差连接和层归一化
x = x + self.attn(self.norm1(x))
x = x + self.ffn(self.norm2(x))
return x
3. 完整ViT模型组装
现在我们可以将所有组件组合成完整的Vision Transformer模型:
class VisionTransformer(nn.Module):
def __init__(self, num_classes=1000, depth=12, emb_size=768, num_heads=12,
patch_size=16, img_size=224, in_channels=3):
super().__init__()
self.patch_embed = PatchEmbedding(in_channels, patch_size, emb_size, img_size)
self.encoder = nn.Sequential(*[
TransformerBlock(emb_size, num_heads) for _ in range(depth)
])
self.classifier = nn.Sequential(
nn.LayerNorm(emb_size),
nn.Linear(emb_size, num_classes)
)
def forward(self, x):
x = self.patch_embed(x)
x = self.encoder(x)
# 使用CLS token进行分类
cls_token = x[:, 0]
return self.classifier(cls_token)
4. 模型训练与调试技巧
4.1 学习率设置与优化器选择
ViT模型通常需要特定的训练策略才能达到最佳性能:
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingLR
model = VisionTransformer(num_classes=10) # 假设是10分类任务
optimizer = AdamW(model.parameters(), lr=3e-5, weight_decay=0.01)
scheduler = CosineAnnealingLR(optimizer, T_max=100) # 余弦退火学习率
criterion = nn.CrossEntropyLoss()
4.2 常见问题与解决方案
在实现和训练ViT过程中,可能会遇到以下典型问题:
-
训练不稳定
- 解决方案:使用梯度裁剪(
nn.utils.clip_grad_norm_) - 添加更多的层归一化
- 尝试更小的学习率
- 解决方案:使用梯度裁剪(
-
过拟合
- 增加Dropout比例
- 使用更强的数据增强
- 添加标签平滑(Label Smoothing)
-
内存不足
- 减小batch size
- 使用混合精度训练
- 梯度累积
# 混合精度训练示例
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for inputs, labels in dataloader:
optimizer.zero_grad()
with autocast():
outputs = model(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
scheduler.step()
4.3 可视化与调试工具
理解模型内部运作对调试至关重要:
# 可视化注意力权重
def visualize_attention(model, img):
model.eval()
with torch.no_grad():
embeddings = model.patch_embed(img)
attn_output, attn_weights = model.encoder[0].attn(
model.encoder[0].norm1(embeddings),
return_attention=True
)
# 绘制注意力热力图
import matplotlib.pyplot as plt
plt.imshow(attn_weights[0, 0].cpu().numpy()) # 第一个头的注意力
plt.colorbar()
plt.show()
5. 进阶优化与变体
5.1 高效ViT变体
原始ViT计算量较大,以下是几种改进方案:
| 变体名称 | 核心改进 | 计算量减少 |
|---|---|---|
| DeiT | 知识蒸馏+数据高效训练 | ~50% |
| Swin Transformer | 分层特征+滑动窗口注意力 | ~30% |
| MobileViT | 混合CNN-Transformer架构 | ~70% |
| CrossViT | 多尺度patch融合 | - |
5.2 混合架构设计
结合CNN和Transformer的优势:
class HybridViT(nn.Module):
def __init__(self):
super().__init__()
# CNN特征提取器
self.cnn_backbone = nn.Sequential(
nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
)
# ViT部分
self.vit = VisionTransformer(
img_size=56, # 经过CNN下采样后的尺寸
patch_size=7,
emb_size=256
)
def forward(self, x):
x = self.cnn_backbone(x) # [B, 64, 56, 56]
return self.vit(x)
5.3 迁移学习实践
使用预训练ViT进行迁移学习的典型流程:
- 加载预训练模型(如Google的ViT-B/16)
- 替换最后的分类头
- 选择性冻结部分层
- 使用较小的学习率微调
from torchvision.models import vit_b_16
# 加载预训练模型
pretrained_vit = vit_b_16(pretrained=True)
# 替换分类头
num_ftrs = pretrained_vit.heads.head.in_features
pretrained_vit.heads.head = nn.Linear(num_ftrs, 10) # 新任务有10类
# 仅训练分类头(可选)
for param in pretrained_vit.parameters():
param.requires_grad = False
for param in pretrained_vit.heads.parameters():
param.requires_grad = True
在完成这个ViT实现项目的过程中,最让我印象深刻的是模型对超参数的敏感性。与CNN不同,学习率、权重衰减和Dropout比例的微小变化都可能显著影响ViT的性能。建议在实际应用中从小的配置开始,逐步扩展模型规模,并保持耐心——ViT通常需要更长的训练时间才能展现出其真正的潜力。
&spm=1001.2101.3001.5002&articleId=154169730&d=1&t=3&u=ab14b3681a31474d8f670423922fedd4)
456

被折叠的 条评论
为什么被折叠?



