1. Opening

本文会一直使用同一个视角:把 attention 写成 tensor contraction,并显式保留每一个 dimension。在这个视角下,MLA 里一些看似分散的现象,可以被更加深入和统一地理解。具体来说,可以得到几个有趣的结论:

  • 在 prefill 中,MLA 减少了 linear projections 的 FLOPs,但不能过早 absorption:quadratic core-attention 项仍然应该在 dk,dvd_k,d_v 上计算,而不是 dcd_c
  • 在 decode 中,最优顺序会改变,因为 cache 已经存在。吸收 key/value projections 后,长度为 tt 的 cache computation 可以经过 latent state,而不是重构出来的 per-head K,VK,V
  • RoPE 的麻烦恰好在于它让 latent-to-key map 依赖 position,从而破坏完整的 key absorption。
  • Q compression 的主要贡献是降低训练时的 activation memory,以及 prefill/training 阶段的 projection FLOPs;它对 decode 的贡献很少。
  • MTP-style speculative decoding 在 MLA 上表现不同,是因为 MLA decode 不再读取完整的 per-head KV cache,而是读取压缩后的 latent cache 和 RoPE cache,并进行额外的 contraction。

2. Tensor View

Takeaway. 相比把 tensor flatten 成二维矩阵,tensor view 会保留每一个有意义的 axis。FLOPs 可以直接从 output 和 contracted axes 读出,而不同 computation orders 也能在同一套 notation 下比较。

本文的 tensor view 只是一句话:每一次 tensor contraction 都由三类轴决定 —— output axes、contracted axes 和 broadcast axes。只要这些轴写清楚,FLOPs 就是 shape accounting。

  • Q(h,T,dk)=X(T,d)WQ(h,d,dk)    einsum("td,sdk->stk",X,WQ).\underset{(h,T,d_k)}{Q}=\underset{(T,d)}{X}\underset{(h,d,d_k)}{W_Q}\;\Longleftrightarrow\;\mathrm{einsum}(\texttt{"td,sdk->stk"},X,W_Q).

  • S(h,T,T)=Q(h,T,dk)K(h,dk,T)    einsum("stk,suk->stu",Q,K).\underset{(h,T,T)}{S}=\underset{(h,T,d_k)}{Q}\underset{(h,d_k,T)}{K^\top}\;\Longleftrightarrow\;\mathrm{einsum}(\texttt{"stk,suk->stu"},Q,K).

一般形式。 对于一个 output axes 为 B1,,Br,m,nB_1,\ldots,B_r,m,n、contracted axis 为 kk 的 batched contraction,FLOPs 是:

FLOPs  =  2B1Brmnk.\mathrm{FLOPs} \;=\; 2 \cdot B_1 \cdots B_r \cdot m \cdot n \cdot k .

换句话说,把所有 output dimensions 乘起来,再乘 contracted dimension,最后乘以 multiply-add 的 22(这是标准的 GEMM 计数约定:一次 fused multiply-add 包含一次乘法和一次加法,计作 2 FLOPs)。因此 Q=XWQQ=XW_Q 的 cost 是 2hTddk2hTdd_k,而 S=QKS=QK^\top 的 cost 是 2hT2dk2hT^2d_k

3. Notation and Preliminaries

Takeaway. 本节通过直接在公式上标注 shape 来约定符号,并给出 MHA、MLA 和后文所需 RoPE rotation 的基本介绍。

关于从 MHA/MQA/GQA 到 MLA 的背景,可以参考 苏剑林的笔记。本文不展开比较这些变体;这里只把 MHA 当作干净的代数 reference,把 MLA 当作主要推导对象,并介绍后文 absorption 讨论所需的 RoPE 恒等式。

为了让推导保持简洁,我们使用以下约定:

  • 逐元素的 attention score 统一写成 Ss,t,uS_{s,t,u}hh 是 head 总数,s=1,,hs=1,\ldots,h 是具体的 head,tt 是 query position,uu 是 key position。也就是说,score matrix 的行对应 tt,列对应 uu
  • 在第 5 节之前,MLA 公式暂时不考虑 RoPE。
  • 全文忽略 causal mask;读者可以直接把完整的 attention matrix 替换成对应的 causal 形式。
  • FLOPs 对比中不统计 softmax。MHA 和 MLA 对相同 shape 的 attention score 执行相同的 softmax,而模型 FLOPs 以占主导的 GEMM 和 tensor contraction 作为可比较指标。PyTorch SDPA 等 optimized implementation 也会把 softmax 融合进 scaled-dot-product attention。第 5 节仍然计算 content score 与 RoPE score 的显式 addition,因为这是 decoupled branch 引入的额外操作,并且随完整 attention matrix 增长。
  • AA^\top 表示只交换最后两维,batch axes 不动。

MHA. Reference 形式是:

Q(h,T,dk)=X(T,d)WQ(h,d,dk),K(h,T,dk)=X(T,d)WK(h,d,dk),V(h,T,dv)=X(T,d)WV(h,d,dv).\underset{(h,T,d_k)}{Q} = \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_Q}, \qquad \underset{(h,T,d_k)}{K} = \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_K}, \qquad \underset{(h,T,d_v)}{V} = \underset{(T,d)}{X} \underset{(h,d,d_v)}{W_V}.

然后 attention 是:

S(h,T,T)=Q(h,T,dk)K(h,dk,T),P(h,T,T)=softmax ⁣(Sdk),\underset{(h,T,T)}{S} = \underset{(h,T,d_k)}{Q} \underset{(h,d_k,T)}{K^\top}, \qquad \underset{(h,T,T)}{P} = \operatorname{softmax}\!\left(\frac{S}{\sqrt{d_k}}\right), O(h,T,dv)=P(h,T,T)V(h,T,dv),Y(T,d)=O(h,T,dv)Wo(h,dv,d).\underset{(h,T,d_v)}{O} = \underset{(h,T,T)}{P} \underset{(h,T,d_v)}{V}, \qquad \underset{(T,d)}{Y} = \underset{(h,T,d_v)}{O} \underset{(h,d_v,d)}{W_o}.

最后一项会同时 contract head axis 和 dvd_v axis。这等价于通常的实现:先 concat 所有 heads,再做一次 output projection。因为把 (h,dv)(h,d_v) flatten 成一个 hdvhd_v axis,并不会改变求和:

Yt,d=s=1hv=1dvOs,t,v(Wo)s,v,d=j=1hdv(Oconcat)t,j(Wo,concat)j,d.Y_{t,d}=\sum_{s=1}^{h}\sum_{v=1}^{d_v}O_{s,t,v}(W_o)_{s,v,d} =\sum_{j=1}^{hd_v}(O_{\mathrm{concat}})_{t,j}(W_{o,\mathrm{concat}})_{j,d}.

两种写法的 FLOPs 也相同。在 tensor contraction 形式中,output axes 是 TTdd,contracted axes 是 hhdvd_v,所以 FLOPs 是 2Tdhdv2Tdh d_v。Concat 之后,矩阵乘的 shapes 是 (T,hdv)(hdv,d)(T,hd_v)(hd_v,d),所以 FLOPs 是 2T(hdv)d=2Thdvd2T(hd_v)d=2Thd_vd


MLA. KV path 把 per-head 的 K,VK,V projection 换成一个 shared latent。为了让推导保持简洁,在第 5 节之前,我们暂时不考虑 MLA 中的 RoPE:

CKV(T,dc)=X(T,d)WDKV(d,dc),K(h,T,dk)=CKV(T,dc)WUK(h,dc,dk),\underset{(T,d_c)}{C^{KV}} = \underset{(T,d)}{X} \underset{(d,d_c)}{W_{DKV}}, \qquad \underset{(h,T,d_k)}{K} = \underset{(T,d_c)}{C^{KV}} \underset{(h,d_c,d_k)}{W_{UK}}, V(h,T,dv)=CKV(T,dc)WUV(h,dc,dv).\underset{(h,T,d_v)}{V} = \underset{(T,d_c)}{C^{KV}} \underset{(h,d_c,d_v)}{W_{UV}}.

DeepSeek-V2 论文声称,query compression 可以降低训练时的 activation memory;但 DeepSeek-V2-Lite 不压缩 query。这正好提供了一个自然的对照:先推导 no-Q MLA,再加入 Q compression,单独看它改变了什么。

不压缩 Q。

Q(h,T,dk)=X(T,d)WQ(h,d,dk).\underset{(h,T,d_k)}{Q} = \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_Q}.

Q compression。

CQ(T,dqc)=X(T,d)WDQ(d,dqc),Q(h,T,dk)=CQ(T,dqc)WUQ(h,dqc,dk).\underset{(T,d_{qc})}{C^Q} = \underset{(T,d)}{X} \underset{(d,d_{qc})}{W_{DQ}}, \qquad \underset{(h,T,d_k)}{Q} = \underset{(T,d_{qc})}{C^Q} \underset{(h,d_{qc},d_k)}{W_{UQ}}.

RoPE。

在 tensor view 中,RoPE 只是在最后一个 tensor dimension 上插入一个 position-dependent linear map,不改变 tensor shape。因为 RoPE 作用于 QQKK,这里用它们的 tensor dimension dkd_k 作为 rotation dimension。对于偶数 dkd_k,standard rotation matrix 是

Rt=diag ⁣([cos(tθ1)sin(tθ1)sin(tθ1)cos(tθ1)],,[cos(tθdk/2)sin(tθdk/2)sin(tθdk/2)cos(tθdk/2)])Rdk×dk.R_t= \operatorname{diag}\!\left( \begin{bmatrix} \cos(t\theta_1) & -\sin(t\theta_1)\\ \sin(t\theta_1) & \cos(t\theta_1) \end{bmatrix}, \ldots, \begin{bmatrix} \cos(t\theta_{d_k/2}) & -\sin(t\theta_{d_k/2})\\ \sin(t\theta_{d_k/2}) & \cos(t\theta_{d_k/2}) \end{bmatrix} \right) \in\mathbb{R}^{d_k\times d_k}.

把 RoPE base 记作 bb,则 θj=b2(j1)/dk\theta_j=b^{-2(j-1)/d_k}。RoFormer 使用 b=10000b=10000;现代语言模型为了支持更长的 context,可能会使用更大的 base。由于本文使用 row vector,把各个 position 的 transposed rotation stack 成一个 tensor:R(T,dk,dk)\underset{(T,d_k,d_k)}{\mathcal R},其中 Rt,:,:=Rt\mathcal R_{t,:,:}=R_t^\top

RoPE 和后续 score contraction 可以统一写成

Q~(h,T,dk)=Q(h,T,dk)R(T,dk,dk),K~(h,T,dk)=K(h,T,dk)R(T,dk,dk),S(h,T,T)=Q~(h,T,dk)K~(h,dk,T).\underset{(h,T,d_k)}{\widetilde Q} = \underset{(h,T,d_k)}{Q} \underset{(T,d_k,d_k)}{\mathcal R}, \qquad \underset{(h,T,d_k)}{\widetilde K} = \underset{(h,T,d_k)}{K} \underset{(T,d_k,d_k)}{\mathcal R}, \qquad \underset{(h,T,T)}{S} = \underset{(h,T,d_k)}{\widetilde Q} \underset{(h,d_k,T)}{\widetilde K^\top}.

两个 TT axis 对应同一个 position,因此输出只保留一个 TTQQKKdkd_k axis 与 R\mathcal R 的第一个 dkd_k axis contraction,第二个 dkd_k axis 保留;R\mathcal R 在 head axis 上 broadcast。因此 Q,KQ,K(h,T,dk)(h,T,d_k) shape 不变。又因为 RtRu=RutR_t^\top R_u=R_{u-t}

Ss,t,u=Qs,t,:RtRuKs,u,:=Qs,t,:RutKs,u,:,S_{s,t,u} = Q_{s,t,:}R_t^\top R_uK_{s,u,:}^\top = Q_{s,t,:}R_{u-t}K_{s,u,:}^\top,

所以每次 rotation 编码 absolute position,而 score 只依赖 relative position。

实际的 elementwise 计算。 上面的 rotation tensor 是代数表示;implementation 不需要构造 dense dk×dkd_k\times d_k matrix。令 m=dk/2m=d_k/2。对于任意 qRdkq\in\mathbb R^{d_k},定义 adjacent-pair coordinates

qj,0:=q2j1,qj,1:=q2j,j=1,,m.q_{j,0}:=q_{2j-1}, \qquad q_{j,1}:=q_{2j}, \qquad j=1,\ldots,m.

沿用前面的 θj\theta_j,对应的 phase tensor 是

Θ^t,2j1=Θ^t,2j=tθj,j=1,,m,Θ^(T,dk).\widehat\Theta_{t,2j-1} = \widehat\Theta_{t,2j} = t\theta_j, \qquad j=1,\ldots,m, \qquad \underset{(T,d_k)}{\widehat\Theta}.

然后直接计算

Q~(h,T,dk)=Q(h,T,dk)cosΘ^(T,dk)+rotate_half ⁣(Q(h,T,dk))sinΘ^(T,dk),\underset{(h,T,d_k)}{\widetilde Q} = \underset{(h,T,d_k)}{Q} \odot \cos\underset{(T,d_k)}{\widehat\Theta} + \operatorname{rotate\_half}\!\left( \underset{(h,T,d_k)}{Q} \right) \odot \sin\underset{(T,d_k)}{\widehat\Theta}, K~(h,T,dk)=K(h,T,dk)cosΘ^(T,dk)+rotate_half ⁣(K(h,T,dk))sinΘ^(T,dk).\underset{(h,T,d_k)}{\widetilde K} = \underset{(h,T,d_k)}{K} \odot \cos\underset{(T,d_k)}{\widehat\Theta} + \operatorname{rotate\_half}\!\left( \underset{(h,T,d_k)}{K} \right) \odot \sin\underset{(T,d_k)}{\widehat\Theta}.

这里 \odot 是 elementwise multiplication;cosine 和 sine tensors 在 head axis 上 broadcast。在上面的 pair notation 中,rotate_half 定义为

rotate_half(q)j,0=qj,1,rotate_half(q)j,1=qj,0.\operatorname{rotate\_half}(q)_{j,0}=-q_{j,1}, \qquad \operatorname{rotate\_half}(q)_{j,1}=q_{j,0}.

它不改变 shape,也不消去任何 axis。这个 elementwise formulation 与上面的 rotation-tensor formulation 完全等价;本文把它纳入 tensor view,但不单独计入后文以主要 tensor contractions 为对象的 FLOPs comparison。

4. Prefill、Training 与 Decode 中的计算顺序和 FLOPs

我们推导每一步 tensor contraction 及其 FLOPs,将 MHA 和 MLA 对比。

4.1 Prefill 与 Training:标准计算顺序

Takeaway. MLA 减少了 linear projections 的 FLOPs。它没有减少 T2T^2 core-attention FLOPs,因为 QKQK^\topPVPV 的 contraction 仍然分别沿 dkd_kdvd_v 进行。


MHA。

对于 MHA prefill,完整的 score 表达式是:

从左到右读这条 chain,可以把它看成是在选择何时产生更小的 representation。Weights 会在任何跨越全部 TT 个 query-key pairs 的操作之前,把 model dimension dd 降到 per-head key dimension dkd_k

S=QK=X(T,d)WQ(h,d,dk)WK(h,dk,d)X(d,T).S = QK^\top = \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_Q} \underset{(h,d_k,d)}{W_K^\top} \underset{(d,T)}{X^\top}.

从 FLOPs 的角度,关键是让占主导的 T2T^2 项系数尽可能小。由于 dk<dd_k<d,计算顺序应该是先构造 Q=XWQQ=XW_QK=XWKK=XW_K,再计算 S=QKS=QK^\top

X(T,d)WQ(h,d,dk)Q(h,T,dk),FLOPs=2Thddk,X(T,d)WK(h,d,dk)K(h,T,dk),FLOPs=2Thddk,Q(h,T,dk)K(h,dk,T)S(h,T,T),FLOPs=2hT2dk.\begin{aligned} \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_Q} &\rightarrow \underset{(h,T,d_k)}{Q}, &\qquad \mathrm{FLOPs}&=2Thdd_k, \\ \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_K} &\rightarrow \underset{(h,T,d_k)}{K}, &\qquad \mathrm{FLOPs}&=2Thdd_k, \\ \underset{(h,T,d_k)}{Q} \underset{(h,d_k,T)}{K^\top} &\rightarrow \underset{(h,T,T)}{S}, &\qquad \mathrm{FLOPs}&=2hT^2d_k. \end{aligned}

前两项对 TT 是 linear 的,因为每个 token 都独立做 projection。只有最后一次 contraction 会把每个 query 与每个 key 连接起来,因此它是 score path 中唯一的 T2T^2 项。

完整的 value 和 output chain 是:

P(h,T,T)X(T,d)WV(h,d,dv)Wo(h,dv,d)Y(T,d).\underset{(h,T,T)}{P} \underset{(T,d)}{X} \underset{(h,d,d_v)}{W_V} \underset{(h,d_v,d)}{W_o} \rightarrow \underset{(T,d)}{Y}.

Value/output order:使用 (P(XWV))Wo(P(XW_V))W_o,而不是 ((PX)WV)Wo((PX)W_V)W_oP(X(WVWo))P(X(W_VW_o))

这里的计算顺序问题是:与 PP 做 quadratic multiplication 时,哪一个 tensor dimension 会被带进去。先构造 VV,可以让 attention 落在 dvd_v 上,而不是让更宽的 model dimension dd 穿过所有 token pairs。

标准顺序是 XWVVXW_V\rightarrow VPVOPV\rightarrow OOWoYOW_o\rightarrow Y

F(P(XWV))Wo=2ThddvXWV+2hT2dvPV+2ThdvdOWo.F_{(P(XW_V))W_o} = \underbrace{2Thdd_v}_{XW_V} +\underbrace{2hT^2d_v}_{PV} +\underbrace{2Thd_vd}_{OW_o}.

先计算 PXPX 时:

F((PX)WV)Wo=2hT2dPX+2Thddv(PX)WV+2ThdvdOWo.F_{((PX)W_V)W_o} = \underbrace{2hT^2d}_{PX} +\underbrace{2Thdd_v}_{(PX)W_V} +\underbrace{2Thd_vd}_{OW_o}.

预先组合 WVWoW_VW_o 时:

FP(X(WVWo))=2hd2dvWVWo+2hTd2X(WVWo)+2hT2dP[X(WVWo)].F_{P(X(W_VW_o))} = \underbrace{2hd^2d_v}_{W_VW_o} +\underbrace{2hTd^2}_{X(W_VW_o)} +\underbrace{2hT^2d}_{P[X(W_VW_o)]}.

后两种顺序都会让 quadratic contraction 落在 dd 而不是 dvd_v 上。由于 dvdd_v\ll d标准顺序更便宜:2hT2dv2hT^2d_v,而不是 2hT2d2hT^2d

把这些项合在一起,dense MHA prefill FLOPs 是:

FMHA,prefill=2Tdh(2dk+dv)XWQ,XWK,XWV+2hT2(dk+dv)QK,PV+2ThdvdOWo.F_{\mathrm{MHA,prefill}} =\underbrace{2Tdh(2d_k+d_v)}_{XW_Q,\,XW_K,\,XW_V} +\underbrace{2hT^2(d_k+d_v)}_{QK^\top,\,PV} +\underbrace{2Thd_vd}_{OW_o}.

MLA。

对于 MLA prefill,首先计算 shared KV latent CKV=XWDKVC^{KV}=XW_{DKV}。不使用 query compression 时,完整的 score 表达式是:

S=X(T,d)WQ(h,d,dk)WUK(h,dk,dc)WDKV(dc,d)X(d,T).S = \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_Q} \underset{(h,d_k,d_c)}{W_{UK}^\top} \underset{(d_c,d)}{W_{DKV}^\top} \underset{(d,T)}{X^\top}.

使用 query compression 时,它变成:

S=X(T,d)WDQ(d,dqc)WUQ(h,dqc,dk)WUK(h,dk,dc)WDKV(dc,d)X(d,T).S = \underset{(T,d)}{X} \underset{(d,d_{qc})}{W_{DQ}} \underset{(h,d_{qc},d_k)}{W_{UQ}} \underset{(h,d_k,d_c)}{W_{UK}^\top} \underset{(d_c,d)}{W_{DKV}^\top} \underset{(d,T)}{X^\top}.

无论以哪种顺序计算 score,output shape 都是 (h,T,T)(h,T,T),所以 hT2hT^2 是固定因子;我们只需要比较 contracted tensor dimension。因此应该让这个 dimension 是 dkd_k(在 DeepSeek-V2 中,dk=128<dc=512<d=5120d_k=128<d_c=512<d=5120)。于是可以把 formula 分成两部分:XWDQWUQXW_{DQ}W_{UQ}XWDKVWUKXW_{DKV}W_{UK}

Q compression 的 projection order:使用 (XWDQ)WUQ(XW_{DQ})W_{UQ},而不是 X(WDQWUQ)X(W_{DQ}W_{UQ})

不使用 Q compression 时,Q=XWQQ=XW_Q 与 MHA 相同(FLOPs:2Thddk2Thdd_k)。它只有一次 contraction,因此 query side 没有计算顺序可选。下面的比较只在把 WQW_Q factorize 成 WDQWUQW_{DQ}W_{UQ} 后才会出现。

使用 Q compression 时,factorized order 先把每个 token 映射到 CQC^Q,再通过 WUQW_{UQ} 展开 head axis:

X(T,d)WDQ(d,dqc)CQ(T,dqc),FLOPs=2Tddqc,CQ(T,dqc)WUQ(h,dqc,dk)Q(h,T,dk),FLOPs=2Thdqcdk.\begin{aligned} \underset{(T,d)}{X} \underset{(d,d_{qc})}{W_{DQ}} &\rightarrow \underset{(T,d_{qc})}{C^Q}, &\qquad \mathrm{FLOPs}&=2Tdd_{qc}, \\ \underset{(T,d_{qc})}{C^Q} \underset{(h,d_{qc},d_k)}{W_{UQ}} &\rightarrow \underset{(h,T,d_k)}{Q}, &\qquad \mathrm{FLOPs}&=2Thd_{qc}d_k. \end{aligned}

如果先合并 weights,两次 contraction 是

WDQ(d,dqc)WUQ(h,dqc,dk)WQeff(h,d,dk),FLOPs=2hddqcdk,X(T,d)WQeff(h,d,dk)Q(h,T,dk),FLOPs=2Thddk.\begin{aligned} \underset{(d,d_{qc})}{W_{DQ}} \underset{(h,d_{qc},d_k)}{W_{UQ}} &\rightarrow \underset{(h,d,d_k)}{W_Q^{\mathrm{eff}}}, &\qquad \mathrm{FLOPs}&=2hdd_{qc}d_k, \\ \underset{(T,d)}{X} \underset{(h,d,d_k)}{W_Q^{\mathrm{eff}}} &\rightarrow \underset{(h,T,d_k)}{Q}, &\qquad \mathrm{FLOPs}&=2Thdd_k. \end{aligned}

因此,两种顺序的 activation-side cost 分别是

F(XWDQ)WUQ=2Tddqc+2Thdqcdk=66.1MT,FX(WDQWUQ)=2Thddk=167.8MT(DeepSeek-V2).\begin{aligned} F_{(XW_{DQ})W_{UQ}} &=2Tdd_{qc}+2Thd_{qc}d_k =66.1\text{M}T, \\ F_{X(W_{DQ}W_{UQ})} &=2Thdd_k =167.8\text{M}T \qquad\text{(DeepSeek-V2)}. \end{aligned}

其中第二个公式没有计算额外的 weight-precomposition cost 2hddqcdk2hdd_{qc}d_k。当

dqc(d+hdk)<hddkd_{qc}(d+hd_k)<hdd_k

时,factorized order 更便宜。收益来自在两次 contraction 之间引入 CQ:(T,dqc)C^Q:(T,d_{qc}),使 per-head contraction 沿 dqcd_{qc} 而不是 dd 进行。

同理,对于 key and score order,先计算 CKV=XWDKVC^{KV}=XW_{DKV},重构 K=CKVWUKK=C^{KV}W_{UK},再计算 S=QKS=QK^\top

代入 DeepSeek-V2 的 dimensions,key-projection order (XWDKV)WUK(XW_{DKV})W_{UK} 的 cost 是 22.0MT22.0\text{M}T,而 X(WDKVWUK)X(W_{DKV}W_{UK}) 的 cost 是 167.8MT167.8\text{M}T;后者没有计算一次性的 weight-composition cost。

Value and output order:使用 (P(CKVWUV))Wo(P(C^{KV}W_{UV}))W_o,而不是 ((PCKV)WUV)Wo((PC^{KV})W_{UV})W_oP(CKV(WUVWo))P(C^{KV}(W_{UV}W_o))

标准顺序在计算 PVPV 之前重构 VV

F(P(CKVWUV))Wo=2ThdcdvCKVWUV+2hT2dvPV+2ThdvdOWo.F_{(P(C^{KV}W_{UV}))W_o} =\underbrace{2Thd_cd_v}_{C^{KV}W_{UV}} +\underbrace{2hT^2d_v}_{PV} +\underbrace{2Thd_vd}_{OW_o}.

如果先计算 PCKVPC^{KV},quadratic term 会变成 2hT2dc2hT^2d_c。如果先合并 WUVWoW_{UV}W_o,它会变成 2hT2d2hT^2d。由于 dv<dc<dd_v<d_c<d,先重构 VV 可以让 quadratic contraction 保持在最小的 tensor dimension dvd_v 上。

不使用 query compression 时,dense MLA prefill FLOPs 是:

FMLA,prefill,noQ=2ThddkQ=XWQ+2TddcCKV=XWDKV+2Thdc(dk+dv)K=CKVWUK,  V=CKVWUV+2hT2(dk+dv)S=QK,  O=PV+2ThdvdY=OWo.\begin{aligned} F_{\mathrm{MLA,prefill,noQ}} &=\underbrace{2Thdd_k}_{Q=XW_Q} \\ &\quad+\underbrace{2Tdd_c}_{C^{KV}=XW_{DKV}} +\underbrace{2Thd_c(d_k+d_v)}_{K=C^{KV}W_{UK},\;V=C^{KV}W_{UV}} \\ &\quad+\underbrace{2hT^2(d_k+d_v)}_{S=QK^\top,\;O=PV} +\underbrace{2Thd_vd}_{Y=OW_o}. \end{aligned}

使用 query compression 时:

FMLA,prefill,Qcomp=2TddqcCQ=XWDQ+2ThdqcdkQ=CQWUQ+2TddcCKV=XWDKV+2Thdc(dk+dv)K=CKVWUK,  V=CKVWUV+2hT2(dk+dv)S=QK,  O=PV+2ThdvdY=OWo.\begin{aligned} F_{\mathrm{MLA,prefill,Qcomp}} &=\underbrace{2Tdd_{qc}}_{C^Q=XW_{DQ}} +\underbrace{2Thd_{qc}d_k}_{Q=C^QW_{UQ}} \\ &\quad+\underbrace{2Tdd_c}_{C^{KV}=XW_{DKV}} +\underbrace{2Thd_c(d_k+d_v)}_{K=C^{KV}W_{UK},\;V=C^{KV}W_{UV}} \\ &\quad+\underbrace{2hT^2(d_k+d_v)}_{S=QK^\top,\;O=PV} +\underbrace{2Thd_vd}_{Y=OW_o}. \end{aligned}

对比。

Prefill FLOPs 可以分成三部分:linear projectionscore-attention computationfinal output projection。下面的表格汇总推导结果并代入 DeepSeek-V2 dimensions;之后的图展示除 RoPE branch 之外的 FLOPs 随 sequence length TT 的变化。

MethodLinear-projection FLOPsCore-attention FLOPs (QKQK^\topPVPV)Output-projection FLOPsDeepSeek-V2 prefill FLOPs(不包括 RoPE branch)
MHA2Tdh(2dk+dv)2Tdh(2d_k+d_v)2hT2(dk+dv)2hT^2(d_k+d_v)2Thdvd2Thd_vd65,536T2+671,088,640T65{,}536T^2+671{,}088{,}640T
MLA without query compression2Tddc+2Thddk+2Thdc(dk+dv)2Tdd_c+2Thdd_k+2Thd_c(d_k+d_v)2hT2(dk+dv)2hT^2(d_k+d_v)2Thdvd2Thd_vd65,536T2+374,341,632T65{,}536T^2+374{,}341{,}632T
MLA with query compression2Tddc+2Tddqc+2Thdqcdk+2Thdc(dk+dv)2Tdd_c+2Tdd_{qc}+2Thd_{qc}d_k+2Thd_c(d_k+d_v)2hT2(dk+dv)2hT^2(d_k+d_v)2Thdvd2Thd_vd65,536T2+272,629,760T65{,}536T^2+272{,}629{,}760T

DeepSeek-V2 attention configuration:d=5120, h=128, dk=dv=128, dc=512, dqc=1536d=5120,\ h=128,\ d_k=d_v=128,\ d_c=512,\ d_{qc}=1536

除 RoPE branch 外,MHA 与 MLA 的 prefill FLOPs 随 sequence length 的变化

4.2 Decode:Latent Cache 与矩阵吸收

Takeaway. 单步 decode 中,current-query length 是 11,cache length 是 tt。更省的顺序就是把这个 11 留在外面,尽早 contraction 掉其他 dimensions,最后再做 expansion。这样无需为全部 tt 个 cached tokens 重构完整的 K,VK,V:MLA FLOPs 随 tt 增长得更快,但读取的 cache data 更少。


MHA。

在 decode 当前位置 tt 之前,历史 K<t,V<tK_{<t},V_{<t} 已经在此前的 decode steps 中计算完成。当前 step 只计算新的 qt,kt,vtq_t,k_t,v_t;append kt,vtk_t,v_t 之后,这一步读取的 cache 是 Kt,VtK_{\le t},V_{\le t}。跨所有 heads,current-token projection cost 是 2dh(2dk+dv)2dh(2d_k+d_v)

与 prefill 不同,这里只有一个新的 query,因此 quadratic T2T^2 interaction 消失了。现在增长的因子是 tt:当前 query 扫过 tt 个 cached keys,再用得到的 probabilities 合并 tt 个 cached values。

Score、value 和 output contractions 是:

q(h,1,dk)Kt(h,dk,t)score(h,1,t),FLOPs=2htdk,p(h,1,t)Vt(h,t,dv)o(h,1,dv),FLOPs=2htdv,o(h,1,dv)Wo(h,dv,d)y(1,d),FLOPs=2hdvd.\begin{aligned} \underset{(h,1,d_k)}{q} \underset{(h,d_k,t)}{K_{\le t}^\top} &\rightarrow \underset{(h,1,t)}{\mathrm{score}}, &\qquad \mathrm{FLOPs}&=2htd_k, \\ \underset{(h,1,t)}{p} \underset{(h,t,d_v)}{V_{\le t}} &\rightarrow \underset{(h,1,d_v)}{o}, &\qquad \mathrm{FLOPs}&=2htd_v, \\ \underset{(h,1,d_v)}{o} \underset{(h,d_v,d)}{W_o} &\rightarrow \underset{(1,d)}{y}, &\qquad \mathrm{FLOPs}&=2hd_vd. \end{aligned}

两个 chain 遵循同一个原则:从 singleton current-token state 出发,选择让每个 intermediate 尽可能小的 contraction order。

把所有 dimensions 显式写出后,完整的 value-output chain 是

y(1,d)=p(h,1,t)Vt(h,t,dv)Wo(h,dv,d).\underset{(1,d)}{y} = \underset{(h,1,t)}{p} \underset{(h,t,d_v)}{V_{\le t}} \underset{(h,d_v,d)}{W_o}.

两种 association 的 cost 不同:

F(pVt)Wo=2htdvpVt:(h,1,dv)+2hdvd(pVt)Wo:(1,d),F_{(pV_{\le t})W_o} = \underbrace{2htd_v}_{pV_{\le t}:(h,1,d_v)} +\underbrace{2hd_vd}_{(pV_{\le t})W_o:(1,d)}, Fp(VtWo)=2htdvdVtWo:(h,t,d)+2htdp(VtWo):(1,d).F_{p(V_{\le t}W_o)} = \underbrace{2htd_vd}_{V_{\le t}W_o:(h,t,d)} +\underbrace{2htd}_{p(V_{\le t}W_o):(1,d)}.

从左到右计算,会让长度为 11 的 query axis 始终保持为 output axis,并在 output projection 之前先消去 cache-length axis。另一种顺序则要先把 WoW_o 应用于全部 tt 个 cached values,使 intermediate 同时带着 ttdd,因此计算代价更高。

Score path 也是同一个原则。把 dimensions 显式写出来:

xt(1,d)WQ(h,d,dk)Kt(h,dk,t)scoret(h,1,t).\underset{(1,d)}{x_t} \underset{(h,d,d_k)}{W_Q} \underset{(h,d_k,t)}{K_{\le t}^{\top}} \rightarrow \underset{(h,1,t)}{\mathrm{score}_t}.

把 singleton query axis 11 保持在外层,在它周围依次 contraction 其他 feature 和 cache axes,而不是先把它们展开。

因此 cached MHA decode FLOPs 是:

FMHA,decode,cached=2dh(2dk+dv)xtWQ,xtWK,xtWV+2ht(dk+dv)qKt,pVt+2hdvdoWo.F_{\mathrm{MHA,decode,cached}} =\underbrace{2dh(2d_k+d_v)}_{x_tW_Q,\,x_tW_K,\,x_tW_V} +\underbrace{2ht(d_k+d_v)}_{qK_{\le t}^\top,\,pV_{\le t}} +\underbrace{2hd_vd}_{oW_o}.

MLA。

对于 MLA decode,cache 存储 shared latent,而不是 full per-head K,VK,V。在这个单独的 decode step 中,新的 latent 是

xt(1,d)WDKV(d,dc)ctKV(1,dc),FLOPs=2ddc,\underset{(1,d)}{x_t} \underset{(d,d_c)}{W_{DKV}} \rightarrow \underset{(1,d_c)}{c_t^{KV}}, \qquad \mathrm{FLOPs}=2dd_c,

append 之后的 cache 是 CtKV=[C<tKV;ctKV]C^{KV}_{\le t}=[C^{KV}_{<t};c_t^{KV}],shape 为 (t,dc)(t,d_c)。下面先把完整 contraction chains 写出来,不预先假定 association。

Score chain。 不失一般性,只考虑带 Q compression 的情况。完整 contraction 是

scoret(h,1,t)=xt(1,d)WDQ(d,dqc)WUQ(h,dqc,dk)WUK(h,dk,dc)(CtKV)(dc,t).\underset{(h,1,t)}{\mathrm{score}_t} = \underset{(1,d)}{x_t} \underset{(d,d_{qc})}{W_{DQ}} \underset{(h,d_{qc},d_k)}{W_{UQ}} \underset{(h,d_k,d_c)}{W_{UK}^{\top}} \underset{(d_c,t)}{(C^{KV}_{\le t})^{\top}}.

先看 tt 的系数。Cache 是唯一带有 tt axis 的 factor。如果先把它左边的 factors 全部 contraction 掉,左侧 intermediate 的 shape 是 (h,1,dc)(h,1,d_c),最后一次 contraction 的 cost 是

2htdc.2htd_c.

这已经是 length-tt term 能达到的最小值:hhttdcd_c 无法消掉,剩下的 query axis 最小就是 11。如果更早使用 CtKVC^{KV}_{\le t},length-tt intermediate 还会保留额外的 feature axis。因此,cache 应该最后 contraction:

(xt(1,d)WDQ(d,dqc)WUQ(h,dqc,dk)WUK(h,dk,dc))(CtKV)(dc,t).\left( \underset{(1,d)}{x_t} \underset{(d,d_{qc})}{W_{DQ}} \underset{(h,d_{qc},d_k)}{W_{UQ}} \underset{(h,d_k,d_c)}{W_{UK}^{\top}} \right) \underset{(d_c,t)}{(C^{KV}_{\le t})^{\top}}.

接下来只需要决定下面这个 four-factor prefix 的 association:

xt(1,d)WDQ(d,dqc)WUQ(h,dqc,dk)WUK(h,dk,dc).\underset{(1,d)}{x_t} \underset{(d,d_{qc})}{W_{DQ}} \underset{(h,d_{qc},d_k)}{W_{UQ}} \underset{(h,d_k,d_c)}{W_{UK}^{\top}}.

下表先列出五种 online binary association,再列出相邻两个或三个 weights 的所有不同 offline precomposition。去掉公共的 final cache contraction 2htdc2htd_c 后,每个 decode step 的 cost 如下;一次性的 offline composition cost 不计入表格。

AssociationCache contraction 之前的 FLOPsDeepSeek-V2
(((xt(1,d)WDQ(d,dqc))WUQ(h,dqc,dk))WUK(h,dk,dc))\left(\left(\left(\underset{(1,d)}{x_t}\underset{(d,d_{qc})}{W_{DQ}}\right)\underset{(h,d_{qc},d_k)}{W_{UQ}}\right)\underset{(h,d_k,d_c)}{W_{UK}^{\top}}\right)2ddqc+2hdqcdk+2hdkdc2dd_{qc}+2hd_{qc}d_k+2hd_kd_c82.8M82.8\text{M}
((xt(1,d)(WDQ(d,dqc)WUQ(h,dqc,dk)))WUK(h,dk,dc))\left(\left(\underset{(1,d)}{x_t}\left(\underset{(d,d_{qc})}{W_{DQ}}\underset{(h,d_{qc},d_k)}{W_{UQ}}\right)\right)\underset{(h,d_k,d_c)}{W_{UK}^{\top}}\right)2hddqcdk+2hddk+2hdkdc2hdd_{qc}d_k+2hdd_k+2hd_kd_c257.9B257.9\text{B}
(xt(1,d)WDQ(d,dqc))(WUQ(h,dqc,dk)WUK(h,dk,dc))\left(\underset{(1,d)}{x_t}\underset{(d,d_{qc})}{W_{DQ}}\right)\left(\underset{(h,d_{qc},d_k)}{W_{UQ}}\underset{(h,d_k,d_c)}{W_{UK}^{\top}}\right)2ddqc+2hdqcdkdc+2hdqcdc2dd_{qc}+2hd_{qc}d_kd_c+2hd_{qc}d_c26.0B26.0\text{B}
xt(1,d)((WDQ(d,dqc)WUQ(h,dqc,dk))WUK(h,dk,dc))\underset{(1,d)}{x_t}\left(\left(\underset{(d,d_{qc})}{W_{DQ}}\underset{(h,d_{qc},d_k)}{W_{UQ}}\right)\underset{(h,d_k,d_c)}{W_{UK}^{\top}}\right)2hddqcdk+2hddkdc+2hddc2hdd_{qc}d_k+2hdd_kd_c+2hdd_c344.3B344.3\text{B}
xt(1,d)(WDQ(d,dqc)(WUQ(h,dqc,dk)WUK(h,dk,dc)))\underset{(1,d)}{x_t}\left(\underset{(d,d_{qc})}{W_{DQ}}\left(\underset{(h,d_{qc},d_k)}{W_{UQ}}\underset{(h,d_k,d_c)}{W_{UK}^{\top}}\right)\right)2hdqcdkdc+2hddqcdc+2hddc2hd_{qc}d_kd_c+2hdd_{qc}d_c+2hdd_c1.06T1.06\text{T}
((xt(1,d)(WDQWUQ)offline(h,d,dk))WUK(h,dk,dc))\left(\left(\underset{(1,d)}{x_t}\underset{(h,d,d_k)}{(W_{DQ}W_{UQ})_{\mathrm{offline}}}\right)\underset{(h,d_k,d_c)}{W_{UK}^{\top}}\right)2hddk+2hdkdc2hdd_k+2hd_kd_c184.5M184.5\text{M}
(xt(1,d)WDQ(d,dqc))(WUQWUK)offline(h,dqc,dc)\left(\underset{(1,d)}{x_t}\underset{(d,d_{qc})}{W_{DQ}}\right)\underset{(h,d_{qc},d_c)}{(W_{UQ}W_{UK}^{\top})_{\mathrm{offline}}}2ddqc+2hdqcdc2dd_{qc}+2hd_{qc}d_c217.1M217.1\text{M}
xt(1,d)(WDQWUQWUK)offline(h,d,dc)\underset{(1,d)}{x_t}\underset{(h,d,d_c)}{(W_{DQ}W_{UQ}W_{UK}^{\top})_{\mathrm{offline}}}2hddc2hdd_c671.1M671.1\text{M}

代入 DeepSeek-V2 dimensions,即使允许所有 offline precomposition,第一行仍然最小。它的直觉也最直接:先用 singleton xtx_t 消去 dd,之后 length-11 axis 会依次经过 dqcd_{qc}dkd_kdcd_c。因此选择从左到右的顺序,总 cost 是

Fscore,Qcomp=2ddqcxtWDQ+2hdqcdkctQWUQ+2hdkdcqtWUK+2htdcqabs(CtKV).F_{\mathrm{score,Qcomp}} =\underbrace{2dd_{qc}}_{x_tW_{DQ}} +\underbrace{2hd_{qc}d_k}_{c_t^QW_{UQ}} +\underbrace{2hd_kd_c}_{q_tW_{UK}^{\top}} +\underbrace{2htd_c}_{q_{\mathrm{abs}}(C^{KV}_{\le t})^{\top}}.

不使用 Q compression 时,只需要删掉 WDQ,WUQW_{DQ},W_{UQ},换成 WQW_Q。Factors 更少,因此相同的 association argument 留给读者;最终的最小值是

Fscore,noQ=2hddkxtWQ+2hdkdcqtWUK+2htdcqabs(CtKV).F_{\mathrm{score,noQ}} =\underbrace{2hdd_k}_{x_tW_Q} +\underbrace{2hd_kd_c}_{q_tW_{UK}^{\top}} +\underbrace{2htd_c}_{q_{\mathrm{abs}}(C^{KV}_{\le t})^{\top}}.

Value 和 output chain。 完整的 value-side chain 是

yt(1,d)=pt(h,1,t)CtKV(t,dc)WUV(h,dc,dv)Wo(h,dv,d).\underset{(1,d)}{y_t} = \underset{(h,1,t)}{p_t} \underset{(t,d_c)}{C^{KV}_{\le t}} \underset{(h,d_c,d_v)}{W_{UV}} \underset{(h,d_v,d)}{W_o}.

先看 tt 的系数。用 ptp_t contraction CtKVC^{KV}_{\le t},会通过 singleton query axis 立刻消掉 cache-length axis,cost 是

2htdc.2htd_c.

这是 length-tt term 能达到的最小值。如果先用 WUVW_{UV} 展开 CtKVC^{KV}_{\le t},仅这一步就需要 2htdcdv2htd_cd_v。因此应该先计算 ptCtKVp_tC^{KV}_{\le t},得到 oc:(h,1,dc)o_c:(h,1,d_c)

接下来只需要决定下面这个 three-factor chain 的 association:

oc(h,1,dc)WUV(h,dc,dv)Wo(h,dv,d).\underset{(h,1,d_c)}{o_c} \underset{(h,d_c,d_v)}{W_{UV}} \underset{(h,d_v,d)}{W_o}.

去掉公共的第一步 contraction 2htdc2htd_c 后,下表包含两种 online association 和 offline-precomposed variant;后者一次性的 composition cost 不计入表格:

AssociationptCtKVp_tC^{KV}_{\le t} 之后的 FLOPsDeepSeek-V2
((oc(h,1,dc)WUV(h,dc,dv))Wo(h,dv,d))\left(\left(\underset{(h,1,d_c)}{o_c}\underset{(h,d_c,d_v)}{W_{UV}}\right)\underset{(h,d_v,d)}{W_o}\right)2hdcdv+2hdvd2hd_cd_v+2hd_vd184.5M184.5\text{M}
oc(h,1,dc)(WUV(h,dc,dv)Wo(h,dv,d))\underset{(h,1,d_c)}{o_c}\left(\underset{(h,d_c,d_v)}{W_{UV}}\underset{(h,d_v,d)}{W_o}\right)2hdcdvd+2hdcd2hd_cd_vd+2hd_cd86.6B86.6\text{B}
oc(h,1,dc)(WUVWo)offline(h,dc,d)\underset{(h,1,d_c)}{o_c}\underset{(h,d_c,d)}{(W_{UV}W_o)_{\mathrm{offline}}}2hdcd2hd_cd671.1M671.1\text{M}

第一行仍然最小。这里仍然是把 length-11 axis 留在外面,从左到右计算更便宜。完整的 value-side cost 是

Fvalue=2htdcptCtKV+2hdcdvocWUV+2hdvdoWo.F_{\mathrm{value}} =\underbrace{2htd_c}_{p_tC^{KV}_{\le t}} +\underbrace{2hd_cd_v}_{o_cW_{UV}} +\underbrace{2hd_vd}_{oW_o}.

合并这些 association 后,完整的 cached-decode cost 是

FMLA,decode,cached,noQ=2ddc+2hddkcurrent-token projections+2hdkdc+4htdcscore absorption 和 latent reductions+2hdcdv+2hdvdvalue 和 output projections,FMLA,decode,cached,Qcomp=2ddc+2ddqc+2hdqcdkcurrent-token projections+2hdkdc+4htdcscore absorption 和 latent reductions+2hdcdv+2hdvdvalue 和 output projections.\begin{aligned} F_{\mathrm{MLA,decode,cached,noQ}} &=\underbrace{2dd_c+2hdd_k}_{\text{current-token projections}} +\underbrace{2hd_kd_c+4htd_c}_{\text{score absorption 和 latent reductions}} +\underbrace{2hd_cd_v+2hd_vd}_{\text{value 和 output projections}}, \\ F_{\mathrm{MLA,decode,cached,Qcomp}} &=\underbrace{2dd_c+2dd_{qc}+2hd_{qc}d_k}_{\text{current-token projections}} +\underbrace{2hd_kd_c+4htd_c}_{\text{score absorption 和 latent reductions}} +\underbrace{2hd_cd_v+2hd_vd}_{\text{value 和 output projections}}. \end{aligned}

对比。

MethodCurrent-token setupScore chainValue + output chainCache elements/token(不包括 RoPE branch)DeepSeek-V2 decode FLOPs(不包括 RoPE branch)DeepSeek-V2 cache elements/token(不包括 RoPE branch)
MHA2dh(2dk+dv)2dh(2d_k+d_v)2htdk2htd_k2htdv+2hdvd2htd_v+2hd_vdh(dk+dv)h(d_k+d_v)65,536t+671,088,64065{,}536t+671{,}088{,}64032,76832{,}768
MLA without query compression2ddc+2hddk2dd_c+2hdd_k2hdkdc+2htdc2hd_kd_c+2htd_c2htdc+2hdcdv+2hdvd2htd_c+2hd_cd_v+2hd_vddcd_c262,144t+374,341,632262{,}144t+374{,}341{,}632512512
MLA with query compression2ddc+2ddqc+2hdqcdk2dd_c+2dd_{qc}+2hd_{qc}d_k2hdkdc+2htdc2hd_kd_c+2htd_c2htdc+2hdcdv+2hdvd2htd_c+2hd_cd_v+2hd_vddcd_c262,144t+272,629,760262{,}144t+272{,}629{,}760512512

DeepSeek-V2 attention configuration:d=5120, h=128, dk=dv=128, dc=512, dqc=1536d=5120,\ h=128,\ d_k=d_v=128,\ d_c=512,\ d_{qc}=1536。FLOPs 和 cache values 都不包括 RoPE branch。

DeepSeek-V2 decode FLOPs 和每层 cache 随 cache length 的增长,两者都不包括 RoPE branch

在这个不考虑 RoPE 的对比中,MLA decode FLOPs 随 tt 的增长速度是 MHA 的 4×4\times,而 cache 只有 MHA 的 1/641/64,cached elements 减少了 98.4%98.4\%

4.3 为什么还要压缩 Q?

上面对两种 MLA variant 的比较,构成了一个关于 Q compression 的直接 ablation。由于它只 factorize query projection,其影响主要落在 query-side activation 和 projection FLOPs 上,而不是 cache path:

  • Training activation memory. 我们不知道 DeepSeek 的内部实现,但我们猜测可以通过 recomputation 和 kernel fusion 来节省 activation memory。
  • Prefill 和 training FLOPs.dqc(d+hdk)<hddkd_{qc}(d+hd_k)<hdd_k 时,factorized projection 更便宜,这一点已经体现在第 4.1 节的对比中。
  • Decode. 只有当前 token 的 query projection 发生变化;latent cache 和随 tt 增长的 attention 项保持不变。因此,与 KV compression 相比,Q compression 对 decode 的贡献很少。

5. 如何让 MLA 支持 RoPE

5.1 为什么 RoPE 与 MLA Compression 不兼容

Takeaway. 没有 RoPE 时,WUKW_{UK} 可以吸收到 query 一侧,得到一个对所有 key positions 复用的 latent query。加入 RoPE 后,WUKW_{UK} 前面多了随 key position 变化的 RuR_u;同一个 absorbed query 因而无法复用于所有 keys。

先直接写出带 RoPE 的 QKQK^\top。固定一个 head ss,令 qs,tC:(1,dk)q_{s,t}^C:(1,d_k) 是 position tt 的 content query,cuKV:(1,dc)c_u^{KV}:(1,d_c) 是 position uu 的 KV latent,并令 WUK,s:(dc,dk)W_{UK,s}:(d_c,d_k)。重构出来的 content key 是 ks,uC=cuKVWUK,sk_{s,u}^C=c_u^{KV}W_{UK,s}。按照第 3 节的 row-vector RoPE convention,query 和 key 分别右乘 RtR_t^\topRuR_u^\top,因此这一对 positions 的 score 是

Ss,t,uRoPE=(qs,tCRt)((cuKVWUK,s)Ru)=qs,tCRtRuWUK,s(cuKV).\begin{aligned} S^{\mathrm{RoPE}}_{s,t,u} &= \left(q_{s,t}^C R_t^\top\right) \left((c_u^{KV}W_{UK,s})R_u^\top\right)^\top \\ &= q_{s,t}^C R_t^\top R_u W_{UK,s}^\top (c_u^{KV})^\top. \end{aligned}

问题只在于这条矩阵乘法链应该怎样加括号。先去掉 RoPE,看普通 MLA 为什么可以做 key absorption。用 TT 表示 query positions 的数量、UU 表示 key positions 的数量,完整的 score chain 是

SC(h,T,U)=QC(h,T,dk)WUK(h,dk,dc)(CKV)(dc,U).\underset{(h,T,U)}{S^C} = \underset{(h,T,d_k)}{Q^C} \underset{(h,d_k,d_c)}{W_{UK}^\top} \underset{(d_c,U)}{(C^{KV})^\top}.

一种顺序是先从 latent 重构所有 keys,再计算 attention:

QC(h,T,dk)(WUK(h,dk,dc)(CKV)(dc,U))=QC(h,T,dk)(KC)(h,dk,U).\underset{(h,T,d_k)}{Q^C} \left( \underset{(h,d_k,d_c)}{W_{UK}^\top} \underset{(d_c,U)}{(C^{KV})^\top} \right) = \underset{(h,T,d_k)}{Q^C} \underset{(h,d_k,U)}{(K^C)^\top}.

另一种顺序是先把 WUKW_{UK} 吸收到 query 一侧:

(QC(h,T,dk)WUK(h,dk,dc))(CKV)(dc,U)=Q^C(h,T,dc)(CKV)(dc,U).\left( \underset{(h,T,d_k)}{Q^C} \underset{(h,d_k,d_c)}{W_{UK}^\top} \right) \underset{(d_c,U)}{(C^{KV})^\top} = \underset{(h,T,d_c)}{\widehat Q^C} \underset{(d_c,U)}{(C^{KV})^\top}.

WUKW_{UK} 是固定 weight,所以 absorbed query Q^C=QCWUK\widehat Q^C=Q^CW_{UK}^\top 与 key position 无关。每个 query 只需要计算一次这个 (h,T,dc)(h,T,d_c) intermediate,随后就能与所有 cached latents 相乘。这正是 key absorption。

现在把 RoPE 放回去。上面的完整 score 可以按“先重构并旋转 key”的顺序计算:

(qs,tCRt)(Ru(WUK,s(cuKV))).\left(q_{s,t}^C R_t^\top\right) \left( R_u \left(W_{UK,s}^\top(c_u^{KV})^\top\right) \right).

也可以尝试像刚才一样,把 key reconstruction 之前的 factors 全部放进左边的括号:

(qs,tCRtRuWUK,s)希望吸收到 query 一侧的部分(cuKV).\underbrace{\left(q_{s,t}^C R_t^\top R_u W_{UK,s}^\top\right)}_{\text{希望吸收到 query 一侧的部分}} (c_u^{KV})^\top.

区别现在一眼就能看到:这个括号里含有 RuR_u。它会随 key position uu 改变,所以不能再得到一个对所有 keys 复用的 Q^C\widehat Q^C。如果把所有 query 和 key positions 一次展开,这个括号产生的 intermediate shape 不是 (h,T,dc)(h,T,d_c),而是

(h,T,U,dc).(h,T,U,d_c).

其中 UU 是 key-position axis:prefill 时 U=TU=T;一步 decode 时 T=1T=1U=tU=t。换句话说,RoPE 没有让矩阵乘法失效,但它让所谓的“absorbed query”变成每个 (t,u)(t,u) position pair 各自一个。这样既没有消掉 key-length axis,也失去了 absorption 的意义。若 cache 只保存 CKVC^{KV},decode 时就只能重新构造并旋转历史 keys,或显式生成这个 (h,1,t,dc)(h,1,t,d_c) intermediate。

因此,普通 RoPE 和 position-independent key absorption 不能放在同一条 feature path 上。第 3 节的 elementwise RoPE 与这里的 rotation-matrix form 完全等价;改写成 elementwise computation 并不会消除对 key position uu 的依赖,因此同样无法实现 position-independent absorption。第 5.2 节的完整 MLA 会把它们拆开:content branch 保持可 absorption,单独的 positional branch 负责 RoPE。

5.2 MLA 的完整形式:Decoupled Rotary Position Embedding

Takeaway. 完整的 MLA 让 content channel 保持无 rotation、可 absorption,同时用一条独立的低维 RoPE channel 承载 position。Cache 保存 CKVC^{KV} 和一份 shared RoPE key,两条 channel 的 scores 在 softmax 前相加。

第 5.1 节说明了,同一条 path 无法同时支持普通 RoPE 和 position-independent key absorption。因此,完整 MLA 不再要求同一个 representation 同时完成这两件事,而是把 feature space 拆成两条 branch:content branch 不做 rotation,专门保持 compression 和 absorption;positional branch 则不再强求 absorption,专门应用 RoPE、承载位置信息。简单来说,一条负责压缩得干净,另一条负责旋转得干净。


MLA with decoupled RoPE.

drd_r 表示 decoupled RoPE branch 的 per-head dimension,对应 DeepSeek-V2 技术报告中的 dhRd_h^R。这里沿用第 3 节的 rotation tensor,只把 tensor dimension 换成 drd_r

R(r)(T,dr,dr),Rt,:,:(r)=(Rt(r)).\underset{(T,d_r,d_r)}{\mathcal R^{(r)}}, \qquad \mathcal R^{(r)}_{t,:,:}=(R_t^{(r)})^\top.

Query path 是

CQ(T,dqc)=X(T,d)WDQ(d,dqc),\underset{(T,d_{qc})}{C^Q} = \underset{(T,d)}{X} \underset{(d,d_{qc})}{W_{DQ}}, QC(h,T,dk)=CQ(T,dqc)WUQ(h,dqc,dk),QR(h,T,dr)=CQ(T,dqc)WQR(h,dqc,dr),\underset{(h,T,d_k)}{Q^C} = \underset{(T,d_{qc})}{C^Q} \underset{(h,d_{qc},d_k)}{W_{UQ}}, \qquad \underset{(h,T,d_r)}{\overline Q^R} = \underset{(T,d_{qc})}{C^Q} \underset{(h,d_{qc},d_r)}{W_{QR}}, QR(h,T,dr)=QR(h,T,dr)R(r)(T,dr,dr).\underset{(h,T,d_r)}{Q^R} = \underset{(h,T,d_r)}{\overline Q^R} \underset{(T,d_r,d_r)}{\mathcal R^{(r)}}.

Query-side RoPE operation 可以保持不变。但 key side 不同:如果 RoPE key 以 compressed form 存储,每个 decode step 都必须重新展开并旋转全部 tt 个历史 keys,产生大量重复计算。因此,完整 MLA 直接缓存 materialized KRK^R,而不是 compressed positional-key latent。

Key-value path 是

CKV(T,dc)=X(T,d)WDKV(d,dc),\underset{(T,d_c)}{C^{KV}} = \underset{(T,d)}{X} \underset{(d,d_c)}{W_{DKV}}, KC(h,T,dk)=CKV(T,dc)WUK(h,dc,dk),V(h,T,dv)=CKV(T,dc)WUV(h,dc,dv),\underset{(h,T,d_k)}{K^C} = \underset{(T,d_c)}{C^{KV}} \underset{(h,d_c,d_k)}{W_{UK}}, \qquad \underset{(h,T,d_v)}{V} = \underset{(T,d_c)}{C^{KV}} \underset{(h,d_c,d_v)}{W_{UV}}, KR(T,dr)=X(T,d)WKR(d,dr),KR(T,dr)=KR(T,dr)R(r)(T,dr,dr).\underset{(T,d_r)}{\overline K^R} = \underset{(T,d)}{X} \underset{(d,d_r)}{W_{KR}}, \qquad \underset{(T,d_r)}{K^R} = \underset{(T,d_r)}{\overline K^R} \underset{(T,d_r,d_r)}{\mathcal R^{(r)}}.

如果每个 head 都保存一份 positional key,每个 token 会额外增加 hdrhd_r 个 cache elements。完整 MLA 改为让所有 hh 个 query heads 共享同一个 KR:(T,dr)K^R:(T,d_r)。因此 RoPE branch 采用的是 MQA-style 结构:positional query 仍然是 per-head 的 QR:(h,T,dr)Q^R:(h,T,d_r),但 positional key 只有一个 shared head。这样额外 cache 保持为每个 token drd_r 个 elements。

完整的 per-head query 和 key,是沿最后一个 tensor dimension 的 concat:

Q(h,T,dk+dr)=[QC(h,T,dk);QR(h,T,dr)],\underset{(h,T,d_k+d_r)}{Q} = \left[ \underset{(h,T,d_k)}{Q^C}; \underset{(h,T,d_r)}{Q^R} \right], K(h,T,dk+dr)=[KC(h,T,dk);KR(T,dr) broadcast over h].\underset{(h,T,d_k+d_r)}{K} = \left[ \underset{(h,T,d_k)}{K^C}; \underset{(T,d_r)}{K^R} \text{ broadcast over }h \right].

为什么 score 会变成两个部分相加?Concat 会把最后一个 tensor dimension 划分成两个互不重叠的 blocks。对于 head ss、query position tt 和 key position uu,完整 dot product 是

Ss,t,u==1dk+drQs,t,Ks,u,=a=1dkQs,t,aCKs,u,aC+b=1drQs,t,bRKu,bR=Ss,t,uC+Ss,t,uR.\begin{aligned} S_{s,t,u} &= \sum_{\ell=1}^{d_k+d_r}Q_{s,t,\ell}K_{s,u,\ell} \\ &= \sum_{a=1}^{d_k}Q^C_{s,t,a}K^C_{s,u,a} + \sum_{b=1}^{d_r}Q^R_{s,t,b}K^R_{u,b} \\ &=S^C_{s,t,u}+S^R_{s,t,u}. \end{aligned}

同一个恒等式也可以写成 block tensor multiplication:

S(h,T,T)=[QC(h,T,dk)QR(h,T,dr)][(KC)(h,dk,T)(KR)(dr,T) broadcast over h]=SC(h,T,T)+SR(h,T,T).\underset{(h,T,T)}{S} = \left[ \begin{array}{cc} \underset{(h,T,d_k)}{Q^C} & \underset{(h,T,d_r)}{Q^R} \end{array} \right] \left[ \begin{array}{c} \underset{(h,d_k,T)}{(K^C)^\top} \\ \underset{(d_r,T)}{(K^R)^\top} \text{ broadcast over }h \end{array} \right] = \underset{(h,T,T)}{S^C} + \underset{(h,T,T)}{S^R}.

这里不存在 QC(KR)Q^C(K^R)^\topQR(KC)Q^R(K^C)^\top 这样的 cross terms:content 和 RoPE coordinates 位于不同的 feature blocks,dot product 只会配对同一个 block 内的 coordinates。因此,两条 score channel 可以分别计算:

SC(h,T,T)=QC(h,T,dk)(KC)(h,dk,T),SR(h,T,T)=QR(h,T,dr)(KR)(dr,T),\underset{(h,T,T)}{S^C} = \underset{(h,T,d_k)}{Q^C} \underset{(h,d_k,T)}{(K^C)^\top}, \qquad \underset{(h,T,T)}{S^R} = \underset{(h,T,d_r)}{Q^R} \underset{(d_r,T)}{(K^R)^\top}, P(h,T,T)=softmax ⁣(SC+SRdk+dr),\underset{(h,T,T)}{P} = \operatorname{softmax}\!\left( \frac{S^C+S^R}{\sqrt{d_k+d_r}} \right), O(h,T,dv)=P(h,T,T)V(h,T,dv),Y(T,d)=O(h,T,dv)Wo(h,dv,d).\underset{(h,T,d_v)}{O} = \underset{(h,T,T)}{P} \underset{(h,T,d_v)}{V}, \qquad \underset{(T,d)}{Y} = \underset{(h,T,d_v)}{O} \underset{(h,d_v,d)}{W_o}.

对于 cached decode,content absorption 和 positional branch 给出

score(h,1,t)=(qC(h,1,dk)WUK(h,dk,dc))(CtKV)(dc,t)+qR(h,1,dr)(KtR)(dr,t),\underset{(h,1,t)}{\mathrm{score}} = \left( \underset{(h,1,d_k)}{q^C} \underset{(h,d_k,d_c)}{W_{UK}^\top} \right) \underset{(d_c,t)}{(C_{\le t}^{KV})^\top} + \underset{(h,1,d_r)}{q^R} \underset{(d_r,t)}{(K_{\le t}^R)^\top}, o(h,1,dv)=(p(h,1,t)CtKV(t,dc))WUV(h,dc,dv).\underset{(h,1,d_v)}{o} = \left( \underset{(h,1,t)}{p} \underset{(t,d_c)}{C_{\le t}^{KV}} \right) \underset{(h,d_c,d_v)}{W_{UV}}.

因此 content key 和 value 仍然留在 CKVC^{KV} 后面,只有 shared rotated key KRK^R 被显式缓存。


Cache 补表。

形式每个 token、每层缓存的 tensors元素数DeepSeek-V2
MHAK,VK,Vh(dk+dv)h(d_k+d_v)3276832768
第 4 节中暂时忽略 RoPE 的 MLACKVC^{KV}dcd_c512512
完整 MLACKV,KRC^{KV},K^Rdc+drd_c+d_r576576

对于 DeepSeek-V2,dc=512d_c=512dr=64d_r=64。因此完整 cache 相对 MHA 的缩减是 32768/57656.9×32768/576\approx56.9\times,而不是只计算 RoPE-free latent 时的 64×64\times


FLOPs 补表。

这个补表沿用第 4 节的 accounting,统计 decoupled branch 新增的 tensor contractions。这里也计算显式的 SC+SRS^C+S^R:每个 score element 做一次 addition,因此 prefill 是 hT2hT^2 FLOPs,cached decode 是 htht FLOPs。RoPE rotation 本身的 structured elementwise operations 不进入这个对比。

阶段新增 projections新增 score computation新增 FLOPs 合计DeepSeek-V2
Prefill / training2Thdqcdr+2Tddr2Thd_{qc}d_r+2Tdd_r2hT2dr+hT22hT^2d_r+hT^22Thdqcdr+2Tddr+2hT2dr+hT22Thd_{qc}d_r+2Tdd_r+2hT^2d_r+hT^225,821,184T+16,512T225{,}821{,}184T+16{,}512T^2
Cached decode2hdqcdr+2ddr2hd_{qc}d_r+2dd_r2htdr+ht2htd_r+ht2hdqcdr+2ddr+2htdr+ht2hd_{qc}d_r+2dd_r+2htd_r+ht25,821,184+16,512t25{,}821{,}184+16{,}512t

对于 cached decode,tt 的 coefficient 现在是 4hdc+2hdr+h=278,6564hd_c+2hd_r+h=278{,}656(DeepSeek-V2)。

6. MLA 如何改变 MTP 的 Speculative-Decoding 收益

Takeaway. MHA speculative verification 会复用 KV-cache 数据,并让新增 query 的计算被 memory-bound 的读取掩盖。MLA 缩小了这次读取,并使 attention 向 compute-bound 一侧移动,因此每个新增 query 的边际成本更高。

DeepSeek-V3 Technical Report 将 MTP 作为辅助训练目标。它的实际配置只使用一个 MTP module(D=1D=1),预测一个额外 token。正常 inference 时可以移除这个 module;用于 speculative decoding 时,它可以 draft 一个 token。报告给出的第二个 token acceptance rate 为 85%85\%-90%90\%,并报告了 speculative decoding 的 1.8×1.8\times TPS。

在 MHA 中,多个 verification query 会复用同一份 KV-cache 数据。由于 decode 是强 memory-bound 的,它们带来的额外 attention 计算可以被 cache 读取的时间掩盖:memory bandwidth 已经成为 bottleneck,而 compute units 仍有余量。MLA 改变了这个平衡。它用 CKVC^{KV} 和 decoupled RoPE key 替代完整的 per-head KV cache,同时 absorption 又增加了围绕 CKVC^{KV} 的 contraction。cache 读取变小,attention 向 compute-bound 一侧移动,能够掩盖额外计算的余量也随之减少。这就是 MLA 会削弱一部分 MTP speculative-decoding 收益的原因。

贡献与致谢

本文以 NonLinear-1 的名义完成。Place holder 参与了推导与写作。感谢 Place holder 提供的讨论与反馈。


如何引用本文。

如果您需要引用本文,请参考:

NonLinear-1. (Jul. 8, 2026). 《MLA, dim by dim》[Blog post]. Retrieved from https://nonlinear1.com/zh/posts/mla-dim-by-dim

@online{nonlinear1-mla-dim-by-dim,
        title={MLA, dim by dim},
        author={{NonLinear-1}},
        year={2026},
        month={July},
        url={\url{https://nonlinear1.com/zh/posts/mla-dim-by-dim}},
}