从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层),每层包含以下关键组件:

  1. 多头自注意力(MSA) :计算不同图像块之间的关系权重
  2. 前馈网络(MLP) :对每个位置的特征进行非线性变换
  3. 层归一化(LayerNorm) :稳定训练过程
  4. 残差连接 :缓解梯度消失问题

多头注意力的计算过程可以表示为:

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)上预训练,然后迁移到下游任务。以下是微调的关键步骤:

  1. 加载预训练权重
model.load_state_dict(torch.load('vit_base_patch16_224.pth'))
  1. 替换分类头
model.head = nn.Linear(model.head.in_features, new_num_classes)
  1. 分层学习率
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等增强技术可以进一步提升模型鲁棒性,特别是在小数据集场景下。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐