论文

PISA:以金字塔块选择实现对数线性稀疏注意力

Block Sparse Attention with Log-Linear Complexity

模型架构Transformer

摘要

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