ARTICLE DETAIL

资讯详情

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

Swin Transformer骨干网络完全指南:从原理到检测分割实战

Swin Transformer骨干网络完全指南:从原理到检测分割实战 做CV项目这几年backbone的选择一直是决定整个任务上限的关键因素。我最早从VGG换到ResNet又从ResNet换到EfficientNet中间也试过纯ViT直接当骨干网络用但训练数据和调参成本实在让人头疼。直到Swin Transformer出来我才真正找到了一个既保留了Transformer的建模能力、又能在密集预测任务里用得舒服的backbone方案。这篇内容不是论文复读而是我把Swin Transformer真正塞进项目、跑通训练、解决各种奇怪问题之后的一些经验和总结。如果你正准备把Swin作为backbone接入自己的分类、检测或者分割任务这篇文章应该能帮你少走不少弯路。1. 换backbone时我为什么最终选了Swin Transformer先说一个我自己的项目背景。当时手头有个任务需要同时做图像分类和实例分割原来的方案是用ResNet-50做骨干网络分类精度和检测的mAP都到了一个瓶颈。我当时想过两条路一条是继续堆CNN的宽度和深度换ResNet-101甚至ResNeSt另一条是直接上ViT-Base靠全局注意力搞大感受野。但ViT-Base在只有几十万张图片的业务数据上表现并不理想它天生缺了CNN那种局部先验小数据下收敛很慢而且输出的特征图是单尺度的对检测和分割这种需要多尺度特征的任务非常不友好。就在这个节骨眼上我注意到了Swin Transformer它的设计正好踩中了我所有的需求点。Swin Transformer全称是Shifted Window Transformer核心卖点是它把Transformer的全局注意力改成了基于窗口的局部注意力并且通过窗口移位来建立跨窗口的信息交互。这么做的直接收益是计算复杂度从ViT的O(n²)降到了O(n)n是序列长度也就是token数量。对于图像任务来说token数量就是H×W/patch_size²分辨率一大ViT的计算量会爆炸而Swin几乎可以线性扩展。这一点在实际工程里太重要了因为检测和分割任务往往都要跑在512甚至1024分辨率上ViT跑到这种分辨率基本就是显卡性能测试工具而Swin还能保持一个体面的训练速度。另外一个让我下决心换Swin的关键原因是它的金字塔式特征结构。ViT输出的是单一分辨率的特征序列你要做检测还得在上面硬接一个neck去恢复多尺度信息但Swin把网络分成了四个stage每个stage的特征图分辨率逐级减半通道数逐级翻倍输出的多尺度特征和FPN、UperNet这些结构的输入要求天然匹配。Meaning你在检测任务里直接用Swin替换ResNet-50剩下的FPN和检测头基本不用大改只有通道数的对齐问题需要处理迁移成本比换成ViT低太多了。还有一点容易被大家忽略就是Swin在不同任务之间的通用性特别好。我后来在分类、检测、语义分割三个任务上都用了同一套Swin-Tiny作为骨干效果都比同量级的ResNet-50好。这种经验上的普适性在做技术选型时很重要因为你不需要为每个任务单独维护一套backbone团队协作和模型迭代的成本都会降下来。我当时还专门对比了很多backbone的ImageNet精度和下游任务表现Swin-Tiny在ImageNet上大概81.3%的top-1精度比ResNet-50的76-78%高了一截而在COCO检测上的mAP优势也稳定在3-5个点这种提升幅度在CV领域已经算非常可观了。不过Swin也不是万能的我后面会详细说它的坑。这里先给一个直观的选型判断标准方便你根据自己的任务情况做决策如果你的任务对分辨率很敏感比如遥感图像、医疗影像、文档版面分析这类细节密集的场景Swin比同参数量的CNN更值得试。如果你的任务本身训练数据很少比如只有几千张图那Swin的收敛难度比ResNet高需要更多的数据增强和更久的训练时间不一定划算。如果你需要在移动端或者嵌入式设备上部署Swin的结构虽然精度高但推理效率不一定比结构优化的CNN好需要谨慎评估。如果你只是需要一个现成的强backbone做快速baselineSwin-Tiny是可以无脑抄作业的选择它几乎是ResNet-50最好的平替。2. Swin的核心设计拆开看patch、分层与窗口注意力既然要把Swin当backbone用你就不能只当它是一个黑盒。它为什么比ViT更适合做骨干网络为什么能在ImageNet上拿高分这些问题都藏在它的结构设计里。我用比较好理解的方式拆一遍保证你听懂之后不仅能跑通代码还能在脑子里构建出它的前向传播图。2.1 Patch Embedding从像素到token的第一站Swin的输入处理方式和ViT类似都要先把图像切成patch做embedding。不过Swin用的patch size是4×4不像ViT那么大。它的做法是用一个卷积核大小和步长都是4的卷积层把H×W×3的输入图像变成H/4×W/4×C的特征图。这里的C对应每个stage的基础通道数Swin-T是96Swin-S是96Swin-B是128Swin-L是192。用4×4小patch的好处很明显patch越小保留的空间细节越多尤其对检测和分割这种像素级任务来说太粗糙的patch会丢掉小目标的特征。ViT用16×16的patch在ImageNet分类上还行但在COCO这种小目标密集的数据集上就明显吃力。Swin这种设计相当于在Transformer的全局建模能力前面先加了一层卷积式的局部感知很像给Transformer装了一个“视网膜”。2.2 四阶段金字塔通道翻倍、分辨率减半Swin总共分四个stage每个stage由Patch Merging和若干个Swin Transformer Block组成。第一个stage从H/4×W/4×C开始第二个stage先做Patch Merging把2×2相邻的patch合并成一个分辨率减半到H/8×W/8通道维度变成2C第三个stage再做一次变成H/16×W/16×4C第四个stage继续变成H/32×W/32×8C。这个逐级下采样的过程和CNN骨干网络非常像ResNet也是这么做的。正因为如此Swin输出的四个stage特征图天然就是金字塔结构对应到输入图像上就是1/4、1/8、1/16、1/32的步长。FPN直接把Swin四个stage的输出接进来就可以用这在检测分割里是巨大的便利。我见过的很多框架比如MMDetection里的SwinFPN配置就是这么干的。2.3 W-MSA与SW-MSA窗口注意力和移位窗口注意力这是Swin结构里最核心的部分。Swin Block的Attention是在窗口内做的每个窗口的大小是7×7个patch。比如输入是224×224经过4倍下采样得到56×56的patch网格56×56会被均匀切成8×8个7×7大小的窗口attention只发生在每个7×7窗口内部。这样做最直接的收益是算力。全局注意力的复杂度是O(n²)n是token数量窗口注意力每个窗口内部是O(49×49)所有窗口加起来是O(49×n)和n成线性关系。对于大分辨率输入这个差异是质的飞跃。我算过一笔账在1024×1024输入下纯ViT-Base的全局注意力序列长度是4096自注意力矩阵是4096×4096大概6700万对token的关系而Swin用7×7窗口后每个窗口只有49个token总计算量小了两个数量级还不止。但窗口注意力有个明显的缺陷窗口和窗口之间完全没有信息交流每个窗口内部只能看到自己那一片区域感受野被锁死在了局部。为了解决这个问题Swin的第二层注意力做了shifted窗口操作也就是把窗口划分整体向右下方偏移偏移量是窗口大小的一半3个patch。这样原本在两个窗口角落的相邻patch第二次注意力时就可能落在同一个窗口里实现了跨窗口的信息流动。Swin Block是W-MSA和SW-MSA交替排列的两个block一组一个用规则窗口一个用移位窗口任务就是让特征既保持局部精细建模又能逐步扩大感受野到全局。2.4 相对位置偏置Swin精度高的隐藏功臣Swin在注意力计算时除了QKV之外还加了一个相对位置偏执项B。这个B是每个head单独学习的它在每个窗口内根据token之间的相对位置关系取一个偏移值加到注意力得分上。说白了就是让注意力在建模的时候能感知到“这两个token间隔多远、在哪个方向”空间位置信息不再依赖绝对位置编码而是用相对位置来刻画泛化能力更强。官方实现里用的是Swin relative position index方案每个head的可学习偏置参数shape是(2×window_size-1)×(2×window_size-1)因为两个patch在窗口内的相对位置偏移范围就是-(window_size-1)到(window_size-1)。window_size7时这个偏置表的边长就是13。这个设计看起来不起眼但实际效果非常好消融实验里去掉这个偏置Swin-T在ImageNet上会掉1.5-2个点非常可观。2.5 Patch Merging降采样的时候在做什么Patch Merging其实就是把四个相邻patch的特征在通道维度上拼接起来然后通过一个线性层把4C的通道压缩成2C。这个操作等价于做了一次2倍下采样和CNN里的stride2卷积在功能上类似但有一点区别Patch Merging没有引入卷积核的空间混合只是纯通道维度的变换。所以我在实际使用中会专门注意Swin对特征图空间混洗的偏好比如数据增强里不要做太强的平移扰动因为Swin的窗口机制天然对平移不太敏感这个后面在训练策略部分再细说。3. 工程接入timm里的Swin、输入尺寸与训练超参调整理论部分讲得差不多了现在说正事——怎么把Swin真正接入你的训练代码。我用的是timm库它把Swin的实现封装得比较干净同时也会暴露一些关键参数。除非你有很强的定制需求否则不建议自己从零手写Swin没人愿意在debug window_index上浪费人生。3.1 用timm创建Swin模型timm里创建Swin模型非常简单import timm # 分类任务224输入 model timm.create_model(swin_tiny_patch4_window7_224, pretrainedTrue, num_classes1000) # 换成你自己的类别数 model timm.create_model(swin_tiny_patch4_window7_224, pretrainedTrue, num_classes10)模型名字里的命名规则其实已经透露了所有关键信息swin_tiny_patch4_window7_224表示Swin-Tiny、patch size为4、窗口大小为7、训练输入分辨率是224。Swin-S、Swin-B、Swin-L对应的模型名分别是swin_small_patch4_window7_224、swin_base_patch4_window7_224、swin_large_patch4_window7_224。如果是下游检测分割任务通常需要Swin输出多尺度特征这时候用features_only模式import timm backbone timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, features_onlyTrue, out_indices(0, 1, 2, 3), )这段代码会让backbone返回四个stage的特征图分辨率分别是输入的1/4、1/8、1/16、1/32通道数对应96、192、384、768Swin-Tiny。这四个特征图可以直接喂给FPN做多尺度融合。这里有个巨坑需要在最开始提醒Swin不像ViT那样有CLS token。ViT分类时是把CLS token拿出来接分类头而Swin是直接对最后一层特征图做全局平均池化然后接一个Linear分类头。所以在做迁移学习时你只需要替换最后一层head就行不要强行去拼接CLS相关的东西我见过有人从ViT代码改过来之后在Swin里硬找CLS token折腾了半天发现根本没有这个设计。3.2 window_size与输入尺寸的匹配最容易翻车的点Swin的窗口大小是固定的输入图像的尺寸必须是patch_size×window_size的整数倍否则会直接报错或者是默默做了resize导致对齐错误。拿224输入举例patch_size4token网格就是56×56窗口大小756÷78刚好整除。如果你有一天想直接拿这个模型跑384×384的输入384÷49696÷7≈13.71没法整除绝对会出问题。解决方案有两个一个是用官方专门为更高分辨率预训练的模型比如swin_tiny_patch4_window12_384它的窗口大小是12384÷49696÷128整除没问题另一个是自己在代码里写一个动态resize逻辑但会丢失预训练权重中的位置信息迁移效果会有损失。我的经验是如果你要在高分辨率下微调直接用官方对应的window_size版本别图省事随便改输入尺寸。这里我整理了一个常用分辨率与窗口参数的匹配表建议收藏输入分辨率patch_sizetoken网格window_size是否整除224×224456×567是384×384496×9612是512×5124128×1288是640×6404160×16010是256×256464×647否会报错512输入下的window_size8和640输入下的window_size10虽然在数学上能整除但官方没有预训练对应权重的模型你需要自己从224的模型上做位置插值或者直接从头训练。实际项目里如果需要512输入我更倾向于自定义一个Swin模型把window_size改成8随机初始化去训练或者用更大的预训练模型做finetune效果比强行插值好得多。3.3 训练超参照搬还是自定义Swin官方在ImageNet上的训练配置是这样的300个epochAdamW优化器初始学习率1e-3batch size1024时weight decay 0.0520个epoch的warmupcosine学习率衰减还有一定的drop path rate。drop path是Swin一个不能不提的超参Swin-T官方设置的drop path rate是0.1Swin-B会用到0.3模型越大、训练数据越多drop path就可以开得越大。它的作用相当于给每个残差分支做随机丢弃可以显著防止过拟合同时也让模型在深层次上更像一个集成模型。我把Swin官方推荐的训练超参整理成一个表格方便你对照参考超参数Swin-TSwin-SSwin-BImageNet-22K预训练优化器AdamWAdamWAdamW初始学习率1e-31e-31e-3weight decay0.050.050.05warmup epochs20205训练epochs30030090batch size102410244096drop path rate0.10.20.3数据增强CutMix, Mixup, RA同左同左如果你是在自己的业务数据上微调我不建议照搬这个大配置。一个很实际的经验当你的数据集只有一万张左右时初始学习率放到1e-4到5e-5之间比较稳妥训练epochs控制在50-100之间warmup可以缩短到5个epoch以内。如果把ImageNet上的1e-3直接拿来用训到第10个epoch你就会看到loss在高位震荡这就是学习率过大的信号。还有一个训练细节容易被忽略Swin对Bias和Norm层的weight decay有特殊的处理习惯。官方是用了weight decay为0的偏置和LayerNorm参数翻译成代码就是设置两个参数组一个是常规的weight decay另一个是bias和norm的weight decay为0。很多用Swin的人直接复制了别人的训练脚本从来没有检查过这个设置最终结果就是模型在验证集上掉的权重不理想。我建议你在写优化器的时候特别注意这个细节它虽然不能让模型立刻涨点但能让训练过程稳定不少。4. 实操中躲不开的坑从维度崩溃到NaNSwin在工程接入时会出现很多让人摸不着头脑的问题。这些问题大部分可以提前避开但如果没避开排查起来会很浪费时间。我把我踩过的、以及在社区里见过的高频问题整理了一下每个都附上了排查方法和解决方案。4.1 预设窗口参数和输入分辨率不符直接崩这是我见过最多的报错没有之一。用户拿着224预训练的模型直接换到384输入程序要么直接异常退出要么在某个中间层显示index out of range的报错。原因就是前面说的window_size7和patch数为96的token网格没法整除。我之前有次就是因为这个问题排查了很久。当时我是在一个开源检测框架里改backbone框架的dataloader会自动把输入resize到1333×800我用Swin-T替换ResNet之后训练到第一个iteration就崩了日志里报错的位置在get_relative_position_index当时还以为是框架的兼容性问题。后来把resize的尺寸打印出来才发现问题所在800/4200200/728.57压根对不上。排查方法很简单在dataloader里打印输入尺寸然后手动算一下除以4之后是否能被window_size整除。不能整除就调整resize的尺寸或者换一个window_size更合适的模型。千万不要指望模型的forward会帮你优雅处理这种尺寸不匹配你只会得到一个充满误导信息的报错。4.2 直接改forward导致relative position index异常Swin的relative position index是在初始化阶段根据window_size构建好的是一个持久化的buffer不会在forward动态更新。有些人想在forward里对特征图做任意尺寸的裁剪或者想做一些非常规的patch操作结果只要输出的token网格尺寸和初始化时的window_size不一致relative position index表查出来的索引就会越界或者完全错乱。这个问题我踩过一次之后就记住了如果要改动Swin的前向逻辑请先检查输出的特征图分辨率是否static。一旦需要动态分辨率你就得自己重写get_relative_position_index的逻辑或者使用更高版本的官方实现来判断尺寸。但更核心的建议是不要把Swin当成一个可以随意动态shape输入的模型来用它就是为固定分辨率预训练的backbone。真需要动态输入直接上CNN或者用支持动态shape的改进版本。4.3 混合精度训练下loss变成NaNSwin在fp16混合精度训练下比较容易出现训练不稳定的情况特别是当模型比较大、batch size比较小、或者learning rate设置不当的时候。这个问题的根源在于layer normalization和softmax操作在fp16下的动态范围有限梯度计算容易溢出。我处理这个问题有三次经验第一次是把GradScaler的init_scale调低大概1e-4附近稍微缓解了NaN问题第二次是把LayerNorm层强制保持在fp32只让Conv和Linear层用fp16第三次最彻底是全程用fp32训练虽然显存占用上去了速度慢了一些但loss曲线非常稳。最终我的建议是训练初期不要急着上混合精度先用fp32验证模型能收敛再逐步开启amp。如果一定要用amp至少在backbone部分使用fp32只在检测头或者分割头部分用fp16。4.4 检测分割任务中neck的通道不匹配Swin和ResNet替换时最容易忽略的其实是通道数对接问题。ResNet-50四个stage输出的通道是256、512、1024、2048Swin-T四个stage输出的是96、192、384、768。如果你直接改backbone的out_channels但它忘了改FPN的输入通道配置你会看到维度匹配的报错。这个问题倒不可怕因为它立即就能被发现。真正麻烦的是FPN内部如果对通道数有归一化的需求比如需要把输入压缩到256维Swin-T的第一层特征只有96个通道信息量可能会少一些需要协商调整。我在项目里习惯的做法是单独封装一个config类把backbone的输出通道写成可配置项检测框架里通过配置动态对齐而不是在代码里hardcode。这样切换不同backbone的时候只要把out_channels列表改掉就行其他不动。4.5 显存OOM和加速训练的小技巧Swin由于窗口注意力机制在相同参数量下比CNN更吃显存原因在于中间的attention score矩阵需要保存虽然窗口尺寸不大但每个stage的窗口数量很多。我实测在batch size8、输入384×384、Swin-T做检测backbone时8GB显存已经非常紧张了。解决显存问题的几个常用手段包括梯度累积gradient accumulation模拟更大的batch激活重计算activation checkpointing来用算力换显存以及减少输入分辨率。timm里的Swin已经支持set_grad_checkpointing但要注意这会明显拖慢训练速度我实测大概会慢30%-50%但对显存的释放效果非常显著。显存实在紧张的时候也可以用更轻量的Swin-T或者直接在检测头的设计上做减法比如把FPN的输出通道从256减到128。5. 下游任务扩展检测、分割与轻量化选型Swin作为backbone的真正价值体现在它接入下游任务后的表现。这里我结合自己的项目把Swin在检测、分割、轻量化选型三个方向上的应用思路讲清楚。5.1 检测任务SwinFPN的经典组合在MMDetection或者Detectron2框架下Swin作为backbone接入Mask R-CNN或者Cascade R-CNN已经是非常成熟的组合。Swin四个stage的特征图直接接FPNFPN负责自顶向下和侧向连接把语义信息和空间信息融合起来。Swin-L配合Cascade Mask R-CNN在COCO test-dev上的box AP能达到58这在两年前是非常惊艳的成绩。即使只用Swin-T和ResNet-50相比box AP也能高出3-5个点。这种提升在业务场景里意味着什么意味着你可以在不增加太多推理时间的前提下让模型的误检和漏检显著减少。如果你是想在现有检测框架里把ResNet换掉我建议先做两层改动第一层是把backbone实例化改成Swin第二层是把FPN的输入通道改成Swin各stage的输出通道。做完这两步先跑通一个1000步的smoke test确认loss在正常下降后再开始完整训练能省掉很多指数级增长的debug时间。5.2 分割任务SwinUperNet的黄金搭档在语义分割任务上Swin的最佳搭档是UperNet。UperNet是一个多层级特征融合的分割框架它会把Swin四个stage的特征都利用起来做一个金字塔池化后再融合。Swin-LUperNet在ADE20K数据集上的mIoU能到53.5左右在当时的SOTA水准。Swin-T在这个数据集上也能跑到44-46比ResNet-50高5个点左右而且参数量几乎持平。分割任务的输入分辨率通常很高动不动就是512或者768因此Swin的显存问题在分割任务里会被放大。我建议分割场景下优先考虑Swin-T或Swin-SLC的显存压力在常规设备上很难承受。另外分割模型的训练通常比较吃GPU显存开activation checkpointing几乎是标配。5.3 轻量化backbone选型Swin-T与轻量改进Swin-T本身有2800万左右的参数量输入224时FLOPs约4.5G这是一个很有竞争力的数字。但注意FLOPs低不代表速度快Swin的窗口attention在GPU上有不少访存开销实际推理速度并不比同FLOPs的CNN快。如果你打算在边缘端部署我会优先建议你选结构针对推理做优化的backbone比如基于卷积的RepVGG、MobileOne或者基于注意力但带硬件友好的FasterViT、MobileViT这些方案。如果你确实要在服务端使用Transformer骨干并且希望进一步轻量化可以考虑Swin后续的改进版Swinv2把窗口注意力的position bias改成了log空间连续取值支持不同分辨率之间更好的迁移Focal Transformer引入了细粒度token和粗粒度token的交互在大分辨率下的效率更高FasterViT则在Swin的结构基础上加了卷积路径的混合GPU实测吞吐要比Swin的原始实现高不少。我自己的选型建议表可以给你参考场景推荐backbone理由服务器端高精度检测Swin-L或Swin-B精度第一显存充足服务器端常规任务Swin-T或Swin-S精度和速度平衡边缘部署分类MobileOne/RepVGG推理速度优先边缘部署分割MobileViT或轻量CNN内存占用敏感高分辨率遥感/医疗Swin-Twindow_size8/10细节保留能力好最后再说一个关于权重的建议。Swin有很多从ImageNet-22K预训练再微调到ImageNet-1K的版本这些版本在收敛速度和最终精度上都比直接从1K训练的好。你在timm里看到的swin_base_patch4_window7_224_22k这类命名就是22K预训练的产物。如果业务数据量不大强烈建议直接用22K预训练的权重作为起点省时省力还能提高精度。我现在每次拿到一个新任务都会先确认输入分辨率和窗口尺寸是否匹配再看训练数据的量级决定要不要做大量数据增强然后才敢放心把Swin放到backbone的位置上。这套流程跑过几个项目之后已经很熟了也让我对Swin的脾气摸得比较清楚。希望这篇内容能帮你把Swin真正用起来而不是停留在跑通一个Demo的阶段。
返回列表