跳到正文
原文
HuggingFace Daily Papers(社区热门论文)·· 3 天前AI 评分36

PISA:基于金字塔 Top-K 选择的块稀疏注意力,实现 Log-Linear 复杂度

Block Sparse Attention with Log-Linear Complexity

AI 导读

研究者提出 PISA,一种采用金字塔 Top-K 选择策略的块稀疏注意力机制,通过构建 O(log N) 层级的键层次结构并逐层用 LogSumExp 打分筛选候选,将整体复杂度降至 O(Nlog N)。团队还开发了面向训练和推理的硬件感知 Triton kernel,无需物化 query-key 分数矩阵。在语言建模任务上,PISA 在常识推理等基准上与基线相当,并在检索任务上取得更好结果。

正文

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(Nlog 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.

来源:HuggingFace Daily Papers(社区热门论文) · arxiv.org