机器学习基础(十五):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}\):
- 分块:每个 Patch 大小 \(P \times P\)
- 展平:每个 Patch 变成 \(P^2 \cdot C\) 维向量
- 线性投影:映射到 \(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 下一步学习建议
- Swin Transformer:窗口注意力与多尺度特征
- CLIP:视觉-语言预训练
- Diffusion Models:基于Transformer的图像生成
- 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

浙公网安备 33010602011771号