acceptodds
Under review as a conference paper at ICLR 2027

LASA: Learnable Adaptive Sparse Attention for Efficient and Accurate Reasoning

Abstract

In long chain-of-thought reasoning, every decoding step reads an ever-growing KV cache, and this memory traffic becomes the inference bottleneck. Sparse attention reduces this cost by accessing only a subset of the KV cache. Most existing methods focus on improving the quality of ranking KV tokens, while using a fixed token count (top-) for budgeting. Some adaptive-budget methods instead vary token counts according to a fixed target for retained attention mass (top-). Both rely on a priori targets shared across layers, KV heads, and queries, rather than learning from model output quality where more or less computation is needed. We introduce (Learnable Adaptive Sparse Attention), which makes budgeting a learned and fine-grained decision. Given a block ranking distilled from dense attention, a M-parameter predictor sets a threshold for each layer, KV head, and decoding step, and uses simple elementwise comparisons to select blocks whose normalized ranker scores exceed it. The predictor is trained end-to-end through output-level distillation under a relaxed token budget, with both the backbone and ranker frozen. Our extensive evaluation shows that substantially improves the accuracy–density Pareto frontier over existing sparse attention methods, retaining near-full-attention accuracy while under 1K average attended tokens. Our efficient kernels and vLLM integration translate these savings into up to attention-kernel and end-to-end decoding speedups over dense attention on H20 GPUs.

Then back it, or bet against it.

Related papers

Open the market on this paper to see 7 more related papers.