OBCache: Optimal Brain KV Cache Pruning for Efficient Long-Context LLM Inference

发表时间: 2025-10 · arXiv:2510.07651 (ICML 2026)

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

Yuzhe Gu, Xiyu Liang, Jiaojiao Zhao, Enmao Diao

速读

一句话结论
本文提出了 OBCache,将大模型 KV Cache 驱逐转化为层级结构化剪枝问题,通过计算移除 Token 对注意力输出的扰动来精准评估其重要性,在不增加显著开销的前提下全面提升了长上下文推理的准确率。

要解决什么问题
大语言模型在处理长上下文时,自回归生成的特性要求必须缓存上下文中所有历史 Token 的 Key 和 Value(KV)状态。这种机制导致显存占用和延迟随序列长度与批次大小呈线性爆炸式增长(例如运行 1M 上下文的 LLaMA-3.1-8B 模型需要超过 120GB 显存,远超单卡容量)。为了缓解显存卡点,现有的 KV Cache 驱逐方法通常利用注意力稀疏性来丢弃冗余 Token。然而,原有做法卡在了一个关键的评估机制上:它们高度依赖启发式地累加注意力权重来对 Token 进行排序,却完全忽略了 Value 状态以及它们对最终模型输出的实际贡献。简单累加注意力分数无法真实反映移除某个 Token 后对注意力输出造成的实际扰动,而注意力输出会直接影响隐藏层状态和下游预测。这种盲区会导致模型误留对输出影响微乎其微的 Token,或者错误丢弃那些注意力权重不高但对最终预测至关重要的 Token,最终在长上下文任务中导致严重的精度下降。

怎么做的
OBCache 的核心思路是将 KV Cache 驱逐严格定义为基于最优脑损伤(Optimal Brain Damage, OBD)理论的层级结构化剪枝问题。它绕开了仅看注意力权重的盲区,转而直接衡量真正的影响:即如果剪枝掉某个 KV 对,会对近期的历史注意力输出造成多大的局部扰动。因为在推理时无法获取未来的状态,作者用这种历史输出的扰动作为真实驱逐误差的有效代理指标。具体而言,OBCache 将 Value 和 Key 矩阵视为剪枝变量,目标是最小化剪枝前后的注意力输出误差。为了避免极高的重算代价,该方法利用二阶泰勒展开对扰动进行近似,推导出了三种闭式(closed-form)的 Token 显著性得分。第一是孤立 Value 剪枝得分,仅考虑扰动 Value 状态,得分由注意力权重的平方和与 Value 向量的 L2 范数平方相乘构成,这在数学上统一了近期基于 Value 范数的启发式方法:

$$ S_p^{\mathrm{value}} = \sum_i |\mathbf{A}_{i,p}|^2 \|\mathbf{v}_p\|^2 $$


第二是孤立 Key 剪枝得分,仅考虑扰动 Key 状态。因为改变 Key 会影响整个注意力分布,其误差通常比剪枝 Value 更大。该得分不仅包含注意力权重和 Softmax 前的 Logits($\mathbf{Z}$),还计算了 Value 向量与注意力输出 $\mathbf{o}_i$ 之间的偏差:

$$ S_p^{\mathrm{key}} = \sum_i |\mathbf{A}_{i,p}\mathbf{Z}_{i,p}|^2 \|\mathbf{v}_p - \mathbf{o}_i\|^2 $$
第三是联合剪枝得分,将 Key 和 Value 作为联合剪枝单元,综合了上述两项以及它们交互作用的交叉项,提供了最全面的扰动估计:
$$ S_p^{\mathrm{joint}} = S_p^{\mathrm{value}} + S_p^{\mathrm{key}} + 2 \sum_i |\mathbf{A}_{i,p}|^2 \mathbf{Z}_{i,p} (\|\mathbf{v}_p\|^2 - \mathbf{v}_p^\top \mathbf{o}_i) $$
通过这套设计,OBCache 不仅引入了输出感知(output-aware)的信号,还能作为插件无缝替换现有方法中的注意力得分。无论是 Prefill 阶段的静态驱逐,还是 Decoding 阶段的动态驱逐,只需计算上述闭式解即可精准定位并保留关键 Token。

效果如何
实验在单张 A100 GPU 上进行,评估了 LLaMA-3.1-8B 和 Qwen-2.5-7B 模型。测试任务涵盖长上下文检索与推理(RULER 4K 和 32K)、真实长文本理解(LongBench,平均 16K 词)以及语言建模(PG19,平均 70K Token)。对比基线分为两类:一类是代表纯注意力路线的 $\mathrm{H_2O}$(全局累加加近期窗口)、TOVA(仅看最新 Query)、SnapKV(近期窗口加池化平滑)和 AdaKV(自适应多头预算分配);另一类是代表 Value 感知路线的 VATP 和 CriticalKV。实验表明,将四种注意力基线中的启发式得分替换为 OBCache 得分后,性能获得全面提升。在 RULER 4K 任务中,结合 OBCache 使得 $\mathrm{H_2O}$ 的平均准确率提升了超过 10%,TOVA 提升了 2% 到 5%。即使是当前最强的基线 AdaKV,在 30% 缓存预算的 Query 不可知(Query-agnostic)严苛设置下,替换为 OBCache 后准确率依然大幅提升了近 15%。在 LongBench 任务中,10% 极低缓存预算下,OBCache 也能为 AdaKV 带来 1.2 到 2.6 的指标增长。在 PG19 的动态解码测试中,OBCache 维持 1024 个 Token 预算时的困惑度显著优于 $\mathrm{H_2O}$。消融实验证明,包含 Key 扰动的得分始终优于仅考虑 Value 的得分。作者也坦诚了该方法的局限性:随着扰动窗口增大,早期 Token 会因为参与了更多 Query 的计算而累积过高的得分,产生结构性偏差,目前仍需像 $\mathrm{H_2O}$ 那样保留一个固定的近期窗口来缓解;此外,联合剪枝得分在动态解码阶段并未表现出优于孤立 Key 剪枝得分的效果,暗示两者的相加组合方式仍有改进空间。

A1 主要贡献

大型语言模型(LLMs)在处理长序列时面临巨大的内存开销,因为自回归生成需要缓存上下文窗口中的所有键值(KV)状态。KV缓存的大小随序列长度和批处理大小线性扩展,导致了严重的内存和延迟瓶颈(例如,具有1M上下文的LLaMA-3.1-8B需要超过120GB的KV缓存)。现有的缓存驱逐方法主要通过利用注意力稀疏性来解决此问题,但它们通常仅基于累积的注意力权重启发式地对Token进行排序,而忽略了Token移除对最终注意力输出的真实影响(即忽略了Value状态的作用)。

为了解决这一问题,本文提出了Optimal Brain Cache (OBCache),这是一个原则性的评分框架,将KV缓存驱逐形式化为逐层结构化剪枝问题。
本文的主要创新点如下:
1. 引入了OBCache评分框架,通过测量驱逐KV Token在注意力输出中引起的扰动来估计Token的显著性,该框架可以无缝集成到现有的驱逐管道中以改善Token选择。
2. 首次基于最优脑损伤(Optimal Brain Damage, OBD)理论,将KV缓存驱逐进行了理论上的形式化。推导出了针对孤立Value、孤立Key以及联合Key-Value对的闭式显著性分数,并证明了现有基于注意力的启发式分数只是该广义框架的特例。
3. 在LLaMA-3.1和Qwen-2.5模型上进行了广泛实验,结果表明,用OBCache的感知输出(output-aware)分数替换现有方法中的启发式注意力分数,能够一致地提高长上下文任务的准确性。

A3 背景知识/关键Observation/设计原则

KV缓存压缩策略: 缓存驱逐方法通过消除冗余的KV状态来直接减少内存占用。例如,StreamingLLM发现注意力下沉现象并保留初始和最新Token;$\mathrm{H_2O}$ 【1,H2O: Heavy-hitter oracle for efficient generative inference of large language models+2023+NeurIPS】跨所有查询位置累积注意力权重;TOVA 【2,Transformers are multi-state rnns+2024+EMNLP】仅考虑最新查询的注意力;SnapKV 【3,Snapkv: Llm knows what you are looking for before generation+2024+NeurIPS】在小观察窗口内聚合分数并应用池化。为了改进纯注意力的指标,VATP 【4,Attention score is not all you need for token importance indicator in kv cache reduction: Value also matters+2024+EMNLP】和CriticalKV 【5,Identify critical kv cache in llm inference from an output perturbation perspective+2025+arXiv】引入了Value状态的范数进行缩放。其他正交方法包括稀疏注意力、缓存合并以及自适应预算分配(如AdaKV 【6,Adakv: Optimizing kv cache eviction by adaptive budget allocation for efficient llm inference+2026+NeurIPS】)。

模型剪枝理论: 经典的剪枝框架如最优脑损伤(OBD)【7,Optimal brain damage+1989+NeurIPS】通过泰勒级数二阶近似来量化剪枝单元对损失函数造成的扰动。由于全局二阶统计在LLM中计算效率低下,近期方法转向逐层公式化以最小化层输出的变化 【8,Sparsegpt: Massive language models can be accurately pruned in one-shot+2023+ICML】。本文将这种逐层静态权重剪枝范式扩展到了动态KV缓存的驱逐中。

符号与预备知识: 在Prefill阶段,注意力输出计算为 $\mathbf{O} = \mathbf{A}\mathbf{V}, \mathbf{A} = \sigma(\mathbf{Z}), \mathbf{Z} = \frac{\mathbf{Q}\mathbf{K}^\top}{\sqrt{d}}$,其中 $\mathbf{K}$ 和 $\mathbf{V}$ 被缓存。在Decoding阶段,第 $t$ 步生成新Token,将其 $\mathbf{k}_t, \mathbf{v}_t$ 附加到缓存中,使得序列长度扩展为 $s = l + t$。当缓存大小 $s$ 超过预定义预算 $N$ 时,触发缓存驱逐,算法从 $\mathbf{K}$ 和 $\mathbf{V}$ 中选择 $s - N$ 行进行永久删除。

A2 方法细节

通过扰动最小化进行缓存驱逐: 缓存驱逐可以看作是逐层结构化剪枝,但由于生成的自回归特性,移除第 $s$ 步的键值对实际上只会影响未来的注意力输出 $\mathbf{o}_{s+1}, \mathbf{o}_{s+2}, \dots$,而这些在驱逐时是不可访问的(称为真实驱逐误差)。因为真实误差不可直接获得,作者观察到可以通过测量修剪相应KV向量时最近历史注意力输出($\mathbf{o}_s, \mathbf{o}_{s-1}, \dots$)的扰动来有效近似。这被称为修剪引起的驱逐误差。基于此,作者将 $\mathbf{V}$ 和 $\mathbf{K}$ 视为剪枝变量。设 $\widehat{\mathbf{V}} = \mathbf{V} + \delta\mathbf{V}$ 和 $\widehat{\mathbf{K}} = \mathbf{K} + \delta\mathbf{K}$,位置 $p$ 的显著性分数定义为当 $\mathbf{v}_p$ 和 $\mathbf{k}_p$ 被剪枝时近期历史注意力输出 $\mathbf{O}$ 的变化:$S_p := \mathcal{L}_{e_p^\top [\widehat{\mathbf{V}} \widehat{\mathbf{K}}] = \mathbf{0}} (\widehat{\mathbf{V}}, \widehat{\mathbf{K}}) = f \bigg( \sigma \big( \frac{\mathbf{Q} \widehat{\mathbf{K}}^\top}{\sqrt{d}} \big) \widehat{\mathbf{V}} \Big|_{e_p^\top [\widehat{\mathbf{V}} \widehat{\mathbf{K}}] = \mathbf{0}} - \sigma \big( \frac{\mathbf{Q} \mathbf{K}^\top}{\sqrt{d}} \big) \mathbf{V} \bigg)$。这里采用 Frobenius 范数平方 $\|\cdot\|_F^2$ 作为范数函数 $f(\cdot)$。
OBCache评分机制概述

二阶泰勒近似求解: 直接为每个位置 $p$ 重新计算上述公式在计算上是不可行的。受OBD理论启发,作者在未扰动点 $(\mathbf{V}, \mathbf{K})$ 附近通过二阶泰勒展开来近似误差。设 $\mathbf{H}^{\nu\nu}, \mathbf{H}^{kk}, \mathbf{H}^{\nu k}$ 分别为 $\mathcal{L}$ 关于 $\widehat{\mathbf{V}}, \widehat{\mathbf{K}}$ 及其交叉项的Hessian矩阵。修剪引起的误差展开为:$\mathcal{L}(\widehat{\mathbf{V}}, \widehat{\mathbf{K}}) = \frac{1}{2} \delta\mathbf{V}^\top \mathbf{H}^{\nu\nu} \delta\mathbf{V} + \frac{1}{2} \delta\mathbf{K}^\top \mathbf{H}^{kk} \delta\mathbf{K} + \delta\mathbf{V}^\top \mathbf{H}^{\nu k} \delta\mathbf{K} + \mathcal{O}\big(\lVert(\delta\mathbf{V}, \delta\mathbf{K})\rVert^3\big)$。当仅修剪 $\mathbf{v}_p$ 和 $\mathbf{k}_p$ 时,一阶项消失,且利用Hessian对角块假设,显著性分数简化为:$S_p \overset{second}{=} \frac{1}{2} \mathbf{v}_p^\top \mathbf{H}_{pp}^{\nu\nu} \mathbf{v}_p + \frac{1}{2} \mathbf{k}_p^\top \mathbf{H}_{pp}^{kk} \mathbf{k}_p + \mathbf{v}_p^\top \mathbf{H}_{pp}^{\nu k} \mathbf{k}_p$。

孤立Value剪枝分数: 当仅将 $\mathbf{V}$ 作为剪枝单元时($e_p^\top \widehat{\mathbf{V}} = \mathbf{0}$),误差简化为分数:$S_p^{\mathrm{value}} = \frac{1}{2} \mathbf{v}_p^\top \mathbf{H}_{pp}^{\nu\nu} \mathbf{v}_p = \sum_i |\mathbf{A}_{i,p}|^2 \Vert\mathbf{v}_p\Vert^2$。该分数计算注意力权重矩阵第 $p$ 列的平方 $\ell_2$ 范数,并由Value向量 $\mathbf{v}_p$ 的平方 $\ell_2$ 范数进行缩放。这与VATP和CriticalKV中提出的感知Value分数相对应(它们使用的是 $\ell_1$ 范数)。

孤立Key剪枝分数: 当仅将 $\mathbf{K}$ 作为剪枝单元时($e_p^\top \widehat{\mathbf{K}} = \mathbf{0}$),误差简化为分数:$S_p^{\mathrm{key}} = \frac{1}{2} \mathbf{k}_p^\top \mathbf{H}_{pp}^{kk} \mathbf{k}_p = \sum_i |\mathbf{A}_{i,p} \mathbf{Z}_{i,p}|^2 \|\mathbf{v}_p - \mathbf{o}_i\|^2$。该分数捕获了Value向量与注意力输出之间的偏差,并由注意力权重和Softmax前的Logits共同加权。因为修剪Key会改变整个注意力分布,其通常比修剪Value引起更大的误差。

联合剪枝分数: 当 $\mathbf{V}$ 和 $\mathbf{K}$ 被作为一个组合剪枝单元时,引入交叉项,分数形式为:$S_p^{\mathrm{joint}} = S_p^{\mathrm{value}} + S_p^{\mathrm{key}} + \mathbf{v}_p^\top [\mathbf{H}^{\nu k}]_{pp} \mathbf{k}_p = S_p^{\mathrm{value}} + S_p^{\mathrm{key}} + 2 \sum_i |\mathbf{A}_{i,p}|^2 \mathbf{Z}_{i,p} (\|\mathbf{v}_p\|^2 - \mathbf{v}_p^\top \mathbf{o}_i)$。这提供了对修剪引起误差最全面的估计。这些分数既可用于Prefill阶段的一次性贪婪驱逐,也可在Decoding阶段通过累积实现动态驱逐。

与现有注意力方法的联系: 引入查询位置窗口 $w \in [1, s]$。如果将最小化目标退化为仅保留历史注意力矩阵行 $\mathbf{A}_{w:s}$ 的扰动,误差简化为 $S_p^{\mathrm{attn}} = \sum_{i=w}^s |\mathbf{A}_{i,p}|$。这正是现有方法的本质:$\mathrm{H_2O}$ 设定 $w=1$ 累积全局历史;TOVA 设定 $w=s$ 仅关注最新查询;SnapKV 使用短窗口 $w \gg 1$。OBCache通过将目标放宽为同一扰动窗口内的输出误差,不仅泛化了这些方法,还引入了感知输出的信号。

修剪引起误差作为代理的有效性: 作者在Needle-In-A-Haystack任务上进行了定性分析。首先建立一个基于"真实驱逐误差"(即在Prefill阶段驱逐Token对第一步解码输出 $\mathbf{o}_{l+1}$ 的扰动)的Oracle基线。实验表明,当扰动窗口选择合适时,精确计算的修剪代理误差能达到Oracle前 $k$ 个选择的85%召回率。OBCache的二阶闭式近似分数实现了与精确代理几乎相同的排名性能,且一致优于纯注意力分数。由于注意力因果关系导致的结构性偏差(早期Token得分偏高),保留一个固定大小的最近窗口可以进一步提高召回率。
Oracle驱逐误差识别Top-40显著Token的召回率

A4 实验环境

  • 数据集:

    • Prefill阶段:RULER(合成基准,包含检索、多跳推理等13个任务,评估4K和32K上下文);LongBench(16个真实世界数据集,涵盖QA、摘要、代码等,平均长度约16K Tokens)。
    • Decoding阶段:PG19(100本书,平均长度70K Tokens),用于评估语言建模困惑度。
  • 模型架构参数: LLaMA-3.1-8B-Instruct 和 Qwen-2.5-7B-Instruct。两者均原生支持128K上下文窗口,并采用分组查询注意力(GQA)。

  • 硬件配置: NVIDIA A100-80GB GPUs。
  • 软件配置: 基于PyTorch和Transformers库,使用KVPress代码库实现。为了提高效率,Prefill阶段的注意力计算通过FlashAttention-2实现。

A5 实验结果

  • RULER基准测试结果: 将OBCache分数集成到四种基线($\mathrm{H_2O}$, TOVA, SnapKV, AdaKV)中。在所有压缩设置(10%, 20%, 30%, 40% KV预算)和模型上,OBCache一致提高了任务准确性。例如,结合 $\mathrm{H_2O}$ 时,在RULER-4K上平均准确率提升超过10%,在32K上提升超过5%。结合最强的AdaKV基线时,在30%预算的query-agnostic设置下获得了近15%的准确率提升。消融实验表明,包含Key剪枝的分数(OBCache-K和OBCache-V&K)始终优于仅Value的变体(OBCache-V)。(支撑数据见表1及Fig 3)。
    在LLaMA-3.1-8B和Qwen-2.5-7B上的Prefill阶段缓存驱逐评估
  • LongBench基准测试结果: 同样观察到一致的性能改进,且在更高压缩率下增益更加明显。对于LLaMA模型在10%预算下,AdaKV结合OBCache-K在query-aware设置中提升+1.2,在query-agnostic设置中提升+2.6。(支撑数据见Fig 3)。
  • 与感知Value基线的比较: 在AdaKV上集成VATP和CriticalKV进行对比。在query-aware设置中,OBCache-K和OBCache-V&K表现最强。OBCache-V优于CriticalKV和VATP。这证明了OBCache的增益不仅来自于Value范数,更来自于捕获输出敏感性的Key感知评分项。(支撑数据见Fig 4)。
    OBCache分数与现有感知Value基线的比较
  • Decoding阶段动态缓存驱逐评估: 在PG19数据集上固定1024 Token预算。Sink基线(固定初始和最近Token)困惑度最高;$\mathrm{H_2O}$ 表现稍好。所有OBCache变体(随时间累积感知输出分数)在所有序列长度上始终优于 $\mathrm{H_2O}$。实验也发现OBCache-V&K并未优于OBCache-K,表明纯加法组合可能还有改进空间。(支撑数据见Fig 5)。
    在PG19测试集上的动态缓存驱逐评估

A6 结论

本文提出了OBCache,这是一个基于最优脑损伤(OBD)理论的KV缓存驱逐原则性评分框架。通过将缓存驱逐转化为逐层结构化剪枝问题,推导出了旨在最小化注意力输出影响的Token显著性分数。在长上下文基准测试上的Prefill和Decoding实验表明,OBCache一致改进了SOTA基线,以极小的开销实现了卓越的性能-压缩权衡。未来,该框架可灵活扩展至通道级KV剪枝,或通过放宽对角线近似应用于缓存合并等任务中。

A7 附录 (理论分析与实现细节)

OBCache目标函数分解: 目标 $\mathcal{L}(\widehat{\mathbf{V}}, \widehat{\mathbf{K}}) := \|\widehat{\mathbf{O}}_{w:s} - \mathbf{O}_{w:s}\|_F^2$ 可以显式分解为所有元素级平方误差的求和形式:$\sum_{i=w}^s \sum_{j=1}^d \mathcal{E}_{i,j}(\widehat{\mathbf{V}}_{:,j}, \widehat{\mathbf{K}})$,其中 $\mathcal{E}_{i,j} = |\mathrm{softmax}\big(\frac{\mathbf{q}_i \widehat{\mathbf{K}}^\top}{\sqrt{d}}\big) \widehat{\mathbf{V}}_{:,j} - \mathrm{softmax}\big(\frac{\mathbf{q}_i \mathbf{K}^\top}{\sqrt{d}}\big) \mathbf{V}_{:,j}|^2$。

二阶泰勒展开推导: 遵循OBD理论,在 $(\mathbf{V}_{:,j}, \mathbf{K})$ 处对 $\mathcal{E}_{i,j}$ 进行二阶泰勒展开。因为在展开点 $\widehat{\mathbf{O}}_{i,j} - \mathbf{O}_{i,j} = 0$,常数项消失。
* 孤立Value剪枝推导: 当仅扰动 $\mathbf{V}$ 时,一阶导数 $\frac{\partial \mathcal{E}_{i,j}}{\partial \widehat{\mathbf{V}}_{:,j}} = 2(\widehat{\mathbf{O}}_{i,j} - \mathbf{O}_{i,j})\hat{\mathbf{a}}_i$ 在展开点为0。Hessian矩阵为 $2\mathbf{a}_i^\top \mathbf{a}_i$。代入得到 $\mathcal{L}^{\mathrm{value}} = \sum_{i=w}^s \sum_{j=1}^d |\mathbf{a}_i \delta\mathbf{V}_{:,j}|^2$。施加行级剪枝约束 $e_p^\top \widehat{\mathbf{V}} = \mathbf{0}$,即可化简得到前文的 $S_p^{\mathrm{value}}$。
* 孤立Key剪枝推导: 当仅扰动 $\mathbf{K}$ 时,一阶导数同样在展开点为0。定义中间变量矩阵 $M_{p,r}^{(i,j)} = \frac{1}{\sqrt{d}} \mathbf{Q}_{i,r} \cdot \mathbf{A}_{i,p} \cdot (\mathbf{V}_{p,j} - \mathbf{O}_{i,j})$。Hessian矩阵推导为 $2 \mathrm{vec}(M^{(i,j)})^\top \mathrm{vec}(M^{(i,j)})$。代入约束 $e_p^\top \widehat{\mathbf{K}} = \mathbf{0}$ 并化简,得到前文的 $S_p^{\mathrm{key}}$。
* 联合剪枝交叉项推导: 交叉二阶导数为 $2\mathbf{a}_i \otimes M^{(i,j)}$。结合 $\mathbf{V}$ 和 $\mathbf{K}$ 的扰动形式,交叉项化简为 $\frac{2}{\sqrt{d}} \sum_{i=w}^s \mathbf{A}_{i,p}^2 \cdot (\mathbf{q}_i \mathbf{k}_p^\top) \sum_{j=1}^d \mathbf{V}_{p,j}(\mathbf{V}_{p,j} - \mathbf{O}_{i,j})$,最终得到 $S_p^{\mathrm{joint}}$。

分组查询注意力(GQA)的分数适配: 在GQA中,多个查询头共享一个Key和Value头。为了支持这种架构,修改目标为最小化同一组内各查询头输出扰动之和:$\mathcal{L} := \sum_{h \in H(g)} \|\widehat{\mathbf{O}}^h - \mathbf{O}^h\|_F^2$。由此,针对第 $g$ 个KV头的OBCache分数会引入一个对其关联查询头 $h \in H(g)$ 的额外求和操作。

OBCache分数的效率评估:
分析表明,OBCache-V的复杂度仅为 $O(W + d_{\mathrm{head}})$,而OBCache-K和OBCache-V&K的复杂度为 $O(W d_{\mathrm{head}})$。由于扰动窗口 $W$ 通常极小(如Prefill中 $W=64$,Decoding中 $W=1$),额外开销非常小。
在A100 GPU上的实证基准测试(Fig 6和Fig 7)证实,OBCache-V的Decoding延迟($<2\mathrm{ms}$)与纯注意力方法几乎相同。OBCache-K/V&K引入了轻微的额外延迟($<15\mathrm{ms}$),但远低于全缓存Decoding的成本,且不随上下文长度线性增加。在Prefill阶段,所有驱逐方法的额外延迟均可忽略不计。
OBCache的Decoding复杂度评估
OBCache的Prefill复杂度评估1
OBCache的Prefill复杂度评估2

OBCache的PyTorch伪代码实现:
OBCache分数计算的实现逻辑如下:

def obcache_score(key_states, value_states, A, Z, O, w):
    # 仅针对最近的扰动窗口
    A, Z, O = A[..., -w:, :], Z[..., -w:, :], O[..., -w:, :]
    # 计算基于注意力的累积得分
    S_attn = A.pow(2).sum(-2)
    # 计算Value剪枝得分
    V_2norm = value_states.pow(2).sum(dim=-1)
    S_value = S_attn * V_2norm
    # 计算Key剪枝得分
    O_2norm = O.pow(2).sum(dim=-1)
    VO = torch.einsum('bhqd,bhpd->bhqp', O, value_states)
    VmO_2norm = O_2norm.unsqueeze(-1) + V_2norm.unsqueeze(-2) - 2 * VO
    S_key = ((A * Z).pow(2) * VmO_2norm).sum(dim=-2)
    # 计算联合剪枝得分
    VVmO = V_2norm.unsqueeze(-2) - VO
    S_joint = (2 * A.pow(2) * Z * VVmO).sum(dim=-2) + S_value + S_key
    return S_joint

缓存驱逐的实现逻辑如下:先通过 torch.topk 选择得分最高的 num_hh 个Token的位置索引,然后通过 torch.gather 提取这些Heavy-hitter Token,最后使用 torch.cat 将它们与固定保留的最近Token(num_recent)拼接起来,返回压缩后的KV缓存。