ViT 刚出现时并不被看好:它没有卷积的局部性先验,没有平移等变性,在小数据集上打不过 ResNet。但当成千上万张图堆上来时,它反超了——归纳偏置不是被证明无用,而是被证明可以被数据替代。随后的自监督预训练进一步把「需要多少标注」这个问题也解开了。这两条线合在一起,构成了当下视觉骨干的主流路线。
视觉 Transformer 的真正转折点不是架构,而是预训练范式的改变。MAE、DINO、CLIP 用无标注数据学出的表征,在下游任务上超过了有监督预训练。架构只是提供了可扩展的容器,数据与目标函数才是关键。
从 CNN 到 ViT
归纳偏置的取舍
| 特性 | CNN | ViT |
|---|---|---|
| 局部性 | 内置(卷积核) | 无,需学习 |
| 平移等变 | 内置 | 无 |
| 感受野 | 逐层扩大 | 全局(第一层就是) |
| 数据需求 | 小数据即可 | 大数据才发挥 |
| 扩展性 | 中等 | 极好 |
结论:数据少时用 CNN 或混合架构,数据多时用 ViT。这也是为什么 ViT 论文必须在 JFT-300M 这种量级上才能打平 ResNet——小数据集上它连收敛都困难。
混合架构的现实价值
ConvNeXt 与 Swin 证明了中间路线依然有效:在浅层保留卷积的局部性,在深层用注意力做全局建模。工程上,如果预训练数据规模不到千万级,混合架构往往是更稳的选择。
补丁嵌入与位置编码
把图像切成序列
ViT 的第一步是把 H×W×C 的图像切成 P×P 的补丁,展平后线性投影成 token:
import torch
import torch.nn as nn
class PatchEmbed(nn.Module):
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
super().__init__()
self.grid = img_size // patch_size
self.n_patches = self.grid ** 2
# 用卷积实现,等价于切块 + 线性投影,且快得多
self.proj = nn.Conv2d(in_chans, embed_dim,
kernel_size=patch_size, stride=patch_size)
def forward(self, x):
x = self.proj(x) # (B, D, G, G)
x = x.flatten(2).transpose(1, 2) # (B, N, D)
return x
16×16 补丁是最常见选择:224/16 = 14,得到 196 个 token。补丁越小,token 越多,精度越高但算力平方增长。
位置编码的几种方案
注意力本身是排列不变的,必须显式注入位置信息:
| 方案 | 形式 | 特点 |
|---|---|---|
| 可学习绝对编码 | 每个位置一个向量 | 简单,但换分辨率要插值 |
| 正弦绝对编码 | 固定三角函数 | 可外推,效果略逊 |
| 相对位置偏置 | 注意力里加偏置 | Swin 用,效果好 |
| RoPE 二维 | 旋转位置编码 | 现代 ViT 主流 |
def interpolate_pos_embed(pos_embed, new_grid, old_grid=14):
"""换分辨率时对位置编码做双三次插值"""
cls_token, rest = pos_embed[:, :1], pos_embed[:, 1:]
rest = rest.reshape(1, old_grid, old_grid, -1).permute(0, 3, 1, 2)
rest = torch.nn.functional.interpolate(
rest, size=(new_grid, new_grid), mode="bicubic", align_corners=False)
rest = rest.permute(0, 2, 3, 1).reshape(1, new_grid * new_grid, -1)
return torch.cat([cls_token, rest], dim=1)
位置编码的插值是 ViT 部署里最常见的坑:预训练是 224,推理用 384,忘了插值就会直接报形状错误或精度暴跌。用 RoPE 可以回避这个问题,因为它是相对编码,天然支持长度外推。
CLS token 与池化
ViT 在序列前加一个 [CLS] token,用它的输出做分类。后续研究发现平均池化(GAP)往往更好,尤其在下游密集预测任务上。DINOv2 同时保留两种用法:CLS 用于分类,patch token 平均用于分割。
注意力机制在视觉中的形态
标准全局注意力
class MultiHeadAttention(nn.Module):
def __init__(self, dim, n_heads, qkv_bias=True):
super().__init__()
self.n_heads = n_heads
self.scale = (dim // n_heads) ** -0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.proj = nn.Linear(dim, dim)
def forward(self, x):
B, N, D = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.n_heads, D // self.n_heads)
q, k, v = qkv.permute(2, 0, 3, 1, 4)
attn = (q @ k.transpose(-2, -1)) * self.scale
attn = attn.softmax(dim=-1)
out = (attn @ v).transpose(1, 2).reshape(B, N, D)
return self.proj(out)
窗口注意力:Swin 的层次化设计
全局注意力的复杂度是 O(N²),高分辨率下不可接受。Swin 把注意力限制在局部窗口内,并做层次化降采样:
Stage 1: 56×56 token, 7×7 窗口
Stage 2: 28×28 token (patch merging 降采样)
Stage 3: 14×14 token
Stage 4: 7×7 token
窗口间信息靠 Shifted Window 传递:下一层的窗口偏移半个窗口大小,让原本不相邻的 token 有机会交互。
def window_partition(x, window_size):
B, H, W, C = x.shape
x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
windows = x.permute(0, 1, 3, 2, 4, 5).contiguous()
return windows.view(-1, window_size, window_size, C)
def window_reverse(windows, window_size, H, W):
B = int(windows.shape[0] / (H * W / window_size / window_size))
x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
return x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
层次化设计让 Swin 天然适配检测与分割——不同 stage 的输出对应不同尺度的特征图,可以直接接 FPN。这也是它比 ViT 更适合密集预测的原因,相关任务可参考 计算机视觉 中的检测与分割章节。
训练配方
ViT 的成败很大程度取决于训练配方,而非架构本身。
数据增强
| 增强 | 作用 | 强度 |
|---|---|---|
| RandomResizedCrop | 尺度不变性 | 强 |
| Mixup / CutMix | 正则化 | 中 |
| RandAugment | 通用增强 | 中强 |
| Random Erasing | 遮挡鲁棒 | 中 |
| 颜色抖动 | 颜色不变性 | 弱 |
强增强 + 长训练是 ViT 的标配。用 ResNet 的轻增强配方训练 ViT,效果会差一大截。
正则化组合
def build_optimizer(model, lr=1e-3, weight_decay=0.05, layer_decay=0.75):
"""分层学习率衰减:浅层学习率小,深层大"""
param_groups = []
n_layers = len(model.blocks)
for name, param in model.named_parameters():
depth = get_layer_depth(name, n_layers)
scale = layer_decay ** (n_layers - depth)
param_groups.append({"params": [param], "lr": lr * scale,
"weight_decay": weight_decay if param.ndim > 1 else 0.0})
return torch.optim.AdamW(param_groups)
三个关键点:
- AdamW 而非 Adam:解耦权重衰减,对 Transformer 更稳。
- bias 与 norm 层不加权重衰减。
- 分层学习率衰减:微调时浅层用小学习率,避免破坏预训练特征。
随机深度与 drop path
class DropPath(nn.Module):
def __init__(self, p=0.1):
super().__init__()
self.p = p
def forward(self, x):
if self.p == 0.0 or not self.training:
return x
keep = 1 - self.p
mask = torch.rand(x.shape[0], 1, 1, device=x.device) < keep
return x * mask / keep
Drop path 对深层 ViT 是必需的——没有它,24 层以上的 ViT 几乎无法收敛。衰减率通常从 0 线性增到 0.1~0.4。
掩码自编码:MAE
MAE 把 NLP 的掩码语言建模搬到视觉,但做了一个关键改动:掩码比例高达 75%。
为什么高掩码比例有效
图像有极强的空间冗余——相邻像素高度相关。如果只掩 15%(BERT 的做法),模型靠邻域插值就能重建,学不到语义。掩到 75% 后,插值不再可行,模型必须理解全局结构。
非对称编码解码
MAE 的另一半创新是只把可见补丁送进编码器:
class MAE(nn.Module):
def __init__(self, encoder, decoder_dim=512, mask_ratio=0.75):
super().__init__()
self.encoder = encoder
self.mask_ratio = mask_ratio
self.decoder_embed = nn.Linear(encoder.embed_dim, decoder_dim)
self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
self.decoder = build_decoder(decoder_dim)
def forward(self, x):
patches = self.encoder.patch_embed(x) # (B, N, D)
B, N, D = patches.shape
n_keep = int(N * (1 - self.mask_ratio))
noise = torch.rand(B, N, device=x.device)
ids_shuffle = torch.argsort(noise, dim=1)
ids_keep = ids_shuffle[:, :n_keep]
visible = torch.gather(patches, 1, ids_keep.unsqueeze(-1).expand(-1, -1, D))
latent = self.encoder.forward_features(visible) # 只算可见部分,省 3/4 算力
# 解码时把 mask token 填回原位
full = self.mask_token.expand(B, N, -1).clone()
full.scatter_(1, ids_keep.unsqueeze(-1).expand(-1, -1, decoder_dim),
self.decoder_embed(latent))
recon = self.decoder(full)
return recon, ids_shuffle, n_keep
编码器只处理 25% 的 token,训练速度比全量编码快约 3 倍——这是 MAE 能扩展到 ViT-Huge 的关键。
损失只算被掩位置
def mae_loss(recon, target, ids_shuffle, n_keep, patch_size):
target = patchify(target, patch_size) # (B, N, p*p*3)
target = normalize_pixels(target)
loss = (recon - target) ** 2
loss = loss.mean(dim=-1) # (B, N)
mask = torch.ones_like(loss)
mask.scatter_(1, ids_shuffle[:, :n_keep], 0.0) # 可见位置置 0
return (loss * mask).sum() / mask.sum()
只在被掩位置计算损失很重要:如果也算可见位置,模型会倾向于学「复制输入」,退化成自编码器。
自蒸馏:DINO 与 DINOv2
DINO 的核心机制
DINO 不需要负样本、不需要重建,靠学生-教师自蒸馏学表征:
- 教师是学生的指数滑动平均(EMA)。
- 同一张图做两种增强,学生看局部、教师看全局。
- 学生预测教师的输出分布,用交叉熵对齐。
@torch.no_grad()
def ema_update(student, teacher, m=0.996):
for ps, pt in zip(student.parameters(), teacher.parameters()):
pt.data.mul_(m).add_(ps.data, alpha=1 - m)
def dino_loss(student_out, teacher_out, temp_s=0.1, temp_t=0.04, center=None):
s = (student_out / temp_s).log_softmax(dim=-1)
t = (teacher_out - center) / temp_t
t = t.softmax(dim=-1)
return -(t * s).sum(dim=-1).mean()
防止塌陷的三个技巧
自蒸馏最容易塌陷——学生和教师一起输出常数。DINO 用三个机制避免:
| 机制 | 作用 |
|---|---|
| 温度锐化 | 教师温度更低,输出更尖锐 |
| 中心化 | 减去教师输出的均值,防止某个维度主导 |
| 多裁剪 | 学生看多个局部裁剪,教师看全局 |
中心化的更新也必须是 EMA:
def update_center(center, teacher_out, momentum=0.9):
return momentum * center + (1 - momentum) * teacher_out.mean(dim=0)
DINOv2 的工程化
DINOv2 在 DINO 基础上做了三件事:更大的数据(LVD-142M 自建数据集)、更强的增强、以及蒸馏到小模型。它的表征在密集任务上表现极好,且无需微调就能直接用——这对工程很有吸引力,省掉了每个下游任务重新训练的环节。
对比学习:从 SimCLR 到 CLIP
三条路线
| 方法 | 负样本来源 | 显存需求 |
|---|---|---|
| SimCLR | 同批次其他样本 | 极大(batch 4096+) |
| MoCo | 队列 + 动量编码器 | 小 |
| BYOL | 无负样本 | 小 |
| CLIP | 图文配对 | 大 |
def nt_xent(z1, z2, temperature=0.5):
"""SimCLR 的 InfoNCE 损失"""
z1 = torch.nn.functional.normalize(z1, dim=-1)
z2 = torch.nn.functional.normalize(z2, dim=-1)
N = z1.shape[0]
z = torch.cat([z1, z2], dim=0) # (2N, D)
sim = z @ z.T / temperature # (2N, 2N)
sim.fill_diagonal_(-1e9)
# 正样本:i 与 i+N 互为对方
labels = torch.arange(N, device=z.device)
labels = torch.cat([labels + N, labels])
return torch.nn.functional.cross_entropy(sim, labels)
CLIP 的双塔与零样本
CLIP 用图文对比学习把图像与文本映射到同一空间,从而实现零样本分类:把类别名做成文本 prompt,选相似度最高的。它的表征也是 多模态部署 的基础组件。CLIP 的局限同样明显:对细粒度分类弱,对计数与空间关系不敏感。
下游微调与知识蒸馏
微调策略选择
| 数据量 | 策略 | 说明 |
|---|---|---|
| 极少(<100) | 线性探针 | 冻结主干,只训分类头 |
| 少(100~10k) | 只调后几层 | 保护浅层通用特征 |
| 中(10k~100k) | 全量微调 + 分层 lr | 标准做法 |
| 多(>100k) | 从头训或全量微调 | 预训练收益递减 |
蒸馏到小模型
部署时往往需要小模型。蒸馏比直接训小模型效果好得多:
def distill_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.5):
hard = torch.nn.functional.cross_entropy(student_logits, labels)
soft = torch.nn.functional.kl_div(
torch.nn.functional.log_softmax(student_logits / T, dim=-1),
torch.nn.functional.softmax(teacher_logits / T, dim=-1),
reduction="batchmean") * (T ** 2)
return alpha * hard + (1 - alpha) * soft
T² 的缩放很重要——它补偿了温度带来的梯度量级变化,让不同温度下的损失可比。蒸馏常与 模型压缩
中的量化、剪枝组合使用。
部署与效率
主要开销
| 环节 | 开销 | 优化 |
|---|---|---|
| 注意力 | O(N²) | FlashAttention、窗口注意力 |
| 高分辨率推理 | token 数平方增长 | 分块推理、动态分辨率 |
| 显存 | 激活值大 | 梯度检查点(训练) |
高分辨率推理的分块
输入 1024×1024 时 token 数达到 4096,全局注意力显存会爆。做法是滑动窗口分块推理再拼接:
@torch.no_grad()
def tiled_inference(model, img, tile=224, stride=168):
"""重叠分块推理,重叠区取平均,缓解接缝"""
B, C, H, W = img.shape
out_sum = torch.zeros(B, model.num_classes, H, W, device=img.device)
count = torch.zeros(1, 1, H, W, device=img.device)
for y in range(0, H - tile + 1, stride):
for x in range(0, W - tile + 1, stride):
patch = img[:, :, y:y + tile, x:x + tile]
pred = model(patch)
out_sum[:, :, y:y + tile, x:x + tile] += pred
count[:, :, y:y + tile, x:x + tile] += 1
return out_sum / count.clamp(min=1)
动态分辨率
现代 ViT(如 NaViT、Qwen-VL)支持把不同分辨率的图打包进同一批次,用块对角注意力掩码隔离不同样本。这样既避免了缩放失真,又保持了批次效率。
排错清单
- 换分辨率后精度暴跌:位置编码没插值,或插值方式与预训练不一致。
- 训练 loss 不降:增强太弱或没有 drop path。ViT 对增强强度非常敏感。
- 小数据集上过拟合:改线性探针或只调后几层,加更强的权重衰减。
- MAE 重建模糊:损失算在了可见位置,模型退化成复制。检查 mask 计算。
- DINO 塌陷:中心化未更新,或温度设置错误。监控教师输出的熵。
- 对比学习不收敛:batch 太小,负样本不足。改用 MoCo 队列或 BYOL。
- 推理显存 OOM:高分辨率全图推理。改分块推理或降低分辨率。
- 注意力图全均匀:位置编码被错误初始化或学习率过高,注意力退化成平均池化。
小结
视觉 Transformer 的演进讲了一个清晰的道理:架构决定上限,数据与目标函数决定能否触及上限。ViT 提供了可扩展的容器,Swin 补上了效率与多尺度,MAE 让训练算力降到可接受,DINO 与 CLIP 让无标注数据变得可用。工程落地时的关键决策——用不用混合架构、选哪条自监督路线、微调还是线性探针、如何蒸馏到小模型——都取决于你的数据规模与部署约束,而非架构本身的先进程度。它与 卷积网络 的关系不是替代而是互补,在多模态系统中更是与语言模型深度耦合。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。