Elastic Threshold Attention: Learned Contextual Sparsity for Long-Context Decoding
cs.LG
Submitted: 2026-09-16
Updated: 2026-09-25
Code: https://github.com/meta-llama/llama3
License: http://creativecommons.org/licenses/by/4.0/
The gist: Massive KV caches can cause severe memory-bandwidth bottlenecks during long-context decoding.
Terminology
Abstract
Massive KV caches can cause severe memory-bandwidth bottlenecks during long-context decoding. Sparse attention methods mitigate this via selective loading, but that comes at a cost: rigid heuristics drop necessary context, leading to quality degradation. We introduce Elastic Threshold Attention (ETA), an end-to-end trainable architecture that achieves hardware-accelerated decoding speed without sacrificing dense model quality. ETA predicts dynamic, contextual thresholds directly from query representations, allowing the model to allocate dense-like context to difficult retrieval or reasoning steps while pruning routine tokens. To learn this policy from scratch without representation collapse, ETA multiplicatively suppresses sub-threshold logits toward zero during training rather than deleting them. Training against this smooth uniform attention floor provides a distributed probability reservoir that causes localized attention sinks on initial tokens to disappear. It also enables the model to hard-prune uninformative KV blocks at inference time and absorb incidental tokens co-admitted by coarse GPU block selection. As a result, a 1.45B pretrained ETA model rivals dense attention across language modeling, commonsense reasoning, and long-context needle retrieval at about 85% training sparsity and about 38% active decode density. At inference time, we implement a custom decode kernel in Triton that screens KV blocks in O(1) time using cached geometric-probabilistic bounds, delivering up to 2.5 times wall-clock decode speedups over FlashAttention-2 on sequences up to 512K tokens. Finally, we introduce an offline calibration algorithm for domain-specific deployments that freezes per-head constant thresholds to eliminate predictor overhead, cutting attention compute by an additional 27%.
Sources
- Generating Long Sequences with Sparse Transformers
- Rethinking Attention with Performers
- Think you have Solved Question Answering? Try ARC, the AI2 Reasoning Challenge
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
- HashAttention: Semantic Sparsity for Faster Inference
- SeerAttention: Learning Intrinsic Sparse Attention in Your LLMs
- Trainable Dynamic Mask Sparse Attention
- Quest: Query-Aware Sparsity for Efficient Long-Context LLM Inference
- HiRE: High Recall Approximate Top-$k$ Estimation for Efficient LLM Inference
- SpargeAttention: Accurate and Training-free Sparse Attention Accelerating Any Model Inference
Related papers
- Polynomial-Augmented Neural Networks (PANNs) with Weak Orthogonality Constraints for Enhanced Function and PDE Approximation
- AIRL-S: Unifying Reinforcement Learning and Search-Based Test-Time Scaling via Adversarial Inverse Reinforcement Learning
- Transformers as Bayesian In-Context Experimenters: Smoothness-Adaptive Efficient ATE Estimation
- Convergence issues in Relational Concept Analysis based on AOC-posets
- Beliefs Beyond Posteriors: Local-Consistency Optimisation for Bayesian Neural Networks
- Understanding Diffusion Models via Ratio-Based Function Approximation with SignReLU Networks