WhiteMatter半缓存优于同层,全缓存打平加深50%层

WhiteMatter: All-to-All Cross-Layer Connections via KV Source Mixing

Wenbo Zhang, Xiang Ren

cs.CL, cs.LG

2026-08-19

WhiteMatter把各层历史动态混成共享KV通道。同等训练量下全缓存打平加深50%的模型,半缓存在1.3B上仍优于同层。

这篇在解决什么

标准 Transformer 每解一个 token,都会在每一层给它留下一套表示。但注意力读历史时,第 ℓ 层只能看第 ℓ 层的 key 和 value。这样训起来、prompt 预填充(prefill)起来都好并行,已经算过的别层信息却用不上。

Feedback Transformer 把一个过去 token 的全部层状态压成一份加权和,再投影成所有层共用的 KV。LCKV 让 condensed 层读最深层的 KV,warmup 层仍读本层。两条路都把多深度信息挤进一份层宽的摘要,这份摘要必须同时伺候所有目标层。LCKV 的实验显示,把 feedback KV 套到每一层,困惑度会明显变差。消融也表明,光把 KV 缓存做大,补不回「不同层要看不同深度」这件事。

对 agent 和长推理轨迹,解码经常占推理时间的大头。跨层复用已经算过的表示,同时尽量保住标准结构的可并行性,是这篇要打的点。

方法

WhiteMatter 把每层独立的 KV 投影换成一个跨层 KV 池。对每个过去的 token 位置:

每层只读一个通道,是为了避免从高带宽显存(HBM)再拉另外 k-1 份 KV。

解码用严格因果注意力:当前 token 只看位置更早的缓存,走完层栈后再构造自己的通道,留给后面的 token。如果当前步读自己的通道,浅层和深层会打成环。

训练和 prefill 不能按 token 串行。LCKV 用 Jacobi 迭代,每一轮所有 token 并行,但只能读上一轮的 KV,信息要等下一轮才能往右传,常常要跑很多遍全序列。WhiteMatter 用 cyclic Gauss-Seidel 迭代:按位置 i mod g 切成交错的 g 组,组内并行,组间顺序刷新。后一组在同一轮就能读到前一组刚写好的 KV,邻居之间传得更快。反向只走最后 ng 轮,前面的轮次只用来靠近不动点。实现上把 FlashAttention 式分块改成交错因果掩码。

主实验 g=8,router 步长 p=2,训练是 1 轮无梯度加 2 轮带梯度。

结果

Qwen3 结构的 decoder,从零训 FineWeb-Edu。小规模宽 512、8B token;大约 1.3B 参数的规模宽 1792、10B token。

小规模 held-out 困惑度(Figure 4):

方法测试 PPLKV 相对 16L
Vanilla 16L21.751.0×
LCKV w=721.460.5×
FusedKV21.590.5×
WhiteMatter k=820.470.5×
WhiteMatter k=1620.081.0×
Vanilla 24L20.181.5×

全缓存相对同层 vanilla 低 7.7%,和多 50% 层的 24L 基本打平。半缓存相对同层低 5.9%。零样本 11 项均分,k=16 是 49.89,高于 16L 的 47.21 和 24L 的 48.52。LAMBADA 困惑度从 127.47 降到 60.73。

1.3B 上半缓存(k=14,1.326B 参数)held-out PPL 12.93,对照 28 层 vanilla(1.351B)的 13.51,相对低 4.3%。零样本均分 59.24 对 57.48。这个规模没有全缓存,也没有 LCKV 或 FusedKV 对照。

4 层、用精确自回归训出来的参考模型上,cyclic g=16 用 4 轮、每序列 7.32 ms 进入参考 PPL 的 1% 误差带;Jacobi 要 53 轮、91.20 ms,快 12.5 倍。1.3B、batch 64、A6000 上四套结构解码都在约 2500 token/s。WhiteMatter 解码峰值显存 10.05 GiB,比 vanilla 的 16.59 GiB 低 39.4%;prefill 峰值 13.44 GiB 对 21.21 GiB。prefill 吞吐只有 vanilla 的 31%,但比 LCKV 快 1.78 倍,比 Feedback Transformer 快 2.92 倍。

小规模训练 FLOPs 是 vanilla 的 2.32 到 2.50 倍,三轮 prefill 是 3.05 到 3.30 倍,解码几乎持平,0.99 到 1.03 倍。

在 3.28 亿 token 的消融里,共享一份混合再配 16 套独立 KV 投影,质量仍差过 k=4 的 WhiteMatter,尽管缓存大 4 倍。静态权重相对动态 router,k=1 时 PPL 高 3.0%,k=16 时高 1.9%。禁掉深到浅的反馈后,全缓存 PPL 比完整 WhiteMatter 高 4.1%。训练日程离不动点越远,继续迭代时质量越容易掉;最强日程相对最弱日程 PPL 低 32%。

为什么重要

对解码占时间、KV 占显存的场景,半缓存换到更好的质量,这条交换是实的。

解码计算量和吞吐跟标准模型几乎一样,显存少约四成。

这不是免费午餐。训练和 prefill 更贵,还要一份适配交错因果掩码的 attention kernel。证据停在从零训到 1.3B、最多 10B token,没有在现成大模型上继续训的结果。如果线上瓶颈是 prefill,比如短生成或 RAG,31% 的吞吐可能直接否决。

架构上,这篇把跨层复用从「压成一份共享摘要」推进到「按目标层、按内容选源」。消融站在这件事这边:缓存容量本身不够,专属混合和动态路由都有贡献。

局限与存疑

论文自己写了三条:迭代训练和 prefill 仍比标准 Transformer 贵;主结果只到 1.3B、10B token;全缓存和其它 KV 共享基线只在小规模比过。

另外几处站得不稳。1.3B 只测了半缓存,不知道全缓存在这个规模还能不能打平「加深 50%」。评测时 WhiteMatter 用 3 轮 cyclic,LCKV 用 9 轮 Jacobi,vanilla 一轮,质量数字里嵌了不同的计算预算。零样本里 OpenBookQA 几乎没动,1.3B 上 PIQA 从 70.46 降到 69.53,均分很大程度靠 LAMBADA、SciQ、SQuAD 拉起来。脑白质类比只在附录,对机制没有约束力。

术语

原文与代码

相关论文

全部论文解读