发表机构
Shanghai Jiao Tong University; Shanghai Innovation Institute(上海交通大学; 上海创新研究院)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
针对自注意力二次方成本限制长上下文的问题,提出PISA块稀疏注意力机制,采用金字塔Top-K选择策略,实现O(N log N)复杂度,性能与基线相当且在检索任务上更优。
AI 中文摘要
将语言模型扩展到长上下文受到自注意力二次方成本的限制。块稀疏注意力提供了一种高效的替代方案,但选择保留的块仍然是一个瓶颈。传统的块选择需要对所有查询-块对进行评分,因此序列长度的复杂度仍然是二次方的。为了解决这个问题,我们提出了PISA,一种采用金字塔Top-$K$选择策略的块稀疏注意力机制。主要思想是逐步在不同层级上缩小候选范围,从而更高效地找到最相关的键。具体来说,我们构建了一个从粗到细的键层级,并从最粗的层级开始进行选择。在每个层级,对有限的候选集应用LogSumExp评分,以选择下一更细层级的候选,直到达到最细层级。通过池化,我们构建了$O(\log N)$层级的键,从而实现了总体复杂度为$O(N\log N)$,其中$N$表示序列长度。我们为训练和推理开发了硬件感知的Triton内核,融合了层级路由和LogSumExp评分,而无需物化查询-键分数矩阵。我们进一步在语言建模任务上评估了我们的方法。与基线相比,我们的方法在常识推理等基准上取得了相当的性能,同时在检索任务上提供了更好的结果。
英文摘要
Scaling language models to long contexts is limited by the quadratic cost of self-attention. Block sparse attention offers an efficient alternative, but selecting the retained blocks remains a bottleneck. Conventional block selection requires scoring all query-block pairs and therefore remains quadratic in sequence length. To address this issue, we propose PISA, a block-sparse attention mechanism that employs a pyramid Top-$K$ selection strategy. The main idea is to gradually narrow down the candidates across different levels, making it more efficient to find the most relevant keys. Specifically, we construct a coarse-to-fine hierarchy of keys and perform selection from the coarsest level. At each level, LogSumExp scoring is applied to a bounded candidate set to select candidates for the next finer level, continuing until the finest level is reached. Through pooling, we construct $O(\log N)$ levels of keys, yielding an overall complexity of $O(N\log N)$, where $N$ denotes the sequence length. We develop hardware-aware Triton kernels for both training and inference, fusing hierarchical routing and LogSumExp scoring without materializing the query-key score matrix. We further evaluate our method on language modeling tasks. Compared with the baseline, our method achieves comparable performance on benchmarks such as commonsense reasoning while delivering better results on retrieval tasks.