☰
TransformerXL相对位置编码详解:原理、公式与PyTorch实现
2026/10/3 5:35:07 网站建设 项目流程

如果你已经熟悉经典 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, 2i+1) = 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))

把它展开,会得到四项:

  1. W_q e_i 与 W_k e_j 的点积:纯内容对纯内容的匹配;
  2. W_q e_i 与 W_k p_j 的点积:查询内容与键位置的关系;
  3. W_q p_i 与 W_k e_j 的点积:查询位置与键内容的关系;
  4. 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_pos);
  • u 和 v 是两个可学习的全局偏置向量。

拆开来看,这四个项的语义是:

  1. q_i^T k_j —— 纯内容项,完全基于语义内容打分;
  2. q_i^T W_kR^T R_{i-j} —— 查询内容与“键的相对位置”之间的得分,模型关注“我当前这个 token 和距离我 i-j 的那个 token 在内容上是否相关”;
  3. u^T k_j —— 全局偏置与键内容的得分,表示模型对每个键内容本身的先验偏好;
  4. 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。我会把代码拆成两个类:

  1. RelativePositionEncoding —— 负责生成相对位置向量矩阵;
  2. 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_len=512): 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, device=self.pos_table.device ).view(-1, 1) - torch.arange( k_len, device=self.pos_table.device ).view(1, -1) # 限制在 [-max_len+1, 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_len=512,那表里就有 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_len=512): 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, mask=None): # 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_len=512, dropout=0.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, mask=None): # 子层 1:相对位置多头注意力 + 残差 + LayerNorm attn_out = self.attn(x, x, x, mask=mask) 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_len=64) 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, mask=mask) 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, device=q_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 相等,那范围就是 [-L+1, L-1],恰好需要表长度 2L-1。但如果是交叉注意力,q_len 不等于 k_len,比如查询 10 个位置,键 5 个位置,i 取 0 到 9,j 取 0 到 4,那么 i-j 的取值范围是 -4 到 9。你的位置向量表必须要能覆盖这个范围。所以我在构造相对位置编码时,把所有潜在范围都 clamp 到 [-max_len+1, 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] 的显式相对位置向量矩阵。假设序列长度 L=512,head_dim=64,那么这个张量的大小是 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 的其他部分,也会尽量多穿插一些工程上的细节。

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

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

立即咨询