ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ViT代码逐行拆解:从Patch Embedding到Transformer Encoder

2026/9/30 10:35:26 拓冰建站 浏览量
ViT代码逐行拆解:从Patch Embedding到Transformer Encoder 最近不少朋友在后台问我ViT 的代码到底该怎么看。说实话Transformer 类的代码初次接触确实有点绕尤其是把图像切成 patch 再送进 encoder 的过程光看论文里的公式图很容易懵。这篇我就直接拿一份可运行的 PyTorch ViT 实现一行一行拆开讲把每个张量的 shape 变化、每个模块在做什么、对应原论文哪张图一次说清楚。这份代码是我在实际项目中调过的版本不是那种只跑通就行的小 demo。你把它吃透了后面再去看 Swin、DeiT、MAE 这些变体会发现核心模块基本都是同一套思路在打转。需要的基础知识也不用太深懂基本的 PyTorch 张量操作了解自注意力的大致概念就能跟下来。1. 整体设计思路为什么 ViT 要把图像切成 patch 再送进 Transformer在动手写代码之前得先把 ViT 的结构思路捋清楚否则代码看完了还是一团浆糊。1.1 从 CNN 到 Transformer 的思路切换传统 CNN 处理图像靠的是卷积核滑动。卷积核天然有局部归纳偏置——相邻像素之间的关联性强远处像素关联弱。这个先验让 CNN 在小数据集上很占便宜因为它不需要学太多东西就能把局部结构提取出来。ViT 的思路是完全反过来的。它认为局部归纳偏置是可以不要的只要数据量够大模型自己能从全局里学出结构与关联。所以 ViT 把一张图切成一堆小方块patch每个 patch 拉平成一个向量然后像处理 NLP 里的 token 一样把这一堆向量送进标准的 Transformer encoder。这个设计对代码的影响是决定性的图像处理任务被打包成了序列建模任务。所以整个 ViT 的主干代码其实就是在写一个 Transformer encoder只有输入端的 patch embedding 和输出端的分类头是新增的。1.2 为什么用 Linear 而不是卷积来做 Patch Embedding原始 ViT 论文里Patch Embedding 的实现方式有两种理解一种是把图像切块后拉平过一个 Linear 层另一种是直接用一个 stride 等于 patch size 的卷积层。Dosovitskiy 等人在代码里用的是Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size)。这两种方式在数学上是完全等价的因为卷积核的权重就是 Linear 权重重排后的形式。实际写代码时我建议用卷积实现原因有两点卷积操作底层高度优化前向速度更快显存占用也更可控。写法更紧凑不需要先手动切块再 reshape少掉一堆形状转换的代码。1.3 为什么需要 class token 和 position embeddingNLP 里的 BERT 会加一个[CLS]tokenViT 照搬了这个思路。在 patch embedding 之后输入序列的最前面再拼接一个可学习的向量。这个向量经过 encoder 之后它对应的输出位置就当作整张图的全局特征接一个分类头即可。这个设计的历史原因是 Transformer 本身不具备序列聚合能力——输出序列长度和输入一致如果不用 class token就得对所有 token 做平均池化效果略差一些。position embedding 则是给模型注入位置信息。Transformer 本身是置换等变的你把它输入的顺序打乱输出也只是跟着换位置语义不会变。但图像的空间结构极其重要所以必须显式地把位置信息编码进去。ViT 用的是可学习的 1D position embedding直接加到 patch embedding 结果上。2. 核心模块代码解析从 Patch Embedding 到 Encoder Block现在开始看代码。我会按照数据流动的顺序逐个模块拆解。每个模块都会给出完整代码、输入输出 shape 变化、以及对应的原论文图解位置。2.1 Patch Embedding把图像切成 token 序列class PatchEmbed(nn.Module): 将 [B, C, H, W] 的图像转换为 [B, num_patches, embed_dim] 的 token 序列。 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): B, C, H, W x.shape x self.proj(x) # [B, embed_dim, H/patch, W/patch] x x.flatten(2) # [B, embed_dim, num_patches] x x.transpose(1, 2) # [B, num_patches, embed_dim] return x这一步做的事情用图来表示就是输入: [B, 3, 224, 224] ↓ Conv2d(3, 768, kernel_size16, stride16) ↓ 中间: [B, 768, 14, 14] ↓ flatten(2) → [B, 768, 196] ↓ transpose(1,2) → [B, 196, 768]这里有两个比较容易踩坑的点我第一次写的时候都栽过。第一个是flatten(2)的语义它表示从第 2 维开始展平所以结果是一个[B, 768, 196]的形状而不是[B, 196, 768]。第二个是 transpose 之后必须用contiguous()吗在 PyTorch 里 transpose 只改变视图不改变内存布局如果后面马上接 view 或 reshape 就容易报错。我在这个模块里没有显式调用contiguous()但后面如果发现shape对的上却报内存不连续的错加一行x x.contiguous()就能解决。不过一般Linear层内部会处理非连续张量所以实际上不调用也能正常工作。2.2 拼接 class token 并加上位置编码class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0, dropout0.1): # ... 其他初始化 ... self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(pdropout) def forward(self, x): B x.shape[0] x self.patch_embed(x) # [B, 196, 768] cls_tokens self.cls_token.expand(B, -1, -1) # [B, 1, 768] x torch.cat([cls_tokens, x], dim1) # [B, 197, 768] x x self.pos_embed # [B, 197, 768] x self.pos_drop(x) # ... 后续送入 encoder ...注意看self.cls_token初始化为全零。这是个非常实用的小技巧一开始让 class token 不携带任何信息训练时模型会自己学到它该干什么。如果随机初始化一个很大的值可能会导致训练初期梯度爆炸。torch.nn.Parameter是必须的这样才能让 PyTorch 把这俩张量当作可训练参数。位置编码的形状是[1, 197, 768]注意这里不是 196 而是 197因为 class token 在最前面也占了一个位置。加位置编码用的是广播机制pos_embed的 batch 维度是 1会自动扩展到 B。这里有个好习惯pos_embed设计成可学习的参数时官方预训练权重能直接加载进来自适应插值方便迁移到不同分辨率的数据集。2.3 Attention 模块ViT 的灵魂所在class Attention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasTrue, attn_drop0.0, proj_drop0.0): 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.unbind(0) 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这一段是很多人看着最头疼的地方我拆开揉碎讲。第一步self.qkv是一个 Linear 层把输入从C维映射到3C维一次性把 Q、K、V 都算出来。这样设计是为了效率矩阵乘法一次搞定比三个 Linear 分开算省一次大矩阵乘法的时间。然后 reshape 成[B, N, 3, num_heads, head_dim]的布局。第二步permute(2, 0, 3, 1, 4)把第 2 维3即 Q/K/V 的区分维度挪到最前面得到[3, B, num_heads, N, head_dim]的张量。这里为什么要这么 permutation因为下一步unbind(0)可以很方便地解出 q、k、v每个的形状是[B, num_heads, N, head_dim]。第三步attn (q k.transpose(-2, -1)) * self.scale。q k.transpose(-2, -1)计算每个位置对所有位置的注意力分数得到[B, num_heads, N, N]的矩阵。N x N矩阵是 Transformer 的计算瓶颈所在序列长度 N 增加时这里的时间复杂度是 O(N²)。乘以scale head_dim ** -0.5是缩放注意力。为什么要缩放当维度 head_dim 变大时点积结果的方差会变大把数值逼近 softmax 的饱和区梯度会消失。用1/sqrt(head_dim)缩放后点积方差维持在 1 左右。补充一下head_dim 也就是 768 / 12 6464 ** -0.5 0.125。第四步softmax 归一化后乘 V再 reshape 回原来的形状最后过一个输出投影层。2.4 MLP Block通道维度的信息融合class Mlp(nn.Module): def __init__(self, in_features, hidden_featuresNone, out_featuresNone, act_layernn.GELU, drop0.0): super().__init__() 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) def forward(self, x): x self.fc1(x) x self.act(x) x self.drop(x) x self.fc2(x) x self.drop(x) return xMLP 里的第一个 Linear 把维度从 768 放大到 768 * 4 3072第二个 Linear 再把维度压回 768。这个先放大再压缩的设计是被实验验证过的两层全连接加非线性激活能让每个 token 在不同特征维度之间充分融合信息。激活函数用 GELU 而不是 ReLUGELU 对负值的处理更平滑论文实验显示在 Transformer 里 GELU 普遍比 ReLU 效果好一截。2.5 Encoder Block把以上组件拼起来class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, drop0.0, attn_drop0.0): super().__init__() self.norm1 nn.LayerNorm(dim, eps1e-6) self.attn Attention(dim, num_headsnum_heads, attn_dropattn_drop, proj_dropdrop) self.norm2 nn.LayerNorm(dim, eps1e-6) self.mlp Mlp(in_featuresdim, hidden_featuresint(dim * mlp_ratio), dropdrop) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x这里有个值得说道的设计细节Pre-LN 结构。也就是先做 LayerNorm再进 Attention/MLP而不是后做。原始 Transformer 论文用的是 Post-LN但 ViT 官方代码用的是 Pre-LN。实践经验表明Pre-LN 训练更稳定梯度传播更顺可以省掉一些 warmup 的麻烦。残差连接x self.attn(...)是必须的如果去掉深层网络的梯度根本无法有效回传。LayerNorm 的eps1e-6也是个容易被忽略的细节。eps 太小在 FP16 混合精度训练时可能出 NaN太大又会影响归一化效果。实测下来1e-6是官方验证过的稳妥值不要乱改。3. 完整模型组装与参数量分析3.1 完整 ViT 模型代码import torch import torch.nn as nn class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4.0, drop0.1, attn_drop0.0): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(pdrop) self.blocks nn.Sequential( *[ Block(dimembed_dim, num_headsnum_heads, mlp_ratiomlp_ratio, dropdrop, attn_dropattn_drop) for _ in range(depth) ] ) self.norm nn.LayerNorm(embed_dim, eps1e-6) self.head nn.Linear(embed_dim, num_classes) self.init_weights() def init_weights(self): # 用截断正态分布初始化权重 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) self.apply(self._init_weights) 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) def forward(self, x): B x.shape[0] x self.patch_embed(x) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) x x self.pos_embed x self.pos_drop(x) x self.blocks(x) x self.norm(x) # 只取 class token 对应的输出 x x[:, 0] x self.head(x) return x注意最后一层x[:, 0]取的是第一个位置class token 所在位置的输出。这一步对应论文 Fig 1 最上方那个MLP Head的输入来源。如果你不用 class token另一种做法是x.mean(dim1)对所有 token 做平均池化DeiT 论文里比较过class token 略优。3.2 前向传播的 shape 变化全览拿一张 224x224 的 RGB 图batch 设为 2过一遍模型各阶段张量形状如下阶段操作输出 shape输入-[2, 3, 224, 224]Patch EmbeddingConv2d flatten transpose[2, 196, 768]拼接 cls tokentorch.cat[2, 197, 768]加位置编码广播加法[2, 197, 768]Encoder Block x12Attention MLP 残差[2, 197, 768]最终 LayerNormnorm[2, 197, 768]提取 class tokenx[:, 0][2, 768]分类头Linear(768, 1000)[2, 1000]打眼一看整个过程中 768 维这个数字始终没变过这也是 Transformer 架构的一个特点特征维度在整个 encoder 中保持恒定变化只发生在序列长度和通道数内部临时扩展的部分。3.3 参数量计算ViT-Base 到底有多少参数我们来手动算一下 ViT-Base/16 的参数量加深对各个模块规模的感知。Patch EmbeddingConv2d 权重输入 3 通道输出 768 通道卷积核 16x16。参数量为 768 * 3 * 16 * 16 589,824。class token position embedding197 * 768 768 152,064。Encoder Block 的 Attentionqkv 权重是 768 * (768*3) 1,769,472输出投影是 768 * 768 589,824加上 bias 共约 2.36M。Encoder Block 的 MLPfc1 是 768 * 3072 ≈ 2.36Mfc2 是 3072 * 768 ≈ 2.36M加上 bias 约 4.72M。每个 Block 里的两个 LayerNorm每个 768 * 2 1,536共 3,072。单个 Block 总参数量约 2.36M 4.72M 0.003M 7.08M。12 个 Block7.08M * 12 85M。分类头768 * 1000 1000 769,000。总和约为 589,824 152,064 85M 769,000 ≈ 86.5M。这个数字和 ViT-Base 官方公布的大约 86M 参数吻合。大部分计算量集中在多头注意力和 MLP 两大块这也是后面优化时最先考虑剪枝或量化的位置。4. 图解 ViT 结构对照代码看官方 Figure 1很多解读文章会直接把 ViT 论文里的结构图贴出来这里我用文字方式把图和代码位置对照一下。ViT 的结构自下而上可以分成四层。4.1 输入端Linear Projection of Flattened Patches官方图最左侧是一个 224 x 224 的图像画成 14 x 14 的网格每个格子 16 x 16 像素。这些格子就是 patch。图中每个 patch 四条不同颜色的线代表被拉平后过了一个 Linear Projection 层。对应代码就是PatchEmbed。Conv2d在这里干的就是切块 线性投影两步。如果你想验证这一步等价性可以手动把 patch 拉平后过nn.Linear(768, 768)结果和卷积是完全一样的。4.2 编码器前位置嵌入与 class token 拼接图中有个[class]的方块插在所有 patch token 前面。这一部分对应代码第 29 行的torch.cat。图中还有一串 Position Embedding 的图标对应self.pos_embed和后面的加法操作。官方图里 Position Embedding 画成一组不同颜色的短线暗示它是可学习的参数。4.3 Transformer Encoder 内部标准 Block 的堆叠图的中间部分是 L x 的矩形框代表堆叠的 Transformer Encoder block。框内从上到下依次是 Multi-Head Attention、Add Norm、MLP、Add Norm。对应代码就是Block里的norm1-attn- 残差加 -norm2-mlp- 残差加。整个框是 L 份叠起来对应nn.Sequential里 12 个 Block。4.4 输出端分类头图的最上方是 MLP Head接的输入是编码器输出的第一个 token 对应的向量即x[:, 0]。这个向量被当作整张图的全局语义表示。注意在预训练阶段MLP Head 通常是先接一个更大维度的隐藏层比如 3072再接分类层但在大多数开源实现里直接就是单层 Linear效果差异不大。5. 训练细节与超参数代码之外的关键光是把前向代码跑通不算完ViT 的训练配置同样重要。下面这些参数都是我在实际训练中验证过或者从官方案例里整理出来的。5.1 数据增强和正则化是 ViT 的命根子ViT 和 CNN 一个本质区别是CNN 有很强的归纳偏置所以数据量小也能硬训ViT 归纳偏置弱全靠数据量或数据增强来补。实际经验是在 ImageNet 上训练 ViT 至少要 90 到 100 个 epoch 才能收敛如果只有 30 个 epoch效果会被同规模 CNN 碾压。常用的增强配置RandAugment随机增强强度 9 到 15Mixupalpha 参数 0.8CutMixalpha 参数 1.0Random Erasing擦除概率 0.25随机裁剪 水平翻转常规操作另外一个容易被忽略的是 Stochastic Depth随机深度。也就是在训练时随机跳过某些 Block 的输出等价于给深层 Transformer 添加正则化。ViT 在 ImageNet 上训练drop path rate 一般设置 0.1如果从头在自建数据集上训练设置 0.2 到 0.3 可能更稳。5.2 优化器与学习率调度ViT 官方采用的优化器是 AdamW这点是出乎很多人意料的因为炼丹师们默认 CV 任务就该用 SGD。但 ViT 的论文和 DeiT 都验证了 AdamW 配上正确学习率收敛速度比 SGD 快得多。常见配置优化器AdamW参数 betas(0.9, 0.999)weight decay0.05学习率初始 0.001配合 warmup 和 cosine anneal。warmup epoch 通常 5 到 10 个Batch Size建议至少 1024 起步如果显存受限可以降到 512效果会略有损失混合精度 AMPfp16 训练可以省一半显存但 LayerNorm 的 eps 建议保持 1e-65.3 预训练权重加载的注意事项如果是从官方预训练权重做微调需要注意位置编码维度的匹配。ViT-B/16 预训练权重里的pos_embed形状是[1, 197, 768]你要是把输入分辨率从 224 改成 384patch 数量会从 196 变成 576pos_embed就对不上了。解决办法是插值。把pos_embed从[1, 197, 768]处理成[1, 769, 768]这种新尺寸需要把普通的 1D 插值换成 2D 插值具体做法是先把位置编码从 197 里拆出 class token剩下的 196 重排成 14x14 的网格用双线性插值缩放到新分辨率对应的网格大小再拼回 class token。PyTorch 官方 timm 库里的resize_pos_embed函数就是这么实现的。直接对整条序列做插值会损失空间结构信息效果明显变差。6. 常见问题与调试实录这个部分把我在跑 ViT 代码时遇到过的、以及身边朋友常踩的坑整理成一张速查表后面附几个典型场景的排查思路。现象可能原因解决方法训练 Loss 不下降学习率过大或过小初始 lr 设为 3e-4 到 1e-3用 warmup 过渡验证集精度远低于训练集过拟合严重增大 drop path、添加更强的数据增强、增大 weight decay显存 OOMpatch 数过多或 batch 过大减小 batch、换大 patch size16→32、开启梯度累积加载预训练权重报 key 不匹配修改了 num_classes用 strictFalse 加载只加载 backbone 部分在自建数据集上效果很差数据量太小用预训练权重微调不要从头训练前向测试时 shape 对不上cls_token 拼接维度错误检查expand(B, -1, -1)是否用对了 batch 维度6.1 案例position embedding 插值之后的精度损失有一次我把 ViT-B/16 从 224 分辨率迁移到 384直接做双线性插值后微调发现下游任务在验证集上不如直接用 384 从头训练数据量充足的模型差距约 1.5%。后来按上面说的方式把 197 拆成1 196重排成 14x14 网格再做插值精度就追回来了。这说明位置编码的空间结构对 ViT 来说真的不是摆设处理不好位置信息的迁移模型能力会打折。6.2 案例attention 输出出现 NaN还有一个经典的坑是混合精度训练时attention 的 softmax 输入中出现 NaN。排查下来发现是q k.transpose(-2, -1)这一步在 FP16 下累加溢出了。解决办法是torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction False或者在 attention 内部把计算临时转回 FP32。ViT 的 attention 部分因为存在大量矩阵乘法在高精度需求上比 CNN 更容易出问题建议始终保留 FP32 前向FP16 只开在卷积和 token 混合那部分实测更稳。6.3 案例自建小数据集上 ViT 完全打不过 ResNet一个朋友拿 2 万张图片做分类ViT 从头训练精度被 ResNet50 吊打。这其实不是代码问题是模型和数据规模的匹配问题。ViT 在数据量少时缺少归纳偏置学不出来。有两个解法一是用 DeiT 的蒸馏策略拿 ResNet 当 teacher让 ViT 学 teacher 的软标签很大程度上弥补归纳偏置缺失的问题二是用预训练权重做微调把 ImageNet 上学到的视觉结构迁移过来效果会好很多。注意如果数据集只有几千张不要直接从头训 ViT。老老实实用预训练模型微调或者直接换 CNN 系模型别跟数据量过不去。7. 后续扩展从 ViT 到更多 Transformer 视觉模型理解了这份 ViT 代码你对其他 Transformer 视觉模型的适应速度会快很多。DeiT数据高效的 ViT引入了知识蒸馏 tokendistillation token代码结构几乎一模一样只是多了一个 token 教学机制。Swin Transformer把全局注意力改成窗口注意力核心思路变了但 Attention 模块 QKV 的打法完全一致你可以照着 ViT 的 Attention 去对比会发现很多相同影子。MAE自监督训练范式encoder 就是 ViT只不过 decoder 更轻量。你只要掌握了 ViT 的前向流程MAE 的 mask 策略就只是中间加了一层操作。ViTDet把 ViT 当检测骨干网络处理的是多尺度特征关键接法来自 ViT 输出的各层特征融合。我的建议是把 ViT 代码吃透之后去官方 timm 库读一遍它的 ViT 实现。timm 里对 ViT 做了很多细节优化比如forward_features和forward_head的分离设计、dynamic_img_size等特性读完之后迁移到新任务会顺手很多。最后分享一个我从实战里养成的习惯拿到一个新视觉 Transformer 模型的代码我会先手动把一个小输入比如 1x3x32x32配合 patch_size8过一遍前向打印每一层的 shape确认没有 mismatch 之后再替换成正式数据。这一步看似麻烦能帮你省下大量在训练一个 epoch 之后才发现 shape 错乱的抓狂时间。