Attention via Black-Box Vector Search
Abstract
Sparse attention mechanisms estimate attention over 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 -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 retrieved keys are both sufficient and necessary. With indices, we give an algorithm that retrieves only 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 retrieved keys. When integrated into LLM inference, our algorithms outperform top- and sampling approaches used in prior work and yield attention approximation that scales favorably to long contexts.
Then back it, or bet against it.
Related papers
Open the market on this paper to see 7 more related papers.