HyperAIHyperAI

Command Palette

Search for a command to run...

具有对数线性复杂度的块稀疏注意力

Bohao Tang Zhen Qin Yuqi Pan Zheng Li Pengfei Liu

摘要

将语言模型扩展到长上下文受限于自注意力的二次方计算代价。块稀疏注意力提供了一种高效的替代方案,但选择被保留的块仍然是一个瓶颈。传统的块选择方法需要对所有查询-块对进行打分,因此在序列长度上仍然是二次方的。为了解决这一问题,我们提出了 PISA,一种采用金字塔式 Top-K 选择策略的块稀疏注意力机制。其主要思想是逐步缩小不同层级上的候选范围,从而更高效地找到最相关的键。具体来说,我们构建了一个由粗到细的键层级结构,并从最粗层级开始进行选择。在每一层级,我们对一个有界的候选集合应用 LogSumExp 打分,为下一更细层级选择候选,直到到达最细层级。通过池化,我们构建了 O(log N) 个键层级,从而获得 O(N log N) 的总体复杂度,其中 N 表示序列长度。我们为训练和推理开发了硬件感知的 Triton 内核,将层级路由与 LogSumExp 打分融合,且无需显式构造查询-键打分矩阵。我们进一步在语言建模任务上评估了该方法。与基线相比,我们的方法在常识推理等基准上取得了相当的性能,同时在检索任务上表现更好。

一句话总结

上海交通大学、上海创新研究院和字节跳动 Seed 的研究者提出 PISA,一种块稀疏注意力机制,它在由粗到细的键层级上使用金字塔式 Top-K 选择,并采用有界的 LogSumExp 评分,以实现 O(N log N) 复杂度;同时开发了面向训练和推理的硬件感知 Triton 内核,在语言建模基准上取得与基线相当的性能,并在检索任务上获得更好的结果。

核心贡献

  • 本文提出 PISA,一种块稀疏注意力机制,它在由粗到细的键层级上使用金字塔式 Top-K 选择,并在每个层级上使用 LogSumExp 评分,将块选择复杂度降低到 O(N log N)。
  • 面向训练和推理的硬件感知 Triton 内核融合了层级路由和 LogSumExp 评分,保持有界候选集,并避免物化稠密的查询-键评分矩阵。
  • 语言建模实验表明,PISA 在常识推理上达到与传统块稀疏注意力基线相当的性能,在检索任务上取得更好的结果,并且在最长 256K tokens 的序列上实现了比单层块选择更低的选择延迟。

引言

将语言模型扩展到更长的上下文十分重要,但全自注意力会随序列长度呈二次增长。块稀疏注意力通过将每个查询限制到选定的键块来降低成本;然而,传统 Top-K 块选择仍会扫描每个候选键块,这在长序列上造成 O(N2/C)O(N^2/C)O(N2/C) 的选择瓶颈。为了解决该问题,作者提出 PISA,一种金字塔稀疏注意力方法,它构建多级池化键层级,并使用 LogSumExp 评分进行由粗到细的 Top-K 选择。这将块选择复杂度降低到 O(Nlog⁡N)O(N \log N)O(NlogN),并使用硬件感知 Triton 内核,在融合层级路由的同时避免稠密的查询-键评分矩阵。

方法

作者提出了 PISA,一种金字塔块稀疏注意力机制。它结合了三个主要组件:块稀疏注意力骨干、建立在由粗到细键金字塔上的分层 top-KKK 块选择器,以及专门的训练、预填充和解码内核。

块稀疏注意力

键和值被划分为连续块。对于查询 qtq_tqt​,块选择器只保留键块的一个子集,并在对应的原始键和值上计算注意力:

ot=∑j∈Ttexp⁡(qt⊤kj/d)vj∑j∈Ttexp⁡(qt⊤kj/d).o_t = \frac{\sum_{j \in \mathcal{T}_t} \exp\left(q_t^\top k_j / \sqrt{d}\right) v_j} {\sum_{j \in \mathcal{T}_t} \exp\left(q_t^\top k_j / \sqrt{d}\right)}.ot​=∑j∈Tt​​exp(qt⊤​kj​/d​)∑j∈Tt​​exp(qt⊤​kj​/d​)vj​​.

一个扁平块稀疏注意力方法对每个查询的 M=⌈N/C⌉M = \lceil N/C \rceilM=⌈N/C⌉ 个块摘要进行评分,这仍然在全序列评分上花费 O(N2d/C)O(N^2 d / C)O(N2d/C)。PISA 则将块选择组织为分层搜索。

金字塔键块层级

键层级由细到粗构建。第 000 级包含作为单元素块的原始键:

M0=N,Ki(0)=[ki].M_0 = N, \qquad K_i^{(0)} = [k_i].M0​=N,Ki(0)​=[ki​].

第一个合并步骤使用 g0=Cg_0 = Cg0​=C 并产生叶子块。对所有更高层级,合并因子为 gℓ=g=2g_\ell = g = 2gℓ​=g=2:

Mℓ+1=⌈Mℓgℓ⌉,M_{\ell+1} = \left\lceil \frac{M_\ell}{g_\ell} \right\rceil,Mℓ+1​=⌈gℓ​Mℓ​​⌉, Ki(ℓ+1)=Concatr∈Chℓ(i)Kr(ℓ).K_i^{(\ell+1)} = \mathrm{Concat}_{r \in \mathrm{Ch}_\ell(i)} K_r^{(\ell)}.Ki(ℓ+1)​=Concatr∈Chℓ​(i)​Kr(ℓ)​.

每个块 Ki(ℓ)K_i^{(\ell)}Ki(ℓ)​ 都有一个摘要向量 kˉi(ℓ)\bar{k}_i^{(\ell)}kˉi(ℓ)​。从 kˉi(0)=ki\bar{k}_i^{(0)} = k_ikˉi(0)​=ki​ 开始,摘要通过均值池化递归计算:

kˉi(ℓ+1)=1∣Chℓ(i)∣∑r∈Chℓ(i)kˉr(ℓ).\bar{k}_i^{(\ell+1)} = \frac{1}{|\mathrm{Ch}_\ell(i)|} \sum_{r \in \mathrm{Ch}_\ell(i)} \bar{k}_r^{(\ell)}.kˉi(ℓ+1)​=∣Chℓ​(i)∣1​r∈Chℓ​(i)∑​kˉr(ℓ)​.

由于该层级在叶子层之后使用固定合并因子 g=2g=2g=2,它包含 L=O(log⁡N)L = O(\log N)L=O(logN) 层。查询在整个选择过程中保持逐 token 处理。

由粗到细的 Top-KKK 选择

选择从最粗层级 LLL 向下进行到叶子层级 111。对于每个查询,最粗层级的候选集初始化为:

At(L)={1}.\mathcal{A}_t^{(L)} = \{1\}.At(L)​={1}.

在层级 ℓ\ellℓ,PISA 只对 At(ℓ)\mathcal{A}_t^{(\ell)}At(ℓ)​ 中的候选块进行评分,选择 top KKK 块,并将保留的块扩展为其子块:

At(ℓ−1)=⋃i∈It(ℓ)Chℓ−1(i).\mathcal{A}_t^{(\ell-1)} = \bigcup_{i \in \mathcal{I}_t^{(\ell)}} \mathrm{Ch}_{\ell-1}(i).At(ℓ−1)​=i∈It(ℓ)​⋃​Chℓ−1​(i).

由于对于 ℓ>1\ell > 1ℓ>1,每个保留块最多有 ggg 个子块,因此下一级候选集大小以 gKgKgK 为界。如果某一层最多包含 KKK 个有效候选,则所有候选无需评分即被保留。最终的叶子层选择集 Tt=Tt(1)\mathcal{T}_t = \mathcal{T}_t^{(1)}Tt​=Tt(1)​ 决定稀疏注意力层使用的原始键和值块。

在评分方面,PISA 使用 LogSumExp 块评分。候选块 Ki(ℓ)K_i^{(\ell)}Ki(ℓ)​ 对查询 qtq_tqt​ 的评分为:

st,i(ℓ)=log⁡∑r∈Chℓ−1(i)exp⁡(qt⊤kˉr(ℓ−1)d),ℓ≥1.s_{t,i}^{(\ell)} = \log \sum_{r \in \mathrm{Ch}_{\ell-1}(i)} \exp\left(\frac{q_t^\top \bar{k}_r^{(\ell-1)}}{\sqrt{d}}\right), \quad \ell \geq 1.st,i(ℓ)​=logr∈Chℓ−1​(i)∑​exp(d​qt⊤​kˉr(ℓ−1)​​),ℓ≥1.

在叶子层,子项是原始键,因此叶子评分就是该块的精确原始键 LogSumExp 评分。在中间层,每个候选由其子摘要评分,每个评分最多归约 ggg 个 logits。这使得评分成本保持较低,同时使叶子层选择与注意力的 log-sum-exp 计算一致。

内核实现

对于训练和预填充,PISA 使用两阶段内核设计。

在第一阶段,内核处理从 LLL 到 222 的中间层级。它在查询位置和键/值头上并行化。同一 GQA 组内的查询头共享相同的选定键块。对于键/值头 hhh,内核为 H(h)\mathcal{H}(h)H(h) 中的所有查询头计算 LSE 评分并求和:

uh,t,i(ℓ)=∑h′∈H(h)sh′,t,i(ℓ).u_{h,t,i}^{(\ell)} = \sum_{h' \in \mathcal{H}(h)} s_{h',t,i}^{(\ell)}.uh,t,i(ℓ)​=h′∈H(h)∑​sh′,t,i(ℓ)​.

然后它对求和后的评分应用 top-KKK 选择,并将保留的块扩展为其子块,以形成下一候选集。候选索引在整个循环中保持程序局部性。

第一阶段之后,每个查询最多有 gKgKgK 个候选叶子块。第二阶段将需要评分相同候选键块的查询分组为簇。每个簇被划分为最多 QtileQ_{\text{tile}}Qtile​ 个查询的瓦片,实现中 Qtile=4Q_{\text{tile}}=4Qtile​=4。每个内核程序一次加载形状为 d×Cd \times Cd×C 的键瓦片,并在形状为 (QtileGQ)×d(Q_{\text{tile}} G_Q) \times d(Qtile​GQ​)×d 的查询瓦片中复用,其中 GQ=∣H(h)∣G_Q = |\mathcal{H}(h)|GQ​=∣H(h)∣。它计算

(QtileGQ)×d×d×C⟶(QtileGQ)×C,(Q_{\text{tile}} G_Q) \times d \times d \times C \longrightarrow (Q_{\text{tile}} G_Q) \times C,(Qtile​GQ​)×d×d×C⟶(Qtile​GQ​)×C,

对每个查询头在键上应用 LSE,然后对 GQG_QGQ​ 个查询头求和。最后一个内核为每个查询和键/值头独立应用 top-KKK 选择。

在解码期间,PISA 缓存均值金字塔,并只针对每个新 token 更新当前叶子均值以及受影响的祖先路径。块选择使用单个融合内核,从最粗层级向下运行到 ℓ=1\ell=1ℓ=1,直接对原始键块进行评分。这避免了每个解码步骤启动多个内核的开销。

复杂度

由于金字塔层级有 O(log⁡N)O(\log N)O(logN) 层,且每个查询在每层最多评分 gKgKgK 个候选块,因此每个查询的选择成本为 O(log⁡N)O(\log N)O(logN)。因此,全序列训练或预填充成本为 O(Nlog⁡N)O(N \log N)O(NlogN),而解码平均每步成本为 O(log⁡N)O(\log N)O(logN)。

选择两阶段训练/预填充内核主要是为了输入/输出效率。通过将共享同一候选键块的查询分组,第二阶段在多个查询间复用每个已加载的键块。对于实现设置 GQ=16G_Q = 16GQ​=16、Qtile=4Q_{\text{tile}} = 4Qtile​=4 和 C=64C = 64C=64,

GQ+CQtile=32<64=C,G_Q + \frac{C}{Q_{\text{tile}}} = 32 < 64 = C,GQ​+Qtile​C​=32<64=C,

因此,与单阶段设计相比,两阶段设计减少了叶子层查询/键 IO。在解码期间,簇只包含一个查询,因此没有跨查询键复用。融合的单阶段内核复用中间选择期间已加载的查询向量,并避免单独的分组和最终选择启动。

实验

实验在匹配训练条件下比较了 PISA 与 Full Attention、BSA、NSA 和 HiLS,规模为 418M、1.47B 和 2.67B,先在 4K 序列长度上使用 100B tokens 训练,再在 16K 上使用 10B tokens 进行持续预训练。PISA 在语言建模和常识推理上与稀疏基线表现相当,在所有规模上的平均包含准确率在稀疏方法中最高,尽管 Full Attention 仍然更高,并在 RULER 上表现出较强的长上下文检索能力。消融实验表明,相对于 PISA-1、PISA-2 和 BSA,基于 LSE 的评分改善了 Top-K 块选择质量和下游性能,而块选择效率在非常长的上下文下相对 BSA 更具优势。

该表比较了稀疏注意力方法的训练设置和计算复杂度。HiP 作为一种免训练方法尤为突出,具有次二次预填充和对数级解码复杂度,而 HISA 和大多数可训练方法使用二次预填充和线性解码复杂度。LLSA 是可训练方法中的例外,具有次二次预填充,但其成本按每次扩散注意力传递报告,而不是按自回归预填充报告。在列出的免训练方法中,HiP 实现了比 HISA 更低的预填充和解码复杂度,具有次二次预填充和对数级解码。大多数可训练方法具有二次预填充和线性解码复杂度;LLSA 具有次二次预填充,但按扩散注意力传递报告,而不是按自回归预填充报告。

在 100B tokens 预训练后,PISA 在语言建模和常识推理上与 BSA、NSA 和 HiLS 表现相当。在稀疏方法中,PISA 在所有评估规模上的包含任务平均准确率最高,而全注意力保持更高的包含平均值。PISA 还表现出比 PISA-1、PISA-2 和 BSA 更低的训练损失和更高的平均包含准确率。PISA 在比较的稀疏方法中跨模型规模达到最佳包含平均值。全注意力在包含准确率上仍然领先,而 PISA 在语言建模和常识推理上与稀疏基线相当。PISA 在训练损失和平均包含准确率上优于 PISA-1、PISA-2 和 BSA。

在 16K 继续预训练后,全注意力在 2.67B 模型上具有最高的平均 RULER 大海捞针准确率。在稀疏方法中,PISA 和 PISA-2 整体领先且几乎持平,PISA-2 平均略领先,二者均高于 NSA 和 BSA。各方法的准确率随着上下文长度下降,尤其是在多键、多查询和多值任务上。全注意力取得了最高的平均准确率,而 PISA-2 和 PISA 是最强的稀疏方法且几乎持平。PISA 变体在若干短上下文多键和多查询设置上优于全注意力,但所有稀疏方法在 16K 时下降更剧烈,尤其是在多键和多值任务上。

实验从效率、100B-token 预训练后的语言建模和常识推理,以及 16K 继续预训练后的长上下文检索方面比较了稀疏注意力方法。PISA 在语言建模和常识推理上与 BSA、NSA 和 HiLS 表现相当,并在稀疏方法中实现了最高的平均包含准确率,尽管全注意力整体上仍然更强。在长上下文 RULER 评估中,PISA 和 PISA-2 是最强的稀疏方法且几乎持平,并在一些短上下文场景中优于全注意力,但所有稀疏方法到 16K 时退化更严重。复杂度结果还表明,免训练的 HiP 具有次二次预填充和对数级解码,而大多数可训练方法使用二次预填充和线性解码。


用 AI 构建 AI

从创意到上线——通过免费 AI 协同编码、开箱即用的环境和最优惠的 GPU 价格,加速您的 AI 开发。

AI 协同编码
开箱即用的 GPU
最优定价

HyperAI Newsletters

订阅我们的最新资讯
我们会在北京时间 每周一的上午九点 向您的邮箱投递本周内的最新更新
邮件发送服务由 MailChimp 提供