返回博客

MLA 如何省下 KV 缓存:矩阵吸收与不变的 softmax 温度

从同一 MLA 的展开与吸收路径推导缓存压缩,解释 RoPE 的障碍、192 与 576 维的缩放陷阱,以及显存节省为何不能直接换成速度。附可复现代数校验。

长上下文推理里,一个经常被混在一起的问题是:少存一些历史信息,是否就能少做一些运算?MLA(Multi-head Latent Attention,多头潜变量注意力)给出的答案很有意思:它可以让历史缓存显著缩小,同时把某些注意力点积移到更宽的空间里。显存和计算量并不沿着同一个方向变化。

本文重读 2024 年 5 月 7 日首发的 DeepSeek-V2 报告(采用 6 月 19 日 v5)与 DeepSeek-V3 报告(2024 年 12 月 27 日首发,采用 2025 年 2 月 18 日 v2),并核对固定版本的官方推理代码。目标是弄清一个具体问题:为什么可以不还原整段历史的 K、V,直接在潜变量上计算,却不能顺手更改 softmax 的温度?这是基础机制解读,不是把旧报告包装成新发布。

1. 先固定比较对象:同一个 MLA,两个计算顺序

这里的“等价”指同一组 MLA 权重的两条前向计算路径:一条显式展开历史 K、V;另一条只保存潜变量,把线性映射挪到查询侧和聚合之后。它不意味着任意已经训练好的普通多头注意力都能无损变成 MLA。联合低秩结构本来就是模型架构约束,语言能力仍需训练和任务评估。

只看一个注意力层,采用列向量。第 \(j\) 个 token 的层输入为 \(x_j\in\mathbb R^D\),共有 \(H\) 个查询头。内容 key/query 的每头维度为 \(d_k\),value 维度为 \(d_v\),联合潜变量维度为 \(r\),单独的位置分支维度为 \(s\)。一个标量占 \(b\) 字节。批量、层数和缓存长度稍后分别记为 \(B,L,T\)。

把归一化也纳入潜变量的定义:

\[\begin{aligned}c_j&=\operatorname{RMSNorm}(W_Dx_j)\in\mathbb R^r,\\ k^C_{j,i}&=U_{K,i}c_j,\qquad v_{j,i}=U_{V,i}c_j,\\ U_{K,i}&\in\mathbb R^{d_k\times r},\qquad U_{V,i}\in\mathbb R^{d_v\times r}.\end{aligned}\]

上标 C 表示不做旋转的内容分支。当前查询 \(q^C_{t,i}\in\mathbb R^{d_k}\) 由模型的查询投影产生。官方实现还可能在查询侧使用低秩投影与归一化,这不改变下文从已生成查询开始的恒等式。尤其不要把 RMSNorm 跨过矩阵随意合并;我们吸收的是潜变量之后的线性映射,归一化本身保留在原位置。固定版本的官方 MLA 类同时提供 naive 与 absorb 路径,可逐步对照。

2. 两次交换顺序,历史 K、V 都不必展开

先暂时忽略位置项。一个查询和某个历史 key 的点积可以改写为:

\[\begin{aligned}(q^C_{t,i})^\top k^C_{j,i}&=(q^C_{t,i})^\top U_{K,i}c_j\\&=\underbrace{(U_{K,i}^\top q^C_{t,i})^\top}_{\widetilde q_{t,i}^{\,\top}}c_j.\end{aligned}\]

关键不是“压缩后相似度大致不变”,而是矩阵乘法结合律给出了相同的标量。对一个当前查询,每头只需算一次 \(\widetilde q_{t,i}\in\mathbb R^r\),就可以与所有历史潜变量点积;无需对每个历史 token 重新乘上 key 的展开矩阵。

设该头经过掩码与 softmax 后的权重为 \(a_{tj,i}\)。对 value 做同样的检查:

\[\begin{aligned}o_{t,i}&=\sum_{j\le t}a_{tj,i}U_{V,i}c_j\\&=U_{V,i}\underbrace{\left(\sum_{j\le t}a_{tj,i}c_j\right)}_{z_{t,i}\in\mathbb R^r},\\y_t&=\sum_{i=1}^{H}O_iU_{V,i}z_{t,i}.\end{aligned}\]

其中 \(O_i\in\mathbb R^{D\times d_v}\) 是输出投影对应第 i 个头的块。可以先聚合历史潜变量,再展开一次 value;也可以在代数上把 value 展开与输出投影合并。是否真的预先存储合并权重,取决于矩阵尺寸、量化格式与内核效率,恒等式不要求一定采用这种物理布局。

共享潜变量并没有消除多头差异:每个头的查询与注意力权重不同,因此每个头有自己的 \(z_{t,i}\)。不能先把各头的权重平均,再只做一次共同聚合。这里的等价是在实数运算意义下成立;浮点加法顺序改变会带来舍入差异,量化与裁剪还会引入另外的误差。

3. 为什么 RoPE 不能直接从等号中间穿过去?

RoFormer 的旋转位置编码用位置相关的旋转矩阵作用于查询和 key。记位置 t 的旋转为 \(R_t\)。如果直接旋转刚才的内容分支,得分会包含:

\[\begin{aligned}(R_tq^C_{t,i})^\top(R_jU_{K,i}c_j)&=(q^C_{t,i})^\top R_t^\top R_jU_{K,i}c_j\\&=(U_{K,i}^\top R_j^\top R_tq^C_{t,i})^\top c_j.\end{aligned}\]

最后一行虽仍是合法恒等式,但新的“查询”依赖历史位置 \(j\)。前面一次查询变换供整段历史复用的好处就丢了。并不是历史位置编码每生成一个 token 都会改变,而是位于两个投影之间的相对旋转,通常不能通过一个与历史位置无关的矩阵挪走。

一个二维反例就够了:令 \(U=\operatorname{diag}(1,2)\)、\(R=\begin{bmatrix}0&-1\\1&0\end{bmatrix}\)、\(q=(1,0)^\top\)、\(c=(1,1)^\top\)。那么 \(q^\top RUc=-2\),而 \(q^\top URc=-1\)。这甚至是在投影前后维度相同的有利情形;矩阵不交换,旋转潜变量不能一般性地替代旋转展开后的 key。具有特殊可交换结构的权重是额外假设,不是通用 MLA 权重的性质。

解耦方案把位置分支单独保留。定义 \(q^R_{t,i}=R_t\bar q^R_{t,i}\) 与 \(k^R_j=R_j\bar k^R_j\),两者都是 \(s\) 维,位置 key 在各头之间共享。于是完整 logits 为:

\[\begin{aligned}\ell_{tj,i}&=\gamma\left(\widetilde q_{t,i}^{\,\top}c_j+(q^R_{t,i})^\top k^R_j\right)+M_{tj},\\a_{tj,i}&=\frac{\exp(\ell_{tj,i})}{\sum_{u\le t}\exp(\ell_{tu,i})}.\end{aligned}\]

基本形式使用 \(\gamma=1/\sqrt{d_k+s}\)。\(M_{tj}\) 对允许读取的位置取零,对未来或无效位置取负无穷。softmax 实现要减去行最大值以稳定数值,并避免构造全部被屏蔽的查询行。历史缓存只需保存归一化的 \(c_j\) 与已经按其绝对位置旋转的 \(k^R_j\);后续查询不应再次旋转旧 key。

MLA 计算示意:历史 token 的共享潜变量和位置 key 进入缓存;当前每头查询变换后计算两项得分,以同一组注意力权重聚合潜变量,再展开 value 与输出。
原创机制示意,非实验数据。缓存跨头共享;查询、softmax 权重和聚合结果仍按头区分。内容项与位置项相加后只做一次 softmax。

4. 192 变成 576,为什么不能换 softmax 缩放?

以官方 V3 配置中的 \(d_k=d_v=128\)、\(r=512\)、\(s=64\) 为例:展开路径的 query/key 宽度是 \(d_k+s=192\),吸收后拼接的 query/cache 宽度是 \(r+s=576\)。但两条路径的未缩放点积完全相同。把缩放从 \(1/\sqrt{192}\) 改成 \(1/\sqrt{576}\),会把所有允许位置的 logits 再乘 \(1/\sqrt3\),通常使分布更平;这已经是另一个注意力算子。

这里不能套用“向量更宽,方差更大”的独立同分布直觉:吸收后的坐标由同一组训练权重线性变换而来,相关性也发生了变化。维度的变化不等于原点积方差被任意重设。只要目标是保持原函数,温度就由原模型约定决定。

这不是纯理论提醒。所核对的 FlashMLA 历史版本接口在没有传入 softmax_scale 时,会默认用查询最后一维的平方根倒数。因此接入吸收后的张量时,要显式传入模型的缩放值。V3 参考实现从原始非旋转与旋转维度之和计算基础缩放,并可能叠加长上下文的 mscale 修正;实际接入应完整保留这些设置,而不是硬编码本文的基础数值。本文固定代码提交,不把不同版本的硬件支持和接口混为一谈。

# Column-vector pseudocode; one current token, one head.
# c_cache: [T, r]; k_rope_cache: [T, s]
q_latent = U_K.T @ q_content       # [r]
logits = model_scale * (
    c_cache @ q_latent + k_rope_cache @ q_rope
)
weights = stable_softmax(logits + absolute_position_mask)
z = weights @ c_cache             # [r], separate for each head
o = U_V @ z                      # [d_v]
# Concatenate head outputs, then apply the model output projection.

分块 prefill 同样要使用绝对位置:已有前缀长度为 \(P\),当前块内第 \(u\) 个查询的零起始位置为 \(P+u\),只能读取满足 \(j\le P+u\) 的 key。仅对当前块画一个不带偏移的三角掩码,会错误屏蔽前缀或放开未来。单 token decode 在缓存仅包含过去和当前 token 时自然全部可见,但 padding 与无效缓存槽仍需排除。

5. 缓存节省有精确账本,速度没有固定兑换率

先比较同一 MLA 的两种缓存布局,展开路径也让位置 key 跨头共享。每层、每个历史 token 的元素数为:

\[\begin{aligned}n_{\mathrm{expanded}}&=H(d_k+d_v)+s,\\n_{\mathrm{latent}}&=r+s,\\\mathrm{Bytes}_{\mathrm{latent}}&=bBLT(r+s).\end{aligned}\]

加入 \(H=128\),上面的配置得到展开 32,832 个元素、潜变量 576 个元素,前者恰为后者 57 倍。若每元素两字节,潜变量缓存为每层每 token 1,152 字节。这是指定布局的张量容量算术,不含页表、分页空洞、量化尺度、临时工作区或分布式复制。官方参考代码的 naive 路径还会把位置 key 复制到各头,不能把那个额外冗余悄悄算成算法必需。

这个 57 倍也不是 V2 摘要里的“减少 93.3%”:后者的比较对象是 DeepSeek 67B,而这里固定同一 MLA 的维度和层,比较展开与吸收两条路径。不同架构、层数和缓存精度之间的部署数字,不能直接从这个小账本推出。

再看只与历史交互相关的乘加次数(MAC,一次乘加计一次),忽略 softmax 和投影,以单请求、单层、单个当前查询为口径:

路径得分与加权聚合的 MAC上述配置:每头每历史 token
展开 K、V\(HT(d_k+s+d_v)\)320
潜变量计算\(HT(2r+s)\)1,088

吸收路径这部分算术反而为 3.4 倍,还需要当前查询变换和聚合后 value 展开的 \(Hr(d_k+d_v)\) 次乘加。展开路径则要在新 token 到达时生成其 K、V,prefill 时还要为每个 token 做对应投影。比较整层必须把两边的工作都计入,而不能只挑一个局部公式。

为什么仍可能更快?潜变量跨头共享,长期缓存小得多;合适的内核可以提高数据复用与矩阵计算利用率。2025 年 4 月 22 日的 FlashMLA 内核说明专门讨论了 MLA decode 在其指定配置下也可能受计算而非带宽限制,并据此设计调度。这里不把供应方某台 GPU 的峰值数字外推成普遍加速倍数。

prefill 同时处理许多查询,朴素 dense 注意力的历史交互工作随长度近似二次增长;展开形式有时更适合成熟的矩阵内核。decode 每步通常只有少量新查询,保留压缩缓存的价值更直接。它们可以使用不同计算路径,但掩码、位置、归一化和缩放必须一致。Flash 类分块算法还能避免显式存储整张二次大小的注意力矩阵,因此“计算量二次”不等于“必须分配二次注意力显存”。

6. 把恒等式变成能失败的实现检查

本文附有可下载的 Python 标准库校验脚本,使用固定随机种子的合成小张量与双精度浮点数。它检验的是计算路径,不加载真实模型、不运行 GPU,也不测语言质量或吞吐。本轮已实际运行 12 组合成案例,正例最大绝对误差小于 \(10^{-14}\),低于预设容差 \(10^{-10}\);错误路径均触发了预期差异。完整输出和断言可在本地运行:

python3 mla_checks.py

校验分成三组:同一 MLA 的显式展开与吸收在 logits、注意力权重和输出上的一致性;整段因果 prefill、逐 token decode 及带前缀分块计算的一致性;故意改错 softmax 缩放、混用各头聚合或错误交换旋转与投影时,反例确实能产生差异。所有“通过”仅适用于脚本声明的数值容差与合成样例,不能据此声称真实低精度内核已验收。

接入生产内核时,建议先固定同一个 checkpoint、精度和绝对位置,从短序列的算子对照开始,再覆盖长上下文、多请求不同长度、非整页尾块及量化。应同时检查每层输出误差与最终任务质量;浮点重排或量化造成的小误差可能经过多层累积。

性能实验需要另开账本:固定 GPU、软件版本、张量并行、请求长度分布和并发,分别报告 prefill 首 token 时间、decode 每 token 时间及其分位数、峰值显存、有效输出吞吐;再在相同延迟约束下寻找可容纳批量。若瓶颈已经转到计算、通信或调度,释放显存并不保证单请求更快。本文没有执行这些 GPU 实验。

MLA 最值得复用的思路是:先找出哪些中间量必须长期保存,再用结合律把昂贵的逐历史 token 变换移到复用度更高的位置。要守住的是同一个函数——同一组权重、归一化、位置、掩码和温度;满足这个条件之后,才有资格讨论缓存和速度分别获得了什么。

参考资料

  1. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model (v5; first submitted 2024-05-07) · 2024-06-19 · 查阅 2026-10-07
  2. DeepSeek-V3 Technical Report (v2; first submitted 2024-12-27) · 2025-02-18 · 查阅 2026-10-07
  3. DeepSeek-V3 official inference model.py: naive and absorb MLA paths (commit b15f0db) · 2025-08-27 · 查阅 2026-10-07
  4. FlashMLA interface: explicit softmax_scale and dense MLA dimensions (commit ba89a34) · 2026-09-15 · 查阅 2026-10-07
  5. A Deep-Dive Into the New Flash MLA Kernel (2025-04-22; pinned snapshot) · 2025-04-22 · 查阅 2026-10-07
  6. RoFormer: Enhanced Transformer with Rotary Position Embedding · 2024-02-01 · 查阅 2026-10-07
利友诚

关于作者

利友诚 · Youcheng Li

北京大学智能学院人工智能专业博士研究生,导师为王立威教授;Isoplex Intelligence(壹索智能)联合创始人兼 CTO。

研究关注医疗人工智能、生成式基础模型、诊断推理与科学智能体。以第一作者或共同第一作者身份在 Nature Biomedical Engineering、Scientific Data、KDD 和 PLOS Computational Biology 发表研究。