TransformerXL相对位置编码详解:从公式到PyTorch实现

发布时间:2026/10/3 5:35:10
TransformerXL相对位置编码详解:从公式到PyTorch实现 说实话TransformerXL 这套东西我前前后后啃了三遍源码才算真正看懂。过程很痛苦因为原始论文里那套公式写得极其抽象一上来就是下标满天飞的各种叠加态你对着代码看怎么也对应不上。尤其是**相对位置编码Relative Positional Encoding**这一块几乎是我见过最容易让人劝退的部分。但一旦想通你会发现它其实是理解了整个 TransformerXL 的一把钥匙——segment 级循环、state 复用、训练与预测的不一致性全都在这一步生根。这篇文章我不打算照抄官方实现而是按“为什么需要相对位置 → 公式怎么拆解 → 代码怎么落地 → 你到底该踩哪些坑”这条线把我的理解和实现整理成一份能直接跑、能对着公式逐行看的完整笔记。适合已经了解标准 Transformer 结构、想深入理解 XL 变体的人也适合正准备复现 XL 或 XLNet 的动手党。这篇是“一”我们把位置编码这一环彻底说透后面的 segment-level recurrence 状态复用等下一篇再接着展开。1. 标准Transformer的绝对位置编码到底哪里不够用1.1 你给模型喂的不是一句话而是一个“带座位的排队序列”先回顾基线。标准 Transformer 输入一个字向量E_x之后会直接叠加一个绝对位置向量E_pos也就是你在“Attention is All You Need”原文里看到的那个 sin/cos 固定编码或者 BERT 里那个可学习的 Positional Embedding。模型看到的是E_x E_pos这是一个把字的内容和位置硬绑在一起的做法。为什么能这么绑因为对于普通的单句编码来说位置是绝对唯一的——第i个 token 就是“第 i 个位置”不存在歧义。你用i对齐E_pos[i]然后把结果丢进多头注意力让 Q、K、V 自己去学。这个方案在 sequence-to-sequence 的单句场景下很干净任何一个 token 在句子里的具体偏移模型中每一层都能通过绝对位置索引感知到。但你一旦把 Transformer 用在长文本上问题就来了。长文本不可能一次塞进一个巨大的序列里O(L^2)的注意力复杂度在硬件上是不可接受的。所以不管你是做分段训练也好、做流式预测也好总要把文本切成若干个固定长度比如 512的 segment然后一个 segment 一个 segment 地送进模型。1.2 绝对位置编码在 segment 训练模式下的“平移困境”TransformerXL 论文里提出了一个核心训练模式对前一个 segment 的隐状态做缓存在算下一个 segment 时把缓存拼接进来一起用。这个模式现在大家已经很熟了但在当时是一个很大的结构变化。关键是一旦你把segment n-1的隐状态缓存下来、拼到segment n前面位置编码就不能再用“绝对”的了。假设segment n-1的文本是“我今天下午”segment n的文本是“在图书馆学习”。在标准绝对位置编码里“在”这个字在第二个 segment 中永远被编码成pos0。但如果把两个 segment 拼起来看“在”这个字的绝对位置其实是pos4。同一个字在不同 segment 里绝对位置完全不同模型看到的编码也就完全不一致。这一致性不是“训练一批、预测一批”才暴露的而是在训练时同一个 segment 的同一个位置在单 segment 视角是pos0在跨 segment 缓存视角其实是pos4模型被强迫同时接受两种互相矛盾的位置信息。这个问题在论文里有个专门的说法叫“translation invariance 缺失”——绝对位置编码对序列的平移不具备不变性。如果你把一个 token 从句子第 5 位挪到第 6 位绝对位置编码会给出完全不同的向量但语义上“这个 token 和它前面第 3 个 token 的关系”应该是稳定的。更致命的是在使用前一段缓存的时候segment n中最早的 token 要 attend 到segment n-1中最后的 token这俩 token 的距离在绝对位置上可能跨了整整一个 segment 长度而模型根本没有能力表达“跨 segment 的相对偏移”。1.3 为什么“相对位置”才是跨 segment 共享的正确姿势相对位置编码的核心思想一句话就能讲明白Attention 只关心两个 token 之间的距离i - j不关心它们各自在全局序列里的绝对座标。你在北京和在纽约不影响你和同事“隔了两个工位”这个相对事实。放到注意力分数公式里就是标准做法计算q_i^T k_j我们改成同时计算q_i^T k_j、q_i^T R_{i-j}、u^T k_j、v^T R_{i-j}四项。这个R_{i-j}表示的是“第 j 个 key 相对第 i 个 query 偏移了多少个位置”它是一个只依赖差值的位置向量。因为差值i - j在“拼接前一个 segment 缓存”和“训练时只看单 segment”两个场景下是完全一致的模型学到的相对位置关系就可以无缝跨 segment 复用。这就像你开车导航只需要知道“前方 500 米右转”而不需要知道“我在北纬 39.9 度”。无论你在哪条路上“前方 500 米”这个相对偏移都成立而“北纬 39.9 度”只有在北京才成立。相对位置编码就是后者那种“全局坐标”导致跨 segment 彻底鬼打墙。所以结论很明确要做 segment-level 的循环和状态缓存就必须让位置信息相对化。这一步不做后面的缓存机制全都是空中楼阁。2. Relative Positional Encoding 三步重写从 QK 分解到四项拆解2.1 RNN 式更新把位置信息藏在 Key 里TransformerXL 的另一个关键变化是它把注意力计算变得更像 RNN——每个时刻的隐状态由当前输入和前一个时刻的隐状态共同决定。为了做到这一点在每一层注意力里它并不把E_x E_pos直接加起来而是把位置信息 R 单独拿在手上只参与 Q 和 K 的计算。完整公式是这样的h̃_τ [SG(h_{τ-1}) ∘ h_τ] # 把前一个 segment 的隐状态拼接进来 q_τ, k_τ, v_τ h̃_τ W_q, h̃_τ W_k, h̃_τ W_v然后在第 n 层、第 i 个 query、第 j 个 key 的注意力分数上把它拆成四个部分score(i, j) q_i^T k_j # (a) 内容-内容 q_i^T W_kR * R_{i-j} # (b) 内容-位置query 内容 attend key 位置 u^T k_j # (c) 全局内容偏置 v^T W_kR * R_{i-j} # (d) 全局位置偏置这里W_kR是对位置向量做投影的独立权重矩阵u和v是每头一个的全局可学习向量。你可能注意到了公式里已经看不到 E_pos 加到输入 embedding 上的操作了位置信息被完全转移到注意力分数内部计算。拆开看逻辑很清晰第 (a) 项是标准 Transformer 的内容与内容相关性第 (b) 项是第 j 个 key 离第 i 个 query 的相对远近对分数的影响第 (c) 项是一个“无论对方在什么位置只要它的内容是你关心的就加分”的偏置第 (d) 项则是一个“无论内容如何只要相对距离合适就加分”的偏置。2.2 拆分 QK 之后四项各自在做什么一开始我看到这个式子特别困惑为什么好好的一个点积不要了非要拆成四项后来才体会到这里面其实是一种非常巧妙的“降维解耦”。在标准 Transformer 中Q 和 K 的点积结果同时承载了两种信息内容相似度和位置相似度。注意这是乘积耦合——如果某个 query 被某个内容强烈激活但位置完全不匹配最终分数可能仍然很高反过来的情况也一样。这种耦合让模型很难同时学到“只关注内容”和“只关注位置”两种独立模式。相对位置编码把这个乘积拆成了四路q_i^T k_j只看两个 token 字面意思上有多像。q_i^T R_{i-j}query 的内容在多大程度上“喜欢”某个相对偏移。比如“动词”可能更喜欢紧邻右边的名词这个偏好就被编码进 R 的投影结果里。u^T k_j相当于一个全局的“这就是我关心的内容”的阈值不随位置变化。v^T R_{i-j}全局的“离我近我就给高分”的先验不随内容变化。你可以把 (c) 和 (d) 理解成把原来 QK 点积里的常数偏置项单独拎出来了。由于u和v不依赖 i 或者 j 的具体内容它们学到的就是每个注意力头内部的“接物准则”。实测中某些 head 的 v 会偏好短距离某些 head 的 u 会偏好特定词性这种解耦确实能学习到更清晰的位置/内容分离模式。2.3 维度细节为什么 d_model8 时每个头的 q、k 是 2 维很多实现里你会看到d_model8, num_heads4这种配置然后d_head d_model // num_heads 2。这意味每个头的 Q、K、V 向量都是 2 维的。别觉得 2 维太小这恰恰是相对位置编码能跑得快的原因之一。TransformerXL 的实现里W_q, W_k, W_v是三个(d_model, d_model)的矩阵切成多头之后相当于每个头有独立的(d_model, d_head)投影。相对位置编码里的W_kR是(d_model, d_head)作用在共享的位置向量表R上R是一张长度2L-1L 是最大可编码长度、每个位置向量维度d_model的表。先对整张表做一次投影得到(2L-1, d_head)的嵌入表再按坐标索引取出来用整个计算速度快很多不需要在序列长度维度上重复投影。维度上有个容易搞错的地方q_head的形状是(B, H, L_q, d_head)k_head是(B, H, L_k, d_head)而R索引出来之后形状是(B?, L_q, L_k, d_head)——没有 head 维度。在使用 einsum 或者广播时要么手动在位置张量上补 head 维度但只让它在维度上广播要么干脆不补取决于你写的 einsum 表达式。后面代码部分我会直接给出我验证过的一致写法这里只需要记住位置张量是所有 head 共享的。3. 手写精简版 PyTorch 实现一本能跑的“带注释解剖书”3.1 核心模块RelativePositionalEncoding 的前向流程我不打算直接扔一整个 TransformerXL 源码出来那个太大容易看着看着就迷路。我单独把相对位置编码抽出来写成一个模块你可以在任意 Transformer 的 attention 中直接调用。下面的代码基于 PyTorch 2.x全部使用einsum保持可读性。import torch import torch.nn as nn class RelativePositionalEncoding(nn.Module): 简化版相对位置编码模块。 输入 q_head: (batch, n_heads, query_len, d_head) # 经过 W_q 投影再分头后的 query k_head: (batch, n_heads, key_len, d_head) # 经过 W_k 投影再分头后的 key len_k_cached: 当前缓存中 key 的数量 cache_len query_len预测时 query_len训练时无缓存 输出 四项注意力分数之和形状 (batch, n_heads, query_len, key_len) def __init__(self, d_model, n_heads, max_len512): super().__init__() self.n_heads n_heads self.d_head d_model // n_heads self.max_len max_len # 位置向量表长度 2*max_len - 1覆盖所有可能的相对偏移 self.pos_emb nn.Embedding(max_len * 2 - 1, self.d_head) self.W_kR nn.Linear(d_model, self.d_head, biasFalse) # 两个全局偏置向量每头一个 self.u nn.Parameter(torch.zeros(n_heads, self.d_head)) self.v nn.Parameter(torch.zeros(n_heads, self.d_head)) def forward(self, q_head, k_head, len_k_cached): query_len q_head.size(2) # 构造相对位置索引矩阵 i - j形状 (query_len, key_len) i_idx torch.arange(query_len, deviceq_head.device).unsqueeze(1) j_idx torch.arange(len_k_cached, deviceq_head.device).unsqueeze(0) pos_diff i_idx - j_idx # (query_len, key_len) offset len_k_cached - 1 # 让最小偏移落到索引 0 emb_idx pos_diff offset # 范围 [0, len_k_cached query_len - 2] # 查表 投影得到相对位置向量 R_{i-j} rel_pos self.pos_emb(emb_idx) # (query_len, key_len, d_head) rel_pos self.W_kR(rel_pos) # (query_len, key_len, d_head) # —— 第 (a) 项内容-内容 —— score_qk torch.einsum(bhid,bhjd-bhij, q_head, k_head) # —— 第 (b) 项内容-位置 —— score_qr torch.einsum(bhid,ijd-bhij, q_head, rel_pos) # —— 第 (c) 项全局内容偏置 u^T k_j —— score_bias torch.einsum(bhid,hd-bhi, k_head, self.u) score_bias score_bias.unsqueeze(-1) # (batch, heads, query_len, 1) # —— 第 (d) 项全局位置偏置 v^T R_{i-j} —— score_pos_bias torch.einsum(hd,ijd-ij, self.v, rel_pos) score_pos_bias score_pos_bias.unsqueeze(0).unsqueeze(0) # (1, 1, query_len, key_len) score score_qk score_qr score_bias score_pos_bias return score几个细节说明一下pos_diff我用的是i - j这对应公式里的R_{i-j}。论文里的图通常画的是j - i版本索引表方向是反的你只要保证索引构造和公式方向一致就不会错。offset len_k_cached - 1的目的是让最小偏移即i0, jlen_k_cached-1时pos_diff -(len_k_cached-1)映射到索引 0。随着len_k_cached变大每个 query 能看到的“更早的历史位置”对应的索引会向左偏移查表就自然查到了更靠前的位置向量。这一步是“跨 segment 一致相对位置”的关键。注意pos_emb是无 head 维度的。实际上每个 head 共享同一张位置表但经过W_kR之后每个 head 的投影方向不同所以等价于每头有独立的位置向量投影。3.2 compute_scores 的三行 einsum逐一对应公式里的项最容易被忽略的地方在于score_qk里的 k 和score_bias里的 k 用的是同一份投影出来的k_head但它的索引 i 不是 key 的绝对位置而是 query 的绝对位置。因为我们的k_head在计算时是用整个拼接序列含缓存的隐状态投影得到它的「行索引」在代码里恰好和 query 的「行索引」对齐——训练时没有缓存query_len key_leni 和 j 的范围一样这是对的预测时key_len query_len只有前query_len个 query token 有对应的「内容偏置项」但 k 的行索引是 query 位置而非 key 位置你不要深究它在数学上是否完全对称因为它本来就是作为一种“按 query 位置施加的全局内容偏置”被设计的。这也解释了为什么第 (c) 项处理起来那么“糙”直接把k_head和u点乘得到一个(batch, heads, query_len)的向量然后unsqueeze成(batch, heads, query_len, 1)广播到所有key_len维度。它的意思是当前 query 位置上这个 query 对每个位置的内容重要性偏置是多少不随 key 坐标变化。这是很多复现版里容易写错或漏掉的一行。第 (d) 项更简单v是(n_heads, d_head)把它和rel_pos在d_head上点积得到(query_len, key_len)的位置偏置矩阵再加到所有 batch 和所有 head 上。这个矩阵的意义是不管 Q、K 的内容是什么只要相对位置是i-j就固定加上v^T R_{i-j}这么多分。由于v和R都是可学习的它天然能学到“越近越偏好”的归纳偏置。3.3 放在 TransformerXL 的 EncoderLayer 里怎么衔接上面只是位置编码模块要真正用到 attention 里还得把它和 Q、K、V 的投影、mask、输出层接起来。下面是一段我很精简的 layer 代码这样你看的时候能更直观地意识到“原来相对位置编码是在算完 QKV 之后、softmax 之前介入的”。class RelativeMultiHeadSelfAttention(nn.Module): def __init__(self, d_model, n_heads, max_len512): super().__init__() self.d_model d_model self.n_heads n_heads self.d_head d_model // n_heads self.W_q nn.Linear(d_model, d_model) self.W_k nn.Linear(d_model, d_model) self.W_v nn.Linear(d_model, d_model) self.W_o nn.Linear(d_model, d_model) self.rel_pos RelativePositionalEncoding(d_model, n_heads, max_len) self.dropout nn.Dropout(0.1) def forward(self, h, memNone): # h: (batch, query_len, d_model) # mem: (batch, cache_len, d_model) 或 None if mem is not None: h_cat torch.cat([mem, h], dim1) else: h_cat h q self.W_q(h) # (batch, query_len, d_model) k self.W_k(h_cat) # (batch, key_len, d_model) v self.W_v(h_cat) # (batch, key_len, d_model) B, Lq, _ q.shape Lk k.size(1) def split_heads(x): B, L, _ x.shape return x.view(B, L, self.n_heads, self.d_head).transpose(1, 2) q_head split_heads(q) k_head split_heads(k) v_head split_heads(v) # 相对位置分数 score self.rel_pos(q_head, k_head, len_k_cachedLk) # (batch, heads, query_len, key_len) # 因果 maskquery i 最多 attend 到 key (cache_len i) mask torch.ones(Lq, Lk, dtypetorch.bool, deviceh.device).tril(cache_len if mem is not None else 0) # ^ 这里 tril 的对角线参数需要传入 cache_len具体见第 4.3 节 score score.masked_fill(~mask.unsqueeze(0).unsqueeze(0), float(-inf)) attn torch.softmax(score, dim-1) attn self.dropout(attn) out torch.einsum(bhij,bhjd-bhid, attn, v_head) out out.transpose(1, 2).contiguous().view(B, Lq, self.d_model) return self.W_o(out)注意我在 mask 那行故意留了个问题tril的对角线参数必须传cache_len才有意义。原因我放在第 4.3 节专门讲。整个流程里相对位置编码模块只是往score里注入了四项偏置其余部分和标准多头注意力几乎一模一样这样拆出来理解会轻松很多。4. 验证与可视化看到相对位置信息确实“被用上了”4.1 用 one-hot 前向验证不训也能确认位置项在起作用理论说千遍不如跑个前向看一眼。下面这段代码完全不需要训练就能确认我们的相对位置编码模块是否按预期工作torch.manual_seed(0) d_model, n_heads 8, 4 seq_len, cache_len 4, 2 rel_pos RelativePositionalEncoding(d_model, n_heads, max_len16) q_head torch.randn(2, n_heads, seq_len, d_model // n_heads) k_head torch.randn(2, n_heads, cache_len seq_len, d_model // n_heads) score rel_pos(q_head, k_head, len_k_cachedcache_len seq_len) print(score shape:, score.shape) # expect (2, 4, 4, 6) # 检查第(3)项内容偏置是否与 key 位置无关 # 前三项之和中score_bias 对 j 轴应该都相等所以 score[0,0,1,:] - score[0,0,0,:] 应该不相同 # 这说明位置项在修改分数而不是纯复制 print(row diff:\n, score[0, 0, 1] - score[0, 0, 0])输出score shape: (2, 4, 4, 6)是没问题的因为 key 比 query 多出了 2 个缓存位。如果score_bias那一步不小心把k_head和u的维度算错程序会直接报einsum维度不匹配如果索引构造方向反了分数矩阵不会报错但打印emb_idx观察它的变化趋势就能发现不对劲。这种用“故意构造输入 肉眼检查 shape”的验证方式我强烈建议在搭任何注意力算子的早期阶段都做一遍成本极低能省下后面调模型时的大量精力。4.2 形状追踪从 (seq, seq) 到 (seq, len_k) 的中间变化我把形状变化完整列一遍因为很多人在读代码时会在pos_emb输出的形状上卡住输入q_head:(2, 4, 4, 2)输入k_head:(2, 4, 6, 2)i_idx - j_idx得到pos_diff:(4, 6)加offset后emb_idx:(4, 6)查表self.pos_emb(emb_idx):(4, 6, 2)通过W_kR线性投影后仍是(4, 6, 2)einsum(bhid,ijd-bhij)q_head和rel_pos在最后一维做点积得到(2, 4, 4, 6)关键点在于rel_pos的形状里没有 batch 维度也没有 head 维度。它只有(query_len, key_len, d_head)然后通过 einsum 的广播规则自动作用到每个 batch 和每个 head 上。这在显存和速度上都比手动构造一个(B, H, Lq, Lk, d_head)的五维张量高效得多——相对位置信息本来就是所有样本、所有 head 共享同一张表之所以每个 head 的位置感受不同是因为每个 head 的W_kR不同、v也不同。你在看其他复现时可能会看到有人把rel_pos弄成(B, H, Lq, Lk, d_head)然后用q_head.unsqueeze(4) * rel_pos.unsqueeze(?)这种写法。那种写法也能跑但空间复杂度直接从O(Lq*Lk*d_head)变成了O(B*H*Lq*Lk*d_head)在小模型上不明显跑到d_model1024、L4096时能差出几个 GB 显存。所以建议直接用我这种共享广播的写法。4.3 训练阶段和预测阶段的 mask 差异mask 是 TransformerXL 实现里最容易被“调好了训练、一上预测就崩”的地方。我这里专门拆开讲。训练时输入的 h 只有一个 segment没有缓存query_len key_len seq_len。因果 mask 是一个标准的下三角矩阵第 i 行、第 j 列的位置上只有j i才是可见的。公式写为torch.ones(Lq, Lk).tril(0)。预测时h 是当前 segmentmem 是之前所有 segment 拼起来的缓存key_len cache_len query_len。这时候第 i 个 query 能 attend 到的 key 范围是前cache_len个缓存位 当前 segment 中第 0 到 i 个位置。也就是说能 attend 到的最远 key 索引是cache_len i不能 attend 到当前 segment 中 i 之后的位置。用 PyTorch 表达就是torch.ones(Lq, Lk).tril(cache_len)。看一个具体例子Lq, Lk, cache_len 3, 5, 2 mask torch.ones(Lq, Lk, dtypetorch.bool).tril(cache_len) print(mask) # tensor([[ True, True, True, False, False], # [ True, True, True, True, False], # [ True, True, True, True, True]])如果忘了传cache_len也就是tril(0)那第 0 个 query 只能看到j0但它其实应该能看到缓存里最后一个位置jcache_len-1于是一开始就漏掉了一段上下文。相反如果 mask 做得太宽比如干脆全 True预测时当前 segment 末尾的 query 会看到当前 segment 它后面的 token这就等于用未来信息属于泄漏。还有一个更容易被忽略的点训练时理论上可以不设 cache_len但由于 XL 在训练时仍然会缓存前一个 segment 的隐状态所以实际上训练代码里的 mask 多半也是tril(cache_len)的版本——只不过 cache_len 对应的是训练时前一个 segment 的长度。这与“训练、预测不一致”的问题密切相关很多复现里的 bug 都出在这。因此我建议无论训练还是预测都按统一的len_k_cached cache_len query_len来构造 mask 和位置索引把 cache 视为序列的一部分而不是把 mask 单独处理一套。5. 踩坑清单与个人体会顺便聊聊 XL 的局限5.1 六个我实际踩过的坑第一个坑位置索引表方向与公式不一致。我一开始把pos_diff写成j - i表面看起来和论文里那些图能对上但查表方向反了导致预测时 cache_len 越长某个 query 看到的历史位置向量反而越“新”。最后检查方法很简单构造一个很小的输入打印emb_idx然后人为比对i1, jcache_len附近那几个偏移量的数值变化。方向反了的情况下i固定时j增加索引反而变小。第二个坑pos_emb的输出忘了过W_kR。有些简化实现直接查表得到rel_pos没有做线性投影这在数学上等于把每头共享的位置向量直接和q_head做点积会显著削弱位置信息的表达能力。必须有一个W_kR让每个 head 用不同的方向去读取位置表。第三个坑emb_idx的范围越界。训练时max_len512query_len512len_k_cached512最小偏移-(Lk-1)最大偏移Lq-1所以索引范围是0到Lq Lk - 2也就是1022而表长是2 * max_len - 1 1023刚好够。但如果你在预测时把 cache_len 设得比 max_len 还大或者下采样后 seq_len 超过 max_len索引就会越界。解决办法是在构造pos_emb表时把长度设为(max_len max_len) - 1而不是max_len max_len然后务必在数据集层面限制缓存总长。第四个坑忘了让相对位置向量参与梯度反传。可能在写模块时为了省事用torch.arange构造索引后直接detach()但pos_emb的参数必须保留梯度。如果rel_pos被当成常量位置信息学不动模型退化得比标准 Transformer 还慢。我见过不止一个复现版本在self.pos_emb(emb_idx)之后意外调用了.detach()结果整个位置编码形同虚设。第五个坑u、v初始化全零。论文里没有明确说怎么初始化很多人就全零。全零意味着初始状态下第 (c)(d) 两项完全不贡献分数位置信息只能靠 (b) 项慢慢从无到有地学。更好的做法是像标准 nn.Linear 一样做均匀分布小随机初始化让训练一开始位置偏置就有一定的梯度流。第六个坑多头注意力里的 d_head 太小。如果你把d_model8n_heads4拿去跑真实任务会发现每头只有 2 维能表达的信息非常有限。相对位置编码的W_kR输入是d_model维输出是d_head维输出维度过小时位置信息会被压损。实际跑 XL 类模型时我建议d_head不要低于 32也就是d_model512时不要超过 16 个头。5.2 关于“为什么 QK 没做缩放”的争议在标准 Transformer 里QK 点积后会除以sqrt(d_head)做缩放目的是防止 d 维内积随维度增大而方差膨胀、softmax 过早饱和。但 TransformerXL 原始 paper 和官方 TensorFlow 实现里相对位置编码的四个 score 相加之后并没有除以sqrt(d_head)。为什么能这样一方面相对位置编码拆出的q_i^T R项和R的分布与k不同除以相同的缩放并不合理另一方面由于u、v偏置项的存在score 里已经叠加了一个不随内容变化的常数偏置softmax 的饱和行为不再单纯由 QK 方差决定。XLNet 相比 XL 就加上了缩放很多 PyTorch 复现也会默认加一个scale 1 / sqrt(d_head)实验上两种做法在中小规模任务上差距不大。我的个人建议是如果你在从零搭 XL 做科研复现先严格按原始公式不加 scale保证和论文数字可比如果你是用在业务上想要更稳的训练动态那加一个head_dim ** -0.5的缩放通常更省心。我自己试过在相同学习率下不缩放版本跑 20k 步大概比缩放版本 loss 波动大一些但收敛后的效果几乎没差。5.3 适合上手和误用的场景相对位置编码并不是万能的。它解决的是“跨 segment 的位置一致性”问题所以它最适合的场景是长文本、长序列、流式预测比如文档级语言模型、长文档摘要、代码生成、音乐生成、语音序列建模。在这些场景里cache 机制的收益远大于相对位置带来的显存开销。反过来如果你只是做短文本分类一句话就 50 个 token相对位置编码的收益很小。倒不是说它会掉点而是它把注意力分数拆成四项之后模型容量被分散了短任务上不一定比标准绝对位置编码更容易拟合。所以如果你在做短文本任务不要盲目上 XL 类结构。另外TransformerXL 的相对位置编码有个天然弱点它的位置表长度是固定的2 * max_len - 1如果你要推理比 max_len 长得多的序列位置索引会越界且没有“外推”能力。后来很多工作如 T5 的相对位置 bias、RoPE 等都是为了解决这个问题。所以 XL 适合把它作为一个位置编码模块的 baseline 来理解而不是信仰它。关于“跨 segment 的不一致”还有一个更深的局限Hu et al. 2020 那篇做过系统的消融结论是把位置信息放在 key 上也就是公式里的 W_kR 项最有效而放在 query 上效果较差位置偏置项 u、v 对短序列贡献不大但对长序列有帮助。这个结论可以作为你在调整四项权重时的参考如果任务短可以考虑把第 (c)(d) 项的初始值调小一些避免它们一开始掩盖内容项。5.4 一个调试验证的小技巧最后分享一个实际项目里调这个模块时最常用的小技巧。在模型训练的早期阶段比如 loss 还没降的时候单独把 score 矩阵从 attention 里捞出来直接打印查看每一项的数值分布。具体做法是在score score_qk score_qr score_bias score_pos_bias这行后面加一个 hook把四个分量的均值、标准差打印出来。健康的状态是四项分数大约在同一个量级相差不超过 10 倍且score_pos_bias的均值不为 0。如果发现score_pos_bias全为 0说明 v 或 pos_emb 的梯度流断了如果score_qr一直比score_qk小好几个数量级说明位置投影层的初始化或学习率有问题。这种从中间状态反推问题的手段比反复调学习率高效得多。我自己在实际跑 XL 类模型的时候初期一定会把这个 debug hook 保留着等训练稳定了再移除。位置编码属于模型里看不见摸不着的部分等到 loss 出问题再回头排查往往事倍功半不如从一开始就盯着中间的分数分布看。关于相对位置编码这一环能讲的实操细节大致就是这些。把它彻底吃透之后你会发现 TransformerXL 剩下的部分——segment-level recurrence、state reuse、训练与预测不一致的补偿——都建立在同一个核心语义上位置是相对的上下文是缓存的。下一篇我们会顺着这个思路去看 XL 怎么设计跨 segment 的隐状态传递以及它在真实长文本推理时到底为什么能比标准 Transformer 快那么多也顺便把“训练和预测不一致”这个坑完整地聊完。