深入理解 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 的核心思想可以概括为三句话:
- 把图像切成小块 (Patch) → 模拟 NLP 中的”词 Token”
- 每个块线性映射成向量 → 相当于”词嵌入”
- 送入标准 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 切分
假设输入图像尺寸为 ,其中 是通道数(RGB 图像为 3)。我们把它切分成 的小块,则有:
例如:, ,则 个 patch。
每个 patch 展平后是 维的向量,例如 维。
2.2 线性嵌入 (Patch Embedding)
每个 patch 的展平向量经过一个可学习的线性投影层,映射到 维( 是 Transformer 的隐层维度):
其中 是可学习的投影矩阵, 是位置编码。
2.3 Patch Embedding 的代码实现
实际上,Patch + Linear Projection 可以优雅地用一个二维卷积完成:
import torchimport 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 Size | Patch 数量 (N) | Token 序列长度 |
|---|---|---|---|
| 224×224 | 16 | 196 | |
| 384×384 | 16 | 576 | |
| 224×224 | 14 | 256 | |
| 1024×1024 | 32 | 1024 |
可以看到,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 相加:
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 通常使用 图像(位置编码表有 个条目)。如果推理时换用 ,则需要 个位置编码——怎么办?
标准做法是 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 完全一致,堆叠 层:
每层包含两个子模块:1. LayerNorm → Multi-Head Self-Attention → Residual2. LayerNorm → MLP (FFN) → Residual4.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 x4.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 变体配置
| 模型 | 隐层 | 头数 | MLP 倍数 | Encoder 层数 | 参数量 |
|---|---|---|---|---|---|
| ViT-Base | 768 | 12 | 4 | 12 | ~86M |
| ViT-Large | 1024 | 16 | 4 | 24 | ~307M |
| ViT-Huge | 1280 | 16 | 4 | 32 | ~632M |
| ViT-Giant | 1408 | 16 | 4.4 | 39 | ~1B+ |
5. 归纳偏置:ViT vs CNN
5.1 什么是归纳偏置 (Inductive Bias)?
归纳偏置 = 模型架构对数据先验知识的”假设”,它由架构本身决定。
例如:
- CNN 的归纳偏置:局部性 (Locality) 和平移不变性 (Translation Invariance)
- Transformer 的归纳偏置:几乎没有——它从数据中学一切
5.2 详细对比
| 维度 | CNN | ViT |
|---|---|---|
| 局部先验 | 强(卷积核只看局部) | 弱(每个 token 看全局) |
| 平移不变性 | 内置 | 需数据学习 |
| 空间层次 | 天然金字塔 (stride 池化) | 需专门设计 (Swin) |
| 数据需求 | 小数据也强 | 需大数据预训练 |
| 可解释性 | 特征图可视化直观 | 注意力图可视化 |
| 感受野 | 浅层小,深层才全局 | 第一层就全局 |
5.3 为什么 ViT 需要大数据?
CNN: ViT:───────────────────── ─────────────────────小卷积核 → 局部特征 全局自注意力 → 全局关系 │ │ ▼ ▼下采样 → 更大感受野 每个 token 位置编码 │ │ ▼ ▼+ 平移不变性 → 极强的归纳偏置 没有这些先验 │ │ ▼ ▼在小数据上也能训练 需要海量数据学会这些先验ViT 论文中一个关键实验:
| 模型 | ImageNet-1k (从头训练) | ImageNet-21k 预训练 | JFT-300M 预训练 |
|---|---|---|---|
| ResNet152 (BiT) | - | 85.3% | 87.5% |
| ViT-L | 76.5% | 85.3% | 87.8% |
结论:在小数据上 BiT > ViT;但当预训练数据足够大(>14M 张图)时,ViT 反超 CNN。
6. 自注意力机制的数学原理
6.1 注意力图 (Attention Map)
对于输入序列 ,自注意力计算每个 token 对其他所有 token 的”关联分数”:
其中 。
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 torchimport 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 的问题
| 问题 | 原因 | 影响 |
|---|---|---|
| 需要海量数据 | 没有归纳偏置 | 小数据集训练困难 |
| 高分辨率贵 | 注意力是 | 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):
- 把图像分成不重叠的局部窗口
- 在窗口内做自注意力 (复杂度降为 )
- 通过 shifted window 让信息跨窗口流动
- 金字塔结构 → 可以直接接 FPN → 适配检测/分割
Swin Transformer 金字塔:
Stage 1: 56×56×96 (每 token 计算量小)Stage 2: 28×28×192Stage 3: 14×14×384Stage 4: 7×7×7689.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 上算 loss9.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): 视觉基础预训练,对齐 CLIP11. 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 架构一致 |
| 劣势 | 数据饥饿、 复杂度 |
| 后继 | Swin (层级化)、MAE (自监督)、EVA (基础) |
12.2 ViT 带给我们的启示
- 架构创新可以跨界:把 NLP 的架构挪到视觉上,居然成功
- 大数据 + 弱归纳偏置 > 小数据 + 强归纳偏置:在算力足够时,让模型自己学
- 通用性比专门性更值钱:统一架构 (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 的关键基石。
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!

