预测式在线剪枝:JKU 把 KV cache 压缩做到了 88% 还几乎不掉点

你有没有这种感觉——长上下文推理时显存哗哗涨,生成到几万 token 就开始 OOM,可是很多"老 token"模型压根没在看?这篇 Sepp Hochreiter 团队(JLU Linz,也就是 LSTM 的那个 Sepp)最新挂出来的 KVpop,给了一个挺漂亮的解法:在驱逐边界直接监督 keep-or-drop 决策,让 scorer 学"未来 attention"信号。

最扎眼的数据是:在 Qwen3-8B 上压掉 88% 的 KV cache,还能拿到 teacher 模型的 99%——平均 Pass@1 从 0.43 几乎不动到 0.43。4B 模型上 75%/88% 压缩分别保留 95%/94%。这个数说实话有点超出我对 learned eviction 的预期。

下面我把这篇论文掰开揉碎讲一遍,重点聊两个我觉得最值得借鉴的设计选择。


核心摘要

KV cache 是长上下文推理的硬瓶颈——内存和带宽随 context 长度线性增长。已有的 KV 驱逐方法,要么靠静态启发式(StreamingLLM、TOVA),要么学个 proxy 分数,但没有一个方法直接监督"未来哪些 token 真正有用"。KVpop 的解法很直接:在驱逐边界(eviction boundary)训练一个轻量 scorer,用 future-attention mass 作为监督目标;为了避免 dense attention 的开销,作者用了一个叫 transposed attention 的小技巧(交换 q 和 k 的角色,用 attention kernel 的 auxiliary LSE 直接拿 per-key 的 future mass)。更进一步,KVpop 引入了基于 mLSTM 的 stateful scorer——不是 token 进 cache 就评分,而是延迟到它即将被驱逐时再评,让 scorer 在保护窗口内积累"近未来证据"再下结论。

效果:Qwen3-4B 在 75% 压缩下保留 98% teacher 性能、88% 压缩下保留 97%;Qwen3-8B 在 88% 压缩下保留 100% teacher 性能(即追平未压缩的全注意力 baseline)。131k 长度推理时峰值显存稳定在 19GB,而 dense attention 已经飙到 36GB。

我自己的判断:这是一篇工程上非常干净的 KV 压缩后训练方案。两个核心选择(boundary-aware supervision + delayed stateful scoring)都站得住脚,不像某些 learned eviction 论文搞了一堆花活最后 baseline 没选好。Sepp 团队把 mLSTM 用进 scorer 算是"主场作战"——这篇文章的真正价值在于它把"未来 attention"这个 target 形式化得足够干净,transposed attention 的实现细节也值得所有做 KV 压缩的人学一下。


论文信息

  • 标题: KVpop -- Key-Value Cache Compression with Predictive Online Pruning
  • 作者: Lukas Hauzenberger, Niklas Schmidinger, Anamaria-Roberta Hartl, David Stap, Thomas Schmied, Sebastian Böck, Günter Klambauer, Sepp Hochreiter
  • 机构: JKU Linz(LIT AI Lab,LSTM 的诞生地)
  • arXiv: 2607.05061
  • 日期: 2026-07-06

问题动机:为什么 KV cache 一定要压缩

Transformer 自回归解码每个 step 都要把历史的 K、V 存下来用于 attention——context 越长,KV cache 越大,推理时既费显存又费带宽。这事儿在 chat 场景没那么明显(context 就几千 token),但一旦上到 code agent、long-doc QA、math reasoning 那种动辄 16k+ 的场景,KV cache 直接成为瓶颈。

社区对 KV 压缩的方案大致分三类:

类别 代表方法 思路
启发式驱逐 StreamingLLM, H2O, TOVA, SnapKV sinks + 滑动窗口 或 按累积 attention 分数驱逐
稀疏检索 Quest, Landmark, Native Sparse Attention, DeepSeek Sparse, TokenButler 保留完整 cache,但每次 query 只读部分 token
Learned 压缩 DMC, DMS 学一个压缩/驱逐策略

KVpop 之前最近的两个 learned 方案是: - DMC:学 layer/head-specific 策略,合并不重要的 token 表示(保留完整 cache 数) - DMS:训练二值驱逐门(可微松弛),滑动窗口内延迟驱逐

KVpop 作者点出了两个关键局限:

  1. 没有方法直接监督"哪些 token 未来有用"——DMS 用 Gumbel-sigmoid 训门,本质是隐式监督;KVpop 直接对未来 attention mass 做监督。
  2. 没有方法延迟"评分决策本身"——DMS 把驱逐操作延后到 token 离开保护窗口,但保留/驱逐的决定在 token 插入时就已经做了。换句话说,token 在保护窗口里积累的那段近未来上下文,DMS 完全没用上。

这两个观察直接催生了 KVpop 的两个核心设计。下面进入方法部分。


方法核心:KVpop 在做什么

Figure 1: KVpop 概览

图 1:KVpop 架构概览。左边蓝色是 Importance Target(Q/K 矩阵),中间是 Wq/Wk/Wv 投影 + Sparse Softmax Attention (Running TopK) + Importance Scorer,右边是产生的 Sparse Attention Pattern(Protected/Top-K/Evicted 三类 token 的空间分布)。每层 attention 的 K/V 一边喂给 scorer 算重要性,一边参与稀疏 attention 计算。

KVpop 的 KV cache 分三类: - Sink tokens:最前面的 s 个 token(默认 s=4) - Protected window:最近的 w 个 token(默认 w=256) - Long-range top-k:剩下的"老 token"里按 scorer 排名取 top-k

每 head 的总预算: $\(B = s + w + k\)$

接下来重点聊两个设计选择。

设计选择 1:在驱逐边界用 future-attention target 监督

这是 KVpop 最核心的贡献——不预测 proxy,而是直接预测"这个 token 未来会被多大 attention mass 覆盖"

设训练序列长度 S,g ∈ {1,...,G} 索引共享 KV head h 的 query heads(GQA)。定义 token t 在 query head (h,g) 上的 mean future-attention mass per group:

\[m^{(h,g)}_t = \frac{1}{N_t} \sum_{d=t+w}^{S-1} p^{(h,g)}_{d \to t}\]

其中 \(p^{(h,g)}_{d \to t}\) 是 dense causal attention probability,\(N_t = \max(1, S - (t+w))\) 是个归一化项。

聚合 per-group log-masses 得到 future-utility target: $\(r^{\text{tgt}}_{h,t} = \text{Agg}_g \left[\log(\epsilon + m^{(h,g)}_t)\right]\)$

默认用 max 聚合("任何一个 query head 强依赖就保留"),也可以换成 mean("奖励广泛有用的 token")。

teacher 的截断位置在 query 位置 q:\(t_{\text{new}} = q - w\)(即当前最新可驱逐的 token),\(t_{\text{bnd}}\) 是 teacher 截断处的 boundary token。teacher 标签就是 +1/-1(保留或驱逐)。

为了让高分 token 不会永久霸占 cache,target 加了一个 temporal decay 项: $\(r^{\text{tgt}}_{h,t}(q) = r^{\text{tgt}}_{h,t} + \left\lfloor \frac{q-t}{n} \right\rfloor \log \gamma_h\)$

\(\gamma_h \in (0,1)\) 是 per-head 衰减因子,n 是衰减步长(默认 n=1)。当 n=1 时退化成 static priority。这一招让老 token 慢慢降权,避免 scorer 卡在"早期重要 token 永远留着"的局部最优。

最终的 boundary-aware retention loss: $\(\mathcal{L}_{\text{score}} = \mathbb{E}_{q,h} \left[\omega_{q,h} \cdot \text{softplus}\left(-y_{q,h} \frac{\hat{r}_{h,t_{\text{new}}}(q) - \hat{r}_{h,t_{\text{bnd}}}(q)}{\tau}\right)\right]\)$

直觉上就是:在驱逐边界做一次"new token vs boundary token"的成对比较,loss 聚焦在"改变 cache 成员的那个单次决策"上,每个采样 query 位置的代价只有 O(1)。\(\omega_{q,h}\) 是 margin weighting(让 teacher 决策不模糊的样本权重更高)+ keep/drop balancing 的组合。

设计选择 2:Transposed attention——避免 dense attention 的开销

问题来了:算 future-attention target \(m^{(h,g)}_t\) 不是要算 dense attention map 吗?那训练时显存不就爆了?

KVpop 的解法很巧妙——转置 q 和 k 的角色

原始公式: $\(\log m^{(h,g)}_t = \log \sum_{d=t+w}^{S-1} \exp\left(\ell^{(h,g)}(d,t) - \text{LSE}^{(h,g)}_d\right) - \log N_t\)$

转置之后: - 原 key 位置 t 变成 query - 原 query 位置 d 变成 key - dot product 数值上恢复原 logits - 减去 per-query LSE normalizer - 施加 block mask \(d \geq t+w\)

Figure 2: Transposed Attention

图 2:Transposed Attention 示意。左边的蓝色 tile 展示 token t 的 future attention mass;转置 q 和 k 后,原本的 per-key column-sum 变成 per-query row-sum,attention kernel 返回的 auxiliary LSE 直接给出第一个项,不需要 materialize dense attention map。

这个转置在 FlashAttention 之类的 kernel 里只是一次调用重排,attention kernel 返回的 auxiliary LSE 直接给出第一个项。第二项的 LSE normalizer 可以用 sparse-query 的 LSE 近似(只对 cache 里的 token 求和,O(B) 复杂度),论文实验验证了和 dense-LSE target 的下游性能匹配。

得到未来 attention target 之后,跑 top-k 用一个 Fenwick tree 做 running topk,O(S log S) 时间、O(S) 空间 per head,可以塞进 FlexAttention 的 sparse mask 里。

设计选择 3:Stateful scorer(mLSTM)+ 延迟评分

这部分是 KVpop 的第二大贡献,也是 Sepp 团队"主场作战"的地方——他们把 mLSTM(matrix LSTM,xLSTM 那套)用进 scorer。

直觉:不要在 token 插入 cache 时就评分,而是等它即将离开保护窗口时再评。这样 scorer 在保护窗口里能积累"近未来上下文"(比如接下来几步的 KV 状态),对"这个 token 是不是真的重要"做出更准确的判断。

每个 KV head 维护一个 mLSTM 记忆。在 query 位置 q: 1. scorer 先用 up to q 的 token 更新 memory 2. 再对刚要离开保护窗口的 token \(t_{\text{new}} = q - w\) 评分

读出用了一个 delayed readout 公式(公式 10-11,论文里就是 mLSTM 的 parallelized readout,跟 xLSTM 那篇论文里的一样)。scorer 的输入是 \(\bm{x}_{h,t} = [\bm{k}_{h,t}; \bm{v}_{h,t}]\),过一个 Hedgehog 特征化(softmax + softmax(-x) 拼接,xLSTM 的标配),再过 mLSTM 记忆 + 最终线性投影出分数。

Figure 3: KVpop stateful eviction policy

图 3:stateful scorer 工作流。右边 mLSTM state 不断用最近的 (k, v) pair 更新(步骤 1),到 token 即将离开 sliding window 时才用 delayed readout 评重要性分(步骤 2),分数进 top-k 排名决定保留/驱逐(步骤 3)。左边是 Sparse KV Cache——按分数高低排好,被驱逐的 token 显示为带删除线的方块。

跟 DMS 的关键区别:DMS 延迟的是"驱逐操作",但保留/驱逐的决策在 token 进入时已经做出;KVpop 延迟的是"评分决策本身",让 scorer 看到近未来再下结论。

论文 Figure 5 的消融验证了这一点:mLSTM scorer + delayed readout 比 immediate scoring(等同于 stateless mLSTM)提高 0.2 个点的 token accuracy。0.2 个点听起来不多,但这是 2000 步训练后的稳定值,作者特意指出它有累积效应。

Stateless vs Stateful 的选择

KVpop 还提供了一个 stateless 版本(KVpop_mlp)——两层 MLP + SiLU 激活,输入是 token 自己的 (k, v),没有 memory。优点是实现简单、没有 recurrent state 带来的额外内存开销。论文里把 MLP 版本和 mLSTM 版本都跑了,结果两个都大幅超过 DMS 等 baseline。


实验结果

主实验:AIME & HMMT 上的 Pass@1

数据集是数学推理领域的 AIME 2024/2025、HMMT 2502/2511——共 4 个 benchmark。模型用 Qwen3-4B-Instruct-2507 和 Qwen3-8B,压缩率 75%(保留 25% KV cache)和 88%(保留 12% KV cache)。

Table 1(Qwen3-4B,Pass@1)

方法 AIME 2024 AIME 2025 HMMT 2502 HMMT 2511 Average
Teacher (full attn) 0.61 0.46 0.30 0.43 0.45
CR=75%
StreamLLM 0.47 0.33 0.24 0.34 0.34
TOVA 0.56 0.38 0.28 0.33 0.33
StreamLLM+ 0.55 0.44 0.30 0.41 0.41
DMS 0.62 0.44 0.28 0.43 0.43
KVpop_mlp 0.62 0.43 0.29 0.44 0.44
KVpop (mLSTM) 0.62 0.44 0.31 0.44 0.44
CR=88%
StreamLLM 0.30 0.23 0.15 0.21 0.21
TOVA 0.36 0.23 0.19 0.26 0.26
StreamLLM+ 0.45 0.33 0.26 0.33 0.33
DMS 0.58 0.41 0.27 0.40 0.40
KVpop_mlp 0.59 0.44 0.27 0.42 0.42
KVpop (mLSTM) 0.61 0.44 0.30 0.44 0.44

注意 88% 压缩下,KVpop 几乎追平 teacher(0.44 vs 0.45),DMS 已经掉到 0.40,TOVA 直接崩到 0.26。

Qwen3-8B,CR=88%

方法 AIME 2024 AIME 2025 HMMT 2502 HMMT 2511 Average
Teacher 0.58 0.49 0.28 0.37 0.43
StreamLLM 0.13 0.11 0.03 0.08 0.08
TOVA 0.08 0.07 0.09 0.08 0.08
StreamLLM+ 0.42 0.30 0.18 0.29 0.29
DMS 0.52 0.38 0.22 0.36 0.36
KVpop_mlp 0.57 0.39 0.31 0.42 0.42
KVpop 0.58 0.44 0.31 0.43 0.43

8B 上 KVpop 8B 88% 压缩的 Average 0.43 追平 teacher 0.43——理论上几乎没有掉点。TOVA 几乎全崩,StreamLLM 崩到 0.08。

域外泛化(Table 2)

训练数据是数学推理(Nemotron-Math v2 high-reasoning 子集),但论文在 GPQA-D(STEM 推理)和 LCB(LiveCodeBench,代码生成)上测了泛化。

Qwen3-4B 域外测试

方法 GPQA-D (75%) LCB (75%) GPQA-D (88%) LCB (88%)
Teacher 0.59 0.35 0.59 0.35
StreamLLM 0.54 0.36 0.49 0.35
TOVA 0.58 0.34 0.54 0.35
DMS 0.55 0.37 0.54 0.35
KVpop_mlp 0.59 0.35 0.57 0.35
KVpop 0.57 0.33 0.56 0.34

域外只比 teacher 少 2-3 个点,比 DMS 强一些。说明 KVpop 学到的"未来重要性"信号在 STEM 和代码任务上也有一定迁移——这是 learned eviction 相对启发式方法的优势。

推理效率(Figure 4)

Qwen3-8B,batch size 1,75% KV 压缩。

Figure 4a: End-to-end latency

图 4a:端到端 latency 对比。横轴是生成长度(16k/32k/64k/128k),纵轴是秒数。Dense attention 在 128k 时飙到 33000+ 秒,KVpop 增长最慢(128k 时大约 2500 秒),DMS 居中(约 6000 秒)。

峰值 VRAM:16k 时 Dense 18GB,131k 时 Dense 36GB(线性增长);DMS 和 KVpop 131k 时都只到 19GB,比 Dense 省了一半还多。长序列下 latency 也比 DMS 低——这部分我猜是 KVpop 的 sparse attention kernel 用了 Fenwick tree 跑 running topk,比 DMS 的 Gumbel-sigmoid 训练+推理的离散化开销更友好。

驱逐模式可解释性(Figure 6, 7, 9)

这部分我觉得是论文最有意思的"副产品"——KVpop 不仅是工具,它学到的策略还能告诉我们什么 token 重要

Figure 6: 驱逐模式热力图

图 6:一条数学推理 trace 的驱逐模式。行=层(0-33),列=token(最后 112 个),颜色=该层保留该 token 的 attention head 数(0-8)。可以看到结构 token("Thus"、"Pattern"、"multiplied"、"work" 等)几乎所有 head 跨所有层都保留(深红),纯数字 token 大多被驱逐。

关键观察: - 纯数字 token(具体数值)更常被驱逐 - 推理结构 token(discourse markers 如 "Thus"、操作词如 "multiplies"、符号 token 如 "=")被多数 head 跨多层保留 - 第一层是个例外——几乎所有 token 都被保留

这跟 human 的数学解题直觉是吻合的:你在解一道应用题时记不住"鸡有 23 只",但肯定记得住"接下来用乘法"——KVpop 学到了类似的优先级。

Figure 7: Top-budget recall vs oracle

图 7:KVpop 的 top-k 决策跟"用未来 attention 当 oracle"的匹配度。Qwen3-4B 36 层,每层是 head-level recall 的 boxplot。Global mean recall = 81%——也就是说 KVpop 选出来的 top-k 跟"真用 future attention 选出来的"重合度 81%,这是相当高的水平。中间层和深层的 recall 更稳,浅层(0-2)和最深层(35)波动稍大。

DMS 的隐藏问题(Figure 8)

论文还顺手"鞭尸"了一下 DMS——画了 Qwen3-4B DMS 在 75% 压缩下 100 条随机序列的中位驱逐比。0 = dense attention,1 = sliding window only。

DMS 展示出高度异质模式:少数 head 接近 dense(驱逐比接近 0,几乎不驱逐),多数 head 塌缩成 sliding-window-only(驱逐比接近 1)。也就是说 DMS 实际上学到的策略是"让少数 head 干活,其他 head 退化成滑动窗口"——这不是 KVpop 想要的"均衡 top-k",而是"少数 head 扛所有责任"。

KVpop 的设计(per-head 独立 budget + boundary-aware loss + keep/drop balancing)显式避免了这个问题。

训练超参数(Table 3)

关键参数
Token budget 2B
训练序列长度 S 16384
Batch size 128
Sliding window w 256
Sink tokens s 4
Top-k budget 2016 (CR=75%) / 4032 (CR=88%)
Decay step n 1
Pairwise loss τ 1.0
Scorer LR cosine, peak 1e-3, 100 steps warmup
Base LR constant 8e-5
KL temperature / Top-k 1 / 256
训练步数 2000
硬件 8 × H100, FSDP

训练时除了 KV cache loss,还对 teacher top-256 logits 加了 KL divergence loss,温度 1,跟 scorer loss 一起联合优化——这跟标准 distillation 一致。


我的判断

先说亮点:

  1. Target 形式化得非常干净。"这个 token 在驱逐之后还会被多大 attention mass 关注"——这个信号比任何 proxy score 都直接,而且用 transposed attention 算出来 O(1) overhead。后续 KV 压缩工作大概率会跟进这种"用未来信号做监督"的设计范式。

  2. Delayed scoring 这件事被严重低估了。DMS 那类方法里大家都默认 token 进 cache 就该立刻决定保不保,但 KVpop 用实验告诉你"等一下,scorer 看到近未来之后能多 0.2 个点"。这种"延迟决策"的思路在 RL 领域其实是老生常谈(offline-to-online、delayed reward shaping),迁移到 KV 压缩挺自然的——之前没人这么做。

  3. mLSTM 的主场优势用得恰到好处。Sepp 团队把 mLSTM 塞进 scorer 不是为了炫技——mLSTM 的 recurrent memory 天然适合"近未来上下文积累",而且 mLSTM 训起来比 Transformer 稳定、推理延迟低。算是一招"用对地方"。

  4. 泛化性靠谱。只在数学推理上训,但在 GPQA-D(STEM)和 LCB(代码)上只掉 2-3 个点。比 DMS 强。说明 future-attention 监督信号本身是任务无关的——重要的不是"数字 token 该不该留",而是"这个 token 未来会不会被 attention 覆盖"。

  5. 代码看起来会开源。Sepp 团队传统是开源的,PyTorch-style pseudocode 也在论文里给了(Algorithm 3),工程上能直接拿来用。

再说问题:

  1. 训练数据只有 2000 步,batch size 128、序列长度 16k——这个训练量对"改变 4B/8B 模型的 attention 模式"来说其实很小。作者靠 KL divergence + boundary-aware loss 的组合,把训练量压到了 2B tokens——这既是亮点(数据效率高)也是隐忧(不知道换更大的模型/更长的 context 训起来是否还这么稳)。

  2. 均匀 per-head cache budget 不一定最优。作者自己在 Limitations 里也承认了——某些 head 可能根本不需要那么多 budget("少数 head 扛所有责任"那种),某些 head 可能要给更多。混合 dense-sparse 的 layer 设计可能更好,但作者没做。

  3. 没有跟"从零训练稀疏 attention"的方法直接对比。KVpop 是个 retrofit(dense attention 的后训练压缩补丁),但 NSA、DeepSeek Sparse Attention 那种从预训练开始就是稀疏的方案,论文只是列在"被排除的 baseline"里。不是说一定要比,但读者会好奇 KVpop 的工程优势能不能掩盖"从零训练"的理论优势。

  4. 88% 压缩的 memory 数字我没完全理解。作者说 stateful variant 因为 recurrent state 额外内存,top-k 预算降低了以匹配 stateless 的内存占用——但论文没给具体数字对比。如果算上 recurrent state,mLSTM 版本的 VRAM 是不是其实跟 KVpop_mlp 差不多?

  5. 未来 attention target 本身有 selection bias。transposed attention 算的是"如果 teacher 用 full attention,这个 token 会拿多少 mass"——但 teacher 自己也不是 ground truth。理论上更合理的 target 应该是"如果只用 cache 里这 k 个 token 算 attention,效果会怎样"——但那要 nested optimization,作者避开了。

对比同期工作:

KVpop 的真正对手是 DMS(同期 learned eviction 里最强的),DMS 的训练是 Gumbel-sigmoid 隐式监督,KVpop 是显式 future-attention 监督。8B 88% 压缩下,KVpop Average 0.43 vs DMS 0.36——7 个点的差距,这是显著差距,不是统计噪声。另一边,SnapKV/PyramidKV 那种 query-aware 的启发式方法在长 context QA 上很强,但在 AIME/HMMT 这种需要"多步推理、跨段引用"的场景里会崩(论文里没列具体数字但 StreamLLM 0.08 已经说明问题)。

工程启发:

如果你也在做长上下文 LLM 推理,KVpop 给了一个很实际的 retrofit 方案——不用重训模型,2B tokens 训个轻量 scorer 就能拿到 95%+ 的 teacher 性能,且支持现有 KV cache 框架(FlexAttention + Fenwick tree 跑 running topk)。mLSTM 的 recurrent state 不大,延迟也低。


收尾

回到开头的那个问题——长上下文推理的瓶颈真的是"显存不够"吗?还是说"我们一直在存一堆没用的 token"?

KVpop 的答案是后者。它的两个核心贡献: - future-attention target:直接监督"未来有用性",不绕道 proxy - delayed stateful scoring:让 scorer 看到近未来再下结论

这两个设计选哪个都不能算颠覆性创新,但组合在一起确实 clean。而且 88% 压缩追平 teacher 这个数据,在我读过的 KV 压缩论文里算是非常硬核的。

如果非要挑刺——它没有解决"从零训练稀疏 attention"的问题。KVpop 假设有个 dense teacher 模型存在,然后用未来 attention 蒸馏。当模型本身就是为稀疏设计的(比如 NSA 那种),KVpop 的 retrofit 还有多少优势?这个问题留给后续工作。

Sepp 团队在 LSTM 发明 30 多年后还能持续输出"硬件级推理优化"的新工作——这事儿本身也挺让人感慨的。


论文链接: https://arxiv.org/abs/2607.05061

觉得有启发的话,欢迎点赞、在看、转发。跟进最新 AI 前沿,关注我。