🤖 AI Summary
This study addresses the substantial KV cache read overhead in long-context inference, where critical tokens vary dynamically with queries. We propose a training-free stochastic attention method that partitions cached keys into groups, samples representative groups via proxy keys, and estimates full attention through importance weight correction. Additionally, we introduce a group-level importance sampling mechanism that efficiently identifies critical subsets without scanning the entire cache, enabling a flexible memory-accuracy trade-off. Combined with GPU kernel optimizations, our approach reads only 16%–22% of the KV cache at 32K context length while retaining 94%–99% of baseline performance, achieving a 1.69× inference speedup.
📝 Abstract
Attention often concentrates on a small subset of tokens in the context, but which subset matters changes from one query to the next. To exploit this changing structure, we introduce SANTA++, a training-free stochastic attention method that uses representative keys for memory-efficient selection without scanning the entire key-value (KV) cache. Cached keys are organized into teams, and the query scores one representative from each team to decide which teams to sample. We compute exact attention scores within the sampled teams and reweight each team's contribution by the inverse of its inclusion probability. This importance sampling correction estimates attention over the full cache, with a sampling budget that lets us trade memory reads for accuracy. Remarkably, with 32 or 64 sampled teams, SANTA++ uses 16% to 22% of dense attention's KV reads and retains 94% to 99% of the dense-attention baseline's scores on LongBench v2 and HELMET's retrieval-augmented generation subset, and 85% to 91% on RULER, with Qwen2.5-7B-Instruct at 32K context. With 31 sampled teams, our GPU implementation delivers a $1.69\times$ attention speedup over the dense FlashAttention baseline at 32K context. By reducing the number of cache entries read, SANTA++ in principle complements architectures with compressed KV representations, such as multi-head latent attention. Our kernels are available at: https://github.com/OPUSLab/santapp-kernel-demo.git.