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

资讯详情

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

Attention Goes Blind:ALiBi在低精度下的数值陷阱

Attention Goes Blind:ALiBi在低精度下的数值陷阱 当模型上下文从 1K 推到 8K、甚至 128K 时一个隐蔽但致命的问题开始浮出水面模型突然“看不见”远端 token 了。不是显存不够也不是收敛失败而是位置编码在低精度计算下悄悄失效。这个问题在 ALiBi 这类基于线性偏置的位置编码中尤为明显我把它称为“Attention Goes Blind”——注意力失明。下面我会从数学原理、PyTorch 复现、指标诊断到修复策略完整拆解这个数值陷阱。1. 背景位置编码为什么重要ALiBi 又是从哪来的1.1 Transformer 无法凭空感知位置Transformer 的核心计算是自注意力Self-Attention它会计算序列中任意两个 token 之间的关联权重。如果不加任何位置信息模型看到的序列和“词袋”没有本质区别把“猫追老鼠”和“老鼠追猫”中的 token 顺序打乱注意力分数完全一致。这对理解自然语言来说是致命的因为语序往往决定语义。所以各类位置编码方案被提出来。早期的 Transformer 使用正弦绝对位置编码后来的 RoPE旋转位置编码把相对位置信息编码进 query 和 key 的旋转角度中T5 的 relative bias 直接给不同相对距离分配可学习偏置。这些方案的目标一致让注意力分数携带“距离远近”的信息让模型知道 token 之间隔了多远。1.2 ALiBi 的设计动机ALiBi 全称是 Attention with Linear Biases由《Train Short, Test Long: Attention with Linear Biases Enables Input Length Extrapolation》提出。它不修改 token embedding也不改变 query 和 key 的投影方式而是直接在注意力 logits 上加一个与相对距离成线性关系的负偏置。距离越远被减掉的分数越多注意力自然更关注近邻。这个操作非常轻量不需要额外参数而且它宣称能做到 train short / test long即训练时序列短、推理时序列长也能保持不错的效果。正因如此ALiBi 被很多开源模型采用尤其适合长文本继续预训练和推理扩展。1.3 什么是 Numerical Failure数值失败听起来很“数学”其实可以通俗理解成在计算机有限精度下计算过程出现了上溢、下溢或精度丢失导致结果偏离真实数学值最终让模型学到错误表示。ALiBi 的偏置是线性增长的上下文长度一长负偏置可能变得非常大。在 FP16半精度浮点下超过一定范围后softmax 的指数函数会直接下溢成 0远端 token 的注意力权重变成 0梯度也不再回传。模型从“能考虑全局”退化成“只看局部”这就是注意力失明的直接表现。理解这个问题需要从 ALiBi 的数学公式和低精度浮点的表达范围讲起。2. ALiBi 原理与计算拆解2.1 标准注意力计算回顾我们以单头注意力为例。输入序列长度为 Lhead 维度为 dquery 矩阵 Q、key 矩阵 K、value 矩阵 V 的维度都是 [L, d]。注意力分数的计算方式是import math import torch import torch.nn.functional as F def standard_attention(q, k, v, maskNone): # q, k, v: [batch, heads, seq_len, head_dim] d q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) probs F.softmax(scores, dim-1) out torch.matmul(probs, v) return out, probs, scores关键点是 softmax 的输入是 logits。logits 的绝对大小会影响 softmax 的输出分布。如果所有 logits 都集中在 [0, 1] 区间输出会比较均匀如果某个 logit 特别大对应概率会被推到接近 1其余概率趋近 0。2.2 ALiBi 的线性偏置公式ALiBi 的做法非常直接。对于第 h 个注意力头给定 query 位置 i、key 位置 j它会在原始注意力分数上增加一个偏置score_alibi(i, j) score_origin(i, j) - m_h * |i - j|其中m_h是该头对应的斜率系数|i - j|是两个位置之间的距离。由于偏置是负数距离越大score 被压得越低。每个头有不同的 m所以不同头对距离的敏感度不同有的头几乎只看临近几个 token有的头则能保留较远距离的信息。这和生物视觉中的多尺度感受野有些类似模型用不同的“分辨率”去扫描序列。如果使用因果掩码还要把未来位置的 logits 置为-inf保证位置 i 只能注意到 j i 的历史 token。2.3 斜率系数 m_h 的生成方式论文中给出的斜率不是随机初始化的而是一个几何级数。以 h 个注意力头为例def get_alibi_slopes(n_heads): 生成 ALiBi 中每个注意力头的斜率系数。 这里使用常见的几何级数2^(-8 * (head_idx 1) / n_heads) def get_slope(head_idx): return 2.0 ** (-8.0 * (head_idx 1) / n_heads) return [get_slope(i) for i in range(n_heads)]当 n_heads 8 时m 依次约为 0.5、0.25、0.125、0.0625、0.03125、0.015625、0.0078125、0.00390625。可以看到第一个头的斜率最大对距离最敏感最后一个头的斜率最小理论上可以照顾更远距离。这里有个容易误用的细节m_h应该作为 logits 层的偏置不是概率层的偏置。有些实现会在 softmax 之后再减效果完全不一样会导致位置信息失效。2.4 与绝对位置编码的差异绝对位置编码是把位置向量加到 token embedding 上RoPE 是把位置编码成旋转矩阵作用于 Q 和 K而 ALiBi 完全不动输入表达只在注意力分数上调整。它的优势在于简单、推理时可以做长度外推但缺点也很明显偏置是固定且无界的。上下文一旦非常长偏置的绝对值就会变得非常大数值风险随之而来。对比之下RoPE 的位置信息是通过旋转角累积的也有自身的精度问题但表现方式完全不同。正因如此遇到 ALiBi 的“失明”问题时不能简单照搬 RoPE 的稳定化手段。3. 数值失败的本质注意力熵与精度坍缩3.1 隐藏的 FP16 陷阱现代大模型训练和推理普遍使用混合精度前向计算中很多张量是 FP16。FP16 的动态范围上限约 65504最小正正规数约 6.1e-5更小的数会下溢成 0。attention logits 经过指数函数后如果 logits 小于约 -11exp(logits) 已经小于 1.7e-5逼近 FP16 的精度极限如果 logits 小于约 -20exp(logits) 已经小于 2e-9在 FP16 下基本就是 0。ALiBi 的偏置第一头是 -0.5 * distance。当 distance20 时偏置就已经是 -10distance40 时偏置是 -20。也就是说在 FP16 下第一个头几乎只能关注到前面 20 到 40 个 token再远的位置softmax 概率直接归零。这就是“远端失明”在数值层面的第一层原因。同时QK^T 的点积也可能很大。head_dim128 时如果 q 和 k 的某个元素在 FP16 下是几十的量级点积很容易超过几百甚至上千softmax 上溢为 inf导致 NaN。不过 ALiBi 场景中更常见的还是负偏置把远端压成 0属于下溢问题。3.2 注意力熵值为什么关键“注意力熵”是描述注意力分布集中程度的指标。假设一个位置对前面 100 个 token 的注意力权重分别为 p1, p2, ..., p100那么注意力熵定义为entropy -sum(p_i * log(p_i))如果注意力分布非常集中比如某个 token 概率接近 1熵值接近 0如果所有 token 概率均匀熵值接近 log(100)。你可以把熵值理解为模型“视野”的量化指标。当 ALiBi 数值下溢发生时远端概率变成 0注意力分布只在近端少数 token 上有非零值熵值会异常低。反过来如果偏置在生产环境中因为精度问题没有生效或者 QK^T 被错误归一化注意力分布可能变得过于均匀熵值异常高。所以训练和推理时监控注意力熵能快速发现“失明”或“无差别注意力”两种极端。3.3 长上下文下的渐进式退化有一个容易忽略的点ALiBi 数值失败并不是突然发生的而是随着上下文长度增加渐进出现的。在短序列中比如 512 token第一头偏置范围是 0 到 -256。虽然远端概率已经很低但还没有完全消失模型还能勉强学到一些长距离信息。可一旦序列长度扩展到 4096 或 8192第一头的远端 token exp(-2048) 在数学上是绝对 0在 FP16 下更加彻底消失。这种退化不是一条平滑曲线而是从“低权重”到“严格零权重”的突变梯度也彻底断开。这就是为什么很多团队做长上下文扩展时会观察到模型“越长越笨”。一开始以为是数据或训练步数不够最后定位到位置编码数值问题损失函数已经无法给远端 token 回传有效梯度了。3.4 填充掩码与偏置的干扰很多实现会同时存在 padding mask 和 causal mask。掩码通常用masked_fill(mask 0, float(-inf))实现。这里要特别小心如果先将 padding 位置置为 -inf再叠加 ALiBi biasbias 的有限值不会覆盖 -inf行为正确但如果顺序反过来先把 bias 加到所有位置再把 padding 位置置为 -inf也没有问题。真正危险的是用 0 代替 -inf会被 softmax 当成有效 logits导致填充位置收到非零注意力。此外在 FP16 下做masked_fill后-inf 依然保留为 -inf但加上一个很大的负 bias 可能出现-inf (-2000) -inf这在 IEEE 浮点下没问题。可如果某个框架将 -inf 表示成极小的负数再叠加偏置后可能变成有限值造成掩码失效这是需要额外注意的实现差异。def build_alibi_bias(n_heads, seq_len, device, causalTrue): 构造 ALiBi bias。 返回形状为 [n_heads, seq_len, seq_len] 的 FP32 张量。 因果场景下未来位置会被置为 -inf。 slopes torch.tensor(get_alibi_slopes(n_heads), devicedevice, dtypetorch.float32) positions torch.arange(seq_len, devicedevice, dtypetorch.float32) row positions.view(1, -1, 1) col positions.view(1, 1, -1) distance (col - row).abs() # [1, seq_len, seq_len] bias -slopes.view(-1, 1, 1) * distance.unsqueeze(0) # [n_heads, seq_len, seq_len] if causal: mask torch.tril(torch.ones(seq_len, seq_len, devicedevice, dtypetorch.bool)) bias bias.masked_fill(~mask, float(-inf)) return bias这里注意把 bias 保持在 FP32。如果你的 attention logits 使用 FP16 计算建议在叠加 bias 之前把 scores 转成 FP32叠加后再决定是否转成 FP16。否则 bias 虽然本身没问题但和低精度 scores 相加时可能引入额外误差。4. 复现实验构造 ALiBi 数值失败4.1 搭建最小注意力模块为了观察数值失败我们先实现一个支持 ALiBi 的通用 decoder attention 模块。这个模块不依赖特定框架只使用 PyTorch可以理解为 seq2seq decoder 中一个 generic attention module 的最小实现。class AlibiAttention(nn.Module): def __init__(self, n_heads, head_dim, dtypetorch.float16): super().__init__() self.n_heads n_heads self.head_dim head_dim self.dtype dtype def forward(self, q, k, v, bias): q, k, v: [batch, n_heads, seq_len, head_dim] bias: [n_heads, seq_len, seq_len] d self.head_dim scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d) # 关键步骤叠加 ALiBi 偏置 scores scores bias.unsqueeze(0) probs F.softmax(scores, dim-1) out torch.matmul(probs, v) return out, probs, scores4.2 在 FP16 下运行长序列下面用随机初始化数据模拟一个 4096 token 的输入。为了让问题更容易复现我们仅观察第一头。torch.manual_seed(42) n_heads 8 head_dim 64 seq_len 4096 batch 1 device cuda if torch.cuda.is_available() else cpu q torch.randn(batch, n_heads, seq_len, head_dim, devicedevice, dtypetorch.float16) k torch.randn(batch, n_heads, seq_len, head_dim, devicedevice, dtypetorch.float16) v torch.randn(batch, n_heads, seq_len, head_dim, devicedevice, dtypetorch.float16) bias build_alibi_bias(n_heads, seq_len, devicedevice, causalTrue) model AlibiAttention(n_heads, head_dim, dtypetorch.float16).to(device) out, probs, scores model(q, k, v, bias)这里要说明随机 q/k 只是为了暴露数值特征不代表真实模型。真实模型中 q/k 是学习出来的分布可能更集中也可能更发散但低精度下的下溢风险依然存在。4.3 观察下溢比例我们要统计 softmax 概率中严格等于 0 的比例这部分 token 在反向传播中梯度为 0不会对输出产生任何影响。def compute_zero_prob_fraction(probs): # 只统计非 mask 位置的概率 zero_frac (probs 0).float().mean().item() return zero_frac zero_frac compute_zero_prob_fraction(probs) print(fAttention prob exactly zero fraction: {zero_frac:.4f})在 FP16 下head 0 因为斜率最大远端概率几乎全部下溢为 0。如果你的随机种子不同数值可能有差异但整体趋势一致序列越长下溢比例越高。将上面的torch.float16换成torch.float32再跑一次你会发现零概率比例明显下降这就是精度影响最直观的证据。4.4 观察注意力熵我们还可以计算不同 head 的注意力熵观察数值失败如何影响“视野”。def attention_entropy(probs): log_probs torch.log(probs 1e-12) entropy -(probs * log_probs).sum(dim-1) return entropy entropy attention_entropy(probs) # [batch, n_heads, seq_len] head_0_entropy entropy[0, 0].mean().item() head_last_entropy entropy[0, -1].mean().item() print(fhead_0 avg entropy: {head_0_entropy:.4f}) print(fhead_last avg entropy: {head_last_entropy:.4f})由于 head 0 的偏置最强注意力集中在非常近的 token 上熵值会比最后一个 head 小很多。如果观察每个位置的熵值随序列长度变化你会看到越后面的位置熵值越低说明远端信息几乎被“剪枝”了。4.5 实验结果解读复现实验告诉我们三件事ALiBi 在数学定义上确实会给远端 token 一个很小的负分数但不是 0在 FP16 下这个很小的分数经过指数函数后变成精确 0。一旦概率为 0反向传播中对应位置的梯度也是 0。这意味着模型完全没有办法从这些远端 token 学习到任何信息。不同 head 受影响程度不同斜率大的 head 失明更早斜率小的 head 还能保留一定长距离能力。所以模型并非完全失明而是部分 head 失明整体表现为“长距离检索能力严重退化”。5. 诊断与排查方法5.1 指标检查清单遇到模型长上下文效果异常时不要急着调学习率或加数据先按下面的指标清单排查指标检查方式异常信号Attention logits 最大/最小值打印每个 head 的 logits 统计出现 inf、NaN 或绝对值超过 5000Softmax 概率零占比统计probs 0的比例比例超过 10% 且随序列长度显著上升注意力熵计算每个 head 平均熵熵值过低说明视野狭窄过高说明位置信息丢失梯度范数检查远端 token 对应位置梯度梯度全为 0 说明远端失明Head 间差异对比不同 head 的注意力分布斜率最大 head 与最小 head 差异过大5.2 逐层打印 q·k 与偏置量级当 logits 异常时需要拆开看是 QK^T 的问题还是 ALiBi bias 的问题。可以在 forward 中临时打印def debug_attention_logits(q, k, bias): d q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d) scores_finite scores[scores ! float(-inf)] bias_finite bias[bias ! float(-inf)] print(QK^T score range: [{:.4f}, {:.4f}].format( scores_finite.min().item(), scores_finite.max().item())) print(ALiBi bias range: [{:.4f}, {:.4f}].format( bias_finite.min().item(), bias_finite.max().item())) combined scores bias.unsqueeze(0) combined_finite combined[combined ! float(-inf)] print(Combined logits range: [{:.4f}, {:.4f}].format( combined_finite.min().item(), combined_finite.max().item()))如果 QK^T 的量级正常而 combined logits 的 min 值远小于 -20那么问题大概率出在 ALiBi bias 的动态范围上。5.3 定位是“溢出”还是“精度丢失”数值问题分两类上溢logits 太大softmax 出现 inf最终产生 NaN。特征是 loss 突然变 NaN检查 logits max 经常超过 65504。下溢logits 太小负得很大softmax 概率被舍入为 0。特征是 loss 没有变 NaN但长上下文效果极差梯度稀疏。上溢通常可以通过 stable softmax 解决在 softmax 前减去每行最大值。下溢问题则更隐蔽stable softmax 只能避免指数计算出 inf不能把极小的概率“找回”成非零值必须从偏置本身入手。5.4 常见错误场景表问题现象常见原因解决思路长上下文 loss 不降、指标下降ALiBi 远端概率下溢为 0降低斜率、限制最大距离、用 FP32 bias训练中出现 NaNQK^T 上溢或梯度爆炸检查 logits 范围使用 stable softmax推理长度外推失效训练长度与推理长度差异过大重新校准斜率或改用 RoPE注意力熵过低所有 head 都只看近邻检查斜率是否被错误放大注意力熵过高bias 没有生效或 type 被转换确认 bias 是否叠加到 logits 而不是概率掩码位置出现非零注意力mask 被覆盖或使用 0 而不是 -inf检查 mask 与 bias 的叠加顺序6. 修复策略与工程实践6.1 调整偏置缩放最简单的修复是限制距离项的最大值。不要直接使用-m * |i - j|而是对距离做截断或对数压缩def build_alibi_bias_clipped(n_heads, seq_len, device, causalTrue, max_distance512): slopes torch.tensor(get_alibi_slopes(n_heads), devicedevice, dtypetorch.float32) positions torch.arange(seq_len, devicedevice, dtypetorch.float32) row positions.view(1, -1, 1) col positions.view(1, 1, -1) distance (col - row).abs() # 关键修改对距离做截断避免线性偏置无限增长 distance torch.clamp(distance, maxmax_distance) bias -slopes.view(-1, 1, 1) * distance.unsqueeze(0) if causal: mask torch.tril(torch.ones(seq_len, seq_len, devicedevice, dtypetorch.bool)) bias bias.masked_fill(~mask, float(-inf)) return bias这种做法会牺牲一部分“严格线性衰减”的性质但在长上下文场景中更稳健。你也可以用log1p(distance)代替线性距离让远端衰减变慢避免偏置过度下探。注意这是工程化的改进不是原版 ALiBi是否采用取决于你的任务对长距离依赖的敏感程度。6.2 分块注意力与 Flash Attention如果问题主要是内存和计算效率Flash Attention 是当前主流选择。它以分块方式计算注意力利用在线 softmax 和重缩放避免完整 [L, L] 分数矩阵驻留显存同时内部统计量通常用更高精度维护。PyTorch 2.0 以后可以直接使用torch.nn.functional.scaled_dot_product_attention传入 additively mask 即可支持类似 ALiBi 的偏置import torch.nn.functional as F def flash_attention_with_alibi(q, k, v, bias): # q, k, v: [batch, n_heads, seq_len, head_dim] scale q.size(-1) ** 0.5 attn_mask bias.unsqueeze(0) # [1, n_heads, seq_len, seq_len] out F.scaled_dot_product_attention( q, k, v, attn_maskattn_mask, dropout_p0.0, is_causalFalse, # 使用自定义 mask避免重复叠加 causal scale1.0 / scale ) return out需要特别强调Flash Attention 解决的是计算效率和中间矩阵爆炸问题并不自动解决 ALiBi 在线性偏置下溢导致的远端概率为 0。它可能在累加时使用更高精度但输入 logits 的动态范围仍然由 ALiBi 的偏置决定。所以在长上下文场景仍然建议配合偏置缩放。6.3 使用更高精度累加混合精度训练中建议把 attention 分数计算保持在 FP32def attention_with_fp32_accumulation(q, k, v, bias): q q.float() k k.float() v v.float() d q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d) scores scores bias.unsqueeze(0) probs F.softmax(scores, dim-1) out torch.matmul(probs, v) return out.half(), probs, scores这种方法能减少 QK^T 累加的数值误差但不会改变 exp 下溢的本质。如果偏置本身已经到了 -2000FP32 下 exp(-2000) 仍然是 0。它更多是“不引入额外误差”而不是“修复偏置动态范围”。6.4 稳定 softmax 的局限与正确用法稳定 softmax 是必备手段def stable_softmax(scores): scores scores.float() max_val scores.max(dim-1, keepdimTrue).values scores scores - max_val exp_scores scores.exp() probs exp_scores / exp_scores.sum(dim-1, keepdimTrue) return probs它的作用是防止 logits 上溢确保最大值对应概率在合理范围内。但注意它只是“平移”了 logits没有改变 logits 之间的差值。远端 token 的 logits 依然远小于近端exp 之后依然可能被舍入为 0。换句话说stable softmax 是安全底线不是修复失明的银弹。6.5 长上下文下的替代方案如果 ALiBi 的数值问题在超长序列中无法通过调参解决可以考虑以下替代方案RoPE旋转位置编码的数值范围相对温和已被大量长文本模型验证。Sparse / Deformable Attention动态选择需要关注的位置避免对所有远端 token 都计算权重。Deformable Attention 通过可学习的采样点聚焦重要区域能缓解线性偏置无界衰减的问题。Sliding Window Global Token用窗口注意力处理近邻配合少量全局 token 保留长距离信息。Coordinate Attention在图像或二维序列任务中把坐标信息显式编码进注意力也能缓解绝对位置注入带来的数值问题。这些方案各有适用场景不要在文章里一次性铺开比较而是要根据你的任务数据分布和推理长度做实验选择。7. 最佳实践与生产建议7.1 配置层面的经验在实际项目中ALiBi 参数通常不参与训练但斜率计算方式和上下文的 max_position 会影响数值稳定性。建议把以下配置显式化attention: type: alibi n_heads: 8 head_dim: 64 bias_dtype: float32 max_distance: null # null 表示使用完整线性距离可以填 512 或 1024 做截断 slope_version: geometric # geometric / learnable / warmup把max_distance、bias_dtype作为可配置项不要硬编码在模型代码里。上线前用不同上下文长度跑一遍注意力指标记录每个 head 的熵和下溢比例形成基准线。7.2 上线前的数值安全测试上线前建议至少做以下几项检查用 FP16 和 FP32 分别跑同一条长文本比较 attention logits 的差异。检查长序列下probs 0的比例是否随序列长度失控。检查每个 head 的平均注意力熵是否落在合理区间。用一个小型下游任务验证长距离 token 的梯度是否正常回传。这些检查不需要很重通常写一个诊断脚本就能跑完。但它们能帮你避免“训练了一周才发现模型根本没有看见远端 token”的尴尬。7.3 监控与报警在模型训练日志中加上 attention 指标输出每 1000 步记录一次def log_attention_health(probs, tag): zero_frac (probs 0).float().mean().item() entropy attention_entropy(probs).mean().item() max_prob probs.max(dim-1).values.max().item() print(f{tag} zero_frac{zero_frac:.4f} entropy{entropy:.4f} max_prob{max_prob:.4f})当zero_frac突然升高时说明可能出现了数值问题或学习率异常当熵值整体过低时说明模型视野收缩需要考虑调整位置编码策略。7.4 安全与权限建议如果你使用的是外部预训练模型库修改位置编码实现时建议先在本地小模型和测试数据集上验证再进入正式训练或微调。不要直接在线上生产模型热更新位置编码逻辑。涉及模型文件覆盖、权重导出等操作务必先备份原始权重并确保有回滚方案。7.5 代码仓库组织建议将一个可复用的 attention 模块独立成文件避免把 ALiBi 逻辑散落在多个地方。推荐目录结构src/ models/ attention/ base_attention.py alibi_attention.py flash_attention.py utils/ attention_diagnostics.py configs/ alibi_test.yaml这样你可以快速切换不同位置编码实现也方便在 A/B 实验中对齐参数。8. 总结ALiBi 是一个简单高效的位置编码方案但它的线性偏置在低精度和超长上下文下存在天然的数值风险。核心机制是距离越大负偏置越大当距离大到一定程度softmax 指数在 FP16 下下溢为 0远端 token 的注意力权重和梯度同时消失模型表现为“注意力失明”。定位这类问题优先检查三个指标attention logits 范围、softmax 概率零占比、注意力熵值。修复手段从轻到重依次是提高偏置计算精度、限制最大距离、使用稳定 softmax、切换到 Flash Attention、最后考虑把位置编码替换成 RoPE 或稀疏注意力。如果你正在用 ALiBi 做长文本训练或推理外推建议把上述诊断脚本和监控指标加入你的工作流。位置编码是模型的“视野”数值问题藏得越深越需要提前用指标把它暴露出来。
返回列表