从ViT出现之后,Transformer在视觉领域的热度就一直没降过。但真正让我觉得“这东西能落地”的转折点,恰恰是Swin Transformer。ICCV 2021的最佳论文,把注意力机制的复杂度从平方级拉回线性级,同时靠层次化设计让模型天然适配检测、分割这类密集预测任务。这篇笔记是系列的第8篇,我会把Swin从动机、原理到代码实现的关键细节完整拆一遍,同时结合我在实际项目里替换backbone时踩过的坑,给准备上手或者正在调参的朋友一些参考。不管你是刚接触视觉Transformer,还是已经在用Swin做下游任务,这篇都能帮你把底层的设计逻辑理清楚。
1. ViT的先天不足:为什么必须引入窗口机制
1.1 全局注意力算不起,本质是分辨率问题
ViT的做法是把图像切成16x16的patch,然后对所有patch做全局自注意力。这个思路在分类任务上效果很好,但如果输入图像是224x224,patch size是4或者8,token数量会急剧膨胀。假设patch size是4,那么序列长度就是56x56=3136个token,全局注意力的复杂度是O(N²D),也就是约一千万级别的计算量,一张卡很难吃得消。
对于检测和分割这类任务,输入分辨率通常更高,比如512、1024甚至更大。如果沿用全局注意力,显存和算力都会爆炸。换句话说,ViT在ImageNet分类上可以刷点,但在需要高分辨率输入的密集预测任务上,直接套用并不现实。
1.2 单尺度特征图让检测分割无从下手
CNN时代的backbone(比如ResNet)天然会产出多层特征图,从高分辨率低语义到低分辨率高语义,FPN、PAN等结构都依赖这种多尺度特征来做目标检测和语义分割。ViT输出的是一整条token序列,如果要恢复成单尺度的特征图,就只有一个分辨率,缺少层级结构。对分割头或者检测头来说,这种“没有金字塔”的特征并不好用。
Swin Transformer的提出,本质上就是同时解决这两个问题:用窗口注意力把复杂度降低到线性,同时通过Patch Merging构建层次化特征图,让Transformer能像CNN backbone一样方便地接入下游任务。这也是它能在COCO检测和ADE20K分割上全面超越CNN架构的重要原因。
1.3 局部性先验的引入与归纳偏置
另一个容易被忽略的点是,ViT的全局注意力意味着每个token都要和所有其他token交互,这对数据量的要求非常高。而CNN天然具备局部性和平移等变性这两种归纳偏置,所以即使数据量不大也能学得不错。Swin把注意力限制在窗口内,实际上就是引入了局部性先验,让模型更接近CNN的归纳偏置,训练难度也因此降低。
窗口尺寸固定为7x7或者8x8时,每个token只和窗口内的49个token交互,计算量大幅减少。不过如果一直固定窗口,不同窗口之间就没有信息交流,这会导致感受野受限。Swin的思路是:在相邻层之间移动窗口,让窗口边界发生变化,从而在深层实现跨窗口的信息传递。
2. Swin的核心设计逻辑:窗口注意力、shift与层次化特征
2.1 窗口注意力到底怎么算
假设输入特征图是HxWxC,窗口大小是MxM。先把特征图均匀划分成若干个不重叠的窗口,每个窗口内有MxM个token,对每个窗口内部做标准的自注意力。这样整体复杂度从O(H²W²)变成O(HW·M²),由于M是一个固定的小常数(常见的是7),复杂度就变成线性的了。
用公式来看更清楚:全局注意力的计算量是4HWC² + 2(HW)²C,窗口注意力是4HWC² + 2M²HWC。当HW远大于M²时,第二项被大幅压缩。这也是为什么Swin能在高分辨率下保持可接受的显存占用。
实现窗口划分最简单的方式是用einops的rearrange,或者用view和permute组合。需要注意的是,窗口划分必须保证H和W都能被M整除,所以在实际代码里通常会做padding,或者特征图尺寸本身就是32的倍数,这样窗口划分就很规整。
2.2 Shifted Window:跨窗口通信的关键
固定窗口注意力的问题是没有跨窗口连接,Swin的做法是交替使用W-MSA和SW-MSA。W-MSA是规则窗口,SW-MSA则是把特征图平移(M//2, M//2)之后再划分窗口。这样上一层的窗口边界在下一层就变成了窗口内部,不同区域的token就有机会交互。
Shift操作在代码里直接用torch.roll实现,非常轻量。但问题在于,平移之后特征图的边界处会拼出一些不完整的窗口,这些窗口并不是规则的MxM。直接对它们做注意力会引入不合理的像素位置关系,所以需要引入mask机制,把不属于同一窗口区域的位置屏蔽掉。这个mask矩阵的生成是Swin实现里比较绕的部分,后面我会专门讲。
2.3 Patch Merging:用空间换深度
Swin的层次化结构通过Patch Merging实现下采样。每两个相邻的patch(2x2的区域)在通道维度上被拼接起来,然后通过一个线性层把通道数翻倍并压缩回原来的维度。实际上就是类似于CNN里stride=2卷积的效果:空间分辨率减半,通道数翻倍。
以Swin-T为例,输入224x224x3,经过Patch Embedding后变成56x56x96,经过一个Stage(两层Swin Transformer Block)输出还是56x56x96,然后经过Patch Merging变成28x28x192,再经过Stage,以此类推。四个Stage输出的分辨率分别是56、28、14、7,通道数分别是96、192、384、768,这就得到了和ResNet类似的多尺度金字塔特征。
有了这个金字塔特征,就可以直接替代ResNet50作为检测、分割模型的backbone。不需要像ViT那样额外设计复杂的特征金字塔结构,FPN可以直接在Swin输出的多个stage上工作。
2.4 Swin Transformer Block的具体结构
一个Swin Transformer Block包含两个连续的Transformer Block,第一个是W-MSA,第二个是SW-MSA。每个Block内部由LayerNorm、多头注意力、MLP和残差连接组成,MLP的隐藏层维度是4倍并采用GELU激活。与ViT的Block相比,主要区别就在注意力部分,其余结构基本一致。
就是这层窗户纸:Swin不是发明了新的注意力机制,而是把注意力限制在局部窗口,同时用移动窗口的方式弥补跨窗口信息缺失。整体模型依然可以当作标准的Transformer来理解,很多在ViT上有效的训练技巧(比如AdamW、余弦退火、Mixup、CutMix等)都可以直接沿用。
3. 相对位置编码和掩码的代码实现细节
3.1 相对位置编码表的构造
Swin使用的是相对位置编码,而不是ViT的绝对位置编码。原因是窗口注意力中token的位置关系是相对固定的,相对位置编码能更好地泛化到不同尺寸的输入。实现时是先构造一个(2M-1)x(2M-1)的可学习参数表,然后根据窗口内每对token的相对坐标索引到这个表中取编码。
具体做法是在构造索引矩阵时,将每个token的坐标设为(行索引, 列索引),两两相减得到相对坐标。然后将相对行坐标偏移(M-1)、列坐标偏移(M-1),再算出行序号和列序号,最后得到一个M²xM²的索引矩阵,用这个矩阵去查位置编码表。
这是一个容易把人绕晕的部分,但代码写起来只要十几行。理解了这点,后面读源码或者改写就非常顺畅。
3.2 Shifted Window Attention的mask矩阵
当窗口经过shift操作后,会出现四个角落的窗口是不完整的情况。如果我们只对这些不完整窗口做padding,计算量会浪费,而且padding引入的像素值会干扰注意力。Swin的解决方案是,仍然把整个特征图划分成规则网格,每个窗口内可能会有来自不同“原来区域”的token,通过mask把不属于同一个区域的token之间的注意力分数置为负无穷。
这个mask矩阵的尺寸是(num_windows, M², M²)。在生成时,首先要根据shift的大小计算每个窗口内token所属区域的行列标签,然后比较每对token的标签是否一致,不一致的位置就mask掉。具体生成逻辑我建议直接看官方代码,但核心思路就是给每个窗口内部的token一个区域id,然后逐对比较id是否相等。
3.3 一个简洁的mask生成思路
很多初学者会被官方实现里那一堆复杂的索引计算劝退,其实可以换一种更容易理解的方式:在shift后的特征图上做一个和窗口大小一致的grid,标记每个位置所属的原始区域编号。然后对每个窗口,取window_size×window_size个编号,构造一个矩阵,比较对应元素是否相同,不同的位置就mask。
用伪代码表示就是:
# H, W为特征图尺寸,M为窗口尺寸,shift为移动距离 # 构造行、列方向编号,模拟区域划分 row_idx = torch.arange(H).unsqueeze(1).repeat(1, W) col_idx = torch.arange(W).unsqueeze(0).repeat(H, 1) # 判断是否属于同一个shift区域(可简化为每个token的行块和列块) # 具体细节需结合shift值计算region_id这么写虽然效率不如官方的高效索引实现,但逻辑非常清楚,适合用来理解原理。实际训练时再用官方的实现,也不会有性能问题。
3.4 实现完整Swin Block时的易错点
我在复现Swin Block时踩过几个坑,列出来供大家参考:
- Window partition之前,必须保证输入tensor是连续的,不然view会报错或者产生额外的拷贝。建议先调用
tensor.contiguous()。 - 做完相对位置编码查询后,维度是(B, num_heads, M², M²),需要reshape成(B*num_heads, M², M²),再和注意力分数相加。
- SW-MSA的mask必须和注意力分数在同一个device上,不然会出现device mismatch。最好在模型初始化时就生成好mask并注册为buffer。
- 不需要用
torch.roll的时候一定要设置shift_size=0,否则默认的roll会引入混乱。
代码写完之后,最好用一个随机输入验证一下W-MSA和SW-MSA的输出shape是否一致,同时检查mask是否起到了作用(可以故意把mask去掉对比loss是否能正常下降)。这样能快速定位问题。
4. 训练配置、下游任务接入与调参心得
4.1 ImageNet上的标准训练配置
Swin论文在ImageNet-1K上训练Swin-T的配置大体是:300个epoch,batch size 1024,AdamW优化器,初始学习率1e-3,weight decay 0.05,余弦学习率衰减,5个epoch的linear warmup。数据增强包括RandAugment、Mixup、CutMix、Random Erasing等。
我们在实际复现时,如果只有单机多卡(比如8张V100),完全可以照搬这套配置。但需要注意的是,batch size变化时,学习率要按比例调整。论文里用的是1024的batch size,如果你的batch size只有256,建议把初始学习率降到2.5e-4或者3e-4,否则训练容易不稳定。
另外,因为Swin的窗口注意力本身是局部计算,它的训练吞吐量在GPU上比ViT高很多,实际训练Swin-T的速度大约是ViT-B的1.3倍以上,显存占用也更低。这也是Swin在实际应用中更受欢迎的原因之一。
4.2 把Swin接入检测和分割框架
在MMDetection或者MMSegmentation里使用Swin,主要是把它作为backbone替换掉ResNet。需要关注的是输出通道数和输出stride。Swin-T的四层输出分别是[96,192,384,768],对应stride是[4,8,16,32],这和ResNet的C2-C5是能对应上的。
检测框架里通常会把最后两个stage送入FPN,所以需要把stage3和stage4的输出通道数配置给neck模型。另外一个容易被忽略的问题是,Swin的stage输出特征图大小依赖输入能否被32整除。如果输入是800x1024这类尺寸,特征图分别会是200x256、100x128、50x64、25x32,总体没问题,但某些梯队可能会多出1个像素,需要预处理时做padding或resize。
我在实际项目里用Swin替换ResNet50后,mAP大概提升了2到3个点,但训练显存和耗时也相应增加了。如果算力有限,建议结合sparse attention蒸馏或者用小模型(Swin-T)起步。
4.3 与FPN和检测头配合时的细节
一个值得注意的点是Swin每个stage内部的通道数和分辨率变化比较规整,FPN可以直接用。但Swin的stage输出特征图都是经过LayerNorm的,建议在接入FPN前先过一层1x1卷积或者对通道数做对齐,避免特征分布不匹配。
另外,Swin在stage的输出上通常会有norm层,在使用torchvision的模型API时,要记得在获取特征时用out_indices参数,并拿到输出的list,不要直接拿最后一个stage的特征做分类以外的任务,否则会丢失浅层细节。
4.4 调参时的观察与经验
- warmup非常重要。Swin对学习率比较敏感,尤其是大batch size下,过长的warmup有助于稳定initial training stage。一般最少5个epoch。
- 数据增强不要过于激进。Swin本来窗口是局部注意力,全局建模能力在某些情况下弱于ViT,因此对遮挡、形变的抗性略差。Mixup和CutMix可以加但不要一起加得太多,否则小模型容易欠拟合。
- 位置编码部分,Swin的相对位置编码表是(2M-1)x(2M-1)大小的,改窗口大小时要同步改这个表,否则会索引越界。实际使用中,我发现增大窗口从7到8换来约0.2-0.4个点的提升,但显存增长约10%,算力有限时不建议盲目加大。
- 训练精度如果出现NAN,多半是attention的mask写错了,或者学习率过大。可以先关掉AMP混合精度测试,再用torch.autograd.detect_anomaly定位。
5. Swin之后:与ViT、ConvNeXt的对比选型建议
5.1 Swin好在哪,差在哪
跟ViT相比,Swin的最大优势是能在高分辨率输入下保持较低的计算量,并且输出多尺度特征。缺点也很明显:窗口注意力限制了全局感受野,在某些需要长距离依赖的任务里(比如全景分割、视频理解)性能不如全局注意力的模型。不过在COCO检测和ADE20K分割上,Swin仍然是很好的baseline,很多state-of-the-art模型都用它做backbone。
跟ResNet相比,Swin在同等参数和FLOPs下精度更高,但推理延迟也更高。原因在于Transformer的算子相对复杂,部署到TensorRT或者ONNX时会遇到不少麻烦。如果只是为了刷点,Swin好用;如果是为了工业部署又对延迟敏感,可能还得优化算子或者换成改进版结构。
5.2 ConvNeXt给我的启示:Swin的设计可以被CNN借鉴
ICCV 2022的ConvNeXt是很有意思的工作,它反向把Swin的设计理念移植到CNN上,比如patchify stem、depthwise conv、倒瓶颈结构、大卷积核、更少的激活函数等。ConvNeXt在同等规模下能达到和Swin接近的性能,但推理速度更快,部署也更方便。
我在实际项目中会把两者都拉出来对比。如果只能用GPU而且对延迟要求不太苛刻,Swin-T在检测任务里表现更好;如果用CPU或者需要边缘部署,ConvNeXt-S是更务实的选择。这个对比也说明,Swin最核心的贡献不只是一个模型,而是一套从局部注意力到层次化设计的方法论,后来者可以从里面学到很多结构设计思路。
5.3 给初学者的一条路线
如果你想快速把Swin用起来,建议先跑通开源代码,再去看官方源码的窗口切分和mask部分,最后用一个小数据集做一次实验,对比一下Swin和ResNet在分割任务上的差异。不要一上来就调参,先把前向推理跑明白。
等你想把Swin改成自己需要的变体时,重点研究以下三个位置即可:Patch Embedding、Stage内部的Block循环、Patch Merging。其余部分跟常规Transformer Block没有本质区别。
我自己的体会是,Swin的代码一度被很多人觉得“难啃”,但只要你亲手把mask索引推导一遍,整个模型就完全通透了。这个推导过程虽然枯燥,但非常值——理解了一次之后,很多后来基于Swin的paper(比如CSWin、Focal Transformer、SwinV2)你都能很快看懂。
6. 从学术到落地:Swin的工程化注意事项
6.1 模型导出、量化与部署中的坑
Swin在PyTorch里跑起来很方便,但要导出到ONNX或者TensorRT,有几个点需要特别注意:
- 动态shape支持不好。窗口划分时如果图像长宽不是窗口尺寸的整数倍,很多实现会做padding或者循环。导出时建议固定输入尺寸,或者用TensorRT的显式动态shape做好充分测试。
- torch.roll在导出到ONNX时可能被拆成一堆Split和Concat操作,延迟增加,而且某些版本ONNX不支持,建议在导出时重写为直接通过索引实现循环位移,或者用crop+concat替代。
- 相对位置索引表如果用非整数可导操作,导出一般没问题,但注意要转成常量,避免被当成模型权重动态更新。
- 量化时要留意attention中的softmax,动态范围变化比较大,PTQ容易掉精度。建议对Attention的输入输出做校准,或者用QAT。
6.2 半精度训练与混合精度
Swin在混合精度的AMP下训练一般比较稳定,但偶尔会出现loss spike。可能的原因是softmax或者layer norm在fp16下的精度不足,特别是当注意力分数很分散时。这时候建议给attention部分的计算强制使用fp32,或者用torch.cuda.amp.autocast中指定忽略某些操作。
另外一个经验是,当batch size很大时,用bf16会比fp16更稳。虽然bf16在较新GPU上有加速,但数值表示范围和fp32相同,省去了很多调loss scaler的心力。
6.3 如何进一步压缩模型
如果你要在移动端跑Swin,原始模型肯定太庞大了。可以考虑:
- 通道剪枝。Swin各stage通道比较规律,对FPN输出影响大的通道保留,其余剪掉。但剪枝后要微调一段时间才能恢复精度。
- 知识蒸馏。用Swin-B作为teacher去蒸馏Swin-T甚至更小的学生模型,通常比直接训练小模型效果更好。蒸馏时temperature设在1-4之间,soft label的loss权重0.5-1.0。
- 用结构重参数化或窗口间的参数共享来减少参数量,但实现比较复杂,收益也不一定稳定。
6.4 现代化改进方向
Swin本身是2021年的工作,现在再从头复现一个Swin可能有些过时。我的建议是,把Swin当作理解视觉Transformer的“经典教材”,但实际项目里可以考虑SwinV2(针对稳定性和大规模训练优化)、CSWin(十字形注意力降低复杂度)、Focal Transformer(多尺度注意力)等改进版本,它们在精度和速度上往往有更好的折中。
不过,无论模型如何进化,Swin奠定的两个设计原则——局部注意力降低复杂度、层次化特征适配多尺度任务——到今天仍然成立。掌握这两个原则,再去理解其他架构就会轻松很多。
7. 最后的实操建议:从零手写一个可用的Swin-T
如果你打算自己动手实现,我建议按这个顺序推进:
- 实现Patch Embedding:把图片变成patch序列,并加一个可学习的绝对位置编码(Swin中其实没有绝对位置编码,但你可以加一个可选的)。
- 实现Patch Merging:把2x2patch按通道concat,然后线性映射为2C维输出。
- 实现WindowAttention类:支持输入窗口内的相对位置编码和mask。这是核心难点,务必逐行理解索引和mask。
- 实现SwinTransformerBlock:包含LayerNorm、WindowAttention、MLP、DropPath。
- 实现SwinTransformerStage:对输入做窗口划分,然后交替调用W-MSA和SW-MSA Block,最后做Patch Merging。
- 把四个Stage串成完整模型,加一个分类头。
- 在小数据集(CIFAR-100)上训练验证正确性,然后迁移到ImageNet继续调参。
手写一遍的效果远胜于直接调用mmclassification,因为你会彻底理解shift窗口的mask机制。我现在对自己团队新同学的培养,也基本是这个路子。
如果你的目标只是做应用,直接用开源的Swin预训练模型是最快的。在MMPretrain里加载权重时,注意源码里的相对位置索引表是否需要随输入分辨率重新初始化。Swin官方代码里提供了一个get_relative_position_index函数,是绝对的硬核,理解了它你就能彻底摆脱对他人实现的依赖。