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

资讯详情

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

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

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_size=patch_size, stride=patch_size)。这两种方式在数学上是完全等价的,因为卷积核的权重就是 Linear 权重重排后的形式。实际写代码时我建议用卷积实现,原因有两点:

  • 卷积操作底层高度优化,前向速度更快,显存占用也更可控。
  • 写法更紧凑,不需要先手动切块再 reshape,少掉一堆形状转换的代码。

1.3 为什么需要 class token 和 position embedding

NLP 里的 BERT 会加一个[CLS]token,ViT 照搬了这个思路。在 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_size=224, patch_size=16, in_chans=3, embed_dim=768): 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_size=patch_size, stride=patch_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_size=16, stride=16) ↓ 中间: [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_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, dropout=0.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(p=dropout) 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], dim=1) # [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_heads=8, qkv_bias=True, attn_drop=0.0, proj_drop=0.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, 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 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 = 64,64 ** -0.5 = 0.125。

第四步,softmax 归一化后乘 V,再 reshape 回原来的形状,最后过一个输出投影层。

2.4 MLP Block:通道维度的信息融合

class Mlp(nn.Module): def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.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 x

MLP 里的第一个 Linear 把维度从 768 放大到 768 * 4 = 3072,第二个 Linear 再把维度压回 768。这个"先放大再压缩"的设计是被实验验证过的:两层全连接加非线性激活能让每个 token 在不同特征维度之间充分融合信息。激活函数用 GELU 而不是 ReLU,GELU 对负值的处理更平滑,论文实验显示在 Transformer 里 GELU 普遍比 ReLU 效果好一截。

2.5 Encoder Block:把以上组件拼起来

class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio=4.0, drop=0.0, attn_drop=0.0): super().__init__() self.norm1 = nn.LayerNorm(dim, eps=1e-6) self.attn = Attention(dim, num_heads=num_heads, attn_drop=attn_drop, proj_drop=drop) self.norm2 = nn.LayerNorm(dim, eps=1e-6) self.mlp = Mlp(in_features=dim, hidden_features=int(dim * mlp_ratio), drop=drop) 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 的eps=1e-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_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0, drop=0.1, attn_drop=0.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(p=drop) self.blocks = nn.Sequential( *[ Block(dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, drop=drop, attn_drop=attn_drop) for _ in range(depth) ] ) self.norm = nn.LayerNorm(embed_dim, eps=1e-6) self.head = nn.Linear(embed_dim, num_classes) self.init_weights() def init_weights(self): # 用截断正态分布初始化权重 nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.cls_token, std=0.02) self.apply(self._init_weights) 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) 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], dim=1) 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(dim=1),对所有 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 Embedding:Conv2d 权重,输入 3 通道输出 768 通道,卷积核 16x16。参数量为 768 * 3 * 16 * 16 = 589,824。
  • class token + position embedding:197 * 768 + 768 = 152,064。
  • Encoder Block 的 Attention:qkv 权重是 768 * (768*3) = 1,769,472,输出投影是 768 * 768 = 589,824,加上 bias 共约 2.36M。
  • Encoder Block 的 MLP:fc1 是 768 * 3072 ≈ 2.36M,fc2 是 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 个 Block:7.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 到 15
  • Mixup:alpha 参数 0.8
  • CutMix:alpha 参数 1.0
  • Random 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 decay=0.05
  • 学习率:初始 0.001,配合 warmup 和 cosine anneal。warmup epoch 通常 5 到 10 个
  • Batch Size:建议至少 1024 起步,如果显存受限可以降到 512,效果会略有损失
  • 混合精度 AMP:fp16 训练可以省一半显存,但 LayerNorm 的 eps 建议保持 1e-6

5.3 预训练权重加载的注意事项

如果是从官方预训练权重做微调,需要注意位置编码维度的匹配。ViT-B/16 预训练权重里的pos_embed形状是[1, 197, 768],你要是把输入分辨率从 224 改成 384,patch 数量会从 196 变成 576,pos_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 size(16→32)、开启梯度累积
加载预训练权重报 key 不匹配修改了 num_classes用 strict=False 加载,只加载 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,引入了知识蒸馏 token(distillation 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_size=8)过一遍前向,打印每一层的 shape,确认没有 mismatch 之后再替换成正式数据。这一步看似麻烦,能帮你省下大量在训练一个 epoch 之后才发现 shape 错乱的抓狂时间。

返回列表