ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

Swin Transformer 2D相对位置编码:原理、实现与工程实践

Swin Transformer 2D相对位置编码:原理、实现与工程实践

1. 项目概述:从绝对位置到相对位置的注意力革命

在视觉Transformer的演进道路上,位置编码一直是个绕不开的核心议题。早期的ViT直接将NLP中的绝对位置编码(APE)搬过来用,给每个图像块(patch)分配一个固定的位置向量。这方法简单直接,但很快就暴露了问题:当模型处理训练时没见过的图像分辨率时,这些固定的位置编码就“对不上号”了,模型性能会显著下降。这就像你背熟了一张固定座位表,突然换到一个更大或更小的教室,你就找不到北了。

Swin Transformer的横空出世,引入了划时代的窗口多头自注意力(W-MSA)移位窗口(SW-MSA)机制,极大地提升了计算效率和建模长距离依赖的能力。但随之而来的,是一个更微妙的位置问题:在固定的、非重叠的窗口内部,模型如何感知像素之间的相对位置关系?Swin Transformer的答案是二维相对位置偏置(2D Relative Position Bias, 2D-RPE)。这不是一个简单的技术点,而是理解Swin为何能在保持线性计算复杂度的同时,实现强大视觉表征能力的关键钥匙。

简单来说,Swin的2D-RPE为注意力机制注入了一种“空间先验”。它不告诉模型“你在第几行第几列”(绝对位置),而是告诉模型“你和我之间,在水平和垂直方向上分别差了多少个像素”(相对位置)。这种设计天生就具备了平移不变性(Translation Invariance)的潜质,也是其能优雅处理多尺度输入的核心原因之一。本文将深入拆解Swin Transformer中2D-RPE的设计思想、实现细节、背后的数学原理,以及在实际应用和模型改进中你可能会遇到的坑与技巧。

2. 核心设计思想与方案选型解析

2.1 为何放弃绝对位置编码(APE)?

在标准Transformer中,APE通常是一个可学习的参数矩阵,其形状为(num_patches + 1, dim),其中+1是为了CLS token。对于图像任务,将二维坐标展平为一维后使用。其根本缺陷在于:

  1. 分辨率敏感:训练时固定了序列长度(即图像块数量)。如果推理时图像分辨率改变,序列长度变化,预训练的APE矩阵无法直接使用。虽然可以通过插值来适应新分辨率,但这会引入误差,并非原生支持。
  2. 缺乏平移不变性:计算机视觉的许多任务(如物体检测、分割)要求模型对物体的平移具有不变性。APE明确编码了绝对位置,与这一先验略有冲突。
  3. 不符合视觉直觉:人类识别物体,更多依赖的是物体部件之间的相对关系(眼睛在鼻子上面,轮子在车身下面),而非其在图像中的绝对坐标。

Swin Transformer的窗口化注意力设计,使得注意力计算被限制在一个局部窗口内。在这个局部上下文中,相对位置信息比绝对位置信息更有意义,也更容易建模。

2.2 相对位置编码(RPE)的范式转变

相对位置编码的核心思想是:在计算查询向量(Query)和键向量(Key)的注意力得分时,额外加入一个偏置项,这个偏置项仅由查询元素和键元素之间的相对位置决定。

公式上,标准注意力计算为:Attention(Q, K, V) = Softmax(QK^T / sqrt(d_k)) V

加入相对位置偏置B后,变为:Attention(Q, K, V) = Softmax(QK^T / sqrt(d_k) + B) V

这里的B就是一个矩阵,其中元素B_{i,j}表示第i个查询(位于某个位置)与第j个键(位于另一个位置)之间的相对位置偏置。在Swin中,ij是同一个窗口内的两个图像块。

2.3 Swin 2D-RPE 的具体方案选型

Swin Transformer的作者们做出了几个关键且巧妙的设计选择:

  1. 参数化与共享:相对位置偏置B被设计为一个可学习的参数,而不是通过正弦余弦函数生成。这意味着模型可以从数据中学习到哪种相对位置关系应该被加强或减弱。更重要的是,这个偏置参数在所有窗口、所有层、所有头之间共享。这是一个很强的归纳偏置,假设“相同的相对位置关系,在任何地方、任何语义层次上都具有相似的重要性”。实践证明,这个假设非常有效,且极大减少了参数量。

  2. 离散化的二维相对坐标:这是Swin RPE最精髓的部分。对于一个大小为M x M的窗口(例如7x7),窗口内共有M^2个图像块。任意两个块之间都有一个二维的相对位移(Δx, Δy)ΔxΔy的取值范围都是[-(M-1), M-1]。Swin的作者将连续的相对坐标离散化,映射到一个有限的索引上。

    • 首先,将ΔxΔy分别加上(M-1),使其范围变为[0, 2M-2]
    • 然后,将这两个维度上的坐标展平为一个一维索引:index = Δx * (2M-1) + Δy。因为ΔxΔy各有(2M-1)种可能,所以总共会有(2M-1)*(2M-1)个独特的相对位置对。
    • 最后,我们初始化一个形状为((2M-1)*(2M-1), num_heads)的可学习参数表relative_position_bias_table。通过计算出的index,我们就可以从这个表中查取出对应所有注意力头的偏置值。
  3. 与注意力头的解耦:偏置参数表最后一维是num_heads。这意味着每个注意力头都有自己独立的一套相对位置偏置。这赋予了模型更大的灵活性:有的头可能更关注局部(小位移)关系,有的头可能更关注窗口内较远的关系,模型可以自行学习。

注意:这里有一个极易混淆的点。许多初学者会认为relative_position_bias_table的形状是(num_heads, (2M-1)*(2M-1))。在代码实现中,两种维度顺序都有可能出现,取决于后续矩阵加法的便利性。关键是要理解其物理意义:它是一个查询表,为每一种可能的二维相对位置关系,存储了所有注意力头对应的偏置值。

3. 核心细节解析与实操要点

3.1 相对位置索引的生成:代码级详解

理解索引的生成是复现RPE的第一步。下面我们以M=7(窗口大小7x7)为例,拆解这个过程。

import torch def generate_relative_position_index(window_size=7): """ 生成用于索引 relative_position_bias_table 的索引矩阵。 返回的索引矩阵形状为 (M*M, M*M) """ M = window_size # 1. 生成每个位置的绝对坐标 (0到M-1) # 使用 meshgrid 生成坐标网格 coords = torch.stack(torch.meshgrid([torch.arange(M), torch.arange(M)])) # 形状 (2, M, M) # 展平为 (2, M*M),每一列代表一个位置的(x, y)坐标 coords_flatten = coords.flatten(1) # 形状 (2, M*M) # 2. 计算所有位置对之间的相对坐标 # coords_flatten[:, :, None] 形状 (2, M*M, 1) # coords_flatten[:, None, :] 形状 (2, 1, M*M) # 相减后得到 relative_coords 形状 (2, M*M, M*M) # relative_coords[0] 是所有对的 Δx # relative_coords[1] 是所有对的 Δy relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 形状 (2, M*M, M*M) # 3. 将相对坐标从 (Δx, Δy) 转换为一维索引 # 首先,将坐标偏移到非负范围 relative_coords += M - 1 # 现在范围是 [0, 2M-2] # 4. 将 x 和 y 坐标展平为一维索引 # 给 x 坐标乘以 (2M-1),然后加上 y 坐标 relative_coords = relative_coords.permute(1, 2, 0).contiguous() # 形状变为 (M*M, M*M, 2) relative_coords[:, :, 0] *= (2 * M - 1) # Δx 分量乘以跨度 relative_position_index = relative_coords.sum(-1) # 形状 (M*M, M*M) return relative_position_index # 示例 index_matrix = generate_relative_position_index(7) print(f"索引矩阵形状: {index_matrix.shape}") print(f"索引取值范围: [{index_matrix.min()}, {index_matrix.max()}]") print(f"理论唯一索引数量: {(2*7-1)*(2*7-1)} = {13*13}")

这段代码的输出会验证:生成的index_matrix是一个49x49的矩阵,里面的每个值都在[0, 168]之间(因为(2*7-1)^2 = 169)。这个矩阵就是后续查询偏置表的“地图”。

实操要点1:permutecontiguous的重要性在步骤4中,permute(1,2,0)是为了将形状从(2, 49, 49)变为(49, 49, 2),以便对最后一维(x,y)进行操作。紧接着调用.contiguous()是PyTorch中的最佳实践。permute操作只改变了张量的视图(stride),并未实际改变内存布局。某些后续操作(如view或作为某些函数的输入)要求张量在内存中是连续的,contiguous()会确保这一点,避免潜在的运行时错误。

3.2 偏置表的初始化与使用

偏置表是一个可学习参数,通常在全模型初始化时被定义。

import torch.nn as nn class WindowAttentionWithRPE(nn.Module): def __init__(self, dim, window_size, num_heads): super().__init__() self.dim = dim self.window_size = window_size self.num_heads = num_heads # 计算唯一相对位置的数量 self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) # 定义可学习的相对位置偏置表 # 形状: (num_relative_distance, num_heads) 或 (num_heads, num_relative_distance) # 这里采用第一种,便于后续广播加和 self.relative_position_bias_table = nn.Parameter( torch.zeros(self.num_relative_distance, num_heads) ) # 生成并注册不参与学习的相对位置索引缓冲区 # 这是一个固定的查找表,不需要梯度 relative_position_index = generate_relative_position_index(window_size[0]) self.register_buffer("relative_position_index", relative_position_index) # ... 其他初始化代码 (qkv投影层, 缩放因子等) ... def forward(self, x, mask=None): """ x: 输入特征,形状为 (num_windows*B, M*M, C) """ B_, N, C = x.shape # B_ = num_windows * B # 1. 计算Q, K, V 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] # 每个形状 (B_, num_heads, N, C//num_heads) # 2. 计算缩放点积注意力分数 attn = (q @ k.transpose(-2, -1)) * self.scale # 形状 (B_, num_heads, N, N) # 3. 关键步骤:加上相对位置偏置 # 从表中根据索引取出偏置 # self.relative_position_index 形状 (N, N) # self.relative_position_bias_table 形状 (num_relative_distance, num_heads) # 索引后得到 relative_position_bias 形状 (N, N, num_heads) 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 ) # 形状 (N, N, num_heads) # 调整维度以匹配attn: (B_, num_heads, N, N) relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # 形状 (num_heads, N, N) attn = attn + relative_position_bias.unsqueeze(0) # 广播加到每个batch和窗口 # 4. 如果存在窗口移位带来的mask,在这里加上mask if mask is not None: # mask 形状 (nW, N, N), nW是窗口数量 attn = attn.view(B_ // mask.shape[0], mask.shape[0], self.num_heads, N, N) attn = attn + mask.unsqueeze(1).unsqueeze(0) attn = attn.view(-1, self.num_heads, N, N) # 5. Softmax和Value加权 attn = attn.softmax(dim=-1) x = (attn @ v).transpose(1, 2).reshape(B_, N, C) x = self.proj(x) return x

实操要点2:register_buffer的妙用relative_position_index是一个固定的、根据窗口大小计算出来的整数张量,它不参与训练。使用self.register_buffer('name', tensor)将其注册为模块的缓冲区。这样做的好处是:

  • 它会被自动转移到正确的设备(GPU/CPU)上,与模型参数同步。
  • 它会被包含在模型的state_dict中,因此保存和加载模型时,这个预计算的索引也会被保存和加载,保证一致性。
  • 它不参与梯度计算,节省了显存和计算量。

实操要点3:视图(view)与维度变换的陷阱forward函数中,从偏置表查取出数据后,有一系列viewpermute操作。这里的顺序和维度必须非常小心。一个常见的错误是维度不匹配导致view操作失败。在view之前使用contiguous()是一个安全的好习惯。建议在编写这部分代码时,使用print(tensor.shape)或调试器逐步检查每个中间张量的形状,确保与预期一致。

4. 实操过程与核心环节实现

4.1 从零实现一个带2D-RPE的窗口注意力模块

让我们整合前面的知识,构建一个完整的、可嵌入到Swin Block中的注意力模块。我们将考虑移位窗口(Shifted Window)所需的注意力掩码(mask)。

import torch import torch.nn as nn import torch.nn.functional as F class ShiftedWindowAttention2D(nn.Module): """ 一个完整的、支持移位窗口和2D-RPE的注意力模块。 假设输入特征图已经被分割成了窗口。 """ def __init__(self, dim, window_size=(7,7), num_heads=8, qkv_bias=True, attn_drop=0., proj_drop=0.): super().__init__() self.dim = dim self.window_size = window_size self.num_heads = num_heads head_dim = dim // num_heads self.scale = head_dim ** -0.5 # 相对位置偏置表 self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) self.relative_position_bias_table = nn.Parameter( torch.zeros(self.num_relative_distance, num_heads) ) # 初始化偏置表,通常使用截断正态分布 nn.init.trunc_normal_(self.relative_position_bias_table, std=.02) # 生成相对位置索引 coords_h = torch.arange(window_size[0]) coords_w = torch.arange(window_size[1]) coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing='ij')) # (2, Wh, Ww) coords_flatten = torch.flatten(coords, 1) # (2, Wh*Ww) relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # (2, Wh*Ww, Wh*Ww) relative_coords = relative_coords.permute(1, 2, 0).contiguous() # (Wh*Ww, Wh*Ww, 2) 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) # (Wh*Ww, Wh*Ww) self.register_buffer("relative_position_index", relative_position_index) # 线性投影层 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) self.proj = nn.Linear(dim, dim) self.proj_drop = nn.Dropout(proj_drop) def forward(self, x, mask=None): """ Args: x: 输入特征,形状为 (num_windows * B, N, C),其中 N = Wh * Ww mask: (可选) 注意力掩码,用于移位窗口,形状为 (nW, N, N) 或 (B*nW, N, N) Returns: 输出特征,形状同输入x """ B_, N, C = x.shape # 生成Q, K, V 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] # 每个形状 (B_, num_heads, N, head_dim) # 计算注意力分数 attn = (q @ k.transpose(-2, -1)) * self.scale # (B_, num_heads, N, N) # 添加相对位置偏置 relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)] relative_position_bias = relative_position_bias.view( self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1 ) # (N, N, num_heads) relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # (num_heads, N, N) attn = attn + relative_position_bias.unsqueeze(0) # 广播到batch维度 # 应用注意力掩码(如果提供) if mask is not None: nW = mask.shape[0] # 掩码的窗口数 # 将attn的batch维度拆分为 实际batch * nW attn = attn.view(B_ // nW, nW, self.num_heads, N, N) attn = attn + mask.unsqueeze(1).unsqueeze(0) # 广播添加掩码 attn = attn.view(-1, self.num_heads, N, N) attn = self.attn_drop(attn.softmax(dim=-1)) else: attn = self.attn_drop(attn.softmax(dim=-1)) # 与Value相乘并输出投影 x = (attn @ v).transpose(1, 2).reshape(B_, N, C) x = self.proj(x) x = self.proj_drop(x) return x def extra_repr(self): return f'dim={self.dim}, window_size={self.window_size}, num_heads={self.num_heads}'

4.2 移位窗口掩码(Shifted Window Mask)的生成

Swin Transformer通过交替使用常规窗口划分和移位窗口划分来建立跨窗口连接。移位后,窗口不再是规则的,有些窗口包含来自原始特征图中不相邻区域的特征块。为了保持自注意力只在每个新窗口内部进行,需要生成一个掩码,在计算注意力时,将不同子窗口之间的注意力权重置为一个极大的负数(如-100),使其经过softmax后接近0。

def create_shift_window_mask(input_resolution, window_size, shift_size): """ 为移位窗口自注意力生成掩码。 Args: input_resolution: (H, W),输入特征图的高和宽。 window_size: (M, M),窗口大小。 shift_size: (shift_h, shift_w),移位大小,通常为 window_size // 2。 Returns: mask: 形状为 (num_windows, M*M, M*M) 的掩码张量。 其中,需要被掩蔽的位置为0,无需掩蔽的位置为 -100(或一个很大的负数)。 """ H, W = input_resolution M = window_size[0] # 确保H和W能被window_size整除(通过padding实现) Hp = int(np.ceil(H / M)) * M Wp = int(np.ceil(W / M)) * M # 1. 生成特征图的坐标图像(每个像素的坐标) img_mask = torch.zeros((1, Hp, Wp, 1)) # 通道为1,方便后续操作 h_slices = (slice(0, -M), slice(-M, -shift_size[0]), slice(-shift_size[0], None)) w_slices = (slice(0, -M), slice(-M, -shift_size[1]), slice(-shift_size[1], None)) # 2. 为移位后属于不同原始窗口的区域分配不同的编号 cnt = 0 for h in h_slices: for w in w_slices: img_mask[:, h, w, :] = cnt cnt += 1 # 3. 将特征图划分为窗口 mask_windows = window_partition(img_mask, window_size) # (nW, M, M, 1) mask_windows = mask_windows.view(-1, M * M) # (nW, M*M) # 4. 计算窗口内任意两点的掩码 # 如果两点属于img_mask中的不同编号区域,则需要被掩蔽 attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2) # (nW, M*M, M*M) attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0)) return attn_mask def window_partition(x, window_size): """ 将特征图分割成不重叠的窗口。 Args: x: (B, H, W, C) window_size: (M, M) Returns: windows: (num_windows*B, M, M, C) """ B, H, W, C = x.shape x = x.view(B, H // window_size[0], window_size[0], W // window_size[1], window_size[1], C) windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size[0], window_size[1], C) return windows

核心环节解析:掩码的逻辑create_shift_window_mask函数是Swin Transformer的精华之一。它的核心思想是:先对移位后的特征图进行“染色”,将来自原始特征图不同连续区域(即移位前属于不同窗口的区域)标记为不同的编号。然后,在划分出的新窗口内,如果两个像素的“颜色”(编号)不同,说明它们在原始图像中距离很远,不应该直接计算注意力,因此需要被掩蔽。通过attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)这个广播减法操作,我们高效地得到了一个矩阵,其中不为0的位置就是需要掩蔽的位置。

5. 常见问题与排查技巧实录

在实际实现和调试Swin Transformer的2D-RPE时,我踩过不少坑。下面是一些最常见的问题及其解决方案。

5.1 维度不匹配错误

这是新手最容易遇到的问题,尤其是在整合RPE到注意力计算时。

  • 症状:运行时错误,提示shape mismatch,broadcasting errorview size is not compatible
  • 排查清单
    1. 检查relative_position_index的形状:它必须是(N, N),其中N = M*M。使用print(relative_position_index.shape)确认。
    2. 检查relative_position_bias_table的形状:它必须是(num_relative_distance, num_heads)num_relative_distance必须等于(2M-1)*(2M-1)。确保你的索引值没有超出这个范围。
    3. 检查索引操作后的形状self.relative_position_bias_table[self.relative_position_index.view(-1)]这一步会得到一个形状为(N*N, num_heads)的张量。随后的.view(N, N, -1)必须能成功还原。
    4. 检查permute和unsqueeze的维度:确保relative_position_bias在加到attn上之前,形状是(num_heads, N, N)(1, num_heads, N, N),而attn的形状是(B_, num_heads, N, N)。广播规则要求从后往前匹配维度。
  • 我的调试技巧:在forward函数的关键步骤后,插入assert语句。例如:
    # 在索引后 rp_bias_flat = self.relative_position_bias_table[self.relative_position_index.view(-1)] assert rp_bias_flat.shape == (self.window_size[0]*self.window_size[1]**2, self.num_heads), f"Error shape: {rp_bias_flat.shape}" # 在view后 rp_bias = rp_bias_flat.view(N, N, -1) assert rp_bias.shape == (N, N, self.num_heads), f"Error shape: {rp_bias.shape}"

5.2 训练不稳定或收敛慢

  • 可能原因1:相对位置偏置表初始化不当
    • 分析:偏置表是直接加到注意力对数(logits)上的。如果初始化值过大(如默认全0,但经过几层后梯度爆炸),会主导注意力分布,导致softmax饱和(梯度消失)或注意力混乱。
    • 解决方案:使用较小的标准差进行初始化。原论文和代码库常用trunc_normal_(std=.02)。对于非常深的网络或特定任务,可以尝试更小的std,如.01.005
  • 可能原因2:与LayerNorm或残差连接的协同问题
    • 分析:Swin Block通常是“LN -> Attention -> Add -> LN -> MLP -> Add”。如果注意力模块的输出幅度与残差路径的幅度不匹配,可能导致训练不稳定。
    • 解决方案:检查并确保注意力输出投影层self.proj的权重初始化是合适的(如使用Xavier或Kaiming初始化)。同时,可以监控注意力模块前后张量的范数(norm)。
  • 实操心得:在训练初期,使用torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)进行梯度裁剪是一个稳定训练的好习惯,可以有效防止因RPE或其他参数导致的梯度爆炸。

5.3 处理可变分辨率输入或窗口大小

  • 问题:训练时使用固定的窗口大小(如7),但推理时想用不同的窗口大小或输入不同分辨率的图像。
  • 解决方案:Swin的2D-RPE本身不支持直接改变窗口大小,因为relative_position_indexrelative_position_bias_table是基于训练时的M构建的。
    • 方法A(推荐,微调):如果新窗口大小M'小于等于训练时的M,可以采用截取插值相结合的方式。relative_position_index需要重新计算。对于relative_position_bias_table,由于它本质上是一个关于(Δx, Δy)的离散函数,我们可以通过双线性插值,将原表( (2M-1)^2, num_heads)插值到新的大小( (2M'-1)^2, num_heads)。PyTorch的F.interpolate函数可以用于此目的,但需要小心处理维度。
    • 方法B(重新训练):如果分辨率或窗口大小变化很大,最稳妥的方法是使用新设置重新训练或微调模型。可以在预训练权重的基础上,初始化新的relative_position_bias_table(其他权重复用),然后用新数据快速微调。
  • 代码片段示例(方法A的插值思路)
    def adapt_relative_bias_table(old_table, old_window_size, new_window_size): """ old_table: ( (2*M_old-1)**2, num_heads) 通过插值适应新的窗口大小。 """ M_old = old_window_size M_new = new_window_size old_side = 2 * M_old - 1 new_side = 2 * M_new - 1 # 将表视为一个“图像”,尺寸为 (old_side, old_side, num_heads) old_table_2d = old_table.view(old_side, old_side, -1).permute(2, 0, 1).unsqueeze(0) # (1, num_heads, old_side, old_side) # 使用双线性插值缩放到新尺寸 new_table_2d = F.interpolate(old_table_2d, size=(new_side, new_side), mode='bilinear', align_corners=False) new_table = new_table_2d.squeeze(0).permute(1, 2, 0).reshape(-1, new_table_2d.shape[1]) return new_table

5.4 显存占用过高

  • 问题:Swin Transformer的显存占用主要来自注意力矩阵attn,其形状为(B_, num_heads, N, N)。当窗口大小M较大或批次较大时,N=M*M会平方级增长。
  • 优化技巧
    1. 使用Flash Attention:如果你的PyTorch版本和硬件支持,使用torch.nn.functional.scaled_dot_product_attention(PyTorch 2.0+)可以大幅降低显存占用并加速计算。它使用了融合内核和更高效的内存访问模式。你需要将计算attn = (q @ k.transpose(-2, -1)) * self.scale以及加偏置、softmax等步骤替换为这个函数调用。注意,你需要将相对位置偏置B作为attn_mask参数传入(但需注意符号,标准mask是加一个很大的负数,而RPE偏置是可学习的)。
    2. 梯度检查点(Gradient Checkpointing):对于非常深的Swin模型(如Swin-L, SwinV2-G),可以在反向传播时重新计算中间激活值,以时间换空间。可以使用torch.utils.checkpoint.checkpoint
    3. 混合精度训练(AMP):使用自动混合精度训练,将大部分计算转换为FP16,可以有效减少显存占用并可能加快训练速度。但要注意,位置偏置表等小参数最好保持在FP32以保证精度。

5.5 可视化理解RPE学到了什么

理解模型学到了什么对于调试和信任模型很重要。

  • 方法:取出训练好的模型中某一层、某一个注意力头的relative_position_bias_table。将其重塑为(2M-1, 2M-1)的二维矩阵。然后,将这个矩阵可视化为热力图。
  • 预期结果:你可能会观察到一些模式。例如:
    • 局部性:靠近中心(Δx=0, Δy=0)的位置通常有较高的正偏置,这意味着模型倾向于关注自身或非常近的邻居。
    • 方向性:水平或垂直方向上的偏置模式可能不同,这可能对应着学习到的水平或垂直边缘偏好。
    • 对称性:由于相对位置(Δx, Δy)(-Δx, -Δy)在注意力计算中是对称的(Q对K和K对Q),学到的偏置表可能近似中心对称。如果不是,可能是因为每个注意力头独立学习,打破了这种对称性。
  • 代码示例
    import matplotlib.pyplot as plt import seaborn as sns # 假设 model 是训练好的Swin Transformer # 获取第一个stage中第一个block的注意力模块的RPE表 rpe_table = model.layers[0].blocks[0].attn.relative_position_bias_table.data num_heads = rpe_table.shape[1] M = 7 # 假设窗口大小是7 side = 2 * M - 1 # 可视化第一个头 head_idx = 0 rpe_2d = rpe_table[:, head_idx].view(side, side).cpu().numpy() plt.figure(figsize=(8,6)) sns.heatmap(rpe_2d, center=0, cmap='RdBu_r', square=True) plt.title(f'Relative Position Bias (Head {head_idx})') plt.xlabel('Δx (shifted)') plt.ylabel('Δy (shifted)') plt.xticks(range(0, side, 2), labels=range(-(M-1), M, 2)) plt.yticks(range(0, side, 2), labels=range(-(M-1), M, 2)) plt.show()

通过这样的可视化,你可以直观地验证RPE是否在按预期工作,并为模型的可解释性分析提供依据。

返回列表