CoMem:大模型理解早于预测,按层切分让128k显存降至18GB

Understanding Is Done Early: A Depth Division of Labor in Large Language Models and Its Use for Unbounded-Context Memory

论文原文 ↗ 论文发布 解读发布 解读:AI前沿分享

大语言模型处理长上下文的能力,长期受制于自注意力的二次方计算开销,以及随着上下文长度线性膨胀的键值缓存(KV Cache)。为了把超长文本塞进有限的显卡显存,学术界和工业界过去几年几乎把所有精力都倾注在 Token 轴 上:要么靠滑动窗口和注意力汇元(Attention Sinks)丢弃旧 Token,要么靠稀疏注意力或检索增强(RAG)挑出少部分 Token,亦或是训练专门的压缩向量。然而,几乎所有这些方案都默认了一个前提——只要某个 Token 被保留下来,它在 Transformer 每一层中计算出的键值对,都必须完整地存放在显存里。

ArXiv URL:https://arxiv.org/abs/2607.28263

最新研究 Understanding Is Done Early: A Depth Division of Labor in Large Language Models and Its Use for Unbounded-Context Memory 提出了一个打破常规的观察:Transformer 的不同层在功能上存在明确的“深度分工”(Depth Division of Labor)。底层和中层网络其实很早就构建出了完整的语义表征,而接近输出的高层网络,则越来越偏向于根据当前 Query 和预测目标对表征进行特化。换言之,语言模型对文本的“理解”早在中间层就已经完成了。

基于这一洞察,作者提出了 CoMem(Comprehension Memory,理解记忆) 架构。该方法放弃了在 Token 轴上硬扛完整深度 KV 的传统思路,转而沿 模型层深(Depth Axis) 对长文本记忆进行重组:在写入阶段,长文本块只向前传播到指定的中间层 $j$,显存仅持久化该层的单条残差张量(Residual Tensor);在读取阶段,外部检索器选出固定数量的候选块,与当前 Query 拼接后,再从第 $j$ 层继续执行上层网络的前向重算。

CoMem 架构与核心概念总览

这一机制打破了在线读取开销与文档总长度绑定的宿命。在单张 96GB NVIDIA H20 显卡、128k 上下文长度的严格对照实验中,CoMem 将预填阶段(Prefill)的显存峰值从完整上下文基线的 89.36 GB 骤降至 18.26 GB,同时取得了 7.83 倍的预填端到端加速。在包含 1,986 道长对话记忆难题的 LoCoMo 基准上,CoMem 取得了 38.27 的表现,显著超越全上下文基准 KV-Direct 的 34.59。这项工作表明,长上下文记忆的组织维度不仅可以横向切分 Token,更可以纵向切分深度。

为什么大模型理解“早于预测”?

要理解 CoMem 的设计动机,必须先看 Transformer 内部的表征演变过程。以往的机制解释性研究与层级探针(Probing)表明,文本的句法结构与核心语义信息通常在模型的中浅层就已经达到可读峰值。随着层数进一步加深,隐藏状态逐渐脱离原始输入的通用语义,转而高度特化于下一个 Token 的分布预测;若上下文伴随特定的提问,高层状态更是深度依赖于 Query 的注意力导向。

这就引出了一对尖锐的矛盾:

如果我们在离线写入文档时,把整个上下文一直计算到最后一层并保存全部 KV,这些高层 KV 实际上是在“未见 Query”的盲目状态下生成的,不仅体积庞大,而且针对当前特定任务的适配度并不高;反过来,如果像传统 RAG 那样只保存原始文本,每次查询又必须从第 0 层开始对检索到的文本进行全量前向计算,浪费了大量的重复底层特征提取计算。

研究人员通过探测实验发现,中间层的每 Token 残差状态 $h_j$ 本身就是极度浓缩且高度可用的语义载体。以 36 层的 Qwen3-8B 为例,对于一个给定的分块,只存储第 $j$ 层的单一残差向量,其占用的字节数仅为存储全模型 36 层完整 bf16 KV Cache 的 $1/18$。

然而,残差特征并不能无限推迟到极深层再进行切分。实验表明,存在一个“零样本可读边界”(Zero-shot Readable Boundary)。当切分深度 $j$ 较浅时,高层网络能够在没有额外微调的情况下,完美恢复下游检索与单针召回任务;但一旦 $j$ 越过某个临界深度(进入高度预测特化的区域),直接截断并拼接入新 Query 就会导致读取准确率发生断崖式下跌。更重要的是,模型参数规模越大,这个可读边界就越向深层推移。这意味着,在浅层与深层之间,存在一个由架构本身决定的、可以通过参数适配进一步拓宽的质量与计算成本折中区间。

CoMem 的三阶段流水线设计

基于“浅层写入语义、深层结合提问重算”的思想,CoMem 构建了一套解耦的流式记忆流水线,划分为写入(Write)、筛选(Select)与读取(Read)三个核心阶段。

CoMem 读写与重计算流水线

写入阶段(Write),长文档被拆解为互不重叠的局部块(Chunk,实验默认采用 512 个 Token)。每个块独立重置其旋转位置编码(RoPE)从 0 开始计数,然后仅通过主干网络的前 $j$ 层(即 $[0 : j]$)。模型提取第 $j$ 层末尾的残差张量 $h_j \in \mathbb{R}^{c \times d}$,并将其存入外部持久化存储,键值为对应文本块的索引标识。因为该阶段完全不计算 $j$ 层之后的网络,写入的长程计算量被固定在浅层,且显存中无需为历史块驻留任何键值对。

筛选阶段(Select),面对用户的输入 Query,系统通过外部轻量检索器(如迭代式 BM25)从历史数据库中挑选出最相关的 $k$ 个候选块。由于检索完全在外部索引中完成,检索候选集的大小与原始文档总长度解耦,保证了后续模型侧的工作集处于严格的有界状态。

读取阶段(Read),系统将检索出的 $k$ 个块对应的残差张量,按照其在原始文档中的物理先后顺序进行排列,并在开头补充预计算好的注意力汇元(Sink),在尾部拼接当前 Query 的第 $j$ 层残差表示。此时,系统为这组拼接后的张量赋予一组全新的、连续的局部 RoPE 位置编码,随后将其输入到上层网络 $[j : L]$ 中,执行具备完整因果注意力的跨块前向计算。

这一设计在模型计算开销上带来了决定性的转变。假定保留 $k=12$ 个块、块大小为 512,再加上 Sink 和提问本身,模型端最终参与在线读取与上层重算的总上下文仅约 6.5k 个 Token。无论外部记忆库中存储了 32k、64k 还是 128k 甚至更长的文本,大模型在线推理时实际加载到 GPU 显存的工作集(Working Set)规模始终恒定。

在自回归解码(Decode)阶段,系统只需要对拼接包的上层 KV 执行一次常规预填,并在此后为新生成的 Token 维护局部的下层与上层缓存,完全无需在每一步生成中反复重放整个数据包。因此,后续生成单个 Token 的延迟与存储的总上下文长度彻底无关。

自蒸馏 LoRA:打通深层截断的保真度瓶颈

如果把切分点设为 $j=0$,CoMem 退化为纯粹的“检索 + 全量重算”系统,在推理质量上与标准模型完全等价,但此时每次查询都要重跑全部 $L$ 层,失去节省计算的意义;若将 $j$ 设为极深层,虽然在线重算的层数极少,但未见提问的离线表征往往无法被上层网络正常解析。

为了尽可能推迟切分深度 $j$ 以压缩在线重算成本,同时避免语义解析能力的衰减,作者引入了极轻量的自蒸馏(Self-Distillation)机制。

这一设计的精妙之处在于它完全不依赖任何下游微调任务数据或检索标注,而是纯粹在普通的无监督长文本语料(PG-19 书籍数据集)上,对冻结的主干网络训练一个 Rank-32 的低秩适配器(LoRA):

优化目标采用截断在教师模型 Top-64 候选 Token 空间上的双向 KL 散度损失:

\[\mathcal{L}_{\mathrm{distill}} = 0.6 \, \mathrm{KL}(p \| q) + 0.4 \, \mathrm{KL}(q \| p)\]

其中 $p$ 为教师分布,$q$ 为学生分布。

通过这种自监督蒸馏,高层的 LoRA 权重迅速学会了如何去“对齐并平滑”那些在第 12 层就已经脱离上下文、缺乏后续全局交互的中间残差状态。消融实验证实,在完全不改动原始预训练模型参数的前提下,仅需这个极轻量的 LoRA,就能将第 12 层切分点的读取保真度拉回到极高水平,为工程落地提供了一个兼顾极致显存优化与高精度的平衡点。

全面基准评测:长文本与对话记忆的真实表现

为了杜绝长文本评测中常见的因 Prompt 微调模板引入的偏差,本研究在 Qwen3-8B 基础模型上执行了统一的纯文本无模板评测协议(Chat-template-free),直接考察模型在原始输入下的泛化与记忆能力。对比对象包括未做长度外推的全上下文基线 KV-Direct、流式滑动窗口 StreamingLLM、基于块检索的 InfLLM 以及参数内置记忆模型 MemoryLLM。

在综合长文本基准 RULER 上,采用 $j=12$ 并搭载自蒸馏 LoRA 的 CoMem 旗舰配置取得了 97.05 的平均分,大幅领先原生未外推全上下文基准的 78.80 分;在 LongEval 的行检索测试中,CoMem 也以 69.0 分超越了全上下文基线的 65.2 分。而在需要广泛聚合多处背景信息的 LongBench 任务上,CoMem 取得了 12.15 的宏观 F1 值,与全上下文基线的 12.17 基本持平。

最能凸显记忆机制差异的是针对长对话记忆的真实评测基准 LoCoMo(包含 1,986 道复杂情景提问)。在该基准下,系统需要跨越长达数万甚至数十万 Token 的多轮历史对话,精准回答涉及过去事件、属性追踪和因果推断的问题。

评测采用 GPT-4o 语义判别器结合本地规则的方式进行打分。结果显示:

  1. 在全部 1,986 道题目中,CoMem 达到了 38.27 的准确率,而全上下文基准 KV-Direct 仅为 34.59。

  2. 在排除了无关对抗样本的 1,540 道客观题上,CoMem 相对全上下文实现了 +4.81 个百分点的净胜优势。通过 10 组对话聚类的 Bootstrap 重采样分析,置信区间稳定在 $[2.34, 7.27]$ 之间,且 10 组聚类中有 8 组明确支持 CoMem 胜出。

  3. 引入独立的 DeepSeek-V3 模型作为第三方裁判,在 200 个分层抽样样本上复核,其与 GPT-4o 裁判的一致性达到 $\kappa=0.626$,且完全保持了 CoMem 优于全上下文基线的排序结论。

这一对话记忆优势的根源,在于全上下文模型在面对极长历史时,不可避免地会受到注意力弥散和无关对话噪声的干扰;而 CoMem 通过有界的局部检索结合跨块因果重算,反而起到了信息过滤与注意力聚焦的作用。

显存与速度的工程账本

除准确率外,工程吞吐与部署显存是检验长文本记忆方案的硬指标。在单张 96GB 显存的 NVIDIA H20 硬件平台上,研究人员对不同上下文长度下的资源消耗进行了端到端实测。

在处理 128k 极端上下文时,全上下文模型的 Prefill 阶段瞬时显存峰值直接飙升至 89.36 GB,几乎触顶显卡物理极限;而无适配器版本的 CoMem 仅占用了 18.26 GB 显存。这意味着原本需要多卡张量并行才能加载的超长上下文,现在单张卡即可轻松承载,甚至还能留出大半空间用于支持大并发批处理。

在预填速度方面,CoMem 将包含所有历史块第 $[0:j]$ 层浅层写入以及最终上层重算在内的全流程耗时打包计算,在 128k 长度下实现了相对完整上下文基线 7.83 倍的端到端预填加速。即便换上挂载了 LoRA 的旗舰版配置,显存开销也仅微幅增加 0.25 GB,预填加速比依然高达 2.74 倍。

在解码阶段,由于上层 KV 缓存已经被固化,模型的单步 Token 生成延迟完全不受 128k 历史上下文拖累,解码吞吐在 64k 范围内与原生稠密模型保持在 11% 差异以内,而在 128k 长度下甚至反超全模型 1.07 倍。这直接印证了论文的核心假设:只要工作集被锁定在固定大小的候选块上,长上下文的计算惩罚就只存在于离线浅层,而不会污染在线推理。

深度消融:什么在真正起作用?

为了避免“黑盒提升”的误导,本文设计了极为详尽的切片消融实验,清晰拆解了检索策略、层深截断与自蒸馏各自扮演的角色。

检索与层深重算的解耦测试。

在控制输入候选块完全一致的前提下,将切分深度强制置为 $j=0$(即在检索出的候选块上跑全量 36 层重算),模型在 LoCoMo 上的得分直接攀升到 41.59。这表明,仅仅依靠“有界块检索”抛弃历史无关文本,就已经比不加过滤的全上下文输入高出了整整 7 个百分点。

当固定冻结主干网络并强行将截断点推深至 $j=6$、$j=9$ 和 $j=12$ 时,模型的准确率逐步下滑至 32.78、29.15 和 24.52。这直观反映出未经适配的深层网络在读取盲目状态时的退化过程。但代价与收益是对应的:在 128k 长度下,上层重算的在线耗时也随之从 1.01 秒线性递减至 0.72 秒。因此,切分深度本质上是一个平滑调节“准确率 vs 在线算力”的系统控制旋钮。

自蒸馏 LoRA 的修复能力。

在相同且较深的 $j=12$ 切分点下,对比未经训练的冻结网络与自蒸馏版本,LoRA 适配器奇迹般地把 LoCoMo 的综合表现从 24.52 暴力拉升了 13.75 分,重回 38.27 的高位;在 BABILong 的长程多跳推理子项 qa1 上,自蒸馏让模型单项提升了 22.2 分。由于 PG-19 预训练语料完全不包含任何下游任务的监督信号,这一增益充分证明:自蒸馏成功填补了层深截断带来的语义表征断层,教会了高层注意力机制如何正确解读离线计算的残差特征。

跨块交互与检索窗口的辩证取舍。

另一个关键发现是,候选块被检索出来后,上层计算必须运行全局因果交叉注意力。如果为了进一步图省事,将候选块之间设为块对角(Block-diagonal)独立重算、禁止跨块交互,模型的检索召回准确率会出现 36 到 62 个百分点的毁灭性暴跌。这表明大模型整合复杂证据链的能力,高度依赖于不同物理切片在高层网络中的相互注意与信息融合。

此外,检索窗口并非越大越好。在 RULER 变量追踪实验中,当把检索预算从 6.5k 盲目扩大到 17k 时,模型在 16k 长度之后的表现反而明显劣于精简窗口。原因在于,过大的检索包引入了大量具有迷惑性的无关变量链条,分散了高层跨块注意力的焦点。

范式转移:长上下文记忆的组织新维度

CoMem 为长期陷入瓶颈的长文本推理架构带来了一种全新的解题视角。过去,优化长文本的思路几乎都在 Token 轴上打转,试图在序列长度维度不断做稀疏、剪枝或池化,却始终不得不让每一个残留的 Token 穿越所有网络层。

CoMem 证明了:大语言模型的深度轴本身就是一个天然的计算与存储分工界面。通过将通用、无偏的语义特征固化在中间层的残差空间,把特定任务导向的、受提问约束的推理计算推迟到极小规模的高层动态网络中,系统成功实现了模型推理内存与上下文存储规模的彻底解耦。

当然,该方案也并非完美无缺。在窗口内短距离强依赖的任务(如 BABILong 的部分简单问答)中,切块与检索依然存在一定的“压缩税”;此外,深层切分的无损性目前仍依赖轻量蒸馏来弥补,面对更加复杂的跨尺度极端推理时,最优层深切分点可能仍需动态调整。但无论如何,这种将大模型浅层作为“离线语义编译器”、高层作为“在线交互推理器”的深度切分范式,无疑为未来构建真正吞吐无限上下文的 Agent 长期记忆系统和高并发推理引擎,指明了一条极具工程可行性的全新路径。