🤖 AI Summary
This study addresses the low computational efficiency and compromised accuracy of sparse attention in long-context inference for large language models by proposing a unified framework based on priority sampling. Methodologically, it designs a multi-index algorithm leveraging maximum inner product search and vector indexing techniques to optimize key retrieval. Theoretically, it establishes a smooth trade-off mechanism between the number of indices and retrieval volume, transcends single-index performance lower bounds through key augmentation, and introduces a near-optimal estimation algorithm. Experimental results demonstrate that the proposed approach significantly outperforms top-k and sampling baselines, substantially reducing retrieval overhead while effectively enhancing attention approximation accuracy in long-context scenarios.
📝 Abstract
Sparse attention mechanisms estimate attention over $n$ tokens using a small subset of keys. Many existing approaches use maximum inner product search (MIPS) to retrieve the heaviest keys, which motivates the following question: given black-box access to a MIPS oracle, how many keys must be retrieved to output an $\varepsilon$-accurate attention estimate?
We answer this question by unifying prior approaches through the framework of priority sampling. With a single MIPS index, we show that $Θ(\sqrt{n}/\varepsilon)$ retrieved keys are both sufficient and necessary. With $Θ(\log n)$ indices, we give an algorithm that retrieves only $O(\log n+1/\varepsilon^2)$ keys and prove that this is near-optimal. More generally, we design algorithms that establish a smooth tradeoff between the number of MIPS indices and number of retrieved keys. We then show that if we allow augmentation of keys and queries, we can bypass the above lower bounds: there exists a simple priority-sampling estimator using a single MIPS index and $O(1/\varepsilon^2)$ retrieved keys. When integrated into LLM inference, our algorithms outperform top-$k$ and sampling approaches used in prior work and yield attention approximation that scales favorably to long contexts.