ALiBi(Attention with Linear Biases)在不少开源大模型里出现过,BLOOM、MPT 这些名字都和它绑定过。它的思路简单:不把位置信息塞进 embedding,而是在 attention logits 上加一个与距离线性相关的负偏置,距离越远,logits 被压得越低。训练短序列,推长序列,这是 ALiBi 早期被反复提及的外推优势。
但这里有一个经常被忽略的问题:这个负偏置是随距离线性增长的。序列短的时候,偏置很小,一切正常;序列一旦拉长,比如从 2K 外推到 32K、128K,某个 head 的斜率乘上最大距离,数值量级会非常吓人。再叠加 fp16、bf16 混合精度训练,logits 可能溢出成-inf,或者内容项被偏置项“吃掉”。
这就是标题里说的 Numerical Failure。注意力不是语义上“看不到”远处 token,而是数值上被 ALiBi 偏置直接压死了,我们可以叫它“注意力变盲”。这篇文章把问题拆开讲:ALiBi 的数值结构、偏置量级推演、低精度下的溢出与精度损失、注意力熵值塌缩现象,以及如何复现、诊断和缓解。
1. 核心能力速览
| 项目维度 | 说明 |
|---|---|
| 主题类型 | 位置编码 / 注意力机制数值分析 |
| 核心对象 | ALiBi Positional Encodings |
| 主要问题 | 长序列下偏置量级失控,低精度下出现 inf/NaN、精度损失、注意力熵值塌缩 |
| 涉及技术 | Attention、ALiBi、Positional Encodings、Numerical Failure、Flash Attention |
| 关联方向 | RoPE、Deformable Attention、Coordinate Attention、注意力熵值 |
| 本文目标 | 分析 ALiBi 数值风险,给出复现代码、诊断方法和缓解思路 |
| 适合读者 | 使用 ALiBi 位置编码做长文本训练/推理、混合精度调优、外推长度验证的算法与工程同学 |
要明确一点:这篇文章不是否定 ALiBi。它在中短序列、fp32 计算下依然是一个低成本、可外推的位置编码方案。只是当序列长度和数值精度同时逼近边界时,它会暴露出结构性的数值缺陷。
2. ALiBi 位置编码核心回顾
ALiBi 的核心公式不复杂。对第i个 query 和第j个 key,在因果掩码下(j <= i),attention logits 计算方式为:
logits[i, j] = (q_i · k_j) / sqrt(d) - m_h * (i - j)其中m_h是第h个注意力头独有的斜率。注意这里(i - j)是非负的,所以偏置项是一个非正数。距离越远,偏置越负,softmax 之后该位置的注意力权重越低。
斜率m_h不是通过训练学出来的,而是按几何级数预先设定:
m_h = 2^(-8h / H)其中h从 1 到H,H是注意力头总数。也就是说,第一个头的斜率最大,对距离最敏感;最后一个头的斜率最小,能保留更长距离的信息。
ALiBi 被提出的初衷很明确:
- 不需要训练位置 embedding,省参数;
- 训练长度和推理长度可以不同,具备长度外推能力;
- 结构简单,容易融合进已有的 attention 实现。
如果只看公式,ALiBi 确实简单。但它的数值行为不是免费的。偏置项m_h * (i - j)是随距离线性增长的,一旦距离大,偏置量级就会远超 QK^T 内容项。
3. Numerical Failure 的两种形态
ALiBi 的数值失败不能一概而论,它至少有两种表现形态。第一种是显式溢出,第二种是隐式“失明”。两种形态最终都会导致注意力机制失效,但失效路径完全不同。
3.1 显式溢出:fp16 下的-inf与 NaN
先看 fp16。fp16 的动态范围大约是[-65504, 65504],超过这个范围就变成-inf或+inf。回到 ALiBi 偏置项:
bias = -m_h * distance假设某个头斜率m_h = 0.84,当序列长度为 131072 时,最大距离是 131071:
bias = -0.84 * 131071 ≈ -110100这个值远小于-65504,在 fp16 下直接溢出为-inf。如果QK^T内容项还是有限值,加上-inf之后整条 logits 也变成-inf。softmax 之后,这一行的注意力权重全部变成 0,更严重的可能产生 NaN。
即使序列长度没到 131072,偏置量级也可能逼近 fp16 边界。下面这句话请记住:在 fp16 下,ALiBi 偏置超过 65504 只是时间问题,序列越长越容易触发。
3.2 bf16 下的精度吞噬
bf16 的动态范围比 fp16 大很多,和 fp32 一样是±3.4e38,所以不会轻易溢出成-inf。但它有一个致命弱点:尾数位数太少,只有 7 位有效数字。
当一个量级为-7000的偏置项,加上一个量级在[-20, 20]的 QK^T 内容项时,bf16 可能连 QK^T 的个位变化都表示不出来。结果是:即便模型在语义上认为某些远距离 token 很重要,bf16 精度下这些内容贡献被偏置项“吃掉”,模型实际上只能看到距离,看不到内容。
这就是所谓的大数吃小数。ALiBi 偏置越大,QK^T 内容项的信息在低精度下越容易被吞掉。注意力不再是“内容 + 距离”,而是只剩下“距离”。
3.3 隐式失明:注意力熵值塌缩与梯度消失
即使没有发生-inf溢出,ALiBi 在长距离上的数值压制也会造成另一种失败。softmax 里exp(-large_value)会下溢到接近 0。在 fp16 下,这个下限更早到来。np.exp(-50) 大约是1.9e-22,fp16 最小正数约6e-8,所以exp(-50)在 fp16 下被认为是 0。
当远距离位置的注意力权重变成 0,反向传播时这些位置的梯度也变成 0。模型无法从远距离 token 学到任何信息,长距离建模能力基本丧失。注意力熵值会显著下降,注意力分布集中到少数近距离 token 上,像是一个固定窗口的局部注意力。
这在数值上表现为:
- 注意力权重矩阵大量为 0;
- 注意力熵值明显低于正常水平;
- 梯度稀疏化,远距离位置参数几乎不更新;
- 外推长度越大,困惑度恶化越明显。
通俗地说,就是模型对远处的内容“失明”了。
4. 数量级推演:什么情况下会出事
下面用具体数字推演一下 ALiBi 偏置的量级。以H=32头为例,最大斜率是:
m_1 = 2^(-8/32) = 2^(-0.25) ≈ 0.8409这个斜率并不小。不同序列长度下的最大偏置如下:
| 序列长度 | 最大距离 | 最大斜率(H=32) | 最大偏置绝对值 | fp16 是否安全 | 风险等级 |
|---|---|---|---|---|---|
| 4096 | 4095 | 0.8409 | 3443 | 可表示,但 softmax 已下溢 | 中 |
| 16384 | 16383 | 0.8409 | 13777 | 可表示,但 logits 量级失衡 | 中高 |
| 65536 | 65535 | 0.8409 | 55113 | 接近 fp16 上限 65504 | 高 |
| 131072 | 131071 | 0.8409 | 110210 | 超出 fp16 上限,溢出为-inf | 极高 |
注意,H=32 不是最极端的情况。如果 H=64,最大斜率是:
m_1 = 2^(-8/64) = 2^(-0.125) ≈ 0.9170偏置只会更大。而 H=8 时,最大斜率是 0.5,虽然小一些,但 131072 序列下最大偏置也达到 65535,刚好卡在 fp16 边界附近。
更有意思的是 fp16 下 softmax 的下溢阈值。exp(-16.6) ≈ 6e-8,已经低于 fp16 最小正数。也就是说,当偏置绝对值超过 16.6 时,单个 logits 的 exp 结果在 fp16 下就有下溢风险。以斜率 0.5 计算,距离超过 33 的 token,其注意力权重在 fp16 下已经无法用正常精度表示。如果分母主要被近距离 token 占据,远距离权重基本归零。
注意,这里不是说 softmax 归一化后远距离概率一定是 0,而是说在低精度计算中,远距离位置的贡献和梯度会被压缩到不可用级别。这是 ALiBi 在长序列、低精度组合下的真实风险。
5. 复现实验:用代码看注意力如何变盲
理论推演之外,写一段小代码可以直观看到不同精度下 ALiBi 对注意力分布的影响。下面是模拟代码,生成一个 QK^T 内容项,叠加上 ALiBi 偏置,分别用 fp32、fp16、bf16 计算 softmax,并统计注意力熵值和非零权重数量。
import math import torch import torch.nn.functional as F def alibi_slopes(n_heads: int): # ALiBi 论文中的几何级数斜率 slopes = [2 ** (-8 * k / n_heads) for k in range(1, n_heads + 1)] return torch.tensor(slopes, dtype=torch.float32) def build_alibi_bias(seq_len: int, n_heads: int): pos = torch.arange(seq_len, dtype=torch.float32) rel = pos.unsqueeze(0) - pos.unsqueeze(1) # [seq_len, seq_len] rel = rel.clamp(min=0) # causal 下只保留 j <= i slopes = alibi_slopes(n_heads) # [n_heads, seq_len, seq_len] bias = -slopes.view(n_heads, 1, 1) * rel.unsqueeze(0) return bias.unsqueeze(0) # [1, n_heads, seq_len, seq_len] def attention_entropy(attn: torch.Tensor): eps = 1e-12 # 对最后一个维度计算熵,再对 batch 和 head 取平均 return -(attn * attn.clamp_min(eps).log()).sum(dim=-1).mean(dim=(0, 1)).item() seq_len = 1024 n_heads = 8 bias = build_alibi_bias(seq_len, n_heads) # 模拟 QK^T / sqrt(d) 内容项,方差约为 1 content = torch.randn(1, n_heads, seq_len, seq_len) * 1.0 # 因果掩码 mask = torch.tril(torch.ones(seq_len, seq_len, dtype=torch.bool)) content = content.masked_fill(~mask.unsqueeze(0).unsqueeze(0), float("-inf")) logits_fp32 = content + bias.to(torch.float32) logits_fp16 = content.to(torch.float16) + bias.to(torch.float16) logits_bf16 = content.to(torch.bfloat16) + bias.to(torch.bfloat16) for name, lg in [ ("fp32", logits_fp32), ("fp16", logits_fp16), ("bf16", logits_bf16), ]: attn = F.softmax(lg, dim=-1) nan_count = torch.isnan(attn).sum().item() zero_count = (attn == 0).sum().item() ent = attention_entropy(attn.float()) print(f"{name}: entropy={ent:.4f}, zero_weights={zero_count}, nan={nan_count}")这段代码不依赖真实模型,只是演示 ALiBi 偏置在不同精度下的数值行为。从数值原理可以预期,fp32 下注意力熵值最高,fp16 下熵值最低且零权重数量最多,bf16 介于两者之间但同样会出现 QK^T 精度被吞噬的情况。实际数值会随机种子和模型参数而波动,读者可以在自己的环境里跑一遍。
如果想观察长序列下的显式溢出,可以把seq_len调大,比如 65536 或更大。但注意build_alibi_bias生成的矩阵尺寸是n_heads * seq_len * seq_len,内存消耗增长很快,需要根据本机显存合理调整。
6. 真实场景中的触发条件
ALiBi 数值失败不是只在极端情况下才出现。下面几个场景是实际工程里最容易被触发的。
6.1 训练短、推理长
ALiBi 的外推能力是优点,但外推本身就是把偏置推向更大量级。如果模型在 4096 长度上训练,推理时直接外推到 32768,最大距离从 4095 变成 32767,偏置量级翻了好几倍。即便 fp32 不溢出,QK^T 内容项也更容易被偏置压过,外推质量会明显下降。
6.2 混合精度训练
AMP 训练中 attention 计算经常走 fp16 或 bf16。偏置项如果直接参与低精度计算,就存在 3.1 和 3.2 说的风险。很多模型训练初期没暴露问题,是因为序列长度短、偏置量级小;一旦长序列微调,数值问题就会冒出来。
6.3 Flash Attention 的融合实现差异
Flash Attention 通常会把 ALiBi bias 融合进 kernel 内部,避免显式构建完整的[seq, seq]偏置矩阵。但不同实现的计算路径不同:
- 有的在 fp32 下累加 QK^T 和 bias;
- 有的在 fp16/bf16 下累加;
- 有的对 bias 做了缩放或截断。
如果 Flash Attention 和原生 attention 的结果不一致,很可能就是 ALiBi bias 的数值路径不同导致的。
6.4 大 head 数模型下的小斜率头被忽略
ALiBi 的小斜率头本意是保留长距离信息。但在低精度下,小斜率头的偏置量级虽然不大,logits 的精度仍会被 QK^T 的量化误差影响。而且大斜率头主导了注意力分布,模型更容易退化为局部窗口模式,小斜率头的作用被稀释。
7. 如何缓解 ALiBi 数值失败
缓解不等于彻底解决。不同场景可以选不同方案,但都要注意对模型行为的实际影响。
7.1 对偏置做截断(Clamp)
最简单的工程 hack:给 ALiBi 偏置设置一个上限,超出部分截断。这样长序列下偏置不会无限增长。
def build_clamped_alibi_bias(seq_len, n_heads, max_bias_abs=128.0): bias = build_alibi_bias(seq_len, n_heads) return bias.clamp(min=-max_bias_abs, max=0.0)截断会改变 ALiBi 的原始语义。原本距离越远惩罚越强,截断后超过一定距离的 token 不再受额外惩罚,相当于提前退化为滑动窗口注意力。如果模型本来就是短序列训练,截断对短序列没有影响,但长序列外推行为会变化。
7.2 让 QK^T 和 bias 在 fp32 下相加
如果必须用 fp16 计算 attention,至少先保证 QK^T 与 bias 的加法在 fp32 中完成,之后再把结果转回目标精度。这样能避免 bf16 大数吃小数的问题。
logits = content_float32 + bias_float32 logits = logits.to(compute_dtype) attn = F.softmax(logits, dim=-1)这是成本最低的修复,但如果logits转回 fp16 之后仍然超出范围,显式溢出还是会发生。
7.3 softmax 前检查整行是否为-inf
softmax 在实现时通常会先减 max,但如果某一行全部为-inf,max也是-inf,inf - inf会产生 NaN。在因果掩码和 ALiBi 偏置叠加时,如果实现边界处理不当,整行全为-inf的情况是可能出现的。排查时可以在 softmax 前检查:
if torch.isinf(logits).any(): print("logits contains inf") # 定位是哪些行、哪些位置 inf_mask = torch.isinf(logits)7.4 改用 RoPE 或 NoPE
RoPE 不引入线性距离偏置,它通过旋转矩阵编码相对位置。长序列下 RoPE 也有自己的问题,比如高频维度周期性和外推衰减,但不会出现 ALiBi 这种“偏置项吃掉内容项”的结构性问题。
NoPE(No Positional Encoding)则完全依赖注意力结构本身和训练数据隐式学习顺序。它适合某些任务,但不是通用替代方案。
7.5 滑动窗口 / 分组注意力
如果模型只需要局部上下文,可以用滑动窗口注意力替换远距离的 ALiBi 偏置。窗口内保留内容项,窗口外直接 mask,避免所有 token 计算远距离偏置。
7.6 与 Deformable Attention 对比
Deformable Attention 用可学习的采样偏移来决定每个 query 关注哪些位置,不直接对距离做线性惩罚。它的数值风险主要来自采样坐标学习不稳定,而不是长距离偏置溢出。两者设计哲学不同,Deformable Attention 更适合需要灵活关注模式的任务。
7.7 Coordinate Attention 的定位差异
Coordinate Attention 是另一个方向的注意力机制,它把坐标信息通过通道权重引入特征图,不是对 QK^T logits 做线性距离惩罚,所以不存在 ALiBi 这种 logits 量级失控问题。这里对比只是说明:不同的位置编码/注意力机制,数值风险特征完全不同。
8. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 训练 loss 出现 NaN | ALiBi 偏置在 fp16 下溢出为-inf,或 softmax 出现inf - inf | 打印 logits 的 min/max,检查是否有 inf/NaN | 将 QK^T 和 bias 的加法放到 fp32,或对 bias 做 clamp |
| 长序列外推时困惑度骤增 | 远距离 token 的 logits 被 ALiBi 偏置压得过低,注意力熵值塌缩 | 对比不同序列长度下 attention logits 分布和注意力熵值 | 截断偏置、限制外推长度、改用 RoPE 或 NoPE |
| fp16 下远距离 token 学不到信息 | 远距离注意力权重在 fp16 下下溢为 0,梯度消失 | 统计注意力权重中值为 0 的比例 | softmax 前用 fp32 计算 logits,或降低偏置斜率 |
| Flash Attention 与原生 attention 结果不一致 | 不同实现里 ALiBi bias 的数值路径不同 | 分别打印 logits 的统计量,比较差异 | 统一计算路径,在 fp32 下累加 QK^T 和 bias |
| 训练时 bf16 下长距离建模能力下降 | bf16 尾数不足,QK^T 内容项被大 bias 吞噬 | 计算logits - bias的残差,检查是否等于 QK^T | 用 fp32 计算残差后再加入低精度模型 |
| attention 输出几乎每个位置都集中在局部 token | ALiBi 偏置主导了 softmax 分布 | 计算注意力熵值,观察是否显著低于正常水平 | 截断偏置或改用窗口注意力 |
| 扩长后模型只关注最近几个 token | 大斜率 head 的偏置在长距离下完全压过内容项 | 对不同 head 分别计算注意力熵值 | 关闭或调小大斜率 head 的偏置范围 |
这些排查点都可以落到代码上。最简单的做法是在 attention 计算之后加一段诊断代码:
def diagnose_attention(logits, attn): print("logits min/max/mean:", logits.min().item(), logits.max().item(), logits.mean().item()) print("logits nan/inf:", torch.isnan(logits).sum().item(), torch.isinf(logits).sum().item()) print("attn zero ratio:", (attn == 0).float().mean().item()) print("attn entropy:", attention_entropy(attn.float()))在短序列和长序列上分别跑一次,对比数值分布变化,通常能快速定位问题。
9. 最佳实践与使用建议
9.1 先做长序列压力测试
如果用 ALiBi 训练了一个模型,上线前不要只看短序列指标。建议在目标最长序列上跑一版前向推理,统计 logits 是否出现 inf/NaN、注意力熵值是否塌缩、远距离注意力权重是否大量为 0。这三项是第一批要看的指标。
9.2 保留一个 fp32 基线
混合精度训练时,至少准备一个短序列的 fp32 推理结果作为基线。如果 fp16/bf16 的输出和基线差异过大,优先怀疑数值问题,而不是模型结构问题。
9.3 控制外推长度
ALiBi 支持外推,但外推不是无限的。外推长度越大,偏置量级越失控。建议提前设定安全外推范围,并在该范围内做效果测试。
9.4 偏置与掩码分开处理
因果掩码和 ALiBi 偏置最好不要混在一个张量里计算。掩码用-inf填充,偏置用有限负值,两者分开构建,最后再合并。这样能避免-inf与有限值的边界问题。
9.5 记录注意力熵值指标
在训练和推理日志中加入注意力熵值统计。熵值突然下降时,大概率是某个 head 的注意力分布被偏置压死了。把它当成一个常规监控指标,比等 loss 出 NaN 再排查要主动得多。
10. 总结与下一步
这篇围绕一个现象展开:ALiBi 位置编码在长序列和低精度计算的组合下,会出现注意力数值失效。它可能表现为 fp16 下-inf溢出,也可能表现为 bf16 下 QK^T 内容项被偏置吞噬,更多时候是远距离注意力权重下溢和注意力熵值塌缩,最终让模型“看不见”远处的 token。
从数值结构看,ALiBi 的风险源头是线性偏置项随距离无上限增长。要验证自己的模型是否踩坑,最快的办法是在不同序列长度下检查 attention logits 的 min/max、是否有 inf/NaN、注意力权重零值比例和注意力熵值。如果想修复,优先把 QK^T 与 bias 的加法放到 fp32,考虑对偏置做截断,或者直接切换到 RoPE、NoPE 这类不依赖线性距离惩罚的方案。
这里给一个可以立刻落地的小实验:选一个用 ALiBi 的模型,分别用 2048、8192、32768 的长度跑前向,打印每个 head 的注意力熵值和 logits 分布。如果 32768 长度下熵值明显低于 2048,并且零权重比例大幅上升,说明注意力已经出现“盲区”。这时候再去调整偏置截断或计算精度,效果会非常直观。
ALiBi 不是不能用,而是要知道它的边界在哪里。先把数值风险摸清楚,再决定在什么长度的任务里使用它。