arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~
arXiv 2609.20888cs.LG

弹性阈值注意力:面向长上下文解码的学习型上下文稀疏性

Elastic Threshold Attention: Learned Contextual Sparsity for Long-Context Decoding

  • Boston University(波士顿大学)

机构由 AI 辅助整理,请以论文原文为准。

Themistoklis Haris, Henry Li, Maryam Karimzadehgan

AI总结:

针对长上下文解码中KV缓存导致的内存带宽瓶颈,提出弹性阈值注意力(ETA),通过端到端学习动态上下文阈值实现硬件加速解码,在不损失稠密模型质量下,以约38%活跃密度媲美稠密注意力,并带来高达2.5倍解码加速。

AI中文摘要:

大规模KV缓存可能在长上下文解码过程中造成严重的内存带宽瓶颈。稀疏注意力方法通过选择性加载来缓解这一问题,但代价是:僵化的启发式规则会丢弃必要的上下文,导致质量下降。我们引入了弹性阈值注意力(ETA),一种端到端可训练的架构,在不牺牲稠密模型质量的前提下实现硬件加速的解码速度。ETA直接从查询表示中预测动态的、上下文相关的阈值,使模型能够为困难的检索或推理步骤分配类似稠密的上下文,同时修剪常规标记。为了从零开始学习该策略而不发生表示崩溃,ETA在训练期间对低于阈值的logits进行乘法抑制(向零方向),而非直接删除它们。针对这一平滑的均匀注意力底层的训练提供了一个分布式的概率储备,使得初始标记上的局部注意力汇聚消失。它还使模型能够在推理时硬性修剪无信息的KV块,并吸收由粗粒度GPU块选择共同引入的偶然标记。因此,一个1.45B预训练的ETA模型在语言建模、常识推理和长上下文针检索方面,以约85%的训练稀疏度和约38%的活跃解码密度,与稠密注意力相媲美。在推理时,我们使用Triton实现了一个自定义解码内核,利用缓存的几何概率界限以O(1)时间筛选KV块,在长达512K标记的序列上,相较于FlashAttention-2提供了高达2.5倍的墙钟解码加速。最后,我们引入了一种针对特定领域部署的离线校准算法,冻结每头恒定阈值以消除预测器开销,将注意力计算额外削减27%。

英文摘要:

Massive KV caches can cause severe memory-bandwidth bottlenecks during long-context decoding. Sparse attention methods mitigate this problem, but often drop necessary context, leading to quality degradation. We introduce \textbf{Elastic Threshold Attention (ETA)}, an end-to-end trainable architecture that rivals dense model quality under hardware-aligned block-sparse decoding. ETA predicts dynamic, contextual thresholds directly from query representations, adjusting context retention depending on the task at hand. To learn this policy from scratch while enabling near lossless KV cache pruning at inference time, we filter attention logits through a shifted SiLU gate during training. We show theoretically and empirically that this creates a smooth, near-uniform attention floor that neutralizes sub-threshold value contributions while simultaneously causing localized attention sinks on initial tokens to disappear. To materialize these advantages, we implement a fused inference-time kernel in Triton that screens KV blocks in $O(1)$ time using cached geometric-probabilistic bounds. Across language modeling, reasoning, and RULER benchmarks, our 1.45B ETA model matches or exceeds dense quality, outperforming alternative fast decoding methods (Quest, H$_2$O, NSA) while achieving higher sparsity levels. Our kernel also achieves up to $2.15\times$ end-to-end speedups over FlashAttention-2 at context lengths of up to $512$K tokens. Finally, we introduce an offline calibration algorithm for domain-specific deployments that freezes per-head constant thresholds, cutting attention compute by an additional 27\% at no quality cost.

↑