ARTICLE DETAIL

资讯详情

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

ALiBi位置编码在长文本低精度训练下的数值失效与排查实践

ALiBi位置编码在长文本低精度训练下的数值失效与排查实践 在排查一个长文本抽取任务时我遇到过一个很反直觉的现象模型在短上下文上指标一切正常但一旦把序列长度推进到 4K 以上loss 就出现台阶式恶化。打开 attention map 之后问题一目了然——所有偏锐的注意力头对超过一定距离的 key 的注意力权重并不是“很小”而是精确等于 0.0。这就是 ALiBi 位置编码在低精度、长距离场景下的一种数值失效注意力不是没学会看远处而是它“物理上”看不见远处了。ALiBi 的全称是 Attention with Linear Biases思路非常朴素不学习绝对位置嵌入而是在 softmax 之前的注意力分数上减去一个与 query-key 距离成正比的线性偏置。这个设计在长度外推上很优雅却因为 softmax 的指数特性埋下了一个容易被忽略的“盲区”问题。本文不打算再科普一遍 ALiBi 的原始公式而是从工程排查的角度聊聊它什么时候会“失明”、怎么确认失明、以及有哪些缓解路径。一个比较本质的判断是ALiBi 的数值失效是线性偏置、softmax 指数归一化和低精度浮点三者叠加的结果。它决定了一个模型的有效上下文长度不取决于“位置编码学过多少长度”而取决于“在多少精度下跑、偏置在什么距离上把 exp 压成 0”。1. ALiBi 看上去很优雅但它在 softmax 里埋了一个“距离税”1.1 先回忆一下 ALiBi 到底做了什么ALiBi 的注意力分数计算可以写成这样score(j, i) (q_j · k_i) / sqrt(d) - m_k * (j - i)其中j是 query 的位置i是 key 的位置m_k是第 k 个注意力头的斜率d是 head 维度。在 causal attention 里只允许i j所以这个偏置始终是一个非正数query 离 key 越远分数被压得越低。每个头有自己固定的斜率。常见的 8 头配置里斜率大致是1/2, 1/4, 1/8, 1/16, 1/32, 1/64, 1/128, 1/256也就是一个几何递减序列。越靠前的头越“近视”只关注邻近 token越靠后的头越“远视”可以给较远 token 留出注意力空间。这种设计的精妙之处在于它完全不需要学习位置嵌入只靠距离和斜率的乘法就完成了“近处优先”的归纳偏置。因为距离是相对量模型在推理时可以自然外推到训练长度之外。这也是 ALiBi 最初火起来的原因训练 512 长度推理推到 2048 甚至更长仍然能保持不错的效果。相比之下绝对位置编码在推理长度超过训练长度时往往会直接失效。1.2 softmax 不会对线性偏置做出线性回应问题出在 ALiBi 把距离直接放进了 softmax 的指数里。稍微变形一下就能看清楚exp(score - m * d) exp(score) * exp(-m * d)也就是说距离每增加一个单位未归一化的注意力权重就额外乘以一个exp(-m)。对于一个斜率m0.5的头距离 10 时权重因子已经降到e^-5 ≈ 0.0067距离 20 时只剩e^-10 ≈ 4.5e-5。换句话说这个头天然就只能看见眼前几十个 token。斜率为1/256的最远视头看起来好一些距离 1024 时权重因子是e^-4 ≈ 0.018仍然有内容竞争的余地。可一旦距离拉到 8192就变成e^-32 ≈ 1.3e-14。在数学上它还有值在实际计算里已经离“看不见”不远了。所以 ALiBi 表面上是线性偏置实际产生的是指数级衰减。这里真正值得关注的不是“远端权重很小”而是“小到什么时候变成 0”以及“变成 0 之后对训练和推理意味着什么”。2. 数值失效的三种现场饱和、下溢、梯度死亡“Attention goes blind”我不是第一次见。它通常不是某一行代码写错而是三种效应叠加出来的结果。把它们拆开看会更清楚该从哪里下手。2.1 现场一softmax 饱和注意力变成 one-hot当某个头的斜率偏大、距离偏长时注意力分数矩阵里会同时出现两类值近端 token 的分数接近 0远端 token 的分数因为叠加了巨大的负偏置而非常小。softmax 本来就是对全体 key 做指数归一化于是近端 token 会吃掉几乎全部概率质量远端 token 的权重趋近于 0。这个状态的注意力熵值会快速坍缩。正常情况下一个 head 的注意力分布应该有信息量也就是“我知道该看哪几个位置”饱和之后变成“我只能看到当前位置和紧邻位置”甚至退化成 one-hot 分布。从设计角度看这不算 bugALiBi 本来就要给每个头分配不同的感受野。但问题在于斜率是固定的而实际上下文长度会变。你按 2K 训练长度配的斜率跑到 8K 时最锐的几个头早就越过了自己的“设计距离”提前进入饱和。它们并不是在“学习远处”只是单纯看不见。2.2 现场二低精度下 exp 被数值范围卡死如果说 softmax 饱和是数学上的失效那低精度浮点造成的下溢就是工程上的“实锤”。常规 softmax 实现会先减去最大值来保证数值稳定所以单个分数的绝对值大小不是关键关键是这个分数与最大值之间的差距。一旦某个 key 的分数比最大值低太多exp(score - max)就会低于当前浮点格式能表示的最小正数最终被刷成 0。用常见精度来感受一下精度exp 下溢的大致分数差阈值说明fp16大约 16 到 20超出后结果变成 0 或只剩 subnormal 精度bf16大约 87指数范围接近 fp32但尾数位数少小值精度很差fp32大约 88 到 90 以上仍可表示极小值不直接为 0但梯度已经衰减到无意义把这个阈值和 ALiBi 的斜率放在一起算就能看到“盲区”出现得非常早。以 bf16 为例分数差超过 87 左右时 exp 下溢为 0。对于斜率m0.5的头距离大约 175 就已经到极限对于m1/64的头距离大约 5.5K 也会到极限。也就是说在 bf16 下训练长文本ALiBi 最锐的头部很早就失明了只有最平缓的几个头能勉强维持较远距离的注意力。这里经常有一个误解有人觉得 FlashAttention 内部用 fp32 做 online softmax就不会有下溢问题。注意FlashAttention 确实会用 fp32 维护统计量但最终注意力输出仍然要落回 fp16/bf16 参与后续计算。远端 token 的p_i * v_i贡献在累加过程中远小于其他项经过量化后一样会被抹掉。失效的路径可能从 exp 下溢变成输出量化截断但结果没有本质区别。2.3 现场三梯度死亡远端 key 的 V 贡献归零数值下溢的影响不只是推理时的 attention map 出现硬零更麻烦的是训练时梯度也跟着断掉。softmax 反向传播里某个 key 的 value 梯度会乘上对应的注意力权重p_i。如果p_i精确等于 0那么 loss 对那个 key 的 value 投影的梯度就是 0如果p_i是1e-30这种量级梯度会小到在混合精度下根本没法更新。这意味着一个很隐蔽的后果训练初期如果某几个头已经进入饱和远端 token 的梯度信号从一开始就是断的模型永远学不会“在上文里找答案”。它不是“不想看”而是“反向传播根本没告诉它这里有什么”。再叠加 ALiBi 的固定斜率这个盲区在训练过程中会自我强化——越看不见越学不到越学不到越看不见。很多长文本模型训练到中途 loss 掉不下去检查输入输出又都正常最后发现是位置编码引入的梯度稀疏化在做怪。2.4 容易被忽略的第四现场bias 矩阵本身ALiBi 还有一个偏工程的问题距离偏置矩阵的尺寸是[H, L, L]它会随序列长度二次膨胀。以 8 个头、序列长度 16K 为例fp16 下 bias 矩阵约 4.3 GB拉到 32K就是 17 GB 量级。很多实现会把 bias 提前算好缓存在显存里方便每个 sample 复用。短序列时这是优化长序列时它可能直接成为显存瓶颈。更隐蔽的是精度不一致。有的实现里 bias 用 fp32 生成再 cast 到 fp16有的实现里 bias 用 int 距离乘上 fp16 斜率有的 kernel 要求 dense bias 传入有的则在 kernel 内部 on-the-fly 生成。这些细节都会导致你本地跑出来的 attention map 和生产环境不一致。排查这类问题最容易陷进去的地方就是“明明代码一样为什么换了个 kernel 结果就变了”。3. 怎么确认你的模型真的“盲”了“怀疑注意力失明”和“确认注意力失明”之间需要一套可操作的排查流程。我一般按四步走。3.1 第一步attention map 扫描取一段足够长的真实输入对某个 query 位置画每个 head 的注意力权重随距离变化的曲线。观察两个点是否存在硬零悬崖权重不是平滑衰减而是在某个距离之后直接变成 0。是否存在多头退化多个 head 的失效距离几乎一样说明不是个别头的问题而是精度或斜率在起作用。如果曲线是平滑衰减到很小值说明是数学上的偏置效应如果出现一根明显的垂直截断线说明是数值下溢。3.2 第二步注意力熵值统计注意力熵值是一个比单条曲线更稳定的指标。每个 head 对每个 query 都可以算一个分布熵entropy -sum(p * log(p))一个健康 head 的熵应该维持在一个合理区间即使远端权重偏小也不会瞬间坍缩到接近 0。如果某个 head 在距离超过一定值后熵值几乎等于 0那就是失明的直接信号。更实用的做法是算“等效参与 token 数”exp(entropy)。它表示这个 head 实际有效关注的 token 数量。假如序列长度 8192某个 head 的等效参与数只有 20那它对长上下文基本是摆设。3.3 第三步fp32 与 bf16/fp16 对照同一份输入、同一份权重分别用 fp32 和 bf16 跑一遍推理然后比较 attention map。如果两者差异非常大尤其是在远端出现一边有值一边为零的情况就可以确认是低精度数值路径导致的失明。这一步成本很低却很有区分度。它能帮你判断问题是不是“重写一个更好的 code”就能解决还是必须动位置编码本身。3.4 第四步梯度探针训练阶段可以用一个更直接的办法构造一个“必须看远端才能答对”的任务或者直接对远端 key 的 value 梯度做统计。在 mixed precision 训练日志里加上“远端 key 的梯度范数”这个指标。如果它长期为 0说明反向传播路径已经断了。这个探针不需要每步都打每几百步打一次即可但它能提前预警模型在训练中持续失明而不是等 eval 阶段问题全面暴露。3.5 一个可复用的排查脚本骨架下面是一个简化的示例结构用来生成 ALiBi 偏置并统计注意力零值比例和熵值。不是某个库的官方实现但思路可以直接套用。import torch import torch.nn.functional as F def build_alibi_bias(max_len, num_heads, dtypetorch.float32): 生成 ALiBi 距离偏置形状示例[num_heads, max_len, max_len] if num_heads 8: # 常见的 8 头斜率 slopes torch.tensor( [1/2, 1/4, 1/8, 1/16, 1/32, 1/64, 1/128, 1/256] ) else: # 几何序列生成方式具体以原始实现为准 base 2 ** (-8 / num_heads) slopes base ** torch.arange(1, num_heads 1, dtypetorch.float32) # 距离矩阵第 j 行第 i 列表示 queryj, keyi 的距离 positions torch.arange(max_len, deviceslopes.device) dist positions.unsqueeze(1) - positions.unsqueeze(0) # j - i dist dist.clamp(min0) # causal 场景只保留 j i dist dist.unsqueeze(0).to(dtype) # [1, L, L] bias -slopes.view(-1, 1, 1) * dist # [H, L, L] return bias def attention_probe(q, k, bias, causal_mask): q/k 形状: [B, H, L, D] 返回 attention probs 和零值比例 scores torch.matmul(q, k.transpose(-1, -2)) scores scores * (q.shape[-1] ** -0.5) scores scores bias scores scores.masked_fill(~causal_mask, float(-inf)) probs F.softmax(scores, dim-1) zero_ratio (probs 0).float().mean(dim-1) return probs, zero_ratio实际排查时先短序列跑通脚本再逐步放大长度。否则你可能还没看到注意力失效就先被[H, L, L]的 bias 矩阵占满显存了。3.6 后续要盯紧的几个指标指标怎么算异常信号硬零占比(probs 0).mean()远端出现连续硬零注意力熵值-sum(p * log(p))远端熵接近 0分数动态范围max(scores) - min(scores)差距超过 15~20fp16或 80~90bf16有效上下文长度注意力质量超过阈值的最远距离明显小于训练长度4. 缓解路径不只有换位置编码一条路如果确认模型“盲”了
返回列表