一个日常类比
你在写一封长邮件。大部分时候,你的注意力只集中在当前这段话——前一句刚写完,下一句顺理成章。但偶尔,你需要翻回上面去确认一个细节:那个数字是多少?那个名字怎么拼?
如果你每写一个字都要把整封邮件从头读一遍,那太慢了。但如果你完全不回头看,写到后面可能就和前面矛盾了。
这就是长上下文大模型面临的困境。
全注意力(Full Attention):每生成一个 token,都要重新看一遍整个上下文。准确,但慢——上下文越长,每步成本是 O(n)。
局部注意力(Local Attention):只看最近 2048 个 token。快,但容易"失忆"——前面提到的重要信息可能被忽略。
On-Demand Attention(ODA):先做局部计算,然后让一个轻量级"召回头"(recall head)判断——"这一步需要回头看吗?"如果需要,再做一次全注意力计算;如果不需要,就用局部结果。
核心问题:模型能预测自己需要什么吗?
这个想法听起来简单,但有一个关键前提:在决定是否调用全注意力之前,模型必须已经"知道"全注意力会不会有帮助。
但这怎么可能?你还没做全注意力计算,怎么知道它的结果会不会比局部计算更好?
论文的洞察是:模型的解码状态本身就包含了这个信息。
具体来说,召回头读取三个信号:
- 当前 token 的 embedding(E(x_t))——这一步在处理什么
- 当前 Local 计算的隐状态(h_t^L)——局部视角看到了什么
- 上一步的隐状态(h_{t-1})——前面累积的上下文
这三个信号经过一个 28.3M 参数的小头部(相比 Qwen3-1.7B 的 1.7B 参数,只占 1.7%),输出一个标量分数 q_t。如果 q_t > 道阈值 θ(默认为 0),就触发全注意力计算。
训练:配对监督
训练数据怎么来?答案是配对计算。
在训练时,对每个位置同时做 Local 和 Full 两种计算,然后比较它们对正确下一个 token 的预测概率。如果 Full 显著好于 Local,召回头就应该学会说"是";如果差不多,就应该说"否"。
具体的目标函数用了 Huber 回归(对大误差不敏感),并加了一个成本惩罚 λ——每次调用 Full 都有代价,所以只在收益超过代价时才调用。
训练数据:196,608 个样本,在 Qwen3-1.7B 上训练。主模型权重完全冻结,只训练召回头。
关键结果
质量恢复
在 Qwen3-1.7B 上:
| 策略 | RULER16K 分数 | Full 调用率 |
|---|---|---|
| Full(全注意力) | 81.94 | 100% |
| ODA | 81.17 | 41.6% |
| Local(局部) | 19.23 | 0% |
ODA 恢复了 Full 和 Local 之间分数差距的 98.8%,但只用了 41.6% 的全注意力调用。
在 LongBench v1 上:
- Full: 37.94
- ODA: 36.82(恢复 92.2% 差距,Full 调用率 70.6%)
- Local: 远低
跨模型泛化
同一个思路在多个模型上都有效:
- Qwen3-8B:91.07 分(Full 92.59,Local 25.43),Full 调用率 47.3%
- Qwen3.5-2B(混合架构):94.08 分(Full 94.35),Full 调用率 44.3%——只限制 6 个全注意力层,其余 18 层 Gated DeltaNet 保持原生计算
- Gemma-4-12B-it:跨模型家族也有效
学到的时机 vs 随机调用
这是最关键的对照实验:不是"调用多少次 Full"重要,而是"什么时候调用"重要。
在 5 个 RULER16K 任务上,ODA 以 41.33% 的调用率得到 88.96 分;随机以 40.99% 的调用率(相同频率)只得到 32.79±0.95 分。Full 基线是 88.63。
也就是说:同样调用 41% 的 Full,学到的时机几乎追平 100% Full,而随机调用和 Local 差不多。
这证明了召回头确实学到了有意义的"时机判断",不是在随机猜测。
计算节省
在 128K 上下文、12.5% Full 调用率的条件下:
- ODA 的 FLOPs 只有 Full 的 24.26%——节省 75.74%,计算比 4.12 倍
- 在 vLLM 实测中,吞吐量从 75.54 提升到 149.52 tokens/s——1.98 倍加速
但注意:在 4K 短上下文下,ODA 反而比 Full 慢 15.8%——因为 Local 的额外开销在短上下文下不划算。ODA 是长上下文专用优化。
更深的洞察
1. "正确性"和"访问收益"是两件事
论文有一个很微妙的发现:在 1342 个 Local 已经预测对 top-1 token 的位置中,Full 在 718 个位置提高了这个 token 的概率,在 437 个位置反而降低了。
也就是说:即使 Local 猜对了 token,Full 仍然可能提供更好的概率分布。反过来,即使 Full 提供了正的 NLL 收益,它也不一定改变 top-1 预测。
这意味着"召回"不只是"纠正错误"——它是优化整个分布,而不仅仅是 top-1 准确率。
2. 召回头学到的不是"不确定性"
一个自然的假设是:召回头只是在检测"Local 不确定"的位置。但论文做了对照:
- 在 675 个 Local 预测错误的位置中,召回头的 benefit-sign AUROC 是 0.643
- Local 的 entropy 的 AUROC 只有 0.486
也就是说,召回头学到的信号和简单的不确定性度量不同。它学到的是"全注意力能带来多少改善",而不是"Local 有多不确定"。
3. 输入消融揭示的决策机制
三个输入信号的消融实验很有意思:
- 只用 H+E(历史 + 当前 token):几乎总是调用 Full(99.73%)——没学到"什么时候不需要"
- 加上 L(当前 Local 结果):Full 调用率降到 52.58%——Local 的计算结果是关键信号
这说明:判断"需不需要回头看"的关键信息,恰恰来自"先快速看一眼"的结果。你必须先做 Local 计算,才能知道这一步是否需要 Full。
这和人类的阅读经验一致:你先快速扫过当前句子,如果觉得和前面有矛盾,才翻回去仔细看。不是每句话都仔细看,也不是完全不看——而是先快看,再决定。
4. 与其他"条件计算"的区别
ODA 和之前的条件计算方法(CoLT5、AHA、L2A)的关键区别:
- 不修改主模型权重——预训练模型完全冻结
- 不修改架构——只加一个外部小头部
- 保留完整 KV cache——Full 调用时能访问全部历史,不是压缩或截断的版本
- 学的是"时机"而非"范围"——不是决定"看多少",而是决定"什么时候看"
局限
- 短上下文不适用:4K 以下 ODA 反而更慢,Local 的额外开销不划算
- 召回头需要训练:虽然小(28.3M),但需要 196K 样本的配对训练数据
- 只测了基础模型:post-training 后的行为没有评估
- 阈值需要调:默认 θ=0,但不同应用场景可能需要不同阈值
启示
"判断-闸门解耦"的又一例
ODA 的架构是"判断-闸门解耦"的典型实例:
- 判断:召回头预测"全注意力是否有帮助"
- 闸门:根据预测决定是否执行全注意力
这和之前讨论过的 DoM(用 hidden state 检测 reward hacking)、NavTrust(用 benchmark 分数预测真实可靠性)是同一条思路:用一个轻量级、专门的判断器,决定是否触发昂贵的计算。
"先快看,再决定"是通用模式
ODA 的"先 Local 再决定 Full"策略,和 speculative decoding(先小模型生成,再大模型验证)是同构的。这种"先低成本试探,再按需深入"的模式,在计算资源紧张的场景下几乎是通用的优化方向。
"模型知道自己需要什么"
这是更深的洞察:模型的中间状态已经包含了"我需要什么"的信息。召回头只是把这个信息提取出来。这和可解释性领域的"线性探针"(linear probe)发现一致——hidden state 里藏着比 CoT 更多、更可靠的信息。
一句话总结:On-Demand Attention 训练了一个 28.3M 参数的召回头,在 Qwen3-1.7B 上以 41.6% 的全注意力调用率恢复了 98.8% 的长上下文性能,在 128K 上下文下实现 1.98 倍解码加速。核心洞察是:先做局部计算,再用 Local 结果判断是否需要全注意力——和人类"先快看,再决定是否回头看"的阅读策略一致。
讨论回复
加载中...正在加载回复...
推荐
智谱 GLM-5 已上线
我正在智谱大模型开放平台 BigModel.cn 上打造 AI 应用,智谱新一代旗舰模型 GLM-5 已上线,在推理、代码、智能体综合能力达到开源模型 SOTA 水平。