KeyDiff: Key Similarity-Based KV Cache Eviction for Long-Context LLM Inference in Resource-Constrained Environments

发表时间: 2025-12 · arXiv:2504.15364 (NeurIPS 2025)

原文: https://arxiv.org/abs/2504.15364

Junyoung Park, Dalton Jones, Matthew J Morse, Raghavv Goel, Mingu Lee, Chris Lott (Qualcomm AI Research)

速读

一句话结论
本文提出了一种免训练的 KV 缓存淘汰方法 KEYDIFF,仅通过计算 Key 向量在几何空间中的余弦相似度来剔除冗余 Token,在严格的显存限制下实现了几乎无损的长上下文推理,并将端到端推理延迟降低了最高 30%。

要解决什么问题
现有的 KV 缓存淘汰方法(如 H2O 或 TOVA)通常需要一次性处理完整的输入提示词,这会导致中间计算过程的显存占用随序列长度线性增长,极易在显存受限的边缘设备上引发内存溢出。为了在整个预填充和生成阶段严格控制显存上限,工程上常采用分块推理策略,即把长输入切分成小块依次送入模型,并在每块计算后执行缓存淘汰。然而,原有的基于注意力分数的淘汰机制在这个场景下遭遇了严重的精度滑坡。其核心卡点在于:在分块计算时,模型只能基于当前块内的 Token 和历史缓存来计算局部注意力分数,无法预见未来的输入。这种局部视野会导致那些在全局上下文中至关重要、但在当前局部块中得分较低的 Token 被错误地提前丢弃,且这种淘汰误差会随着分块的推进不断累积。此外,依赖注意力分数意味着必须在显存中显式实例化庞大的注意力矩阵,这不仅消耗额外资源,还阻碍了底层算子层面的加速。

怎么做的
核心思路是彻底放弃使用注意力分数,转而利用 Key 向量自身的几何分布特征来评估 Token 的重要性。作者观察到一个关键现象:在注意力机制中,那些与其他 Key 的平均余弦相似度越低(即在几何空间中最独特、最孤立)的 Key,往往能在全局获得越高的注意力权重。这意味着 Key 的多样性本身就是全局重要性的强力替代指标,且该指标完全独立于未来的 Query,完美绕开了分块处理时的局部视野盲区。基于此,作者设计了 KEYDIFF。为了避免两两计算相似度带来的二次方时间复杂度,KEYDIFF 采用了一种高效的线性复杂度实现,主要由两个关键部件构成。首先,计算当前缓存中所有归一化 Key 向量的经验均值,将其作为“锚点向量”,代表当前缓存的整体冗余方向:

$$ \mu(\hat{K}) = \frac{1}{n} \sum_{i=1}^{n} \hat{k}_i $$


其中 $\hat{k}_i$ 是归一化后的 Key 向量。接着,计算每个 Key 与该锚点向量的余弦相似度,并对其取负值进行排序,保留得分最高的 $N$ 个 Token(即剔除与整体分布最相似的冗余 Token,保留最独特的 Token):

$$ S = \mathrm{topk}(-\mathrm{CosSim}(\mu(\hat{K}), k_i), N) $$
由于整个过程完全不需要访问注意力权重矩阵,KEYDIFF 不仅在分块推理时不会产生误差累积,还能无缝兼容 FlashAttention 等不实例化注意力矩阵的优化机制。对于代码或推理等高度依赖近期上下文的任务,该方法还可以零成本结合滑动窗口机制,强制保留最新生成的少量 Token。

效果如何
实验在 NVIDIA A100 80GB GPU 上进行,评测了 Llama 3.1-8B、Llama 3.2-3B、Qwen 2.5 系列以及 DeepSeek-R1-Distill 系列模型。对比基线涵盖了代表性的注意力淘汰路线:累积历史注意力权重的 H2O、利用最后一个 Token 注意力的 TOVA、基于滑动窗口平滑注意力的 SnapKV,以及保留初始 Token 的 Sink Attention。在分块大小为 128 的严格内存限制设置下,KEYDIFF 在 LongBench 长文本基准测试中表现优异:在 8K 缓存预算(约压缩 23%)下,精度较不淘汰的全量基线下降不到 0.04%;在 6K 预算(约压缩 33%)下,精度下降不超过 1.5%,全面超越上述基线。在 Math-500 推理任务上,结合滑动窗口的 KEYDIFF 同样取得了接近全量基线的成绩。在效率方面,由于成功接入了 FlashAttention,KEYDIFF 的首字生成延迟比 TOVA 和 SnapKV 降低了最高 30%。作者也坦诚了该方法的局限性:对于像 Qwen 这样原生 GQA 比例较低、KV 缓存本身已经高度压缩的模型,剔除 Token 带来的精度损失会比 Llama 更敏感,因为每个 Token 承载的信息密度更大;此外,该方法目前专为 GQA 架构设计,尚未适配 MLA 等其他注意力变体。

主要贡献

大语言模型(LLM)在处理长上下文应用(如文档摘要、代码生成、问答、检索增强生成和推理)时,键值(KV)缓存的内存占用会随着输入长度线性增长。在计算、内存和功耗受限的边缘设备等资源受限环境中,这成为了一个严重的瓶颈。现有的KV缓存驱逐策略虽然可以通过移除不重要的KV(通常通过注意力分数衡量)来限制内存开销,但它们通常一次性处理整个提示词(Prompt),在中间计算过程中会违反严格的内存限制。

为了在提示词预填充(Prefill)和词元生成(Generation)的完整推理阶段强制执行严格的内存边界,本文采用了一种分块推理策略:将输入提示词划分为较小的块,由模型顺序处理。在处理完每个块后,根据评估每个KV的驱逐策略驱逐部分缓存的KV。然而,在现有驱逐方法应用于此设置时,会导致准确率下降,因为现有方法假设可以访问全提示词的注意力,而在分块处理时,注意力仅使用当前块的词元计算,无法访问未来的块。

本文的核心创新点如下:
* 关键发现:在推理过程中,几何上具有显著特征的键(即键之间平均成对余弦相似度较低的键)往往具有较高的注意力分数。这表明,即使没有未来词元的信息,键的多样性也可以作为全局词元重要性的强大代理。
* 提出KEYDIFF方法:提出了一种无需训练、完全基于键相似度的KV缓存驱逐方法KEYDIFF。与其他依赖注意力分数的驱逐方法不同,KEYDIFF可以在严格的资源限制下处理任意长度的提示词,并且由于不依赖注意力分数,它允许使用如FlashAttention等优化的注意力机制。
* 理论基础:通过将键的多样性与注意力分数联系起来,为KEYDIFF提供了理论依据,证明KEYDIFF解决了一个最大化键多样性的最优子集选择问题。
* 卓越性能:在严格的内存预算下,KEYDIFF在Llama 3.1-8B和Llama 3.2-3B上的LongBench基准测试中,使用8K缓存预算(约减少23%的KV缓存)时,与不驱逐的基线相比,性能差距不到$0.04\%$。在Math-500推理基准测试中,Deepseek-R1-Distill-Llama-8B也观察到了接近基线的性能。
* 推理效率:与其他词元驱逐方法相比,KEYDIFF将端到端推理延迟降低了高达$30\%$。

带有KV缓存驱逐的分块提示词处理示例。长度为7的输入提示词被分为三个块,LLM中的Transformer层通过以下步骤处理每个块:(1)从输入计算键值状态,(2)计算注意力,(3)计算驱逐分数,(4)根据驱逐分数执行驱逐以满足内存限制(例如,缓存中最多可驻留4个词元)。在每次块处理后,KV缓存被更新并传递到下一轮块处理,从而满足对KV缓存施加的内存限制。
带有KV缓存驱逐的分块提示词处理示例。长度为7的输入提示词被分为三个块,LLM中的Transformer层通过以下步骤处理每个块:(1)从输入计算键值状态,(2)计算注意力,(3)计算驱逐分数,(4)根据驱逐分数执行驱逐以满足内存限制(例如,缓存中最多可驻留4个词元)。在每次块处理后,KV缓存被更新并传递到下一轮块处理,从而满足对KV缓存施加的内存限制。

背景知识与关键观察

Transformer与KV缓存机制:Transformer架构【33,Attention is all you need,2017】通过一系列Transformer块处理输入数据。因果注意力算子将输入词元投影为键($K$)、查询($Q$)和值($V$)矩阵,并计算注意力输出:$O^{\mathrm{attn}} = \mathrm{Softmax}\left(QK^\top / \sqrt{d} + M\right)V = AV$。为了避免在处理新词元时重复计算先前的KV状态,模型将计算过的KV存储在KV缓存$\mathcal{C} = (K, V)$中,从而显著降低延迟。然而,KV缓存的大小随处理的词元数线性增加,在长上下文应用中主导了内存占用。

KV缓存驱逐方法:为了限制内存占用,设定一个固定的缓存预算$N$。当新KV加入且缓存大小超过$N$时,驱逐策略$\pi_N(\mathcal{C})$会从$\mathcal{C}$中驱逐一个子集,返回包含最多$N$个KV的新缓存$\mathcal{C}^\prime$。基于注意力的驱逐策略$\pi_N^{\mathrm{attn}}$使用聚合的注意力值来对KV的相对重要性进行排序,并保留得分最高的KV。这通常需要显式地具体化注意力权重矩阵$A$,这可能是资源密集型的。

资源受限环境下的分块处理挑战:现有的驱逐策略【43,H2o: Heavy-hitter oracle for efficient generative inference of large language models,2024】【24,Transformers are multi-state rnns,2024】通常一次性处理整个输入提示词。但在资源受限环境下,中间缓存会增长到输入提示词的大小,导致内存溢出。解决方案是将输入分割成大小为$B$的非重叠块,并迭代更新缓存。然而,这种分块提示词处理为基于注意力的驱逐带来了挑战:在处理当前块时,基于注意力的驱逐方法仅保留基于当前及历史块计算出的高注意力权重的KV,这可能会过早地驱逐在未来块中具有高权重的KV,导致驱逐错误随时间累积。

键相似度与注意力分数的负相关性(关键观察):为了克服基于注意力的驱逐方法的缺陷,本文回顾了“注意力Sink”现象:LLM通常为前几个词元分配高注意力权重,无论输入是什么【35,Efficient streaming language models with attention sinks,2024】。然而,Sink词元的索引在不同头和层之间可能有所不同,并且可能位于序列中较深的位置。这引发了一个假设:高注意力分数可能由键的内在属性决定,而不是由键和查询的特定组合决定。通过检查注意力块内计算的键之间的余弦相似度,发现与其他键的余弦相似度较低的键,无论查询如何选择,都表现出较高的相对注意力分数。由于键的成对余弦相似度仅是缓存中键的函数,独立于输入查询,这种负相关性表明,这些独特的键本质上恢复了注意力Sink现象。

键的余弦相似度与注意力权重。测量自Llama 3.2-3B-Instruct以及LongBench中NarrativeQA数据集的第一个样本。为了可视化截断到前64个词元。
键的余弦相似度与注意力权重。测量自Llama 3.2-3B-Instruct以及LongBench中NarrativeQA数据集的第一个样本。为了可视化截断到前64个词元。

方法细节

基于键相似度的KEYDIFF驱逐策略:基于前述的观察结果,本文提出了KEYDIFF方法,该方法根据键的相似度从KV缓存中驱逐词元。如果缓存$\mathcal{C}$的中间大小为$n$,且预算为$N$(其中$n > N$),则策略$\pi_{\mathrm{KEYDIFF}}$定义为首先计算键矩阵$K$的成对余弦相似度矩阵$\mathrm{CosSim}(K)$,然后将其与全1向量$\mathbf{1}$相乘并取负值,最后通过$\mathtt{topk}$函数选择得分最高的$N$个索引,并据此收集保留的键和值。具体的公式表达为:$S = \mathtt{topk}(-\mathrm{CosSim}(K)\mathbf{1}, N)$,接着执行$K^\prime = \mathtt{gather}(K, S)$和$V^\prime = \mathtt{gather}(V, S)$。其中$K \in \mathbb{R}^{n \times d}$和$V \in \mathbb{R}^{n \times d}$是缓存的键和值,$\mathrm{CosSim}(K) \in \mathbb{R}^{n \times n}$是键的成对余弦相似度矩阵,其元素为$\mathrm{CosSim}(K)_{ij} = \frac{k_i \cdot k_j}{\|k_i\| \|k_j\|}$。

KEYDIFF的高效变体实现:与基于注意力的驱逐策略不同,KEYDIFF不需要访问注意力权重$A$,这使得可以使用不具体化$A$的优化注意力内核,例如FlashAttention【8,Flashattention: Fast and memory-efficient exact attention with io-awareness,2022】。然而,直接计算成对余弦相似度的时间复杂度为$\mathcal{O}(n^2)$。为了提高效率,本文将每个词元的得分计算优化为$\mathcal{O}(n)$的时间复杂度。具体做法是引入一个锚点向量(Anchor Vector)$\mu(\hat{K}) = \frac{1}{n} \sum_{i=1}^n \hat{k}_i$,其中$\hat{k}_i = \frac{k_i}{\|k_i\|}$。优化后的得分计算公式为:$S = \mathrm{topk}(-\mathrm{CosSim}(\mu(\hat{K}), \hat{k}_i), N)$。在实际实验中发现,将归一化的锚点向量$\mu(\hat{K})$替换为未归一化的均值向量$\mu(K)$,并使用未归一化的键$k$进行计算,不会损失准确性。因此,高效版本的KEYDIFF通过计算键与锚点向量的余弦相似度来决定驱逐哪些词元,保留相似度最低的KV对。

KEYDIFF概览。(1) KEYDIFF首先通过取KV缓存中键的平均值来计算锚点向量,(2) 计算键与锚点之间的余弦相似度,得出驱逐分数,颜色强度指示分数大小,(3) 保留相似度最低的KV对。
KEYDIFF概览。(1) KEYDIFF首先通过取KV缓存中键的平均值来计算锚点向量,(2) 计算键与锚点之间的余弦相似度,得出驱逐分数,颜色强度指示分数大小,(3) 保留相似度最低的KV对。
PCA可视化。(a, b, c) 使用Sink、TOVA和KEYDIFF管理的键缓存的二维PCA可视化。保留的词元为蓝色,驱逐的词元为橙色。键取自Llama3.2-3B-Instruct的第5层第3头,使用NarrativeQA数据集生成。(d) 每种KV缓存驱逐方法保留的键的PCA可视化。
PCA可视化。(a, b, c) 使用Sink、TOVA和KEYDIFF管理的键缓存的二维PCA可视化。保留的词元为蓝色,驱逐的词元为橙色。键取自Llama3.2-3B-Instruct的第5层第3头,使用NarrativeQA数据集生成。(d) 每种KV缓存驱逐方法保留的键的PCA可视化。

结合滑动窗口的KEYDIFF扩展:在推理和代码生成等任务中,最近的词元通常非常重要。为此,本文对KEYDIFF及其高效变体进行了增强,允许使用缓存预算的一定百分比来充当滑动窗口【6,Longformer: The long-document transformer,2020】。这种被称为“带滑动窗口的KEYDIFF”的扩展方法不会引入额外的复杂度或内存开销,并且在某些特定任务上取得了比普通KEYDIFF更好的结果。

KEYDIFF有效性的理论证明:为了巩固KEYDIFF的理论基础,并证明KEYDIFF最终选择的是与查询最对齐的键,本文提出了两个定理。首先,通过Lemma 3.1验证了余弦相似度与注意力分数之间的关系。假设对于固定的查询词元$q$($\|q\|=1$),存在一组键词元$\{k_i\}_{i=1}^n$使得$\|k_i\|_2^2 < M$。假设$k^*$是一个不在该集合中且具有注意力权重$w > 0$的新输入键,且$\|k^*\|_2^2 < M$。当$n \to \infty$时,存在边界:$\frac{-\log(1-w)}{2M} - 1 \leq \mathrm{CosSim}(k^*, q)$。
接着,Theorem 3.2建立了$k^*$、$q$的余弦相似度与先前键的均值$\bar{k}$之间的关系。假设$\mathrm{CosSim}(k^*, q) = \beta_q > 0$且$\mathrm{CosSim}(\bar{k}, q) = \alpha_q < 0$,则有不等式:$\mathrm{CosSim}(\bar{k}, k^*) \leq 1 + \alpha_q\beta_q - 0.5\alpha_q^2 - 0.5\beta_q^2$。
通过结合这两个定理,建立了注意力权重$w$与KEYDIFF得分$\mathrm{CosSim}(\bar{k}, k^*)$之间的关系。当$\mathrm{CosSim}(\bar{k}, q)$减小且$\mathrm{CosSim}(k^*, q)$增加(伴随注意力权重$w$增加)时,$\mathrm{CosSim}(\bar{k}, k^*)$趋向于$-1$。这意味着KEYDIFF会选择与$q$最对齐的独特键。

Llama 3.2 3B中键和查询的PCA嵌入
Llama 3.2 3B中键和查询的PCA嵌入

实验环境

实验结果

大海捞针 (Needle In a Haystack) 测试
实验评估了在6K缓存预算和$B=128$的情况下,Llama3.2-3B-Instruct在不同文档长度和“针”插入深度下的召回准确率。结果表明,对于较短的文档,KEYDIFF的表现与TOVA、SnapKV和Sink attention相似;但随着文档长度的增加,KEYDIFF的表现优于所有这三种方法(参考图6和图19)。
大海捞针测试中不同文档长度和针深度的准确率。缓存大小为6K,B=128。

LongBench 基准测试
在Llama 3.1-8B-Instruct和Llama 3.2-3B-Instruct上,使用2K、4K、6K和8K缓存预算进行测试。KEYDIFF在大多数任务中均优于其他驱逐策略,甚至在较小的缓存预算下也表现出更好的性能。特别是在测试长提示中长期依赖关系的PassageRetrieval-en (PR-en) 数据集上,KEYDIFF展现了显著的改进,即使在最小预算下也达到了接近全上下文模型的性能。使用6K缓存预算(约33%压缩率)时,准确率下降不超过$1.5\%$;使用8K缓存预算(约23%压缩率)时,准确率下降不超过$0.04\%$。相比之下,基于注意力的方法在分块处理($B=128$)时性能下降严重,因为它们无法完全具体化词元级别的注意力权重。

Math-500 推理基准测试
推理任务通常包含较短的提示词和极长的生成过程。使用DeepSeek-R1-Distill-Qwen-7B和Llama-8B模型测试了KEYDIFF(带有保留20%缓存预算的滑动窗口)和SnapKV。结果显示,配备KEYDIFF和中等KV缓存预算的Llama模型的表现与不驱逐的基线相当,甚至略好,并且优于SnapKV。

消融实验
对KEYDIFF的主要参数进行了消融研究。结果表明,锚点向量的选择(使用成对余弦相似度、所有归一化键的均值或键的中位数)对基准测试准确率影响不大。此外,在相似度度量方面,使用余弦相似度作为驱逐标准明显优于点积和欧几里得距离,这表明同时考虑键的方向和幅度对于识别要驱逐的词元至关重要。

延迟与效率 (Latency and Complexity)
测量了Llama 3.2-3B模型在不同块大小和缓存策略下的首个词元生成时间(TTFT)。由于KEYDIFF不需要具体化注意力权重,因此可以无缝结合FlashAttention使用。实验结果表明,与TOVA和SnapKV相比,KEYDIFF将端到端推理延迟降低了高达$30\%$(参考图7)。
使用Flash Attention的Llama 3.2-3B在块提示处理大小为64、128和256时,不同驱逐策略的首词元时间(TTFT)。

结论

本文提出了一种基于键相似度的免训练KV缓存驱逐方法KEYDIFF,使大语言模型能够在内存和计算受限的环境中高效运行。通过最小化KV缓存中键的成对余弦相似度,KEYDIFF最大化了键的多样性,并证明了其能够有效识别最重要的词元。在同等内存限制下,KEYDIFF显著优于最先进的KV缓存驱逐方法,在LongBench上实现33%和23%的KV缓存内存减少时,准确率仅比不驱逐基线下降1.5%和0.04%。未来的工作计划是将KEYDIFF扩展到其他注意力变体,例如多头潜在注意力(Multi-Head Latent Attention, MLA)。

附录细节

运行时间和内存复杂度分析
对于给定的块大小$B$和缓存预算$N$,KEYDIFF需要$\mathcal{O}(N + B)$的运行时间和内存。TOVA同样需要$\mathcal{O}(N + B)$,因为它只计算注意力矩阵的最底行。Sink attention需要$\mathcal{O}(N)$的内存和运行时间。SnapKV由于要在大小为$L$的滑动窗口上计算注意力,其复杂度和内存为$\mathcal{O}((N + B)L)$。H2O需要累积所有词元的注意力权重,复杂度为$\mathcal{O}(NB + B^2)$。
在详细的FLOP计算中,KEYDIFF计算锚点向量和余弦相似度所需的总FLOP数为$(12d + 97)n + 3d + 94$(其中$d$为隐藏层维度,$n$为键的数量)。这相对于注意力算子的二次复杂度而言,是关于键数量$n$的线性复杂度。

从优化角度推导KEYDIFF
为了利用键差异性的观察结果,可以将问题转化为最小化保留在缓存中的每个键的成对余弦相似度之和。这可以表述为一个带有预算$N$的约束组合优化问题。通过将键$k_i$归一化为$\hat{k}_i$,并用子集均值$\mu(\hat{K}_S)$近似全局均值$\mu(\hat{K}) = \frac{1}{n} \sum_{i=1}^n \hat{k}_i$,该目标被松弛为:
$\underset{\boldsymbol{S}}{\mathrm{minimize}} \sum_{i \in \boldsymbol{S}} \hat{k}_i \cdot \mu(\hat{K})$,约束条件为$|\boldsymbol{S}| = N$。
这个松弛问题的最优解可以通过根据词元与$\mu(\hat{K})$的余弦相似度进行排序,并选择最小的$N$个来找到,这正是KEYDIFF算法的理论推导过程。通过计算Gram矩阵的对数行列式$\log(\det(KK^T))$(代表键在空间中跨越的体积),验证了KEYDIFF保留的键比TOVA和Sink跨越了更大的空间体积。

Qasper数据集中log det(K K^T)的分布。值越大意味着键缓存跨越了更多的键空间。KEYDIFF保留的键跨越的周围空间体积大于TOVA或Sink attention。
Qasper数据集中log det(K K^T)的分布。值越大意味着键缓存跨越了更多的键空间。KEYDIFF保留的键跨越的周围空间体积大于TOVA或Sink attention。

注意力Sink与近似共线性的几何解释
研究表明,大多数键和查询在欧几里得空间中近似共线(即$\mathrm{CosSim}(x_i, x_j) \gg 0$),大多数键和查询与它们的均值之间的角距离很小。然而,在所有注意力头中,平均键和平均查询具有负的余弦相似度。此外,键和查询的$L_2$范数紧密聚集在一个相对固定的值附近,这意味着方向对注意力分数的影响大于范数大小。
这些观察表明,大多数键和查询组合会产生均匀较小的注意力逻辑值(Logits)。注意力头通过将少数键对齐到平均查询的方向,选择性地增加这些键的注意力激活,从而防止过度混合(Over-mixing)。因此,键词元的重要性可以通过键与平均键之间的角距离来衡量。与平均键具有最大角度差异的键与平均查询对齐,从而获得非常大的注意力权重。

各头和层中平均键和平均查询的余弦相似度。
各头和层中平均键和平均查询的余弦相似度。

相关性分析
测量了Llama-3.2-3B-Instruct模型每一层中键的余弦相似度与注意力分数之间的Spearman等级相关性($\rho$)。结果显示,所有层中平均存在持续的高度相关性($\rho \approx 0.94$),这表明几何上具有显著特征的键(即成对余弦相似度低的键)与获得较高注意力分数的词元强烈对齐。

扩展实验评估
* 电话簿查找(检索关键型评估):在要求模型从长列表中检索与查询名称对应的电话号码的任务中,随着上下文长度的增加,KEYDIFF的准确率下降比基于注意力的基线方法更加平缓。
* RULER基准测试:在社区维护的KVPress RULER基准测试中,KEYDIFF在Llama-3.2-3B和Qwen-3-8B架构上的得分始终具有竞争力或优于TOVA、SnapKV等方法。
* 设备端延迟评估:在Android智能手机上测量了计算键驱逐分数的延迟。KEYDIFF在小缓存大小时与竞争方法表现相当,而在缓存大小增加时(如8192),其评分延迟显著低于SnapKV和H2O,证明了其在边缘硬件上的可扩展性和极小的开销。