尧图网站设计 尧图网站设计YAOTU DESIGN
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 是 batchConv2d 切块(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_size224, patch_size16, in_chans3, embed_dim768): 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_sizepatch_size, stridepatch_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×3768 维的 patch 映射到 768 维」吗卷积核的权重就是那个线性层的权重只不过被组织成了卷积的形式。这么做的好处有两个一是省掉了 unfold/reshape 这些显式操作二是卷积在现代硬件和推理引擎上的优化程度远高于等价的矩阵乘法组合导出到 ONNX 或 TensorRT 时也更友好。注意patch_size 必须能整除 img_size否则会丢边或者报错。224/1614 没问题但如果你想把输入改成 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.02nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02)官方代码里对所有 Linear 层和 LayerNorm 层也有类似的初始化。我第一次自己复现时偷懒全用了默认初始化结果 loss 从 6.9 卡着不动排查了半天才发现是初始化的问题。ViT 对初始化比 CNN 敏感得多这部分别省。还有一个进阶细节等你做高分辨率微调时会遇到位置编码是按 14×14 的网格生成的如果你把输入改成 384×384patch 数变成 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), modebicubic, align_cornersFalse) grid_pos grid_pos.permute(0, 2, 3, 1).reshape(1, -1, dim) return torch.cat([cls_pos, grid_pos], dim1)2.3 Multi-Head Self-Attention 的维度变换全过程这是整个 ViT 里 reshape 最密集的地方也是初学者最容易绕晕的部分。我把它拆成五步看。class Attention(nn.Module): def __init__(self, dim768, num_heads12, qkv_biasTrue, attn_drop0.0, proj_drop0.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, biasqkv_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×64768。第三步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、残差与 LayerNormPre-LN 为什么更稳Attention 之后接的是 MLP。ViT 用的是两层全连接中间维度放大 4 倍激活函数是 GELUhidden 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.77Mproj 是 768×768 ≈ 0.59M合计 2.36M。所以你在做模型剪枝或者算显存的时候别只盯着注意力。接下来是 LayerNorm 和残差的位置。原始 Transformer 用的是 Post-LN先做子层再残差最后归一化也就是x LN(x sublayer(x))。但 ViT 和后来的绝大多数 Transformer 都改成了 Pre-LNx 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_prob0.0, trainingFalse): 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_ratio4.0, drop0.0, attn_drop0.0, drop_path0.0): super().__init__() self.norm1 nn.LayerNorm(dim, eps1e-6) self.attn Attention(dim, num_heads, attn_dropattn_drop, proj_dropdrop) self.norm2 nn.LayerNorm(dim, eps1e-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_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0, drop_rate0.0, attn_drop_rate0.0, drop_path_rate0.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, eps1e-6) self.head nn.Linear(embed_dim, num_classes) self.apply(self._init_weights) nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) def _init_weights(self, m): if isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std0.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), dim1) 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_size224, is_trainingTrue, color_jitter0.4, auto_augmentrand-m9-mstd0.5-inc1, interpolationbicubic, re_prob0.25, re_modepixel, re_count1, ) val_transform create_transform(input_size224, is_trainingFalse)这里每个参数都有理由。RandAugment 的 m9 表示做 9 次随机增强操作mstd0.5 控制强度抖动这套在 ViT 上比单纯的翻转裁剪明显更有效。Random Erasing 的概率 0.25 是为了模拟遮挡逼模型不要依赖单个局部区域。bicubic 插值比默认的双线性在 ViT 上表现略好这是官方实验里验证过的。训练时还要加 Mixup 和 CutMix二选一或按概率切换mixup_fn Mixup(mixup_alpha0.8, cutmix_alpha1.0, prob1.0, switch_prob0.5, label_smoothing0.1, num_classesnum_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(), lr3e-4, weight_decay0.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_norm1.0) scaler.step(optimizer) scaler.update() scheduler.step()梯度裁剪的 max_norm 设成 1.0。有人觉得 Transformer 不需要梯度裁剪其实在混合精度训练下梯度溢出是常见现象裁剪是很便宜的保险。现在算一下参数量这个对你判断显存和模型规模很有用。ViT-Base 的配置是 embed_dim768、depth12、num_heads12组件计算式参数量Patch Embedding3×16×16×768 768约 0.59MPosition Embedding197×768约 0.15M单个 Block 的 Attention768×2304 2304 768×768 768约 2.36M单个 Block 的 MLP768×3072 3072 3072×768 768约 4.72M12 个 Block 合计(2.36 4.72) × 12约 85.0M分类头768×1000 1000约 0.77M总计约 86.6M和官方公布的 ViT-Base 86M 对得上。顺带说ViT-Large 是 embed_dim1024、depth24参数量约 307MViT-Huge 是 embed_dim1280、depth32约 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_stateTrue。调小输入分辨率或者调大 patch_size 是另一个维度的手段。前面算过patch 从 16 改成 32token 数变成原来的四分之一注意力矩阵变成十六分之一。做原型验证或者跑对比实验时用 128×128 输入配 patch 8或者 224 配 patch 32能让你在单卡上把流程先跑通。优化器换成 SGD 或者 8-bit Adam 也能省一部分因为 AdamW 要为每个参数保存一阶和二阶动量占参数量的两倍。8-bit Adam 能把这块压到四分之一。4.3 常见问题速查表现象高概率原因快速验证方式处理方式Loss 从 6.9 开始不动位置编码或 cls token 初始化错误打印这两个参数的标准差加 trunc_normal_(std0.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, pretrainedTrue, num_classes0) # 去掉分类头 model model.eval()num_classes0 让 timm 直接返回 768 维的特征自己接一个 Linear 做下游任务。这样方便你做线性探测——先冻结整个 backbone只训分类头看看到底有多少信息量可提取。第二步先做线性探测再全量微调。线性探测通常几十个 epoch 就能收敛能帮你判断是特征质量问题还是微调策略问题。如果线性探测的准确率就明显低于预期说明预训练特征跟你的任务域差距大得考虑换个预训练数据源。第三步全量微调时把学习率调小。ViT 微调的标准学习率是 1e-4 到 5e-5 这个区间比从头训练小一个数量级。同时用较小的 weight_decay0.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.1DropPath 设成 0.1 到 0.2这两个在小数据集上几乎是无脑开。有个细节值得说微调时不要随便改输入分辨率。预训练权重的位置编码是按 224 生成的你如果直接用 128 输入训练要么插值位置编码要么就接受精度损失。稳妥做法是保持 224通过调整 batch 和其他参数来适配显存。5.3 推理部署阶段的几个关键点模型训好了部署还有几个坑。ONNX 导出时注意力里的 permute 和 reshape 组合对某些推理引擎不太友好。我建议导出前先把模型包装一层固定输入形状batch1尺寸固定这样能避免动态形状带来的额外开销和算子支持问题dummy torch.randn(1, 3, 224, 224) torch.onnx.export(model, dummy, vit.onnx, input_names[input], output_names[logits], opset_version13)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 变体实现之后最实在的一条经验。
返回列表