资讯动态

Vision Transformer复现笔记:Pytorch手写VIT全流程与调试指南

发布时间:2026/9/29 7:19:53 来源:尧图企业网站定制
很早之前就想写一篇VIT的复现笔记因为工作中频繁用到Transformer结构做图像特征提取后来翻到自己最早在Pytorch下手写的Vision Transformer代码觉得从“看懂论文图”到“真正让模型跑起来、收敛、出结果”这个过程还是很值得完整记录一遍的。Vision Transformer简称VIT是2020年Google团队提出来的经典结构。它的核心贡献很直白把图像当作“词的序列”来处理用标准的Transformer Encoder做视觉特征提取彻底抛开卷积。当年这篇论文出来的时候很多人第一反应是“这也能行”但实测结果证明在数据量足够大的前提下VIT的精度完全不输给同量级的CNN甚至在大规模预训练后迁移到下游任务时表现更好。这篇博文就围绕VIT做一个从零开始的Pytorch复现包含每行代码的解释、训练配置的选择、以及复现过程中我踩过的坑。如果你已经在用Pytorch做分类任务或者对Transformer在NLP以外领域的应用感兴趣这篇文章基本可以当成一份带注释的源码笔记直接参考。即便你对注意力机制的数学细节还不太熟我也会尽量用生活化的方式把关键环节讲清楚——毕竟动手复现一个模型理解“为什么这么写”比“写出来能跑”重要得多。1. 动手前先想清楚VIT的整体设计与复现思路1.1 图像如何变成“句子”Patch Embedding的核心思想在开始写代码之前先把VIT处理图像的逻辑捋一遍。Transformer最早是为自然语言设计的输入是一串token每个token有明确的语义边界。图像不一样一张224x224x3的图就是一个巨大的数值矩阵哪有“词”的概念VIT的做法很直接把图像切成固定大小的patch比如16x16那么224x224的图会被切成14x14196个patch每个patch拉平后是16x16x3768维的向量。这196个向量就当作196个“词”加上一个用于分类的class token变成197个token序列。为了保留空间位置信息再叠加一个可学习的位置编码向量。这里有个实现上的小技巧论文里说的是Linear Projection of Flattened Patches也就是先切patch再拉平然后过一个全连接层把768维映射到隐藏维度D。但实际写代码时直接用一个大卷积核、stride等于patch size的Conv2d就能一步搞定卷积核大小kernel_sizepatch_sizestridepatch_size输入3通道输出D通道。这样输出特征图的每个空间位置恰好对应原图的一个patch既高效又简洁。class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() 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: [B, 3, H, W] x self.proj(x) # [B, embed_dim, H/patch, W/patch] x x.flatten(2).transpose(1, 2) # [B, num_patches, embed_dim] return x1.2 从CNN到TransformerVIT为什么能抛弃归纳偏置传统CNN有两个很强的先验假设局部性和平移等变性。卷积核只在局部窗口内滑动所以CNN天生认为“相邻像素关系更紧密”。这对自然图像很合理但也限制了模型捕捉全局依赖要靠堆叠很多层才能把感受野扩到全图。VIT放弃了这些假设让每个patch都可以直接和整张图的所有patch做注意力计算一步到位建立全局关系。这种设计带来了两个直接后果第一在数据量足够大时模型能学到更加灵活的特征表达不再受限于局部窗口第二在数据量不足时模型缺少CNN那套内置先验很容易过拟合训练集。这也是为什么VIT原论文反复强调要在JFT-300M这种超大数据集上预训练再迁移到下游任务。我在复现时深刻体会到这个差异拿一个小VIT在CIFAR-10上从零训练同样的训练轮数下精度很难追上同参数量的ResNet但一旦加了强数据增强如Mixup、CutMix、RandAugment差距就迅速缩小。所以在复现时不要一上来就纠结精度而是先把结构和训练流程跑通再逐步调优。1.3 复现VIT的完整技术路线整个复现过程我分成了四个阶段这也是我建议你跟着走的顺序阶段一把VIT的模块拆开实现包括Patch Embedding、Multi-Head Self-Attention、Transformer Encoder Block、分类头。阶段二组装完整模型在前向传播时检查每一层的输出shape确保维度匹配。阶段三在CIFAR-10上用一个小配置的VIT训练验证模型能正常收敛。阶段四调参优化加入数据增强、学习率调度等技巧提升最终精度。选择CIFAR-10而不是ImageNet原因很简单VIT原版是针对224x224甚至384x384的大图设计的在个人电脑或单卡环境下复现ImageNet级别的训练并不现实。CIFAR-10分辨率只有32x32单张卡跑起来非常快适合用来验证代码正确性和做实验对比。当然输入尺寸变了patch size也要跟着调整我在后文会给出具体配置。2. 核心模块拆解VIT每一层到底在做什么2.1 Multi-Head Self-Attention从QKV到注意力权重的完整计算Multi-Head Self-Attention多头自注意力是VIT最核心的组件。它的计算过程可以拆成三步。第一步生成Query、Key、Value。输入是上一层的特征xshape为[B, N, D]其中N是token数量D是隐藏维度。通过三个线性层权重不共享分别得到Q、K、V每个的shape都是[B, N, D]。直观理解Q代表“我在找什么”K代表“我是什么”V代表“我携带的信息”。第二步计算注意力权重。先把Q和K做点积得到[B, N, N]的相似度矩阵每个元素表示第i个token对第j个token的关注程度。为了防止点积值过大导致softmax梯度消失要除以sqrt(d_k)其中d_k是每个头的维度。然后对最后一行做softmax就得到归一化的注意力权重。第三步用注意力权重加权求和V。得到每个token融合全图信息后的新特征。多头的意思是把D维特征切成head份每份单独做上面的注意力计算最后拼回去再过一层线性层。这样每个头可以关注不同的关系有的头关注局部纹理有的头关注全局形状。在实现时为了利用GPU并行能力通常不会真的写成循环而是通过reshape把多个头合并到batch维度里一次算完。class Attention(nn.Module): def __init__(self, dim, num_heads8, qkv_biasTrue, 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) q, k, v qkv.permute(2, 0, 3, 1, 4) 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 x2.2 Transformer Encoder Block残差、LayerNorm、MLP的配合顺序有了多头注意力接下来就是标准的Transformer Encoder Block。VIT使用的Block结构是先LayerNorm再进Attention或MLP然后加残差。这个顺序在论文里叫Pre-LN和原始Transformer的Post-LN正好相反。为什么要用Pre-LN这是我早期复现时特别困惑的地方。查了资料和实验对比后发现Pre-LN的优势在于训练更稳定每层的输入都经过归一化梯度在反向传播时更平滑允许用更大的学习率而且对初始化和warmup的敏感度更低。Post-LN在深层模型中容易出现梯度爆炸需要非常精细的warmup策略。所以VIT、DeiT这些视觉Transformer模型默认都用Pre-LN。MLP部分也很有讲究。它是一个两层的全连接中间用GELU激活函数。第一层把维度从D放大到D*mlp_ratio通常mlp_ratio4也就是把隐藏层扩展到4倍第二层再压缩回D。论文中实验表明MLP的宽度对性能影响显著比深度更值得增加。class Mlp(nn.Module): def __init__(self, in_features, hidden_featuresNone, act_layernn.GELU, drop0.): super().__init__() 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, in_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 class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., qkv_biasTrue, drop0., attn_drop0., act_layernn.GELU, norm_layernn.LayerNorm): super().__init__() self.norm1 norm_layer(dim) self.attn Attention(dim, num_headsnum_heads, qkv_biasqkv_bias, attn_dropattn_drop, proj_dropdrop) self.norm2 norm_layer(dim) mlp_hidden_dim int(dim * mlp_ratio) self.mlp Mlp(in_featuresdim, hidden_featuresmlp_hidden_dim, act_layeract_layer, dropdrop) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x2.3 位置编码与Class Token两个易被忽视却至关重要的细节Visual Transformer里有两个容易被忽略的组件位置编码和Class Token。先讲位置编码。Transformer本身是序列模型但它对token的顺序不敏感——你把“猫在狗旁边”改成“狗在猫旁边”注意力计算的结果在结构上完全一样。为了让模型知道patch的空间位置需要显式加上位置信息。VIT使用的是可学习的一维位置编码也就是随机初始化一个[num_tokens, embed_dim]的矩阵训练时作为参数一起更新。这个矩阵的维度跟序列长度绑定所以如果训练时用的是224x224196个patch 1个class token 197推理时想用384x384的图就需要对位置编码做插值这里很容易踩坑后面我会专门说。再讲Class Token。为什么非得在序列里加一个额外的token用于分类而不是把所有patch的特征做全局平均池化VIT论文里有实验对比直接平均池化效果差一些。一种解释是平均池化会强迫每个patch的特征都往“分类友好”的方向优化限制了特征的多样性而class token相当于一个可学习的“查询向量”它只需要关注最终分类所需的信息其他patch的表示可以保留更丰富的语义。你可以把它理解成队伍里的“队长”——别人负责表达自己队长负责汇总信息做最终决策。2.4 完整模型组装从Patch到分类输出的全过程把上面所有模块拼起来VIT的完整前向流程就是卷积嵌入得到patch序列拼接class token加上位置编码依次经过若干层Transformer Encoder Block最后取出class token对应的特征经过LayerNorm和分类头得到预测概率。我自己在组装的时候习惯把完整模型做成一个类方便统一管理参数和后续扩展。这个骨架代码里包含了前向传播、中间特征提取接口和分类头配置复现其他变体时直接改PatchEmbed和Block参数就行。class VisionTransformer(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12, mlp_ratio4., qkv_biasTrue, drop_rate0., attn_drop_rate0., norm_layernn.LayerNorm): super().__init__() self.num_features self.embed_dim embed_dim 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_rate) self.blocks nn.Sequential(*[ Block(dimembed_dim, num_headsnum_heads, mlp_ratiomlp_ratio, qkv_biasqkv_bias, dropdrop_rate, attn_dropattn_drop_rate, norm_layernorm_layer) for _ in range(depth)]) self.norm norm_layer(embed_dim) self.head nn.Linear(embed_dim, num_classes) if num_classes 0 else nn.Identity() def forward(self, x): x self.patch_embed(x) B x.shape[0] 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) x x[:, 0] x self.head(x) return x组装完之后我强烈建议先用一个小batch的随机张量跑一次前向确认输出维度符合预期。比如输入[2, 3, 32, 32]输出应该是[2, num_classes]。这一步能过滤掉绝大多数维度错误比如transpose写反、reshape参数不匹配等低级问题。3. 完整实操Pytorch环境下从零复现VIT并完成训练3.1 环境配置与数据集准备细节这次复现我用的环境是Pytorch 2.x CUDAPython 3.9。其实VIT的训练对Pytorch版本不算敏感关键是torchvision的数据处理和GPU驱动要正常。如果你是用Anaconda管理环境建议单独建一个新的虚拟环境conda create -n vit python3.9 conda activate vit pip install torch torchvision pytorch-cuda12.1 -c pytorch -c nvidia pip install timm tqdm tensorboard安装timm不是必须的但timm里有大量现成的视觉模型实现包括VIT和它的各种改进版用来对照检查自己的复现是否正确非常方便。我经常用timm里同一个配置的VIT输出做对比如果我们的模型在各层输出的shape一致、数值接近说明实现基本没问题。数据集选用CIFAR-10torchvision里自带下载。注意一点CIFAR-10默认的图片格式是PIL图像需要手动做归一化。这里有一个经验CIFAR-10的均值和方差不要自己算直接用官方推荐值mean[0.485, 0.456, 0.406]std[0.229, 0.224, 0.225]效果就很稳。如果你在别的数据集上训练才需要重新统计均值方差。3.2 针对CIFAR-10的模型配置与参数选择原版VIT是针对224x224输入设计的直接套到32x32的CIFAR-10上会出问题patch16的话32x32只分成2x24个patch整个序列才5个tokenTransformer根本学不动。所以复现时针对小图做适当缩放是必要的。我使用的配置如下输入32x32x3patch_size 4得到(32/4)^264个patch加1个class token共65个tokenembed_dim 256depth 8num_heads 4每个头维度64mlp_ratio 4MLP中间维度为1024分类头输出10类这个配置大约有300万参数在CIFAR-10上从零训练比较现实。如果资源允许可以把embed_dim调到384或depth加到12但训练时间会明显变长对于理解模型结构来说没必要。超参数方面我建议优化器用AdamW而不是普通Adam。AdamW把权重衰减和梯度更新解耦对Transformer这类模型更友好。学习率设为1e-3weight decay设为0.05batch size为128。需要注意的是Transformer对学习率非常敏感lr1e-3通常和上述batch size匹配如果batch size变化明显最好按比例调整学习率。训练轮数我设置了100个epoch并配合warmup cosine退火的学习率调度。前10个epoch做warmup让学习率从1e-4线性升到1e-3之后按照cosine曲线逐渐降到接近0。之所以需要warmup是因为Transformer在训练初期很不稳定如果一开始就用大学习率梯度方向混乱容易在前期就陷入坏的最优解后面很难恢复。3.3 训练脚本的关键代码实现训练部分的代码涉及数据加载、数据增强、训练循环、验证循环四个部分。数据增强我用了RandomCrop和RandomHorizontalFlip这两个基础操作如果你想追求更高的精度可以再上Mixup或CutMix不过要在理解训练逻辑之后再加否则出了问题不好排查。transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) train_dataset torchvision.datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue) model VisionTransformer(img_size32, patch_size4, in_chans3, num_classes10, embed_dim256, depth8, num_heads4, mlp_ratio4., drop_rate0.1, attn_drop_rate0.1) optimizer optim.AdamW(model.parameters(), lr1e-3, weight_decay0.05) total_steps len(train_loader) * 100 warmup_steps len(train_loader) * 10 def lr_lambda(current_step): if current_step warmup_steps: return float(current_step) / float(max(1, warmup_steps)) progress float(current_step - warmup_steps) / float(max(1, total_steps - warmup_steps)) return 0.5 * (1. torch.cos(torch.tensor(progress * 3.1415926))) scheduler LambdaLR(optimizer, lr_lambda) criterion nn.CrossEntropyLoss()训练循环本身不复杂唯一要注意的是model.train()和model.eval()的切换。Transformer里Dropout层不少如果eval模式下忘了切验证精度会被随机噪声干扰看起来忽高忽低第一次复现时很容易被这个坑玩到。训练的GPU显存占用不高batch128在这个小模型上大概占用3-4GB显存。如果显存不够可以把batch size降到64或32学习率也相应调低一些。3.4 训练结果记录与精度分析我用上面的配置完整训练了100个epoch。训练过程大概经历了三个阶段第一个是前期快速上升期1-20个epoch。测试精度快速从40%上升到70%左右这个阶段模型主要学到的是颜色、纹理等低级特征损失下降非常明显。第二个是中期平稳上升期20-60个epoch。精度提升速度明显放缓每个epoch只涨0.1-0.2个点这是注意力机制逐步学到更精细语义特征的过程。第三个是后期收敛期60-100个epoch。学习率经过cosine退火后降得很低损失和精度都趋于平稳最终测试精度稳定在90%左右。需要说明的是90%这个数字并不是VIT的“真实水平”。同样的数据量下一个调好参数的ResNet18轻松能到92%以上这是因为CNN的归纳偏置在小数据集上确实是优势。VIT的优势要在大规模数据预训练后才能体现出来。所以在CIFAR-10上复现VIT不要把目标定在“超过CNN”而应该关注结构是否正确、训练是否稳定、各模块是否按预期工作。4. 复现路上最常踩的坑问题与排查技巧实录4.1 训练不收敛或Loss震荡的排查思路复现VIT后第一次训练时Loss下不去或者大幅震荡是最常见的问题。如果你的Loss一直在1.5-2.0附近徘徊不下降按下面的顺序逐个排查。第一优先检查学习率。Transformer不像CNN那么“皮实”学习率过大时模型会直接发散。我建议从1e-3开始如果训练不稳定就降到5e-4或3e-4观察几个batch的loss变化再决定。反过来如果loss降得很慢也可能是学习率太小。第二检查是否加了warmup。我在没有warmup的情况下测试过精度大约会掉2-3个点而且前期训练波动特别大。建议固定使用warmup通常给总步数的10%到20%就行。第三检查数据预处理。数据没有做Normalize或者Normalize的均值和方差不对都会导致模型训练困难。尤其是自实现数据加载时容易漏掉这一步输出特征数值范围不一致梯度更新就会混乱。第四检查标签是否正确。听起来很基础但类别的索引错位、数据增强导致标签错乱这种问题表现出来的现象就是loss正常下降但验证精度始终很低。可以先拿一小批数据过训练循环手动核对预测结果和标签。4.2 显存不足与批量大小的平衡策略小配置的VIT在CIFAR-10上运行时并不怎么吃显存但如果你的输入是224x224patch16embed_dim768depth12那显存需求会急剧上涨。注意力的复杂度是O(N^2)其中N是token数输入分辨率翻倍意味着token数翻四倍显存和计算量都暴涨。我遇到过显存溢出通常有两个解决办法。第一个是梯度累积也就是把大batch拆成几个小batch分别算出梯度后累加达到预设batch量再更新一次参数。这样能模拟大batch的效果但只节省显存并不减少总计算量。第二个是混合精度训练Pytorch里就是AMP简单几行代码就能让显存占用大约减半还能让训练速度明显提升。不过要注意AMP模式下有些操作会变成FP16对loss scale的设定有讲究使用torch.cuda.amp.GradScaler自动管理就行。还有一个更根本的思路降低patch数量。比如输入224x224时patch从16改到32token数从196降到49显存和计算量直接变成原来的四分之一。当然patch太大丢失细节精度会受影响这个取舍要自己权衡。4.3 复现论文精度对不上时的排查清单这是所有复现者最痛苦的问题明明按论文写了代码训练完成后精度就是比论文报告低好几个点。出现这个问题时先不用怀疑自己的代码有问题因为论文里的实验条件通常很复杂很难百分百复现。我会按下面这个清单排查随机种子是否固定。Pytorch里要同时设置torch.manual_seed、numpy.random.seed如果用了DataLoader的多进程还要设置DataLoader的generator否则每次训练的结果波动会非常大。数据增强策略是否一致。VIT论文在ImageNet上用了RandAugment、Mixup、CutMix等一整套增强策略如果你只是用RandomCrop和RandomFlip精度差3-5个点是非常正常的。预训练模型是否有影响。VIT原论文最强大的版本是在JFT-300M上预训练过的而ImageNet上从零训练的VIT效果并没有那么惊艳。复现时如果不用预训练权重就别拿ImageNet上的精度做对比。优化器参数是否一致。AdamW的weight decay在论文里设了0.1很多人默认用1e-4或1e-5差一个数量级效果差别非常大。训练epoch数和学习率曲线。VIT论文很多实验训练了300个epoch你只跑100个epoch精度自然对不上。如果以上全检查过还是对不上建议用timm的预训练模型做一次对比实验确认自己模型的训练pipeline没问题再排查模型结构的细节。4.4 位置编码插值与分辨率适配问题在真实项目里经常遇到训练和推理分辨率不一致的情况。比如训练用224x224部署或测试时想用384x384的图。因为位置编码是学习出来的维度是[1, num_patches1, embed_dim]输入图片分辨率变了patch数量就变了位置编码的shape就对不上了。最简单的办法是插值。把原来的位置编码从二维空间上按比例放大到目标尺寸。这里有一个细节常被忽略需要先把class token对应的位置编码单独拆出来只对patch部分的位置编码做插值处理完再和class token的位置编码拼接回去。如果直接把整个位置编码一起插值class token的位置信息也会被插值污染分类的性能会受影响。def interpolate_pos_embed(pos_embed_ckpt, new_num_patches, embed_dim): # pos_embed_ckpt: [1, old_num_patches1, embed_dim] cls_pos pos_embed_ckpt[:, :1, :] patch_pos pos_embed_ckpt[:, 1:, :] old_h old_w int((patch_pos.shape[1]) ** 0.5) new_h new_w int(new_num_patches ** 0.5) patch_pos patch_pos.reshape(1, old_h, old_w, embed_dim).permute(0, 3, 1, 2) patch_pos F.interpolate(patch_pos, size(new_h, new_w), modebicubic, align_cornersFalse) patch_pos patch_pos.permute(0, 2, 3, 1).reshape(1, new_num_patches, embed_dim) return torch.cat((cls_pos, patch_pos), dim1)4.5 小数据集下过拟合的应对策略在CIFAR-10上训练VIT30个epoch左右就可能出现过拟合训练loss持续下降验证loss开始上升验证精度停滞不涨。这验证了我前面说的“Transformer没有归纳偏置”的问题。应对方案从最有效的开始排列第一是数据增强Mixup、CutMix、RandAugment都可以用随机性和强度要够这是提升VIT小数据性能最直接的手段。第二是增加Dropout和DropPathVIT中drop_rate和attn_drop_rate通常设为0.1如果过拟合严重可以提高到0.2甚至0.3。第三是使用更大的权重衰减AdamW的weight decay可以调到0.1这本身也是Transformer训练的标准配置。第四是模型裁剪降低depth或embed_dim让模型容量匹配数据规模。我自己的测试显示从0增强 drop0.0的配置到最后加上RandAugment Mixup drop0.2的配置在CIFAR-10上精度差距可以超过10个点。这也说明VIT这类模型对训练策略的依赖程度远高于CNN。5. 从复现到理解常见问题速查与后续扩展思路5.1 一个简洁的VIT问题速查表问题可能原因解决方案Loss不下降学习率过大或过小调整学习率建议从1e-3开始用warmup训练震荡严重缺少warmup加10%-20%步数的线性warmup验证精度低数据预处理不完整检查Normalize、数据增强、标签对齐显存溢出token数过多或batch过大梯度累积、混合精度、增大patch推理分辨率变化报错位置编码维度不匹配插值位置编码注意class token单独处理过拟合数据集过小或模型过大增强数据、增大dropout、增大weight decay精度和论文对不上训练条件变化核对随机种子、增强策略、epoch数、lr曲线5.2 复现完成后值得尝试的扩展方向把VIT复现跑通之后可以顺着下面几个方向继续深入每次改一个点体会会很深。方向一改成DeiT结构。DeiT的核心改动是在训练时加了一个蒸馏token和对应的蒸馏损失利用CNN教师网络来指导学生网络训练。这个改动让Transformer在ImageNet上不再依赖JFT那种超大数据集只靠ImageNet-1k就能训练到接近CNN的效果。复现DeiT会接触到更多损失函数设计和训练技巧。方向二改成Swin Transformer。Swin把VIT的全局注意力改成了窗口注意力先在小窗口里做自注意力再通过Shift操作让窗口之间产生信息交互。这样就恢复了多尺度特征计算复杂度也从O(N^2)降到了O(N)很适合做检测、分割等密集预测任务。方向三尝试MAE自监督预训练。MAE的思想很简单把图像的一部分patch遮住让Transformer重建被遮住的像素。自监督预训练优势非常明显尤其在数据量不足的场景下。用MAE做预训练再微调分类在中等规模数据集上往往比直接有监督训练效果更好。5.3 关于本次复现我最后想说的几句话从搭框架到跑通训练再到调优和排坑整个VIT复现过程花了我不少时间但收获确实很大。以前看VIT论文注意力图、位置编码这些概念都停留在抽象层面真到自己写代码后才明白每个设计背后的权衡。一个比较深的体会是复现模型不要只追求“跑起来”。跑起来只是第一步关键是把模型每个模块的输入输出shape、梯度流动方向、不同超参数的影响都搞清楚。比如你试试把Pre-LN改成Post-LN看看训练会不会崩把位置编码删掉看精度掉多少把attention的scale去掉试试有没有NaN和震荡这种对比实验做多了对模型的理解就逐步深入了。如果让我给一个起步建议先用小配置在CIFAR-10或类似的小数据集上把整个流程跑通确认实现无误后再换大图、加深模型。也不用着急加载预训练权重先自己从零训练几次感受一下模型在不同学习率和增强策略下的表现差异等到对整个训练动态有感觉了再去玩预训练和迁移学习会顺利很多。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑