共计 2426 个字符,预计需要花费 7 分钟才能阅读完成。
背景介绍
在计算机视觉领域,卷积神经网络(CNN)长期以来占据主导地位。然而,CNN 存在一些固有局限性,例如局部感受野限制了全局信息的捕获,固定的卷积核难以适应不同尺度的特征。2020 年,Vision Transformer(ViT)的提出打破了这一局面,将自然语言处理中成功的 Transformer 架构引入视觉任务,实现了从局部到全局建模的跨越。

ViT 的核心创新在于完全摒弃了卷积操作,转而使用纯注意力机制处理图像数据。这种架构不仅能捕捉长距离依赖关系,还具有更强的可解释性——模型可以直观展示哪些图像区域获得了更多关注。
核心原理
自注意力机制在视觉中的应用
自注意力机制通过计算查询(Query)、键(Key)和值(Value)之间的关系来分配注意力权重。在视觉任务中:
- 每个图像块被视为一个 ” 词 ”
- 通过计算块间相似度得到注意力分布
- 最终输出是各块值的加权和
这种机制使模型能够动态关注与当前任务最相关的图像区域,无论它们在图像中的物理距离如何。
图像分块与位置编码
- 图像分块(Patches):
- 将输入图像划分为固定大小的非重叠块(如 16×16 像素)
- 每个块展平后通过线性投影得到特征向量
-
这些向量相当于 NLP 中的词嵌入
-
位置编码(Position Embedding):
- 由于 Transformer 本身不包含位置信息,需显式添加位置编码
- 常用可学习的位置向量,与图像块特征相加
- 确保模型理解图像的空间结构
代码实现
以下是一个精简版 ViT 的 PyTorch 实现(省略了部分辅助函数):
import torch
import torch.nn as nn
class PatchEmbedding(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, C, H, W] -> [B, E, H/P, W/P]
x = x.flatten(2).transpose(1, 2) # 展平为序列 [B, N, E]
return x
class VisionTransformer(nn.Module):
def __init__(self, num_classes=1000):
super().__init__()
self.patch_embed = PatchEmbedding()
self.pos_embed = nn.Parameter(torch.zeros(1, 196 + 1, 768)) # 可学习位置编码
self.cls_token = nn.Parameter(torch.zeros(1, 1, 768)) # 分类令牌
self.blocks = nn.TransformerEncoder(nn.TransformerEncoderLayer(d_model=768, nhead=12),
num_layers=12)
self.head = nn.Linear(768, num_classes)
def forward(self, x):
# 1. 分块嵌入
x = self.patch_embed(x) # [B, 196, 768]
# 2. 添加分类令牌和位置编码
cls_tokens = self.cls_token.expand(x.shape[0], -1, -1)
x = torch.cat((cls_tokens, x), dim=1) # [B, 197, 768]
x = x + self.pos_embed
# 3. 通过 Transformer 编码器
x = self.blocks(x)
# 4. 取分类令牌对应的输出做预测
x = x[:, 0]
x = self.head(x)
return x
实战建议
不同数据规模的调优策略
- 大数据集(100 万 + 图像):
- 直接使用原始 ViT 架构
- 可尝试更大的 patch 尺寸(如 32×32)减少计算量
-
学习率可适当增大(如 3e-4)
-
中等数据集(10 万 -100 万):
- 使用预训练模型微调
- 添加随机裁剪、颜色抖动等数据增强
-
考虑混合架构(如 CNN+ViT)
-
小数据集(<1 万):
- 优先考虑轻量级变体(如 DeiT)
- 冻结大部分 Transformer 层
- 使用强正则化(Dropout=0.5)
计算资源有限时的方案
- 减小模型尺寸:
- 减少 Transformer 层数(如从 12 层到 6 层)
- 降低嵌入维度(如从 768 到 512)
-
减少注意力头数
-
优化技巧:
- 使用梯度检查点(checkpointing)
- 混合精度训练
- 分布式数据并行
避坑指南
常见实现错误
- 位置编码错误 :
- 忘记添加分类令牌对应的位置编码
-
解决方案:确保位置编码维度为
[1, num_patches+1, embed_dim] -
注意力掩码问题 :
- 错误地将 padding 掩码应用于图像序列
- 解决方案:ViT 通常不需要注意力掩码
训练问题排查
- 损失不下降 :
- 检查学习率是否过小
- 验证输入图像是否正常归一化
-
确认位置编码是否正确添加
-
显存溢出 :
- 减小 batch size
- 使用梯度累积
- 尝试更小的 patch 尺寸
性能考量
与传统 CNN 对比
| 指标 | ViT | CNN |
|---|---|---|
| 计算效率 | 较高 FLOPs | 优化良好 |
| 内存占用 | 较大 | 较小 |
| 数据需求 | 需要大数据预训练 | 中等数据即可 |
| 长距离建模 | 优秀 | 受限 |
优化建议
- 内存优化 :
- 使用内存高效的注意力实现(如 FlashAttention)
-
分块处理超大图像
-
计算加速 :
- 稀疏注意力机制
- 知识蒸馏到小模型
思考问题
- ViT 完全抛弃了卷积归纳偏置,这是优势还是劣势?在小数据场景下如何弥补?
- 如何设计更适合视频处理的时空注意力机制?
- 当图像分辨率变化时(如从 224×224 到 512×512),ViT 应该如何调整才能保持效率?
希望这篇指南能帮助你顺利入门 Vision Transformer。建议从简单的图像分类任务开始实践,逐步探索更复杂的视觉应用。
