Efficient attention method speeds up large context processing

SANTA++: Sampling Attention through Representative Keys

Machine LearningComputation and Language

Summary

Processing long sequences of information in AI models can be slow because they have to look at everything at once. The authors introduce SANTA++, which cleverly picks smaller representative parts to pay attention to without having to check everything. This approach keeps results very close to full processing but reads much less data, making it faster and more efficient. Their tests show it works well on tasks requiring understanding of very long context.

What this means in practice

  • For nlp engineers: Implement SANTA++ to reduce memory and speed up handling of very long input sequences in transformer-based language models.
  • For cloud platform teams: Integrate SANTA++ attention kernels for cost-effective deployment of large-context language models by lowering resource demands.

Authors

Kyle Lee, Christian Z. Pratt, Ruoyu Fang, Heekyung Lee, Avinash Lohitsa, Ryan Modafe, Kerem Y. Camsari

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.