深入理解 Vision Transformer (ViT):图像识别的范式革新

5091 字
25 分钟
深入理解 Vision Transformer (ViT):图像识别的范式革新

1. ViT 的核心思想#

1.1 为什么要把 Transformer 引入视觉?#

在 2020 年之前,计算机视觉领域几乎被 卷积神经网络 (CNN) 一统天下——从 AlexNet 到 ResNet,再到 EfficientNet,CNN 在图像分类、检测、分割等任务上不断刷新 SOTA。

而 Transformer 早在 2017 年就已在 NLP 领域大放异彩,却迟迟未能”征服”视觉。一个最直接的问题摆在研究者面前:

图像是一个二维的、空间结构极强的信号,而 Transformer 处理的是一维序列。把”擅长语言的架构”硬塞到”图像”上,能行得通吗?

2020 年底,Google 的论文 “An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale” 给出了肯定的答案——Vision Transformer (ViT) 横空出世,并在 ImageNet 等基准上以纯 Transformer 架构击败了当时最强的 CNN。

1.2 ViT 的核心思想#

ViT 的核心思想可以概括为三句话:

  1. 把图像切成小块 (Patch) → 模拟 NLP 中的”词 Token”
  2. 每个块线性映射成向量 → 相当于”词嵌入”
  3. 送入标准 Transformer Encoder → 完全复用 NLP 那套 Transformer
┌─────────────────────────────────────────────────────────────────┐
│ ViT 核心流程图 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 输入图像 (H × W × C) │
│ │ │
│ ▼ │
│ ┌──────────────┐ │
│ │ Patch 切分 │ ──▶ N = (H×W)/(P×P) 个块 │
│ └──────┬───────┘ │
│ │ │
│ ▼ │
│ ┌──────────────┐ │
│ │ 线性投影层 │ ──▶ 每个块 → D 维向量 │
│ │ Linear Embed │ │
│ └──────┬───────┘ │
│ │ │
│ ▼ │
│ ┌──────────────┐ │
│ │ 拼接 [CLS] │ ──▶ 特殊分类 token │
│ └──────┬───────┘ │
│ │ │
│ ▼ │
│ ┌──────────────┐ │
│ │ 位置编码 │ ──▶ 加入位置信息 │
│ └──────┬───────┘ │
│ │ │
│ ▼ │
│ ┌──────────────┐ │
│ │ Transformer │ ──▶ 多层自注意力 │
│ │ Encoder │ │
│ └──────┬───────┘ │
│ │ │
│ ▼ │
│ ┌──────────────┐ │
│ │ MLP Head │ ──▶ 分类结果 │
│ └──────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────┘

1.3 一个生活化的比喻#

把 ViT 想成一群学生集体讨论一张画

  • 把画切成 16×16 的小块(类比 NLP 中的每个”词”)
  • 每位学生只看一小块(这就是 patch)
  • 但所有学生围成一圈讨论,互相询问”你看到的那块和我这块有没有关联”(这就是自注意力)
  • 经过多轮讨论后,每个学生都”知道”了整幅画的内容

这个机制绕开了 CNN 那种”局部→全局”的层级卷积,让模型从第一层就能建立任意两个 patch 之间的关联

1.4 ViT 发展简史#

2017: Transformer 诞生 (NLP)
2018-2019: 图像领域初步尝试
│ - Image Transformer (Parmar et al.)
│ - DETR (目标检测,Carion et al.)
│ 都受限于像素级注意力,计算量爆炸
2020.10: ViT 论文发表
│ - Google Brain / Google Research
│ - Dosovitskiy et al.
│ - 关键突破: patch token + 大规模预训练
2020.12: DeiT (Data-efficient Image Transformer)
│ - Facebook AI
│ - 用知识蒸馏让 ViT 在小数据集也能训
2021: Swin Transformer
│ - 引入层级化 + 窗口注意力
│ - 横扫下游任务 (检测、分割)
2022+: 各种改进
│ - PVT (金字塔 ViT)
│ - MViT (多尺度 ViT)
│ - BeiT (视觉 BERT,预训练范式)
│ - MAE (掩码自编码器)
2023-2025: 多模态基础
│ - CLIP (图文对齐)
│ - SAM (分割一切)
│ - GPT-4V, LLaVA (多模态 LLM)

2. 图像预处理:Patch 切分与线性嵌入#

2.1 Patch 切分#

假设输入图像尺寸为 H×W×CH \times W \times C,其中 CC 是通道数(RGB 图像为 3)。我们把它切分成 P×PP \times P 的小块,则有:

N=H×WP×PN = \frac{H \times W}{P \times P}

例如:H=W=224H=W=224, P=16P=16,则 N=224×22416×16=196N = \frac{224 \times 224}{16 \times 16} = 196 个 patch。

每个 patch 展平后是 P×P×CP \times P \times C 维的向量,例如 16×16×3=76816 \times 16 \times 3 = 768 维。

2.2 线性嵌入 (Patch Embedding)#

每个 patch 的展平向量经过一个可学习的线性投影层,映射到 DD 维(DD 是 Transformer 的隐层维度):

z0i=Expatchi+eposi,i=1,,Nz_0^i = E \cdot x_{\text{patch}}^i + e_{\text{pos}}^i, \quad i = 1, \ldots, N

其中 ER(P2C)×DE \in \mathbb{R}^{(P^2 C) \times D} 是可学习的投影矩阵,eposie_{\text{pos}}^i 是位置编码。

2.3 Patch Embedding 的代码实现#

实际上,Patch + Linear Projection 可以优雅地用一个二维卷积完成:

import torch
import torch.nn as nn
class PatchEmbed(nn.Module):
"""将图像切分为 patch,并通过卷积实现线性投影。"""
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
super().__init__()
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = (img_size // patch_size) ** 2
# 用 kernel_size=stride=patch_size 的卷积实现"切块 + 线性映射"
self.proj = nn.Conv2d(
in_channels=in_chans,
out_channels=embed_dim,
kernel_size=patch_size,
stride=patch_size,
)
def forward(self, x):
"""
输入:
x: (B, C, H, W) 图像批次
输出:
(B, N, D) N = num_patches, D = embed_dim
"""
B, C, H, W = x.shape
assert H == self.img_size and W == self.img_size, "输入尺寸不匹配"
# 卷积: (B, C, H, W) -> (B, D, H/P, W/P)
x = self.proj(x)
# 把空间维度展平到序列维度
x = x.flatten(2) # (B, D, N)
x = x.transpose(1, 2) # (B, N, D)
return x

小技巧:用 kernel=stride=patch_size 的卷积一次性完成”切块 + 嵌入”,既高效又能复用 GPU 优化。

2.4 序列长度计算#

图像尺寸Patch SizePatch 数量 (N)Token 序列长度
224×22416142=19614^2 = 196196
384×38416242=57624^2 = 576576
224×22414162=25616^2 = 256256
1024×102432322=102432^2 = 10241024

可以看到,patch 越小,token 越多,计算量以二次方增长——这是 ViT 处理高分辨率图像时的一个核心挑战。

3. 位置编码 (Positional Encoding)#

3.1 为什么 ViT 需要位置编码?#

Transformer 的自注意力机制有一个特性:它是置换不变的 (Permutation Invariant)——即无论输入 token 的顺序如何打乱,输出结果在数学上是相同的(注意力权重会相应调整)。

这带来了一个严重问题:

如果不给每个 patch 加上位置信息,模型根本无法分辨”天空在左、草地在右”和”天空在右、草地在左”。

3.2 ViT 的位置编码方案#

ViT 使用可学习的 1D 位置编码 (Learned 1D Positional Embedding)

self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
nn.init.trunc_normal_(self.pos_embed, std=0.02)

每个 patch token 都有一个可学习的向量,与它的 embedding 相加:

z0=[xcls;xpatch1E;xpatch2E;;xpatchNE]+Eposz_0 = [x_{\text{cls}}; x_{\text{patch}}^1 E; x_{\text{patch}}^2 E; \ldots; x_{\text{patch}}^N E] + E_{\text{pos}}

3.3 [CLS] Token:借鉴 BERT 的设计#

ViT 借鉴了 BERT 的设计,在 patch sequence 的开头预置一个特殊的可学习 token [CLS]

输入序列:
[CLS] | patch_1 | patch_2 | ... | patch_N
│ │ │ │
│ │ │ │
▼ ▼ ▼ ▼
position_0 pos_1 pos_2 ... pos_N

经过 Transformer Encoder 后,[CLS] token 对应的输出向量作为整张图像的表征,送入 MLP 分类头:

class ViT(nn.Module):
def __init__(self, ...):
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
...
def forward(self, x):
x = self.patch_embed(x) # (B, N, D)
cls_tokens = self.cls_token.expand(x.shape[0], -1, -1)
x = torch.cat([cls_tokens, x], dim=1) # (B, N+1, D)
x = x + self.pos_embed # 加位置编码
x = self.blocks(x) # Transformer Encoder
cls_output = x[:, 0] # 取 [CLS]
return self.head(cls_output)

3.4 位置编码的插值问题#

预训练时 ViT 通常使用 224×224224 \times 224 图像(位置编码表有 196+1196+1 个条目)。如果推理时换用 384×384384 \times 384,则需要 576+1576+1 个位置编码——怎么办?

标准做法是 2D 插值

def interpolate_pos_encoding(self, x, w, h):
"""对位置编码进行 2D 插值,以适应不同分辨率。"""
npatch = x.shape[1] - 1
N = self.pos_embed.shape[1] - 1
if npatch == N:
return self.pos_embed
class_emb = self.pos_embed[:, :1]
pos_embed = self.pos_embed[:, 1:]
dim = x.shape[-1]
# 把 pos_embed 重塑为 2D 网格
pos_embed = nn.functional.interpolate(
pos_embed.reshape(1, int(math.sqrt(N)), int(math.sqrt(N)), dim)
.permute(0, 3, 1, 2),
size=(int(math.sqrt(npatch)), int(math.sqrt(npatch))),
mode='bicubic',
)
pos_embed = pos_embed.permute(0, 2, 3, 1).flatten(1, 2)
return torch.cat([class_emb, pos_embed], dim=1)

4. Transformer Encoder#

4.1 整体结构#

ViT 的 Encoder 与 NLP Transformer 的 Encoder 完全一致,堆叠 LL 层:

每层包含两个子模块:
1. LayerNorm → Multi-Head Self-Attention → Residual
2. LayerNorm → MLP (FFN) → Residual

4.2 单层 Encoder 块#

class Block(nn.Module):
"""Transformer Encoder 单层。"""
def __init__(self, dim, num_heads, mlp_ratio=4.0, dropout=0.0):
super().__init__()
self.norm1 = nn.LayerNorm(dim)
self.attn = nn.MultiheadAttention(dim, num_heads, dropout=dropout, batch_first=True)
self.norm2 = nn.LayerNorm(dim)
self.mlp = Mlp(dim, hidden_dim=int(dim * mlp_ratio))
def forward(self, x):
# 自注意力 + 残差
x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]
# FFN + 残差
x = x + self.mlp(self.norm2(x))
return x
class Mlp(nn.Module):
"""前馈网络: 两个全连接 + 一个激活 (默认 GELU)。"""
def __init__(self, dim, hidden_dim, dropout=0.0):
super().__init__()
self.fc1 = nn.Linear(dim, hidden_dim)
self.act = nn.GELU()
self.fc2 = nn.Linear(hidden_dim, dim)
self.drop = nn.Dropout(dropout)
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x

4.3 前向传播示意#

输入 z_ℓ: (B, N+1, D)
├── LayerNorm → z_ℓ'
├── Multi-Head Self-Attention → attn_out
├── z_ℓ + attn_out (残差连接)
├── LayerNorm → z_ℓ''
├── MLP → mlp_out
├── z_ℓ + mlp_out (残差连接)
输出 z_{ℓ+1}: (B, N+1, D)

4.4 经典 ViT 变体配置#

模型隐层 DD头数 hhMLP 倍数Encoder 层数 LL参数量
ViT-Base76812412~86M
ViT-Large102416424~307M
ViT-Huge128016432~632M
ViT-Giant1408164.439~1B+

5. 归纳偏置:ViT vs CNN#

5.1 什么是归纳偏置 (Inductive Bias)?#

归纳偏置 = 模型架构对数据先验知识的”假设”,它由架构本身决定。

例如:

  • CNN 的归纳偏置:局部性 (Locality)平移不变性 (Translation Invariance)
  • Transformer 的归纳偏置:几乎没有——它从数据中学一切

5.2 详细对比#

维度CNNViT
局部先验强(卷积核只看局部)弱(每个 token 看全局)
平移不变性内置需数据学习
空间层次天然金字塔 (stride 池化)需专门设计 (Swin)
数据需求小数据也强需大数据预训练
可解释性特征图可视化直观注意力图可视化
感受野浅层小,深层才全局第一层就全局

5.3 为什么 ViT 需要大数据?#

CNN: ViT:
───────────────────── ─────────────────────
小卷积核 → 局部特征 全局自注意力 → 全局关系
│ │
▼ ▼
下采样 → 更大感受野 每个 token 位置编码
│ │
▼ ▼
+ 平移不变性 → 极强的归纳偏置 没有这些先验
│ │
▼ ▼
在小数据上也能训练 需要海量数据学会这些先验

ViT 论文中一个关键实验:

模型ImageNet-1k (从头训练)ImageNet-21k 预训练JFT-300M 预训练
ResNet152 (BiT)-85.3%87.5%
ViT-L76.5%85.3%87.8%

结论:在小数据上 BiT > ViT;但当预训练数据足够大(>14M 张图)时,ViT 反超 CNN。

6. 自注意力机制的数学原理#

6.1 注意力图 (Attention Map)#

对于输入序列 XR(N+1)×DX \in \mathbb{R}^{(N+1) \times D},自注意力计算每个 token 对其他所有 token 的”关联分数”:

Attention(Q,K,V)=softmax(QKdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right) V

其中 Q=XWQ,K=XWK,V=XWVQ = X W_Q, K = X W_K, V = X W_V

6.2 ViT 第一层可视化#

论文中一个经典的可视化:在 ViT 的第一层,某些 attention head 已经学会了:

  • 找相邻 patch (类似局部卷积)
  • 找整张图的远端 patch (全局关联)
  • 找同一物体的 patch
ViT 第一层某些 head 的注意力图可视化:
Head 1 (类似 conv): Head 7 (全局):
┌──────────┐ ┌──────────┐
│ ░░░█░░░░ │ │ █░█░█░█░ │
│ ░░░█░░░░ │ │ ░░░░░░░░ │
│ ░░░█░░░░ │ │ █░█░█░█░ │
│ ░░░░░░░░ │ │ ░░░░░░░░ │
└──────────┘ └──────────┘
只看纵向邻居 关注所有同列

惊人发现:ViT 没有内置局部性,但它在底层却自己学会了局部特征——再次说明”在大数据面前,归纳偏置可以被学习出来”。

6.3 多头注意力让 patch 关系更丰富#

def multi_head_attention(x, num_heads):
"""
把 D 维拆成 num_heads 个子空间,每个子空间学习一种"关系类型"。
"""
B, N, D = x.shape
d_head = D // num_heads
Q = self.q_proj(x).view(B, N, num_heads, d_head).transpose(1, 2) # (B, h, N, d_head)
K = self.k_proj(x).view(B, N, num_heads, d_head).transpose(1, 2)
V = self.v_proj(x).view(B, N, num_heads, d_head).transpose(1, 2)
# 每个 head 独立算注意力
attn = (Q @ K.transpose(-2, -1)) / (d_head ** 0.5) # (B, h, N, N)
attn = attn.softmax(dim=-1)
out = (attn @ V).transpose(1, 2).reshape(B, N, D) # 拼接所有头
return self.out_proj(out)

每个 head 可能负责不同语义关系——有的学颜色关联,有的学纹理关联,有的学物体部件关联。

7. 完整的 ViT 实现 (PyTorch)#

7.1 模型主体#

import torch
import torch.nn as nn
class VisionTransformer(nn.Module):
"""完整的 ViT 实现 (简化版)。"""
def __init__(
self,
img_size=224,
patch_size=16,
in_chans=3,
num_classes=1000,
embed_dim=768,
depth=12,
num_heads=12,
mlp_ratio=4.0,
drop_rate=0.0,
):
super().__init__()
self.num_features = embed_dim
self.num_patches = (img_size // patch_size) ** 2
# 1) Patch 嵌入
self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim)
# 2) [CLS] token
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
nn.init.trunc_normal_(self.cls_token, std=0.02)
# 3) 位置编码
self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches + 1, embed_dim))
nn.init.trunc_normal_(self.pos_embed, std=0.02)
self.pos_drop = nn.Dropout(drop_rate)
# 4) Transformer Encoder
self.blocks = nn.ModuleList([
Block(embed_dim, num_heads, mlp_ratio)
for _ in range(depth)
])
# 5) 最后的 LayerNorm
self.norm = nn.LayerNorm(embed_dim)
# 6) 分类头
self.head = nn.Linear(embed_dim, num_classes)
def forward(self, x):
B = x.shape[0]
# Patch 嵌入: (B, C, H, W) -> (B, N, D)
x = self.patch_embed(x)
# 拼接 [CLS] token
cls_tokens = self.cls_token.expand(B, -1, -1) # (B, 1, D)
x = torch.cat([cls_tokens, x], dim=1) # (B, N+1, D)
# 加位置编码
x = self.pos_drop(x + self.pos_embed)
# Transformer Encoder
for blk in self.blocks:
x = blk(x)
x = self.norm(x)
# 取 [CLS] token 输出
cls_out = x[:, 0]
return self.head(cls_out)

7.2 使用示例#

if __name__ == "__main__":
# ViT-Base/16 配置
model = VisionTransformer(
img_size=224,
patch_size=16,
embed_dim=768,
depth=12,
num_heads=12,
num_classes=1000,
)
# 测试前向
images = torch.randn(2, 3, 224, 224)
logits = model(images)
print(f"Output shape: {logits.shape}") # (2, 1000)
# 统计参数量
n_params = sum(p.numel() for p in model.parameters())
print(f"Total params: {n_params / 1e6:.1f}M")

7.3 FLOPs 与显存分析#

def calc_flops(model, img_size=224):
"""粗略估算 ViT 的 FLOPs。"""
patch_size = model.patch_embed.patch_size
num_patches = (img_size // patch_size) ** 2
dim = model.num_features
depth = len(model.blocks)
N = num_patches + 1 # 含 [CLS]
# QKV 线性投影: 3 * N * D^2
qkv_flops = depth * 3 * N * dim * dim
# 注意力矩阵乘法: 2 * N^2 * D (Q@K^T + Attn@V)
attn_flops = depth * 2 * N * N * dim
# FFN: 2 * N * D * 4D
ffn_flops = depth * 2 * N * dim * 4 * dim
total = (qkv_flops + attn_flops + ffn_flops) * 2 # MACs -> FLOPs
print(f"Total FLOPs: {total / 1e9:.2f}G")

8. 训练策略#

8.1 ViT 的训练配方#

阶段 1: 大规模预训练
├── 数据: ImageNet-21k (14M) 或 JFT-300M (300M)
├── 优化: AdamW, β1=0.9, β2=0.999
├── 学习率: 1e-3 ~ 3e-3 (warmup + cosine)
├── Weight Decay: 0.1
├── Batch Size: 4096 (跨多机)
├── Epochs: 7 (JFT) 或 90 (ImageNet-21k)
└── 数据增强: RandAugment, MixUp, CutMix
阶段 2: 中等分辨率微调
├── 输入: 224 -> 384 (位置编码插值)
├── 数据: 下游任务的小数据集
└── 学习率: 1e-4 左右

8.2 位置编码为什么可以”二维插值”?#

ViT 中的 1D 位置编码虽然是一维的,但在初始化时作者用 1D 网格的 2D 位置编码初始化

def _init_pos_embed(cls_token_pe, patch_pe_1d):
"""将 1D 序列位置编码 reshape 成 2D 网格,再 flatten。"""
# patch_pe_1d: (1, N, D)
grid_size = int(math.sqrt(N))
patch_pe_2d = patch_pe_1d.reshape(1, grid_size, grid_size, D)
return patch_pe_2d

这样得到的 1D 编码隐含了 2D 空间结构,因此插值时有意义。

8.3 数据增强的重要性#

必备增强:
✓ RandAugment
✓ MixUp (α=0.2)
✓ CutMix (α=1.0)
✓ Random Erasing
✓ Color Jitter
丢弃增强:
✗ 大量水平翻转 (对 ViT 影响小)
✗ 训练抖动不足会显著降低精度

关键经验:ViT 对数据增强极其敏感——没有强增强,精度会暴跌。

9. ViT 的局限与改进#

9.1 原版 ViT 的问题#

问题原因影响
需要海量数据没有归纳偏置小数据集训练困难
高分辨率贵注意力是 O(N2)O(N^2)384×384 已难以承受
单尺度所有 patch 同尺寸不适合密集预测任务
无层次结构token 始终等大难以替代 CNN backbone

9.2 后续重大改进#

9.2.1 DeiT (Data-efficient Image Transformer)#

Facebook 提出的改进,重点解决”小数据训练”。核心技巧:

  • 蒸馏 token (Distillation Token):让 ViT 学习来自 CNN 教师模型的输出
  • 更强的数据增强
  • 在 ImageNet-1k 上仅用单卡就能训出 84% 的 ViT
# DeiT 的特殊设计
class DistilledViT(VisionTransformer):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# 增加一个蒸馏 token
self.dist_token = nn.Parameter(torch.zeros(1, 1, self.num_features))
nn.init.trunc_normal_(self.dist_token, std=0.02)
# 双头: 一个用于 GT,一个用于教师蒸馏
self.head_dist = nn.Linear(self.num_features, self.num_classes)

9.2.2 Swin Transformer (ICCV 2021 Best Paper)#

引入层级化窗口注意力 (Window Attention)

  • 把图像分成不重叠的局部窗口
  • 窗口内做自注意力 (复杂度降为 O(N)O(N))
  • 通过 shifted window 让信息跨窗口流动
  • 金字塔结构 → 可以直接接 FPN → 适配检测/分割
Swin Transformer 金字塔:
Stage 1: 56×56×96 (每 token 计算量小)
Stage 2: 28×28×192
Stage 3: 14×14×384
Stage 4: 7×7×768

9.2.3 PVT (Pyramid Vision Transformer)#

类似 Swin,但用池化做降采样,结构更简单。

9.2.4 MAE (Masked Autoencoders, 2021)#

Kaiming He 提出的自监督预训练方法

  • 随机 mask 75% 的 patch
  • 让模型重建被 mask 掉的像素
  • 这种 pretrain-finetune 范式让 ViT 再也不需要 JFT 这种巨型数据集
class MAE(nn.Module):
"""简化的 MAE 结构。"""
def __init__(self, encoder, decoder_dim=512, mask_ratio=0.75):
super().__init__()
self.encoder = encoder # ViT
self.decoder = Transformer(...) # 轻量 decoder
self.mask_ratio = mask_ratio
def forward(self, images):
# 1. 随机 mask 75% 的 patch
patches = self.patch_embed(images)
N = patches.shape[1]
num_keep = int(N * (1 - self.mask_ratio))
# 保留 25%,mask 75%
ids_shuffle = torch.randperm(N)
keep_ids = ids_shuffle[:num_keep]
mask_ids = ids_shuffle[num_keep:]
# 2. 只对保留的部分过 encoder
x = self.encoder(patches[:, keep_ids])
# 3. decoder 重建
pred = self.decoder(x)
return pred, mask_ids # 只在被 mask 的 patch 上算 loss

9.3 ViT 改进谱系#

2020 ViT (Google) ──┐
├──> 2021 Swin (Microsoft, 层级化) ──┐
│ │
├──> 2021 DeiT (Meta, 数据高效) │
│ ├──> ViT 在 CV 全面
├──> 2022 MAE (Meta, 掩码自监督) │ 取代 CNN
│ │
├──> 2022 BeiT (微软) │
│ │
2021 Swin V2 ──────┘
2023 EVA, EVA-02 ───────────────────── 视觉大模型预训练
2023 DINOv2 (Meta) ──────────────────── 无监督视觉基础模型

10. ViT 的下游应用#

10.1 图像分类 (ViT 原任务)#

  • ImageNet 1k: 88.5% (ViT-G/14)
  • ImageNet 21k (38k 类): ViT-L 都能训
  • 关键:用更大的 patch_size(如 14) 替代 16,推理时切小 patch 做 fine-tune

10.2 目标检测#

  • DETR / Deformable DETR: 用 ViT 做 backbone + transformer 解码器
  • Mask2Former: ViT backbone + mask 分类
  • Swin: 主流检测 backbone

10.3 语义分割#

  • Mask2Former, SAM: 都基于 ViT backbone
  • 分割需要 dense prediction → 需要多尺度 → 多用 Swin 类

10.4 多模态大模型 (MLLM)#

最火的方向:

CLIP (2021): ViT + 文本 Transformer → 图文对齐
BLIP, BLIP-2: 引入 Q-Former 桥接视觉与语言
LLaVA: 把 ViT 直接接到 LLM 上
GPT-4V, Qwen-VL: ViT 作为 LLM 的"眼睛"
class MultiModalLLM(nn.Module):
"""简化的 ViT + LLM 多模态架构。"""
def __init__(self, vit, llm, projector):
super().__init__()
self.vit = vit # ViT-B/16 提取图像特征
self.projector = projector # 视觉 -> LLM 维度
self.llm = llm # 大语言模型
def forward(self, images, prompts):
# 图像编码
img_features = self.vit(images) # (B, N+1, D_vit)
img_tokens = self.projector(img_features[:, 1:]) # 去掉 [CLS]
# 与文本拼接,送入 LLM
return self.llm(prompts, vision_tokens=img_tokens)

10.5 自监督与基础模型#

DINO (2021): 自监督 ViT,attention map 自动分割物体
DINOv2 (2023): 7B 训练数据 → 通用视觉基础模型
SAM (2023): Segment Anything Model,ViT-H 解码器
EVA (2022): 视觉基础预训练,对齐 CLIP

11. ViT 与多模态时代#

11.1 ViT 是 LLM 时代的关键拼图#

2017 以来深度学习范式的演进:

2017 Transformer 架构统一 NLP
2020 ViT 让 Transformer 统一视觉
2021 CLIP 让视觉-语言对齐
2023 GPT-4V/LLaVA 让视觉成为 LLM 的"感官"
未来 统一的多模态基础模型 (Language + Vision + Audio + Action)

11.2 为什么 ViT 能成为多模态的”标配”?#

需求CNN 的短板ViT 的优势
与文本对齐图像/文本特征空间差异大token 序列结构一致
处理可变长度输入CNN 需固定尺寸ViT 天然支持 patch 数量变化
集成到 LLM模型架构差异大都是 Transformer block
注意力可视化主要是特征图attention map 可对齐到 patch

核心洞见:CLIP 能成功,关键之一就是 图像端也用了 Transformer——这样视觉和文本可以”无缝拼接”。

11.3 视觉大模型时代#

当前主流架构:
视觉编码器: ViT / Swin
视觉投影器: MLP / Q-Former
语言模型: LLaMA / Qwen / Mistral
输出端: 直接接 LLM 输出头

12. 总结与展望#

12.1 核心要点回顾#

维度关键要点
核心创新把图像切成 patch 当 token 处理
架构几乎与 NLP Transformer 相同
关键设计[CLS] token + 可学习 1D 位置编码
训练必须大数据预训练 + 中等数据微调
优势全局感受野、易扩展、与 LLM 架构一致
劣势数据饥饿、O(N2)O(N^2) 复杂度
后继Swin (层级化)、MAE (自监督)、EVA (基础)

12.2 ViT 带给我们的启示#

  1. 架构创新可以跨界:把 NLP 的架构挪到视觉上,居然成功
  2. 大数据 + 弱归纳偏置 > 小数据 + 强归纳偏置:在算力足够时,让模型自己学
  3. 通用性比专门性更值钱:统一架构 (Transformer) 推动了 LLM、ViT、多模态的协同发展

12.3 未来方向#

当前热点:
├── 线性复杂度注意力 (Mamba / Linear Attention) 加速 ViT
├── 视觉大模型预训练 (DINOv3, EVA-03)
├── 视觉-语言-动作统一模型 (VLA, 机器人)
├── 视频 ViT (TimeSformer, ViViT)
├── 3D 点云 Transformer
└── 端侧小 ViT (MobileViT, EfficientFormer)

12.4 推荐学习资源#

论文:
- ViT (2020): "An Image is Worth 16x16 Words"
- DeiT (2021): "Training data-efficient image transformers"
- Swin (2021): "Hierarchical Vision Transformer using Shifted Windows"
- MAE (2021): "Masked Autoencoders Are Scalable Vision Learners"
代码:
- timm 库 (HuggingFace 的 PyTorch Image Models)
- Hugging Face transformers (ViTForImageClassification)
- OpenMMLab (MMSegmentation / MMDetection)

一句话总结:Vision Transformer 通过”图像分块 + 标准 Transformer”这一看似简单的设计,成功把 NLP 的架构范式扩展到了视觉,开启了视觉大模型时代,也奠定了多模态 AI 的关键基石。

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

深入理解 Vision Transformer (ViT):图像识别的范式革新
https://aiattnstudio.link/posts/vision-transformer/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签
1
1. ViT 的核心思想
1.1 为什么要把 Transformer 引入视觉?
1.2 ViT 的核心思想
1.3 一个生活化的比喻
1.4 ViT 发展简史
2
2. 图像预处理:Patch 切分与线性嵌入
2.1 Patch 切分
2.2 线性嵌入 (Patch Embedding)
2.3 Patch Embedding 的代码实现
2.4 序列长度计算
3
3. 位置编码 (Positional Encoding)
3.1 为什么 ViT 需要位置编码?
3.2 ViT 的位置编码方案
3.3 [CLS] Token:借鉴 BERT 的设计
3.4 位置编码的插值问题
4
4. Transformer Encoder
4.1 整体结构
4.2 单层 Encoder 块
4.3 前向传播示意
4.4 经典 ViT 变体配置
5
5. 归纳偏置:ViT vs CNN
5.1 什么是归纳偏置 (Inductive Bias)?
5.2 详细对比
5.3 为什么 ViT 需要大数据?
6
6. 自注意力机制的数学原理
6.1 注意力图 (Attention Map)
6.2 ViT 第一层可视化
6.3 多头注意力让 patch 关系更丰富
7
7. 完整的 ViT 实现 (PyTorch)
7.1 模型主体
7.2 使用示例
7.3 FLOPs 与显存分析
8
8. 训练策略
8.1 ViT 的训练配方
8.2 位置编码为什么可以”二维插值”?
8.3 数据增强的重要性
9
9. ViT 的局限与改进
9.1 原版 ViT 的问题
9.2 后续重大改进
9.2.1 DeiT (Data-efficient Image Transformer)
9.2.2 Swin Transformer (ICCV 2021 Best Paper)
9.2.3 PVT (Pyramid Vision Transformer)
9.2.4 MAE (Masked Autoencoders, 2021)
9.3 ViT 改进谱系
10
10. ViT 的下游应用
10.1 图像分类 (ViT 原任务)
10.2 目标检测
10.3 语义分割
10.4 多模态大模型 (MLLM)
10.5 自监督与基础模型
11
11. ViT 与多模态时代
11.1 ViT 是 LLM 时代的关键拼图
11.2 为什么 ViT 能成为多模态的”标配”?
11.3 视觉大模型时代
12
12. 总结与展望
12.1 核心要点回顾
12.2 ViT 带给我们的启示
12.3 未来方向
12.4 推荐学习资源