说实话,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 中永远被编码成pos=0。但如果把两个 segment 拼起来看,“在”这个字的绝对位置其实是pos=4。同一个字,在不同 segment 里绝对位置完全不同,模型看到的编码也就完全不一致。这一致性不是“训练一批、预测一批”才暴露的,而是在训练时,同一个 segment 的同一个位置,在单 segment 视角是pos=0,在跨 segment 缓存视角其实是pos=4,模型被强迫同时接受两种互相矛盾的位置信息。
这个问题在论文里有个专门的说法叫“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_model=8 时每个头的 q、k 是 2 维
很多实现里你会看到d_model=8, num_heads=4这种配置,然后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-1(L 是最大可编码长度)、每个位置向量维度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_len=512): 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, bias=False) # 两个全局偏置向量,每头一个 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, device=q_head.device).unsqueeze(1) j_idx = torch.arange(len_k_cached, device=q_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的目的是让最小偏移(即i=0, j=len_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_len,i 和 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_len=512): 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, mem=None): # h: (batch, query_len, d_model) # mem: (batch, cache_len, d_model) 或 None if mem is not None: h_cat = torch.cat([mem, h], dim=1) 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_cached=Lk) # (batch, heads, query_len, key_len) # 因果 mask:query i 最多 attend 到 key (cache_len + i) mask = torch.ones(Lq, Lk, dtype=torch.bool, device=h.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_len=16) 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_cached=cache_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_model=1024、L=4096时能差出几个 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 是当前 segment,mem 是之前所有 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, dtype=torch.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 只能看到j=0,但它其实应该能看到缓存里最后一个位置j=cache_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,然后人为比对i=1, j=cache_len附近那几个偏移量的数值变化。方向反了的情况下,i固定时j增加,索引反而变小。
第二个坑:pos_emb的输出忘了过W_kR。有些简化实现直接查表得到rel_pos,没有做线性投影,这在数学上等于把每头共享的位置向量直接和q_head做点积,会显著削弱位置信息的表达能力。必须有一个W_kR,让每个 head 用不同的方向去读取位置表。
第三个坑:emb_idx的范围越界。训练时max_len=512,query_len=512,len_k_cached=512,最小偏移-(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_model=8n_heads=4拿去跑真实任务,会发现每头只有 2 维,能表达的信息非常有限。相对位置编码的W_kR输入是d_model维,输出是d_head维,输出维度过小时位置信息会被压损。实际跑 XL 类模型时,我建议d_head不要低于 32,也就是d_model=512时不要超过 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 快那么多,也顺便把“训练和预测不一致”这个坑完整地聊完。