ARTICLE DETAIL

资讯详情

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

VIT源码逐行解析:PatchEmbed、多头注意力与训练调参

VIT源码逐行解析:PatchEmbed、多头注意力与训练调参 1. 先搞清楚VIT到底在解决什么问题1.1 从CNN的局限说起如果你之前一直在做图像分类大概率用的是ResNet、EfficientNet这类卷积网络。卷积的核心思路是局部感受野加上权重共享一层一层地把局部特征往上抽象。这套逻辑统治了计算机视觉将近十年效果也确实好。但卷积有个天然的约束它的感受野是逐步扩大的。浅层只看得到几个像素的范围中层看到几十个像素要到很深的层才能“看到”整张图。这意味着一张图里相距很远的两个区域之间的关系卷积网络要经过很多层才能建立起来。对于“猫的耳朵和尾巴在图像两端”这种情况卷积得堆够深度才行。VITVision Transformer直接把自然语言处理那一套Transformer搬到了图像上。它的做法极其暴力也极其优雅把图片切成一个个小方块patch每个patch拉平后当成一个“词”然后整个序列丢进Transformer的编码器里。自注意力机制让每个patch从第一层就能和所有其他patch交互全局关系一步到位。我第一次读VIT论文的时候最震撼的不是它的效果而是它的简洁——模型结构几乎没有针对图像做任何特殊设计纯靠数据量堆出来的性能。但简洁不等于简单代码里有很多细节值得反复琢磨。这篇内容我就把自己啃VIT源码的过程完整记录下来从整体架构到每一行关键代码尽量讲透。1.2 这份代码分析适合谁看我写这份分析的时候预设的读者是这样的你已经用过PyTorch写过训练循环大概知道Transformer是什么东西但没仔细读过VIT的实现想知道每一行代码为什么要这么写。如果你是纯小白连卷积都没写过建议先把PyTorch的官方图像分类教程跑一遍再来看。如果你已经是老手可以直接跳到第3、4节看核心模块的代码拆解。代码版本我用的是timm库里的VIT实现timm/models/vision_transformer.py这是目前工程上最常用的版本之一比原论文的JAX代码和Facebook的DETR仓库里的版本更清晰、更易读。提示阅读源码时建议打开一个能跳转定义的IDE边看边跳效率会高很多。纯看文本很容易迷路。2. VIT整体架构拆解——先看地图再进森林2.1 数据在VIT里的完整流转路径在钻进代码之前先建立一张全局地图。一张224x224的RGB图片进入VIT后经历的过程大致如下切分成16x16的小块共(224/16)² 196个patch每个patch拉平成16×16×3 768维向量通过一个线性层映射到模型维度比如768加上一个可学习的位置编码在最前面拼接一个可学习的分类tokenCLS token进入L层Transformer Encoder每层包含多头自注意力和MLP取出CLS token对应的输出经过LayerNorm和线性分类头得到类别预测整个流程可以用一句话概括图像分类问题被转化成了序列分类问题。这是理解VIT所有代码的钥匙。关键点在于VIT没有使用任何卷积操作除了patch切分本身所有的空间关系都靠位置编码和自注意力来学习。这也是为什么VIT在小数据集上表现不如CNN——它缺少卷积那种内置的归纳偏置局部性、平移等变性必须靠大量数据来学习这些规律。2.2 核心模块的职责划分打开vision_transformer.py你会看到几个核心类类名职责代码位置PatchEmbed把图片切成patch并线性映射文件前部Attention多头自注意力计算中间部分Mlp两层全连接加激活Attention之后Block组合Attention和Mlp加残差紧随Mlp之后VisionTransformer顶层模型串联所有模块文件后部我刚开始看的时候犯过一个错误试图从VisionTransformer.forward往下逐行读结果被层层嵌套搞晕了。后来发现更好的方式是先从最底层的PatchEmbed和Attention单独看懂再回头看顶层如何串联。这就像拼乐高先把每块积木搞清楚再按照图纸拼。注意timm版本和原论文有一个细节差异——timm默认使用了qkv_biasTrue而原论文是没有的。这个小差异会影响你复现论文精度时的结果后面在注意事项里会详细说。3. 逐行拆解PatchEmbed与位置编码3.1 PatchEmbed把图片切成“词”的核心实现PatchEmbed的代码其实很短但每一行都有讲究。简化后的实现大概是这样class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.img_size img_size self.patch_size patch_size 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, embed_dim, H/P, W/P) x x.flatten(2) # (B, embed_dim, num_patches) x x.transpose(1, 2) # (B, num_patches, embed_dim) return x这里最精妙的地方在于作者用了一个nn.Conv2d来实现patch切分和线性映射。为什么因为当你把卷积核大小设为patch大小、步长也设为patch大小时每个卷积核恰好覆盖一个patch输出就是每个patch的线性变换结果。这样写比手动切分再reshape再过全连接层要高效得多也简洁得多。我实测过一个细节如果你手动实现patch切分用unfold再matmul速度大概慢15%到20%。所以这个Conv2d的写法不是炫技是真正有工程意义的优化。对于224x224的输入patch_size16时num_patches 196。这个数字会直接影响后面位置编码的参数量也会影响计算复杂度——自注意力的复杂度是O(n²)n从196变成784patch_size8时意味着计算量涨了16倍。这是选择patch大小时必须权衡的核心参数。3.2 位置编码给每个patch发一张“座位票”位置编码的代码看起来简单但背后有一个容易被忽略的知识点self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim))注意那个1——这是给CLS token预留的位置。CLS token放在序列最前面它的位置编码也是可学习的。为什么要加位置编码因为自注意力本身是排列不变的permutation-invariant。如果不加位置信息把patch顺序打乱模型的输出完全一样。位置编码就是告诉模型“这个patch在左上角”“那个patch在右下角”。VIT用的是可学习的绝对位置编码而不是Transformer原文的正弦位置编码。这两种方式的区别在实践中有影响正弦编码可以外推到训练时没见过的序列长度但可学习编码不行。所以如果你的输入分辨率从224变成384位置编码就需要插值处理。timm里有专门的resize_pos_embed函数来做这件事用的是双线性插值。我踩过的坑直接把224训练好的模型改成384输入如果不做位置编码插值模型精度会暴跌。不是因为模型学不会新位置而是因为位置编码的维度对不上直接报错。提示初始化位置编码时timm用的是截断正态分布std0.02不是全零。这个细节在小数据集上对收敛速度有肉眼可见的影响。3.3 CLS Token与分类头设计CLS token的设计直接借用了BERT的做法。它的初始化也是一个可学习参数self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim))在forward里它被复制batch_size份后拼接到patch序列前面cls_tokens self.cls_token.expand(x.shape[0], -1, -1) x torch.cat((cls_tokens, x), dim1)为什么要有CLS token一种解释是self-attention会让每个token的输出混合所有token的信息但最终我们需要一个全局表示来做分类。CLS token经过多层注意力后它本身就变成了整张图的聚合表示。另一种解释是如果不加CLS token用所有patch输出的平均池化也能达到类似效果而且在某些任务上平均池化甚至更好。我自己的实验结论是小数据集上用平均池化更稳大数据集上CLS token的天花板更高。timm里通过class_token和global_pool参数都支持可以灵活切换。分类头的结构也很简单self.head nn.Linear(embed_dim, num_classes)但在最终的forward里分类前会先过一个LayerNormx self.norm(x) # LayerNorm x x[:, 0] # 取CLS token x self.head(x) # 分类头这个LayerNorm很关键。训练深层Transformer时如果不加最后的归一化预训练的模型直接拿来微调容易数值不稳定。4. Transformer Encoder内部机制与代码实现4.1 多头自注意力VIT最核心的20行代码Attention类是VIT里最值得逐行细读的部分。核心代码简化后如下class Attention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasFalse, attn_drop0., proj_drop0.): super().__init__() self.num_heads num_heads head_dim dim // num_heads self.scale head_dim ** -0.5 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 qkv self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) qkv qkv.permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] attn (q k.transpose(-2, -1)) * self.scale attn attn.softmax(dim-1) attn self.attn_drop(attn) x (attn v).transpose(1, 2).reshape(B, N, C) x self.proj(x) x self.proj_drop(x) return x这20行代码里有几个设计细节值得说透第一QKV用一个线性层一次性算出来。nn.Linear(dim, dim*3)比三个独立的线性层快因为GPU对大矩阵乘法的利用率更高。这是一个纯工程优化不影响数学等价性。第二scale因子是head_dim ** -0.5。为什么要缩放因为点积的方差随维度增长不做缩放的话softmax会进入饱和区梯度趋近于零。假设q和k的每个元素独立同分布、均值为0方差为1那么q·k的方差就是head_dim。除以sqrt(head_dim)后方差变成1softmax的输入保持在合理范围。第三attn矩阵的形状是(B, num_heads, N, N)。对于224输入、patch16、num_heads12这个矩阵是(12, 197, 197)。这个矩阵就是可解释性的来源——你可以把它可视化看到每个patch在关注哪些其他patch。我实际调试时发现一个现象在浅层注意力矩阵对角线附近的值比较大说明patch主要关注自己附近的邻居。到了深层注意力变得非常分散CLS token会关注整张图的各个区域。这符合直觉——浅层学局部特征深层做全局整合。4.2 MLP层与Dropout的工程细节Block里的MLP部分看起来平平无奇但有个隐藏的比例参数class Mlp(nn.Module): def __init__(self, in_features, hidden_featuresNone, out_featuresNone, act_layerGELU, drop0.): out_features out_features or in_features hidden_features hidden_features or in_features self.fc1 nn.Linear(in_features, hidden_features) self.act act_layer() self.fc2 nn.Linear(hidden_features, out_features) self.drop nn.Dropout(drop)在VIT-Base里hidden_features通常是in_features的4倍768变3072。这个4倍比例是从原版Transformer继承来的后来很多工作证明它不是最优的比如SwiGLU结构用2/3的比例效果更好。但VIT作为经典实现沿用了4倍。激活函数用的是GELU而不是ReLU。GELU在零点附近是平滑的梯度不会突变对深层网络的训练稳定性有好处。我实测过把GELU换成ReLU小模型上差异不大但VIT-Large这种大模型上GELU的收敛明显更稳。4.3 残差连接与LayerNorm的位置选择Block的完整结构是class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., ...): self.norm1 nn.LayerNorm(dim) self.attn Attention(dim, num_headsnum_heads, ...) self.norm2 nn.LayerNorm(dim) self.mlp Mlp(in_featuresdim, hidden_featuresint(dim * mlp_ratio), ...) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x注意LayerNorm的位置——它在Attention和MLP之前这叫Pre-Norm结构。原版Transformer用的是Post-Normx norm(x attn(x))但VIT用的是Pre-Norm。为什么因为Pre-Norm的梯度流更顺畅不需要学习率预热就能稳定训练深层网络。Post-Norm在层数多了之后必须配合warmup否则很容易梯度爆炸。VIT-Base有12层用Pre-Norm可以省去很多调参的麻烦。但Pre-Norm也有代价它的表达能力略弱于Post-Norm。在一些对精度要求极高的场景下Post-Norm配合精心设计的warmup能拿到更好的最终精度。这是一个典型的工程权衡。注意如果你自己从零实现VITLayerNorm的位置搞错会让训练完全跑不起来。我见过有人把norm放在残差之后结果loss完全不动。5. 完整训练流程与关键参数配置5.1 数据增强与预处理流水线VIT对数据增强非常敏感。原论文用了RandAugment、Mixup、CutMix、Random Erasing一套组合拳在小数据集上微调时这些增强策略直接决定了模型能不能收敛。我常用的训练预处理from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.3, 0.3, 0.3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), transforms.RandomErasing(p0.25), ])RandomResizedCrop的scale下限我调到0.6而不是默认的0.08。为什么因为VIT没有卷积的平移等变性过度裁剪会让patch内容变得过于破碎模型很难从中学到有意义的模式。CNN可以用很激进的裁剪VIT不行。Mixup和CutMix在VIT上的效果特别明显。我做过对比实验不用任何mix策略CIFAR-10上VIT-Base只能到93%左右加上Mixupalpha0.8和CutMixalpha1.0交替使用后能稳定到97%以上。原因还是那个——VIT缺少归纳偏置需要更强的正则化来防止过拟合。5.2 优化器与学习率调度策略VIT原论文用的是AdamWweight_decay0.05配合余弦退火和warmup。这套配置我用了很多次基本可以直接抄optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxepochs)关键点在warmup。VIT对初始学习率很敏感前几个epoch必须从很小的学习率慢慢爬上去。我通常设5到10个epoch的线性warmupbase_lr设为1e-3微调时用1e-4到5e-5。如果跳过warmup直接上1e-3loss大概率会炸。Layer-wise learning rate decay是另一个微调时的技巧浅层用更小的学习率深层用更大的。因为浅层学到的是通用的低级特征微调时不需要大改。timm里有现成的param_groups分层实现用起来很方便。我实测在数据量小于10万张时加上layer decay能把精度提升1到2个百分点。5.3 推理加速与模型导出训练完之后推理阶段的优化空间也很大。VIT的计算瓶颈在自注意力的O(n²)复杂度上。几个实用的加速手段方法原理加速比精度损失减少patch数量用更大的patch_size约4倍1-3%注意力近似用线性注意力替换softmax2-3倍0.5-2%模型量化FP16或INT82倍左右0.1-1%ONNX导出图优化算子融合1.3-1.8倍无我平时最常用的是FP16推理加上ONNX导出组合起来在T4上大概能提升2.5倍吞吐精度基本无损。导出ONNX时要注意把动态维度设好特别是batch_size和序列长度这两个轴不设的话只能跑固定尺寸。6. 踩坑记录与常见问题排查6.1 训练不收敛的排查思路这是问我最多的问题。VIT训练不收敛90%的情况是下面几个原因之一学习率太大。VIT比CNN对学习率敏感得多。如果你用训练ResNet的1e-2去训VIT必炸。从1e-3甚至1e-4开始试。没有warmup。前面说过了这是硬性要求。数据增强不够。VIT在小数据集上极容易过拟合。如果你看到训练loss一直降但验证loss很早就开始涨那就是增强太弱。LayerNorm位置写错。自己实现的时候特别容易犯检查一遍代码。位置编码没加或加错。如果忘了加位置编码模型也能训练但精度会明显偏低而且注意力可视化会很难看。我一般用排除法先在一个极小的子集比如100张图上过拟合如果能到100%精度说明模型结构和代码没问题问题出在训练策略上。如果连100张都过拟合不了那就是代码有bug。6.2 显存不足时的优化策略VIT的显存占用主要来自注意力矩阵。对于224输入、patch16每个样本的注意力矩阵是12×197×197大概1.8MBFP32。batch_size64时就是116MB再加上中间激活值显存很快就满了。几个立竿见影的办法用梯度检查点torch.utils.checkpoint显存降60%左右速度慢20%混合精度训练AMP显存降40%速度还能快一点减小batch_size但用梯度累积保持等效batch冻结浅层只训练后几层和分类头我最常用的是AMP加梯度检查点的组合在单张24G卡上能跑到batch_size128。6.3 从零复现VIT的经验总结最后分享几个我复现过程中总结的经验都是文档里不会写的不要一上来就训ImageNet。先用CIFAR-10或自己手头的小数据集验证代码正确性确认能收敛后再上大规模数据。预训练权重是好朋友。除非你的数据量超过千万级否则从预训练权重微调几乎总是比从零训练好。timm上一行代码就能加载。patch_size的选择要看数据。小分辨率图片比如32x32用patch_size4中等分辨率用8或16高分辨率才用32。patch太大信息损失严重patch太小计算量爆炸。注意timm版本和原论文的差异。qkv_bias、位置编码初始化方式、分类头的dropout率都有细微不同。复现论文精度时务必对齐这些细节我在这上面浪费过至少两天时间。可视化注意力矩阵是调试利器。如果注意力矩阵看起来像噪声没有明显的结构说明模型没学好。正常的注意力应该有清晰的模式比如关注同类物体区域或边缘。分类头之前的LayerNorm不能省。这个我在小模型上试过去掉训练前期看不出问题但后期精度会掉0.5到1个点而且微调时数值不稳定。数据集类别极不平衡时CLS token不如平均池化。这是我做医疗图像项目时发现的CLS token会偏向多数类平均池化更鲁棒。具体原因可能是CLS token在训练中被多数类的梯度主导了。我个人在实际操作中的体会是VIT的代码量其实不大核心逻辑加起来不到300行但它背后涉及的设计决策非常多。每读懂一个细节你对自己的模型在做什么就多一分把握。这种把握在调参和排查问题时特别值钱——你不再是盲目地试而是知道该往哪个方向调。
返回列表