当模型上下文从 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 标准注意力计算回顾
我们以单头注意力为例。输入序列长度为 L,head 维度为 d,query 矩阵 Q、key 矩阵 K、value 矩阵 V 的维度都是 [L, d]。注意力分数的计算方式是:
import math import torch import torch.nn.functional as F def standard_attention(q, k, v, mask=None): # 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 小于约 -11,exp(logits) 已经小于 1.7e-5,逼近 FP16 的精度极限;如果 logits 小于约 -20,exp(logits) 已经小于 2e-9,在 FP16 下基本就是 0。
ALiBi 的偏置第一头是 -0.5 * distance。当 distance=20 时,偏置就已经是 -10;distance=40 时,偏置是 -20。也就是说,在 FP16 下,第一个头几乎只能关注到前面 20 到 40 个 token,再远的位置,softmax 概率直接归零。这就是“远端失明”在数值层面的第一层原因。
同时,QK^T 的点积也可能很大。head_dim=128 时,如果 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 bias,bias 的有限值不会覆盖 -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, causal=True): """ 构造 ALiBi bias。 返回形状为 [n_heads, seq_len, seq_len] 的 FP32 张量。 因果场景下,未来位置会被置为 -inf。 """ slopes = torch.tensor(get_alibi_slopes(n_heads), device=device, dtype=torch.float32) positions = torch.arange(seq_len, device=device, dtype=torch.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, device=device, dtype=torch.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, dtype=torch.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, device=device, dtype=torch.float16) k = torch.randn(batch, n_heads, seq_len, head_dim, device=device, dtype=torch.float16) v = torch.randn(batch, n_heads, seq_len, head_dim, device=device, dtype=torch.float16) bias = build_alibi_bias(n_heads, seq_len, device=device, causal=True) model = AlibiAttention(n_heads, head_dim, dtype=torch.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(f"Attention 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(f"head_0 avg entropy: {head_0_entropy:.4f}") print(f"head_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 或绝对值超过 5000 |
| Softmax 概率零占比 | 统计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 |
| 训练中出现 NaN | QK^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, causal=True, max_distance=512): slopes = torch.tensor(get_alibi_slopes(n_heads), device=device, dtype=torch.float32) positions = torch.arange(seq_len, device=device, dtype=torch.float32) row = positions.view(1, -1, 1) col = positions.view(1, 1, -1) distance = (col - row).abs() # 关键修改:对距离做截断,避免线性偏置无限增长 distance = torch.clamp(distance, max=max_distance) bias = -slopes.view(-1, 1, 1) * distance.unsqueeze(0) if causal: mask = torch.tril(torch.ones(seq_len, seq_len, device=device, dtype=torch.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_mask=attn_mask, dropout_p=0.0, is_causal=False, # 使用自定义 mask,避免重复叠加 causal scale=1.0 / scale ) return out需要特别强调:Flash Attention 解决的是计算效率和中间矩阵爆炸问题,并不自动解决 ALiBi 在线性偏置下溢导致的远端概率为 0。它可能在累加时使用更高精度,但输入 logits 的动态范围仍然由 ALiBi 的偏置决定。所以在长上下文场景,仍然建议配合偏置缩放。
6.3 使用更高精度累加
混合精度训练中,建议把 attention 分数计算保持在 FP32:
def 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 下溢的本质。如果偏置本身已经到了 -2000,FP32 下 exp(-2000) 仍然是 0。它更多是“不引入额外误差”,而不是“修复偏置动态范围”。
6.4 稳定 softmax 的局限与正确用法
稳定 softmax 是必备手段:
def stable_softmax(scores): scores = scores.float() max_val = scores.max(dim=-1, keepdim=True).values scores = scores - max_val exp_scores = scores.exp() probs = exp_scores / exp_scores.sum(dim=-1, keepdim=True) 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 做长文本训练或推理外推,建议把上述诊断脚本和监控指标加入你的工作流。位置编码是模型的“视野”,数值问题藏得越深,越需要提前用指标把它暴露出来。