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

📅 2026/8/10 3:49:26
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。对于图像任务将二维坐标展平为一维后使用。其根本缺陷在于分辨率敏感训练时固定了序列长度即图像块数量。如果推理时图像分辨率改变序列长度变化预训练的APE矩阵无法直接使用。虽然可以通过插值来适应新分辨率但这会引入误差并非原生支持。缺乏平移不变性计算机视觉的许多任务如物体检测、分割要求模型对物体的平移具有不变性。APE明确编码了绝对位置与这一先验略有冲突。不符合视觉直觉人类识别物体更多依赖的是物体部件之间的相对关系眼睛在鼻子上面轮子在车身下面而非其在图像中的绝对坐标。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中i和j是同一个窗口内的两个图像块。2.3 Swin 2D-RPE 的具体方案选型Swin Transformer的作者们做出了几个关键且巧妙的设计选择参数化与共享相对位置偏置B被设计为一个可学习的参数而不是通过正弦余弦函数生成。这意味着模型可以从数据中学习到哪种相对位置关系应该被加强或减弱。更重要的是这个偏置参数在所有窗口、所有层、所有头之间共享。这是一个很强的归纳偏置假设“相同的相对位置关系在任何地方、任何语义层次上都具有相似的重要性”。实践证明这个假设非常有效且极大减少了参数量。离散化的二维相对坐标这是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我们就可以从这个表中查取出对应所有注意力头的偏置值。与注意力头的解耦偏置参数表最后一维是num_heads。这意味着每个注意力头都有自己独立的一套相对位置偏置。这赋予了模型更大的灵活性有的头可能更关注局部小位移关系有的头可能更关注窗口内较远的关系模型可以自行学习。注意这里有一个极易混淆的点。许多初学者会认为relative_position_bias_table的形状是(num_heads, (2M-1)*(2M-1))。在代码实现中两种维度顺序都有可能出现取决于后续矩阵加法的便利性。关键是要理解其物理意义它是一个查询表为每一种可能的二维相对位置关系存储了所有注意力头对应的偏置值。3. 核心细节解析与实操要点3.1 相对位置索引的生成代码级详解理解索引的生成是复现RPE的第一步。下面我们以M7窗口大小7x7为例拆解这个过程。import torch def generate_relative_position_index(window_size7): 生成用于索引 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。这个矩阵就是后续查询偏置表的“地图”。实操要点1permute与contiguous的重要性在步骤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, maskNone): 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实操要点2register_buffer的妙用relative_position_index是一个固定的、根据窗口大小计算出来的整数张量它不参与训练。使用self.register_buffer(name, tensor)将其注册为模块的缓冲区。这样做的好处是它会被自动转移到正确的设备GPU/CPU上与模型参数同步。它会被包含在模型的state_dict中因此保存和加载模型时这个预计算的索引也会被保存和加载保证一致性。它不参与梯度计算节省了显存和计算量。实操要点3视图view与维度变换的陷阱在forward函数中从偏置表查取出数据后有一系列view和permute操作。这里的顺序和维度必须非常小心。一个常见的错误是维度不匹配导致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_heads8, qkv_biasTrue, attn_drop0., proj_drop0.): 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], indexingij)) # (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, biasqkv_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, maskNone): 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 fdim{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 error或view size is not compatible。排查清单检查relative_position_index的形状它必须是(N, N)其中N M*M。使用print(relative_position_index.shape)确认。检查relative_position_bias_table的形状它必须是(num_relative_distance, num_heads)。num_relative_distance必须等于(2M-1)*(2M-1)。确保你的索引值没有超出这个范围。检查索引操作后的形状self.relative_position_bias_table[self.relative_position_index.view(-1)]这一步会得到一个形状为(N*N, num_heads)的张量。随后的.view(N, N, -1)必须能成功还原。检查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), fError shape: {rp_bias_flat.shape} # 在view后 rp_bias rp_bias_flat.view(N, N, -1) assert rp_bias.shape (N, N, self.num_heads), fError 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_norm1.0)进行梯度裁剪是一个稳定训练的好习惯可以有效防止因RPE或其他参数导致的梯度爆炸。5.3 处理可变分辨率输入或窗口大小问题训练时使用固定的窗口大小如7但推理时想用不同的窗口大小或输入不同分辨率的图像。解决方案Swin的2D-RPE本身不支持直接改变窗口大小因为relative_position_index和relative_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), modebilinear, align_cornersFalse) new_table new_table_2d.squeeze(0).permute(1, 2, 0).reshape(-1, new_table_2d.shape[1]) return new_table5.4 显存占用过高问题Swin Transformer的显存占用主要来自注意力矩阵attn其形状为(B_, num_heads, N, N)。当窗口大小M较大或批次较大时NM*M会平方级增长。优化技巧使用Flash Attention如果你的PyTorch版本和硬件支持使用torch.nn.functional.scaled_dot_product_attentionPyTorch 2.0可以大幅降低显存占用并加速计算。它使用了融合内核和更高效的内存访问模式。你需要将计算attn (q k.transpose(-2, -1)) * self.scale以及加偏置、softmax等步骤替换为这个函数调用。注意你需要将相对位置偏置B作为attn_mask参数传入但需注意符号标准mask是加一个很大的负数而RPE偏置是可学习的。梯度检查点Gradient Checkpointing对于非常深的Swin模型如Swin-L SwinV2-G可以在反向传播时重新计算中间激活值以时间换空间。可以使用torch.utils.checkpoint.checkpoint。混合精度训练AMP使用自动混合精度训练将大部分计算转换为FP16可以有效减少显存占用并可能加快训练速度。但要注意位置偏置表等小参数最好保持在FP32以保证精度。5.5 可视化理解RPE学到了什么理解模型学到了什么对于调试和信任模型很重要。方法取出训练好的模型中某一层、某一个注意力头的relative_position_bias_table。将其重塑为(2M-1, 2M-1)的二维矩阵。然后将这个矩阵可视化为热力图。预期结果你可能会观察到一些模式。例如局部性靠近中心(Δx0, Δy0)的位置通常有较高的正偏置这意味着模型倾向于关注自身或非常近的邻居。方向性水平或垂直方向上的偏置模式可能不同这可能对应着学习到的水平或垂直边缘偏好。对称性由于相对位置(Δ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, center0, cmapRdBu_r, squareTrue) plt.title(fRelative Position Bias (Head {head_idx})) plt.xlabel(Δx (shifted)) plt.ylabel(Δy (shifted)) plt.xticks(range(0, side, 2), labelsrange(-(M-1), M, 2)) plt.yticks(range(0, side, 2), labelsrange(-(M-1), M, 2)) plt.show()通过这样的可视化你可以直观地验证RPE是否在按预期工作并为模型的可解释性分析提供依据。