HiRouter: Efficient Maximum Inner Product Search for Top- Attention via Hierarchical Routers
Abstract
Attention enables parallel token retrieval but at a quadratic memory and computation cost relative to sequence length, limiting its applicability to long contexts. Sparse top- attention alleviates this issue by reducing computation while maintaining competitive accuracy. However, existing sparse top- methods rely on k-means clustering or locality-sensitive hashing (LSH) to avoid computing the full attention matrix. While GPU-friendly, these approaches employ approximate and often decoupled partition-based routing, which can limit retrieval accuracy. To this end, we propose Hierarchical Router (HiRouter), a sparse top- attention mechanism that that leverages a learned routing function within a hierarchical structure, coupling retrieval with attention by co-training queries, keys, and routing functions end-to-end. To make the learned router effective, we design two routing rules: relevant tokens should co-route with the query, and bucket occupancy should remain balanced to avoid collapse and support uniform GPU workloads. By optimizing two surrogate losses over a learned multi-level hierarchy, HiRouter realizes these rules, partitioning tokens into discrete buckets in linear time and enabling accurate per-sequence retrieval with GPU-friendly bucket layouts. On standard benchmarks, HiRouter matches full-attention performance, while matching/surpassing full-attention accuracy on sequences beyond 4K tokens with up to 2 faster inference than FlashAttention.
Then back it, or bet against it.
Related papers
Open the market on this paper to see 7 more related papers.