ViT: An Image is Worth 16x16 Words

论文: An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale (ICLR 2021)
作者: Alexey Dosovitskiy, Lucas Beyer, Alexander Kolesnikov, Dirk Weissenborn, Xiaohua Zhai, Thomas Unterthiner, Mostafa Dehghani, Matthias Minderer, Georg Heigold, Sylvain Gelly, Jakob Uszkoreit, Neil Houlsby (Google Research)
代码: google-research/vision_transformer

一句话总结:ViT 把一张图切成若干 patch,每个 patch 当成一个”词”,喂给标准 Transformer Encoder,用 CLS token 做分类。 它证明了 NLP 的 Transformer 直接搬到 CV 也能work,只要数据够大–从此 ViT 成为视觉领域的通用骨干,支撑了 CLIP、LLaVA 等 VLM、VLA 与多模态模型。


一、背景

1.1 为什么要有 ViT

Transformer 在 NLP 取得统治地位后,CV 仍以 CNN(ResNet、EfficientNet)为主。把 Transformer 用到图像的难点在于:图像的”序列”是什么?文本天然是 token 序列,图像是 2D 像素网格,直接把每个像素当 token 会让序列长度爆炸(224×224 图 = 5 万 token,attention 是 $O(n^2)$,算不动)。

早期尝试(iGPT、ViT 前的非局部网络等)效果一般。ViT 的解法很直接:把图切成 16×16 的 patch,每个 patch 经线性投影成一个 token–这样 224×224 图只有 $14\times14=196$ 个 token,和 NLP 句子长度相当,直接复用 Transformer Encoder。

1.2 关键结论:数据够大时,Transformer 干掉 CNN

ViT 论文最核心的实验结论是:

  • 在中等数据量(ImageNet-1k,1.2M 图)上从头训,ViT 弱于 ResNet–因为 CNN 自带”平移等变 + 局部性”两种归纳偏置,小数据上更省;
  • 但在超大数据集(JFT-300M,3 亿图 / ImageNet-21k,1400 万图)上预训练后,ViT 超越 同等规模的 CNN,且随规模增大优势扩大。

这说明 CNN 的归纳偏置是”小数据的拐杖”,数据够大时反而成了限制。ViT 用更少的先验换取更强的可扩展性,开启了 CV 的 scaling law 时代。

1.3 广泛影响:视觉通用骨干

ViT 现在是几乎所有视觉/多模态模型的视觉编码器:

领域 代表 ViT 的角色
视觉-语言模型 (VLM) CLIP、LLaVA、InternVL 图像编码器,把图变成 token 喂给 LLM
多模态大模型 GPT-4V、Gemini 视觉前端
VLA / 机器人 OpenVLA、π0 视觉编码器,输入相机图像
自监督预训练 MAE、BEiT、DINOv2 ViT 作骨干
视频理解 ViViT、TimeSformer 时序扩展

现代 VLM(如 LLaVA)的本质就是”ViT 提取图像 token + LLM 处理文本 token + 投影层对齐“。理解 ViT 是理解整个多模态架构的基础。


二、整体架构

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
图像 (H×W×C, 如 224×224×3)
↓ 切 patch + 线性投影
[patch_1, patch_2, ..., patch_N] N = (H/P)×(W/P), P=16 -> N=196
↓ 每个 patch -> D 维向量 (D=768)
↓ 预置 [CLS] token 在最前
[CLS, patch_1, ..., patch_N] (N+1, D)
↓ + 位置编码 (learned)
[CLS, patch_1, ..., patch_N] + PE (N+1, D)

Transformer Encoder × L (L=12) 标准 Transformer Encoder (MHA + MLP + LN + 残差)

取 [CLS] 位置的输出

MLP Head (Linear) 分类 logits

Softmax -> 类别概率

ViT 本质就是**”标准 Transformer Encoder + 一个 Patch Embedding 前端 + 一个 CLS 分类头”**。Encoder 部分与 NLP 的 Transformer Encoder 完全一致(详见 [[20.notes/感知算法/骨干网络/Transformer/Transformer 详解|Transformer 详解]] 第四章),创新在于”图像如何变成 token 序列”。

骨架四句话:

  1. Patch Embedding:图像切 patch,每个 patch 线性投影成 D 维 token;
  2. CLS token + 位置编码:prepend 一个可学习的 [CLS] token,加上可学习位置编码;
  3. Transformer Encoder:L 层标准 Encoder(无 causal mask,双向);
  4. 分类头:取 [CLS] 输出过 MLP 做分类。

三、输入处理 ⭐

这是 ViT 与 NLP Transformer 最大的不同–图像如何变成模型能吃的 token 序列。

3.1 Patch Embedding:图像切块 + 线性投影

把 $H\times W\times C$ 的图像切成 $N=(H/P)\times(W/P)$ 个 $P\times P\times C$ 的 patch($P=16$),每个 patch 拉平成 $P^2\cdot C$ 维向量,再线性投影到 $D$ 维。

实现上用 Conv2d 一步完成”切块 + 投影”(kernel=stride=patch_size 的卷积等价于不重叠切块再投影):

1
2
3
4
5
6
7
8
9
10
11
class MoonVisionPatchEmbed(nn.Module):
def __init__(self, out_dim, in_dim=3, patch_size=(14,14), pos_emb_height=14, pos_emb_width=14):
super().__init__()
self.patch_size = patch_size
self.proj = nn.Conv2d(in_dim, out_dim, kernel_size=patch_size, stride=patch_size) # ★ 切块+投影
self.pos_emb = Learnable2DInterpPosEmb(height=pos_emb_height, width=pos_emb_width, dim=out_dim)

def forward(self, x, grid_hws):
x = self.proj(x).view(x.size(0), -1) # [N_patches, D]
x = self.pos_emb(x, grid_hws) # + 位置编码
return x

原版 ViT 也可写成 nn.Linear(P*P*C, D) 作用于每个展平的 patch,数学上与 Conv2d 等价,但 Conv2d 更高效。MoonViT(NVIDIA/LocateAnything-3B 用的视觉编码器)即采用 Conv2d 写法,patch_size=14。

以 224×224 图、P=16 为例:

1
2
224×224×3  --Conv2d(k=16, s=16)-->  14×14×768  --flatten-->  196×768
(N=196 个 token, D=768)

3.2 CLS Token:分类的”汇聚位”

prepend 一个可学习的 [CLS] token到序列最前面:

1
[CLS, patch_1, patch_2, ..., patch_196]   (197, 768)

CLS token 本身不对应任何 patch,它的作用是”全局信息的汇聚位”–经过 L 层 self-attention 后,CLS 通过与所有 patch 交互累积整图信息,最后用它做分类。

这继承自 BERT 的 [CLS] 设计。也可不用 CLS,改用对全部 patch 输出做 mean pooling 做分类(GAP),效果相近。

3.3 位置编码:给 patch 注入 2D 顺序

Patch Embedding 后的 token 序列本身无顺序感(打乱 patch 顺序结果不变,同 Transformer 的排列等变性),必须加位置编码。

方式 说明 代表
可学习 1D PE 原版 ViT:一张 (N+1, D) 的可学习查表,按 patch 序号取行 ViT、DeiT
可学习 2D 插值 PE 按 (行,列) 2D 网格组织,支持不同分辨率插值 MoonViT (Learnable2DInterpPosEmb)
2D RoPE 把旋转位置编码扩展到 2D,注入 Q/K MoonViT (Rope2DPosEmb)、NaViT

原版 ViT 用可学习 1D PE(按 patch 拉平后的序号 0..196 编码)。它虽然不显式利用 2D 网格结构,但实验发现模型自己能学到 2D 邻近性(相邻 patch 的 PE 接近)。

现代变体多改用 2D 位置编码(2D 插值 PE 或 2D RoPE),更自然地表达图像的 2D 结构,且支持任意分辨率。MoonViT 同时提供了 Learnable2DInterpPosEmbRope2DPosEmb 两种。


可学习位置编码可视化:位置越接近编码越相似,且呈现行列结构

最终输入 = Patch Embedding + CLS token(prepend)+ 位置编码(相加)

$$
\mathbf{z}0 = [,\mathbf{x}{cls},;, \mathbf{x}_p^1 \mathbf{E},;, \cdots,;, \mathbf{x}p^N \mathbf{E},] + \mathbf{E}{pos}
$$

其中 $\mathbf{x}p^i$ 是第 $i$ 个展平 patch,$\mathbf{E}$ 是投影矩阵,$\mathbf{x}{cls}$ 是 CLS token,$\mathbf{E}_{pos}$ 是位置编码。

timm 中的组装forward_features):

1
2
3
4
5
6
7
8
def forward_features(self, x):
x = self.patch_embed(x) # (B,C,H,W) -> (B,N,E)
cls_token = self.cls_token.expand(x.shape[0], -1, -1)
x = torch.cat((cls_token, x), dim=1) # (B,N,E) -> (B,1+N,E) 预置 CLS
x = self.pos_drop(x + self.pos_embed) # + 可学习位置编码, 再 dropout
x = self.blocks(x) # L 层 Encoder Block
x = self.norm(x)
return self.pre_logits(x[:, 0]) # 取 [CLS] 位置的输出

四、Transformer Encoder

Encoder 部分与 NLP Transformer 完全一致,L 层堆叠,每层两个子层 + 残差 + LayerNorm:

$$
\begin{aligned}
\mathbf{z}’\ell &= \text{MSA}(\text{LN}(\mathbf{z}{\ell-1})) + \mathbf{z}{\ell-1} \
\mathbf{z}
\ell &= \text{MLP}(\text{LN}(\mathbf{z}’\ell)) + \mathbf{z}’\ell
\end{aligned}
\quad \ell=1\ldots L
$$

  • MSA:Multi-Head Self-Attention(双向,无 causal mask);
  • MLP:两层线性 + GELU,中间升维 4×;
  • Pre-LN:ViT 用 Pre-LN(LN 在子层前),比原 Transformer 的 Post-LN 更稳定;
  • 无降采样:ViT 全程保持 patch 数 N 不变(不像 CNN 那样逐层降分辨率),所有 token 在同一分辨率交互。

细节详见 [[20.notes/感知算法/骨干网络/Transformer/Transformer 详解|Transformer 详解]] 第四章(自注意力、多头、FFN、Add&Norm)。MoonViT 的 MoonVitEncoderLayer 额外用了 QKV packed + Flash Attention + 2D RoPE,是工程优化版。


timm 标准实现Attention + Block,Pre-LN):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
class Attention(nn.Module):                       # 多头注意力
def __init__(self, dim, num_heads=8, qkv_bias=False):
self.num_heads = num_heads
self.scale = (dim // num_heads) ** -0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) # QKV packed 一次投影
self.proj = nn.Linear(dim, dim)
def forward(self, x): # x: (B, N, C)
B, N, C = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
q, k, v = qkv.unbind(0) # 各 (B, num_heads, N, head_dim)
attn = (q @ k.transpose(-2, -1)) * self.scale
attn = attn.softmax(dim=-1)
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
return self.proj(x)

class Block(nn.Module): # Encoder Block (Pre-LN)
def __init__(self, dim, num_heads, mlp_ratio=4., drop_path=0., ...):
self.norm1 = nn.LayerNorm(dim)
self.attn = Attention(dim, num_heads, ...)
self.drop_path = DropPath(drop_path) if drop_path > 0 else nn.Identity()
self.norm2 = nn.LayerNorm(dim)
self.mlp = Mlp(in_features=dim, hidden_features=int(dim * mlp_ratio), act_layer=nn.GELU)
def forward(self, x):
x = x + self.drop_path(self.attn(self.norm1(x))) # MHA + 残差
x = x + self.drop_path(self.mlp(self.norm2(x))) # MLP + 残差
return x
1
2
3
4
5
6
7
8
9
class Mlp(nn.Module):                           # Position-wise FFN
def __init__(self, in_features, hidden_features=None, act_layer=nn.GELU, drop=0.):
self.fc1 = nn.Linear(in_features, hidden_features or in_features)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_features or in_features, in_features)
self.drop = nn.Dropout(drop)
def forward(self, x):
x = self.drop(self.act(self.fc1(x))) # 升维 4× + GELU
return self.drop(self.fc2(x)) # 降回原维度


Encoder Block(Pre-LN):LN->MHA->残差->LN->MLP->残差

Drop Path(随机深度):ViT 用 DropPath 替代部分 Dropout–按概率随机跳过整个 Block(该层输出置零),相当于训练时随机”删掉”一层,起正则与加深鲁棒性作用。drop_path_rate 通常随层线性递增。

ViT-Base 配置:

超参
Layers (L) 12
Hidden size (D) 768
MLP size 3072 (4×D)
Heads (h) 12
Patch size (P) 16
Params ~86M

(另有 ViT-Large: L=24/D=1024/307M;ViT-Huge: L=32/D=1280/632M)


论文 Table 1:ViT-Base / Large / Huge 三档配置(Layers / Hidden Size / MLP Size / Heads)


五、分类头

取最后一层 [CLS] 位置的输出 $\mathbf{z}_L^0$,过一层 Linear 投影到类别数:

$$
y = \text{LN}(\mathbf{z}L^0), W{\text{head}}
$$

  • 预训练时通常用更复杂的头(MLP),fine-tune 时换成单层 Linear;
  • 推理时对 $y$ 过 softmax 得类别概率。

不用 CLS 的话,可对全部 patch 输出做 mean pooling 再分类(GAP),效果相近且无需 CLS token。


六、训练

6.1 数据是关键

ViT 几乎不带归纳偏置,严重依赖大数据

预训练数据 规模 ImageNet fine-tune 后
ImageNet-1k 1.2M ViT 弱于 ResNet(数据不够)
ImageNet-21k 14M ViT 持平/略超 CNN
JFT-300M 300M ViT 显著超越 CNN

这正是 ViT 论文的核心卖点:数据够大时,更少先验的 Transformer 反而更强。小数据从头训 ViT 会过拟合、打不过 CNN。


ViT vs ResNet vs Hybrid:训练 epoch 少时 Hybrid 占优,epoch 增大后 ViT 反超


ViT 迁移学习指标:ImageNet 88.55%、ImageNet-ReaL 90.72%、CIFAR-100 94.55% 等

6.2 训练范式

  1. 大数据预训练(JFT-300M / ImageNet-21k):图像分类任务,学通用视觉表示;
  2. 小数据 fine-tune(ImageNet-1k 等):换分类头,用更小学习率精调。

6.3 优化细节

  • Adam (β₁=0.9, β₂=0.999);
  • 大量数据增强(rand augmentment、mixup、cutmix);
  • 学习率 warmup + cosine 衰减;
  • 高分辨率 fine-tune(如 384,比预训练的 224 更高)。

七、与 CNN 对比 / 优缺点

7.1 vs CNN

维度 CNN (ResNet) ViT
归纳偏置 平移等变 + 局部性(强先验) 几乎无(全局 attention)
小数据表现 ✅ 好(先验帮忙) ❌ 差(易过拟合)
大数据表现 增长放缓 ✅ 持续受益,反超 CNN
感受野 需堆层才扩到全局 第一层就全局
可扩展性 ✅ 强(scaling law 友好)
计算复杂度 O(n)(局部卷积) O(n²)(attention,n=patch 数)

7.2 优点 ✅

  1. 可扩展性强:随数据/参数/算力增大持续受益,scaling law 友好;
  2. 全局建模:第一层就全局感受野,长程依赖强;
  3. 统一架构:与 NLP 同一套 Transformer,便于多模态融合(VLM 的基础);
  4. 少先验:不依赖手工设计的卷积核,更通用。

7.3 局限 ❌

  1. 数据 hungry:小数据从头训打不过 CNN,需大数据预训练;
  2. O(n²) 复杂度:高分辨率图 patch 多,计算/显存大(同 Transformer,详见 [[20.notes/感知算法/骨干网络/Transformer/Transformer 详解|Transformer 详解]] 第八章);
  3. 缺少 2D 归纳偏置:平移/局部性需从数据学,效率低;
  4. 变长分辨率处理:原版位置编码固定 N,换分辨率需插值。

八、变体与影响

ViT 之后涌现大量改进,主要方向一览:

变体 改进点 代表
DeiT 数据高效:强增强 + 蒸馏,ImageNet-1k 从头训也能打 DeiT
Swin Transformer 层级结构 + 移窗 attention,引入 CNN 式多尺度,降复杂度 Swin
BEiT / MAE 自监督预训练(掩码 patch 重建),摆脱对标注的依赖 MAE、BEiT、DINOv2
NaViT 原生支持任意分辨率/宽高比(patch packing) NaViT
SigLIP sigmoid 损失替代 softmax 的对比学习视觉编码器 SigLIP
MoonViT 2D RoPE + Flash Attention + Patch Merger,现代 VLM 视觉编码器 LocateAnything-3B

8.1 DeiT:让 ViT 摆脱大数据依赖

原版 ViT 需 JFT-300M 级数据才能打过 CNN,普通实验室玩不转。DeiT 的贡献是在 ImageNet-1k(1.2M)上从头训 ViT 也能达到 SOTA,靠三件事:

  • 强数据增强:RandAugment + Mixup + CutMix + Random Erasing,弥补小数据;
  • 知识蒸馏:以 CNN teacher(RegNet)作监督,额外引入一个 distillation token(与 CLS token 并列),向 teacher 学 CNN 的归纳偏置;
  • 推理融合:CLS 头与 distill 头的预测取平均。

timm 代码里的 dist_token 即 DeiT 的蒸馏 token(forward_featuresdist_token is None 分支对应原版 ViT,非 None 对应 DeiT)。效果:ImageNet-1k 从头训达 80%+ top-1。

8.2 Swin Transformer:层级 + 移窗,支持密集预测

ViT 全局 attention 是 $O(n^2)$,高分辨率算不动;且无多尺度结构(CNN 的特征金字塔对检测/分割很重要)。Swin 的两个关键设计:

  • 层级结构(Hierarchical):通过 Patch Merging 逐层 2× 降采样(类似 CNN 的 stride-2 下采样),形成多尺度特征金字塔,适合检测/分割等密集预测任务;
  • 移窗注意力(Shifted Window MSA):attention 只在局部窗口(如 7×7 patch)内计算,复杂度从 $O(n^2)$ 降到 $O(n)$;相邻层窗口错位(shifted),让信息跨窗口流通。

效果:线性复杂度 + 多尺度,成为 Swin-Object 等检测/分割模型的骨干。ViT 是”单尺度全局 attention”,Swin 是”多尺度局部+移窗”。

8.3 MAE:自监督预训练,摆脱标注依赖

ViT 依赖大量标注数据。MAE(Masked Autoencoders)用自监督方式预训练:

  • 随机 mask 掉 75% 的 patch,只编码可见的 25%;
  • 用轻量 decoder 重建被 mask patch 的像素值
  • 高 mask 比例是关键:75% 才能去掉图像冗余、让任务有意义(mask 太少任务太简单);
  • encoder 只处理可见 patch(省算力),decoder 处理全量。

效果:自监督预训练的 ViT 在 ImageNet fine-tune 后超越有监督 ViT,且 linear probing 与下游检测/分割大幅提升。意义:让 ViT 摆脱对标注的依赖,成为自监督视觉表示主力(同期 BEiT 用预测 token、DINOv2 用自监督+蒸馏,思路相近)。

8.4 其他现代变体

  • NaViT:原生支持任意分辨率/宽高比(patch packing,把不同尺寸图的 patch 打包进一个 batch),打破 ViT 固定 224 的限制;
  • SigLIP:对比学习视觉编码器,用 sigmoid 损失(独立二分类)替代 InfoNCE 的 softmax,更易扩展 batch;
  • DINOv2:自监督 + 自蒸馏,产出强通用视觉特征,无需 fine-tune 即可用于检索/分类/分割;
  • MoonViT(你的代码):2D RoPE + Flash Attention + Patch Merger,面向现代 VLM 的视觉编码器。

趋势:从”固定分辨率分类骨干”走向”任意分辨率 + 自监督 + 高效 attention”的通用视觉前端,服务于 VLM/VLA。


九、总结

ViT 的精髓是把图像当成一句话:切 patch 当词、加位置编码、喂 Transformer Encoder、用 CLS 分类。它牺牲了 CNN 的归纳偏置,换来了更强的可扩展性与跨模态统一性–在大数据预训练下超越 CNN,并成为 CLIP、LLaVA 等 VLM/VLA 的视觉骨干。

理解 ViT 的关键三件事:① Patch Embedding(Conv2d 切块+投影,把图变 token 序列);② CLS token + 位置编码(汇聚全局 + 注入顺序);③ 标准 Transformer Encoder(与 NLP 一致,靠大数据撑起无先验架构)。


相关链接

剪藏来源(10.clippings/感知算法/骨干网络/ViT/):

  • 📋 [[10.clippings/感知算法/骨干网络/ViT/Transformer_CV Vision Transformer(ViT)重點筆記]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/ViT( Vision Transformer)]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/ViT解读 — 深入浅出PyTorch]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/ViT(Vision Transformer)解析]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/Visual Transformer (ViT)模型详解-CSDN博客]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/全网最强ViT (Vision Transformer)原理及代码解析]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/神经网络算法 - 一文搞懂ViT(Vision Transformer)]]
  • 📋 [[10.clippings/感知算法/骨干网络/ViT/神经网络算法 – 一文搞懂ViT(Vision Transformer) – 人工智能 – 白盒子]]