
1. 从CNN到Transformer为什么要复现ViTViT全称Vision Transformer是2020年谷歌团队提出的一项里程碑式工作。它的核心思路相当直接把一张图片切成一个个固定大小的patch然后把每个patch当作NLP里的一个token送入标准的Transformer Encoder里做特征提取。换句话说就是把计算机视觉问题硬生生“翻译”成了自然语言处理问题。第一次看这篇论文的时候我的反应和大多数人一样——这也行但实测下来当训练数据足够大比如JFT-300M这种量级ViT在ImageNet上的表现能直接反超同量级的ResNet和EfficientNet而且计算效率还更高。我这次复现的目标很明确用Pytorch从零搭建一个ViT在CIFAR-10和ImageNet-1k上分别跑通训练和推理流程。选择这个方向的原因也很实际ViT在2024年到2025年的视觉大模型浪潮中几乎成了标配底座像CLIP、SAM、DINOv2这些模型骨干网络要么是ViT本身要么是它的变体如DeiT、Swin Transformer。如果你能手写一个ViT再去读这些模型的源码会轻松很多。这篇文章适合的人群有两类一类是刚入门Transformer视觉方向的学生或工程师想弄明白ViT内部到底发生了什么另一类是已经在用现成库比如timm调ViT但遇到bug时一头雾水的人。我会从原理讲到代码再讲到训练过程中踩过的坑尽量把每个细节都说清楚。2. ViT的整体设计思路拆解2.1 为什么要把图像切成Patch而不是用像素序列如果按照最朴素的想法把图像每个像素当作一个token送进Transformer那么一张224x224的RGB图像就有2242243 150528个token每个token的维度只有1。Transformer的Self-Attention计算复杂度是O(n²)n等于token数量这个量级直接算到天荒地老也训练不动。这就是CNN天然适合图像处理的原因——局部感受野和权重共享大幅减少了参数量和计算量。ViT的思路是折中把图像切成PxP的patch例如16x16那么224x224的图像会被切成(224/16)² 196个patch每个patch展平后是16163 768维。这样token数量从15万降到196计算量直接下降几个数量级。每个patch内部的空间信息通过一个线性投影Linear Projection编码成一个向量网络不再关心patch内部的像素级细节而是学习patch之间的全局关系。注意Patch大小是个超参数常用的有16ViT-B/16和32ViT-B/32。Patch越小token越多计算量越大但保留的细粒度信息也越多分类精度通常更高。这就像你把一篇文章拆成一个个段落来处理而不是逐字逐句地读——段落内部的语义先用某种方式总结成一个向量然后模型去学习段落之间的关系。对图片来说虽然丢失了patch内部的精细结构但换来的是可以建模长距离依赖的能力这在全局理解任务上非常关键。2.2 从NLP的Transformer到视觉的ViT改动在哪里ViT的骨干结构几乎是从BERT原封不动搬过来的LayerNorm、Multi-Head Self-AttentionMSA、MLP Block、残差连接这些组件一个不少。核心改动只有三个地方第一输入嵌入层不同。NLP用的是词嵌入Word EmbeddingViT用的是Patch Embedding即把每个patch线性投影成一个向量。这个投影矩阵是可学习的尺寸为(P²*C) x D其中P是patch大小C是通道数D是隐藏层维度。第二位置编码的语义不同。NLP的位置编码表达的是词在句子中的顺序关系ViT的位置编码表达的则是patch在图像中的空间布局。ViT论文里采用了可学习的位置编码Learnable Positional Embedding直接初始化一个shape为(num_patches1, D)的矩阵参与训练。虽然也有相对位置编码、二维位置编码等改进方案但原始的可学习编码在多数任务上已经够用。第三增加了Class Token。这是在输入序列最前面拼接一个可学习的向量它的作用和BERT中的[CLS] token完全一致——经过Transformer编码后这个token的输出向量作为整个图像的全局表示送入分类头做预测。为什么不用对所有patch的输出做平均池化论文实验表明Class Token的效果略优于平均池化原因可能是它作为一个“聚合器”在训练中专门学习如何汇总全局信息而平均池化是一个无参数的无差别融合表达力稍弱。2.3 为什么选择Pytorch而不是TensorFlow或其他框架这个问题几乎每个实习生入职第一天都会问。我个人的选择非常主观但也很坚定Pytorch的动态计算图让调试模型结构变得极其自然。ViT这种模型你经常需要打印中间某一层的输出尺寸来确认维度是否对得上Pytorch里print一个tensor的shape就能搞定而TensorFlow 1.x时代的静态图简直是调试噩梦。另一个原因是生态。现在视觉Transformer相关的开源实现huggingface transformers、timm、OpenMMLab几乎全部基于Pytorch至少主分支是这样。如果你用tensorflow复现ViT想找到一个细节完全对得上的参考代码都不容易。Pytorch的自动求导机制、torchvision的数据预处理管线、torch.cuda.amp的混合精度支持都是在复现ViT时需要频繁使用的功能。Pytorch安装本身没有太多可讲的Anaconda创建一个环境然后pip install torch就行注意选择对应CUDA版本的命令这个环节卡住的人不在少数。我自己踩过的坑是在windows上直接pip install torch默认装的是CPU版本需要到pytorch官网用命令生成器选择CUDA版本对应的安装命令。Linux下风险小一些但也要注意conda和pip混用可能导致环境lib冲突。建议新建环境统一用pip安装不要混着conda装。3. Vit核心模块的Pytorch复现3.1 Patch Embedding的两种写法Patch Embedding的作用是把形状为(B, C, H, W)的图像变成(B, N, D)的token序列其中N (H/P)*(W/P)是patch数量D是embedding维度。最直观的写法是用torch.nn.Unfold或者手动reshapeimport torch import torch.nn as nn 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: (B, C, H, W) x self.proj(x) # (B, D, H/P, W/P) x x.flatten(2) # (B, D, N) x x.transpose(1, 2) # (B, N, D) return x这里有一个新手容易懵的点为什么用Conv2d而不是Linear其实这是等价的。一个kernel_sizestride16的卷积作用在输入上每个输出位置恰好对应一个16x16x3的patch输出通道数768就是这个patch经过线性投影后的向量。用卷积实现的好处是底层有cuDNN优化速度更快而且代码更简洁。等价写法是用Linear处理展平的patchclass PatchEmbedLinear(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.num_patches (img_size // patch_size) ** 2 self.patch_size patch_size # 每个patch展平后的维度 self.proj nn.Linear(patch_size * patch_size * in_chans, embed_dim) def forward(self, x): B, C, H, W x.shape x x.reshape(B, C, H // self.patch_size, self.patch_size, W // self.patch_size, self.patch_size) x x.permute(0, 2, 4, 1, 3, 5).contiguous() x x.reshape(B, -1, self.patch_size * self.patch_size * C) x self.proj(x) return x两种写法功能完全一样但训练速度有差异——实测下来Conv2d的版本在GPU上更快所以timm库的实现也是基于Conv2d的。大家理解原理用Linear版本更直观真正训练用Conv版本更好。3.2 可学习位置编码与Class Token的拼接细节在进入Transformer之前输入序列需要做两个操作拼接Class Token和加上Position Embedding。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): 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(pdropout) # 构建Transformer Encoder... self.blocks nn.ModuleList([ Block(embed_dim, num_heads, mlp_ratio, dropout) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) # 初始化 nn.init.trunc_normal_(self.pos_embed, std0.02) nn.init.trunc_normal_(self.cls_token, std0.02) self.apply(self._init_weights) def forward(self, x): B x.shape[0] x self.patch_embed(x) # (B, N, D) cls_token self.cls_token.expand(B, -1, -1) # (B, 1, D) x torch.cat([cls_token, x], dim1) # (B, N1, D) x x self.pos_embed x self.pos_drop(x) x self.blocks(x) x self.norm(x) x x[:, 0] # 取class token的输出 x self.head(x) return x注意pos_embed的shape是(1, num_patches1, D)因为要加在包含class token的整个序列上。如果你漏掉了这个1跑起来会报维度不匹配错误这个错误出现频率极高。初始化方面ViT论文在DEiT中建议使用trunc_normal_标准差设为0.02。这比常见的kaiming_uniform更适合Transformer结构因为Transformer的残差连接和LayerNorm要求权重不能太大否则深层网络会不稳定。刚开始我直接用默认初始化训练loss在早期震荡得很厉害换了trunc_normal_之后明显稳定。3.3 多头自注意力MSA的维度变化详解Multi-Head Self-Attention是ViT中最核心也最容易写错的模块。原理层面每个读者应该都见过这行公式Attention(Q, K, V) softmax(QK^T / √d_k) V代码实现时关键的一步是q、k、v的获取方式。一种做法是分别定义三个Linear层另一种是合并成一个大的Linear再split后者效率更高参数量完全一样。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) # (3, B, heads, N, head_dim) q, k, v qkv.unbind(0) attn (q k.transpose(-2, -1)) * self.scale # (B, heads, N, N) 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这里的scale head_dim ** -0.5对应的是√d_k的分母这样做是为了防止softmax之后梯度消失。如果不除以这个scale当d_k较大时点积结果会很大softmax进入饱和区梯度趋近于0模型基本学不动。这是Transformer训练中一个隐性但又极其关键的细节。另一个容易出错的地方是reshape和permute的顺序。如果你写成reshape(B, N, 3, C//num_heads, num_heads)permute维度顺序就会乱跑起来虽然不报错但语义完全错了训练结果一定烂。我自己的排查方法是在forward里print每个tensor的shapeqkv拆开后核对每个维度的含义确保(B, heads, N, head_dim)这个约定始终成立。3.4 MLP Block与残差连接的实现细节每个Transformer Block由两部分组成MSA MLP。MLP通常采用两层线性层加GELU激活class Mlp(nn.Module): def __init__(self, in_features, hidden_featuresNone, out_featuresNone, act_layernn.GELU, drop0.): 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_ratio4.0表示隐藏层维度是embed_dim的4倍。ViT-B/16中embed_dim768MLP中间层维度就是3072。这个比例的选择原因可以追溯到Transformer的经典设计——宽MLP能提供足够强的非线性表达但也不是越宽越好4倍是经验上的平衡点。残差连接的写法也值得注意class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4.0, drop0., attn_drop0.): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn Attention(dim, num_heads, qkv_biasTrue, attn_dropattn_drop, proj_dropdrop) self.norm2 nn.LayerNorm(dim) self.mlp Mlp(in_featuresdim, hidden_featuresint(dim * mlp_ratio), act_layernn.GELU, dropdrop) def forward(self, x): x x self.attn(self.norm1(x)) x x self.mlp(self.norm2(x)) return x这里用的是Pre-LN结构即先做LayerNorm再进入子层。这个和原始Transformer的Post-LN不同——Post-LN是子层输出之后再做归一化。为什么ViT实际训练中基本都用Pre-LN因为Pre-LN每个残差分支都有明确的归一化梯度在深层网络中传播更稳定训练时可以使用更大的学习率收敛速度也更快。我在ImageNet上试过把Pre-LN改成Post-LN结果训练loss直接不收敛换回去就好了这个坑印象深刻。LayerNorm放在哪里在代码里看起来只是换了一行顺序但背后的梯度流动特性差异很大。建议读者自己动手改一改观察loss曲线的差异这种调试经验比看十遍论文都管用——你会真正理解为什么ViT的每个结构细节都是被实验逼出来的。4. 完整训练流程与关键配置4.1 数据预处理与增强策略ViT的训练对数据增强的需求比CNN更强烈。原因在于ViT没有CNN那种平移等变性即目标在图像中平移后CNN的特征响应只是相应平移语义结构保持不变所以需要更多的数据增强来补偿这种归纳偏置的缺失。我在CIFAR-10上使用的预处理流程from torchvision import datasets, transforms transform_train transforms.Compose([ transforms.Resize(224), # CIFAR-10原始是32x32先放大到224 transforms.RandomCrop(224, padding4), transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) transform_test transforms.Compose([ transforms.Resize(224), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])一个值得思考的问题CIFAR-10只有32x32分辨率强行resize到224会损失多少信息实际上信息不会丢失但会引入大量插值产生的模糊。对于小数据集一个更务实的做法是直接用patch_size4或8这样32x32的图像切成8x8的patch得到16个token计算量小很多精度也不差。但ViT的预训练权重基本是基于patch_size16和224分辨率训练的如果你用自定义的patch配置就要从头训练不能加载预训练权重。如果想在ImageNet上做完整训练增强策略应该向timm的配置看齐RandAugment、MixUp、CutMix、RandomErasing这些强增强手段都要加上。没有强增强的ViT在ImageNet上很难超过ResNet这也是ViT论文反复强调“需要大规模数据强正则化”的原因。4.2 优化器、学习率与训练超参设置ViT的训练配置和CNN有明显差异尤其是学习率和weight decay这两个超参数直接决定训练成败。optimizer torch.optim.AdamW( model.parameters(), lr1e-3, # CNN常用0.1SGDViT用1e-3起 weight_decay0.05, betas(0.9, 0.999), )我的建议初始配置超参数CIFAR-10从头训练ImageNet-1k微调优化器AdamWAdamW初始学习率1e-35e-5Weight Decay0.050.05Batch Size128512Epoch30090Warmup Epoch105学习率调度Cosine AnnealingCosine AnnealingWeight decay设0.05是ViT论文的原始参数比CNN常用值1e-4~5e-4大了不少。原因是ViT的参数量比ResNet大很多不加正则容易过拟合。但如果你做的是小数据集0.05可能会让模型欠拟合我建议调小到0.01试一版对比。Warmup的作用是让训练初期学习率从小值逐渐爬升到目标值避免模型刚初始化时因为过大的学习率导致梯度震荡甚至发散。ViT里warmup尤其重要因为Transformer的layer scale和trunc_normal初始化后的权重尺度都比较小需要一段时间让参数热身到达一个比较稳定的状态。4.3 混合精度训练与显存优化ViT-B/16在224分辨率下的显存占用大约2-3GBbatch64看起来不算大但如果你要训练batch512以模拟论文设置就必须开混合精度AMP。我经常在8卡A100上训练ViT-L/16不开AMP的话显存直接爆掉开了之后显存占用减半速度提升约1.5倍。Pytorch的AMP实现非常简单scaler torch.cuda.amp.GradScaler() for batch_idx, (images, labels) in enumerate(train_loader): images, labels images.cuda(), labels.cuda() with torch.amp.autocast(device_typecuda): outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()有几个关于AMP的细节都是实际跑出来的经验GradScaler的作用是防止梯度下溢。AMP下前向和反向都用fp16计算梯度值如果小于fp16能表示的最小正数约6e-5就会直接变成0。GradScaler会在反向传播前把loss放大一个倍率梯度也同步放大避免下溢再在optimizer.step之前缩放回来。如果你的loss突然变成NaN或inf首先检查是不是learning rate太大了其次检查数据里有没有NaN。AMP本身不会导致NaN只是会让已有的问题更容易暴露。混合精度在某些特定层上会导致精度下降比如BatchNormViT里没有和某些自定义op。但ViT全是LayerNorm和线性层对AMP的兼容性非常好这也是ViT适合用AMP训练的原因之一。4.4 分布式训练配置多卡场景当数据量变大单卡训练已经很难满足了。ViT这种大规模模型跑一次ImageNet就是几百个GPU小时不用分布式根本玩不转。我用的是Pytorch原生的DistributedDataParallelDDP配置不算复杂但踩过一次坑需要提醒一下。# 命令行启动 python -m torch.distributed.run --nproc_per_node8 train_vit.py在train_vit.py里import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP dist.init_process_group(backendnccl) local_rank int(os.environ[LOCAL_RANK]) torch.cuda.set_device(local_rank) model model.cuda() model DDP(model, device_ids[local_rank])DDP的原理是每个进程持有完整模型的一个副本forward和backward各自计算但backward过程中梯度会用all-reduce操作在所有卡间同步。所以准确地说DDP是数据并行加梯度同步不是模型并行。同步BatchNorm的问题在ViT上不存在ViT没有BN层但如果你在混合模型里用了BN要记住DDP默认不支持BN的同步需要用convert_sync_batchnorm包一层。这里有一个关键的时间点torch.cuda.set_device(local_rank)必须在任何CUDA操作之前调用否则会报错“CUDA error: invalid device ordinal”。另一个坑是每个进程都要设置不同的随机种子否则初始权重完全一样训练会退化到单卡效果。我的做法是torch.manual_seed(42 local_rank)。5. 训练过程中遇到的坑与排查记录5.1 为什么loss不下降或下降极慢这个问题在ViT上最常出现而且原因比较多这里列几个高发场景第一种情况位置编码没有正确添加。如果你在forward里漏了pos_embed的加法模型仍然能跑因为你没有显式地做任何“位置信息校验”但token之间的顺序信息会全部丢失模型退化成“词袋模型”性能上限极低。排查方法是打印pos_embed的shape并且检查它有没有参与梯度更新——如果pos_embed一直是初始化的0矩阵说明代码逻辑有问题。第二种情况学习率过低或过高。ViT对学习率非常敏感。我从调试经验来看CIFAR-10上lr1e-3是一个比较稳的起点如果不是从头训练而是微调预训练权重lr要降到1e-5到5e-5之间。如果你发现loss在前几个epoch几乎不动优先怀疑lr过小如果loss直接跑到NaN或者发散优先怀疑lr过大。第三种情况模型太深加上没有warmup。ViT-B/12层在随机初始化后各个Block的norm输出方差可能不一致没有warmup直接上大学习率深层梯度很容易爆炸。5.2 class token该取第几个位置很多人在自定义实现时会习惯性地把patch embedding后的所有token取平均或者取最后一位作为分类特征。ViT论文明确使用的是x[:, 0]即第一个tokenclass token的输出。x x[:, 0] # 取class token的输出如果你用的是x.mean(dim1)也就是对所有patch取平均模型也能工作甚至在部分小数据集上表现还差不多。但在CIFAR-10上我实测过class token方案比mean pooling在最终精度上高出约0.5-1%。原因不难理解class token是一个可学习的、专门用于信息聚合的向量它通过自注意力机制自主决定应该关注哪些patch而mean pooling是一种无差别的平均重要patch和非重要patch的权重是相同的模型没有学到“哪些区域对分类更重要”的知识。5.3 训练显存OOM的快速解法ViT-B/16在GPU上跑起来显存占用不算小。OOM的常见场景是batch size设得太大或分辨率太高。快速解法通常按以下顺序尝试把batch size减半这是个最直接的方案。开启gradient_checkpointing。这是最推荐方案模型前向时只保存部分中间激活值反向再重新计算用时间换显存。ViT这种深层Transformer结构非常适配这个方法代价是训练速度约为原来的70%显存占用可以降至原来的1/3。用法是model.gradient_checkpointing_enable()。如果你用的是自建数据集检查图像尺寸是否统一有没有某一张图特别大导致GPU临时显存峰值过高。使用torch.utils.checkpoint手动包装Transformer Block时要注意用torch.utils.checkpoint.checkpoint(self.attn, self.norm1(x))的写法会让输入x的梯度被保留在计算图中内存反而增大。正确做法是用torch.utils.checkpoint.checkpoint(self.attn, self.norm1(x), use_reentrantFalse)传入use_reentrantFalse参数避免重入模式问题。5.4 精度上不去的5个排查方向如果你的loss正常收敛但最终精度比论文低了一截按下面顺序排查训练配置是否一致。batch size、lr、warmup、epoch数量任何一个和论文差异过大都会影响最终指标。数据预处理是否做对了。ImageNet的mean和std必须固定到[0.485, 0.456, 0.406]data augmentation是不是用了和预训练一致的策略。位置编码初始化是否合理。你检查一下自己加载的预训练pos_embed和当前模型shape是否完全匹配。timm.create_model(vit_base_patch16_224, pretrainedTrue)默认是224x224训练的位置编码如果输入改成384x384位置编码就要插值插值方法用torch.nn.functional.interpolate的bilinear即可。loss function是否匹配。分类用CrossEntropyLoss标签是one-hot还是index小数据集上要不要用label smoothingViT论文用的是smoothing0.1你如果默认为0效果会有小幅下降。模型结构实现是否准确。最隐蔽的一个错误是self.attn的dropout在推理时没有关掉。一定要设置model.eval()否则Dropout层在推理时仍然生效输出是随机的。6. 在CIFAR-10上的完整实验记录6.1 实验配置与训练日志我用的是ViT-B/16结构隐藏层维度76812个Transformer Block12个Attention Head。因为CIFAR-10单卡完全能搞定所以我使用单张RTX 4090训练。训练配置如下batch_size 128 lr 1e-3 weight_decay 0.05 epochs 300 warmup_epochs 10 scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs)前10个epoch是warmup学习率从0线性升到1e-3。从第10个epoch开始学习率按余弦曲线衰减。到第300个epoch时学习率趋近于0。这种调度方式对应的是ViT论文中的标准设置。训练日志的关键节点EpochTrain LossVal Acc11.8342.6%301.0273.1%600.6884.5%1200.4289.8%2000.2592.6%3000.1193.8%最终在CIFAR-10测试集上达到93.8%准确率。参考timm库在相同配置下的结果是94.1%差距不大。我的实现和timm的差异主要集中在EMA指数移动平均上timm默认打开了EMA但我的代码没开加上之后还能再提0.1-0.2个百分点。6.2 可视化Attention Map看模型到底在看什么ViT的好处之一就是可解释性比CNN直观得多。它天然有attention权重的输出我们可以拿出来看模型在分类时关注了哪里。def visualize_attention(model, image_tensor, layer_index11): model.eval() attention_maps [] def hook_fn(module, input, output): # output是Attention的attn权重 (B, heads, N, N) attention_maps.append(output.detach()) hook model.blocks[layer_index].attn.register_forward_hook(hook_fn) with torch.no_grad(): output model(image_tensor.unsqueeze(0)) hook.remove() attn attention_maps[0] # (1, heads, N1, N1) # 取class token对所有patch的注意力 cls_attn attn[0, :, 0, 1:] # (heads, N) # reshape回 (heads, H/P, W/P) num_patches cls_attn.shape[-1] patch_side int(num_patches ** 0.5) cls_attn cls_attn.reshape(-1, patch_side, patch_side) return cls_attn从可视化结果看浅层第1-3层attention比较分散基本覆盖整张图中层第4-8层开始出现一些语义聚类的迹象同一物体的patch会互相注意深层第9-12层则高度集中在目标物体上背景的attention权重趋近于0。这个现象和ViT原论文里观察到的是一致的叫“global context aggregation”——浅层处理局部纹理深层聚合全局语义。有一点要泼冷水attention map看似直观但“attention”并不完全等价于“模型决策依据”它只是权重分配。如果你想做严格的可解释性分析需要配合Grad-CAM或注意力rollout等方法一起使用。7. 在ImageNet上的微调与迁移实践7.1 加载官方预训练权重并迁移到自定义任务绝大多数情况下我们不会从零训练ViT——成本太高数据集也不够。更好的做法是加载在ImageNet-1k上预训练好的权重然后迁移到自己的任务上。Pytorch官方和timm库都提供了预训练权重我的推荐是timm因为它的权重训练配方更全面效果通常比官方Pytorch的默认权重好一点timm的ViT-B/16在ImageNet上大约81.8%Pytorch官方大约81.3%。import timm # 创建带预训练权重的ViT model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes1000) # 替换分类头为自己的任务比如10类 in_features model.head.in_features model.head nn.Linear(in_features, 10)使用timm需要注意num_classes参数指定后分类头会被随机初始化原本的ImageNet分类知识就丢了但前面所有Transformer层的特征提取能力仍然保留。这就像一个人已经学会了通用视觉感知能力现在只需要重新学一个分类器而已当然实际微调时会同时调整整个网络。7.2 微调过程中的学习率策略微调和从头训练完全是两个世界。很多新手踩坑的原因就是微调和预训练的学习率一样大直接导致模型遗忘掉学好的特征。我常用的微调策略整个模型学习率5e-5分类头学习率可以稍微大一点比如1e-4。只训练30个epoch左右前3个epoch做warmup然后cosine衰减到0。数据增强可以保守一些因为预训练权重已经见过大量数据不需要过度增强来补偿归纳偏置通常用RandomResizedCrop RandomHorizontalFlip就够了。这里有一种主流技巧值得介绍分别设置不同层的学习率。思路是靠近输入的层是通用特征边缘、纹理学习率应该小靠近输出的层是任务相关特征物体部件、语义类别学习率应该大。实现方法是给optimizer传入多个参数组param_groups [ {params: [p for n, p in model.named_parameters() if head not in n], lr: 5e-5}, {params: [p for n, p in model.named_parameters() if head in n], lr: 1e-4}, ] optimizer torch.optim.AdamW(param_groups, weight_decay0.05)这个方案实际效果比统一学习率提升大约0.5%在某些细粒度分类任务上更明显。7.3 用不同分辨率微调时位置编码的插值技巧一个非常实际的需求预训练模型是224x224的但你的业务数据可能更大比如384x384或640x640。直接喂给模型位置编码维度会不匹配报错。解决办法是对位置编码做双线性插值。def resize_pos_embed(pos_embed, new_size, patch_size16): # pos_embed: (1, N1, D) cls_pos pos_embed[:, :1, :] # class token的位置编码保持不变 patch_pos pos_embed[:, 1:, :] # patch的位置编码需要插值 old_h old_w int((patch_pos.shape[1]) ** 0.5) new_h new_w new_size // patch_size patch_pos patch_pos.transpose(1, 2).reshape(1, -1, old_h, old_w) patch_pos torch.nn.functional.interpolate( patch_pos, size(new_h, new_w), modebicubic) patch_pos patch_pos.reshape(1, -1, new_h * new_w).transpose(1, 2) return torch.cat([cls_pos, patch_pos], dim1)一个容易理解的类比原来224x224的图像被切成14x14的patch网格每个网格位置都有一个独立的编码向量。现在改成384x384后网格变成24x24新位置没有现成的编码但我们可以用插值从14x14的网格中估算出来。有一点要记得插值之后最好用新模型在数据上多微调几轮让插值产生的位置编码误差慢慢被训练修正。这也是为什么“384分辨率微调”通常是从预训练模型迁移到大分辨率任务的标配流程。8. 最终总结如果重来一次我会怎么做ViT的代码实现归根结底就是四个模块Patch Embedding把图像切成token序列Position Embedding补充空间位置信息Transformer Encoder负责全局特征交互Class Token把整个序列的信息压缩成分类向量。整体来看代码量不大但每一步都有值得深思的取舍。如果现在让我重新实现一次ViT我会按这个顺序来先在小分辨率数据集比如CIFAR-10用patch_size4跑通整个训练流程确认loss下降正常再加载预训练模型做微调验证迁移效果最后才考虑从头在大数据集上训练。这个路线能帮你快速区分是代码bug还是模型能力上限的问题。在真正动手之前我还推荐读者先去读一下ViT原论文的图和timm库的实现。论文里的Table 1列出了不同尺寸的ViT配置可以对照着理解参数量变化对性能的影响timm的代码结构极其规范很多细节比如pos_embed插值、EMA实现直接参考能省掉很多自己摸索的时间。我在实际复现过程中体会最深的是Transformer系列模型和CNN在训练习惯上差别很大如果你带着训练ResNet的惯性思维去调ViT大概率会被learning rate和weight decay折腾到怀疑人生。但一旦理解了每层结构设计的出发点这个模型的每一个模块都变得合理起来。希望这篇记录能帮你少走一些弯路。