一篇论文把「Transformer 每一层都存 KV」这件事给否了,省了 5 倍显存
arXiv 昨晚挂了一篇有意思的工作,标题起得挺学术——《Understanding Is Done Early: A Depth Division of Labor in Large Language Models and Its Use for Unbounded-Context Memory》。但核心发现一句话就能概括:Transformer 的深度层不是均匀工作的,前半截在下中层做语义理解,后半截的上层专门做预测。也就是说,模型读一段话的时候,前面十几层已经基本"看懂"了,后面的层更像是在把理解转化成输出。
这话听起来有点反直觉。平时跑长上下文,默认做法就是给每一层都保留 KV Cache,恨不得把从第一层到最后一层的信息都原样存下来。显存就这样被撑爆了。
作者顺着这个观察做了个方案,叫 CoMem(Comprehension Memory)。做法很直接:每个上下文块只写到中间层就停,不继续往上算了;缓存固定数量的残差状态;等真正有查询来了,再根据查询重算上层。对于固定的检索预算,模型侧的读取算力和内存占用跟"存了多少上下文"解耦了——从 32k 加到 128k,CoMem 这边的推理开销不会跟着线性涨。
数字上给得挺硬。RULER 跑到 97.05,LoCoMo 上 38.27,而老老实实给全量 KV 做缓存的方案(Full-Context KV-Direct)只有 34.59。更落地的是效率账:在 NVIDIA H20 上跑 128k 长度,CoMem 只吃 18.26 GB 显存,全基线要 89.36 GB——差了接近 5 倍。Prefill 加速比 7.83x。原来只能塞一条 128k 对话的卡,现在能塞五条,首 token 出得更快。
实验基于 Qwen3-8B 做 backbone,主干冻住不动,只 train 了一个 rank-32 的 self-distillation LoRA,数据用的是 PG19。还单独报了一个不用 adapter 的控制臂。没有花里胡哨的工程,但思路的穿透力很强。
论文也交代了局限:bounded retrieval 带来了 in-window compression tax——你压缩了层维度,窗口内仍然有压缩代价。另外 depth sweep 显示,层数缓存得越深,重计算量越低,但 fidelity 有损失,自蒸馏可以大幅修复但没完全消除。
看完有一个没解的问题让我睡不着:如果模型的理解和预测在深度上是可解耦的,那现在的推理架构——KV Cache、Prefill-Decode 分离、MoE 路由——是不是都默认了"每一层同等重要"这个前提?谁先把这个前提撤掉,谁可能就能省下一大笔钱。