机器学习基础(十五):Vision Transformer

一、引言

上一篇我们学习了 Transformer 在自然语言处理中的威力。2020年,Google 发表论文《An Image is Worth 16x16 Words》,将 Transformer 引入计算机视觉——Vision Transformer(ViT)诞生了。

这一突破打破了 CNN 在视觉领域长达十年的垄断,证明了:

纯注意力机制也能在图像任务上取得SOTA

本文目标:

  • 理解 ViT 的核心思想:图像即序列
  • 掌握 Patch Embedding 和位置编码
  • 从零实现 ViT
  • 对比 ViT 与 CNN 的优劣

二、为什么图像需要 Transformer?

2.1 CNN 的局限

CNN 通过卷积核提取局部特征,层层堆叠获得全局感知:

问题 说明
局部感受野 需要多层才能看到全局信息
归纳偏置 平移等变性是优势也是限制
难以建模长距离依赖 远距离像素关系需要很多层传递

2.2 ViT 的核心洞察

如果把图像切成小块(Patch),按顺序排列,不就是一个序列吗?

图像 [224×224×3]
    ↓ 分块(16×16)
14×14=196 个 Patch
    ↓ 展平 + 线性投影
196 个 token,每个 768 维
    ↓ + 位置编码 + CLS
输入 Transformer Encoder

关键优势:

  • 自注意力直接建模任意两个Patch的关系
  • 没有卷积的归纳偏置,数据驱动学习
  • 可以无缝迁移 NLP 的预训练技术

三、图像分块(Patch Embedding)

3.1 分块策略

假设输入图像 \(x \in \mathbb{R}^{H \times W \times C}\):

  1. 分块:每个 Patch 大小 \(P \times P\)
  2. 展平:每个 Patch 变成 \(P^2 \cdot C\) 维向量
  3. 线性投影:映射到 \(D\) 维(模型维度)

计算:

  • Patch 数量:\(N = \frac{H}{P} \times \frac{W}{P}\)
  • 输出维度:\([N, D]\)

3.2 PyTorch 实现

import torch
import torch.nn as nn

class PatchEmbedding(nn.Module):
    def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768):
        super().__init__()
        self.img_size = img_size
        self.patch_size = patch_size
        self.num_patches = (img_size // patch_size) ** 2
        
        # 使用卷积实现分块 + 投影(高效)
        self.proj = nn.Conv2d(
            in_channels, 
            embed_dim, 
            kernel_size=patch_size, 
            stride=patch_size
        )
        
    def forward(self, x):
        """
        x: [batch, channels, height, width]
        output: [batch, num_patches, embed_dim]
        """
        x = self.proj(x)  # [batch, embed_dim, H/P, W/P]
        x = x.flatten(2)  # [batch, embed_dim, num_patches]
        x = x.transpose(1, 2)  # [batch, num_patches, embed_dim]
        return x

# 测试
x = torch.randn(2, 3, 224, 224)
patch_embed = PatchEmbedding()
patches = patch_embed(x)
print(f"输入图像: {x.shape}")
print(f"Patch Embedding: {patches.shape}")  # [2, 196, 768]
print(f"Patch 数量: {patch_embed.num_patches}")  # 196

四、ViT 整体架构

4.1 完整流程

输入图像 [B, 3, 224, 224]
    ↓
┌─────────────────────────────────────────┐
│           Patch Embedding               │
│   Conv2d(3, 768, kernel=16, stride=16)  │
└─────────────────────────────────────────┘
    ↓
[B, 196, 768]  (196=14×14 个 Patch)
    ↓
┌─────────────────────────────────────────┐
│  + CLS Token [B, 1, 768]                │
│  + Position Embedding [197, 768]        │
└─────────────────────────────────────────┘
    ↓
[B, 197, 768]  (197 = 196 patches + 1 CLS)
    ↓
┌─────────────────────────────────────────┐
│      Transformer Encoder (×L)           │
│  - Multi-Head Self-Attention            │
│  - MLP (Feed Forward)                   │
│  - Layer Norm + Residual                │
└─────────────────────────────────────────┘
    ↓
[B, 197, 768]
    ↓
取 CLS Token 输出 [B, 768]
    ↓
MLP Head (分类层) [B, num_classes]

4.2 与 NLP Transformer 的对比

特性 NLP Transformer ViT
输入 词序列 图像 Patch 序列
Embedding Word Embedding Patch Embedding (Conv2d)
位置编码 正弦/可学习 可学习 1D/2D
特殊 Token [CLS] 可选 [CLS] 必须(分类)
预训练 大规模文本 需要大规模图像 (JFT-300M)

五、位置编码的变体

5.1 1D 位置编码

最简单的方式:给每个 Patch 一个可学习的向量

self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))

5.2 2D 位置编码

考虑图像的二维结构,分别学习行和列的位置编码:

# 分别学习行和列的位置编码
self.pos_embed_x = nn.Parameter(torch.zeros(1, num_patches_x, embed_dim // 2))
self.pos_embed_y = nn.Parameter(torch.zeros(1, num_patches_y, embed_dim // 2))

5.3 相对位置编码

Swin Transformer 使用相对位置编码,更适合局部窗口注意力。

5.4 实现

class ViT(nn.Module):
    def __init__(
        self,
        img_size=224,
        patch_size=16,
        in_channels=3,
        num_classes=1000,
        embed_dim=768,
        depth=12,
        num_heads=12,
        mlp_ratio=4.0,
        dropout=0.1
    ):
        super().__init__()
        
        # Patch Embedding
        self.patch_embed = PatchEmbedding(
            img_size, patch_size, in_channels, embed_dim
        )
        num_patches = self.patch_embed.num_patches
        
        # CLS Token
        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.dropout = nn.Dropout(dropout)
        
        # Transformer Encoder
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim,
            nhead=num_heads,
            dim_feedforward=int(embed_dim * mlp_ratio),
            dropout=dropout,
            activation='gelu',
            batch_first=True
        )
        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=depth)
        
        # 分类头
        self.norm = nn.LayerNorm(embed_dim)
        self.head = nn.Linear(embed_dim, num_classes)
        
        # 初始化
        nn.init.normal_(self.cls_token, std=0.02)
        nn.init.normal_(self.pos_embed, std=0.02)
        
    def forward(self, x):
        batch_size = x.shape[0]
        
        # Patch Embedding
        x = self.patch_embed(x)  # [B, N, D]
        
        # 添加 CLS Token
        cls_tokens = self.cls_token.expand(batch_size, -1, -1)  # [B, 1, D]
        x = torch.cat([cls_tokens, x], dim=1)  # [B, N+1, D]
        
        # 添加位置编码
        x = x + self.pos_embed
        x = self.dropout(x)
        
        # Transformer
        x = self.transformer(x)  # [B, N+1, D]
        
        # 取 CLS Token 输出分类
        x = self.norm(x[:, 0])  # [B, D]
        x = self.head(x)  # [B, num_classes]
        
        return x

# 测试
model = ViT(img_size=224, patch_size=16, num_classes=10)
x = torch.randn(2, 3, 224, 224)
out = model(x)
print(f"ViT 输出: {out.shape}")  # [2, 10]

六、CLS Token 与全局表示

6.1 为什么需要 CLS Token?

BERT 中使用 [CLS] token 的表示作为句子级别的特征,ViT 借鉴了这一设计:

  • CLS Token 与所有 Patch 都有注意力交互
  • 通过自注意力聚合全局信息
  • 比全局平均池化更灵活

6.2 替代方案:全局平均池化

# 不使用 CLS Token,对所有 Patch 取平均
x = x.mean(dim=1)  # [B, D]

实验表明两者性能相近,但 CLS Token 更便于与 NLP 预训练模型统一。


七、ViT 训练技巧

7.1 大数据预训练是关键

数据集大小 ImageNet-1k (1.3M) JFT-300M (300M)
ViT-Base 77.9% 84.2%
ResNet-50 76.2% -

结论: ViT 需要大规模数据才能超越 CNN。

7.2 知识蒸馏(DeiT)

Data-efficient Image Transformer (DeiT) 使用知识蒸馏让小数据集也能训练好 ViT:

# 蒸馏损失 = 软标签损失 + 硬标签损失
distillation_loss = soft_loss(student_logits, teacher_logits)
classification_loss = hard_loss(student_logits, labels)
total_loss = 0.9 * distillation_loss + 0.1 * classification_loss

7.3 数据增强

from torchvision import transforms

train_transform = transforms.Compose([
    transforms.RandomResizedCrop(224),
    transforms.RandomHorizontalFlip(),
    transforms.ColorJitter(0.4, 0.4, 0.4),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406],
                        std=[0.229, 0.224, 0.225])
])

八、ViT vs CNN 对比实验

8.1 CIFAR-10 完整训练代码

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

# 数据准备
train_transform = transforms.Compose([
    transforms.Resize(224),
    transforms.RandomHorizontalFlip(),
    transforms.ToTensor(),
    transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
])

test_transform = transforms.Compose([
    transforms.Resize(224),
    transforms.ToTensor(),
    transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
])

train_dataset = datasets.CIFAR10(
    root='./data', train=True, download=True, transform=train_transform
)
test_dataset = datasets.CIFAR10(
    root='./data', train=False, download=True, transform=test_transform
)

train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2)
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers=2)

# 训练函数
def train_epoch(model, loader, optimizer, criterion, device):
    model.train()
    total_loss = 0
    correct = 0
    total = 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()
        _, predicted = outputs.max(1)
        total += labels.size(0)
        correct += predicted.eq(labels).sum().item()
    
    return total_loss / len(loader), 100. * correct / total

def evaluate(model, loader, criterion, device):
    model.eval()
    total_loss = 0
    correct = 0
    total = 0
    
    with torch.no_grad():
        for images, labels in loader:
            images, labels = images.to(device), labels.to(device)
            outputs = model(images)
            loss = criterion(outputs, labels)
            
            total_loss += loss.item()
            _, predicted = outputs.max(1)
            total += labels.size(0)
            correct += predicted.eq(labels).sum().item()
    
    return total_loss / len(loader), 100. * correct / total

# 训练 ViT
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

# 使用较小的 ViT 配置适配 CIFAR-10
vit_model = ViT(
    img_size=224,
    patch_size=16,
    in_channels=3,
    num_classes=10,
    embed_dim=384,      # 小模型
    depth=6,            # 6层
    num_heads=6,
    dropout=0.1
).to(device)

criterion = nn.CrossEntropyLoss()
optimizer = optim.AdamW(vit_model.parameters(), lr=1e-3, weight_decay=0.05)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50)

print("开始训练 ViT...")
for epoch in range(50):
    train_loss, train_acc = train_epoch(vit_model, train_loader, optimizer, criterion, device)
    test_loss, test_acc = evaluate(vit_model, test_loader, criterion, device)
    scheduler.step()
    
    if (epoch + 1) % 10 == 0:
        print(f"Epoch {epoch+1}: Train Acc={train_acc:.2f}%, Test Acc={test_acc:.2f}%")

8.2 对比结果

模型 参数量 CIFAR-10 准确率
ResNet-18 11M ~93%
ViT-Tiny 5.7M ~90% (需更多epoch)
ViT-Small 22M ~92%

观察:

  • ViT 收敛较慢,需要更多 epoch
  • 小数据集上 CNN 仍有优势
  • ViT 的优势在大规模数据上更明显

九、ViT 的变体

9.1 DeiT (Data-efficient Image Transformer)

  • 使用知识蒸馏减少对大数据的依赖
  • 在 ImageNet 上训练即可达到很好效果

9.2 Swin Transformer

  • 分层特征金字塔:多尺度特征
  • 移位窗口注意力:降低计算复杂度
  • 成为目标检测、分割的新基准
Swin Transformer
├── Patch Partition (4×4)
├── Stage 1: 嵌入 + 窗口注意力 (H/4 × W/4)
├── Stage 2: 合并 + 窗口注意力 (H/8 × W/8)
├── Stage 3: 合并 + 窗口注意力 (H/16 × W/16)
└── Stage 4: 合并 + 窗口注意力 (H/32 × W/32)

9.3 ConvNeXt

将 Transformer 的设计思想(大kernel、LayerNorm、GELU)应用到 CNN:

# ConvNeXt Block
class ConvNeXtBlock(nn.Module):
    def __init__(self, dim):
        super().__init__()
        self.dwconv = nn.Conv2d(dim, dim, kernel_size=7, padding=3, groups=dim)
        self.norm = nn.LayerNorm(dim)
        self.pwconv1 = nn.Linear(dim, 4 * dim)
        self.act = nn.GELU()
        self.pwconv2 = nn.Linear(4 * dim, dim)
    
    def forward(self, x):
        input = x
        x = self.dwconv(x)
        x = x.permute(0, 2, 3, 1)  # [B, H, W, C]
        x = self.norm(x)
        x = self.pwconv1(x)
        x = self.act(x)
        x = self.pwconv2(x)
        x = x.permute(0, 3, 1, 2)  # [B, C, H, W]
        return input + x

十、何时用 ViT,何时用 CNN?

10.1 选择指南

场景 推荐 原因
小数据集 (<10K) CNN 归纳偏置帮助泛化
大数据集 (>1M) ViT 全局注意力优势显现
需要快速推理 CNN 计算效率更高
需要可解释性 ViT 注意力可视化直观
多模态任务 ViT 与文本Transformer统一架构
边缘设备部署 CNN 模型更小、更快

10.2 混合架构

CoAtNet:结合卷积和注意力

CoAtNet = Conv + Transformer
├── 早期层:卷积(局部特征)
└── 后期层:Transformer(全局关系)

十一、总结

11.1 核心要点

Vision Transformer
├── 核心思想
│   └── 图像 = Patch 序列 (16×16 words)
├── 关键组件
│   ├── Patch Embedding: Conv2d 实现
│   ├── CLS Token: 全局表示
│   ├── Position Embedding: 可学习 1D/2D
│   └── Transformer Encoder: 标准结构
├── 训练要点
│   ├── 需要大数据预训练
│   ├── 知识蒸馏提升小数据表现
│   └── 强数据增强
├── 优势
│   ├── 全局感受野
│   ├── 与NLP统一架构
│   └── 可解释性强
└── 局限
    ├── 数据 hungry
    ├── 计算量大
    └── 小数据集不如CNN

11.2 学习路径回顾

机器学习基础系列
├── (一) 绪论
├── (二) 线性回归
├── ...
├── (十) 卷积神经网络CNN
├── (十一) 过拟合与正则化
├── (十二) 优化器大全
├── (十三) 循环神经网络RNN与LSTM
├── (十四) Transformer与注意力机制
└── (十五) Vision Transformer ← 你在这里
    └── 下一步: 多模态/目标检测/生成模型

11.3 下一步学习建议

  1. Swin Transformer:窗口注意力与多尺度特征
  2. CLIP:视觉-语言预训练
  3. Diffusion Models:基于Transformer的图像生成
  4. DETR:用Transformer做目标检测

附录:ViT 配置速查表

模型 Patch Layers Hidden Heads 参数量 ImageNet Top-1
ViT-Tiny 16×16 12 192 3 5.7M 75.4%
ViT-Small 16×16 12 384 6 22M 81.4%
ViT-Base 16×16 12 768 12 86M 84.2%
ViT-Large 16×16 24 1024 16 307M 85.2%

本文代码可在 GitHub 获取:https://github.com/paywqiao-max/ML-Basics

posted @ 2026-04-09 23:14  YZG5N  阅读(93)  评论(0)    收藏  举报