资讯动态

TransformerXL相对位置编码:原理、推导与PyTorch实现

发布时间:2026/10/3 5:35:08 来源:尧图企业网站定制
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 的绝对位置向量。把括号展开会得到四项E_{x_i}^T W_q^T W_k E_{x_j}词 i 内容与词 j 内容的交互E_{x_i}^T W_q^T W_k U_j词 i 内容与词 j 绝对位置的交互U_i^T W_q^T W_k E_{x_j}词 i 绝对位置与词 j 内容的交互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. 代码实现一个自带相对位置编码的Attention3.1 代码整体结构我这里写一个精简但完整的实现聚焦在相对位置编码的注意力计算上。不包含完整的 TransformerXL 层归一化、FFN 等工程细节但核心的四个 term 都有。这个类可以直接插入一个 Transformer 层里使用把原本的 nn.MultiheadAttention 换掉就行。import torch import torch.nn as nn import math class RelPartialLearnableMultiHeadAttn(nn.Module): def __init__(self, d_model128, n_head4, d_head32, max_len512): 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, biasFalse) self.k_proj nn.Linear(d_model, n_head * d_head, biasFalse) self.v_proj nn.Linear(d_model, n_head * d_head, biasFalse) # 相对位置向量投影对应论文中的 W_k^R self.r_proj nn.Linear(d_model, n_head * d_head, biasFalse) # 输出投影 self.o_proj nn.Linear(n_head * d_head, d_model, biasFalse) # 全局 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, devicedevice) 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, maskNone): x: [batch, length, d_model] mask: [batch, length] 或 None1 表示有效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) # batch2, seq_len10, d_model128 attn_layer RelPartialLearnableMultiHeadAttn(d_model128, n_head4, d_head32, max_len512) 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但推理时输入长度超过 512rel_idx 可能超出 pos_emb 表的范围。解决办法是 clamp然后把超出部分都映射到最远距离。上面代码里已经加了 clamp但很多开源代码的早期版本没加推理长文本时会报 IndexError。经验pos_emb 表的长度应当比训练时最大长度再大一点比如训练长度 512 时表长至少 10232*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 长度改成 Lmems 长度query 长度保持 L位置索引矩阵变成 [L, Lmems]。计算时取 i 为当前片段位置j 覆盖全部长度。这样相对距离才能正确反映跨片段的位置关系。这个细节在官方源码里写得很清楚但它散落在循环逻辑里很多人第一次读源码时根本注意不到。4.4 训练长度与推理长度不一致相对位置编码天然支持长度外推但前提是位置表覆盖范围要够大。如果训练时 max_len256推理时突然输入 4096即使 clamp 了远距离的区分度会下降。建议训练时就保留一个较大的 max_len比如 1024 或 2048只让实际训练序列短一些。代价是 pos_emb 表参数略多但通常可以接受。在实际项目里我会先统计训练语料的长度分布把 max_len 设成覆盖 99% 样本的值再额外加 20% 余量。这样既不会让位置表过大也能保证推理时不会频繁触发截断。4.5 初始化导致早期训练不稳代码里我用了 0.02 标准差初始化 pos_emb 和 u/v。这个值在 d_head32 时表现正常。如果 d_head 变大比如 1280.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这三步改动环环相扣少一个都不行。希望这篇能帮你省下一些绕弯的时间。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑