共计 2634 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点:为什么需要 Transformer
传统卷积神经网络(CNN)在计算机视觉(CV)领域取得了巨大成功,但其核心设计存在一个根本限制:局部感受野。CNN 通过滑动窗口的方式逐步提取特征,这意味着:
- 浅层卷积只能捕捉像素级的局部模式(如边缘、纹理)
- 深层网络通过堆叠卷积层才能建立远距离特征关系
- 这种间接的远程建模方式效率低下,且难以处理长距离依赖
而 Transformer 最初在自然语言处理(NLP)中展现出强大的 全局上下文建模能力,其核心是通过 Self-Attention 机制直接计算任意两个位置的关系。2020 年《An Image is Worth 16×16 Words》论文首次将纯 Transformer 架构(Vision Transformer/ViT)成功应用于图像分类,引发了 CV 领域的架构革命。
核心原理对比
CNN 的局部特征提取
典型 CNN 层的计算可表示为:
$$\text{Output}(x,y) = \sum_{i=-k}^{k}\sum_{j=-k}^{k} \text{Input}(x+i,y+j) \cdot \text{Kernel}(i,j)$$
其中 kernel size $k$ 通常为 3 或 5,这种设计带来两个特性:
- 平移不变性:相同模式在不同位置会被同样检测
- 局部性:每个输出仅依赖附近 $k\times k$ 区域的输入

(示意图:CNN 逐层扩大感受野)
Transformer 的全局注意力
Transformer 的核心是 Multi-Head Self-Attention(多头自注意力),其关键计算步骤:
- 将输入投影到 Query(Q)、Key(K)、Value(V)空间
- 计算注意力权重:
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$ - 多头结果拼接后通过线性层融合
(示意图:Transformer 可直接建模任意两像素关系)
实战代码对比
CNN 图像分类器实现
import torch
import torch.nn as nn
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
# 卷积层 1:3 输入通道,16 输出通道,3x3 卷积核
self.conv1 = nn.Conv2d(3, 16, kernel_size=3, stride=1, padding=1)
# 最大池化,2x2 窗口,步长 2
self.pool = nn.MaxPool2d(2, 2)
# 卷积层 2:16->32 通道
self.conv2 = nn.Conv2d(16, 32, 3, 1, 1)
# 全连接层
self.fc = nn.Linear(32 * 8 * 8, num_classes) # 假设输入为 32x32 图像
def forward(self, x):
x = self.pool(torch.relu(self.conv1(x))) # 输出 16x16x16
x = self.pool(torch.relu(self.conv2(x))) # 输出 8x8x32
x = x.view(-1, 32 * 8 * 8) # 展平
return self.fc(x)
ViT 关键模块实现
class PatchEmbedding(nn.Module):
"""将图像分割为 patches 并嵌入"""
def __init__(self, img_size=32, patch_size=4, in_chans=3, embed_dim=64):
super().__init__()
self.num_patches = (img_size // patch_size) ** 2
# 使用卷积实现分块和投影
self.proj = nn.Conv2d(in_chans, embed_dim,
kernel_size=patch_size,
stride=patch_size)
# 可学习的分类 token
self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))
def forward(self, x):
# x 形状: [B, C, H, W]
x = self.proj(x) # [B, E, H/P, W/P]
x = x.flatten(2).transpose(1, 2) # [B, num_patches, E]
# 添加 class token
cls_tokens = self.cls_token.expand(x.shape[0], -1, -1)
x = torch.cat((cls_tokens, x), dim=1)
return x
三大常见陷阱与解决方案
- 显存溢出问题
- 错误:直接对高分辨率图像(如 224×224)计算全局注意力
-
解决:必须先将图像分块(如 16×16 patches),ViT 中典型设置:
patch_size = 16 # 224x224 图像→14x14 个 patch -
位置信息丢失
- 错误:未添加位置编码(Positional Encoding),导致模型无法理解空间结构
-
解决:标准 ViT 使用可学习的 1D 位置编码:
self.pos_embed = nn.Parameter(torch.randn(1, num_patches+1, embed_dim)) -
训练不稳定性
- 错误:直接使用 CNN 常用的学习率(如 0.1)导致梯度爆炸
- 解决:Transformer 通常需要更小的学习率和预热(warmup):
optimizer = AdamW(model.parameters(), lr=3e-5) scheduler = get_cosine_schedule_with_warmup(optimizer, warmup_steps=500)
性能对比实验(CIFAR-10)
| 指标 | ResNet-18 | ViT-Tiny |
|---|---|---|
| 参数量 | 11M | 5M |
| 训练时间 /epoch | 45s | 68s |
| 显存占用 | 1.2GB | 2.1GB |
| 测试准确率 | 93.5% | 91.2% |
注:ViT 在小数据集上需要更强的数据增强和正则化
融合创新方向
现代架构正尝试结合两者的优势:
- Swin Transformer:引入局部窗口注意力 + 跨窗口连接,兼具 CNN 的局部性和 Transformer 的全局性
- ConvNeXt:用 CNN 结构实现类似 Transformer 的特性
- 混合架构:前端使用 CNN 提取低级特征,后端用 Transformer 建模全局关系
建议初学者先从 PyTorch 官方 Vision Transformer 教程入手,逐步理解模块设计背后的动机。记住:没有绝对优越的架构,只有适合特定任务的解决方案。
