ARTICLE DETAIL

资讯详情

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

ALiBi位置编码数值失效:长文本混合精度下的排查与修复

ALiBi位置编码数值失效:长文本混合精度下的排查与修复 在长文本场景里ALiBiAttention with Linear Biases因为实现简单、不增加参数、天然支持一定程度的长度外推经常被当成位置编码的默认选择。它的核心思路是把位置信息从 token embedding 挪到注意力分数上用一组随 head 变化的负斜率对距离较远的 query-key 对施加线性惩罚。这个设计让很多在 2048 长度上训练的模型可以外推到 4096、8192甚至更长。但真正把序列拉到数万长度、并配合 fp16/bf16 混合精度训练时注意力可能出现一种“失明”现象远端位置的权重整体归零内容相关信号完全被偏置项吞掉softmax 输出退化远端梯度消失严重时日志里直接出现 Inf/NaN。这篇文章从 ALiBi 的数学形式入手分析数值失效的成因然后用一个最小 PyTorch 实现复现问题再给出逐层排查链路和工程修复方案。1. 先理解 ALiBi 的位置信息注入方式才能定位失效点1.1 ALiBi 解决的原始问题长度外推经典的绝对位置编码会在 token embedding 后加上一组位置向量比如正弦编码或可学习位置向量。模型在训练时见过的最大位置 id 是固定的推理时一旦超出这个范围要么截断要么使用未经训练的 embedding效果都会垮掉。相对位置编码虽然能部分缓解这个问题但实现复杂度高而且很多变体仍然需要训练阶段覆盖足够多的相对距离。ALiBi 的思路完全不同它不给 token 加任何位置向量而是在注意力分数上直接减去一个与 token 间距离成正比的惩罚项。因为模型在训练阶段已经学会“近距离优先”推理时即使遇到更长的序列也只是把这种线性惩罚向更远处延伸不需要额外的位置编码参数。这也是它名字里“Linear Biases”的来源位置信息以线性偏置的形式进入注意力计算。1.2 ALiBi 的数学形式与 head 斜率假设 query 在位置 ikey 在位置 j注意力分数的计算可以写成score(i, j) (q_i · k_j) / sqrt(d_k) - m_h * |i - j|其中 d_k 是 query/key 维度m_h 是第 h 个 head 的斜率。在因果语言模型里j 只能取 0 到 i所以 |i - j| 就是 i - j。距离越远扣分越多。slope 不是随便设的而是按几何序列生成。论文里的通用形式是m_h 2^(-8h / H), h 1, 2, ..., H其中 H 是 head 总数。以 H8 为例8 个 head 的斜率分别是2^-1, 2^-2, 2^-3, 2^-4, 2^-5, 2^-6, 2^-7, 2^-8也就是从 0.5 一路衰减到约 0.0039。斜率大的 head 偏向局部斜率小的 head 负责捕获更远距离的信息。这种分工让 ALiBi 在较短训练长度下也能保持一定的长距离建模能力。1.3 偏置量会随着序列长度快速变大ALiBi 的问题在于惩罚项是线性增长的序列一长偏置的绝对值会迅速把内容分数甩开。下面以一个 H8 的模型为例计算不同 head 在不同距离下的偏置大小。head斜率 m_h距离 512距离 4096距离 32768距离 131072head 12^-1 0.5-256-2048-16384-65536head 42^-4 ≈ 0.0625-32-256-2048-8192head 82^-8 ≈ 0.0039-2-16-128-512即使是最“全局”的 head 8在 131072 长度下偏置也到了 -512已经足以让该位置的 softmax 概率趋近于 0。而 head 1 在 131072 长度下偏置是 -65536这个数字已经低于 fp16 能表示的最小值在 fp16 计算中会直接变成 -Inf。除了数值本身还有一个工程问题如果直接把完整的偏置矩阵构建出来显存开销是 O(H·L²)。L32768、H8、fp16 时偏置矩阵大约需要 16GBL131072 时直接到几百 GB根本无法落地。所以 ALiBi 的实际实现要么分段计算要么使用能按需生成偏置的融合 kernel。2. 数值失效的三个机制为什么失明不是模型学坏了而是数值算崩了2.1 softmax 饱和远端权重归零梯度消失注意力权重的计算是softmax(score)而 softmax 对非常大的负输入非常敏感。exp 函数在输入低于约 -87 时在 fp32 下已经接近最小正数再往下就逐步下溢为 0。对 -4096、-16384 这样的输入无论用 fp32、fp16 还是 bf16exp 的结果都是 0。这意味着只要距离足够远远端的 key 在 softmax 里获得的分母贡献就是 0权重也是 0。从数学上看这不是错误而是 ALiBi 设计里“局部优先”的正常结果。问题在于当序列长度远超训练长度时连内容相关分数很高的远端 token 也拿不到任何权重模型想做“全局检索”就完全不可能了。更隐蔽的是梯度问题。注意力权重 w_j 对应 key 位置 j 的权重大约为 0 时输出对 k_j 的梯度也会趋近于 0。结果就是远端 token 不仅在推理时被忽略在训练时也几乎收不到任何学习信号模型对长距离依赖彻底失明。2.2 低精度加法偏置项把内容分数“吃掉”fp16/bf16 混合精度训练下另一个更直接的数值问题发生在加法阶段。假设某个 query-key 对的内容分数是 2.5而 ALiBi 偏置是 -4095那么真正的注意力分数应该是 -4092.5。问题在于浮点数的精度是相对的。bf16 只有大约 8 位有效二进制位在 4096 这个量级上相邻可表示数的间隔是 32fp16 有大约 11 位有效二进制位在 4096 量级上间隔是 4。也就是说2.5 (-4095) -4092.5在 bf16 下会被舍入到接近 -40962.5 这个内容信号完全丢失在 fp16 下会被舍入到接近 -4092内容信号的精度也所剩无几。越是长序列、越是大距离偏置绝对值越大内容分数被“吃掉”得越干净。这就是为什么很多模型在短序列上用 ALiBi 没问题一上长序列就 loss 不降、训练不稳定。问题不一定是模型结构而是低精度下偏置项把有意义的内容相关信号全部覆盖了。2.3 超长上下文中的 Inf/NaN 与显存爆炸fp16 的表示范围大约是 ±65504。head 1 的斜率是 0.5当距离达到 131072 时偏置就是 -65536已经低于 fp16 的最小值计算时直接变成 -Inf。如果这一行里某个位置的偏置是 -Inf、而该行最大值也是 -Infsoftmax 里logit - max会出现-Inf - (-Inf)结果就是 NaN。NaN 通常不会出现在完全正常的因果 attention 里因为每行至少有一个当前 token 自己它的分数通常是有限值。但一旦叠加 padding mask、跨 attention 或自定义 mask整行被 mask 成 -Inf 的情况就可能出现再叠加 ALiBi 的 -Inf 偏置NaN 就会悄悄产生并且随着训练传播到所有参数。显存爆炸则是另一个层面的“失败”。L32768 时单是偏置矩阵就有约 16GB任何模式下都无法直接塞进显存即使勉强塞进去attention score 矩阵和 softmax 中间变量又会再翻几倍。这个限制是 ALiBi 在超长上下文场景落地时真正的拦路虎。注意发现 Inf/NaN 时不要先怀疑数据先检查偏置是否在低精度下越界、mask 是否产生了全 -Inf 行。3. 用最小 PyTorch 实现复现“注意力失明”3.1 环境准备以下复现代码不依赖特殊库CPU 上就能观察 dtype 精度问题。如果要对照 flash-attn 的 dtype 报错再准备 GPU 环境。依赖版本建议说明Python3.9 及以上无特殊要求PyTorch2.0 及以上使用 matmul、softmax、autogradCUDA11.8 及以上可选CPU 也能复现精度问题flash-attn2.x可选用于对照 kernel 行为推荐在干净的虚拟环境里安装 PyTorch避免和其他项目依赖冲突。下面的实验先在 CPU 上跑因为 dtype 转换行为遵循 IEEE 标准CPU 和 GPU 对 fp16/bf16 的表示能力是一致的。3.2 最小 ALiBi 注意力模块先实现两个函数一个生成 slopes一个生成偏置矩阵。import math import torch import torch.nn.functional as F def build_alibi_slopes(num_heads: int) - torch.Tensor: # ALiBi 论文的几何序列m_h 2^(-8h/H), h 从 1 开始 h torch.arange(1, num_heads 1, dtypetorch.float32) slopes torch.pow(2.0, -8.0 * h / num_heads) return slopes def build_alibi_bias(num_heads: int, seq_len: int, devicecpu, dtypetorch.float32) - torch.Tensor: slopes build_alibi_slopes(num_heads).view(-1, 1, 1).to(device) positions torch.arange(seq_len, devicedevice, dtypetorch.float32) # rel_dist[i, j] i - j rel_dist positions.view(1
返回列表