拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Vision Transformer 源码分析:张量形状与 PyTorch 实战

Vision Transformer 源码分析:张量形状与 PyTorch 实战

做算法这行的应该都有过这种体验:论文读完了,公式看着也懂,但打开 Vision Transformer 的源码,看到那一堆 reshape、permute、transpose 就卡住了,不知道某个维度到底代表什么。这篇 Vision Transformer(下面统一叫 ViT)的代码分析,就是把我自己从「论文能读、代码发懵」到「能默写、能改、能调优」这一路的东西整理出来。内容偏向保姆级,会从张量形状这条主线讲起,把 Patch Embedding、Class Token、位置编码、多头自注意力、Pre-LN 残差块这些模块一个个拆开,然后给一份能直接跑的完整实现,再补上训练超参、参数量估算、显存占用、踩坑排查这些文档里基本不会写的部分。适合刚接触 ViT 想读懂源码的同学,也适合已经在用 timm 但要改结构、做小数据集微调的工程师。全文代码基于 PyTorch,不依赖任何特定训练框架,你复制到自己项目里稍微改改就能跑。

1. 读懂 ViT 源码之前,先建立三个基本认知

很多人读 ViT 代码读不下去,根本原因不是 Python 不熟,而是脑子里缺一张「形状地图」。ViT 的代码量其实很小,核心实现两百行出头,但它对张量维度的操作密度极高,一个 forward 里可能连着七八次 reshape 和 permute。所以这一节先把认知框架搭起来,后面读代码会顺很多。

1.1 ViT 到底把「卷积」换成了什么

卷积神经网络处理图像时,隐含了两条很强的先验:局部性(相邻像素相关)和平移等变性(物体挪个位置,特征图跟着挪)。这两个先验叫归纳偏置(inductive bias),它让 CNN 在数据量不大时也能学得不错。

ViT 的做法是把这两条先验几乎全部丢掉。它把图片硬切成固定大小的小方块(patch),每个 patch 拉直成一个向量,当成一个「词」,然后整套 Transformer Encoder 原封不动搬过来。注意力机制是全局的,第一个 patch 可以直接跟最后一个 patch 交互,中间没有任何局部性约束。

这个取舍带来的直接后果,也是读代码时必须记住的一句话:

ViT 的强项是「数据够多时上限高」,弱项是「数据少时容易学偏」。所以官方代码里那套重增强、长训练、强正则的配置不是可选项,而是结构决定的必需品。

理解了这一点,你再看代码里的 DropPath、Mixup、Label Smoothing、权重衰减 0.05 这些设置,就不会觉得是作者随手加的,而是对「缺先验」这件事的补偿。

1.2 张量形状变化是读代码的主线

我给自己的规矩是:读 ViT 源码时,只盯一个东西——张量形状。每个模块进去什么形状,出来什么形状,中间为什么变,全写下来。以标准的 ViT-Base、输入 224×224 为例,整条链路是这样:

阶段张量形状含义
输入(B, 3, 224, 224)原始图片,B 是 batch
Conv2d 切块(B, 768, 14, 14)196 个 patch,每个压成 768 维
flatten + transpose(B, 196, 768)变成序列,196 个 token
拼接 cls token(B, 197, 768)多一个全局 token
加位置编码(B, 197, 768)形状不变,只是数值相加
进入 Block(B, 197, 768)12 个 Block 形状都不变
LayerNorm 后取 [:, 0](B, 768)只取 cls token 作为图像表示
分类头(B, num_classes)输出 logits

这张表背下来,读任何 ViT 变体的代码都能找到锚点。你会发现在 ViT 里,除了最开始那次「图像变序列」和最后那次「序列变向量」,中间 12 个 Block 的形状是完全不动的——(B, 197, 768) 从头贯穿到尾。这一点跟 CNN 里特征图逐层变小、通道逐层变多完全不同,也是 ViT 代码看起来「平」的原因。

1.3 完整流程与模块清单

把上面的形状链路翻译成模块,一个 ViT 其实只有五个可复用的零件:

  • PatchEmbed:用一次 Conv2d 完成切块加线性投影,是整个模型里唯一跟图像空间结构打交道的部分。
  • Class Token 与 Position Embedding:两个可学习参数,负责给序列加「全局信息位」和「位置信息」。
  • Attention:多头自注意力的全部实现,包含 QKV 投影、缩放点积、输出投影。
  • MLP:两层全连接加 GELU,中间维度放大 4 倍。
  • Block:把 Attention 和 MLP 用 Pre-LN 残差串起来。

再加上一个最终 LayerNorm 和分类头,就是全部。我第一次按这个清单把代码重写一遍之后,再回头看 timm 的实现,基本上一眼就能对上——它无非是多了几层封装和一堆配置开关。

2. 核心模块逐行拆解:从 Patch Embedding 到 Encoder

这一节开始抠细节。我会按数据流动的顺序讲,每个模块都给出关键代码,并且解释「为什么这么写」而不是「写了什么」。这些「为什么」往往就是面试和生产环境里真正会出问题的地方。

2.1 Patch Embedding:一行 Conv2d 完成切图与线性投影

论文里描述 Patch Embedding 是「把图像切成不重叠的 patch,再对每个 patch 做线性映射」。如果严格按论文写,会是这样:先 reshape 成 (B, 196, 16×16×3),再过一个 Linear(768, 768)。但官方实现和 timm 都用了一个等价但更高效的小技巧:

class PatchEmbed(nn.Module): def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): super().__init__() assert img_size % patch_size == 0, "图像尺寸必须能被 patch 尺寸整除" self.num_patches = (img_size // patch_size) ** 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, 3, 224, 224) -> (B, 768, 14, 14) x = x.flatten(2) # (B, 768, 196) x = x.transpose(1, 2) # (B, 196, 768) return x

为什么用一个卷积就能替代「切块+线性映射」?关键在于卷积核尺寸等于步长等于 patch 尺寸,而且没有 padding。此时卷积核在图上滑动时,恰好一次覆盖一个 16×16 的不重叠区域,每个区域输出 768 个通道值——这不就是「把 16×16×3=768 维的 patch 映射到 768 维」吗?卷积核的权重就是那个线性层的权重,只不过被组织成了卷积的形式。

这么做的好处有两个:一是省掉了 unfold/reshape 这些显式操作,二是卷积在现代硬件和推理引擎上的优化程度远高于等价的矩阵乘法组合,导出到 ONNX 或 TensorRT 时也更友好。

注意:patch_size 必须能整除 img_size,否则会丢边或者报错。224/16=14 没问题,但如果你想把输入改成 220,就会出问题。改输入分辨率时先算一下整除关系,能省掉半小时的 debug 时间。

2.2 Class Token 与位置编码:两个容易写错的细节

序列准备好了,接下来要补两样东西。

第一样是 Class Token。它是一个形状为 (1, 1, 768) 的可学习参数,在序列最前面拼上去,变成 197 个 token。经过 12 层注意力之后,只取这个位置(索引 0)的输出送进分类头。为什么不在最后对所有 token 做平均池化?论文里试过,效果跟加 cls token 差不多,但 cls token 实现更简单、参数量更少。后来很多工作(比如 DeiT)干脆用平均池化,两者都行,别在这上面纠结。

第二样是位置编码。注意力机制本身是排列不变的——把 197 个 token 打乱顺序,结果一样。可是图像是有空间结构的,左上角的 patch 和右下角的 patch 不该被同等对待,所以要显式注入位置信息。

ViT 用的是可学习的一维位置编码,形状 (1, 197, 768),跟 cls token 一起相加。这里有两个细节特别容易踩:

第一,位置编码要和 cls token 一起参与拼接顺序。代码里是先 cat cls token 得到 197 个 token,再加 197 个位置编码。顺序反了就会报形状不匹配。

第二,位置编码的初始化不能用默认的。PyTorch 里 nn.Parameter 默认是均匀分布,而 ViT 需要用截断正态分布,标准差 0.02:

nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.02)

官方代码里对所有 Linear 层和 LayerNorm 层也有类似的初始化。我第一次自己复现时偷懒全用了默认初始化,结果 loss 从 6.9 卡着不动,排查了半天才发现是初始化的问题。ViT 对初始化比 CNN 敏感得多,这部分别省。

还有一个进阶细节,等你做高分辨率微调时会遇到:位置编码是按 14×14 的网格生成的,如果你把输入改成 384×384,patch 数变成 576,位置编码对不上了。这时候要做二维双三次插值,把 14×14 的网格放大到 24×24,再拉平。下面这段是必备工具函数:

def interpolate_pos_encoding(self, x, w, h): npatch = x.shape[1] - 1 N = self.pos_embed.shape[1] - 1 if npatch == N: return self.pos_embed dim = x.shape[-1] patch_size = self.patch_embed.proj.kernel_size[0] w0, h0 = w // patch_size, h // patch_size cls_pos = self.pos_embed[:, :1, :] grid_pos = self.pos_embed[:, 1:, :] grid_pos = grid_pos.reshape(1, int(N ** 0.5), int(N ** 0.5), dim) grid_pos = grid_pos.permute(0, 3, 1, 2) grid_pos = nn.functional.interpolate( grid_pos, size=(h0, w0), mode='bicubic', align_corners=False) grid_pos = grid_pos.permute(0, 2, 3, 1).reshape(1, -1, dim) return torch.cat([cls_pos, grid_pos], dim=1)

2.3 Multi-Head Self-Attention 的维度变换全过程

这是整个 ViT 里 reshape 最密集的地方,也是初学者最容易绕晕的部分。我把它拆成五步看。

class Attention(nn.Module): def __init__(self, dim=768, num_heads=12, qkv_bias=True, attn_drop=0.0, proj_drop=0.0): super().__init__() self.num_heads = num_heads self.head_dim = dim // num_heads # 768 / 12 = 64 self.scale = self.head_dim ** -0.5 # 1 / sqrt(64) = 0.125 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) self.proj = nn.Linear(dim, dim) self.proj_drop = nn.Dropout(proj_drop) def forward(self, x): B, N, C = x.shape # (B, 197, 768) qkv = self.qkv(x) # (B, 197, 2304) qkv = qkv.reshape(B, N, 3, self.num_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, 12, 197, 64) q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * self.scale # (B, 12, 197, 197) attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) x = (attn @ v) # (B, 12, 197, 64) x = x.transpose(1, 2).reshape(B, N, C) # (B, 197, 768) x = self.proj(x) x = self.proj_drop(x) return x

第一步,用一个 Linear(768, 2304) 一次性算出 Q、K、V。三个矩阵合成一次矩阵乘,比分开三次快,这是工程上的惯例,不是数学上的必要。

第二步,reshape 成 (B, 197, 3, 12, 64)。这里的 12 是 head 数,64 是每个 head 的维度,12×64=768。

第三步,permute 成 (3, B, 12, 197, 64)。为什么把 3 提到最前面?因为这样才能用 qkv[0]、qkv[1]、qkv[2] 一次性拆开,比 split 更直观。把 head 维提到 batch 之后,是因为后续矩阵乘法要在最后两维上做,每个 head 独立计算。

第四步,q @ k.transpose(-2, -1),得到 (B, 12, 197, 197) 的注意力矩阵。注意最后两维是 197×197,表示每个 token 对其他所有 token 的注意力权重。这个矩阵是 ViT 显存占用的主要来源之一。

第五步,乘以 scale 再 softmax。scale 是 1/sqrt(head_dim),这里等于 0.125。为什么要缩放?因为 Q 和 K 的每个元素方差大致是 1,做 64 维点积后方差会变成 64,数值过大会让 softmax 进入饱和区,梯度趋近于零。除以 sqrt(64)=8 把方差拉回 1。

提醒:有些实现会把 self.scale 写成 self.head_dim ** -0.5,有些写成 1.0 / math.sqrt(head_dim),数值上一样。但如果你看到有人用 dim ** -0.5(即 768 的负 0.5 次方),那是 bug,会让注意力分布过于平滑。这个错我见过不止一次。

2.4 MLP、残差与 LayerNorm:Pre-LN 为什么更稳

Attention 之后接的是 MLP。ViT 用的是两层全连接,中间维度放大 4 倍,激活函数是 GELU:

hidden = int(dim * 4.0) # 768 -> 3072 self.mlp = nn.Sequential( nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(drop), nn.Linear(hidden, dim), nn.Dropout(drop), )

这个 4 倍是个经验值,后来有工作(比如 Swin)证明 2 倍也能用,但对 ViT 这个原始结构,4 倍是标配。值得注意的是,MLP 的参数量实际上比 Attention 还大:768×3072 + 3072×768 ≈ 4.7M,而 Attention 里 qkv 是 768×2304 ≈ 1.77M,proj 是 768×768 ≈ 0.59M,合计 2.36M。所以你在做模型剪枝或者算显存的时候,别只盯着注意力。

接下来是 LayerNorm 和残差的位置。原始 Transformer 用的是 Post-LN:先做子层,再残差,最后归一化,也就是x = LN(x + sublayer(x))。但 ViT 和后来的绝大多数 Transformer 都改成了 Pre-LN:

x = x + drop_path(self.attn(self.norm1(x))) x = x + drop_path(self.mlp(self.norm2(x)))

为什么改?Post-LN 在深层网络里,输出层的方差会随着深度累积放大,训练必须靠精心调过的 warmup 才能收敛,稍微换个学习率就崩。Pre-LN 把归一化放在子层之前,残差路径上是一条干净的恒等映射,梯度可以直接从最后一层回传到第一层,深层训练稳定得多。代价是理论上表达力略弱,但在 ViT 这种 12 到 24 层的规模下完全不是问题。

还有个小细节:ViT 里 LayerNorm 的 eps 是 1e-6,而 PyTorch 默认是 1e-5。虽然差别不大,但既然要复现,就按官方的来。

Block 里还有一个容易忽略的组件:DropPath,也叫随机深度(stochastic depth)。它跟普通 Dropout 不一样,普通 Dropout 是随机丢神经元,DropPath 是随机把整个残差分支的输出置零。作用是在深网络里给每个 Block 一个「跳过」的机会,起到正则和加速收敛的作用。

def drop_path(x, drop_prob=0.0, training=False): if drop_prob == 0.0 or not training: return x keep_prob = 1 - drop_prob shape = (x.shape[0],) + (1,) * (x.ndim - 1) mask = x.new_empty(shape).bernoulli_(keep_prob) return x.div(keep_prob) * mask

注意里面的x.div(keep_prob),这是为了保证训练和推理时的期望一致,跟 Dropout 的 inverted dropout 是同一个思路。漏了这一步,训练和推理的输出尺度会差一个系数,表现为训练集准确率还行、验证集掉点。

而且 DropPath 的概率不是每层都一样,官方用的是从 0 线性增加到 0.1 的调度:

dpr = torch.linspace(0, drop_path_rate, depth).tolist()

越靠后的 Block 丢弃概率越大。逻辑是浅层学到的是通用特征,不该丢;深层学到的是任务相关的细节,过拟合风险高,可以多丢一些。

3. 从零手写一个能跑通的 ViT(附完整代码)

前面拆完了零件,这一节把它们装成一台能转的机器。我会给出一份完整可运行的实现,并且把环境、数据、训练循环、超参都写清楚。这部分代码我实际在单卡 24G 显存上跑过 CIFAR-100 和自定义的小数据集,能收敛。

3.1 环境准备与依赖版本选择

依赖就三样:

pip install torch torchvision pip install timm # 只用来做数据增强和加载预训练权重

PyTorch 版本建议 1.10 以上,因为用到了 torch.cuda.amp 和较稳定的 LayerNorm 实现。timm 是必须要装的,哪怕你要自己写模型——它里面的 RandAugment、Mixup、CutMix 实现都经过大量验证,自己写容易出细节问题。数据增强这块我强烈建议别重复造轮子,效果差异往往来自这些不起眼的地方。

至于要不要直接读 timm 的源码?我的建议是:先用自己写的版本跑通一遍,再去读 timm。因为 timm 为了兼容上百个模型变体做了大量抽象,第一次读很容易迷失在继承关系里。自己写完再看,会发现它其实就是把你这几百行代码参数化了。

3.2 模型代码:拆成五个组件写

把 2.x 节的内容组装起来,完整模型如下:

import torch import torch.nn as nn class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio=4.0, drop=0.0, attn_drop=0.0, drop_path=0.0): super().__init__() self.norm1 = nn.LayerNorm(dim, eps=1e-6) self.attn = Attention(dim, num_heads, attn_drop=attn_drop, proj_drop=drop) self.norm2 = nn.LayerNorm(dim, eps=1e-6) hidden = int(dim * mlp_ratio) self.mlp = nn.Sequential( nn.Linear(dim, hidden), nn.GELU(), nn.Dropout(drop), nn.Linear(hidden, dim), nn.Dropout(drop), ) self.drop_path_rate = drop_path def forward(self, x): x = x + drop_path(self.attn(self.norm1(x)), self.drop_path_rate, self.training) x = x + drop_path(self.mlp(self.norm2(x)), self.drop_path_rate, self.training) return x class ViT(nn.Module): 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, attn_drop_rate=0.0, drop_path_rate=0.1): super().__init__() self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim) n = self.patch_embed.num_patches self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, n + 1, embed_dim)) self.pos_drop = nn.Dropout(drop_rate) dpr = torch.linspace(0, drop_path_rate, depth).tolist() self.blocks = nn.ModuleList([ Block(embed_dim, num_heads, mlp_ratio, drop_rate, attn_drop_rate, dpr[i]) for i in range(depth) ]) self.norm = nn.LayerNorm(embed_dim, eps=1e-6) self.head = nn.Linear(embed_dim, num_classes) self.apply(self._init_weights) nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.02) def _init_weights(self, m): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.LayerNorm): nn.init.zeros_(m.bias) nn.init.ones_(m.weight) def forward_features(self, x): x = self.patch_embed(x) cls = self.cls_token.expand(x.shape[0], -1, -1) x = torch.cat((cls, x), dim=1) x = x + self.pos_embed x = self.pos_drop(x) for blk in self.blocks: x = blk(x) x = self.norm(x) return x[:, 0] def forward(self, x): return self.head(self.forward_features(x))

几点使用说明。num_classes 换成你自己的类别数就行,分类头是唯一需要改维度的地方。drop_path_rate 在小数据集上可以调到 0.1 到 0.2,大数据集上 0.0 到 0.1 即可。如果你想拿它做特征提取器,直接调用 forward_features 拿到 (B, 768) 的向量,接个下游头就能做检测或者检索。

3.3 数据管线与增强策略配置

ViT 对数据增强的依赖比 CNN 重得多,这是结构决定的。我常用的配置是这样:

from timm.data import create_transform, Mixup, CutMix train_transform = create_transform( input_size=224, is_training=True, color_jitter=0.4, auto_augment='rand-m9-mstd0.5-inc1', interpolation='bicubic', re_prob=0.25, re_mode='pixel', re_count=1, ) val_transform = create_transform(input_size=224, is_training=False)

这里每个参数都有理由。RandAugment 的 m9 表示做 9 次随机增强操作,mstd0.5 控制强度抖动,这套在 ViT 上比单纯的翻转裁剪明显更有效。Random Erasing 的概率 0.25 是为了模拟遮挡,逼模型不要依赖单个局部区域。bicubic 插值比默认的双线性在 ViT 上表现略好,这是官方实验里验证过的。

训练时还要加 Mixup 和 CutMix,二选一或按概率切换:

mixup_fn = Mixup(mixup_alpha=0.8, cutmix_alpha=1.0, prob=1.0, switch_prob=0.5, label_smoothing=0.1, num_classes=num_classes)

心得:小数据集(一万到十万张这个量级)上,Mixup 和 CutMix 的收益非常明显,经常能让验证集准确率涨 3 到 5 个点。但要注意它们会拉长收敛时间,训练轮数得相应增加,别训练 30 个 epoch 看效果不好就放弃了。

3.4 训练循环、超参与参数量估算

优化器用 AdamW,不是 SGD。这一点跟 CNN 的习惯不同,但 ViT 对优化器很敏感,用 SGD 往往收敛得很慢甚至不收敛。

import math epochs, warmup_epochs = 100, 5 optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.05, betas=(0.9, 0.999)) def lr_lambda(epoch): if epoch < warmup_epochs: return (epoch + 1) / warmup_epochs progress = (epoch - warmup_epochs) / max(1, epochs - warmup_epochs) return 0.5 * (1 + math.cos(math.pi * progress)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda) scaler = torch.cuda.amp.GradScaler()

为什么必须有 warmup?因为训练初期模型输出极不稳定,注意力权重接近均匀分布,此时梯度方向噪声大。如果一上来就用 3e-4 的学习率,容易把参数推到坏区域。用 5 个 epoch 线性爬坡,让模型先找到一个大致的下降方向,再进入余弦衰减。这套组合在 ViT 上几乎是默认答案。

训练循环的关键部分:

for epoch in range(epochs): model.train() for images, targets in train_loader: images, targets = images.cuda(), targets.cuda() images, targets = mixup_fn(images, targets) with torch.cuda.amp.autocast(): outputs = model(images) loss = criterion(outputs, targets) optimizer.zero_grad() scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update() scheduler.step()

梯度裁剪的 max_norm 设成 1.0。有人觉得 Transformer 不需要梯度裁剪,其实在混合精度训练下,梯度溢出是常见现象,裁剪是很便宜的保险。

现在算一下参数量,这个对你判断显存和模型规模很有用。ViT-Base 的配置是 embed_dim=768、depth=12、num_heads=12:

组件计算式参数量
Patch Embedding3×16×16×768 + 768约 0.59M
Position Embedding197×768约 0.15M
单个 Block 的 Attention768×2304 + 2304 + 768×768 + 768约 2.36M
单个 Block 的 MLP768×3072 + 3072 + 3072×768 + 768约 4.72M
12 个 Block 合计(2.36 + 4.72) × 12约 85.0M
分类头768×1000 + 1000约 0.77M
总计约 86.6M

和官方公布的 ViT-Base 86M 对得上。顺带说,ViT-Large 是 embed_dim=1024、depth=24,参数量约 307M;ViT-Huge 是 embed_dim=1280、depth=32,约 632M。显存紧张的话优先动 patch_size,从 16 改成 32 能让 token 数从 196 降到 49,注意力矩阵从 197×197 降到 50×50,显存降幅接近一个数量级,代价是精度会掉一些。

4. 实操踩坑实录:ViT 训练中那些反直觉的现象

代码写完只是开始,真正花时间的是调通。这一节记录的是我自己踩过的坑,以及后来帮别人排查时反复遇到的几类问题。里面有的在文档里根本找不到,但确实很浪费生命。

4.1 Loss 不降、准确率卡住:先查这五处

遇到 loss 不降,按下面顺序排查,能覆盖八成情况。

第一,检查位置编码和 cls token 的初始化。这是最高频的原因。忘了 trunc_normal_ 初始化,用默认的均匀分布,loss 经常从 6.9 附近纹丝不动,或者降到 4.6 就下不去了。加两行初始化代码重新跑,通常就能动。

第二,检查 scale 系数。确认是 head_dim ** -0.5 而不是 dim ** -0.5 或者漏乘。这个错误特别隐蔽,因为模型能跑、loss 也能降一点,只是最后准确率明显偏低。

第三,检查 LayerNorm 的位置。如果把 norm 写成了x = self.norm1(x + self.attn(x))这种 Post-LN 形式,模型学到一半发散的概率会增加不少。区分方法很简单:看残差路径上有没有归一化。

第四,检查学习率和 warmup。学习率超过 1e-3 在 ViT-Base 上基本必崩;完全没有 warmup 也很容易在头几个 epoch 出现 loss 突然飙升到 nan。

第五,检查数据增强是不是过强。这个情况有点反直觉——增强本来是为了泛化,但如果你在只有几千张图的数据集上把 RandAugment 开到 m9 再加 Mixup,模型可能连训练集都拟合不了,表现为训练和验证准确率都很低。这时候先关掉 Mixup,把增强强度调下来,确认模型能过拟合一个小批次(比如 20 张图),再逐步加回去。

这里有个通用的调试手段值得记下来:从训练集里取 8 到 16 张图,关掉所有增强,用大学习率训练几百步,看能不能把训练准确率打到 100%。如果打不到,说明是模型代码本身有问题,跟数据增强、学习率调度都无关。这一步能帮你快速把问题范围缩小一半。

4.2 显存不够、跑不动:几种有效的降显存手段

ViT 的显存开销主要来自三块:激活值(尤其是注意力矩阵)、中间张量、优化器状态。按性价比排序,可以这样处理。

降低 batch size 最直接,但会拖慢训练并影响 BN 类统计(ViT 用 LayerNorm,倒是没这个问题)。混合精度(AMP)能省 30% 到 40% 显存,几乎是必开的,同时速度也更快。

梯度检查点(gradient checkpointing)省显存效果最猛,能到 50% 以上,代价是训练速度慢 20% 到 30%。用法就一行:

model = torch.utils.checkpoint.checkpoint_wrapper(model)

不过要注意,用 checkpoint 之后 drop_path 里的随机性需要额外的随机数保存机制,不然两次前向的结果不一致,会影响训练。稳妥一点的做法是自己在 Block 的 forward 里用torch.utils.checkpoint.checkpoint,并给preserve_rng_state=True。

调小输入分辨率或者调大 patch_size 是另一个维度的手段。前面算过,patch 从 16 改成 32,token 数变成原来的四分之一,注意力矩阵变成十六分之一。做原型验证或者跑对比实验时,用 128×128 输入配 patch 8,或者 224 配 patch 32,能让你在单卡上把流程先跑通。

优化器换成 SGD 或者 8-bit Adam 也能省一部分,因为 AdamW 要为每个参数保存一阶和二阶动量,占参数量的两倍。8-bit Adam 能把这块压到四分之一。

4.3 常见问题速查表

现象高概率原因快速验证方式处理方式
Loss 从 6.9 开始不动位置编码或 cls token 初始化错误打印这两个参数的标准差加 trunc_normal_(std=0.02)
Loss 突然变 nan学习率过大、无 warmup、AMP 梯度溢出关掉 AMP 用小 lr 重跑加 warmup、开梯度裁剪 1.0
训练准确率高、验证集低DropPath 或 Mixup 配置不当、过拟合关掉 Mixup 看验证集变化调高 weight_decay 和 drop_path
训练到后期突然崩学习率衰减到接近 0 时数值不稳观察 lr 曲线加最小 lr 下限或者调短总轮数
换分辨率后报形状错误位置编码长度不匹配检查 pos_embed 的 shape[1]加插值函数
显存占用远高于预期batch 过大、未开 AMP、未设 no_gradnvidia-smi 看峰值开 AMP 和梯度检查点
训练速度慢得离谱数据加载成瓶颈、没用 pin_memory单独测一个 epoch 的加载耗时num_workers 设成 8,开 pin_memory
多卡训练结果变差学习率没随总 batch 放大对比单卡和多卡配置lr 按 sqrt 或线性缩放

表里最后一行顺便展开说一句。多卡时总 batch 变大,学习率一般要跟着调,ViT 上常见的做法是线性缩放(batch 翻倍,lr 翻倍),但也有实验表明用 sqrt 缩放更稳。我的经验是先用线性缩放,如果前几个 epoch 有发散迹象就改用 sqrt。

5. 从 ViT 往外延伸:结构选型与落地场景

把 ViT 跑通只是第一步,真正到项目里会遇到「该不该用它」的问题。这一节聊几个我实际做过选型判断的场景,包括跟 CNN 系的对比、小数据集的打法,以及部署阶段的注意点。

5.1 ViT 与 CNN 系(EfficientNetV2 等)怎么选

先给结论:数据量小于十万张、没有预训练权重、又要求推理速度,优先考虑 EfficientNetV2 这类 CNN 或者混合结构;数据量足够大或者能拿到大规模预训练权重,ViT 系上限更高。

下面这张表是我自己踩过坑之后总结的对比,参考的是公开实验结论和实际项目体感:

维度ViT-BaseEfficientNetV2-M
归纳偏置几乎没有强(卷积局部性)
小数据表现差,容易过拟合好,收敛快
大数据上限高中等
单张推理延迟较高,受 token 数影响较低
显存占用高,注意力是平方复杂度中等
迁移到新任务需要较长时间微调微调快,改动小
对增强的敏感度高,必须配强增强中等

举个具体场景。假设你要做一个面部表情识别的任务,数据是几万张标注好的人脸图,类别七类左右。这个数据量对 ViT 来说偏小,从零训练很容易过拟合,验证集准确率可能在 60% 多就卡住。EfficientNetV2 在同样数据上往往能更轻松地到 65% 以上。但如果换成从 ImageNet-21k 预训练的 ViT 权重开始微调,加上 Mixup 和 Label Smoothing,结果通常能反超 CNN,能到 68% 到 70% 这个区间。

所以选型的关键不在于哪个结构更先进,而在于你手上的数据规模和可用预训练权重。这也是为什么近两年很多工作走混合路线:前面几层用卷积做下采样和局部特征提取,后面接 Transformer 做全局建模。Swin、Convolutional Vision Transformer 这些都属于这个思路,本质上是用卷积的局部性补上 ViT 缺的先验。

5.2 小数据集场景下的迁移学习与调参思路

如果你手上就是几万张图,又确实想用 ViT,下面这套流程我试过几次都有效。

第一步,加载预训练权重。用 timm 加载是最省事的:

import timm model = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=0) # 去掉分类头 model = model.eval()

num_classes=0 让 timm 直接返回 768 维的特征,自己接一个 Linear 做下游任务。这样方便你做线性探测——先冻结整个 backbone,只训分类头,看看到底有多少信息量可提取。

第二步,先做线性探测再全量微调。线性探测通常几十个 epoch 就能收敛,能帮你判断是特征质量问题还是微调策略问题。如果线性探测的准确率就明显低于预期,说明预训练特征跟你的任务域差距大,得考虑换个预训练数据源。

第三步,全量微调时把学习率调小。ViT 微调的标准学习率是 1e-4 到 5e-5 这个区间,比从头训练小一个数量级。同时用较小的 weight_decay(0.05 降到 0.01),因为预训练权重的分布已经很好了,没必要再施加太大的正则。

第四步,如果过拟合依然严重,先把浅层冻结。ViT 的浅层学到的是边缘、纹理这类通用特征,冻结它们能显著减少可训练参数量:

for name, param in model.named_parameters(): if 'blocks.0' in name or 'blocks.1' in name or 'patch_embed' in name: param.requires_grad = False

第五步,Label Smoothing 设成 0.1,DropPath 设成 0.1 到 0.2,这两个在小数据集上几乎是无脑开。

有个细节值得说:微调时不要随便改输入分辨率。预训练权重的位置编码是按 224 生成的,你如果直接用 128 输入训练,要么插值位置编码,要么就接受精度损失。稳妥做法是保持 224,通过调整 batch 和其他参数来适配显存。

5.3 推理部署阶段的几个关键点

模型训好了,部署还有几个坑。

ONNX 导出时,注意力里的 permute 和 reshape 组合对某些推理引擎不太友好。我建议导出前先把模型包装一层,固定输入形状(batch=1,尺寸固定),这样能避免动态形状带来的额外开销和算子支持问题:

dummy = torch.randn(1, 3, 224, 224) torch.onnx.export(model, dummy, 'vit.onnx', input_names=['input'], output_names=['logits'], opset_version=13)

opset 建议 13 以上,因为低版本对 LayerNormalization 和 GELU 的支持不够好,会被拆成十几个基础算子,推理速度明显变慢。

量化方面要小心。ViT 对量化比 CNN 敏感,尤其是 LayerNorm 和最后的分类头。做 INT8 量化时,我一般只量化卷积和全连接,把 LayerNorm 和 softmax 留在浮点。直接整网量化,精度掉 3 到 5 个点是常事。

另外,ViT 的推理延迟跟输入分辨率是近似平方关系(token 数线性增长,注意力是平方)。如果你的场景是实时视频流,输入用 224 而且 patch_size 设成 16,单帧延迟在主流显卡上大概十几毫秒,还行;但如果你的输入是 448 或者 512,延迟会翻好几倍,这时候要么换 patch_size,要么考虑混合结构。

还有一些工程上的小技巧。批推理比单张推理吞吐高得多,能做批就做批。序列长度固定时可以把位置编码直接烘焙进模型,省掉一次相加。如果只是做特征提取,注意关掉 Dropout 和 DropPath(调用 model.eval() 就够了),否则每次推理结果都不一样,做检索时会出问题。

我个人在实际项目里的体会是,ViT 这类模型的价值不在于「比 CNN 强多少」,而在于它提供了一个统一的、可迁移的建模框架。你在图像上验证过的注意力结构,换个输入表征就能用到别的模态上,这种横向迁移能力才是它真正被广泛采用的原因。至于代码层面,把张量形状这条主线抓住,再多的变体也只是在这个骨架上做加减法——这大概是我读完十几份 ViT 变体实现之后最实在的一条经验。

返回列表