相对位置编码深度实战:从绝对位置嵌入、T5 相对注意力、ALiBi 到 DeBERTa 解耦位置表示
自 Attention 诞生起,位置信息的注入方式就决定了模型能否理解"顺序"。本文从绝对位置嵌入的失败模式出发,拆解相对位置编码的第一性原理,逐一实现 T5 相对注意力分桶、ALiBi 线性偏置与 DeBERTa 解耦位置表示,并给出长上下文外推、生产陷阱与可落地的 PyTorch 代码。
一、为什么位置信息不能是"事后添加"的装饰
Transformer 的自注意力对 token 顺序是置换不变的:把输入序列打乱,注意力矩阵只会在行/列上跟着重排,输出集合不变。这意味着如果不显式注入顺序信号,模型看到 "猫 吃 鱼" 与 "鱼 吃 猫" 会得到完全相同的语义表征——这在实际任务里是灾难。
位置编码的演化路径,本质上是在回答两个问题:编码什么(绝对位置 vs 相对距离),以及在哪里注入(加到 embedding 上 vs 直接进注意力分数)。绝对位置嵌入是"先编码再加",相对位置编码是"让注意力分数自己感知距离"。
二、绝对位置嵌入:两种范式与它们的天花板
2.1 可学习绝对位置嵌入(Learned Absolute)
BERT 的做法最简单:维护一张形状为 [max_len, d_model] 的查找表,第 i 个 token 直接取出第 i 行加到词嵌入上。
import torch
import torch.nn as nn
class LearnedAbsolutePositionalEmbedding(nn.Module):
def __init__(self, d_model: int, max_len: int = 512):
super().__init__()
self.weight = nn.Embedding(max_len, d_model)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [B, T, d_model]
pos = torch.arange(x.size(1), device=x.device).unsqueeze(0) # [1, T]
return x + self.weight(pos) # 直接广播相加
它的第一个硬伤是外推性为零:推理时长度一旦超过 max_len,查找表越界,模型直接失效。第二个更隐蔽的缺陷在下面。
2.2 正弦绝对位置(Sinusoidal)
原版 Transformer 用固定频率的 sin/cos 生成位置向量,好处是长度外推能力稍好、且不同位置可线性组合逼近。
def sinusoidal_embedding(t: int, d_model: int, device="cpu"):
pos = torch.arange(t, device=device).unsqueeze(1) # [t, 1]
div = torch.exp(torch.arange(0, d_model, 2) * (-torch.log(torch.tensor(10000.0)) / d_model))
pe = torch.zeros(t, d_model)
pe[:, 0::2] = torch.sin(pos * div)
pe[:, 1::2] = torch.cos(pos * div)
return pe # [t, d_model]
2.3 绝对位置的系统性局限
把"绝对位置 i"作为固定偏移项加进 q_i、k_i 后,注意力分数里会多出四项交叉项,其中 q_i · p_j(查询位置 dotted with 键位置)显式依赖绝对坐标而非相对距离 j - i。这带来两个后果:
- 泛化脆弱:训练时见过的绝对位置组合,推理换一个偏移就失真。
- 距离单调性不成立:位置 1↔100 与位置 50↔51 在"绝对位置空间"里没有任何距离约束,模型要自己从数据里学出"相邻更相关",既低效又不稳定。
直觉结论:下游任务关心的是"两个 token 隔了多远",而不是"它们各自在第几个位置"。相对位置编码正是把这一先验直接焊进注意力公式。
| 维度 | 可学习绝对 | 正弦绝对 | 相对位置编码 |
|---|---|---|---|
| 长度外推 | 无(越界失效) | 弱 | 取决于方案(ALiBi 强) |
| 距离感知 | 间接学习 | 间接学习 | 直接建模 |
| 参数/计算 | 一张表 | 无参数 | 偏置表或解析项 |
| 典型代表 | BERT | 原始 Transformer | T5 / ALiBi / DeBERTa |
三、相对位置编码的第一性原理
相对位置编码的核心思想(Shaw et al., 2018)是:不再把位置信息塞进 q/k 向量,而是让它直接修正注意力分数。
最干净的起点是带边偏置(edge bias)的注意力:
score(i, j) = (q_i · k_j) / sqrt(d) + a_{clip(j - i, -k, k)}
其中 a 是一个可学习的相对位置偏置表,下标是裁剪后的相对距离 j - i。i 是查询位置,j 是键位置,当 j > i 时相对距离为正(键在查询之后),反之为负(键在查询之前)。
这套形式有三个关键性质,请刻进脑子里:
- 平移等变:只依赖差值
j - i,与绝对起点无关,天然具备平移不变性。 - 可裁剪:距离过远时裁剪到
±k,远距信息用同一个"远端桶"统一表示,既省参数又防过拟合。 - 双向对称可控:可以通过对称共享
a_m = a_{-m}或解耦前后向,分别适配双向(编码)与单向(自回归解码)场景。
四、T5 相对注意力:分桶 + 缩放
T5 没有用绝对位置嵌入,而是把相对位置做成分桶(bucketing)的相对偏置,并额外做一个 sqrt(3/d_k) 的缩放,把偏置量级拉到与 q·k/sqrt(d_k) 同一数量级,防止偏置项在训练初期压垮主注意力。
4.1 相对位置分桶
T5 的关键技巧:近处用精确距离,远处用对数分桶,把可能上千的相对距离压缩到固定数量(默认 32)的桶里。
import math
def _relative_position_bucket(relative_position, bidirectional=True, num_buckets=32, max_distance=128):
# relative_position: 形状任意,值为 j - i
ret = 0
n = torch.abs(relative_position)
if bidirectional:
num_buckets //= 2
ret += (relative_position > 0).to(torch.long) * num_buckets
n = torch.where(relative_position > 0, n, num_buckets - n)
else:
n = num_buckets - 1 - n # 单向:越靠后桶越小
# 近处精确、远处对数
max_exact = num_buckets // 2
is_small = n < max_exact
val = num_buckets - max_exact + (
torch.log(n.float() / max_exact) / math.log(max_distance / max_exact) * (num_buckets - max_exact)
).long().clamp(0, num_buckets - max_exact - 1)
ret += torch.where(is_small, n, val)
return ret
4.2 注入注意力分数
def t5_relative_bias(rel_pos, bias_weight, num_buckets=32, max_distance=128):
# bias_weight: [num_heads, num_buckets],可学习
bucket = _relative_position_bucket(rel_pos, num_buckets=num_buckets, max_distance=max_distance)
# 每个 head 独立取偏置,形状 [H, T, T]
return bias_weight[:, bucket] # bias_weight 已按 [H, num_buckets] 索引展开
在注意力计算里:
scores = q @ k.transpose(-1, -2) / math.sqrt(d_k) # [H, T, T]
scores = scores + t5_relative_bias(rel_pos, bias_weight) # 直接加偏置
attn = torch.softmax(scores, dim=-1)
T5 的做法之所以稳,是因为偏置与内容解耦:内容注意力负责"语义匹配",相对偏置负责"距离先验",二者线性相加,梯度互不污染。
五、ALiBi:没有位置嵌入的位置编码
ALiBi(Attention with Linear Biases)走了一条更激进的路:完全不往输入加任何位置信号,而是给每个注意力头一条斜率为负的线性偏置线,距离越远惩罚越大。
bias(i, j) = -m * (j - i) # m 为每个 head 专属的斜率
score(i, j) = q_i · k_j / sqrt(d) - m * (j - i)
5.1 斜率 m 的分配
为了让不同 head 关注不同距离,斜率 m 按几何级数分配。原论文用 2 的负幂次再分档:
import numpy as np
def alibi_slopes(num_heads: int):
# 经典 8 头划分:从 2^(-1/2) 起的几何级数
def get_slopes(n):
if n <= 0:
return []
if n <= 2:
return [2 ** (-i) for i in range(1, n + 1)] # 1/2, 1/4
# 奇数档 + 偶数档插值
closest = 2 ** math.ceil(math.log2(n))
slopes = get_slopes(closest)
return slopes[0::2][:n]
return get_slopes(num_heads)
5.2 构造偏置矩阵
def build_alibi_bias(num_heads, t, device="cpu"):
m = torch.tensor(alibi_slopes(num_heads), device=device) # [H]
# 相对距离矩阵:行=query i, 列=key j, 值=j-i
rel = torch.arange(t, device=device)[None, :] - torch.arange(t, device=device)[:, None] # [T, T]
# bias[h, i, j] = -m[h] * (j - i)
bias = -m.view(-1, 1, 1) * rel.unsqueeze(0) # [H, T, T]
return bias
ALiBi 的杀手锏是长度外推:因为它只用相对距离做线性惩罚,推理时即便 T 超过训练长度,偏置矩阵照样能生成,且"远距离更不受关注"的先验在更长序列上依然成立。实测中 ALiBi 模型在 2×~10× 训练长度上仍可保持合理 perplexity,这使其成为长上下文训练与推理部署的性价比方案。
六、DeBERTa 解耦位置:内容与位置两路
DeBERTa 提出"解耦注意力"(Disentangled Attention):传统注意力把内容与位置耦合进同一个 q/k;DeBERTa 把 query/key 拆成"内容向量"和"位置向量"两路,注意力分数由四项组成:
A_{i,j} = <H^c_i, H^c_j> # 内容-内容
+ <H^c_i, P^v_j> # 内容-相对位置(键侧)
+ <P^h_i, H^c_j> # 相对位置(查询侧)- 内容
+ <P^h_i, P^v_j> # 位置-位置
实现上是把相对位置偏置换成"相对位置向量",并与内容向量分别做点积后再相加:
def disentangled_attention(content_q, content_k, pos_q, pos_k, d_k):
# 全部 [H, T, d]
ac = content_q @ content_k.transpose(-1, -2) / math.sqrt(d_k) # 内容-内容
ar = content_q @ pos_k.transpose(-1, -2) / math.sqrt(d_k) # 内容-相对位置
rp = pos_q @ content_k.transpose(-1, -2) / math.sqrt(d_k) # 相对位置-内容
scores = ac + ar + rp # 省略位置-位置项以突出重点
return torch.softmax(scores, dim=-1)
DeBERTa 的洞察是:位置信息应作为"修饰"而非"主体"。内容向量负责语义,相对位置向量只刻画 i、j 之间的距离关系,二者解耦后模型对位置扰动的鲁棒性显著提升,这也是它在 GLUE 上反超 BERT/RoBERTa 的关键之一。
七、相对 PE 与 RoPE 的关系
旋转位置编码(RoPE)常被拿来与相对 PE 并列,但二者机制不同:RoPE 把位置信息编码进 q、k 向量本身(通过旋转矩阵让 q_i·k_j 只依赖 i-j),属于"把相对性焊进向量";而本文的 T5/ALiBi/DeBERTa 是"在分数层面直接加相对偏置",属于"把相对性焊进分数"。
二者不是替代关系,而是互补:RoPE 让内积天然携带相对距离,相对偏置表则显式可学习、可解释。实践中大模型(如 LLaMA 系)偏好 RoPE 以获得更好的外推与硬件友好性;而 T5/ALiBi 在Encoder、长上下文推理这类场景仍有极强生命力。
八、工程落地:一个可复用的相对偏置模块
把上述方案收敛成一个生产可用的 RelativeBias 模块,支持 T5 分桶偏置与 ALiBi 两种模式:
class RelativeBias(nn.Module):
def __init__(self, num_heads, mode="t5", num_buckets=32, max_distance=128, alibi_t=512):
super().__init__()
self.num_heads = num_heads
self.mode = mode
if mode == "t5":
self.bias = nn.Embedding(num_buckets, num_heads)
else: # alibi
self.alibi_t = alibi_t
m = torch.tensor(alibi_slopes(num_heads))
self.register_buffer("m", m)
def forward(self, t):
if self.mode == "t5":
# 构造相对位置矩阵并分桶取偏置
rel = torch.arange(t)[:, None] - torch.arange(t)[None, :] # [T, T]
bucket = _relative_position_bucket(rel, num_buckets=self.bias.num_embeddings)
b = self.bias(bucket) # [T, T, H]
return b.permute(2, 0, 1) * math.sqrt(3 / 64.0) # [H, T, T],含 T5 缩放
else:
rel = torch.arange(t)[None, :] - torch.arange(t)[:, None]
return -self.m.view(-1, 1, 1) * rel.unsqueeze(0) # [H, T, T]
# 用法
bias = RelativeBias(num_heads=12, mode="t5")
scores = q @ k.transpose(-1, -2) / math.sqrt(d_k)
scores = scores + bias(t=q.size(-2)) # 直接广播加偏置
注意 T5 的 sqrt(3/d_k) 缩放不要漏——它保证偏置项与主体注意力同量级,否则训练会极不稳定。
九、长上下文外推的代价与 NTK-aware 提示
相对位置偏置虽然天然利于外推,但仍有两个现实约束:
- ALiBi 斜率固定:外推到极长序列时,远端偏置可能过负导致注意力几乎全集中在近端。可在推理时按比例缩放
m,等价于"压缩"位置尺度来换取更长有效上下文。 - T5 桶上限:
max_distance之外的距离都被压进同一个远端桶,超过该距离后模型对"谁更远"失去分辨力。增大max_distance或改用对数桶是常见修复。
NTK-aware 思路(源自 RoPE 社区)也可迁移到相对偏置:通过非线性地重标定距离刻度,让模型在未见过的更长序列上保留相对距离的对数结构,而非简单线性拉伸。
十、生产陷阱清单
| 陷阱 | 现象 | 修复 |
|---|---|---|
漏掉 T5 缩放 sqrt(3/d_k) |
训练 loss 不降或震荡 | 偏置乘缩放系数后再相加 |
| ALiBi 斜率全头相同 | 所有 head 关注同样距离,表达退化 | 用几何级数分档分配 m |
推理长度超过 max_distance |
远端位置全塌进同一桶 | 增大 max_distance 或换对数桶 |
| 单向/双向桶混淆 | 自回归生成质量骤降 | 解码场景用单向分桶(不共享前后向) |
| 偏置未随 batch 设备移动 | CUDA 报错或静默错误 | 用 register_buffer 保证设备一致 |
| 与绝对位置嵌入叠加 | 双重位置信号互相干扰 | 相对 PE 方案应移除输入级位置嵌入 |
结语
相对位置编码不是"又一个 trick",而是把"距离即先验"这一朴素直觉,从数据里硬学变成结构上强约束。T5 用可学习分桶偏置换来稳定与可解释,ALiBi 用零参数线性偏置换来惊艳的外推,DeBERTa 用解耦双路换来鲁棒——三条路殊途同归:让注意力分数自己"看见"距离。当你下一次为长上下文、外推性或推理成本头疼时,不妨回到这一层,往往比堆参数更划算。

发表评论 取消回复