尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

ALiBi数值陷阱:长序列低精度下注意力如何失明?

ALiBi数值陷阱:长序列低精度下注意力如何失明? ALiBiAttention 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 iattention 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 到HH是注意力头总数。也就是说第一个头的斜率最大对距离最敏感最后一个头的斜率最小能保留更长距离的信息。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 时最大距离是 131071bias -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-22fp16 最小正数约6e-8所以exp(-50)在 fp16 下被认为是 0。当远距离位置的注意力权重变成 0反向传播时这些位置的梯度也变成 0。模型无法从远距离 token 学到任何信息长距离建模能力基本丧失。注意力熵值会显著下降注意力分布集中到少数近距离 token 上像是一个固定窗口的局部注意力。这在数值上表现为注意力权重矩阵大量为 0注意力熵值明显低于正常水平梯度稀疏化远距离位置参数几乎不更新外推长度越大困惑度恶化越明显。通俗地说就是模型对远处的内容“失明”了。4. 数量级推演什么情况下会出事下面用具体数字推演一下 ALiBi 偏置的量级。以H32头为例最大斜率是m_1 2^(-8/32) 2^(-0.25) ≈ 0.8409这个斜率并不小。不同序列长度下的最大偏置如下序列长度最大距离最大斜率H32最大偏置绝对值fp16 是否安全风险等级409640950.84093443可表示但 softmax 已下溢中16384163830.840913777可表示但 logits 量级失衡中高65536655350.840955113接近 fp16 上限 65504高1310721310710.8409110210超出 fp16 上限溢出为-inf极高注意H32 不是最极端的情况。如果 H64最大斜率是m_1 2^(-8/64) 2^(-0.125) ≈ 0.9170偏置只会更大。而 H8 时最大斜率是 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, dtypetorch.float32) def build_alibi_bias(seq_len: int, n_heads: int): pos torch.arange(seq_len, dtypetorch.float32) rel pos.unsqueeze(0) - pos.unsqueeze(1) # [seq_len, seq_len] rel rel.clamp(min0) # 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, dtypetorch.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_abs128.0): bias build_alibi_bias(seq_len, n_heads) return bias.clamp(min-max_bias_abs, max0.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 前检查整行是否为-infsoftmax 在实现时通常会先减 max但如果某一行全部为-infmax也是-infinf - inf会产生 NaN。在因果掩码和 ALiBi 偏置叠加时如果实现边界处理不当整行全为-inf的情况是可能出现的。排查时可以在 softmax 前检查if torch.isinf(logits).any(): print(logits contains inf) # 定位是哪些行、哪些位置 inf_mask torch.isinf(logits)7.4 改用 RoPE 或 NoPERoPE 不引入线性距离偏置它通过旋转矩阵编码相对位置。长序列下 RoPE 也有自己的问题比如高频维度周期性和外推衰减但不会出现 ALiBi 这种“偏置项吃掉内容项”的结构性问题。NoPENo 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 出现 NaNALiBi 偏置在 fp16 下溢出为-inf或 softmax 出现inf - inf打印 logits 的 min/max检查是否有 inf/NaN将 QK^T 和 bias 的加法放到 fp32或对 bias 做 clamp长序列外推时困惑度骤增远距离 token 的 logits 被 ALiBi 偏置压得过低注意力熵值塌缩对比不同序列长度下 attention logits 分布和注意力熵值截断偏置、限制外推长度、改用 RoPE 或 NoPEfp16 下远距离 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 输出几乎每个位置都集中在局部 tokenALiBi 偏置主导了 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 不是不能用而是要知道它的边界在哪里。先把数值风险摸清楚再决定在什么长度的任务里使用它。
返回列表