Swin Transformer代码逐行解析:从Patch Embedding到Window Attention的PyTorch实现细节
Swin Transformer代码精解:从Patch Embedding到Window Attention的工程实现
在计算机视觉领域,Transformer架构正逐渐取代传统CNN成为主流骨干网络。Swin Transformer作为其中的佼佼者,通过创新的滑动窗口机制和分层特征提取策略,在多项视觉任务中展现了卓越性能。本文将深入剖析Swin Transformer的PyTorch实现细节,特别关注Patch Embedding、Window Attention等核心模块的工程实现技巧。
1. 架构总览与初始化设计
Swin Transformer的整体架构采用分层金字塔结构,包含四个主要阶段(stage),每个阶段通过Patch Merging逐步降低分辨率同时增加通道数。这种设计使其能够处理不同尺度的视觉特征,从局部细节到全局语义都能有效捕捉。
模型初始化时需要注意几个关键参数配置:
def __init__(self, patch_size=4, in_chans=3, num_classes=1000,
embed_dim=96, depths=(2, 2, 6, 2),
num_heads=(3, 6, 12, 24), window_size=7,
mlp_ratio=4., qkv_bias=True, drop_rate=0.,
attn_drop_rate=0., drop_path_rate=0.1,
norm_layer=nn.LayerNorm, patch_norm=True,
use_checkpoint=False, **kwargs):
其中值得关注的参数设计包括:
- depths:各阶段的Transformer Block数量,通常后期阶段需要更多层来学习高级语义
- num_heads:多头注意力机制的头数随深度增加,与通道数增长保持比例
- drop_path_rate:采用线性递增的随机深度衰减策略,从0逐渐增加到0.1
模型参数初始化采用截断正态分布:
def _init_weights(self, m):
if isinstance(m, nn.Linear):
nn.init.trunc_normal_(m.weight, std=.02)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
这种初始化方式能有效避免训练初期的梯度爆炸问题,同时保证各层的输出分布稳定。
2. Patch Embedding的工程实现细节
Patch Embedding模块负责将输入图像转换为序列化的patch嵌入,其实现远比简单的图像分块复杂。核心类PatchEmbed通过卷积操作实现高效的分块与嵌入:
class PatchEmbed(nn.Module):
def __init__(self, patch_size=4, in_c=3, embed_dim=96, norm_layer=None):
super().__init__()
self.proj = nn.Conv2d(in_c, embed_dim,
kernel_size=patch_size,
stride=patch_size)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
实际工程实现中需要处理几个关键问题:
- 动态填充策略:当输入尺寸不是patch_size的整数倍时,自动进行右下方填充
def forward(self, x):
_, _, H, W = x.shape
pad_input = (H % self.patch_size[0] != 0) or (W % self.patch_size[1] != 0)
if pad_input:
x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1],
0, self.patch_size[0] - H % self.patch_size[0],
0, 0))
- 维度变换技巧:将2D特征图转换为序列形式时,采用
flatten+transpose组合而非直接reshape,保证内存连续性
x = x.flatten(2).transpose(1, 2) # [B, C, H, W] -> [B, HW, C]
- 归一化选择:默认使用LayerNorm而非BatchNorm,更适合序列数据且对batch大小不敏感
提示:在实际部署中,可以预先计算输入尺寸并调整padding策略,避免推理时的动态计算开销。
3. 滑动窗口机制的完整实现
Swin Transformer的核心创新在于其滑动窗口注意力机制,通过常规窗口(W-MSA)和移位窗口(SW-MSA)的交替使用,实现跨窗口信息交互。这一机制的实现涉及多个关键技术点。
3.1 窗口划分与还原
窗口划分函数window_partition将输入特征图划分为不重叠的局部窗口:
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)
windows = x.permute(0, 1, 3, 2, 4, 5).contiguous()
.view(-1, window_size, window_size, C)
return windows
对应的窗口还原函数window_reverse则是其逆过程:
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()
.view(B, H, W, -1)
return x
3.2 移位窗口的高效实现
移位窗口通过torch.roll实现,相比直接重新计算窗口划分更高效:
if self.shift_size > 0:
shifted_x = torch.roll(x,
shifts=(-self.shift_size, -self.shift_size),
dims=(1, 2))
移位后会产生新的窗口划分,需要特殊处理边缘区域。实践中采用mask机制来区分真实窗口和虚拟窗口:
def create_mask(self, x, H, W):
Hp = int(np.ceil(H / self.window_size)) * self.window_size
Wp = int(np.ceil(W / self.window_size)) * self.window_size
img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device)
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会用于注意力计算,确保不同区域的token不会错误交互。
4. Window Attention的优化实现
窗口注意力是Swin Transformer的计算核心,其实现需要考虑计算效率和内存占用的平衡。
4.1 相对位置编码的紧凑表示
Swin Transformer采用相对位置偏置来注入位置信息,其实现非常巧妙:
self.relative_position_bias_table = nn.Parameter(
torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads))
coords_h = torch.arange(self.window_size[0])
coords_w = torch.arange(self.window_size[1])
coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing="ij"))
coords_flatten = torch.flatten(coords, 1)
relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]
这种参数化方式将位置偏置从O(N²)压缩到O(2M-1)²,大幅减少了参数量。
4.2 注意力计算的内存优化
标准的注意力计算需要存储中间矩阵,内存占用为O(B·N²)。在窗口注意力中,通过分步计算和及时释放中间变量来优化:
q = q * self.scale
attn = (q @ k.transpose(-2, -1))
attn = attn + relative_position_bias.unsqueeze(0)
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)
attn = self.attn_drop(attn)
实际部署时还可以采用以下优化手段:
- 使用混合精度训练减少显存占用
- 对大型输入启用checkpointing机制
- 采用Flash Attention等优化实现
5. Patch Merging的下采样策略
Swin Transformer通过Patch Merging实现特征图下采样,其设计类似于CNN中的池化层但更加灵活:
class PatchMerging(nn.Module):
def __init__(self, dim, norm_layer=nn.LayerNorm):
super().__init__()
self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
self.norm = norm_layer(4 * dim)
其前向传播过程包含几个关键步骤:
- 空间重组:通过间隔采样将特征图分为四个子图
- 通道拼接:将四个子图沿通道维度拼接
- 线性投影:使用全连接层压缩通道数
def forward(self, x, H, W):
B, L, C = x.shape
x = x.view(B, H, W, C)
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) # [B, H/2, W/2, 4*C]
x = x.view(B, -1, 4 * C)
x = self.norm(x)
x = self.reduction(x) # [B, H/2*W/2, 2*C]
这种设计既保留了CNN局部连接的特性,又保持了Transformer的全连接优势,在降维的同时增强了特征表达能力。
6. 工程实践中的调优经验
在实际项目中应用Swin Transformer时,有几个关键点值得注意:
- 窗口大小选择:7×7窗口在大多数CV任务中表现良好,但对于高分辨率输入可适当增大
- drop path配置:建议采用线性递增的随机深度衰减,从0逐渐增加到0.1-0.3
- 混合精度训练:使用AMP自动混合精度可显著减少显存占用
- 自定义注意力掩码:对于不规则输入(如分割任务),需要调整mask生成逻辑
以下是一个典型的多GPU训练配置示例:
model = SwinTransformer(
embed_dim=128,
depths=[2, 2, 18, 2],
num_heads=[4, 8, 16, 32],
window_size=7,
drop_path_rate=0.3
)
scaler = torch.cuda.amp.GradScaler()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
with torch.cuda.amp.autocast():
outputs = model(inputs)
loss = criterion(outputs, targets)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
在模型部署阶段,可以考虑以下优化:
- 将模型转换为TorchScript格式
- 使用TensorRT等推理加速引擎
- 对Window Attention进行算子融合
更多推荐


所有评论(0)