深度学习进阶(二十三)偏置型 RPE
深度学习进阶(二十三)偏置型 RPE
在Transformer架构中,位置编码(Positional Encoding)是模型理解序列顺序的关键。经典的绝对位置编码(如Sinusoidal)和相对位置编码(如T5的Relative Bias)各有优劣。而偏置型RPE(Relative Positional Encoding with Bias) 是一种将相对位置信息以可学习偏置形式注入注意力机制的方法,它兼顾了灵活性与计算效率。本文将从基础概念出发,逐步深入其原理与实现。## 一、为什么需要偏置型RPE?在标准多头注意力中,注意力分数计算为:[\text{Score}(Q,K) = \frac{QK^T}{\sqrt{d_k}}]这里完全没有位置信息。如果直接加上绝对位置编码(如Transformer原始做法),模型只能感知绝对位置,难以泛化到更长的序列。而相对位置编码(RPE)能显式建模“两个token之间的相对距离”,更符合语言规律。偏置型RPE的核心思想是:在注意力分数上添加一个可学习的偏置矩阵,该矩阵的每个元素只依赖于查询和键之间的相对位置。这样既保留了相对位置信息,又避免了复杂的向量拼接或变换,计算开销极小。## 二、基础实现:手动计算偏置我们先从最简单的情况开始:假设序列长度为 ( L ),我们定义一个形状为 ( (2L-1, num_heads) ) 的可学习参数 bias,其中索引 ( k ) 对应相对距离 ( k - (L-1) )。在计算注意力分数时,为每个头提取对应的偏置值并加到分数上。下面是一个简化的PyTorch实现:pythonimport torchimport torch.nn as nnimport mathclass SimpleBiasRPE(nn.Module): def __init__(self, num_heads, max_len=512): super().__init__() self.num_heads = num_heads # 可学习偏置,形状: (2*max_len-1, num_heads) # 索引 i 对应相对距离 i - (max_len-1) self.bias = nn.Parameter(torch.zeros(2 * max_len - 1, num_heads)) nn.init.uniform_(self.bias, -0.1, 0.1) def forward(self, attn_scores): """ attn_scores: (batch, heads, L, L) 已经除以 sqrt(d_k) 的注意力分数 返回加上偏置后的分数 """ B, H, L, _ = attn_scores.shape # 生成相对位置索引矩阵 (L, L) # 例如 L=4, 索引矩阵为: # [[0,1,2,3], # [-1,0,1,2], # [-2,-1,0,1], # [-3,-2,-1,0]] row = torch.arange(L, device=attn_scores.device).unsqueeze(1) # (L,1) col = torch.arange(L, device=attn_scores.device).unsqueeze(0) # (1,L) rel_pos = col - row # (L, L) 注意这里的顺序: col - row 得到 key相对于query的偏移 # 映射到 [0, 2L-2] 范围 rel_pos_idx = rel_pos + (L - 1) # 偏移量,使得索引从0开始 # 提取偏置: (L, L, H) bias_val = self.bias[rel_pos_idx] # (L, L, H) # 转置为 (H, L, L) 并加到注意力分数上 bias_val = bias_val.permute(2, 0, 1).unsqueeze(0) # (1, H, L, L) return attn_scores + bias_val# 测试L = 4H = 2model = SimpleBiasRPE(num_heads=H, max_len=L)attn = torch.randn(2, H, L, L) # (batch, heads, L, L)out = model(attn)print("输入形状:", attn.shape)print("输出形状:", out.shape)代码解释:- rel_pos = col - row 计算了每个位置对 (i, j) 的相对距离 j - i,即键相对于查询的位置。- 将相对距离加上 L-1 映射到非负索引,用于从 bias 参数表中查找。- self.bias[rel_pos_idx] 利用张量索引,得到形状 (L, L, H) 的偏置值。- 最后加到注意力分数上,实现偏置注入。## 三、进阶:多头可分离偏置与缩放上述实现中,偏置对所有查询-键对共享同一个参数表,但实际中不同位置关系的重要性可能不同。更高级的做法是将偏置分解为头相关的可学习向量,甚至加入可学习的缩放因子。另一种常见变体是在偏置上乘以一个可学习的标量,或者将偏置分为两部分:一部分用于“相对距离”本身,另一部分用于“方向”(前向或后向)。以下是一个更贴近实际应用的实现,它支持不同头有不同的偏置模式,并允许偏置与注意力分数进行逐元素相乘后再相加(类似门控机制)。pythonimport torchimport torch.nn as nnclass GatedBiasRPE(nn.Module): def __init__(self, num_heads, max_len=512, scale=True): super().__init__() self.num_heads = num_heads self.max_len = max_len self.scale = scale # 可学习偏置,形状: (2*max_len-1, num_heads) self.bias = nn.Parameter(torch.zeros(2*max_len-1, num_heads)) # 可学习门控标量,每个头一个 self.gate = nn.Parameter(torch.ones(num_heads)) # 可学习缩放因子(可选) if scale: self.logit_scale = nn.Parameter(torch.ones(num_heads) * math.log(1.0)) else: self.register_parameter('logit_scale', None) def forward(self, attn_scores): B, H, L, _ = attn_scores.shape device = attn_scores.device # 构建相对位置索引矩阵 row = torch.arange(L, device=device).unsqueeze(1) col = torch.arange(L, device=device).unsqueeze(0) rel_pos = col - row # (L, L) rel_pos_idx = rel_pos + (self.max_len - 1) # 确保索引在 [0, 2*max_len-2] # 提取偏置: (L, L, H) bias_val = self.bias[rel_pos_idx] # (L, L, H) bias_val = bias_val.permute(2, 0, 1) # (H, L, L) # 应用门控 gate = torch.sigmoid(self.gate).view(H, 1, 1) # (H,1,1) bias_val = bias_val * gate # 如果启用缩放,则对偏置进行缩放 if self.scale: scale = self.logit_scale.exp().view(H, 1, 1) # (H,1,1) bias_val = bias_val * scale # 扩展到batch维度 bias_val = bias_val.unsqueeze(0) # (1, H, L, L) return attn_scores + bias_val# 测试L = 8H = 4model = GatedBiasRPE(num_heads=H, max_len=L)attn = torch.randn(3, H, L, L)out = model(attn)print("输出形状:", out.shape)print("门控参数:", torch.sigmoid(model.gate).data)代码解释:- gate 通过sigmoid函数将值压缩到(0,1),控制每个头对偏置的敏感度。- logit_scale 通过指数函数得到正值,用于调整偏置的幅度。- 这种设计让模型能够自适应地决定位置信息的重要程度,而不仅仅是加一个固定偏置。## 四、在完整Transformer中的应用在实际Transformer中,偏置型RPE通常嵌入到多头注意力模块中。以下是集成到完整注意力层的示例(省略了投影权重初始化):pythonclass RPETransformerBlock(nn.Module): def __init__(self, d_model, num_heads, max_len=512, dropout=0.1): super().__init__() self.d_model = d_model self.num_heads = num_heads self.head_dim = d_model // num_heads self.wq = nn.Linear(d_model, d_model) self.wk = nn.Linear(d_model, d_model) self.wv = nn.Linear(d_model, d_model) self.wo = nn.Linear(d_model, d_model) self.rpe = GatedBiasRPE(num_heads, max_len) self.dropout = nn.Dropout(dropout) self.norm1 = nn.LayerNorm(d_model) def forward(self, x): B, L, _ = x.shape Q = self.wq(x).view(B, L, self.num_heads, self.head_dim).transpose(1, 2) K = self.wk(x).view(B, L, self.num_heads, self.head_dim).transpose(1, 2) V = self.wv(x).view(B, L, self.num_heads, self.head_dim).transpose(1, 2) # 计算注意力分数 scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim) # 注入偏置 scores = self.rpe(scores) attn_weights = torch.softmax(scores, dim=-1) attn_weights = self.dropout(attn_weights) out = torch.matmul(attn_weights, V) # (B, H, L, head_dim) out = out.transpose(1, 2).contiguous().view(B, L, -1) out = self.wo(out) out = self.norm1(out + x) # 残差连接 return out# 测试完整块d_model = 32num_heads = 4block = RPETransformerBlock(d_model, num_heads, max_len=16)x = torch.randn(2, 16, d_model)y = block(x)print("输入:", x.shape, "输出:", y.shape)工作流程:1. 将输入线性投影到Q、K、V。2. 计算原始注意力分数。3. 调用GatedBiasRPE为分数添加相对位置偏置。4. 进行Softmax和加权求和。5. 输出投影并残差连接。## 五、总结偏置型RPE是一种高效且强大的位置编码方式。它通过可学习的偏置表直接注入相对位置信息,避免了复杂的相对位置向量拼接,在保持模型表达能力的同时显著减少了计算量。本文从最简单的偏置加法开始,逐步引入了门控机制和可学习缩放,最后展示了如何嵌入到完整的Transformer块中。在实际应用中,这种编码方式在机器翻译、文本生成等任务上已被证明优于绝对位置编码,尤其擅长处理长序列和长度外推问题。关键要点回顾:- 相对位置:偏置值仅依赖于查询和键之间的相对距离,而非绝对位置。- 可学习性:偏置表是端到端训练得到的,无需手动设计。- 灵活性:通过门控和缩放,模型可以自适应地控制位置信息的强度。- 高效性:只需一次索引查表操作,即可为所有注意力头添加偏置。掌握偏置型RPE,你就能在Transformer架构中灵活地注入位置信息,为下游任务带来更优的性能。
更多推荐

所有评论(0)