大模型读长文,终于不用把每个字都看一遍
大模型读长文时,得把每个词和所有其他词都算一遍关系,所以越长越慢,成本是平方级上涨。这篇提出一种新算法:先粗后细地筛,像查地图先看省再看市,只挑最相关的部分来算,把成本从平方级压到接近线性。在常识推理上不输原版,在检索任务上还更好。它不是你明天能用上的,但这是让大模型真正读得动整本书、整份合同的关键一步。
📄 原文摘要(英文)
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.