SwinTransformer

发布时间:2026/8/10 9:12:48
SwinTransformer 一、Swin Transformer 核心前置认知1.1 ViT 的两大硬伤在讲Swin之前先搞清楚ViT有什么问题这样你才能理解Swin的每一个设计都是为了解决什么问题。ViTVision Transformer的做法很简单把图像切成一个个patch把每个patch当成一个token然后扔进标准的Transformer里。但它有两个硬伤计算量爆炸自注意力的复杂度是O(N²)N是token数量。分辨率越高patch越多计算量平方级增长。一张224×224的图切成16×16的patch有196个token要是512×512的图呢1024个token注意力矩阵就是1024×1024根本跑不动。多尺度能力弱ViT从头到尾只有一个特征尺度而检测、分割等任务需要不同层级的特征图类似CNN的特征金字塔。1.2 Swin 的核心思路两句话概括Swin Transformer的核心思想就是两个词窗口注意力 层级结构。窗口注意力Window Attention不在整张图上做注意力而是把特征图分成一个个小窗口只在窗口内做注意力。计算量从O(N²)降到O(N)线性增长。移位窗口Shifted Window光在窗口内做注意力不够窗口之间没有信息交互。每隔一层把窗口偏移一下让相邻窗口的token能互相看见实现跨窗口连接。层级结构Hierarchical类似CNN每经过一个stage分辨率减半、通道数翻倍天然形成特征金字塔方便下游任务使用。记住这三点后面所有的代码都是在实现这三个想法。二、整体架构一览先有全局观先看Swin-Ttiny版本的完整结构有个全局印象输入图像 224×224×3 │ ▼ Patch Embedding4×4卷积切patch │ 56×56×96 ▼ Stage 12个 Swin Block │ 56×56×96 ▼ Patch Merging下采样 │ 28×28×192 ▼ Stage 22个 Swin Block │ 28×28×192 ▼ Patch Merging下采样 │ 14×14×384 ▼ Stage 36个 Swin Block │ 14×14×384 ▼ Patch Merging下采样 │ 7×7×768 ▼ Stage 42个 Swin Block │ 7×7×768 ▼ LayerNorm → AdaptiveAvgPool → Linear → 分类输出几个关键数字Swin-T 配置embed_dim 96初始通道数depths [2, 2, 6, 2]每个stage的block数量num_heads [3, 6, 12, 24]每个stage的注意力头数window_size 7窗口大小每个窗口7×7个token是不是和ResNet很像4个stage分辨率越来越低通道数越来越高。这就是层级式的含义。三、核心模块源码逐行解析最核心部分以下代码均来自微软官方Swin-Transformer仓库的 models/swin_transformer.py我按数据流顺序从输入到输出逐模块拆解。3.1 Patch Embedding把图像切成patch第一步把图像转换成token序列。ViT用的是直接reshapeSwin用的是卷积效果一样但更高效。class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size4, in_chans3, embed_dim96, norm_layerNone): super().__init__() img_size to_2tuple(img_size) patch_size to_2tuple(patch_size) patches_resolution [img_size[0] // patch_size[0], img_size[1] // patch_size[1]] self.img_size img_size self.patch_size patch_size self.patches_resolution patches_resolution self.num_patches patches_resolution[0] * patches_resolution[1] self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) if norm_layer is not None: self.norm norm_layer(embed_dim) else: self.norm None def forward(self, x): B, C, H, W x.shape x self.proj(x).flatten(2).transpose(1, 2) # B Ph*Pw C if self.norm is not None: x self.norm(x) return x原理讲解用一个 kernel_size4, stride4 的卷积来切patch同时把3通道映射到96维。224×224×3 的输入 → 56×56×96 的特征图 → flatten成序列 (B, 3136, 96)3136就是token数量56×56每个token是96维向量。新手常问为什么用卷积而不是直接reshape效果完全等价但卷积实现更高效也方便后续改成重叠patch。3.2 MLP两层全连接 GELUclass 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 x标准的 Transformer FFN 结构先升维再降维中间用 GELU 激活。mlp_ratio4 表示隐层维度是输入的 4 倍。3.3 Window Attention窗口内的自注意力核心创新这是Swin最核心的创新。先看窗口怎么划分def window_partition(x, window_size): x: (B, H, W, C) 返回: (num_windows*B, window_size, window_size, C) B, H, W, C x.shape x x.view(B, H // window_size, window_size, W // window_size, window_size, C) windows x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C) return windows def window_reverse(windows, window_size, H, W): windows: (num_windows*B, window_size, window_size, C) 返回: (B, H, W, C) 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().view(B, H, W, -1) return x原理讲解把 (B, H, W, C) 的特征图切成 H/M × W/M 个窗口每个窗口大小 M×MM7。以56×56为例56/78一行8个窗口总共64个窗口。输出形状是 (num_windows*B, M, M, C)把batch和窗口数量混在一起后面做注意力时可以并行计算。关键技巧用 view permute contiguous view 实现无拷贝的窗口切分这是PyTorch里处理多维张量分块的标准写法。接下来是窗口注意力的主体注意里面的相对位置偏置class WindowAttention(nn.Module): def __init__(self, dim, window_size, num_heads, ...): super().__init__() self.num_heads num_heads head_dim dim // num_heads self.scale head_dim ** -0.5 # 相对位置偏置表(2*Wh-1) * (2*Ww-1) 种相对位置 self.relative_position_bias_table nn.Parameter( torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 计算相对位置索引预计算存在buffer里 coords_h torch.arange(window_size[0]) coords_w torch.arange(window_size[1]) coords torch.stack(torch.meshgrid([coords_h, coords_w])) 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[0] - 1 relative_coords[:, :, 1] window_size[1] - 1 relative_coords[:, :, 0] * 2 * window_size[1] - 1 relative_position_index relative_coords.sum(-1) self.register_buffer(relative_position_index, relative_position_index) self.qkv nn.Linear(dim, dim * 3) self.proj nn.Linear(dim, dim) def forward(self, x, maskNone): B_, N, C x.shape # 生成QKV qkv self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) q, k, v qkv[0], qkv[1], qkv[2] # 计算注意力分数 q q * self.scale attn (q k.transpose(-2, -1)) # 加上相对位置偏置从表里查表 relative_position_bias self.relative_position_bias_table[ self.relative_position_index.view(-1)].view( self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) relative_position_bias relative_position_bias.permute(2, 0, 1).contiguous() attn attn relative_position_bias.unsqueeze(0) # SW-MSA时用mask屏蔽不相关的区域 if mask is not None: nW mask.shape[0] attn attn.view(B_ // nW, nW, self.num_heads, N, N) mask.unsqueeze(1).unsqueeze(0) attn attn.view(-1, self.num_heads, N, N) attn self.softmax(attn) x (attn v).transpose(1, 2).reshape(B_, N, C) x self.proj(x) return x重点理解相对位置偏置Relative Position BiasViT用的是绝对位置编码直接加到patch embedding上。Swin用的是相对位置偏置加到注意力分数上Attention(Q, K, V) SoftMax(QK^T / √d B) V为什么用相对位置更符合视觉直觉图像中像素的相对位置比绝对位置更重要A在B的左上方比A在第3行第5列更有意义。参数更少窗口大小7×749个token相对位置的取值范围是(-6~6, -6~6)总共13×13169种相对位置。查表法先算出每对token的相对坐标映射到一维索引从表里查表取值。高效又优雅。3.4 Swin Transformer BlockW-MSA SW-MSA 交替一个Swin Block就是标准的Transformer Block结构只是把全局注意力换成了窗口注意力。但Swin有两种Block交替出现偶数层shift_size0普通窗口注意力W-MSA奇数层shift_sizewindow_size//2移位窗口注意力SW-MSAclass SwinTransformerBlock(nn.Module): def __init__(self, dim, input_resolution, num_heads, window_size7, shift_size0, mlp_ratio4., qkv_biasTrue, qk_scaleNone, drop0., attn_drop0., drop_path0., act_layernn.GELU, norm_layernn.LayerNorm): super().__init__() self.dim dim self.input_resolution input_resolution self.num_heads num_heads self.window_size window_size self.shift_size shift_size self.mlp_ratio mlp_ratio # 如果特征图比窗口还小就不用分窗口了 if min(self.input_resolution) self.window_size: self.shift_size 0 self.window_size min(self.input_resolution) assert 0 self.shift_size self.window_size self.norm1 norm_layer(dim) self.attn WindowAttention( dim, window_sizeto_2tuple(self.window_size), num_headsnum_heads, qkv_biasqkv_bias, qk_scaleqk_scale, attn_dropattn_drop, proj_dropdrop) self.drop_path DropPath(drop_path) if drop_path 0. else nn.Identity() self.norm2 norm_layer(dim) mlp_hidden_dim int(dim * mlp_ratio) self.mlp Mlp(in_featuresdim, hidden_featuresmlp_hidden_dim, act_layeract_layer, dropdrop) # 计算 SW-MSA 的 attention mask if self.shift_size 0: H, W self.input_resolution img_mask torch.zeros((1, H, W, 1)) h_slices (slice(0, -self.window_size), slice(-self.window_size, -self.shift_size), slice(-self.shift_size, None)) w_slices (slice(0, -self.window_size), slice(-self.window_size, -self.shift_size), slice(-self.shift_size, None)) 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, self.window_size) mask_windows mask_windows.view(-1, self.window_size * self.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)) else: attn_mask None self.register_buffer(attn_mask, attn_mask) def forward(self, x): H, W self.input_resolution B, L, C x.shape assert L H * W shortcut x x self.norm1(x) x x.view(B, H, W, C) # 循环移位 if self.shift_size 0: shifted_x torch.roll(x, shifts(-self.shift_size, -self.shift_size), dims(1, 2)) else: shifted_x x # 划分窗口 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, maskself.attn_mask) # 合并窗口 attn_windows attn_windows.view(-1, self.window_size, self.window_size, C) shifted_x window_reverse(attn_windows, self.window_size, H, W) # 反移位 if self.shift_size 0: x torch.roll(shifted_x, shifts(self.shift_size, self.shift_size), dims(1, 2)) else: x shifted_x x x.view(B, H * W, C) # 残差 FFN x shortcut self.drop_path(x) x x self.drop_path(self.mlp(self.norm2(x))) return x重点理解移位窗口 Mask 机制这是Swin最巧妙的设计。直接说原理普通窗口W-MSA窗口之间没有信息交互每个窗口各自算各自的。移位窗口SW-MSA把整张图往左上偏移半个窗口的距离再重新划分窗口。这样原来不相邻的窗口现在挨在一起了信息就能跨窗口流动。但移位之后边界上的窗口是由原来不相邻的区域拼起来的这些区域之间不应该有注意力连接。怎么办用attention mask同一区域的位置mask0正常注意力不同区域的位置mask-100softmax后趋近于0相当于屏蔽。用torch.roll做循环移位再配合mask是非常优雅的实现方式。如果用padding的方式来移位会引入额外的计算和边界处理效率更低。3.5 Patch Merging层级下采样Swin的层级特性就是靠Patch Merging实现的类似CNN里的池化层class PatchMerging(nn.Module): def __init__(self, input_resolution, dim, norm_layernn.LayerNorm): super().__init__() self.reduction nn.Linear(4 * dim, 2 * dim, biasFalse) self.norm norm_layer(4 * dim) def forward(self, x): H, W self.input_resolution B, L, C x.shape x x.view(B, H, W, C) # 隔点采样取2×2相邻的四个位置 x0 x[:, 0::2, 0::2, :] # 左上 x1 x[:, 1::2, 0::2, :] # 左下 x2 x[:, 0::2, 1::2, :] # 右上 x3 x[:, 1::2, 1::2, :] # 右下 x torch.cat([x0, x1, x2, x3], -1) # 通道拼接4C x x.view(B, -1, 4 * C) x self.norm(x) x self.reduction(x) # 线性层降到2C return x原理讲解把2×2相邻的4个token拼在一起通道数变成4倍。然后用一个线性层把4C降到2C。效果分辨率减半H/2, W/2通道数翻倍2C。以Stage 1到Stage 2为例输入(B, 56×56, 96)输出(B, 28×28, 192)这和CNN里stride2的卷积/池化效果一样但用的是可学习的线性变换。3.6 整体组装SwinTransformer 类class SwinTransformer(nn.Module): def __init__(self, img_size224, patch_size4, in_chans3, num_classes1000, embed_dim96, depths[2, 2, 6, 2], num_heads[3, 6, 12, 24], window_size7, ...): super().__init__() # Patch Embedding self.patch_embed PatchEmbed(img_sizeimg_size, patch_sizepatch_size, in_chansin_chans, embed_dimembed_dim, ...) # 随机深度衰减规则从浅到深逐渐增大 dpr [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # 构建4个stage self.layers nn.ModuleList() for i_layer in range(self.num_layers): layer BasicLayer( dimint(embed_dim * 2 ** i_layer), input_resolution(patches_resolution[0] // (2 ** i_layer), patches_resolution[1] // (2 ** i_layer)), depthdepths[i_layer], num_headsnum_heads[i_layer], window_sizewindow_size, drop_pathdpr[sum(depths[:i_layer]):sum(depths[:i_layer 1])], downsamplePatchMerging if (i_layer self.num_layers - 1) else None, ...) self.layers.append(layer) self.norm norm_layer(self.num_features) self.avgpool nn.AdaptiveAvgPool1d(1) self.head nn.Linear(self.num_features, num_classes) def forward_features(self, x): x self.patch_embed(x) x self.pos_drop(x) for layer in self.layers: x layer(x) x self.norm(x) x self.avgpool(x.transpose(1, 2)) x torch.flatten(x, 1) return x def forward(self, x): x self.forward_features(x) x self.head(x) return x整体数据流Swin-T输入 224×224×3 → PatchEmbed → 56×56×96 → Stage 12个block→ 56×56×96 → PatchMerging → 28×28×192 → Stage 22个block→ 28×28×192 → PatchMerging → 14×14×384 → Stage 36个block→ 14×14×384 → PatchMerging → 7×7×768 → Stage 42个block→ 7×7×768 → 全局平均池化 → 分类头 → 输出四、训练配置与工程细节涨点关键Swin能取得好效果除了结构设计训练技巧也功不可没。从官方config.py里整理一下4.1 优化器与学习率优化器AdamW权重衰减0.05基础学习率5e-4学习率策略余弦退火cosine训练轮数300 epoch热身20个epoch学习率从5e-7线性升到5e-44.2 数据增强三件套RandAugment自动数据增强策略Random Erasing随机擦除概率0.25Mixup CutMix混合增强两者按概率切换使用4.3 正则化技巧Stochastic Depth随机深度drop path rate 0.1从浅到深逐渐增大训练深层网络的标配Label Smoothing标签平滑0.1防止过拟合Gradient Clipping梯度裁剪最大范数5.0五、计算复杂度窗口注意力到底省多少算力最后来算一笔账直观感受一下窗口注意力的优势。假设特征图大小是H×W通道数C窗口大小M全局自注意力ViT总复杂度O(HW × C² (HW)² × C)关键项是 (HW)²分辨率翻倍计算量翻16倍。窗口自注意力Swin总复杂度O(HW × C² HW × M² × C)关键项是 HW × M²M7是固定的计算量随分辨率线性增长。以56×56的特征图为例全局注意力56×563136个token注意力矩阵 3136² ≈ 980万窗口注意力7×749个token64个窗口注意力矩阵 64 × 49² ≈ 15万差了64倍这就是Swin能处理高分辨率图像的根本原因。六、总结Swin Transformer 的三大贡献最后用三句话总结Swin Transformer的核心贡献窗口自注意力W-MSA把计算复杂度从平方级降到线性级让高分辨率图像的Transformer成为可能。移位窗口SW-MSA巧妙解决了窗口之间信息不流通的问题用循环移位 mask实现跨窗口连接。层级结构Patch Merging类似CNN的多尺度特征金字塔让Transformer能轻松适配检测、分割等下游任务。从代码角度看Swin的实现非常优雅window_partition / window_reverse 用张量操作实现无拷贝窗口划分torch.roll attention mask 实现移位窗口简洁高效相对位置偏置用查表法参数少效果好Swin Transformer之后几乎所有视觉Transformer的工作都在它的基础上改进。说它是视觉Transformer的里程碑一点不为过。参考资料论文Swin Transformer: Hierarchical Vision Transformer using Shifted Windows官方代码https://github.com/microsoft/Swin-Transformer