Universal Attention 用复合衰减门实现超 10 倍 KV-cache 压缩,并在 RULER 上超过未剪枝 Llama3
核心概要
该工作提出 Universal Attention,把查询感知、键感知与相似度三类衰减机制以几何平均组合成注意力矩阵的加性偏置,在保留 RoPE 与 Softmax 的前提下让模型自适应剪枝 KV-cache;在 1B、50B token 训练的 Llama3 架构上实现超过 10 倍压缩,RULER 平均分 50.3/49.3 高于未剪枝 Llama3 的 40.8,缓存仅 21.4/12.9 MB。
Figure 1: Overview of Universal Attention. ( Left ) Standard softmax attention with a fixed causal mask 𝐌 \mathbf{M} . ( Middle ) UA augments the attention logits 𝐐𝐊 ⊤ \mathbf{QK}^{\top} with a learned, content-aware decay bias 𝚲 \mathbf{\Lambda} , while applying L 2 L_{2} -normalization to keys, and a low-rank sigmoid gate to the attention output. ( Right ) Construction of 𝚲 \mathbf{\Lambda} from three decay mechanisms: similarity-based f sim f_{\text{sim}} , query-aware f q f_{\text{q}} , and key-aware f k f_{\text{k}} . These are combined via a geometric mean, transformed to log-space, and accumulated causally to produce the decay matrix.
arXiv深度剖析
提出统一框架,把多种上下文感知衰减机制以几何平均合成单一衰减率,作为 Softmax 注意力的加性偏置,从而在保留 RoPE 与 Softmax 的同时获得比单一遗忘门更丰富的衰减模式。 此前的遗忘门(如 FoX)只用单一标量衰减,实践中退化为软滑窗;该工作把查询感知、键感知与相似度衰减组合起来,并指出几何平均保留“否决”性质:任一机制判定不应衰减,合成衰减即被压向零。 以公式推导给出三类机制的构造,并在 10B token 训练的 1B 模型上做消融:去掉键感知衰减在 16k 使 RULER S1–S3 下降 36.20,去掉相似度衰减下降 31.76,而单独使用查询感知(FoX 式)在 16k 仅 40.0,完整 UA 为 58.4。
学习到的衰减本身构成自适应剪枝准则:推理时维护每个键值对的标量偏置,掩码低于阈值即逐出,且偏置单调下降使被逐出 token 不会重新变重要。 H2O、SnapKV、KeyDiff 等按预设预算(如保留 10%)剪枝,而 UA 由模型按头自行决定保留哪些、保留多少 token,压缩率随输入与长度变化。 在 50B token 模型上,阈值 0.15 时 LongBench-E(16k)平均压缩 21.38、单序列最大 113.95,分数与困惑度变化均小于一个点;训练中缓存先迅速稀疏化,随后随语言建模能力提升重新纳入信息,并在后半程缓慢下降。
在 1B 与 3B 规模上同时改善质量与内存:UA 在 RULER 上超过未剪枝 Llama3 与 Llama3-G,缓存远低于线性注意力与混合基线。 相对线性注意力与混合模型,UA 直接在 Softmax 注意力上做衰减,避免常数规模循环状态带来的内存下界;相对剪枝基线,UA 在更紧的缓存预算下保持更高分数。 1B 模型 RULER(4k)平均 50.3(p=0.05,21.4 MB)与 49.3(p=0.15,12.9 MB),对比 Llama3 的 40.8(143.7 MB)、Llama3-G 的 41.9、GDN-H 的 44.1(46.0 MB);3B 规模 RULER 平均 52.6 对 Llama3-3B 的 48.8,缓存从 341.8 MB 降至 46.6 MB。
长上下文泛化:UA 的缓存随长度趋于平台,在 16k 仍保持可用分数,而全注意力基线缓存线性增长。 在训练长度之外(1K 至 16K)评估,UA 在更长上下文上仍处于内存—性能前沿,且缓存保持在 25 MB 以下。 附录 I 的逐任务 RULER 结果显示,UA 在 16k 的缓存为 22.7/13.6 MB,而 Llama3 为 299.0 MB;GDN-H 加 YaRN 在 8K/16K 有提升,但 UA 仍以更小缓存取得更高总分。
启示与展望
该结果面向以 Llama3 架构、RoPE 与 GQA 为基础、在 4k 上下文上预训练的 1B 与 3B 模型,训练数据为 Dolma v1.7 并按 Chu 等的数据配比重加权,长上下文通过 YaRN 扩展。适用场景是需要高检索与追踪能力、且能接受按头自适应稀疏缓存的长上下文推理;论文的混合实验显示,把 UA 与全注意力或 Mamba 层组合可提升聚合类任务,因此 UA 更适合作为更大长上下文架构中的缓存高效注意力组件。阈值可按内存—质量权衡选择,论文以 0.05 与 0.15 作为较温和与较激进的代表工作点。
聚合类任务上 UA 相对较弱,论文将其归因于需要整合广泛分布弱证据的任务特性,并指出混合架构可缓解;构造衰减与剪枝掩码带来额外计算,降低这部分开销仍是系统优化目标。当前实现尚未包含完全融合、面向推理优化的动态缓存管理器,附录 H 的延迟收益来自朴素实现,因此动态稀疏缓存能否在真实服务中兑现收益仍待专门稀疏内核与系统支持。从已有检查点继续训练适配可行,但与从头预训练仍有明显差距。此外,正文中若干公式与数值在解析文本中缺失,若需精确复现超参数与压缩比,应查阅原文附录与表格。
