
第一次跑通 Swin Transformer 的代码是在我用 ViT 做目标检测被多尺度问题折磨了很久之后。那时候我对 Transformer 做视觉这件事其实有点怀疑ViT 在 ImageNet 上刷分确实猛但一放到 COCO 这种需要多尺度特征的任务上就明显不如 ResNet 系的老牌 backbone。Swin Transformer 出现在 ICCV 2021 上拿了当年的 Best Paper理由其实很简单——它把 CNN 最拿手的两个本事也就是层级下采样和局部归纳偏置重新装回了 Transformer 里而且几乎没牺牲 Transformer 本身的表达能力。这个组合让它能在图像分类、目标检测、语义分割三个方向同时刷新 SOTA这也是我后来愿意深入去啃它的原因。这篇文章不打算复读论文的 abstract我会把 Swin 从设计动机、核心机制、完整架构到代码实现的关键细节串起来讲尤其是那些论文里没明说、但实际训练中一定会遇到的坑一次说透。适合正在做视觉 Transformer 研究、想复现 Swin或者打算把它当 backbone 替换进自己项目里的朋友。1. 先说清楚Swin Transformer 到底被什么逼出来的1.1 ViT 在视觉任务上留下的两个硬伤ViT 的思路很直接把图像切成 16×16 的 patch拉平后当成 NLP 里的 token然后直接做全局自注意力。这个思路好在架构足够简洁坏处也恰恰藏在这个直接里。第一个硬伤是计算复杂度。全局自注意力的计算量和 patch 数量的平方成正比。一张 224×224 的图切成 16×16 的 patch会得到 196 个 token这个规模还好。但你想做检测、分割输入分辨率往往要提到 512 甚至 1024patch 数量直接变成 4 倍、16 倍注意力的计算量就变成 16 倍、256 倍。这种增长速度放在工程上是很难接受的。第二个硬伤更致命ViT 不管堆多少层特征图始终保持 1/16 分辨率。这在 ImageNet 分类上无所谓但检测和分割任务对多尺度特征有强需求——大目标要靠浅层的高分辨率特征小目标要靠语义更强的深层特征。ResNet 通过四个 stage 逐步下采样得到了 1/4、1/8、1/16、1/32 的金字塔特征FPN 直接拿来用而 ViT 给不了这个想做检测分割还得额外设计模块去凑特征金字塔既别扭又低效。1.2 Swin 的核心回应局部窗口加层级金字塔Swin Transformer 的回应可以用两句话概括第一句注意力只在固定的局部窗口内做窗口大小固定为 7×7计算量随图像分辨率线性增长第二句网络分成四个 stage每个 stage 之间做一次 2×2 的 Patch Merging把分辨率降一半、通道数翻一倍这样得到的特征金字塔和 ResNet 完全对齐。这两点单独看都不是新东西。局部注意力之前就有工作做过金字塔结构更是 CNN 的老传统。Swin 真正的贡献在于它找到了一个特别干净的方案同时拿下了这两点并且用 ICML 2021 的 Shifted Window 机制解决了只在窗口内做注意力会导致窗口间信息无法交互的问题。1.3 这篇文章会讲到什么程度我不会把 Swin 的每个参数都念一遍而是按下面这条线走先拆解 Shifted Window 这个核心机制把它的数学本质讲透然后过一遍完整的网络数据流搞清楚每个 stage 的维度变化接着给出可直接运行的核心代码包括窗口切分、相对位置偏置和 attention mask 的生成再看 Swin 在分类、检测、分割上的实际收益最后讲训练和复现过程中最容易踩的坑。每个部分我都会带上自己在实现和调试时的体会。2. 拆解移动窗口机制为什么 Shifted Window 是关键一战2.1 窗口自注意力的复杂度对比先看一组复杂度对比。设输入特征图大小为 h×w通道数为 C全局自注意力 MSAMulti-head Self-Attention的复杂度大约是4·h·w·C² 2·(h·w)²·C前半部分来自 QKV 的线性投影和输出投影后半部分来自注意力矩阵的计算。问题出在第二项它是 h·w 的平方。如果输入分辨率翻倍这一项直接变 4 倍分辨率变 512 的时候算力开销已经很难看了。Swin 的做法是把注意力限制在一个 M×M 的窗口内M 默认取 7。这样每个窗口内只有 M² 个 patch全局注意力矩阵换成了本地注意力矩阵复杂度变成4·h·w·C² 2·M²·h·w·C注意第二项M 是固定常量 7所以这一项随 h·w 线性增长。这就是为什么 Swin 能在高分辨率输入下依旧保持可行的根本原因。2.2 Shifted Window 如何打通窗口间的信息流动如果只是把注意力限制在局部窗口问题也来了窗口之间完全没有信息交换。第 l 层窗口左上角的一个 patch感受野永远被限制在它所在的窗口里不管你堆多少层都一样。这种情况下的 Transformer 其实退化成了一堆相互独立的局部小网络表达能力和我们想要的全局建模差了很远。Swin 的解决方案很巧妙窗口不是固定的而是交替使用的。第 l 层用普通窗口 W-MSA第 l1 层就把窗口整体向右下平移 (⌊M/2⌋, ⌊M/2⌋) 个像素也就是 3 个像素M7 时然后用新窗口做 SW-MSA。这么一移原本在窗口边界两边的 patch在新窗口里就处在了同一个注意力计算范围里。每个 patch 都有机会接触到原来窗口之外的邻居。两个连续的 Swin Block 组合起来就等价于一次跨窗口的信息传播。这种设计不需要额外增加计算量只是把窗口重新钉了个网格就实现了窗口间的信息交互。2.3 循环位移的实现逻辑与 Attention Mask 的作用这里有个工程实现上的细节容易把人绕晕。如果直接平移窗口你会发现两个问题第一窗口数量变了原来 h/M × w/M 个窗口平移后多出不少边缘碎片第二边缘会出现不完整的窗口没法直接做矩阵运算。论文的做法是循环位移cyclic shift也就是用 torch.roll 把特征图沿两个方向整体滚动。这样原本散落在边缘的碎片会被拼回完整窗口窗口数量保持不变。但代价是循环位移后同一个窗口里会混入从图像另一端滚过来的、原本不相邻的位置。这些位置之间不应该建立注意力连接否则语义就乱套了。所以必须用一个 attention mask 来遮住这些非法连接。mask 的生成逻辑是给特征图上每个位置打一个编号表示它属于循环位移前的哪个区域窗口切分之后同一个窗口内编号不同的位置之间注意力分数会被置为一个很大的负数经过 softmax 之后权重趋近于零。这样既享受了循环位移带来的整齐窗口划分又保证了内容上的正确性。3. 完整数据流从 Patch Partition 到四个 Stage 的维度变化3.1 起始的 Patch EmbeddingSwin 的输入处理比 ViT 更细。ViT 用 16×16 的 patchSwin 用 4×4 的 patch通过一个 kernel4、stride4 的卷积实现。这一步把输入从 H×W×3 变成 H/4 × W/4 × C其中 C 是 embedding 维度不同配置取值不同Tiny 是 96Small 是 96Base 是 128Large 是 192。为什么要用 4×4 而不是 16×16因为 Swin 要构建金字塔结构起始分辨率越细后面能下采样的次数就越多。16×16 的 patch 直接让特征图变成 14×14再下采样两次就到 4×4 了金字塔层级不够用。4×4 起步经过三次 patch merging可以得到 1/4、1/8、1/16、1/32 四个分辨率这个设计是跟检测分割任务的需求对齐的。3.2 Swin Transformer Block 里发生了什么Swin 的每个 Block 内部结构是这样的先是 LayerNorm然后接 W-MSA 或 SW-MSA两者交替残差连接再 LayerNorm接 MLP残差连接。如果你熟悉 ViT 的 Block会发现结构几乎一样唯一的区别就是把全局注意力换成了窗口注意力。一个 Stage 里有偶数个 Block成对出现。前一个用 W-MSA后一个用 SW-MSA这样交替保证了每个位置在经过一对 Block 后能接触到更广范围的上下文。论文给出的层数配置是Swin-T 和 Swin-S 为 [2, 2, 6, 2]Swin-B 和 Swin-L 为 [2, 2, 18, 2]。注意第三个 Stage 的层数明显更深这跟 ResNet 的设计逻辑一致——中间层分辨率适中、通道数够多是网络表达力的主要承载部分。3.3 Patch Merging 构建金字塔结构Stage 与 Stage 之间的下采样靠 Patch Merging 实现。它的操作很朴素把一个 2×2 邻域内的四个 patch 在通道维度直接拼接分辨率减半通道数变 4C然后过一个线性层把通道压回 2C。这个操作和卷积里的 stride2 卷积效果类似但它是无参数的拼接加线性变换信息保留更完整。经过三次 Patch Merging特征图从 56×56 变成 28×28再变 14×14最后到 7×7通道数从 96 逐步翻倍到 384Tiny 配置。到此Swin 输出的是 1/32 分辨率的最高层特征。这个层级结构和 ResNet 的 C2、C3、C4、C5 几乎等价所以 FPN、UperNet 这些原本为 CNN 设计的分割检测头可以直接在 Swin 上工作不需要任何调整。这是 Swin 能在检测分割任务上快速推广的一个重要工程原因。3.4 相对位置编码Swin 和 ViT 在位置信息上的分歧ViT 用的是一维绝对位置嵌入加在输入 token 上。Swin 的做法不一样它在注意力计算里加了一个可学习的相对位置偏置relative position bias加到 attention score 上而不是加到特征上。具体来说注意力公式变成Attention(Q, K, V) Softmax(Q·Kᵀ/√d B)·V这个 B 是相对位置偏置矩阵维度是 (2M−1) × (2M−1)每个 head 一组。窗口内任意两个位置之间的相对坐标差范围是 −(M−1) 到 M−1两个维度组合起来共有 (2M−1)² 种可能所以偏置表的大小是 (2M−1)²查表后得到一个 M²×M² 的偏置矩阵加到注意力分数上。为什么相对位置比绝对位置更适合视觉两个原因。第一相对位置天然具有平移等变性——图像里的目标无论出现在左上角还是右下角两个 patch 之间的相对位置关系不变这跟卷积的权值共享逻辑一致。第二窗口本身是局部的绝对位置在窗口内意义不大反而是相对偏移更能表达空间结构。4. 代码落地Swin Block 的核心实现与关键维度推演原理聊完直接上代码。我平时复现 Swin 的时候会把核心逻辑拆成四块来写窗口切分、相对位置偏置、Shifted Window 的 mask、Block 组装。4.1 窗口切分与恢复窗口切分就是把 (B, H, W, C) 的特征图切成 (B, H/M × W/M, M, M, C)然后 reshape 成 (B × nW, M, M, C)后面注意力直接在窗口维度上做def window_partition(x, window_size): B, H, W, C x.shape x x.view(B, H // window_size, window_size, W // window_size, window_size, C) x x.permute(0, 1, 3, 2, 4, 5).contiguous() x x.view(-1, window_size, window_size, C) return x def window_reverse(windows, window_size, H, W): B int(windows.shape[0] / (H * W / window_size / window_size)) x windows.view(B, H // window_size, W // window_size, window_size, window_size, -1) x x.permute(0, 1, 3, 2, 4, 5).contiguous() x x.view(B, H, W, -1) return x注意这里的 permute 顺序很多人在这一步把维度搞错。核心逻辑是先把 (H, W) 拆成 (H/M, M, W/M, M)然后通过 permute 把两个 M 放到中间最终得到窗口维度的排布。4.2 相对位置偏置表的生成相对位置偏置的代码是复现时最容易出错的地方核心在于索引的生成def get_relative_position_index(window_size): coords_h torch.arange(window_size) coords_w torch.arange(window_size) coords torch.stack(torch.meshgrid([coords_h, coords_w], indexingij)) coords_flatten torch.flatten(coords, 1) relative_coords coords_flatten[:, :, None] - coords_flatten[:, None, :] relative_coords relative_coords.permute(1, 2, 0).contiguous() relative_coords[:, :, 0] window_size - 1 relative_coords[:, :, 1] window_size - 1 relative_coords[:, :, 0] * 2 * window_size - 1 relative_position_index relative_coords.sum(-1) return relative_position_index这段代码的核心思路是先构建窗口内所有位置的坐标网格然后两两相减得到相对坐标因为相对坐标可能是负数先加上 window_size−1 把它变成非负然后把两个维度的索引合并成一个一维索引。归一化后的索引范围是 0 到 (2M−1)²−1正好对应偏置表的行数。实际使用时偏置表是一个可学习的参数self.relative_position_bias_table nn.Parameter(torch.zeros((2M−1)², num_heads))前向时按索引查表得到 (M², M², num_heads)再 permute 到 (num_heads, M², M²) 加到注意力分数上。4.3 Shifted Window Mask 的生成Mask 的逻辑是把循环位移前不同区域的位置编号在窗口内比较编号是否一致不一致就遮住def generate_mask(window_size, shift_size, input_resolution): H, W input_resolution, input_resolution img_mask torch.zeros((1, H, W, 1)) h_slices (slice(0, -window_size), slice(-window_size, -shift_size), slice(-shift_size, None)) w_slices h_slices cnt 0 for h in h_slices: for w in w_slices: img_mask[:, h, w, :] cnt cnt 1 mask_windows window_partition(img_mask, window_size) mask_windows mask_windows.view(-1, window_size * window_size) attn_mask mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2) attn_mask attn_mask.masked_fill(attn_mask ! 0, float(-100.0)).masked_fill(attn_mask 0, float(0.0)) return attn_maskh_slices和w_slices分别把特征图在 H 和 W 方向切成三段正常区域、从另一端滚过来的区域、被平移的区域。三层括号循环覆盖 3×39 个区域每个区域给一个不同的整数编号。窗口切分后把窗口内的编号展平两两相减编号相同说明是原本相邻的位置差为 0编号不同说明经过了循环位移原本不相邻差不为 0就把对应位置的注意力分数置为 −100。提示−100 是经验值目的是让 softmax 之后权重接近 0。如果你的模型是在 fp16 下训练的偶尔会遇到极端情况可以改用更小的值如 −10000。4.4 组装成完整的 Swin Block把上面几块组合起来一个完整的 Swin Block 长这样class SwinBlock(nn.Module): def __init__(self, dim, num_heads, window_size7, shiftFalse): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn WindowAttention(dim, num_heads, window_size) self.shift shift self.window_size window_size def forward(self, x, H, W): B, L, C x.shape shortcut x x self.norm1(x) x x.view(B, H, W, C) if self.shift: shifted_x torch.roll(x, shifts(-self.window_size // 2, -self.window_size // 2), dims(1, 2)) attn_mask self.mask else: shifted_x x attn_mask None x_windows window_partition(shifted_x, self.window_size) x_windows x_windows.view(-1, self.window_size * self.window_size, C) attn_windows self.attn(x_windows, maskattn_mask) x window_reverse(attn_windows, self.window_size, H, W) x x.view(B, H * W, C) if self.shift: x torch.roll(x, shifts(self.window_size // 2, self.window_size // 2), dims(1, 2)) x shortcut x x x self.mlp(self.norm2(x)) return x这里有几个维度细节需要特别强调。Swin 的输入输出都是 (B, L, C) 的序列形式其中 LH×W方便和预训练权重对接但窗口划分和循环位移都要求在 (B, H, W, C) 的空间形式下操作所以每次要来回 view。torch.roll的位移方向必须和 mask 的h_slices/w_slices定义保持一致方向反了或者 mask 写反了模型也能跑但精度会掉得很明显这是复现时最容易翻车的地方。5. Swin 在 ImageNet 和 COCO 上带来多少增量5.1 图像分类ImageNet-1K 的结果在 ImageNet-1K 上Swin-T 的 top-1 准确率约 81.3%Swin-S 约 83.0%Swin-B 约 83.5%Swin-L 约 86.3%使用 ImageNet-22K 预训练后微调可以到 87% 以上。单纯看这几个数字和同期的一些 CNN 模型相比提升不算夸张。Swin 真正的价值不在分类而在于它证明了Transformer 也能像一个正统的视觉 backbone 那样工作。ViT 在 ImageNet 上也能刷到不错的分数但一换任务就露馅Swin 是第一个在分类、检测、分割三个方向都全面超越同量级 CNN 的纯 Transformer 架构这让它具备很强的通用性。5.2 检测分割为什么 Swin 在这里赢面更大在 COCO 目标检测上用相同的 Mask R-CNN 框架ResNet-50 backbone 的 box AP 约 38 左右换用 Swin-T 后提升到 42 左右涨了约 4 个点。在 ADE20K 语义分割上用 UperNetResNet-50 的 mIoU 约 40 上下Swin-S 跑到 49 左右。对比之下分类任务只有一两个点的差距检测分割的涨点要明显得多。这也印证了设计动机层级金字塔结构和线性复杂度让 Swin 天然适配密集预测任务。检测需要多尺度特征Swin 给得出高分辨率输入下计算量可控Swin 扛得住相对位置编码提供了平移等变性对目标位置不敏感Swin 学得更稳。5.3 不同规模配置的对比我把常用配置整理成了一张表方便对照模型Embedding 维度层数参数量计算量224² 输入ImageNet top-1Swin-T96[2, 2, 6, 2]28M4.5 GFLOPs81.3%Swin-S96[2, 2, 18, 2]50M8.7 GFLOPs83.0%Swin-B128[2, 2, 18, 2]88M15.4 GFLOPs83.5%Swin-L192[2, 2, 18, 2]197M34.5 GFLOPs86.3%选择哪个配置取决于任务。分类和中小规模检测场景 Swin-T 是性价比之王计算量只有 ViT-B 的约 1/4大规模预训练迁移到检测分割用 Swin-B 起步比较稳Swin-L 适合资源充裕、追求极限精度的场景。6. 训练和调优 Swin 时的工程细节与常见坑6.1 分辨率与窗口尺寸的约束关系Swin 对输入分辨率有一个比较麻烦的约束每个 stage 的特征图边长必须能被窗口大小 7 整除。224×224 输入下patch 之后 56×5656 能被 7 整除没问题但如果你把输入调到 384×384特征图是 96×9696÷7 是 13.71整除不了窗口划分就会出问题。在工程里两个常用方案一是把输入调到 448×448 或 224×224 这种能被 4 和 7 同时整除的尺寸二是重新实现窗口划分逻辑对边缘做 padding。这里提醒一下如果直接加载官方预训练权重window_size 相关的位置偏置表是固定的换分辨率后必须重新插值或重新训练否则会崩。6.2 混合精度和优化器的坑我遇到过 Swin 在混合精度训练下 loss 变成 NaN 的情况。排查下来主要是两个原因一是注意力分数加上相对位置偏置后直接做 softmax浮点溢出二是 LayerNorm 在 fp16 下的数值稳定性不够。稳妥的做法是把 LayerNorm 固定在 fp32softmax 保持 fp32 计算。PyTorch 的 autocast 默认会在某些算子中间保留 fp32但不是所有版本都保险自己加一个强制 fp32 的 wrapper 更放心。优化器建议用 AdamWbetas(0.9, 0.999)weight decay 设 0.05。初始学习率按 batch size 缩放参考公式是lr 1e-3 × batch_size / 1024。Batch size 1024 时用 1e-3batch size 256 时就降到 2.5e-4 左右。Warmup 设 20 个 epoch后面用 cosine decay 衰减到接近 0。数据增强方面直接用 DeiT 那套配置就行RandAugment、MixUp、CutMix、Random Erasing 全套上Swin 对这套增强的适应性很好。6.3 复现时最容易翻车的三个细节第一个坑是 torch.roll 的方向和 mask 定义必须配套。很多复现版本代码是从官方仓库改的roll 方向和 mask 切片方向写反了模型也能 forward但指标会掉 1 到 2 个点而且很难查。排查方法很简单打印出第一个 Block 的 attention 权重的方差如果和官方模型差距很大大概率就是这里出了问题。第二个坑是torch.meshgrid的indexing参数。PyTorch 新版本默认indexingij早期版本默认xy两种模式生成的坐标顺序不一样直接影响相对位置索引的顺序。如果你的复现代码相对位置矩阵乱掉先查这个。第三个坑是 checkpoint 的转换。Swin 官方权重里所有 Block 交替标记了shift属性加载时是按 (W-MSA, SW-MSA, W-MSA, SW-MSA) 的顺序来的。如果你自定义的 Block 顺序和官方不一致加载权重的 name 对不上微调效果会大打折扣。用一个简单的 key 对齐脚本检查一下每个 Transformer Block 的shift标识就好。最后再说一个和 Input Resolution 相关的调试经验Swin 在低于训练分辨率下推理时相对位置索引表会越界因为(2M−1)²是根据训练时的窗口大小生成的。如果你做多尺度测试比如 TTA需要为每个尺度重新生成 position bias 索引不能直接复用 224×224 训练时的那套参数。这个细节在我第一次做多尺度推理时卡了很久最后查源码才意识到问题出在索引的静态绑定上。希望这篇拆解能帮你少走几步弯路把精力真正花在模型设计和实验验证上。