摘要:本文系统介绍了 Vision Transformer (ViT) 的核心原理与实现。首先分析了将 Transformer 引入视觉领域的动机,对比了 CNN 的局部归纳偏置与 Transformer 的全局建模能力。接着详细解析了 ViT 的整体架构,包括图像分块、Patch Embedding、位置编码、CLS Token 以及 Transformer Encoder 中的自注意力、多头注意力和 MLP 模块。文章提供了完整的 PyTorch 实现代码,可在 CIFAR-10 上训练,并总结了训练要点和常见问题。最后讨论了 ViT 的局限性、后续改进工作(如 DeiT、Swin Transformer、MAE 等)及其在多模态大模型中的基石地位。
关键词:Vision Transformer、ViT、自注意力机制、图像分类、PyTorch、计算机视觉
本文带你从 0 到 1 理解 ViT 的设计思想,推导多头注意力的计算过程,并用 PyTorch 手写一个能在 CIFAR-10 上训练的 ViT 模型。全文约 4500 字,建议收藏后阅读。
一、为什么要把 Transformer 搬进视觉领域?
2017 年 Transformer 横空出世后,迅速统一了 NLP 领域。而在计算机视觉领域,CNN(卷积神经网络)统治了将近十年。CNN 的两个核心先验是:
- 局部性(Locality):相邻像素相关性高,卷积核只在局部窗口内滑动;
- 平移等变性(Translation Equivariance):同一个特征出现在图像任何位置,都应被同一个卷积核检测到。
这些归纳偏置在小数据上是优势,但也带来一个根本局限:CNN 的感受野增长缓慢,对长距离依赖建模能力弱。虽然可以通过堆叠层数、空洞卷积扩大感受野,但信息在层间传递时仍会逐层衰减。
Transformer 恰恰相反——它几乎没有归纳偏置,自注意力机制让任意两个 token 之间的距离都是 1。2020 年 Google 提出的 ViT(Vision Transformer) 证明:当数据量足够大(JFT-300M)时,纯 Transformer 可以超越所有 CNN,并且计算效率更高。
核心思想一句话概括:
把图像切成一个个 Patch,把每个 Patch 当作 NLP 里的"单词",直接套用标准的 Transformer Encoder。
二、ViT 的整体架构
ViT 的前向流程可以拆成 5 步:
- 图像分块(Patch Partition)
- 线性映射 + 位置编码(Patch Embedding + Position Embedding)
- 拼接分类 token(CLS Token)
- Transformer Encoder(L 层堆叠)
- MLP 分类头
2.1 图像分块与 Embedding
设输入图像 ,把它切成 N 个
的小块,则:
以 ViT-Base/16 为例:输入 224×224×3,Patch 大小 16×16,得到 N = 14×14 = 196 个 Patch。
每个 Patch 展平成向量 (即 16×16×3=768 维),再通过一个可学习的线性投影
映射到 D 维(ViT-Base 中 D=768):
这里有两处关键设计:
- CLS Token:借鉴 BERT,在序列最前面拼接一个可学习的分类
,它经过多层注意力后聚合全局信息,最终用于分类;
- 位置编码 E_{pos}:自注意力本身是"置换不变"的,打乱 Patch 顺序结果不变,因此必须显式注入位置信息。ViT 使用可学习的 1D 位置编码,后续实验证明它与正弦编码效果相当。
工程技巧:Patch 划分 + 展平 + 线性映射这三步,可以等价地用一个步长等于核大小的卷积一步完成:
nn.Conv2d(3, 768, kernel_size=16, stride=16)。
2.2 Transformer Encoder
每个 Encoder Block 由两部分组成,均采用 Pre-Norm(先 LayerNorm 再进子层)结构:
自注意力(Self-Attention)
对输入序列 ,计算三组投影:
缩放点积注意力:
为什么要除以 ?因为点积的方差随维度
线性增长,softmax 会进入梯度极小的饱和区,缩放后保持梯度稳定。
多头注意力(Multi-Head Attention)
单头注意力只能在一个表示子空间里计算相关性。多头机制把 D 维拆成 h 个头(ViT-Base 有 12 个头,每头 64 维),各自独立计算注意力再拼接:
多头的好处是不同头可以关注不同模式——实验可视化表明,有的头关注局部纹理(类似卷积),有的头关注全局轮廓。
MLP 块
每个 Block 中的前馈网络是两层全连接 + GELU 激活:
隐藏层维度通常是 4D(ViT-Base 为 3072)。
2.3 分类头
取最后一层 CLS Token 对应的输出 ,经 LayerNorm 后接一个线性层即可:
三、PyTorch 手写 ViT
下面是一个可在 CIFAR-10 上直接训练的精简版 ViT(输入 32×32,Patch 4×4,共 64 个 Patch)。
import torch
import torch.nn as nn
class PatchEmbedding(nn.Module):
"""图像分块 + 线性投影,用卷积一步实现"""
def init(self, img_size=32, patch_size=4, in_chans=3, embed_dim=384):
super().init()
self.n_patches = (img_size // patch_size) ** 2
self.proj = nn.Conv2d(in_chans, embed_dim,
kernel_size=patch_size, stride=patch_size)
def forward(self, x): # x: (B, 3, 32, 32)
x = self.proj(x) # (B, D, 8, 8)
x = x.flatten(2).transpose(1, 2) # (B, 64, D)
return x
class MultiHeadAttention(nn.Module):
def init(self, dim=384, num_heads=6, dropout=0.1):
super().init()
assert dim % num_heads == 0
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.scale = self.head_dim ** -0.5
self.qkv = nn.Linear(dim, dim * 3)
self.proj = nn.Linear(dim, dim)
self.dropout = nn.Dropout(dropout)
def forward(self, x): # x: (B, N+1, D)
B, N, D = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, H, N, head_dim)
q, k, v = qkv.unbind(0)
attn = (q @ k.transpose(-2, -1)) * self.scale # (B, H, N, N)
attn = attn.softmax(dim=-1)
attn = self.dropout(attn)
x = (attn @ v).transpose(1, 2).reshape(B, N, D)
return self.proj(x)
class Block(nn.Module):
"""Pre-Norm Transformer Block"""
def init(self, dim=384, num_heads=6, mlp_ratio=4, dropout=0.1):
super().init()
self.norm1 = nn.LayerNorm(dim)
self.attn = MultiHeadAttention(dim, num_heads, dropout)
self.norm2 = nn.LayerNorm(dim)
self.mlp = nn.Sequential(
nn.Linear(dim, dim * mlp_ratio),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(dim * mlp_ratio, dim),
nn.Dropout(dropout),
)
def forward(self, x):
x = x + self.attn(self.norm1(x)) # 残差连接 1
x = x + self.mlp(self.norm2(x)) # 残差连接 2
return x
class ViT(nn.Module):
def init(self, img_size=32, patch_size=4, in_chans=3, num_classes=10,
embed_dim=384, depth=6, num_heads=6, mlp_ratio=4, dropout=0.1):
super().init()
self.patch_embed = PatchEmbedding(img_size, patch_size, in_chans, embed_dim)
n_patches = self.patch_embed.n_patches
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
self.pos_embed = nn.Parameter(torch.zeros(1, n_patches + 1, embed_dim))
self.pos_drop = nn.Dropout(dropout)
self.blocks = nn.Sequential(
*[Block(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth)]
)
self.norm = nn.LayerNorm(embed_dim)
self.head = nn.Linear(embed_dim, num_classes)
nn.init.trunc_normal_(self.cls_token, std=0.02)
nn.init.trunc_normal_(self.pos_embed, std=0.02)
def forward(self, x):
B = x.shape[0]
x = self.patch_embed(x) # (B, 64, D)
cls = self.cls_token.expand(B, -1, -1) # (B, 1, D)
x = torch.cat([cls, x], dim=1) # (B, 65, D)
x = self.pos_drop(x + self.pos_embed)
x = self.blocks(x)
x = self.norm(x)
return self.head(x[:, 0]) # 取 CLS token
测试
model = ViT()
x = torch.randn(2, 3, 32, 32)
print(model(x).shape) # torch.Size([2, 10])
print(f"参数量: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M")
训练要点(踩坑总结)
-
ViT 非常吃数据。在小数据集上从头训练,ViT 通常打不过同规模 ResNet。补救方法:
- 强数据增强:RandAugment、Mixup、CutMix、Random Erasing;
- 或者干脆用
timm库加载 ImageNet 预训练权重做微调:
import timm model = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=10) -
优化器与正则:原论文使用 AdamW + 余弦退火学习率 + Warmup。ViT 对超参数比 CNN 敏感得多,建议 warmup 至少 5~10 个 epoch。
-
位置编码插值:微调时如果输入分辨率变了(如 224 → 384),Patch 数变化导致位置编码形状不匹配,需要对位置编码做双线性插值。
四、ViT 的局限与后续发展
| 局限 | 改进工作 |
|---|---|
| 缺少归纳偏置,小数据上表现差 | DeiT(蒸馏 + 强增强,只用 ImageNet-1K 训练) |
| 自注意力复杂度 O(N^2),高分辨率下昂贵 | Swin Transformer(窗口注意力,O(N),层级结构) |
| 单一尺度,不适合检测/分割等稠密任务 | PVT、Swin、MViT(金字塔多尺度特征) |
| 需要超大数据集预训练 | MAE(掩码自监督预训练,重建像素) |
ViT 更大的意义在于统一了视觉与语言的架构。此后 CLIP(图文对齐)、DALL·E(图像生成)、Flamingo(多模态对话)乃至今天的 GPT-4o、Qwen-VL 等多模态大模型,视觉侧大多沿用 ViT 作为图像编码器。可以说,不理解 ViT,就很难理解当下多模态模型的技术底座。
五、小结
- ViT = Patch 切分 + 线性 Embedding + 标准 Transformer Encoder,思想极简;
- 自注意力的全局建模能力是其优势,O(N^2) 复杂度和数据饥渴是其代价;
- 工程上推荐
timm库直接调预训练模型,自己手写一遍则是理解细节的捷径; - ViT 是当今多模态大模型的视觉基石,值得投入时间吃透。
参考资料
- Dosovitskiy et al. An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale. ICLR 2021.
- Vaswani et al. Attention Is All You Need. NeurIPS 2017.
- Touvron et al. Training data-efficient image transformers & distillation through attention (DeiT). ICML 2021.
- Liu et al. Swin Transformer: Hierarchical Vision Transformer using Shifted Windows. ICCV 2021.
- He et al. Masked Autoencoders Are Scalable Vision Learners (MAE). CVPR 2022.

3万+

被折叠的 条评论
为什么被折叠?



