视觉 Transformer 与自监督预训练

视觉 Transformer 把注意力机制引入图像,用自监督预训练摆脱了标注依赖。本文系统讲解 ViT 的补丁嵌入与位置编码、Swin 的层次化窗口注意力、MAE 掩码重建、DINO 自蒸馏与对比学习三条自监督路线、下游微调与知识蒸馏策略,以及部署效率与常见踩坑。

ViT 刚出现时并不被看好:它没有卷积的局部性先验,没有平移等变性,在小数据集上打不过 ResNet。但当成千上万张图堆上来时,它反超了——归纳偏置不是被证明无用,而是被证明可以被数据替代。随后的自监督预训练进一步把「需要多少标注」这个问题也解开了。这两条线合在一起,构成了当下视觉骨干的主流路线。

视觉 Transformer 的真正转折点不是架构,而是预训练范式的改变。MAE、DINO、CLIP 用无标注数据学出的表征,在下游任务上超过了有监督预训练。架构只是提供了可扩展的容器,数据与目标函数才是关键。

从 CNN 到 ViT

归纳偏置的取舍

特性CNNViT
局部性内置(卷积核)无,需学习
平移等变内置无
感受野逐层扩大全局(第一层就是)
数据需求小数据即可大数据才发挥
扩展性中等极好

结论:数据少时用 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 让无标注数据变得可用。工程落地时的关键决策——用不用混合架构、选哪条自监督路线、微调还是线性探针、如何蒸馏到小模型——都取决于你的数据规模与部署约束,而非架构本身的先进程度。它与 卷积网络 的关系不是替代而是互补,在多模态系统中更是与语言模型深度耦合。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. 排序学习与搜索召回排序系统
  2. 数据版本控制与血缘:DVC 与 LakeFS
  3. 模型可解释性:SHAP、LIME 与注意力归因