发表时间: 2025-07 · ACL 2025
原文: https://aclanthology.org/2025.acl-long.1211
作者/机构:Fangyuan Xu (New York University), Tanya Goyal (Cornell University), Eunsol Choi (New York University)
一句话结论 提出 RefreshKV 推理方法,通过在长文本生成中交替使用全局注意力并动态刷新局部小 KV Cache,在保持与现有 KV 驱逐方法同等加速比的前提下,解决了原有方法在长输出任务上性能崩溃的问题。
要解决什么问题 大语言模型在处理长上下文时,KV Cache 的显存占用随长度线性增长,注意力计算量呈二次方增长,导致推理延迟极高。为了加速,现有的 KV Cache 压缩方法(如 StreamingLLM、H2O、SnapKV)通常会在生成过程中永久性地驱逐被认为不重要的历史 token,只保留一个较小的 KV Cache。这种做法在生成短文本时效果尚可,但在长文本生成任务中会遭遇严重的卡点。卡点的核心机制在于:大多数 KV 压缩是一次性或不可逆的,它们基于早期的注意力分数做出了硬性裁剪。然而在长输出任务中(例如将长 HTML 转换为 TSV,或多跳检索式的生成),随着解码的不断推进,模型当前的注意力模式会发生剧烈变化,后续解码往往需要用到早期被判定为“不重要”并被丢弃的 token。由于被永久驱逐的 token 无法恢复,导致模型在长输出场景下出现严重的信息丢失和性能断崖式下跌,甚至完全无法完成任务。
怎么做的 核心思路是放弃“永久驱逐”,改为在推理时维护完整的全局 KV Cache,但在绝大部分解码步中只对一个动态构建的局部小 KV Cache 进行注意力计算,并定期通过全局注意力来“刷新”这个小 Cache。这种设计绕开了永久驱逐导致的不可恢复问题,既利用小 Cache 降低了计算量和数据搬运延迟,又通过定期回看全局保证了信息的完整性。方法由三个关键机制构成。第一是预填充与初始化,在 Prefill 阶段对完整输入做全局注意力,得到最后一个 token 对历史 token 的注意力分数,通过最大池化选取分数最高的 K 个 token 初始化局部缓存 $C_p$。第二是动态触发机制,为了决定何时进行高代价的全局注意力,方法没有采用固定的步长,而是设计了基于查询向量相似度的自适应调度。每隔 $S$ 步,计算当前步输入 token 的查询向量与上一次全局注意力步的查询向量之间的余弦相似度:
$$\text{CosineSim}(q_t, q_{\text{last\_full}}) < s$$其中 $q_t$ 是当前层所有注意力头平均后的查询向量。若相似度低于设定的阈值 $s$,说明注意力模式发生了偏移,此时触发全局注意力;否则继续使用局部缓存解码。第三是交替解码与缓存刷新,在局部解码步,模型仅使用 $C_p$ 生成下一个 token $y_t \sim M(C_p)$,并将新 token 的 KV 存入 $C_p$,为了维持 $C_p$ 大小不变,会剔除掉在上次全局注意力中得分最低的 token;在全局解码步,模型首先将局部解码产生的新 token 同步到全局缓存 $C_f$ 中,随后使用 $C_f$ 生成下一个 token $y_t \sim M(C_f)$,并获取新的全局注意力分数 $\mathbf{a_L}$,最后根据 $\mathbf{a_L}$ 重新选出 top-K 的 token 来彻底刷新局部缓存 $C_p$。
效果如何 实验在单张 80GB A100 GPU 上进行,结合 Flash Attention 测试了 Llama-3.1-8B 和 Qwen2-7B 两款支持 128K 上下文的模型。对比基线包括:代表全量计算的 Vanilla attention,以及代表永久驱逐路线的 StreamingLLM(保留 Sink 和最近 token)、H2O(保留高累计分数 token)和 SnapKV(基于提示词末尾注意力选 token)。在 16K 语言建模任务中,RefreshKV 达到了与 SnapKV 相当的推理速度(例如 Llama 生成耗时从全局的 7.5 秒降至 6.3 秒),且困惑度更低。最能说明问题的量化结果来自长输入长输出任务,在 18K 输入的 HTML 转换为 TSV 任务中,SnapKV 和 H2O 等驱逐方法完全失效,F1 分数直接掉到 0,而 RefreshKV 成功恢复了 Vanilla 全局注意力 52% 的性能;在作者构造的 Chain-of-key 任务中,驱逐基线无法生成超过 2 个 key 的有效链条,准确率低于 20%,而 RefreshKV 是唯一能持续生成长链条的方法。此外,使用 RefreshKV 的设定对模型进行继续预训练可以进一步提升性能。该方法的局限性在于:它必须在显存中始终保留完整的全局 KV Cache,因此完全没有降低长上下文推理的显存占用峰值,仅优化了计算量和访存延迟,且动态相似度阈值目前是全局统一设定的,未来仍需探索更精细的逐层调节策略。
大型语言模型(LLMs)在处理长上下文输入并生成长序列时,面临着极高的计算和内存开销。随着上下文长度的增加,存储键值(KV)缓存的内存线性增长,而注意力计算呈二次方扩展,导致推理延迟极高。现有的主流加速方法(如基于驱逐的方法)通过构建较小的KV缓存来缓解这一问题,这些方法在生成短序列时表现良好,但在长文本生成任务中性能会迅速下降。其根本原因在于,大多数KV压缩是一次性完成的,过早地移除了在后续生成中可能仍然有用的token,且一旦移除便无法恢复。
为了解决这一问题,本文提出了一种新的推理方法——RefreshKV。该方法在生成过程中灵活地在“完整上下文注意力”和“输入token子集注意力”之间交替进行。在每次执行完整注意力计算后,RefreshKV会根据整个输入的注意力模式来更新较小的KV缓存。具体而言,该方法具有以下核心创新点:
1. 揭示了现有KV缓存驱逐方法在面临具有挑战性的长文本生成任务时的失效问题,并提出了一个新任务(Chain-of-key,键链生成)来暴露这种弱点,该任务要求模型更全面地记住输入上下文。
2. 引入了RefreshKV推理方法,该方法在整个推理过程中保留完整的KV缓存(不减少内存占用),但通过动态构建的小KV缓存执行注意力计算以实现推理加速。该方法不采用固定的更新计划,而是通过比较当前步与上一次完整注意力步的查询向量(Query embedding)相似度,在相似度较低时动态触发完整注意力步并刷新小缓存。
3. 将该方法应用于现成的LLMs(Llama-3.1-8B和Qwen2-7B),在各种长文本生成任务中实现了与基于驱逐的方法相当的加速比,同时大幅提高了性能(例如在HTML转TSV任务中恢复了大量性能)。此外,实验表明,在RefreshKV推理设置下对模型进行持续预训练(Continued Pretraining)可以带来进一步的性能提升。
背景与设置:假设 $M$ 是一个语言模型,$\mathtt{x}$ 是一个输入token序列,$x = x_1, \cdot\cdot\cdot x_L$。在推理时,$M$ 分两个阶段生成输出token序列 $\hat{y} = y_1, \cdots y_N$:(1)预填充阶段(Pre-filling stage),在此阶段 $M$ 摄入输入并为所有 $L$ 个token构建KV缓存;(2)生成阶段(Generation stage),在此阶段模型每次从条件分布 $P_M(y_i | x, y_1 \cdot\cdot\cdot y_{i-1})$ 中采样一个token $y_i$。在每一步中,模型关注KV缓存中的token,并更新缓存以包含当前token的键值对。
延迟增加的原因:生成阶段推理延迟增加有两个主要原因。首先,注意力计算随输入长度 $L$ 呈二次方增加。其次,较大的 $L$ 需要维护过去token的大型KV缓存,由于完整的KV缓存需要从GPU的高带宽内存(HBM)中移动,这会引发延迟。
现有方法的局限性与设计原则:先前的方法,如 $H_2O$ 【38, H2o: Heavy-hitter oracle for efficient generative inference of large language models+2024+NeurIPS】 和 SnapKV 【17, SnapKV: LLM knows what you are looking for before generation+2024+NeurIPS】,通过在解码过程中永久驱逐“不重要”的token以保持较小的KV缓存来解决这个问题。虽然这些方法已被证明对“大海捞针”(NIAH)等短文本生成任务有效,但其潜在的缺点是会过早地移除对后续生成步骤有用的token。基于这一观察,作者没有采用这种严格的策略,而是提出定期执行对所有token的完整注意力计算,并基于注意力模式构建小缓存,从而定期更新小KV缓存。由于缓存只是偶尔更新,该方法通过关注小缓存,同时减少了注意力计算量和数据移动量。
算法概述:RefreshKV的伪代码如图2所示。该算法接收语言模型 $M$ 和输入token序列 $x_1, ..., x_L$ 作为输入。首先,模型使用输入序列进行预填充。接着,算法在完整注意力(full attention)和部分注意力(partial attention)之间交替进行。该方法维护两个独立的KV缓存:$C_f$ 和 $C_p$,分别对应于完整注意力步骤和部分注意力步骤中使用的KV缓存。
预填充阶段(Prefilling stage):给定输入 $x_1, ..., x_L$,首先对 $M$ 进行带有完整注意力的预填充,并初始化包含 $L$ 个token的完整KV缓存 $C_f$。同时,获取最后一个token $x_L$ 的注意力分数 $\mathbf{a_L}$。为了确定要保留的前 $K$ 个token(topK),遵循先前工作 SnapKV 【17, SnapKV: LLM knows what you are looking for before generation+2024+NeurIPS】 的做法,对周围token的注意力分数采用最大池化(max pooling)操作,而不是直接使用原始注意力分数,以此来保持信息的完整性。
决定何时使用完整缓存解码(Deciding when to decode with full cache):算法需要决定何时在对所有token执行注意力和对较小缓存执行注意力之间进行切换。一种直接的方法是使用固定计划,即每 $S$ 步执行一次完整注意力。然而,这会对所有层和输入文本强制执行相同的计划。相反,作者提出了一种基于当前步的查询向量(query vector)与最近一次完整注意力步的查询向量之间相似度的自适应计划。直观地说,如果当前步某一层和特定头(head)的查询向量与该层最近一次完整注意力步的查询向量相似,那么它们的注意力模式也应该相似。因此,只有当这种相似度低于某个阈值时,才执行完整注意力步骤。
动态调度具体实现:具体而言,在每第 $S$ 个解码步骤,对于每一层 $l$,首先决定是否需要执行完整注意力。计算输入token $t$ 在层 $l$ 中所有查询头平均后的查询向量,与该层最近一次完整注意力步的平均查询向量之间的余弦相似度。如果相似度高于阈值 $s$,则使用部分缓存 $C_p$ 进行解码;否则,在层 $l$ 使用 $C_f$ 进行解码。为了最小化相似度检查的计算开销,仅每隔 $S$ 步执行一次此操作;这被称为查询比较(query comparison, QC)步长。
使用部分缓存解码(Decoding with partial cache):在每个部分注意力步骤中,使用 $C_p$ 计算注意力以生成下一个token $y_t \sim M(C_p)$,并存储输入token的KV缓存。这减少了注意力计算的浮点运算次数(FLOPs)以及由于KV缓存移动带来的延迟(因为只需要移动较小的KV缓存 $C_p$,而不是较大的完整KV缓存 $C_f$,其中 $|C_p| \ll |C_f|$)。随着解码每个额外的token并用这个新生成的token更新KV缓存,为了维持 $C_p$ 的大小,会从 $C_p$ 中移除在完整注意力步骤中注意力分数最低的token对应的KV。需要注意的是,如果在预填充后部分缓存从未被刷新,那么使用 $C_p$ 解码等同于 SnapKV 【17, SnapKV: LLM knows what you are looking for before generation+2024+NeurIPS】。
使用完整缓存解码(Decoding with full cache):在每个完整注意力步骤中,首先用使用 $C_p$ 解码的token的键值对来更新完整KV缓存 $C_f$。接着,使用完整KV缓存 $C_f$ 生成下一个token $y_t \sim M(C_f)$,并获取注意力分数 $\mathbf{a_L}$。最后,根据 $\mathbf{a_L}$ 基于前 $K$ 个token刷新部分缓存 $C_p$。
内存与时间需求(Memory and Time requirements):该方法的内存需求与标准注意力(vanilla attention)相似,因为并没有从KV缓存中永久驱逐任何token。然而,正如后续实验所示,其解码延迟与其他KV缓存驱逐方法处于同一水平。
数据集名称、规模及用途:
模型架构关键参数:采用两个长上下文语言模型 Llama-3.1-8B 和 Qwen2-7B,两者均支持高达128K tokens的输入。实验中部分缓存大小 $K$ 统一设置为输入长度 $L$ 的1/8(对于最长的NovelSumm 100K输入,设置 $K=4096$,即 $1/25 L$)。对于RefreshKV,查询比较步长(QC stride)设置为 $\{5, 10\}$,Llama-3.1-8B的相似度阈值 $s$ 设为0.85,Qwen2-7B设为0.95。所有任务均使用贪婪解码(greedy decoding)。
本文提出了RefreshKV,这是一种推理时方法,通过从一个小型的、动态的KV缓存中解码来加速长上下文输入的长文本生成,该缓存基于相邻token的注意力模式进行更新。与先前从上下文中永久驱逐token的工作相比,RefreshKV保持了完整的KV缓存,并在针对完整KV缓存和小型KV缓存的推理之间交替进行。将该方法应用于两个现成的长上下文模型表明,与基于驱逐的方法相比,该方法在长文本生成任务上减少了推理的挂钟时间,同时更好地保持了性能。最后,研究表明,使用RefreshKV对模型进行持续预训练可以进一步改善性能与效率的权衡。未来的工作可以探索其他的调度策略(如按层设置不同的阈值),并将其扩展到其他模态(如视觉Transformer)。
实施细节
与Flash Attention的兼容性:FlashAttention 【11, FlashAttention-2: Faster attention with better parallelism and work partitioning+2024+ICLR】 通过直接产生注意力块的输出而不存储 $O(L^2)$ 的注意力矩阵,从而减少了GPU上的数据移动,显著提高了标准注意力计算的效率。然而,RefreshKV需要依赖这些注意力分数在完整注意力步骤中选择前 $K$ 个token并构建部分KV缓存 $C_p$。为了使该方法与Flash Attention兼容,实现了一个额外的步骤,在完整注意力步骤中重新计算注意力分数。由于并非在每个生成步骤都执行完整注意力,这不会引入显著的开销。对于需要访问注意力分数的方法(如 $H_2O$),同样应用此过程使其与Flash Attention兼容。
基线设置:对于 StreamingLLM 【32, Efficient streaming language models with attention sinks+2023+ArXiv】,遵循原始论文,维护一个包含4个sink token和 $K-4$ 个最近token的缓存。对于 $H_2O$ 【38, H2o: Heavy-hitter oracle for efficient generative inference of large language models+2024+NeurIPS】,遵循原始论文将heavy hitter大小和最近缓存大小各设置为 $K/2$。对于 SnapKV 【17, SnapKV: LLM knows what you are looking for before generation+2024+NeurIPS】,遵循原始论文将RefreshKV和SnapKV的观察窗口大小设为1,卷积核大小设为7。对于具有分组查询注意力(GQA)的模型,SnapKV和 $H_2O$ 应用相同的聚合方法(在所有查询头上取最大值)。
持续预训练:从 RedPajama 数据集的 Arxiv 划分中随机采样了 200k 个序列的子集,并过滤掉少于 8192 个 token 的序列。对 Llama-3.1-8B 进行了 1 个 epoch 的训练,全局批大小为 64,学习率为 5e-6。使用 20 个预热步(warm-up steps)以及权重衰减为 0 的线性调度。优化器使用 AdamW。使用完全分片数据并行(FSDP)【39, Pytorch fsdp: Experiences on scaling fully sharded data parallel+2023+Proc. VLDB Endow.】 和 8-bit 优化器 【12, 8-bit optimizers via block-wise quantization+2021+CoRR】 来提高训练效率。训练在 4 张 H100 80 GB GPU 上完成。
内存与时间需求比较:
比较了RefreshKV与基线方法的内存和注意力计算需求。将部分缓存的大小设置为与基于驱逐的方法的完整缓存大小相同。在这种设置下,RefreshKV需要比基于驱逐的基线更大的KV缓存内存($L+K$ 对比 $K$),但与标准注意力相似($L+K$ 对比 $L$,其中 $K \ll L$)。然而,其解码延迟与基线方法相当。RefreshKV的效率取决于两个超参数:部分缓存大小 $K$,以及决定执行完整注意力频率的QC步长和 $s$。通过设置 $K \ll L$ 和一个较大的 $S$,可以实现与基于KV驱逐的基线相似的挂钟时间。
GQA模型的注意力分数聚合:
对于具有分组查询注意力(GQA)的模型,比较了不同注意力分数聚合方法在语言建模任务上的结果。实验表明,在同一组的查询头上进行聚合比仅使用其中一个头的注意力分数效果更好,其中取最大值(max)的表现略好于取平均值(mean)。
为查询相似度调度调整阈值s:
为了为动态调度选择相似度阈值 $s$,在 RedPajama 数据集 Book 划分的 50 个保留样本上运行了 RefreshKV。评估了 $QC$ 步长为 $\{5, 10\}$,阈值 $s$ 为 $\{0.80, 0.85, 0.90, 0.95\}$ 的情况。结果显示,对于 Llama-3.1-8B,设置 0.85 的阈值在两个步长下均取得了与 0.90 和 0.95 相似的性能。相反,Qwen2-7B 的性能从阈值 0.80 到 0.95 持续提升。因此,将 Llama-3.1-8B 的阈值设定为 0.85,Qwen2-7B 设定为 0.95。
有效步长(Effective stride):
图4绘制了Llama-3.1-8B和Qwen2-7B在三个任务中的跨层有效步长。利用查询相似度为两个模型实现了跨层的动态步长。观察到两个模型有不同的模式,Llama-3.1-8B在前几层具有较大的步长,而Qwen2-7B在中间层具有较大的步长。不同任务的模式也略有不同,表明该方法能够根据上下文实现灵活调度。
键链(Chain-of-key)任务设置:
任务设置:为模型提供一个包含多个键的长列表,每个键由 $W$ 个单词组成,例如:apricot-waggish,其中 $W = 2$。任务要求模型生成一个包含上下文中的 $T$ 个键的序列,使得下一个键的第一个单词是当前键的最后一个单词。例如:waggish-fishery, fishery-mosquito 等,其中 $T = 5$。该任务要求模型根据之前生成的内容在上下文中查找信息,类似于多跳检索。
数据生成:首先生成一个英语单词列表。然后将每个单词与另一个单词配对形成键列表。确保对于上下文中的每个键 $k_1$,存在且仅存在一个满足约束的键 $k_2$(即 $k_2$ 的第一个单词是 $k_1$ 的最后一个单词)。这些键在上下文中被随机打乱。
评估:通过有效链的长度除以 $T$ 来评估生成输出的正确性。有效链必须满足两个标准:(a)所有的键必须存在于上下文中;(b)当前键的第一个单词必须是前一个键的最后一个单词。
LongProc短输入任务结果:
任务设置:报告了 LongProc 基准测试中另外 4 个任务的结果:Path Traversal、Travel Planning、Countdown 和 Theory-of-mind tracking。这些任务的输入少于 10K 个 token。
评估与结果:遵循原始论文的评估实践,使用基于规则的验证器报告最终解决方案的正确性或准确率。结果观察到与 HTML to TSV 任务相似的趋势——大多数基线方法在任务上完全失败。而带有 $QC = 5$ 的 RefreshKV 分别为 Llama-3.1-8B 和 Qwen2-7B 恢复了完整注意力 50% 和 60% 的性能。
RULER详细结果:
遵循了包含13个任务的评估套件,将其按类型分组为:单键NIAH、多键NIAH(带干扰键)、多值NIAH、多查询NIAH、变量跟踪(多跳追踪)、常见词/高频词提取,以及问答任务。详细的各任务结果显示,对于两个模型,最好的基线(SnapKV)在需要短文本输出的任务(如单键NIAH)上取得了与RefreshKV相当的结果。然而,对于需要更长输出的任务,如多键和多值NIAH,RefreshKV优于所有基线方法。