OjaKV: Context-Aware Online Low-Rank KV Cache Compression

发表时间: 2026-07 · ACL 2026 Findings

原文: https://aclanthology.org/2026.findings-acl.494

Yuxuan Zhu, David H. Yang, Mohammad Mohammadi Amiri, Keerthiram Murugesan, Tejaswini Pedapati, Pin-Yu Chen
(Rensselaer Polytechnic Institute, IBM Research)

速读

一句话结论本文提出了一种名为 OjaKV 的上下文感知在线低秩 KV 缓存压缩框架,通过混合存储策略保留关键 token 的全秩表示,并利用 Oja 算法在线动态更新低秩投影子空间,在大幅降低长上下文推理显存占用的同时,显著提升了模型在长文本生成与复杂推理任务中的准确率。

要解决什么问题大型语言模型在处理长上下文时,自回归生成所需的 KV 缓存会引发严重的显存瓶颈。例如,在批处理大小为 4 的情况下,Llama-3.1-8B 模型处理 32K tokens 的提示词需要消耗约 16 GB 的 KV 缓存显存,这甚至超过了模型权重本身的体积。为了缓解这一问题,低秩近似技术被广泛应用,它通过将键和值向量投影到低维子空间来压缩缓存。然而,现有的低秩压缩方法严重依赖于从特定校准数据集中离线学习到的静态正交投影基。这种静态基隐含了一个假设,即推理时的提示词分布与校准数据一致。但在实际应用中,模型往往会面临数据分布偏移(例如从日常对话切换到代码生成,或者从简短提示词过渡到极长的推理思维链),这会导致静态投影基的近似误差不断累积,使得注意力机制的计算严重失真,最终导致生成质量出现断崖式下降。

怎么做的OjaKV 摒弃了全局统一且静态的压缩范式,设计了一套结合选择性全秩存储与在线子空间自适应的混合框架。首先,该方法认识到并非所有 token 对压缩误差的容忍度都相同。OjaKV 采用了一种基于误差的混合存储策略,它会计算每个键向量在当前低秩基下的重建残差,并结合近期查询的注意力权重得出误差分数。对于那些重建误差最高、难以被低秩空间拟合的关键 token,系统会保留其全秩表示,以此作为注意力机制的高保真锚点,从而在理论上限制了最坏情况下的注意力扰动。其误差分数的计算公式为:

$$e_t = \frac{1}{|\mathcal{W}_t|} \sum_{h, q \in \mathcal{W}_t} \left| \frac{\pmb{q}_{h, q}^\top \pmb{r}_t}{\sqrt{d_h}} \right|$$

其中 $\pmb{r}_t$ 为全秩键与低秩重建键之间的残差。其次,对于绝大多数被压缩的中间 token,OjaKV 引入了 Oja 算法(一种在线主成分分析算法)来动态更新投影基 $\pmb{U}$,使其始终与不断演变的上下文分布保持对齐。更新公式为:

$$\pmb{U} \gets \pmb{U} + \eta (\pmb{x} \pmb{x}^\top \pmb{U} - \pmb{U} \pmb{U}^\top \pmb{x} \pmb{x}^\top \pmb{U})$$
该公式的第一项将基拉向方差最大的方向,第二项则隐式维持正交性。这一自适应过程分为两阶段:在预填充阶段,利用池化后的提示词特征对投影基进行一次全面的批量更新;在解码阶段,则将新生成的键值对存入缓冲区,每隔一定步数执行一次轻量级的周期性更新。最后,为了与 FlashAttention 等现代优化算子兼容,OjaKV 在显存中仅存储压缩后的低秩缓存,而在传入注意力算子前,会在计算图内实时将其重建为全维张量,从而在不修改底层内核的前提下实现了无缝集成。

效果如何实验在 NVIDIA H100 GPU 上进行,评估了 Llama-2-7B、Llama-3.1-8B 以及 LongChat-7b-v1.5-32k 等模型,对比基线包括全秩 KV 缓存(性能上限)、Eigen-N(原生低秩实现,不兼容 FlashAttention)、StaticPCA(带实时重建的静态基方法)以及 Palu(权重分解方法)。在极具挑战性的 RULER 16K 长上下文检索任务中,Eigen-N 因无法使用 FlashAttention 直接爆显存,Palu 准确率跌至 0%,StaticPCA 在 0.6 倍压缩率下准确率仅为 23.44%,而 OjaKV 的实用变体 OjaKV-PF 则达到了 50.7%。在 AIME 2025 复杂数学推理任务中(提示词极短但生成的思维链极长,分布偏移剧烈),所有静态压缩基线全部失效(准确率为 0%),而 OjaKV 依然维持了 13.0% 的准确率(全秩基线为 43.3%),证明了在线自适应的必要性。在效率方面,对于 32K tokens 的输入,OjaKV 将显存占用从 16 GB 降至 11.6 GB,代价是首字延迟从 2102 毫秒增加到 2801 毫秒,解码延迟增加了 10.9% 到 13.4%。此外,该方法还能与 SnapKV 等 token 筛选技术正交结合,实现乘数级的显存节省。作者也承认了方法的局限性:引入了学习率、更新间隔等固定超参数,可能需要针对不同任务进行微调,且基于重建误差的启发式选择策略未必能完美捕捉所有下游任务的 token 重要性。

主要贡献

大语言模型(LLM)在长上下文推理过程中,GPU内存成为了关键的瓶颈。这种内存占用主要来源于模型权重以及自回归解码过程中所需的键值(KV)缓存。例如,Llama-3.1-8B模型在处理32K token的提示词(批处理大小为4)时,需要消耗约16 GB的KV缓存,这甚至超过了模型本身的权重大小。

为了缓解这一挑战,现有的方法主要包括量化、Token选择、卸载以及低秩近似。本文聚焦于低秩近似方向,该方向通过将每个键和值向量从高维投影到低维子空间来减少内存。然而,现有的低秩压缩方法依赖于从校准数据集中学习到的静态、离线的子空间。当面临数据分布偏移(例如从对话切换到代码生成)时,这些静态方法的近似效果会显著恶化,从而损害生成质量。

为克服现有静态低秩方法的局限性,本文提出了OjaKV,这是一种新颖的KV缓存压缩框架。该框架的核心创新点如下:
1. 混合存储策略:OjaKV认识到并非所有Token对压缩的容忍度都相同。它策略性地将具有高重构误差的关键Token保留为全秩状态,从而为注意力机制维持高保真度的锚点。
2. 在线子空间自适应:对于大部分中间Token,OjaKV利用Oja的增量主成分分析(PCA)算法对投影基进行在线自适应更新。该机制在预填充(Prefill)阶段进行全面的更新,而在解码(Decoding)阶段执行轻量级的周期性更新,确保低秩子空间始终与不断演变的上下文保持对齐。
3. 硬件与架构兼容性:该框架与FlashAttention等现代注意力模块完全兼容,确保了在真实长上下文推理中的实用性。
4. 正交扩展性:该方法与Token选择(剪枝)方法兼容,二者结合能够实现复合的内存节省效果。

背景知识与设计原则

低秩注意力机制基础。标准注意力机制将输入序列投影为查询($Q$)、键($K$)和值($V$)。低秩近似的核心思想是将$K$和$V$投影到低维子空间中。具体而言,定义两个正交基:用于键和查询的$U_k \in \mathbb{R}^{d_h \times r_k}$,以及用于值的$U_v \in \mathbb{R}^{d_h \times r_v}$,其中$r_k, r_v \ll d_h$。不再缓存全秩的$K$和$V$,而是存储其压缩表示:$\tilde{K} = K U_k$ 以及 $\tilde{V} = V U_v$。这种方式将每个Token的KV缓存存储需求从$2d_h$降低到了$r_k + r_v$。为了与FlashAttention等优化内核保持兼容,系统会存储压缩后的KV缓存,并在传入FlashAttention之前,通过$\hat{K} = \tilde{K} U_k^\top$和$\hat{V} = \tilde{V} U_v^\top$实时重构出原始空间的张量。

动态分布偏移的挑战与应对。离线计算的低秩基通常拟合于校准数据的分布,在推理时遇到领域或任务偏移时,会产生错位,从而增加键和值的投影误差。为了验证这一假设,作者进行了一项实验:从通用语料库(WikiText-2)计算初始基$U_{\text{cal}}$,并在不同领域(MultiNews)的长上下文摘要任务上进行评估。实验采用残差能量比(RER)来衡量投影误差,并使用子空间重叠度(SO)来量化与测试集Oracle基$U_{\text{test}}$的对齐程度。实验结果表明,静态校准基在分布偏移下泛化能力差,其RER在MultiNews任务上增加到了0.255。而通过Oja规则在短前缀上进行在线更新得到的自适应基$U_{\text{adapt}}$,成功将RER降低至0.097,并将SO从0.597提升至0.653。这一发现证实了在线更新能够有效抵消分布偏移的影响。

方法细节

OjaKV框架概述。OjaKV引入了一种用于内存高效推理的混合策略。该策略结合了选择性全秩保留(旨在为关键Token保留高保真表示)与针对剩余序列的在线自适应低秩压缩。大部分键和值向量通过学习到的投影矩阵投影到一个紧凑的子空间中,这些投影矩阵在推理过程中不断自适应,以保持与演变中的上下文对齐。该框架围绕三个核心组件构建:混合KV缓存存储策略、轻量级初始化过程【1,Wiki-40b: Multilingual language model dataset+2020+LREC】、以及使用Oja算法的两阶段在线更新方案。

图1:OjaKV工作流程概述。左上方面板展示了使用全秩KV缓存的标准注意力机制。我们的方法(底部面板)引入了一条低秩路径,其中键和值在缓存前使用投影矩阵$(U_k, U_v)$进行压缩。右上方的插图说明了核心机制:这些投影矩阵在预填充和解码阶段动态更新,以适应上下文。
图1:OjaKV工作流程概述。左上方面板展示了使用全秩KV缓存的标准注意力机制。我们的方法(底部面板)引入了一条低秩路径,其中键和值在缓存前使用投影矩阵$(U_k, U_v)$进行压缩。右上方的插图说明了核心机制:这些投影矩阵在预填充和解码阶段动态更新,以适应上下文。

基于重构误差的Token选择机制。在长上下文生成中,不同Token对下游预测的贡献不同,且受低秩近似误差的影响也不同。混合存储策略动态识别哪些Token应保留全秩表示。首先计算重构残差:对于每个键Token,计算全秩键与其低秩重构之间的残差 $r_t = k_t - U_k U_k^\top k_t$。这个残差直接控制了注意力分数的扰动,通过柯西-施瓦茨不等式,压缩引入的点积误差$\Delta$被界定为 $|\Delta| \leq \|q\| \|r_t\| / \sqrt{d}$。接着评估下游影响并筛选:为了评估对下游的影响,利用最近查询窗口内的注意力对残差进行加权。给定滑动窗口内的查询$Q_{\text{win}}$,计算查询加权误差分数 $e_t = \frac{1}{|\mathcal{W}_t|} \sum_{h, q \in \mathcal{W}_t} \left| \frac{q_{h,q}^\top r_t}{\sqrt{d_h}} \right|$。在可选的池化操作后,系统保留具有最高误差分数的Top-$k$个Token为全秩,同时压缩所有其他Token。最后建立误差到生成质量的理论联系:这种基于残差的选择规则提供了端到端的误差解释。如果压缩Token的最大残差范数为$E$,通过标准的Softmax稳定性分析,注意力输出的误差被界定为 $\|\hat{o} - o\|_2 \leq O(V Q E / \sqrt{d_h})$。如果后续网络是$L$-Lipschitz连续的,最终Logit扰动受到 $\|\Delta \ell\|_\infty \leq L \|\hat{o} - o\|_2$ 的控制。因此,降低残差范数能直接改善生成保真度。

利用Oja算法进行两阶段在线更新。静态基仅在提示词上训练,随着上下文的增长会累积近似误差。为了解决这个问题,使用Oja规则【2,Simplified neuron model as a principal component analyzer+1982+Journal of Mathematical Biology】在线更新投影基。给定新样本$x$,Oja规则将基$U$更新为:$U \gets U + \eta (x x^\top U - U U^\top x x^\top U)$。这相当于在隐式正交性约束下对方差进行随机梯度上升。首先是预填充(Prefill)阶段的更新:在预填充期间,使用提示词中的全量键和值向量初始化基。为了减少相邻Token的冗余并降低计算成本,在更新前应用局部平均池化:$X_{\text{pooled}} = \text{AvgPool}(X, b)$,然后执行更新 $U \gets U + \eta (C_{\text{pooled}} U - U U^\top C_{\text{pooled}} U)$,其中$C_{\text{pooled}}$是经验协方差。更新后,通过QR分解重新正交化以确保数值稳定性。随后是解码(Decoding)阶段的更新:在解码期间,新生成的Token会产生可能位于当前子空间之外的键值对。将这些向量累积在缓冲区$B$中。每隔$T$步,对缓冲的特征应用Oja规则:$U \gets U + \eta (C_B U - U U^\top C_B U)$,并使用保守的学习率以平衡自适应与稳定性。更新后,重新正交化基并清空缓冲区。此外引入了实用变体OjaKV-PF:在实际应用中,预填充阶段可以使用原始的全秩键和值计算注意力,同时仍使用提示词特征更新低秩基并填充压缩的KV缓存。解码阶段则继续使用重构的键和值,这保留了解码时的内存节省,同时减少了预填充时不必要的近似开销。

实验环境

  • 硬件配置:单张 NVIDIA H100 NVL GPU。
  • 软件配置:PyTorch 2.6.0, Transformers 4.44.0, FlashAttention 2.7.4.post1,采用float16精度运行。
  • 评估模型:Llama-2-7B, Llama-3.1-8B, LongChat-7b-v1.5-32k, Deepseek-R1-Distill-Llama3-8b。
  • 数据集与基准测试:WikiText-2(用于校准)、RULER(16K长上下文检索测试)、LongBench(多任务长上下文理解测试)、AIME 2025(数学与推理生成任务)、lm-eval-harness(短上下文多项选择任务)。

实验结果

  1. RULER基准测试(长上下文检索):在16K输入序列的极端内存压力下,Eigen-N基线由于不兼容FlashAttention直接导致OOM(内存溢出);Palu基线完全失败,所有任务准确率为0%;StaticPCA基线在0.8x和0.6x压缩下平均准确率骤降至28.37%和23.44%。相比之下,OjaKV-PF在0.8x和0.6x压缩下分别保持了54.19%和50.7%的平均准确率,验证了动态上下文自适应框架的有效性。(引自原表3)
  2. LongBench基准测试(长上下文理解):在文档问答、少样本学习和代码生成等任务中,OjaKV在两种模型和不同压缩率下均优于StaticPCA。虽然优势不如在RULER中显著(因为LongBench测试的是对稳定上下文的理解),但OjaKV依然提供了最稳健的性能。(引自原表2)
  3. AIME 2025(长推理生成任务):推理任务具有“极短提示词+极长生成过程”的特点。在Deepseek-R1-Distill模型上,静态压缩方法(Eigen-N, StaticPCA, Palu)完全失效,准确率均为0%。而OjaKV通过在线Oja更新持续自适应解码过程,是唯一保持有效推理能力的压缩方法,达到了13.0%的准确率。(引自原表4)
  4. 效率分析:在Llama-3.1-8B处理32K Token时,OjaKV(60%压缩)将内存消耗从16 GB显著降低至11.6 GB。虽然在线更新引入了适度的延迟(首字延迟TTFT从2102 ms增加到2801 ms,解码延迟增加约10.9-13.4%),但这使得在相同硬件预算下处理更长输入成为可能。(引自原图2及表5)
  5. 消融实验:在0.6x压缩率下的RULER测试中,相比于StaticPCA(38.9%),仅加入“混合存储”将准确率提升至57%,这证实了全秩保留高误差Token的关键作用。进一步加入“在线更新”(即完整的OjaKV)将准确率提升至60%,证明了自适应子空间跟踪的补充优势。(引自原表6)

结论

本文提出了OjaKV框架,有效解决了长上下文LLM中KV缓存的内存瓶颈问题。该框架创新性地结合了保留关键Token全秩状态的混合存储策略,以及基于Oja规则的轻量级在线更新方案。实验证明,OjaKV在极具挑战性的生成类长上下文任务中表现优异,有效缓解了分布偏移问题,并且与FlashAttention完全兼容。未来的工作将探索把固定的更新频率替换为基于重构误差触发的动态更新策略,以进一步优化自适应质量与运行时开销之间的平衡。

附录

A.1 Oja的更新算法
OjaKV在线更新的完整过程被整合为以下算法逻辑。预填充阶段对提示词进行自适应,解码阶段处理新生成的Token。

# 算法1:OjaKV 伪代码逻辑
# 输入: 低秩投影矩阵 Uk, Uv; 学习率 eta; 更新缓冲区大小 T; 池化大小 p; 选择的top-k
# 预填充阶段 (Prefilling Phase):
# 1. 形成矩阵 K, V; 应用大小为 p 的平均池化
# 2. K_tilde = Uk.T @ K; Uk = Uk + eta * (K - Uk @ K_tilde) @ K_tilde.T
# 3. V_tilde = Uv.T @ V; Uv = Uv + eta * (V - Uv @ V_tilde) @ V_tilde.T
# 4. (Uk, Uv) = Orthonormalise(Uk, Uv)
# 5. K_hat = Uk @ Uk.T @ K; E = abs(Q_window.T @ (K - K_hat)) # 误差加权选择
# 6. T_full = Top-k(Normalize(E))

# 解码阶段 (Decoding Phase):
# for step t = 1, 2, ... :
#     1. 生成新的 (k_t, v_t) 并追加到缓冲区 B_k, B_v
#     2. if t % T == 0:
#            K_tilde = Uk.T @ K; Uk = Uk + eta * (K - Uk @ K_tilde) @ K_tilde.T
#            V_tilde = Uv.T @ V; Uv = Uv + eta * (V - Uv @ V_tilde) @ V_tilde.T
#            (Uk, Uv) = Ortho(Uk, Uv)
#            重置缓冲区 B_k, B_v

A.2 从子空间误差到注意力与生成误差
首先定义基本残差:设$k_t \in \mathbb{R}^{d_h}$为全秩键,$\hat{k}_t = k_t + e_t$为其压缩重构版本,$e_t$为重构残差。对于查询$q$,其注意力Logit的扰动为 $|\hat{z}_t - z_t| = \frac{|q^\top e_t|}{\sqrt{d_h}} \leq \frac{Q \|e_t\|_2}{\sqrt{d_h}}$。接着推导注意力权重的变化:定义最大残差$E = \max_t \|e_t\|_2$,通过Softmax稳定性可得 $\|\hat{\alpha} - \alpha\|_1 \leq 2 \|\hat{z} - z\|_\infty \leq \frac{2 Q E}{\sqrt{d_h}}$。然后推导注意力输出的误差:假设$\|v_t\|_2 \leq V$,则输出误差界定为 $\|\hat{o} - o\|_2 \leq V \|\hat{\alpha} - \alpha\|_1 \leq \frac{2 V Q E}{\sqrt{d_h}}$。最后关联至最终预测:若后续网络是$L$-Lipschitz的,最终Logit向量的扰动为 $\|\hat{\ell} - \ell\|_\infty \leq \frac{2 L V Q E}{\sqrt{d_h}}$。这表明累积退化与解码步骤中最大残差$E_i$的总和成正比。混合存储策略通过全秩保留高误差Token直接限制了单步的$E_i$,而Oja在线更新则控制了残差随时间的增长。

A.3 低秩子空间初始化
对于注意力头$i$,从$n_s$个采样序列中收集查询、键和值的激活矩阵$R_i^Q, R_i^K, R_i^V$。首先构建查询-键基:为了鼓励共享表示,拼接查询和键矩阵 $R_i^{KQ} = [R_i^Q, R_i^K]$。应用紧凑SVD分解 $R_i^{KQ} = U \Sigma V^\top$。根据能量标准 $\frac{\|(R_i^{KQ})_r\|_F^2}{\|R_i^{KQ}\|_F^2} \geq \epsilon_{\text{th}}$ 选择最小秩$r$,$U$的前$r$列定义了基$U_k$。接着构建值基:对$R_i^V$应用相同的SVD过程获得$U_v$。最后,将有效秩设置为该层中所有注意力头观察到的最大$r$值。

A.4 与全秩FlashAttention的等价性及成本比较
首先证明Logit等价性:在低秩内核中,点积计算为$\tilde{Q} \tilde{K}^\top$。在FlashAttention兼容模式中,重构$\hat{K} = \tilde{K} U_k^\top$。由于 $\hat{K}^\top = U_k \tilde{K}^\top$,因此 $Q \hat{K}^\top = Q (U_k \tilde{K}^\top) = (Q U_k) \tilde{K}^\top = \tilde{Q} \tilde{K}^\top$。接着证明输出等价性:利用结合律,$\text{softmax}(\frac{\tilde{Q} \hat{K}^\top}{\sqrt{d_h}}) \hat{V} = \text{softmax}(\frac{\tilde{Q} \hat{K}^\top}{\sqrt{d_h}}) (\tilde{V} U_v^\top) = (\text{softmax}(\frac{\tilde{Q} \hat{K}^\top}{\sqrt{d_h}}) \tilde{V}) U_v^\top$。这证明了重构后计算与在降维空间计算是数值等价的。最后进行复杂度分析:FlashAttention兼容路径保留了全秩内核的时间复杂度$O(m n d_h)$,但通过仅存储$\tilde{K}, \tilde{V}$保留了内存优势。内存节省比例为 $1 - \frac{r_k + r_v}{2 d_h}$。

A.5 & A.8 详细实验设置与默认超参数
默认超参数设置为:Oja更新学习率 $\eta = 0.10$;解码更新周期 $T = 32$(步);重要性窗口大小 $W = 32$(查询数)。

A.6 Lm-eval-harness
在短上下文零样本基准测试(如PiQA, WinoGrande等)中,实验观察到Eigen-N和StaticPCA基线产生完全相同的结果,这从经验上验证了A.4节中关于重构路径与原生低秩内核数值等价的分析。此外,在短上下文中,OjaKV和带有注意力槽的StaticPCA-H都表现出接近全秩基线的准确率,表明短上下文对压缩更具鲁棒性。

A.7 定性分析与案例研究
在MultiNews的长文档摘要任务中,输入文档前半部分重点讨论洛杉矶的抗议活动,后半部分提及费城的平静情况。静态基线(StaticPCA)的表现:由于未能适应语义偏移,其生成的摘要完全局限于费城,丢失了洛杉矶的核心事件,导致事实不完整。OjaKV的表现:成功捕捉到了两个地点的关键事件,合成了连贯且全面的概述。这证明了在线子空间自适应能够随着新信息的引入动态更新主成分,避免了灾难性的信息丢失。

A.9 与序列长度压缩的兼容性
OjaKV压缩的是特征维度($d \to r$),这与压缩序列长度($n \to m$)的Token淘汰或选择技术是正交且兼容的。理论分析:对于Token选择矩阵$S$,低秩投影完美结合 $U_k^\top (K S) = (U_k^\top K) S$。两者结合的总压缩比为 $\text{CR} = (d / r) \times (n / m)$。实验验证:将OjaKV(0.6x秩压缩)与SnapKV(50% Token保留率)结合使用。实验结果显示,复合方法将KV缓存内存使用量降低至30%,同时在LongBench上保持了43.33%的准确率,证明了特征维度压缩与序列长度压缩的互补性。

补充细节

局限性。本方法存在几个局限性。首先,OjaKV引入了固定的超参数(学习率、缓冲区大小、Top-$k$),不同模型或任务可能需要重新调参。其次,尽管Oja更新是轻量级的,但与静态压缩方法相比,仍会产生额外的计算开销。最后,基于重构误差的混合选择策略可能无法完美捕捉所有下游任务中Token的真实重要性。


方法细节中的参考文献汇总

  • 【1】Wiki-40b: Multilingual language model dataset (2020, LREC)。引用段落:方法细节 - OjaKV框架概述。原文描述用于说明轻量级初始化过程所依赖的小型通用校准语料库。
  • 【2】Simplified neuron model as a principal component analyzer (1982, Journal of Mathematical Biology)。引用段落:方法细节 - 利用Oja算法进行两阶段在线更新。原文描述用于引入Oja规则,即一种能够递增追踪非平稳数据主子空间的流式算法。