面向块稀疏注意力的不确定性门控选择

arXiv cs.LG 论文

摘要

提出了一种不确定性门控路由器,对于截止边际不确定的查询,将选中的关键块数量加倍,从而提升长上下文语言模型中块稀疏注意力的召回率和准确率,并在多种架构上得到验证。

arXiv:2607.07724v1 公告类型:新 摘要:块稀疏注意力通过用每个查询的关键块 top-k 选择替代 O(N^2) softmax 来扩展长上下文语言模型。这种截止是短视的:当第 k 块和第 (k+1) 块的得分几乎持平时,选择器会在不增加预算的情况下做出决定,而一个携带答案证据的被丢弃块在下游将无法恢复。我们提出了一种信息价值路由器,它能衡量每个查询的 top-k 截止判定有多果断,并针对差距最小的那些查询将保留集合加倍;该规则与骨干网络无关,并且可以与现有的块评分方法(如 Quest)叠加使用。在 LongBench-v2 medium 的 n=215(整个数据集子集)上,路由器加持的 Quest 达到了 0.75 的配对召回率,而 top-k 仅为 0.47 —— 比 SSA 风格的基线提高了 28 个百分点(McNemar p<0.01)—— 并且在相同上下文下的 RULER NIAH 多键测试中,其效果与密集注意力相差 2 个百分点以内。这种提升在来自三种架构的四个模型(Qwen2.5、Mistral-Nemo、Qwen3.6)上得到了复现。在 128K 上下文下,该路由器在 Qwen2.5-7B-1M 和 Qwen3.6 上分别保留了密集注意力准确率的 0.81 和 0.89(而 SSA 风格的 top-k 在前者上仅为 0.09),同时融合的选择加核流水线运行时间为密集注意力的 0.62 倍和 0.80 倍(挂钟时间)。
查看原文
查看缓存全文

缓存时间: 2026/07/10 06:13

# 不确定性门控选择用于块稀疏注意力

代码、数据和复现脚本:https://github.com/ThomasRossi/uncertainty-gated-block-sparse-attention。永久存档:doi:10.5281/zenodo.20630587(概念DOI;始终解析到最新版本)。基于信息价值的SSA风格选择器视图,来源:https://arxiv.org/html/2607.07724

###### 摘要

块稀疏注意力通过将 \(O(N^2)\) softmax 替换为每个查询对键块的 top-\(k\) 选择来扩展长上下文语言模型。这种截断是短视的:当第 \(k\) 个和第 \(k+1\) 个块的得分几乎相同时,选择器会在不花费额外预算的情况下做出决定,而包含答案证据的被丢弃块在下游无法恢复。我们提出了一种基于信息价值的*路由器*(router) ,它针对每个查询衡量 top-\(k\) 截断的决定性程度,并对那些差距最小的查询加倍保留的块集;该规则与骨干网络无关,并且可以叠加在现有的块评分方法(如 Quest)之上。在 LongBench-v2 medium 数据集(\(n=215\),即整个数据集子集)上,router-on-Quest 的成对召回率达到 0.75,而 top-\(k\) 为 0.47——比 SSA 风格基线提升 +28 个百分点(McNemar \(p<0.01\))——并且在相同上下文的 RULER NIAH multikey 任务中,性能与全注意力仅差 2 个百分点。该提升在来自三种架构的四个模型(Qwen2.5、Mistral-Nemo、Qwen3.6)上均可复现。在 128K 上下文下,路由器在 Qwen2.5-7B-1M 和 Qwen3.6 上保留了全注意力准确率的 0.81 和 0.89(而 SSA 风格 top-\(k\) 在前者上仅为 0.09),同时融合的选择加核流水线在全注意力 wall time 的 0.62 倍和 0.80 倍下运行。

## 1 引言

长上下文语言模型越来越多地使用*块稀疏注意力*作为 \(O(N^2)\) softmax 的即插即用替代方案。这一想法被 Quest (Tang et al., 2024)、H2O (Zhang et al., 2023)、SnapKV (Li et al., 2024)、MInference (Jiang et al., 2024)、NSA (Yuan et al., 2025)、MoBA (Lu et al., 2025) 和 Subquadratic 的 SSA (Subquadratic, 2025) 所共享:一个廉价的每查询*选择器*从 \(N_B\) 个键块中挑选 \(k\) 个;然后精确注意力仅在选中的集合上运行。选择器是控制杠杆——它决定了模型关注的位置。

这是一个短视的决定。当第 \(k\) 个和第 \(k+1\) 个块的得分几乎相同时,top-\(k\) 会默不作声地打破平局并继续;如果被丢弃的块包含证据令牌,答案就丢失了,并且任何下游层都无法恢复它。这种失败模式是结构性的,而非噪声性的:它对*多跳*和*查询潜在*检索的打击最为严重,因为一个块的相关性取决于同一前向传播中先前学习到的内容,而选择器表面的查询-键匹配无法看到这一点。SSA 自己报告的 NIAH 多键召回率随键数量增长而急剧下降 (Subquadratic, 2025)。

**本文。** 我们将 top-\(k\) 截断视为一个*信息价值*(Value-of-Information, VoI) 决策,并在其之上添加一层策略。对于每个 Q-tile 和注意力头,我们计算归一化的截断间隔 \(\sigma = \frac{s_{(k-1)} - s_{(k)}}{s_{(0)} - s_{(k)}} \in [0,1]\),其中 \(s_{(\cdot)}\) 是排序后的块得分。小的 \(\sigma\) 意味着截断是高风险;然后我们将该 tile *路由*到 \(2\times\) 扩展的 kv_idx——它能够关注更多块——而置信的 tile 保持基准预算 \(k_{\text{budget}}\)。扩展是有选择地付出的:我们每层只触发底部 \(q\) 分位的 tile,因此平均关注集每行仅增加 \(1+q\) 个块。该路由器与骨干网络无关。

截断间隔 \(\sigma\) 是排序后块得分的函数;它不依赖于这些得分如何计算。现有的选择器在块评分骨干网络方面有所不同——SSA 使用平均池化的键(\(\bar{k}_b = \frac{1}{B_n} \sum_j k_j\)),Quest 使用 min/max 上界(\(s = \sum_d \max(q_d k_{b,d}^{\max}, q_d k_{b,d}^{\min})\))。路由器可以放在任何方式之上。这将贡献从*对 top-\(k\) 的替换*转变为*一个通用的预算分配层,叠加在任何最适合任务的评分骨干网络上*。我们通过实验验证了这一点:更好的评分(Quest)和更好的预算分配(路由器)是正交方向;将两者结合在我们测试的两个基准上都严格优于单独使用任何一个。

**贡献。**

1. 一个 VoI 形式的选择器截断公式,它为每个 tile 添加一个标量,每层添加一个分位数阈值,与块评分骨干网络无关。无需重新训练,无需额外参数,无需每行的保持张量。

2. 实验证明路由器可以与两种不同的评分骨干网络(SSA 风格的 K-均值 和 Quest 的 K-min/K-max)组合,并且,在面板中的每个模型上,*胜出*骨干网络的路由器提升版本严格优于另一种骨干网络的未提升版本。该结果在两个标准化基准(RULER NIAH 多键和 LongBench-v2 medium)上、来自三种架构类别的四个模型上以及从 32K 到 128K 的上下文中均成立。哪个骨干网络胜出取决于模型——QK-Norm 将胜出者从 Quest 的 K-max 翻转为 SSA 风格的 K-均值——但路由器会提升任何胜出的骨干网络。完整面板见第 4 节。

3. 一个融合的选择加核实现,直接为每个 Q-tile 生成 kv_idx,并在同一代码路径中处理路由扩展,保持核调度的形状统一。所有四种稀疏策略(top-\(k\)、路由器、Quest、router-on-Quest)都在同一个核上运行;它们仅在选择步骤上有所不同。在 Qwen2.5-7B-1M 上,wall-time 剖面在 32K 和 64K 之间越过全注意力(64K 时为全注意力的 0.87 倍,128K 时为 0.62 倍),在混合型 Qwen3.6 上则在 64K 和 128K 之间(128K 时为 0.80 倍);通过预填充的 Amdahl 分解刻画了交叉区域。

4. 一个自定义诊断基准,*指针追逐草垛*(Pointer-Chase Haystack, PCH)(附录 B),在方法开发期间用于将选择器质量与模型能力隔离。

5. 一个阴性对照结果(LongBench-v1),它限定了*何时*路由器有帮助:仅当选择器的每查询预算(它保留的块)相对于答案相关证据所在位置较小时。

## 2 背景与相关工作

### 2.1 块稀疏注意力

标准的解码器层 (Vaswani et al., 2017) 将隐藏状态 \(h^{(L)} \in \mathbb{R}^{N \times d_{\text{model}}}\) 映射为 \(h^{(L+1)}\),通过:
\[
h' = h^{(L)} + \operatorname{Attn}\bigl(\operatorname{LN}(h^{(L)})\bigr), \qquad h^{(L+1)} = h' + \operatorname{MLP}\bigl(\operatorname{LN}(h')\bigr).
\]
注意力块,对于查询位置 \(i\) 和头 \(h\),计算:
\[
\operatorname{Attn}(x)_{i,h} = \sum_{j \leq i} \frac{\exp\bigl(q_{i,h} \cdot k_{j,h} / \sqrt{d_{\text{head}}}\bigr)}{\sum_{j' \leq i} \exp\bigl(q_{i,h} \cdot k_{j',h} / \sqrt{d_{\text{head}}}\bigr)} v_{j,h}.
\]
*块稀疏注意力*将 \(\{j \leq i\}\) 替换为一个小的选定子集 \(S_i\),该子集在推理时基于冻结的权重对每个查询进行选择。

### 2.2 选择器全景

具体的选择器在 (a) 如何将键池化为逐块摘要,以及 (b) 如何对块进行评分(相对于查询)方面有所不同,但**它们都在每查询选择步骤上归结为对块得分的 top-\(k\) 规则**:

- **SSA** (Subquadratic, 2025) (Subquadratic Sparse Attention) – 每个 Q-tile 的 top-\(k\) 作用于平均池化的键,\(s_b = q \cdot \bar{k}_b\),其中 \(\bar{k}_b = \frac{1}{B_n} \sum_j k_j\)。计算简单;平均池化会模糊单一的强键,这正好是上述讨论的多键 NIAH 退化的原因。
- **Quest** (Tang et al., 2024) – 逐块元素级 \(K_b^{\min}, K_b^{\max}\) 摘要;得分是 \(\max_{j \in b} q \cdot k_j\) 的上界,按坐标计算(见公式 (2))。能够恢复平均池化丢失的单一强键信号。
- **H2O** (Zhang et al., 2023), **SnapKV** (Li et al., 2024), **MInference** (Jiang et al., 2025) – top-\(k\) 作用于基于学习或注意力历史的块得分。
- **NSA** (Yuan et al., 2025), **MoBA** (Lu et al., 2025) – 端到端学习的门控,但仍然在每查询选择步骤上归结为对块得分的 top-\(k\)。

这些方法都没有将截断视为一个*不确定性下的决策*:当 \(s_{(k-1)} \approx s_{(k)}\) 时,选择器在不花费额外预算的情况下做出决定。这个步骤正是本文所操控的杠杆。因为杠杆作用于截断(而不是如何计算 \(s\)),它可以与上述任何评分骨干网络组合。在实验中,我们在 SSA 风格的 K-均值骨干网络和 Quest 的 K-max 上界基础上评估了它。

### 2.3 长上下文评估

已发表的基准分为两类。

**合成/诊断**:RULER (Hsieh et al., 2024)(NIAH, VT)、BABILong、MRCR。设计用于对长上下文召回进行受控的压力测试。它们的失败模式是可解释的,但**不能预测**下游性能,正如 HELMET (Yen et al., 2024) 所记录的那样。

**真实任务**:LongBench (Bai et al., 2023, 2024)、HELMET (Yen et al., 2024)、NoCha。涵盖多跳问答、摘要、代码等。LongBench-v2 特别有一个 *medium* 子集,其原生提示长度通常 \(\geq 100K\) 词,旨在压力测试长上下文选择。我们使用 RULER NIAH(合成,标准化)和 LongBench-v1 + v2 medium(真实任务,标准化)作为主要基准,并在方法开发期间使用自定义诊断基准(PCH,附录 B)。

## 3 方法

我们以端到端的方式描述工作方法。每一步都由前一步驱动,并解决一个特定的弱点。完整的逐方程推导见附录 A;这里我们给出读者理解实验所需的路径。

### 3.1 第一步:块评分

键被分组为连续的块,块大小为 \(\text{BLOCK}_N = 64\) 个令牌。我们在两种块评分规则之上评估路由器。

第一种是 SSA 风格的均值池化键内积:在层 \(L\),选择器通过下式对块 \(b \in \{0, \dots, N_B-1\}\) 进行评分:
\[
s^{(L)}[i,h,b] = \frac{q_i^{(L,h)} \cdot \bar{k}_b^{(L,h)}}{\sqrt{d_{\text{head}}}}, \qquad \bar{k}_b^{(L,h)} = \frac{1}{\text{BLOCK}_N} \sum_{j \in \text{block } b} k_j^{(L,h)}. \tag{1}
\]
均值池化会模糊单一键信号:一个包含一个强键且邻近噪声的块,其得分像全是噪声的块,从而被丢弃。

第二种是 Quest 的 K-min/K-max 上界 (Tang et al., 2024),它将 \(\bar{k}_b\) 替换为元素级对 \((K_b^{\min}, K_b^{\max})\),并评分为:
\[
s_b^{\text{quest}} = \sum_{d=1}^{d_{\text{head}}} \max\bigl(q_d \cdot K_{b,d}^{\max},\; q_d \cdot K_{b,d}^{\min}\bigr), \tag{2}
\]
这是一个按坐标计算的 \(\max_{j \in b} q \cdot k_j\) 上界,保留了均值所平均掉的单一键信号。

我们将公式 (1) 和 (2) 之间的选择视为一个*骨干网络超参数*;下游的所有内容(每个 tile 的选择、截断间隔 \(\sigma\)、触发条件)都是相同的。

### 3.2 第二步:每 tile 选择

一个朴素的逐行 top-\(k\) 会输出一个形状为 \([B, H, M, N_B]\) 的布尔 keep 张量,下游必须对其进行排序和合并(针对 Q-tile 的行);这在长上下文下成为主要成本。遵循 SSA 的做法,我们将 \(\text{BLOCK}_M = 64\) 个连续的查询行分组为一个 *Q-tile* \(t\),并针对每个 tile 做出一个选择决策,由该 tile 中的所有行共享:
\[
\text{tile\_score}[t, h, b] = \max_{r \in \text{tile } t} s[r, h, b], \qquad \text{kv\_idx}[t, h] = \operatorname*{top-}k\bigl(\text{tile\_score}[t, h, \cdot]\bigr). \tag{3}
\]
槽块 0 和 tile 自身的块通过将其得分*在 top-\(k\) 之前*添加 \(+\infty\) 而强制纳入 kv_idx,这避免了后续的拼接和去重操作。输出是一个形状为 \([B, H, Q_t, k_{\text{budget}}]\) 的单一整数张量(其中 \(Q_t = M / \text{BLOCK}_M\)),直接传递给注意力核,无需逐行内部掩码。选择操作的 wall time 复杂度为 \(O(N N_B / \text{BLOCK}_M)\),而非逐行版本的 \(O(N N_B)\)。这在结构上比逐行更粗糙:一个被 tile 中某行强烈需要的块可能被多个行弱需要的块超越,因为 tile 得分是内积在行上的最大值。路由器(第三步至第五步)解决了这一问题——它在 top-\(k\) 截断模糊的 tile 上花费可控的额外预算,同时保持置信的 tile 不变。

### 3.3 第三步:截断是一个决策——读取其不确定性

公式 (3) 中的 top-\(k\) 是一个决策:保留块 \((k-1)\),丢弃块 \((k)\)。其*质量*自然由该排序的决定性程度来衡量。令 \(s_{(0)} \ge s_{(1)} \ge \cdots\) 为 tile \(t\) 和头 \(h\) 的排序后 tile 得分。我们定义归一化的截断间隔:
\[
\sigma[t, h] = \frac{s_{(k-1)} - s_{(k)}}{s_{(0)} - s_{(k)}} \in [0,1]. \tag{4}
\]
直接的解释:
- \(\sigma \to 1\) – 保留集远高于被拒绝的尾部;截断是明确的。
- \(\sigma \to 0\) – 保留集的最后一个元素与第一个被拒绝的元素几乎相同;截断相当于抛硬币。

\(\sigma\) 通过一个 top-\((k+1)\) 部分排序计算(比产生 kv_idx 的 top-\(k\) 多一个元素),无需额外渐近成本。

#### 为什么这是一个信息价值信号

如果我们决定采用 top-\(k\),截断决策的期望损失由以下因素界定:

相似文章

学习跳跃块:自我发现的超度量路由用于硬件加速稀疏注意力

Reddit r/artificial

本文介绍了动态超度量注意力(Dynamic Ultrametric Attention),这是一个框架,其中Transformer在训练期间学习每头块稀疏路由拓扑,然后在推理时将这些拓扑卸载到自定义的Triton块稀疏内核上,与密集注意力相比,实现了高达28倍的加速和98.4%的内存减少。

学习重点:使用因果证据集监督稀疏注意力路由

arXiv cs.LG

本文测试了注意力权重能揭示模型输出实际依赖内容的假设,发现注意力与因果依赖常不一致。作者提出将干预掩码获得的因果证据集作为稀疏注意力路由器的监督信号,在注意力蒸馏路由器失败的检索任务上实现了近乎完美的准确率。

基于压缩内容选择的无参数自适应稀疏注意力

arXiv cs.LG

本文提出了一种无参数的自适应稀疏注意力方法,利用gzip压缩比动态选择非冗余块进行长程注意力,在PG-19语言建模上相较于固定和学习的稀疏注意力基线取得了显著的困惑度提升。