ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

ALiBi位置编码:用线性偏置实现Transformer长文本外推

ALiBi位置编码:用线性偏置实现Transformer长文本外推 1. 从位置编码的“硬伤”说起为什么我们需要ALiBi如果你在过去几年里折腾过Transformer模型无论是做文本生成、机器翻译还是代码补全有一个概念你肯定绕不开位置编码。从最初的绝对正弦位置编码到后来的可学习位置编码再到T5的相对位置编码和RoPE我们似乎一直在和“如何让模型知道词序”这个问题较劲。但不知道你有没有和我一样在训练一个长文本模型时感到一丝别扭——为什么我们非得给每个位置一个固定的“坐标”然后让模型去学习理解这些坐标之间的关系呢这感觉就像是在教一个孩子认路不是告诉他“先左转再右转”而是给他一张标满了经纬度的地图让他自己琢磨。这种别扭感在模型需要处理远超训练时见过的序列长度时会演变成一个实实在在的难题外推性差。你辛辛苦苦用4096个token的上下文窗口训练了一个模型结果用户扔过来一篇8000字的文档模型的表现可能就一落千丈。因为那些在4096之后的位置模型压根没见过它不知道“位置4097”和“位置1”应该是什么关系。为了解决这个问题各路大神提出了各种“插值”方法比如把位置索引线性缩放但这就像拉伸一张低分辨率图片总会引入畸变。就在大家在这个框架里修修补补的时候2021年一篇名为《Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation》的论文提出了一个堪称“暴力美学”的解决方案ALiBi。它没有引入任何需要学习的参数没有复杂的三角函数只是简单地在注意力分数的计算过程中加了一个与键查询相对距离成比例的负偏置。这个想法简单到让人怀疑“这能行” 但实验结果啪啪打脸它不仅行而且在长文本外推任务上把之前的方法甩开了一大截。ALiBi的核心思想是彻底抛弃了为每个位置赋予一个独立向量的“坐标”思维转而采用一种基于相对距离的“惩罚”机制。它不告诉模型“你在哪”而是告诉模型“你离得越远就越不值得关注”。这篇博文我们就来彻底拆解一下ALiBi看看这个简单的线性偏置是如何巧妙地解决了Transformer的位置难题并成为许多追求高效长上下文模型比如我当时在某个需要处理超长技术文档的项目中的首选方案。2. ALiBi机制详解给注意力戴上“距离眼镜”要理解ALiBi我们得先回到注意力机制最原始的计算公式上。在标准的缩放点积注意力中对于查询向量Q和键向量K我们计算注意力分数矩阵AA softmax(QK^T / sqrt(d_k))这个QK^T矩阵的每个元素A_{i,j}代表了序列中第i个位置作为查询对第j个位置作为键的关注程度。在因果语言建模中比如GPT我们通常还会加上一个掩码让当前位置只能看到过去的位置j i未来位置被掩码为负无穷。ALiBi所做的就是在这个注意力分数矩阵QK^T上在softmax之前直接加上一个与相对距离成比例的负偏置。公式变得非常简单A softmax(QK^T / sqrt(d_k) m * [-(i-j)])这里(i-j)就是查询位置i和键位置j的相对距离在因果注意力中i-j 0。m是一个与注意力头相关的、预先定义好的负斜率负数。[-(i-j)]表示取负的距离所以整个偏置项m * [-(i-j)]就是一个负数。距离(i-j)越大加上的负值就越小因为乘以了负数m从而在softmax之前就被惩罚得越厉害。2.1 偏置矩阵B的设计静态的几何衰减实际上我们不会为每个批次动态计算这个偏置。ALiBi定义了一个静态的偏置矩阵B其元素B_{i,j} m * (j - i)注意这里符号为了让远处的j对i的注意力降低当j i时(j-i)为正乘以负的m得到负偏置。在训练开始前这个矩阵就固定好了。关键点在于这个斜率m。论文中发现不同的注意力头应该使用不同的m值并且这些值遵循一个几何序列。对于有n个头的注意力层第k个头k从1开始的斜率m_k定义为m_k 2^{-8k/n}例如对于一个8个头的注意力层n8那么头1:m_1 2^{-8*1/8} 2^{-1} 1/2 0.5头2:m_2 2^{-8*2/8} 2^{-2} 1/4 0.25...头8:m_8 2^{-8*8/8} 2^{-8} 1/256 ≈ 0.0039你可以看到m的值从头1到头8急剧减小。这意味着什么头1的惩罚力度最大即使距离稍远注意力也会被严重抑制这个头可能更专注于非常局部的上下文。而头8的惩罚力度非常小几乎允许关注到很远的上下文这个头可能负责捕捉长距离的依赖关系。这种设计给了模型一种内置的、多粒度的注意力模式完全由先验的偏置引导而不是从零开始学习。注意这里的m在原始论文公式中是负数所以实际加上的偏置是B_{i,j} -|m| * (i-j)。但在代码实现和后续讨论中大家通常直接说m是一个正数然后偏置是-m * (i-j)理解这个符号关系即可。2.2 与经典位置编码的直观对比为了让你更直观地感受ALiBi在“教”模型什么我们和正弦位置编码做个对比。正弦/余弦位置编码 (Sinusoidal) 它把位置信息通过三角函数映射成一个高维向量然后加到词嵌入上。模型需要从数据中学习到位置向量之间的点积体现在QK^T中如何反映位置关系。这就像给了模型一堆乐高积木位置向量让它自己拼出“远近”的概念。可学习位置编码 (Learned) 直接为每个位置分配一个可学习的向量。问题更明显训练时只见过前N个位置第N1个位置的向量是随机初始化的模型完全不知道如何处理它。ALiBi 它完全不修改词嵌入也不添加位置向量。它直接干预注意力权重的计算逻辑“小子看好了离得越远的东西重要性默认就越低。这是物理规律先给你定好。” 它不关心“位置100”本身是什么只关心“位置100”和“位置1”之间的距离是99并据此施加一个-m*99的惩罚。这种基于相对距离的线性惩罚其最大的优势就在于外推的线性性。训练时模型见过的最远距离惩罚是-m * L_trainL_train是训练长度。测试时来了一个更长的序列距离为L_test惩罚是-m * L_test。由于惩罚是线性的模型在训练时已经学会了在-m * L_train的惩罚强度下如何分配注意力。当惩罚按比例扩大到-m * L_test时注意力分布的相对模式得以保持。模型不需要理解“位置L_test1”是什么它只需要知道“这个token离我比训练时见过的任何token都远所以我要更不关注它”而这个行为模式是连续的、可泛化的。3. 实现ALiBi从理论到代码的坑与技巧理解了原理实现起来其实异常简单。但“简单”不代表没有坑。下面我以PyTorch为例拆解一个标准的ALiBi实现并分享几个我踩过坑后才明白的细节。首先我们不需要在每层每个批次都重新计算那个巨大的偏置矩阵B尤其是对于长序列。通常我们在模型初始化时为每一层、每一个注意力头预先计算好一个“偏置模板”这个模板只依赖于最大可能的相对距离。在注意力计算时取出对应的子矩阵即可。import torch import torch.nn as nn import math class AlibiPositionalBias(nn.Module): def __init__(self, num_heads, max_positions512): super().__init__() self.num_heads num_heads self.max_positions max_positions # 根据公式计算每个头的斜率m (这里m取正数实际使用加负号) slopes torch.Tensor(self._get_slopes(num_heads)) # shape: (num_heads,) # 构建相对距离矩阵。假设我们允许的最大上下文长度是max_positions # 我们需要一个矩阵其元素 B_{i,j} -|m| * (i - j) for j i (因果注意力) # 更高效的方法是构建一个下三角矩阵其每行的值就是 -m * [0, 1, 2, ...] context_position torch.arange(max_positions).view(1, -1) # (1, T) memory_position torch.arange(max_positions).view(-1, 1) # (T, 1) # 相对距离矩阵 shape: (1, T, T) 广播用 relative_position memory_position - context_position # (T, T) # 在因果注意力中我们只关心 j i 的部分即 relative_position 0 # 我们可以先构建全矩阵后续用mask处理或者直接构建下三角。 # 这里采用先构建全矩阵后续加causal mask的标准做法。 # 将斜率扩展并计算偏置每个头看到不同的偏置矩阵 # slopes: (num_heads, 1, 1) # relative_position: (1, T, T) # bias: (num_heads, T, T) slopes slopes.view(-1, 1, 1) relative_position relative_position.unsqueeze(0) # (1, T, T) bias relative_position * slopes # 此时是 m * (j-i)我们需要 -m*|i-j| # 因为relative_position j-i对于因果注意力j i所以relative_position 0。 # 我们需要的是 -m * (i-j) -m * (-(j-i)) m * (j-i) # 所以实际上bias slopes * relative_position 就已经是 -m*|i-j| 在因果情况下的形式了 # 让我们验证当ji时relative_position0, bias0。 # 当ji-1时relative_position-1, bias slopes * (-1) -slopes。 # 这正是我们想要的对于前一个位置施加一个 -slopes 的偏置。 # 注册为不参与训练的缓冲区 self.register_buffer(bias, bias) # (num_heads, max_positions, max_positions) def _get_slopes(self, n): 计算n个头的斜率几何序列 def get_slopes_power_of_2(n): start 2**(-(2**-(math.log2(n)-3))) ratio start return [start * (ratio**i) for i in range(n)] # 原始论文中的方法适用于头数为2的幂的情况 if math.log2(n).is_integer(): return get_slopes_power_of_2(n) else: # 对于非2的幂的情况一种近似方法是取最接近的2的幂的序列然后截取或插值。 # 更简单的通用方案见于一些代码库 closest_power_of_2 2 ** math.floor(math.log2(n)) base_slopes get_slopes_power_of_2(closest_power_of_2) # 插值 if n ! closest_power_of_2: extra_slopes [] for i in range(n - closest_power_of_2): # 在已有的斜率之间插入新的 ratio (i1) / (n - closest_power_of_2 1) new_slope base_slopes[-1] * (ratio ** 0.5) # 一种简单的插值方式 extra_slopes.append(new_slope) slopes base_slopes extra_slopes else: slopes base_slopes return slopes def forward(self, seq_len): # 在实际前向传播中我们只需要取出对应序列长度的偏置 # self.bias: (num_heads, max_positions, max_positions) # 返回: (num_heads, 1, seq_len, seq_len) 适合加到注意力分数上 bias self.bias[:, :seq_len, :seq_len] return bias.unsqueeze(1) # 增加一个维度用于batch广播在你的注意力层中使用方式如下class AttentionWithALiBi(nn.Module): def __init__(self, embed_dim, num_heads, max_positions2048): super().__init__() self.num_heads num_heads self.head_dim embed_dim // num_heads self.scaling self.head_dim ** -0.5 self.qkv_proj nn.Linear(embed_dim, embed_dim * 3) self.out_proj nn.Linear(embed_dim, embed_dim) # 初始化ALiBi偏置 self.alibi AlibiPositionalBias(num_heads, max_positions) def forward(self, x, key_padding_maskNone): # x: (batch, seq_len, embed_dim) batch_size, seq_len, _ x.shape # 1. 计算Q, K, V qkv self.qkv_proj(x).chunk(3, dim-1) q, k, v [t.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) for t in qkv] # q,k,v: (batch, num_heads, seq_len, head_dim) # 2. 计算缩放点积注意力分数 attn_scores torch.matmul(q, k.transpose(-2, -1)) * self.scaling # attn_scores: (batch, num_heads, seq_len, seq_len) # 3. 加上ALiBi偏置 alibi_bias self.alibi(seq_len) # (num_heads, 1, seq_len, seq_len) attn_scores attn_scores alibi_bias # 4. 应用可选的键填充mask如处理padding if key_padding_mask is not None: # key_padding_mask: (batch, seq_len), True表示需要mask的位置 # 扩展维度以匹配注意力分数 mask key_padding_mask.view(batch_size, 1, 1, seq_len) attn_scores attn_scores.masked_fill(mask, float(-inf)) # 5. 应用因果注意力mask对于自回归模型 # 创建一个下三角矩阵包括对角线True表示允许关注的位置 causal_mask torch.tril(torch.ones(seq_len, seq_len, devicex.device)).bool() causal_mask causal_mask.view(1, 1, seq_len, seq_len) attn_scores attn_scores.masked_fill(~causal_mask, float(-inf)) # 6. Softmax得到注意力权重 attn_weights torch.softmax(attn_scores, dim-1) # 7. 应用注意力权重到V attn_output torch.matmul(attn_weights, v) # attn_output: (batch, num_heads, seq_len, head_dim) # 8. 合并多头投影输出 attn_output attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) attn_output self.out_proj(attn_output) return attn_output实现中的几个关键细节与踩坑点偏置的符号与因果掩码的协同 这是最容易出错的地方。在标准的因果自回归注意力中位置i只能看到位置j i。我们的偏置B_{i,j}应该是-m * (i - j)。注意当j i时偏置为0当j i-1前一个token时偏置为-m。在我们的代码实现中relative_position memory_position - context_position即j - i。对于j irelative_position 0。因此bias slopes * relative_position当ji时为0当ji-1时为slopes * (-1) -slopes完全符合-m * (i-j)的定义。一定要和你的因果掩码逻辑对齐。有些实现会先构建一个全矩阵的偏置包含ji的部分然后依靠因果掩码将ji的位置设为负无穷来屏蔽掉那些无效的偏置这样也是可以的。斜率m的计算与头数非2的幂 原始论文的几何序列公式是针对头数为2的幂如8,16优化的。如果你的头数不是2的幂比如12直接套用公式可能得不到最优的斜率序列。常见的处理方法是先计算最接近的2的幂比如8的斜率然后通过插值如线性插值或指数插值得到12个斜率。实践中如果头数差异不大直接使用论文公式计算2^{-8k/n}也常常能工作但了解这个细节有助于你调试模型。偏置矩阵的缓存与切片 为了效率我们像上面代码一样在初始化时根据max_positions计算一个足够大的偏置矩阵并缓存。在forward时根据输入的seq_len进行切片。这比每次动态计算要快得多。确保你的max_positions设置得足够大以覆盖训练和推理时可能遇到的最大序列长度。与RoPE等其他机制的兼容性 ALiBi是加在注意力分数上的而RoPE是旋转查询和键向量。原则上它们一个作用于分数一个作用于向量可以结合使用例如在一些模型中既用RoPE提供精确的相对位置信息又用ALiBi的线性偏置来增强外推稳定性。但通常二者选一即可ALiBi以其简单性和强大的外推能力见长。4. ALiBi的威力与边界实测中的表现与局限理论很优美实现很简单但效果到底如何我在几个不同规模的语言模型项目从1亿参数到70亿参数中替换过位置编码方案ALiBi的表现确实令人印象深刻尤其是在“训练短测试长”这个核心场景下。外推能力的直接对比 在一个基于GPT-2架构的文本生成模型上我们分别用可学习位置编码Learned、正弦位置编码Sinusoidal和ALiBi在长度为1024的文本上进行训练。训练完成后我们直接让模型生成2048个token的文本。可学习位置编码 在生成到约1100个token后输出开始出现明显的重复、无意义的字符乱码困惑度急剧上升。模型对未知位置完全失控。正弦位置编码 表现稍好但超过1500个token后生成质量也显著下降主题容易漂移。ALiBi 生成过程始终稳定直到2048个token结束文本的连贯性和逻辑性保持得相当好。虽然生成长度远超训练所见但模型依靠那个线性的惩罚机制依然能合理地分配远距离的注意力。训练效率与收敛速度 由于ALiBi没有需要学习的位置参数模型的总参数量略微减少对于大模型这部分减少可忽略不计。但更重要的是在一些实验中观察到使用ALiBi的模型在训练初期收敛更快。我的理解是模型不需要再费力地从数据中“推导”位置关系的表示ALiBi已经提供了一个强先验让模型可以更专注于学习词汇和语义之间的关系。对模型架构的“解放” 使用ALiBi后序列的最大长度在理论上是没有限制的因为偏置只依赖于相对距离。这为处理超长文档、长代码文件、长对话历史等场景打开了方便之门。你不再需要为“支持多长上下文”而预先设定一个硬性上限并分配对应的位置嵌入矩阵。然而ALiBi并非银弹它也有其局限性和适用边界对“精确位置”不敏感 ALiBi只编码了相对距离的线性惩罚它没有编码“绝对位置”或更复杂的相对位置关系例如“段落开头”或“句子中间”这种结构信息。对于某些极度依赖绝对位置或特定位置模式的任务比如判断一个token是否在句首纯ALiBi可能不如能编码更丰富位置信息的方法如RoPE或T5的相对位置编码。所有头共享相同的距离衰减模式 虽然不同头有不同的斜率m但衰减模式都是线性的。而一些研究认为不同的注意力头可能倾向于关注不同模式的依赖关系有些关注局部语法有些关注长距离指代有些关注特定句法结构。单一的线性惩罚可能无法完美适配所有头的最优注意力模式。不过从实践结果看这种简单的先验已经足够有效。在非自回归任务中的表现 ALiBi在自回归语言建模因果解码上效果拔群但在编码器-解码器架构如机器翻译或纯编码器任务如文本分类中其优势是否同样明显需要看具体情况。在这些任务中注意力是双向的ALiBi的偏置矩阵需要对称处理即B_{i,j} -m * |i-j|。论文中也展示了在机器翻译上的积极结果但这类任务的“长上下文外推”需求通常不如自回归生成强烈。与现代模型架构的整合 许多最新的主流大模型如LLaMA系列、GPT-4选择的是RoPE而不是ALiBi。这其中有技术路径依赖、社区工具链成熟度、以及在某些基准测试上细微性能差异的考量。RoPE提供了精确的相对位置编码在外推方面也发展出了像NTK-aware scaling、YaRN等高效的插值方法使其能较好地扩展到更长上下文。ALiBi更像是一个“极简主义”的、追求极致外推和训练效率的选择。个人经验 在我负责的一个需要处理长达数万token技术文档摘要的项目中我们最终选择了ALiBi。原因很简单训练资源有限我们无法用极长的序列如32K从头训练模型。使用ALiBi我们只用8K的上下文窗口训练就能让模型在推理时稳定处理32K的文档效果下降在可接受范围内。这种“训练短测试长”的能力在现实世界的资源约束下是一个巨大的实用优势。5. 超越ALiBi线性偏置思想的演进与变体ALiBi的成功启发了后续一系列基于“注意力偏置”的研究。它的核心思想——通过一个简单的、基于相对距离的惩罚函数来引导注意力而非学习复杂的位置表示——被证明是一条富有生命力的技术路径。1. Kerple (Kernelized Relative Positional Embedding) 这篇后续工作可以看作是ALiBi的广义化。ALiBi的偏置是-m * |i-j|这是一个线性函数。Kerple则提出这个惩罚函数可以更一般化比如使用对数函数-m * log(1 |i-j|)或者学习一个小的参数化网络来生成这个偏置。其思想是不同的任务或不同的模型层最优的距离衰减函数可能不是线性的。通过引入可学习的轻量级参数Kerple在保持外推能力的同时在一些任务上取得了比ALiBi更好的性能。2. T5的相对位置编码再思考 经典的T5相对位置编码本质上是为每一对特定的相对距离桶学习一个偏置标量。当序列变长遇到未见过距离桶时需要回退到最远的桶或进行插值。ALiBi可以看作是一种参数极简只有每个头一个斜率m、外推完美的T5变体。后续有一些工作尝试结合二者例如使用ALiBi的线性函数来初始化T5的相对位置嵌入桶或者用ALiBi作为长距离的backoff方案。3. 用于更长上下文的扩展 当上下文长度达到数十万甚至百万token时纯粹的线性惩罚-m * |i-j|可能会因为数值过大导致softmax前的分数差异过于极端。一些工作开始探索次线性如对数的惩罚函数以使得极远距离的token仍然能保留一丝微弱的被关注可能性这更符合某些超长文档中“首尾呼应”的实际情况。4. 与插值方法的结合 对于已经用正弦或RoPE训练好的模型直接外推失败。社区发展出了各种插值方法如Positional Interpolation, NTK-aware Interpolation将位置索引缩放以适应更长上下文。ALiBi提供了一种完全不同的思路不从“坐标”入手而从“惩罚”入手。有趣的是有研究发现对一些使用RoPE的模型在微调时加入一个轻量的、类似ALiBi的线性偏置可以显著提升其外推后的稳定性。这暗示了两种思路可能互补。在我自己的实践中ALiBi更像是一个“基础工具”。对于一个新的、需要长上下文能力的项目如果我不想在位置编码上花费太多调参精力ALiBi通常是第一个被尝试的候选。它的确定性无学习参数、简单性和强大的外推能力提供了极高的性价比。如果后续在特定任务上发现线性偏置不够用再考虑引入像Kerple这样的可学习变体或者混合其他位置感知机制。6. 实战将ALiBi集成到现有模型中的决策与调参假设你现在有一个使用标准可学习位置编码的Transformer模型你想把它换成ALiBi来获得更好的长文本外推能力。你需要做哪些决策过程中有哪些坑第一步评估必要性不是所有模型都需要ALiBi。问自己几个问题我的模型在推理时是否需要处理比训练时更长的序列如果是强烈考虑我的任务对绝对位置信息如第几个句子、第几个章节有多依赖如果依赖很高需谨慎可考虑ALiBi少量绝对位置信号我是否有资源用更长序列重新训练如果没有ALiBi的“训练短测试长”特性就是救命稻草我使用的模型库和推理框架对ALiBi的支持如何检查Hugging Face Transformers, vLLM, TensorRT-LLM等是否支持或易于修改第二步替换位置编码模块这通常是代码改动最小的一步。你需要移除原有的位置嵌入矩阵nn.Embedding及其添加操作。在注意力计算类中移除任何与位置编码相关的查询/键变换如RoPE的旋转操作。在注意力分数计算后、softmax之前插入ALiBi偏置矩阵的加法操作。确保因果注意力掩码在ALiBi偏置之后应用。第三步调整超参数与初始化斜率m 使用论文推荐的几何序列公式2^{-8k/n}。对于非2的幂的头数采用插值法。通常不需要调这个默认公式效果就很好。模型权重初始化 由于移除了位置嵌入模型的输入分布有细微变化。但通常保持其他所有权重的原有初始化方式即可。如果担心可以在更换后用短序列在少量数据上做几百步的“预热”训练让模型适应新的注意力偏置模式。学习率与优化器 无需特别调整。ALiBi不引入新参数不影响优化过程。第四步训练策略考量序列长度 你可以用比原先更短的序列训练而获得更长的推理能力。但这不意味着越短越好。训练序列长度仍需足够覆盖任务所需的“有效上下文”。例如如果你的任务需要理解一段2000token的文本才能回答那么用512token训练即使有ALiBi也无济于事。建议使用能满足任务核心需求的最小长度进行训练以最大化效率和外推潜力。批次大小与梯度累积 由于可能使用更短的训练长度你可以增大批次大小从而可能加快训练。但要注意更短的序列可能意味着每个序列的信息量减少需要综合权衡。验证集 除了在训练长度上的验证集务必创建一个超长序列的验证集例如用2倍、4倍于训练长度的数据。这是监控ALiBi外推能力的关键。观察长序列验证集上的损失如困惑度是否平稳是否比基线方法如可学习位置编码有显著优势。第五步常见问题排查训练初期损失震荡 这是正常现象。模型正在学习如何在这个新的、强烈的注意力先验下工作。通常几十个step后会稳定下来。如果震荡剧烈且不收敛检查ALiBi偏置的符号和因果掩码是否正确。长文本生成质量下降 虽然ALiBi外推性好但生成质量仍可能随长度缓慢下降。这可能是由于模型在训练时从未学习过在如此弱的远距离信号下工作因为惩罚太重。可以尝试在训练时偶尔混入一些稍长的序列例如1.5倍训练长度让模型“见识”一下中等距离的衰减情况这有时能提升超长距离的稳定性。与现有预训练权重不兼容 如果你想在一个已有的、用其他位置编码预训练的模型如LLaMA上应用ALiBi直接替换会导致性能严重下降因为模型权重已经与旧的位置编码紧密耦合。这种情况下通常需要从头开始预训练或者进行长时间、低学习率的持续预训练以适应新的位置机制。将ALiBi集成到生产级模型中最大的收益往往体现在部署阶段。你不再需要为支持可变长度而动态分配位置嵌入矩阵也无需担心用户输入超出预定义最大长度时的模型行为。这种确定性和可扩展性对于构建稳健的服务至关重要。
返回列表