☰
TransformerXL相对位置编码:原理、推导与PyTorch实现
2026/10/3 5:35:06 网站建设 项目流程

TransformerXL 这个名字,乍一听容易联想到显卡,其实它是 Transformer 在长文本建模里的一个关键改进版本。2019 年放出来的时候,正赶上各种预训练语言模型刷榜的时期,很多人只把它当成一个“能处理更长上下文”的工具来用。我最初也是冲着“能处理长文本”去读的,但真正读到相对位置编码那一节,才发现这里的改动比“长度变长”本身更有意思。今天这篇就来把这部分掰开揉碎,讲清楚相对位置编码的来龙去脉,附带一份能直接跑起来的 PyTorch 示例代码,让你看完不仅知道思路,还能自己动手改代码。

1. 先理解TransformerXL到底解决了什么

1.1 绝对位置编码在长文本里踩的坑

Transformer 最原始的设定里,每个 token 的输入是 “词向量 + 位置向量” 相加。位置向量可以是固定的正弦余弦,也可以是可学习的 embedding。模型通过这个位置向量感知“这个词出现在句子的第几位”。这个设计在单段文本上没问题,可一旦文本长度超过模型窗口,就麻烦了。

标准做法是把超长文本切分成多个片段,按顺序逐个送入模型。问题是:每个片段的内部位置编号都会从头开始。第二片段的第 1 个词,和第一片段的第 1 个词,位置编码完全一样。模型看到的是“两个位置编号相同的词”,根本不知道它们之间隔了多少个词。这就好比你把一本小说拆成十份,每份重新编页码,第三份的第 5 页和第一份的第 5 页看起来没有区别,阅读者自然就晕了。

更麻烦的是,长距离依赖信息全断了。比如一句话在第 3 个片段里出现了一个“它”,指代的对象在第 1 个片段里,标准 Transformer 在处理第 3 个片段时根本没有办法看到第 1 个片段的内容,连位置信息都没有,更谈不上建模这种跨片段的依赖关系。

1.2 片段级递归:把旧记忆接回来

TransformerXL 的核心创新之一是片段级递归机制(Segment-Level Recurrence)。思路很朴素:处理当前片段的时候,把上一个片段的隐藏状态也拼接进来,让当前层可以 attend 到之前片段的输出。就像你读书翻到第 10 页,桌上还摊着前 9 页,可以随时回看,不需要重新从头翻。

这个机制确实把有效上下文扩展到了多片段。但引入记忆之后,绝对位置编码的问题被放大了。拼接过来的旧片段隐藏状态,如果继续使用当前片段的新位置编号,那位置信息就是错的;如果使用旧片段原有的位置编号,那位置向量的范围会越来越大,后面第 1000 片段的第 1 个词,它的位置编号大到难以处理,而且绝对位置编号本身对很多任务没有意义。模型真正需要知道的,其实是“这个旧片段里的词,相对于当前片段里的词,中间隔了多远”。

1.3 换个问法:从“你在哪”变成“你离我多远”

相对位置编码的核心,就是不再关心某个词绝对处于第几个位置,而是关心两个词之间的相对距离。比如“猫”和“沙发”在句子里差了几个词、谁在谁前面,这些信息对理解语义其实比绝对位置更重要。对 attention 计算来说,query 和 key 之间只要有相对距离的表示就足够了。这正好绕开了分段和拼接带来的位置编号混乱。

同时,相对位置编码引入了一个更合理的归纳偏置:平移不变性。同样的词序关系,无论出现在文章开头还是结尾,编码都一样。自然语言中“主谓宾”的结构不会因为出现在第 50 句还是第 500 句就发生变化,所以相对位置比绝对位置更符合语言本身的规律。

2. 相对位置编码到底改了什么

2.1 从原始Attention分数说起

先回到标准 Transformer 的注意力计算。给定 query 位置 i 和 key 位置 j,原始分数可以写成这样:

score(i, j) = (W_q (E_{x_i} + U_i))^T W_k (E_{x_j} + U_j)

其中 E_{x_i} 是词 i 的内容向量,U_i 是词 i 的绝对位置向量。把括号展开,会得到四项:

  1. E_{x_i}^T W_q^T W_k E_{x_j}:词 i 内容与词 j 内容的交互
  2. E_{x_i}^T W_q^T W_k U_j:词 i 内容与词 j 绝对位置的交互
  3. U_i^T W_q^T W_k E_{x_j}:词 i 绝对位置与词 j 内容的交互
  4. U_i^T W_q^T W_k U_j:两个绝对位置的交互

问题就出在 U_i 和 U_j 上。这两个向量是绝对位置编号,一旦分段或拼接记忆,它们就变得不稳定:同一个词在第 1 段和第 5 段出现时,绝对位置编号完全不同,但模型对它的语义理解不应当因为这个变化而改变。绝对位置编码自己很难处理好这种“位置编号变化但内容不变”的情况。

2.2 TransformerXL的四项替换

TransformerXL 的处理很巧妙。它把 U_j 直接换成了 R_{i-j},也就是 i 和 j 的相对距离向量。而 U_i 那一项,整体替换成两个可学习的向量 u 和 v。论文中的公式可以写成:

score(i, j) = E_{x_i}^T W_q^T W_k E_{x_j} + E_{x_i}^T W_q^T W_k^R R_{i-j} + u^T W_k E_{x_j} + v^T W_k^R R_{i-j}

对照原式,四个 term 的改动如下:

项原始绝对位置编码TransformerXL 相对位置编码含义
(a)E_{x_i}^T W_q^T W_k E_{x_j}E_{x_i}^T W_q^T W_k E_{x_j}query 内容与 key 内容
(b)E_{x_i}^T W_q^T W_k U_jE_{x_i}^T W_q^T W_k^R R_{i-j}query 内容与相对位置
(c)U_i^T W_q^T W_k E_{x_j}u^T W_k E_{x_j}全局 query 偏置与 key 内容
(d)U_i^T W_q^T W_k U_jv^T W_k^R R_{i-j}全局 query 偏置与相对位置

注意看,U_i 在原来第 3、第 4 项里都出现了,现在被 u 和 v 替代。U_j 被换成了 R_{i-j}。整个公式里已经没有任何绝对位置向量了。

2.3 为什么第3、4项可以替换成常量u、v

这是理解相对位置编码最绕但也最关键的一步。在原公式里,U_i 是 query 的绝对位置。既然现在决定不再依赖绝对位置,那 query 这边的位置信息用什么来表示?TransformerXL 的答案是:不用位置,用偏置。

具体来说就是,对于任意一个 query 位置 i,它对 key 内容的偏好由 u 决定,对相对距离的偏好由 v 决定。u 和 v 是每个 attention head 各有一组的可学习向量,形状都是 [d_head]。不管 query 在第 1 位还是第 100 位,它对内容匹配的倾向性是全局统一的,对相对距离的倾向性也是全局统一的。

这个设计的合理性在于:自然语言中,一个词去 attend 另一个词时,起作用的往往不是“我在第几个位置”,而是“我们之间隔了多远”。比如“我昨天买了一只猫,它很可爱”,这句话里“它”和“猫”之间隔了 4 个词,不管这句话出现在文章开头还是第 1000 句,这种 4 个词的间隔关系是稳定的。用一个全局向量 v 来表达“query 对相对距离的偏好”,比每 1000 个位置各学一个 query 位置向量要稳定得多,参数也更少。

2.4 相对位置索引的截断

R_{i-j} 的取值范围是 [-(seq_len-1), seq_len-1]。训练时如果最长见过 512,推理时想处理更长文本,距离超出训练范围怎么办?TransformerXL 的做法是截断(clamp)。把相对距离限制在一个窗口内,比如最大距离是 L,那么所有 i-j 绝对值大于 L 的情况都当作 L 处理。

为什么能放心截断?因为自然语言中,两个词相隔太远时,精确的距离信息几乎没意义。比如“这本书”出现在第 2 句,“书”这个实体在第 50 句被再次提到,中间隔着上百个词,模型只需要知道“它们离得很远”就够了,不需要精确区分隔了 100 个词还是 101 个词。截断还带来了一个额外好处:模型见过的距离类别有限,不会因为序列长度变化而遇到完全没见过的距离模式。

代码实现时,通常会把相对距离偏移到非负索引,方便查表。比如 rel_idx = i - j + max_len - 1,然后 clamp 到 [0, 2 * max_len - 2] 范围内。

3. 代码实现:一个自带相对位置编码的Attention

3.1 代码整体结构

我这里写一个精简但完整的实现,聚焦在相对位置编码的注意力计算上。不包含完整的 TransformerXL 层归一化、FFN 等工程细节,但核心的四个 term 都有。这个类可以直接插入一个 Transformer 层里使用,把原本的 nn.MultiheadAttention 换掉就行。

import torch import torch.nn as nn import math class RelPartialLearnableMultiHeadAttn(nn.Module): def __init__(self, d_model=128, n_head=4, d_head=32, max_len=512): super().__init__() self.d_model = d_model self.n_head = n_head self.d_head = d_head self.max_len = max_len self.scale = 1.0 / math.sqrt(d_head) # Q / K / V 的内容投影 self.q_proj = nn.Linear(d_model, n_head * d_head, bias=False) self.k_proj = nn.Linear(d_model, n_head * d_head, bias=False) self.v_proj = nn.Linear(d_model, n_head * d_head, bias=False) # 相对位置向量投影,对应论文中的 W_k^R self.r_proj = nn.Linear(d_model, n_head * d_head, bias=False) # 输出投影 self.o_proj = nn.Linear(n_head * d_head, d_model, bias=False) # 全局 query 偏置,对应论文中的 u 和 v self.u = nn.Parameter(torch.randn(n_head, d_head) * 0.02) self.v = nn.Parameter(torch.randn(n_head, d_head) * 0.02) # 相对位置编码查找表,距离范围从 -(max_len-1) 到 (max_len-1) # 总共 2 * max_len - 1 个距离 self.pos_emb = nn.Parameter(torch.randn(2 * max_len - 1, d_model) * 0.02) def _get_rel_emb(self, length, device): """生成长度为 length 的相对距离索引矩阵,并取出对应位置向量""" pos = torch.arange(length, device=device) rel_dist = pos.unsqueeze(1) - pos.unsqueeze(0) # [L, L] rel_idx = rel_dist + self.max_len - 1 # 越界处理:超过 max_len 时截断到最远距离 rel_idx = rel_idx.clamp(0, 2 * self.max_len - 2) rel_emb = self.pos_emb[rel_idx] # [L, L, d_model] return rel_emb def forward(self, x, mask=None): """ x: [batch, length, d_model] mask: [batch, length] 或 None,1 表示有效,0 表示无效 返回: [batch, length, d_model] """ batch, length, _ = x.size() # 内容投影 q = self.q_proj(x).view(batch, length, self.n_head, self.d_head).transpose(1, 2) k = self.k_proj(x).view(batch, length, self.n_head, self.d_head).transpose(1, 2) v = self.v_proj(x).view(batch, length, self.n_head, self.d_head).transpose(1, 2) # q/k/v 均为 [batch, n_head, length, d_head] # 相对位置向量投影 rel_emb = self._get_rel_emb(length, x.device) # [L, L, d_model] rel_emb = self.r_proj(rel_emb) # [L, L, n_head * d_head] rel_emb = rel_emb.view(length, length, self.n_head, self.d_head).permute(2, 0, 1, 3) # rel_emb: [n_head, L, L, d_head] # term (a): query 内容 与 key 内容 attn_a = torch.matmul(q, k.transpose(-2, -1)) * self.scale # [batch, n_head, L, L] # term (b): query 内容 与 相对位置 # 用 einsum 对齐维度: batch, n_head, i, d 与 n_head, i, j, d attn_b = torch.einsum('bhid,hijd->bhij', q, rel_emb) * self.scale # [batch, n_head, L, L] # term (c): 全局 u 与 key 内容 # u: [n_head, d_head] attn_c = torch.einsum('hd,bhjd->bhj', self.u, k) # [batch, n_head, L] attn_c = attn_c.unsqueeze(2).expand_as(attn_a) # 广播到每个 query 位置 # term (d): 全局 v 与 相对位置 attn_d = torch.einsum('hd,hijd->hij', self.v, rel_emb) # [n_head, L, L] attn_d = attn_d.unsqueeze(0).expand_as(attn_a) # 广播到 batch attn = attn_a + attn_b + attn_c + attn_d if mask is not None: mask = mask.unsqueeze(1).unsqueeze(2) # [batch, 1, 1, L] attn = attn.masked_fill(mask == 0, -1e9) attn_prob = torch.softmax(attn, dim=-1) output = torch.matmul(attn_prob, v) # [batch, n_head, L, d_head] output = output.transpose(1, 2).contiguous().view(batch, length, -1) return self.o_proj(output)

这段代码就是相对位置编码最核心的部分。先花 30 秒理解结构:q/k/v 都只来自内容 x,相对位置向量单独做一次投影,然后拆成四个分数相加。值得注意的地方有两个:一是位置向量不是加到词向量上,而是作为“key 方位置通道”单独参与注意力计算;二是 u 和 v 与 batch 无关,所有样本共享。

3.2 为什么计算时用了两次不同的投影

代码里和位置相关的参数有两个:pos_emb 和 r_proj。pos_emb 是原始的相对位置向量表,r_proj 负责把位置向量投影到 attention 空间,相当于论文里的 W_k^R。

这样做的好处是,位置编码不是固定的、也不是简单地加到词向量里,而是以“key 方位置投影”的形式参与注意力计算。W_k^R 可以学习“对于不同相对距离,key 应当以一种怎样的方式被 query 检索”。这个设计比直接把位置向量加到词向量上更灵活,因为它允许模型对内容和位置分开建模,从参数上就避免了内容和位置相互干扰。

u 和 v 也是一样的逻辑。u 对应 query 侧对 key 内容的全局偏置,v 对应 query 侧对相对位置的全局偏置。它们的学习目标完全不同:u 要学到的是“不管 query 在哪,它更倾向 attend 内容的哪一部分”,v 要学到的是“不管 query 在哪,它更偏好多大的相对距离”。把这两个职责分开,训练更稳定。

3.3 每一项的维度推演

这一步非常容易踩坑,我把维度完整推一遍:

  • x 的输入形状是 [B, L, D]
  • q/k/v 经过投影和 reshape,形状是 [B, H, L, d_head]
  • rel_emb 经过 r_proj 和 reshape,形状是 [H, L, L, d_head]

attn_a 用 matmul 计算:q [B, H, L, Dh] 与 k.transpose [B, H, Dh, L] 得到 [B, H, L, L],这是标准的点积注意力。

attn_b 用 einsum 计算:'bhid,hijd->bhij' 中,bhid 的 i 是 query 位置,hijd 的 i 同样是 query 位置,j 是 key 位置,d 是 head 维度,对 d 求和,结果就是 [B, H, L, L]。

attn_c 中,u [H, Dh] 与 k [B, H, L, Dh] 在 d 维点积,结果 [B, H, L],代表每个 batch、每个 head、每个 key 位置的内容偏置。然后 unsqueeze 到 [B, H, 1, L],再 expand 成 [B, H, L, L],让每个 query 位置都加上同一个 key 内容偏置。

attn_d 中,v [H, Dh] 与 rel_emb [H, L, L, Dh] 点积,结果 [H, L, L],与 batch 无关,因为所有 batch 共享同一个相对位置偏置。然后 unsqueeze 到 [1, H, L, L],expand 成 [B, H, L, L]。

到这里四个 term 全部对齐,可以相加。mask 在四个 term 加总之后统一设置,这点很重要,后面会提。

3.4 验证一个小例子

用一个小 tensor 跑通前向,顺便验证输出形状:

torch.manual_seed(0) x = torch.randn(2, 10, 128) # batch=2, seq_len=10, d_model=128 attn_layer = RelPartialLearnableMultiHeadAttn(d_model=128, n_head=4, d_head=32, max_len=512) out = attn_layer(x) print(out.shape) # 期望输出: torch.Size([2, 10, 128])

能正常输出就说明维度没问题。如果想把这份代码接到现有模型里,直接把原来的多头注意力替换成这个类就行,注意传入的 x 要经过输入层归一化。

3.5 和标准Attention的差异对照

标准 MultiheadAttention 在实现时,位置信息通常是提前加到 x 里的,比如使用正弦位置编码表,attention 内部不再关心位置。相对位置编码的 attention 则把位置信息移动到 attention 计算内部,分为内容通道和位置通道。这样带来的直接变化是参数分配方式不同:标准 attention 的 W_k 同时承担内容和位置的映射,相对位置 attention 把位置映射交给了单独的 W_k^R。

另一个重要差异是长度外推能力。标准绝对位置编码遇到比训练时更长的序列,需要插值或直接截断,效果都不理想。相对位置编码因为不依赖绝对位置,只要位置表范围覆盖住,处理更长序列时能把超出训练范围的距离截断到“非常远”这一档,损失很小。这在推理长文本时优势很明显。

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

4.1 相对距离索引越界

当训练时 max_len 设为 512,但推理时输入长度超过 512,rel_idx 可能超出 pos_emb 表的范围。解决办法是 clamp,然后把超出部分都映射到最远距离。上面代码里已经加了 clamp,但很多开源代码的早期版本没加,推理长文本时会报 IndexError。

经验:pos_emb 表的长度应当比训练时最大长度再大一点,比如训练长度 512 时,表长至少 1023(2*512-1)。如果还想更保险,可以按预期最大推理长度来建表。我用 2048 长度的表训练 512 的序列,推理 4096 时也不会崩,只是远端距离区分度会下降。

4.2 mask 忘了应用在全部四项

如果你只对 attn_a 做 mask,而 attn_b、attn_c、attn_d 直接加进来,mask 位置还是有可能获得较大分数。因为加性的偏置项没有被 mask 掉。正确做法是先把四个 term 加总,再统一 mask,再 softmax。我在初版代码里就犯过这个错,当时只对 attn_a 做了 masked_fill,结果 padding 位置的 attention 分数里仍然混入了位置偏置,训练 loss 一直不掉。

顺带说一句,如果使用 PyTorch 的 scaled_dot_product_attention 接口,也要注意它默认假设你传入的 attn_mask 是作用在加总后的分数上,不能只 mask 其中一部分。

4.3 和缓存mems配合时的位置计算

TransformerXL 的 memory 机制会拼接上个片段的隐藏状态。这时当前片段的 attention 矩阵大小是 [当前长度, 当前长度+memory长度]。注意相对位置编码计算范围要覆盖到 memory 那边的距离,比如当前片段第 0 个词和最远的 memory 词之间的相对距离可能超过当前长度。如果复用上面的 _get_rel_emb 函数,只传 length 是不够的,还要传入 total_length(当前片段长度+memory长度),并让 query 位置取当前片段部分、key 位置取拼接后的全部。

代码调整方向:把 q/k/v 对应的 key 长度改成 L+mems 长度,query 长度保持 L,位置索引矩阵变成 [L, L+mems]。计算时取 i 为当前片段位置,j 覆盖全部长度。这样相对距离才能正确反映跨片段的位置关系。这个细节在官方源码里写得很清楚,但它散落在循环逻辑里,很多人第一次读源码时根本注意不到。

4.4 训练长度与推理长度不一致

相对位置编码天然支持长度外推,但前提是位置表覆盖范围要够大。如果训练时 max_len=256,推理时突然输入 4096,即使 clamp 了,远距离的区分度会下降。建议训练时就保留一个较大的 max_len(比如 1024 或 2048),只让实际训练序列短一些。代价是 pos_emb 表参数略多,但通常可以接受。

在实际项目里,我会先统计训练语料的长度分布,把 max_len 设成覆盖 99% 样本的值,再额外加 20% 余量。这样既不会让位置表过大,也能保证推理时不会频繁触发截断。

4.5 初始化导致早期训练不稳

代码里我用了 0.02 标准差初始化 pos_emb 和 u/v。这个值在 d_head=32 时表现正常。如果 d_head 变大,比如 128,0.02 可能会导致早期梯度消失,可以改成 0.02/sqrt(d_head) 之类的缩放。TransformerXL 原始实现中,参数初始化也有专门的策略,直接简单正态初始化也行,但要注意观察 loss 曲线。

我试过用 Xavier 初始化 u 和 v,效果差别不大。真正影响大的反而是 q_proj/k_proj/v_proj 的初始化,因为它们的缩放直接决定 attention 分数初始量级。如果初始量级太大,softmax 会退化成 one-hot,训练早期难以恢复。

5. 一点实操心得

最后分享几个我在项目里用相对位置编码时总结的经验,不保证放之四海皆准,但至少能帮你少踩坑。

第一,如果你只是在短文本上用 Transformer,把绝对位置换成相对位置,未必会看到明显提升。相对位置编码的优势主要体现在长文本、段落式输入或需要外推的场景。所以在小规模试水前,先确认自己的任务确实存在“位置编码难以表达”的问题。比如短文本情感分类,绝对位置编码已经够用,换成相对位置可能没区别。

第二,实现相对位置 attention 时,尽量用 einsum 而不是手写维度变换。四个 term 的维度对应关系很容易绕晕,用 einsum 一次性点积和广播,既清晰又不容易错。想快速改实验时,einsum 的改动成本也低。比如想在 attn_d 里加一个可学习的缩放系数,只需要在 einsum 结果上乘一个参数,比手写 matmul 方便得多。

第三,做消融实验比直接上完整模型更有价值。你可以把 attn_b、attn_c、attn_d 分别置零,看看每个 term 对最终效果的贡献。我之前在某份长文本分类数据上试过,去掉 attn_b(内容-位置项)后,准确率下降最明显,说明模型主要靠“内容结合相对距离”来建立位置感知。而去掉 attn_c 和 attn_d 的影响相对较小,说明全局偏置项更多是锦上添花。

第四,如果要读原论文和源码,建议从 TransformerXL 官方实现或者各类复现项目里找那个叫 RelPartialLearnableMultiHeadAttn 的类,你会发现它和我这里的代码核心逻辑一致,但工程细节更多,比如 layer_norm 位置、dropout、bias 设置等。把这些工程细节拼到本文的代码上,就是一个能真正训练的 TransformerXL 层了。

相对位置编码并不难,难在把“为什么这样改”想透。U_i 换成 u 和 v、U_j 换成 R_{i-j}、W_k 拆成 W_k 和 W_k^R,这三步改动环环相扣,少一个都不行。希望这篇能帮你省下一些绕弯的时间。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询