ARTICLE DETAIL

资讯详情

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

ViT与Swin Transformer:从原理到代码的图像识别新范式

ViT与Swin Transformer:从原理到代码的图像识别新范式 几年前我在做图像分类项目时还在一门心思调ResNet的深度和宽度那时候如果有人告诉我NLP那套Transformer架构会“反攻”CV甚至成为主流我大概率是不信的。但后来的事情大家都知道了ViTVision Transformer凭借一套纯注意力机制在ImageNet上杀进了SOTA行列紧接着Swin Transformer又用一套分层窗口的设计把Transformer的通用性和效率一起拉高连目标检测、分割这类密集预测任务也没放过。这篇博文就好好梳理一下Transformer在CV里的进化脉络从原理到代码从不足到突破尽量把每个关键决策背后的“为什么”讲透文末附上可以直接改来用的PyTorch代码。1. 从NLP到CVTransformer凭什么跨界1.1 NLP的成功给CV带来了什么启示2017年《Attention Is All You Need》提出Transformer时大家的注意力都在机器翻译上。它的核心卖点简单粗暴不依赖循环结构直接用自注意力Self-Attention建模序列中任意两个位置的关系。这意味着长距离依赖不再是难题而且整个计算可以高度并行化训练效率远超当时的LSTM、GRU。这个设计在NLP领域迅速开花结果从BERT到GPT系列预训练大模型成为了标配。而CV这边很长一段时间仍然是CNN的天下。卷积核天然带有局部性和平移等变性这让CNN在图像任务上非常高效也符合图像的某些固有属性。但CNN也有天花板一个很尴尬的问题是卷积核的感受野是有限的虽然可以通过堆叠层数来扩大但本质上是“由局部逐步组成全局”对全局关系的建模不够直接。Transformer的出现给了CV研究者一个全新的视角既然句子里的词可以通过注意力机制两两交互图像为什么不行图像也可以被看成一组像素或一组patch的序列patch与patch之间的关系类比词与词之间的关系。于是用Transformer做图像分类就成了一个自然的想法。1.2 ViT出现前的那些过渡尝试在ViT之前其实已经有不少人尝试把注意力机制引入CV比较典型的是SENet、Non-local Network。SENet对通道维做注意力Non-local Network则是对空间位置做全局建模。这些工作有一个共同点它们不是“取代”CNN而是给CNN“加装”注意力模块起到锦上添花的作用。这种思路在当年很务实因为CNN的归纳偏置非常强直接在完整图像上跑自注意力计算量会爆炸。Non-local虽然能建模长距离依赖但计算复杂度是O(N^2)N为空间位置数量比如输入是56x56特征图N3136这个复杂度根本不敢往深了堆。所以ViT真正的贡献不在于“第一个把注意力用在图像上”而在于它给出了一个极简且有效的方案把图像切成patch用标准Transformer Encoder来处理完全不需要卷积。这个“减法”做得非常彻底也正因为足够简洁才让后续大量基于Transformer的CV模型得以快速发展。2. ViT精讲把图像切成“单词”2.1 Patch Embedding是怎么把图像变成序列的ViT的输入处理非常直接。假设输入图像是224x224x3设定patch size为16x16那么图像会被切分成(224/16)^2 196个patch。每个patch展平后是一个16x16x3768维的向量这个维度刚好可以和Transformer的隐层维度对齐。但这里有个细节值得注意直接把展平的像素向量送进Transformer等于完全抛弃了像素之间的局部空间结构。作者的做法是加一个可训练的线性投影层Linear Projection把每个patch的768维向量映射到D维的embedding空间这个过程和NLP里词嵌入Word Embedding本质上是一回事。代码实现简洁到让人意外import torch.nn as nn 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, 224, 224) x self.proj(x) # (B, embed_dim, 14, 14) x x.flatten(2) # (B, embed_dim, 196) x x.transpose(1, 2) # (B, 196, embed_dim) return x看到这个Conv2d了吗用卷积来实现patch切分是一个很巧妙的工程技巧卷积核大小等于patch size步长也等于patch size这样输出特征图的每个位置就对应一个patch的embedding一次前向就完成了切块和线性映射两件事。2.2 Class Token与Position Embedding的细节标准Transformer的输入是序列输出也是序列。对于分类任务ViT借鉴了BERT的做法在序列开头加入一个特殊的可学习class token这个token不来自任何patch它最终对应的输出向量经过分类头后就是整个图像的类别预测。选择class token而不是简单的全局平均池化作者在论文里做了实验对比发现两者的效果非常接近。但class token有一个潜在优势在预训练阶段模型可以适应不同数量的输入patchclass token充当了“全局信息汇总者”的角色这对后面做迁移学习和微调更友好。Position Embedding是另一个容易被忽略但极其重要的细节。由于Transformer本身是置换不变的如果不加位置信息196个patch的顺序就没有任何意义模型会把图像完全打乱看待。ViT使用的是可学习的1D位置编码形状是(1961, 768)加在patch embedding和class token上一起送入Transformer Encoder。这里有个实操中的细节如果预训练时用的是224x224分辨率微调时换成384x384patch数量会从196变成576预训练的位置编码就装不下了。常见做法是对位置编码做二维插值bilinear interpolationViT作者在论文里验证了这个做法有效但我在实际项目中发现插值后最好还是让模型在新分辨率上再训练一小段时间否则位置信息会有一段适应期。3. Swin Transformer精讲分层与窗口带来的改变3.1 ViT的短板和Swin的解题思路ViT在ImageNet这种千万级数据集上确实能发挥出很强的性能但训练数据不足时它的效果往往不如ResNet或EfficientNet。原因也很直接CNN的归纳偏置局部性、平移等变是内置的而ViT把这些先验全抛掉了所有模式都必须从数据里学。数据量跟不上模型就“学不动”。另外一个让人头疼的问题是多尺度。CNN天然有金字塔结构浅层特征分辨率高、语义弱深层特征分辨率低、语义强这非常契合目标检测和分割这类任务。而ViT在所有层都保持全局建模、固定分辨率比如14x14虽然语义信息强但缺少多尺度特征直接拿来做检测和分割并不顺手。Swin Transformer的核心设计就冲着这两个问题去的一是引入分层结构Hierarchical让深层特征分辨率逐层降低像CNN那样形成特征金字塔二是引入窗口注意力Window-based Attention把自注意力限制在局部窗口内大幅降低计算复杂度。3.2 窗口注意力、移位窗口和相对位置编码窗口注意力的思想非常朴素图像被均匀划分成若干个不重叠的小窗口比如7x7的patch区域每个窗口内部独立做自注意力。这样一来计算复杂度从ViT的O(N^2)降成了O(N/M)其中N是patch总数M是窗口内的patch数。以224x224输入、patch size 4为例Swin第一层的窗口注意力计算量比全局注意力少了一个数量级这是它能堆出深层模型的重要原因。但窗口注意力有一个天然缺陷窗口与窗口之间没有信息交流某个位置的token只能“看到”自己窗口内的内容感受野被锁死了。Swin的解法非常妙它引入了一个移位窗口机制Shifted Window当前层用规则窗口下一层就把窗口整体偏移比如向右下各偏移M/2个patch这样上一层在不同窗口里的token在下一层就有机会进入同一个窗口间接实现了跨窗口信息流动。交替使用规则窗口和移位窗口的设计还有一个额外的工程优势每个窗口内的patch数量保持不变计算效率高而且大部分子窗口可以并行处理。再说相对位置编码。Swin里没有用ViT那种大尺寸绝对位置编码而是给每对像素位置生成一个相对位置偏置Relative Position Bias加到注意力分数上。这样做的好处有两个一是参数量少二是有更强的平移不变性。代码实践上这个步骤通常通过一个可学习的偏置表来实现PyTorch里可以用nn.Parameter预生成表然后用索引取出来加在注意力矩阵上面很多开源实现里都有写成bias表的做法。3.3 分层架构和特征金字塔的复现Swin的分层结构跟ResNet的stage概念非常像。整体分为4个阶段Stage每个Stage内部的Transformer Block数量不同常见配置如下Swin-T每个Stage的Block数为[2, 2, 6, 2]通道数为[96, 192, 384, 768]Swin-SBlock数为[2, 2, 18, 2]通道数同Swin-TSwin-BBlock数为[2, 2, 18, 2]通道数为[128, 256, 512, 1024]在每个Stage之间有一个Patch Merging层作用类似于CNN里的下采样。以第一个Stage到第二个Stage为例会把2x2邻域的4个patch在通道维度上拼接起来得到一个4C维的向量再通过一个线性层压缩到2C维。这样一来序列长度减少了4倍通道数增加了2倍正好呼应了CNN中“空间分辨率减半、通道数翻倍”的设计范式。这种结构带来的直接好处是Swin天然可以输出多尺度的特征图。比如输入224x224经过4个Stage后可以得到56x56、28x28、14x14、7x7四种分辨率的特征这正好可以直接喂给FPN、PAN等检测分割头部使用。所以Swin在COCO目标检测任务上不需要像ViT那样额外设计复杂的特征提取逻辑直接用分层特征就能取得很好的效果。4. ViT与Swin一张表看懂核心差异很多同学问过我项目里到底选ViT还是Swin这个问题没有标准答案但可以借助对比表来判断对比维度ViTSwin Transformer核心结构标准Transformer Encoder全局自注意力分层Transformer窗口自注意力序列长度固定不变如196个patch逐阶段减半计算复杂度随分辨率平方增长随分辨率线性增长归纳偏置几乎没有需大量数据引入局部性和平移不变的先验多尺度特征不支持/需要改造天然支持适合检测分割适合场景大数据集分类、预训练、多模态中小数据量分类、检测、分割、密集预测典型精度ImageNet上需要JFT-300M预训练才能碾压CNN纯ImageNet-1K预训练即可达到高精度这里想多说一句很多人看到ViT在ImageNet-21K上预训练后表现很好就觉得ViT“就是更强”但实际上这更像是一个“数据规模和模型容量”匹配的问题。如果你手里只有一两万张图片直接上ViT很可能会过拟合换Swin-T或者干脆用卷积模型会更稳妥。另外需要注意Swin虽然计算量更小但因为窗口划分、移位和Mask等操作的实现复杂度高实际训练时往往会引入一些额外开销。在一些较小的数据集上Swin的训练速度和显存占用未必比ViT有绝对优势这个我在后面的实战部分会再细说。5. 实战用timm库快速实现并微调ViT与Swin5.1 环境准备与数据组织实战部分我直接基于PyTorch和timm来演示timm这个库把ViT、Swin以及大量Transformer变体都封装成了统一接口非常适合快速验证想法。我用一个花卉分类的小数据集来演示迁移学习数据集包含5个类别每个类别约200张图片训练集占80%验证集占20%。先安装依赖pip install torch torchvision timm opencv-python scikit-learn tqdm数据加载部分直接用torchvision的ImageFolder配合标准的预处理流程。这里有个值得注意的点ViT和Swin在预处理上基本一致都要求输入归一化到ImageNet的mean/std且训练时通常用RandomResizedCrop和RandomHorizontalFlip做增强验证时用CenterCrop。import torch from torchvision import transforms, datasets train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) train_dataset datasets.ImageFolder(data/flower_5/train, transformtrain_transform) val_dataset datasets.ImageFolder(data/flower_5/val, transformval_transform) train_loader torch.utils.data.DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader torch.utils.data.DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4)5.2 一键切换ViT与Swin模型timm里加载模型非常简单两行代码的事import timm # 加载ViT-B/16在ImageNet-21k上预训练过 # vit_model timm.create_model(vit_base_patch16_224_in21k, pretrainedTrue, num_classes5) # 加载Swin-T在ImageNet-1k上预训练 swin_model timm.create_model(swin_tiny_patch4_window7_224, pretrainedTrue, num_classes5)这里有两个细节提醒一下第一timm的模型命名非常讲究比如swin_tiny_patch4_window7_224这串名字的意思是tiny配置、patch size为4、窗口大小为7x7、输入分辨率224。建议用之前先跑一下timm.list_models(*vit*)或timm.list_models(*swin*)看看有哪些可用变体避免名字记错。第二如果你用的是在ImageNet-21k上预训练的模型原始分类头是21841类迁移到自己的5类数据集时timm会自动替换分类头这一点不用操心。但要注意如果切换的模型使用了不同的输入分辨率比如vit_base_patch16_384那么你的预处理也要跟着改成384x384否则性能会打折扣。5.3 训练流程与损失函数选择训练逻辑跟传统CNN训练几乎一样这也正是Transformer模型在工程落地上最方便的地方换模型结构不影响训练流程。我习惯用AdamW优化器配上一个简单的余弦退火学习率调度器这在Transformer类模型上是公认比较稳定的组合。from timm.scheduler import CosineLRScheduler from timm.optim import create_optimizer_v2 optimizer create_optimizer_v2(swin_model, optadamw, lr1e-4, weight_decay0.05) scheduler CosineLRScheduler(optimizer, t_initial30, warmup_t3, warmup_lr_init1e-6) criterion torch.nn.CrossEntropyLoss() for epoch in range(30): swin_model.train() for images, labels in train_loader: images, labels images.to(cuda), labels.to(cuda) optimizer.zero_grad() outputs swin_model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() scheduler.step(epoch) # 验证代码省略练习时可以把模型切换成ViT用完全相同的训练配置跑一轮你会直观感受到两者的收敛速度和最终精度差异。5.4 手写一个精简版ViT核心模块timm虽然好用但如果你想深入理解ViT的运行机制手写一遍核心模块会非常有帮助。下面这个精简版实现把patch embedding、class token、position embedding和Transformer Encoder都串起来了去掉了dropout和LayerNorm的一些细节保留主干逻辑。import torch import torch.nn as nn class VitForClassification(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes1000, embed_dim768, depth12, num_heads12): super().__init__() self.patch_size patch_size self.num_patches (img_size // patch_size) ** 2 # patch embedding 展平 self.patch_embed nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) # class token 和位置编码 self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, self.num_patches 1, embed_dim)) # Transformer Encoder encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforwardembed_dim * 4, dropout0.1, activationgelu, batch_firstTrue) self.encoder nn.TransformerEncoder(encoder_layer, num_layersdepth) 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) def forward(self, x): B x.size(0) x self.patch_embed(x).flatten(2).transpose(1, 2) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) x x self.pos_embed x self.encoder(x) x self.norm(x[:, 0]) return self.head(x)把这段代码跟timm里的ViT对比你会发现大体框架是一致的只是timm实现里做了更多工程细节上的优化比如用2D位置编码插值、Stochastic Depth、LayerScale等。6. 训练Transformer模型时踩过的坑与调参经验6.1 学习率与Batch Size的匹配Transformer模型对学习率非常敏感这个跟CNN有比较明显的差别。CNN训练时用SGD加大学习率往往也能跑但Transformer一开始学习率设大了经常会看到loss不降反升甚至直接NaN。我的经验是ViT和Swin用AdamW时基础学习率设置在1e-4到2e-4之间比较稳batch size每翻一倍学习率也跟着适度上调但幅度不用跟线性缩放完全一致1.5倍左右即可。如果出现loss震荡或NaN第一件事不是调网络结构而是把学习率降到原来的十分之一或者把warmup步数拉长。ViT作者在原文里也比较强调warmup的作用前几轮让学习率从很小的值慢慢升上去会有效避免早期不稳定的问题。6.2 数据增强和正则化策略Transformer在中小数据集上的过拟合问题比CNN更明显。除了基本的RandomResizedCrop之外我建议再加一点MixUp或CutMix这两个增强策略对Transformer的提升比CNN更大。原因也比较直观Transformer是弱归纳偏置模型样本多样性的补充能显著缓解过拟合。另外Weight Decay的设置在Transformer上也要重新调。CNN里常用的weight_decay1e-4换成Transformer后可以考虑调到5e-2因为AdamW对weight decay的处理方式与SGD不同稍大一点的weight decay能起到更好的正则效果。Dropout率、Stochastic Depth的drop rate也都是可以从默认值往上调的尤其是在数据量不大的情况下。6.3 Position Embedding插值与迁移学习我在一个遥感图像分类项目里遇到过这样一个问题模型在ImageNet上预训练的时候用的是224x224但遥感图像往往是512x512甚至更大直接resize到224会丢失大量细节。于是我把模型输入调成了384x384这时就需要把预训练的position embedding从196个patch插值到576个patch。实际操作中用timm的resize_pos_embed函数可以很轻松完成这个操作但效果好坏还取决于插值后是否给了模型足够的微调轮数。我的经验是分辨率切换后至少需要正常微调轮次的1.5倍到2倍模型才能把位置信息重新学稳。如果在切换分辨率后只微调很少的轮次就评估精度往往比直接小分辨率还差这会让很多人误以为“模型不支持大分辨率”。6.4 Swin的窗口注意力和显存优化Swin在显存占用方面其实并不比ViT“天然低多少”。窗口注意力虽然降低了计算量但PyTorch等深度学习框架在实现自注意力时为了计算梯度会保存很多中间变量显存峰值依然不低。如果你训练Swin时爆显存了有以下几个实用优化方法开启torch.utils.checkpoint梯度检查点用计算换显存速度变慢但显存需求能降一半以上。减小batch size并增加梯度累积步数效果等同于原batch size。用混合精度训练AMP这是所有Transformer模型的标配不仅省显存还有大概率加速。混合精度这块PyTorch原生支持得很好加一个torch.cuda.amp.autocast()和GradScaler()就行。我在实际项目里跑Swin-L时如果不开AMP一张24G显存的卡连batch size 16都跑不了开了之后可以跑到32甚至更大。7. 从ViT到Swin之后的进化方向聊完ViT和Swin忍不住想简单展望一下这条技术线路的后续发展因为对我自己选择模型时有很大参考价值。Swin之后CV里出现了大量基于Swin的设计做进一步优化的模型比如ConvNeXt就把Swin的设计思想“反哺”给了CNN用标准的ResNet架构配上Swin的训练策略和数据增强最终精度能和Swin打平甚至略高而推理速度更快。这件事很值得深思很多时候架构本身的贡献也许没有想象中大训练策略、数据增强、优化器选择这些细节同样关键。还有一条线是让Transformer和CNN进一步融合比如一些混合架构把卷积用在stem和降采样阶段把Transformer用在高层语义建模阶段这样既保留了CNN的低层纹理建模优势又拿到了Transformer的全局建模能力。另一个值得关注的方向是把大核卷积重新推到台前RepLKNet这类工作证明了大卷积核也能获得接近Swin的性能但部署上对硬件更友好。从纯工程角度说选型时不要盲目追新。如果项目里对推理延迟要求严苛又希望模型效果好先不要直接上Swin或ViT先把ConvNeXt、RepLKNet这些“CNN改进型”测一遍。如果数据集特别大、算力也充足再考虑大规模Transformer模型加自监督预训练的路子效果上限往往会更高。最后分享一个小经验。我在实际调试过程中经常遇到有人问我“ViT和Swin到底哪个更好”我的答案通常是先搞清楚你的任务类型、数据规模、推理约束这三件事再谈模型选型。分类任务、数据量大、训练资源充裕ViT完全值得一试检测分割、中小数据集、有多尺度需求Swin大概率不会让你失望。至于中间地带的场景与其纠结理论优劣不如把你候选的两三个模型都跑一遍小规模实验看验证集上的收敛速度和最终精度这个实测出来的结论永远比任何paper里的对比表更有说服力。这篇梳理从原理讲到了代码细节一些参数和配置也都是我实际跑过的希望能帮你少走一点弯路。
返回列表