
如果你已经熟悉经典 Transformer那一定知道位置编码是绕不开的一个基础模块。不过到了 TransformerXL 这里位置编码的玩法完全变了它不再是给词向量加一个绝对位置向量而是把位置信息全部揉进注意力计算里用相对位置编码替代原来的绝对位置编码。这篇我就拿相对位置编码开刀把公式、推导、PyTorch 实现一次讲透全部带代码保证你看完能直接用在自己的项目里。这个系列的文章都偏实战不适合那种只想大概了解概念就走的读者。如果你是要做长文本建模、做语言模型预训练或者想把 Transformer 的死板位置编码换成更灵活的方案那这篇文章非常适合你。我会把推理过程写得尽量细公式也不会只放结果不解释由来每一步都告诉你为什么要这么做。在动手写代码之前我们先把原理捋清楚。相对位置编码并不是简单地把位置向量换成相对距离TransformerXL 的论文里其实做了两处关键改动一是去掉输入层的位置编码加法改成在注意力分数计算时用相对位置向量二是引入了两组可学习的全局偏置向量 u 和 v。这两处改动到底解决了什么问题以及代码里是怎么落地的我们一个一个来拆。1. 从绝对位置编码到相对位置编码TransformerXL 到底改了什么1.1 经典 Transformer 的位置编码是怎么工作的先回顾一下原始 Transformer 的位置信息是怎么进来的。当时的做法是把每个位置编号 t 编码成一个向量 p_t然后直接加到词嵌入 e_w 上输入给模型的就是 x_t e_w p_t。p_t 的生成公式是论文里那个经典的 sin/cos 函数PE(t, 2i) sin(t / 10000^(2i/d_model)) PE(t, 2i1) cos(t / 10000^(2i/d_model))这个设计有两个明显的作用第一它让模型在输入端就感知到 token 的绝对顺序第二sin/cos 的组合让 p_t 之间存在线性关系理论上模型可以通过学习捕捉到相对位置信息。但注意这只是“理论上”实际上模型要真的从相加后的向量里抽出相对距离信息得多学一层非线性变换代价不小。更麻烦的是绝对位置编码在序列长度面前有一个硬伤训练时只见过 512 个位置推理时喂给它 600 个 token后面那些位置的编码向量完全是没见过的。这就导致模型在长序列上的表现经常突然崩溃。这个问题在很多实际场景里非常致命比如文档级别的文本建模你能保证每段话长度都不超过训练长度吗不能。所以 TransformerXL 的作者在着手解决这个问题时第一刀就砍在了位置编码上。他们的核心思路是既然模型真正需要的其实是“当前位置和每个历史位置的相对距离”那不如直接把这个相对距离作为输入给模型而不是让模型从绝对坐标里去反推。这就引出了相对位置编码的核心设计。1.2 为什么相对距离比绝对坐标更适合语言建模想象一个句子“小明昨天去了超市他买了一瓶水。”模型在预测“他”指代谁的时候它真的需要知道“小明”出现在第 1 个位置、“昨天”在第 2 个位置吗其实不需要。它需要知道的是“小明”出现在当前 token 前面大概多远的位置以及它在句中扮演什么语法角色。这种“距离多少”的信息恰好是注意力机制最擅长利用的。如果用绝对位置编码注意力分数计算时查询和键都带了各自的位置信息模型要判断两个 token 之间的关系得先解算出它们位置的差。这个解算过程对模型来说既不直接也不稳定。如果改用相对位置编码那么查询在计算某个键的时候拿到的直接就是一个“对方离我多远”的向量这相当于把模型的一部分工作直接外包给了位置编码模块。在 TransformerXL 的论文里作者保留了经典 Transformer 的 Q、K、V 结构但对位置信息的使用方式做了替换不再把位置向量加到输入上而是在注意力分数公式里把原来的绝对位置项拿掉换成一项显式的相对位置编码同时还加入两组可学习的偏置向量。这两组偏置向量的作用后面会细说现在先记住一个结论经典 Transformer 的位置信息是“加法注入”TransformerXL 的位置信息是“乘法注入”通过点积计算进注意力分数。这一步改动看着不大但长序列能力提升非常明显。而且因为位置信息不再依赖绝对编号模型天然就具备了一定的外推能力只要相对距离在训练时见过具体出现在第几个位置其实无所谓。这也就是为什么 TransformerXL 能被用来处理比训练长度更长的序列而经典 Transformer 做同样的事情往往会掉点。2. 相对位置编码公式拆解四个项的物理意义这一节是全文的骨架。我会把公式从经典 Transformer 开始一步步改写到最后的样子。只有真正看懂了这四个项各自负责什么写代码的时候才知道每个矩阵相乘在干嘛。2.1 原始注意力公式回顾经典 Transformer 的单头注意力假设不使用缩放点积的简化写法第 i 个查询和第 j 个键之间的注意力分数可以写作score(i, j) (W_q (e_i p_i))^T · (W_k (e_j p_j))把它展开会得到四项W_q e_i 与 W_k e_j 的点积纯内容对纯内容的匹配W_q e_i 与 W_k p_j 的点积查询内容与键位置的关系W_q p_i 与 W_k e_j 的点积查询位置与键内容的关系W_q p_i 与 W_k p_j 的点积纯位置对位置的关系。这种展开看起来没什么问题但仔细想想四项里面有一半是在处理“位置与内容”的交叉关系。模型并不是很关心“一个 token 的内容和另一个 token 的位置”之间有什么相互作用它关心的主要是内容-内容、位置-位置这两类信息。然而公式把所有东西都混在一起学位置信息又要通过绝对坐标去隐式表达相对距离学习压力全堆在参数上了。2.2 TransformerXL 的四个注意力项TransformerXL 论文提出的相对位置编码公式长这样score_rel(i, j) q_i^T k_j q_i^T W_kR^T R_{i-j} u^T k_j v^T W_kR^T R_{i-j}这里我稍微统一一下记号方便你对照代码q_i 是第 i 个查询向量已经经过 W_q 投影k_j 是第 j 个键向量已经经过 W_k 投影但是注意k_j 里不再包含位置编码R_{i-j} 是一个相对位置向量表示距离为 i-j 的位置编码W_kR 是专门给相对位置向量用的投影矩阵对应代码里的 w_k_posu 和 v 是两个可学习的全局偏置向量。拆开来看这四个项的语义是q_i^T k_j —— 纯内容项完全基于语义内容打分q_i^T W_kR^T R_{i-j} —— 查询内容与“键的相对位置”之间的得分模型关注“我当前这个 token 和距离我 i-j 的那个 token 在内容上是否相关”u^T k_j —— 全局偏置与键内容的得分表示模型对每个键内容本身的先验偏好v^T W_kR^T R_{i-j} —— 全局偏置与相对位置的得分表示模型对“某个相对距离”的全局先验。跟原版四项对比一下最大的变化在哪里取消了 W_q p_i 这个项。也就是说查询本身不再携带位置信息位置信息只在计算“查询内容到相对位置键”和“全局偏置到相对位置键”这两项时参与。这让模型可以单独为“内容和位置的关系”建模比原来混合在一块的方式干净得多。这里有一个很容易疑惑的地方为什么保留针对键的位置信息却不给查询加位置我的理解是在语言建模这种自回归场景里查询的位置是固定的我要预测当前 token 时我用它当前位置的查询向量而键的位置才是真正需要扫过整个历史上下文的。你可以把查询想象成一个站在当前位置的人键是历史里一排排的档案他需要知道每份档案离他多远但自己的坐标其实不重要。这个直觉在写代码时也会体现出来相对位置矩阵的形状是 [q_len, k_len]而不是在 q 和 k 上各做一套。2.3 两个偏置向量 u 和 v 到底在干嘛初次接触 TransformerXL 的人看到 u 和 v 通常会有两个问题为什么需要两个为什么它们是全局的而不是每个位置都有一个先回答第二个问题。如果每个位置都有一个单独的偏置向量那其实就退化成了某种绝对位置信息模型会去记住第 3 个位置偏好匹配第 5 个位置这样的绝对模式。而把偏置设为全局可学习的模型学到的是“无论你在什么绝对位置只要内容或距离满足某个语义模式就给你加分”。这样既保留了位置信息又不会让模型过度依赖绝对坐标。那为什么是两个而不是一个我们仔细看公式第三项和第四项。第三项 u^T k_j 只和键的内容有关它相当于一个内容门控如果某个键的内容本身很重要不管当前查询是什么都值得获得一个基础分数。第四项 v^T W_kR^T R_{i-j} 只和相对距离有关它相当于一个距离先验模型可以学到“距离为 1 的 token 之间的注意力权重通常要更高”这样的规律。如果把 u 和 v 合并成一个内容门控和距离先验就纠缠在一起了模型没法分别调节这两类先验的强度。论文里在很多实验中都验证了这两个偏置向量能提升稳定性所以别看它们只是两个向量博文里的小白读者也不要觉得这两个参数无关紧要。2.4 矩阵化改写从单元素到可并行的矩阵乘法上面一直是单元素写法工程上当然不能这么计算。TransformerXL 论文给出了等价的矩阵形式这里我直接写成代码更容易实现的样子score_rel (Q u) · K^T (Q v) · (W_kR · R)^T其中Q 是查询矩阵 [q_len, d_k]K 是键矩阵 [k_len, d_k]R 是相对位置向量矩阵 [q_len, k_len, d_k]它索引的是 R_{i-j} 向量u 和 v 分别广播到每个查询位置W_kR 对 R 做一个线性投影。注意一下这个式子里的 Q 出现了两次这是矩阵形式的核心第一次 Q 和 u 组合用来和键的内容部分算分第二次 Q 和 v 组合和相对位置编码算分。前面拆开的四项在这里通过矩阵运算自然合并成了两项但语义还是那四项。写代码时如果直接按这个矩阵形式实现效率会好很多。到这里原理部分就打通了。接下来进入代码实战我会手写一个相对位置编码模块和一个相对多头注意力层把上面的公式翻译成 PyTorch 代码。3. PyTorch 实现从零开始写相对位置编码注意力层3.1 代码设计与环境准备建议使用 Python 3.8 以上PyTorch 1.10 以上。整个实现不依赖额外第三方库只用到 torch 和 torch.nn。我会把代码拆成两个类RelativePositionEncoding —— 负责生成相对位置向量矩阵RelMultiheadAttention —— 负责完整的多头相对注意力计算。最后再把它们拼成一个 Transformer 块并跑一个最简单的输出形状测试。下面每一段我都会先放代码再解释关键逻辑。3.2 构建相对位置向量矩阵相对位置编码模块最重要的任务是给定查询长度 q_len 和键长度 k_len生成一个形状为 [q_len, k_len, head_dim] 的张量其中第 (i, j, :) 个向量表示位置 i 和位置 j 之间的相对距离编码。注意这里的 head_dim 是每个注意力头的维度也就是说相对位置编码是在每个 head 的维度空间里做的。import torch import torch.nn as nn import torch.nn.functional as F import math class RelativePositionEncoding(nn.Module): def __init__(self, head_dim, max_len512): super().__init__() self.head_dim head_dim self.max_len max_len # 位置向量表索引范围覆盖 -(max_len-1) 到 (max_len-1) # 所以总长度是 2 * max_len - 1 self.pos_table nn.Parameter( torch.randn(2 * max_len - 1, head_dim) * 0.02 ) def forward(self, q_len, k_len): # 构造相对位置索引矩阵 [q_len, k_len] # pos_idxs[batch_i, j] i - j pos_idxs torch.arange( q_len, deviceself.pos_table.device ).view(-1, 1) - torch.arange( k_len, deviceself.pos_table.device ).view(1, -1) # 限制在 [-max_len1, max_len-1] 范围内 pos_idxs pos_idxs.clamp(-(self.max_len - 1), self.max_len - 1) # 平移到非负索引 pos_idxs pos_idxs (self.max_len - 1) # 查表得到相对位置编码矩阵 # 输出形状: [q_len, k_len, head_dim] rel_emb self.pos_table[pos_idxs] return rel_emb这个类虽然短但里面有三个点值得展开讲。第一为什么表的长度是 2 * max_len - 1因为相对距离的取值范围是 [-(max_len - 1), max_len - 1]左闭右闭一共 2 * max_len - 1 个整数。如果你设置了 max_len512那表里就有 1023 个可学习的向量。索引 idx i - j (max_len - 1)这一步很关键因为 PyTorch 的索引只能是非负整数。第二这里的位置向量表是在初始化时随机生成的然后作为 nn.Parameter 参与训练。这不是 sin/cos 固定编码而是学习出来的相对位置向量比固定公式更灵活。TransformerXL 论文里的相对位置编码也是可学习的。第三forward 里的 q_len 和 k_len 是分开传入的因此这个模块天然支持查询和键长度不一致的交叉注意力场景。后面写注意力层的时候你只需要在调用时传入 q.size(2) 和 k.size(2) 就可以。3.3 相对位置多头注意力层的完整实现有了相对位置矩阵接下来就是把公式翻译成多头注意力计算。这个类相对复杂我先放完整代码再逐段拆解。class RelMultiheadAttention(nn.Module): def __init__(self, d_model, n_heads, max_len512): super().__init__() assert d_model % n_heads 0, d_model 必须能被 n_heads 整除 self.d_model d_model self.n_heads n_heads self.head_dim d_model // n_heads self.max_len max_len # 内容投影 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_k_pos nn.Parameter( torch.randn(n_heads, self.head_dim, self.head_dim) * 0.02 ) # 两个全局可学习偏置向量每个头有自己的 self.u nn.Parameter(torch.randn(n_heads, self.head_dim) * 0.02) self.v nn.Parameter(torch.randn(n_heads, self.head_dim) * 0.02) # 相对位置编码生成器 self.pos_enc RelativePositionEncoding(self.head_dim, max_len) # 输出投影 self.out_proj nn.Linear(d_model, d_model) def forward(self, q, k, v, maskNone): # q, k, v: [B, seq_len, d_model] B, q_len, _ q.shape _, k_len, _ k.shape # 1. 线性投影并拆分成多头 q self.w_q(q).view(B, q_len, self.n_heads, self.head_dim).permute(0, 2, 1, 3) k self.w_k(k).view(B, k_len, self.n_heads, self.head_dim).permute(0, 2, 1, 3) v self.w_v(v).view(B, k_len, self.n_heads, self.head_dim).permute(0, 2, 1, 3) # q, k, v 形状: [B, n_heads, seq_len, head_dim] # 2. 内容相关项: (Q u) K^T # u 形状为 [n_heads, head_dim]这里广播到每个 batch 和每个位置 q_with_u q self.u[None, :, None, :] # [B, n_heads, q_len, head_dim] content_score torch.matmul(q_with_u, k.transpose(-2, -1)) # content_score: [B, n_heads, q_len, k_len] # 3. 相对位置相关项 # 生成相对位置编码矩阵 [q_len, k_len, head_dim] rel_emb self.pos_enc(q_len, k_len) # 对位置向量进行逐头投影: [q_len, k_len, head_dim] - [q_len, k_len, n_heads, head_dim] pos_emb_proj torch.einsum( qkd,hde-qkhe, rel_emb, self.w_k_pos ) # 与 (Q v) 做点积 q_with_v q self.v[None, :, None, :] # [B, n_heads, q_len, head_dim] pos_score torch.einsum( bhqd,qkhd-bhqk, q_with_v, pos_emb_proj ) # pos_score: [B, n_heads, q_len, k_len] # 4. 合并内容分数和位置分数除以 sqrt(head_dim) 进行缩放 attn_score (content_score pos_score) / math.sqrt(self.head_dim) # 5. 可选 mask阻止注意力看到未来位置语言模型场景 if mask is not None: attn_score attn_score.masked_fill(mask 0, float(-inf)) # 6. softmax 归一化并加权求和 attn_prob F.softmax(attn_score, dim-1) # 7. 与 value 相乘 out torch.matmul(attn_prob, v) # [B, n_heads, q_len, head_dim] # 8. 合并多头并做输出投影 out out.permute(0, 2, 1, 3).contiguous().view(B, q_len, self.d_model) out self.out_proj(out) return out我们来逐段看代码里的坑和设计考虑。先从投影开始。w_q、w_k、w_v 都是标准的 Linear 层把 d_model 维向量映射到 d_model 维然后通过 view 切分成 n_heads 个头再用 permute 把头的维度挪到第 1 维。这里要注意view 和 reshape 的区别view 要求内存连续permute 之后不能直接 view所以我先调用 contiguous()在最后合并多头时才需要。新手在这一步最容易踩坑报错通常是 “view size is not compatible with input tensor’s size and stride” 之类的。然后是内容分数。这里的 U 是 [n_heads, head_dim]通过 self.u[None, :, None, :] 变成 [1, n_heads, 1, head_dim]然后加在 q 上相当于给每个查询向量都加上了各自的全局内容偏置。这一步是公式里 q_i^T k_j u^T k_j 两个项合并后的结果我把 u 先加到 q 上再和 k 做矩阵乘。如果你对公式的符号比较敏感会发现这是从 (Q u) K^T 推导出来的。相对位置分数部分是我最想说清楚的地方。rel_emb 形状是 [q_len, k_len, head_dim]它代表每一对位置 (i, j) 有一个绝对对应距离的向量。w_k_pos 是 [n_heads, head_dim, head_dim]每个头一个投影矩阵对应论文里的 W_k^R。这里用 einsum 一次性把位置向量投影到每个头的空间得到 [q_len, k_len, n_heads, head_dim]。然后 q_with_v 是 [B, n_heads, q_len, head_dim]再和 pos_emb_proj 做 einsum 点积得到 [B, n_heads, q_len, k_len] 的位置分数。为什么位置分数要把 v 加到 q 上因为我们在公式里要算 q_i^T W_kR^T R_{i-j} v^T W_kR^T R_{i-j}合并同类项就是 (q_i v)^T W_kR^T R_{i-j}。这里 v 的维度是 [n_heads, head_dim]广播到每个位置。最后归一化。scale 用的是 sqrt(head_dim)不是 sqrt(d_model)。多头注意力中每个头的维度是 head_dim所以缩放因子要按 head_dim 来。这个细节很多初学者容易写错会导致训练不稳定。3.4 把注意力层拼成 Transformer 块有了相对多头注意力剩下的就是把它堆成一个可以使用的 Transformer 编码块。为了保持代码简单我这里提供一个最精简的块包含一个多头注意力、一个前馈网络和两层 LayerNorm。class TransformerXLBlock(nn.Module): def __init__(self, d_model, n_heads, d_ff, max_len512, dropout0.1): super().__init__() self.attn RelMultiheadAttention(d_model, n_heads, max_len) self.ff nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model), ) self.ln1 nn.LayerNorm(d_model) self.ln2 nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): # 子层 1相对位置多头注意力 残差 LayerNorm attn_out self.attn(x, x, x, maskmask) x self.ln1(x self.dropout(attn_out)) # 子层 2前馈网络 残差 LayerNorm ff_out self.ff(x) x self.ln2(x self.dropout(ff_out)) return x这个块是标准的 Post-Norm 结构和 TransformerXL 论文原始结构保持一致。你可以拿这个块堆叠多层再接入下游任务头。为了验证代码能跑通我们可以做一个简单的形状测试if __name__ __main__: torch.manual_seed(42) B 2 seq_len 10 d_model 128 n_heads 8 d_ff 256 block TransformerXLBlock(d_model, n_heads, d_ff, max_len64) x torch.randn(B, seq_len, d_model) # 自回归 mask保证位置 i 只能看到前 i 个 token mask torch.tril(torch.ones(seq_len, seq_len)).bool() out block(x, maskmask) print(输入形状:, x.shape) print(输出形状:, out.shape)跑通之后输出形状应该还是 [B, seq_len, d_model]说明维度一路保持正确。3.5 位置分数的另一种实现方式显式 for 循环上面的 einsum 写法简洁但第一次接触的人可能不太容易立刻看懂“为什么这么一转就对了”。我在调试时通常还会写一个非常朴素的 for 循环版本做数值对拍确保 einsum 没有写错。这里也分享给你也方便你理解相对位置分数的计算过程def compute_pos_score_naive(q_with_v, pos_emb_proj): # q_with_v: [B, n_heads, q_len, head_dim] # pos_emb_proj: [q_len, k_len, n_heads, head_dim] B, n_heads, q_len, head_dim q_with_v.shape k_len pos_emb_proj.size(1) pos_score torch.zeros(B, n_heads, q_len, k_len, deviceq_with_v.device) for i in range(q_len): for j in range(k_len): # 每个位置对 (i,j) 单独求内积 pos_score[:, :, i, j] ( q_with_v[:, :, i, :] * pos_emb_proj[i, j][None, :, :] ).sum(dim-1) return pos_score两个版本的输出应该在数值上完全一致。如果你在自己改代码强烈建议新建一个小测试函数把两种方式对齐一遍确认矩阵乘法的维度映射没有偏差再继续往下走。4. 调试心得与常见坑代码能跑通只是第一步真正把相对位置编码用到自己的项目里你还会遇到各种坑。我挑几个自己实际踩过的写在这里。4.1 相对位置矩阵的方向和范围到底怎么定这是最容易错的地方。你可能会想i-j 还是 j-i其实二者只是互为转置的关系如果你最终结果不对把生成索引的公式从 i-j 改成 j-i 再试一次就知道。关键是索引偏移量i-j 最小是 -(k_len-1)最大是 q_len-1。如果你默认是 q_len 和 k_len 相等那范围就是 [-L1, L-1]恰好需要表长度 2L-1。但如果是交叉注意力q_len 不等于 k_len比如查询 10 个位置键 5 个位置i 取 0 到 9j 取 0 到 4那么 i-j 的取值范围是 -4 到 9。你的位置向量表必须要能覆盖这个范围。所以我在构造相对位置编码时把所有潜在范围都 clamp 到 [-max_len1, max_len-1]虽然这样位置超出时多个相对距离会共享同一个向量但至少不会越界报错。实际使用中建议 max_len 设置得比训练序列长度大出一定余量。4.2 mask 的时机先 mask 还是先和位置分数相加语言模型场景必须设置因果 mask也就是当前位置不能看到未来的 token。常规做法是把注意力分数矩阵的上三角位置填充为 -inf。这里有一个细节位置分数和内容分数是分开算的但最终合并后要先加上位置分数再施加 mask还是先施加 mask 再加位置分数结论是先合并再 mask。因为 mask 的作用是把非法位置的注意力权重在 softmax 之前强制变成 0这个操作必须在合并所有分数之后统一进行否则如果先 mask 掉 content_score再加上 pos_score那些未来位置又会通过 pos_score 偷看到信息产生泄漏。4.3 相对位置编码表要不要参与残差连接不需要。相对位置编码表只服务于注意力分数计算不会像词嵌入那样在输入端和输出端参与残差。有些读者可能一开始会把位置表和 token 嵌入搞混以为也要把位置向量加到词向量上这是绝对错误的。你只要把位置表和 Q、K、V 的权重矩阵当成同类参数来理解就行。4.4 初始化策略对收敛的影响我在上面的代码里用的是torch.randn(...) * 0.02这个 0.02 的经验值主要来自 GPT 系列常用的参数初始化策略。如果你用的是更大的模型建议把位置向量表初始化的方差调小一些否则早期训练阶段注意力分布会过于集中导致某些头退化。更稳妥的做法是有条件的话把相对位置编码表初始化为普通的正态分布然后在训练前 2000 步观察 loss 是否快速下降如果 loss 出现明显震荡把初始化方差再除以 2 试试。4.5 性能问题显式相对位置矩阵真的够用吗我在代码里用了最直观的方式直接构造 [q_len, k_len, head_dim] 的显式相对位置向量矩阵。假设序列长度 L512head_dim64那么这个张量的大小是 512 * 512 * 64 16,777,216约 1677 万个浮点数在 float32 下占用约 67MB 内存。这只是一个头的量。多头情况下会翻倍。如果你在 8 个头、batch size 为 8 的规模下跑这张相对位置矩阵如果单独处理内存压力会比较大但考虑到我们并没有把它广播到 batch 维度所以还勉强能接受。如果序列长度到 1024这个矩阵就是 6700 万浮点约 268MB这就有点吃紧了。TransformerXL 论文为了效率用了一个巧妙的滑动窗口技巧通过把位置向量表做 shift 和截取再和 Q 做某种类似卷积的操作避免显式构造 [L, L, d] 的大矩阵。实际工程里这个方法几乎必用因为长文本场景本来就是 TransformerXL 的主场。不过为了可读性和教学目的我在这篇文章故意保留显式实现它更容易验证正确性。你自己在项目里想上长序列建议参照论文附录 A.2 的实现方式对相对位置分数的计算做优化。4.6 和绝对位置编码相比结果怎么验证如果你把相对位置编码实现好自然会想和经典 Transformer 对比效果。我的建议是不要上来就跑大数据集先用一个小的语言建模任务做对照比如中英文都不超过 50MB 的纯文本语料训练相同步数观察 validation perplexity 的差别。通常相对位置编码在序列长度超过 256 之后优势会越来越明显如果序列很短比如小于 50两者的差距可能很小甚至绝对位置编码还会略占优势因为它的归纳偏置在短序列场景更简单。所以做验证时记得把序列长度调大一点这才是相对位置编码发挥威力的场景。5. 从代码回到论文位置编码之外还有什么很多人在看完 TransformerXL 的相对位置编码后会误以为这就是全部创新点。其实不是。TransformerXL 的另一个核心设计是 segment-level recurrence也就是在相邻两个 segment 之间传递隐层状态让模型能利用极长历史的信息。相对位置编码之所以重要很大一部分原因正是在这种 segment 机制下绝对位置编码会带来混乱两个 segment 的绝对位置会冲突模型没法分清“第 1 段第 5 个 token”和“第 2 段第 5 个 token”的位置关系相对位置编码因为不依赖绝对坐标天然就能迁移到 segment 间的状态复用。我之前刚接触 TransformerXL 只盯着相对位置编码看后来把论文读透之后发现这两块是相辅相成的。没有相对位置编码segment recurrence 的效果会大打折扣没有 segment recurrence相对位置编码长序列优势也发挥不完。所以如果你打算把这段代码接到自己的长文本任务里我建议下一步一定要把 segment 级别的状态缓存加上也就是每一层在计算当前 segment 的注意力时能把上一个 segment 的隐状态拼进来一起算这样就完整复现了 TransformerXL。这个我打算在系列下一篇里详细实现。单看相对位置编码本身用它替换掉经典 Transformer 的绝对位置编码在很多任务上就能获得不错的效果增益。即使你暂时不打算完整实现 TransformerXL把这块抽出来用在别的模型里也是很常见的操作。代码最后我想再给大家一个我在实际项目里的经验相对位置编码在面对“训练长度短、推理长度长”的场景时比绝对位置编码要稳很多但不是完全无痛的外推。如果你想让模型在 2048 长度上推理而训练时最长只有 512那你在训练时最好还是随机截取一些更长片段一起训练或者使用位置表插值。这个观点受限于我自己的实验范围不同的数据集表现会有差异但总体趋势是相对位置编码的外推空间比绝对位置编码大但不要指望它可以任意长度无脑泛化。到这儿TransformerXL 相对位置编码的完整原理和代码实现就都过了一遍。我知道网上关于相对位置编码的讲解不少但很多都停留在公式展示缺少“为什么要这么做”的推导以及“代码里具体怎么写”的落地。这篇希望能帮你把这两块补齐。你在实际实现中碰到过哪些奇怪的问题欢迎带着具体报错和现象来交流我后面继续写 TransformerXL 的其他部分也会尽量多穿插一些工程上的细节。