1. Opening
本文会一直使用同一个视角:把 attention 写成 tensor contraction,并显式保留每一个 dimension。在这个视角下,MLA 里一些看似分散的现象,可以被更加深入和统一地理解。具体来说,可以得到几个有趣的结论:
在 prefill 中,MLA 减少了 linear projections 的 FLOPs,但不能过早 absorption:quadratic core-attention 项仍然应该在 d k , d v d_k,d_v d k , d v 上计算,而不是 d c d_c d c 。
在 decode 中,最优顺序会改变,因为 cache 已经存在。吸收 key/value projections 后,长度为 t t t 的 cache computation 可以经过 latent state,而不是重构出来的 per-head K , V K,V K , 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 , d k ) = X ( T , d ) W Q ( h , d , d k ) ⟺ e i n s u m ( "td,sdk->stk" , X , W Q ) . \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). ( h , T , d k ) Q = ( T , d ) X ( h , d , d k ) W Q ⟺ einsum ( "td,sdk->stk" , X , W Q ) .
S ( h , T , T ) = Q ( h , T , d k ) K ⊤ ( h , d k , T ) ⟺ e i n s u m ( "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). ( h , T , T ) S = ( h , T , d k ) Q ( h , d k , T ) K ⊤ ⟺ einsum ( "stk,suk->stu" , Q , K ) .
一般形式。 对于一个 output axes 为 B 1 , … , B r , m , n B_1,\ldots,B_r,m,n B 1 , … , B r , m , n 、contracted axis 为 k k k 的 batched contraction,FLOPs 是:
F L O P s = 2 ⋅ B 1 ⋯ B r ⋅ m ⋅ n ⋅ k . \mathrm{FLOPs}
\;=\;
2 \cdot B_1 \cdots B_r \cdot m \cdot n \cdot k . FLOPs = 2 ⋅ B 1 ⋯ B r ⋅ m ⋅ n ⋅ k .
换句话说,把所有 output dimensions 乘起来,再乘 contracted dimension,最后乘以 multiply-add 的 2 2 2 (这是标准的 GEMM 计数约定:一次 fused multiply-add 包含一次乘法和一次加法,计作 2 FLOPs)。因此 Q = X W Q Q=XW_Q Q = X W Q 的 cost 是 2 h T d d k 2hTdd_k 2 h T d d k ,而 S = Q K ⊤ S=QK^\top S = Q K ⊤ 的 cost 是 2 h T 2 d k 2hT^2d_k 2 h T 2 d k 。
3. Notation and Preliminaries
Takeaway. 本节通过直接在公式上标注 shape 来约定符号,并给出 MHA、MLA 和后文所需 RoPE rotation 的基本介绍。
关于从 MHA/MQA/GQA 到 MLA 的背景,可以参考 苏剑林的笔记 。本文不展开比较这些变体;这里只把 MHA 当作干净的代数 reference,把 MLA 当作主要推导对象,并介绍后文 absorption 讨论所需的 RoPE 恒等式。
为了让推导保持简洁,我们使用以下约定:
逐元素的 attention score 统一写成 S s , t , u S_{s,t,u} S s , t , u :h h h 是 head 总数,s = 1 , … , h s=1,\ldots,h s = 1 , … , h 是具体的 head,t t t 是 query position,u u u 是 key position。也就是说,score matrix 的行对应 t t t ,列对应 u u u 。
在第 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 增长。
A ⊤ A^\top A ⊤ 表示只交换最后两维,batch axes 不动。
MHA. Reference 形式是:
Q ( h , T , d k ) = X ( T , d ) W Q ( h , d , d k ) , K ( h , T , d k ) = X ( T , d ) W K ( h , d , d k ) , V ( h , T , d v ) = X ( T , d ) W V ( h , d , d v ) . \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}. ( h , T , d k ) Q = ( T , d ) X ( h , d , d k ) W Q , ( h , T , d k ) K = ( T , d ) X ( h , d , d k ) W K , ( h , T , d v ) V = ( T , d ) X ( h , d , d v ) W V .
然后 attention 是:
S ( h , T , T ) = Q ( h , T , d k ) K ⊤ ( h , d k , T ) , P ( h , T , T ) = softmax ( S d k ) , \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), ( h , T , T ) S = ( h , T , d k ) Q ( h , d k , T ) K ⊤ , ( h , T , T ) P = softmax ( d k S ) ,
O ( h , T , d v ) = P ( h , T , T ) V ( h , T , d v ) , Y ( T , d ) = O ( h , T , d v ) W o ( h , d v , 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}. ( h , T , d v ) O = ( h , T , T ) P ( h , T , d v ) V , ( T , d ) Y = ( h , T , d v ) O ( h , d v , d ) W o .
最后一项会同时 contract head axis 和 d v d_v d v axis。这等价于通常的实现:先 concat 所有 heads,再做一次 output projection。因为把 ( h , d v ) (h,d_v) ( h , d v ) flatten 成一个 h d v hd_v h d v axis,并不会改变求和:
Y t , d = ∑ s = 1 h ∑ v = 1 d v O s , t , v ( W o ) s , v , d = ∑ j = 1 h d v ( O c o n c a t ) t , j ( W o , c o n c a t ) 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}. Y t , d = s = 1 ∑ h v = 1 ∑ d v O s , t , v ( W o ) s , v , d = j = 1 ∑ h d v ( O concat ) t , j ( W o , concat ) j , d .
两种写法的 FLOPs 也相同。在 tensor contraction 形式中,output axes 是 T T T 和 d d d ,contracted axes 是 h h h 和 d v d_v d v ,所以 FLOPs 是 2 T d h d v 2Tdh d_v 2 T d h d v 。Concat 之后,矩阵乘的 shapes 是 ( T , h d v ) ( h d v , d ) (T,hd_v)(hd_v,d) ( T , h d v ) ( h d v , d ) ,所以 FLOPs 是 2 T ( h d v ) d = 2 T h d v d 2T(hd_v)d=2Thd_vd 2 T ( h d v ) d = 2 T h d v d 。
MLA. KV path 把 per-head 的 K , V K,V K , V projection 换成一个 shared latent。为了让推导保持简洁,在第 5 节之前,我们暂时不考虑 MLA 中的 RoPE:
C K V ( T , d c ) = X ( T , d ) W D K V ( d , d c ) , K ( h , T , d k ) = C K V ( T , d c ) W U K ( h , d c , d k ) , \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}}, ( T , d c ) C K V = ( T , d ) X ( d , d c ) W D K V , ( h , T , d k ) K = ( T , d c ) C K V ( h , d c , d k ) W U K ,
V ( h , T , d v ) = C K V ( T , d c ) W U V ( h , d c , d v ) . \underset{(h,T,d_v)}{V}
=
\underset{(T,d_c)}{C^{KV}}
\underset{(h,d_c,d_v)}{W_{UV}}. ( h , T , d v ) V = ( T , d c ) C K V ( h , d c , d v ) W U V .
DeepSeek-V2 论文 声称,query compression 可以降低训练时的 activation memory;但 DeepSeek-V2-Lite 不压缩 query。这正好提供了一个自然的对照:先推导 no-Q MLA,再加入 Q compression,单独看它改变了什么。
不压缩 Q。
Q ( h , T , d k ) = X ( T , d ) W Q ( h , d , d k ) . \underset{(h,T,d_k)}{Q}
=
\underset{(T,d)}{X}
\underset{(h,d,d_k)}{W_Q}. ( h , T , d k ) Q = ( T , d ) X ( h , d , d k ) W Q .
Q compression。
C Q ( T , d q c ) = X ( T , d ) W D Q ( d , d q c ) , Q ( h , T , d k ) = C Q ( T , d q c ) W U Q ( h , d q c , d k ) . \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}}. ( T , d q c ) C Q = ( T , d ) X ( d , d q c ) W D Q , ( h , T , d k ) Q = ( T , d q c ) C Q ( h , d q c , d k ) W U Q .
RoPE。
在 tensor view 中,RoPE 只是在最后一个 tensor dimension 上插入一个 position-dependent linear map,不改变 tensor shape。因为 RoPE 作用于 Q Q Q 和 K K K ,这里用它们的 tensor dimension d k d_k d k 作为 rotation dimension。对于偶数 d k d_k d k ,standard rotation matrix 是
R t = diag ( [ cos ( t θ 1 ) − sin ( t θ 1 ) sin ( t θ 1 ) cos ( t θ 1 ) ] , … , [ cos ( t θ d k / 2 ) − sin ( t θ d k / 2 ) sin ( t θ d k / 2 ) cos ( t θ d k / 2 ) ] ) ∈ R d k × d k . 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}. R t = diag ( [ cos ( t θ 1 ) sin ( t θ 1 ) − sin ( t θ 1 ) cos ( t θ 1 ) ] , … , [ cos ( t θ d k /2 ) sin ( t θ d k /2 ) − sin ( t θ d k /2 ) cos ( t θ d k /2 ) ] ) ∈ R d k × d k .
把 RoPE base 记作 b b b ,则 θ j = b − 2 ( j − 1 ) / d k \theta_j=b^{-2(j-1)/d_k} θ j = b − 2 ( j − 1 ) / d k 。RoFormer 使用 b = 10000 b=10000 b = 10000 ;现代语言模型为了支持更长的 context,可能会使用更大的 base。由于本文使用 row vector,把各个 position 的 transposed rotation stack 成一个 tensor:R ( T , d k , d k ) \underset{(T,d_k,d_k)}{\mathcal R} ( T , d k , d k ) R ,其中 R t , : , : = R t ⊤ \mathcal R_{t,:,:}=R_t^\top R t , : , : = R t ⊤ 。
RoPE 和后续 score contraction 可以统一写成
Q ~ ( h , T , d k ) = Q ( h , T , d k ) R ( T , d k , d k ) , K ~ ( h , T , d k ) = K ( h , T , d k ) R ( T , d k , d k ) , S ( h , T , T ) = Q ~ ( h , T , d k ) K ~ ⊤ ( h , d k , 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}. ( h , T , d k ) Q = ( h , T , d k ) Q ( T , d k , d k ) R , ( h , T , d k ) K = ( h , T , d k ) K ( T , d k , d k ) R , ( h , T , T ) S = ( h , T , d k ) Q ( h , d k , T ) K ⊤ .
两个 T T T axis 对应同一个 position,因此输出只保留一个 T T T 。Q Q Q 或 K K K 的 d k d_k d k axis 与 R \mathcal R R 的第一个 d k d_k d k axis contraction,第二个 d k d_k d k axis 保留;R \mathcal R R 在 head axis 上 broadcast。因此 Q , K Q,K Q , K 的 ( h , T , d k ) (h,T,d_k) ( h , T , d k ) shape 不变。又因为 R t ⊤ R u = R u − t R_t^\top R_u=R_{u-t} R t ⊤ R u = R u − t ,
S s , t , u = Q s , t , : R t ⊤ R u K s , u , : ⊤ = Q s , t , : R u − t K s , 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, S s , t , u = Q s , t , : R t ⊤ R u K s , u , : ⊤ = Q s , t , : R u − t K s , u , : ⊤ ,
所以每次 rotation 编码 absolute position,而 score 只依赖 relative position。
实际的 elementwise 计算。 上面的 rotation tensor 是代数表示;implementation 不需要构造 dense d k × d k d_k\times d_k d k × d k matrix。令 m = d k / 2 m=d_k/2 m = d k /2 。对于任意 q ∈ R d k q\in\mathbb R^{d_k} q ∈ R d k ,定义 adjacent-pair coordinates
q j , 0 : = q 2 j − 1 , q j , 1 : = q 2 j , j = 1 , … , m . q_{j,0}:=q_{2j-1},
\qquad
q_{j,1}:=q_{2j},
\qquad
j=1,\ldots,m. q j , 0 := q 2 j − 1 , q j , 1 := q 2 j , j = 1 , … , m .
沿用前面的 θ j \theta_j θ j ,对应的 phase tensor 是
Θ ^ t , 2 j − 1 = Θ ^ t , 2 j = t θ j , j = 1 , … , m , Θ ^ ( T , d k ) . \widehat\Theta_{t,2j-1}
=
\widehat\Theta_{t,2j}
=
t\theta_j,
\qquad
j=1,\ldots,m,
\qquad
\underset{(T,d_k)}{\widehat\Theta}. Θ t , 2 j − 1 = Θ t , 2 j = t θ j , j = 1 , … , m , ( T , d k ) Θ .
然后直接计算
Q ~ ( h , T , d k ) = Q ( h , T , d k ) ⊙ cos Θ ^ ( T , d k ) + rotate_half ( Q ( h , T , d k ) ) ⊙ sin Θ ^ ( T , d k ) , \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}, ( h , T , d k ) Q = ( h , T , d k ) Q ⊙ cos ( T , d k ) Θ + rotate_half ( ( h , T , d k ) Q ) ⊙ sin ( T , d k ) Θ ,
K ~ ( h , T , d k ) = K ( h , T , d k ) ⊙ cos Θ ^ ( T , d k ) + rotate_half ( K ( h , T , d k ) ) ⊙ sin Θ ^ ( T , d k ) . \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}. ( h , T , d k ) K = ( h , T , d k ) K ⊙ cos ( T , d k ) Θ + rotate_half ( ( h , T , d k ) K ) ⊙ sin ( T , d k ) Θ .
这里 ⊙ \odot ⊙ 是 elementwise multiplication;cosine 和 sine tensors 在 head axis 上 broadcast。在上面的 pair notation 中,rotate_half 定义为
rotate_half ( q ) j , 0 = − q j , 1 , rotate_half ( q ) j , 1 = q j , 0 . \operatorname{rotate\_half}(q)_{j,0}=-q_{j,1},
\qquad
\operatorname{rotate\_half}(q)_{j,1}=q_{j,0}. rotate_half ( q ) j , 0 = − q j , 1 , 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。它没有减少 T 2 T^2 T 2 core-attention FLOPs,因为 Q K ⊤ QK^\top Q K ⊤ 和 P V PV P V 的 contraction 仍然分别沿 d k d_k d k 和 d v d_v d v 进行。
MHA。
对于 MHA prefill,完整的 score 表达式是:
从左到右读这条 chain,可以把它看成是在选择何时产生更小的 representation。Weights 会在任何跨越全部 T T T 个 query-key pairs 的操作之前,把 model dimension d d d 降到 per-head key dimension d k d_k d k 。
S = Q K ⊤ = X ( T , d ) W Q ( h , d , d k ) W K ⊤ ( h , d k , 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}. S = Q K ⊤ = ( T , d ) X ( h , d , d k ) W Q ( h , d k , d ) W K ⊤ ( d , T ) X ⊤ .
从 FLOPs 的角度,关键是让占主导的 T 2 T^2 T 2 项系数尽可能小。由于 d k < d d_k<d d k < d ,计算顺序应该是先构造 Q = X W Q Q=XW_Q Q = X W Q 和 K = X W K K=XW_K K = X W K ,再计算 S = Q K ⊤ S=QK^\top S = Q K ⊤ :
X ( T , d ) W Q ( h , d , d k ) → Q ( h , T , d k ) , F L O P s = 2 T h d d k , X ( T , d ) W K ( h , d , d k ) → K ( h , T , d k ) , F L O P s = 2 T h d d k , Q ( h , T , d k ) K ⊤ ( h , d k , T ) → S ( h , T , T ) , F L O P s = 2 h T 2 d k . \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} ( T , d ) X ( h , d , d k ) W Q ( T , d ) X ( h , d , d k ) W K ( h , T , d k ) Q ( h , d k , T ) K ⊤ → ( h , T , d k ) Q , → ( h , T , d k ) K , → ( h , T , T ) S , FLOPs FLOPs FLOPs = 2 T h d d k , = 2 T h d d k , = 2 h T 2 d k .
前两项对 T T T 是 linear 的,因为每个 token 都独立做 projection。只有最后一次 contraction 会把每个 query 与每个 key 连接起来,因此它是 score path 中唯一的 T 2 T^2 T 2 项。
完整的 value 和 output chain 是:
P ( h , T , T ) X ( T , d ) W V ( h , d , d v ) W o ( h , d v , 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}. ( h , T , T ) P ( T , d ) X ( h , d , d v ) W V ( h , d v , d ) W o → ( T , d ) Y .
Value/output order:使用 ( P ( X W V ) ) W o (P(XW_V))W_o ( P ( X W V )) W o ,而不是 ( ( P X ) W V ) W o ((PX)W_V)W_o (( P X ) W V ) W o 或 P ( X ( W V W o ) ) P(X(W_VW_o)) P ( X ( W V W o )) 。
这里的计算顺序问题是:与 P P P 做 quadratic multiplication 时,哪一个 tensor dimension 会被带进去。先构造 V V V ,可以让 attention 落在 d v d_v d v 上,而不是让更宽的 model dimension d d d 穿过所有 token pairs。
标准顺序是 X W V → V XW_V\rightarrow V X W V → V 、P V → O PV\rightarrow O P V → O 、O W o → Y OW_o\rightarrow Y O W o → Y :
F ( P ( X W V ) ) W o = 2 T h d d v ⏟ X W V + 2 h T 2 d v ⏟ P V + 2 T h d v d ⏟ O W o . F_{(P(XW_V))W_o}
=
\underbrace{2Thdd_v}_{XW_V}
+\underbrace{2hT^2d_v}_{PV}
+\underbrace{2Thd_vd}_{OW_o}. F ( P ( X W V )) W o = X W V 2 T h d d v + P V 2 h T 2 d v + O W o 2 T h d v d .
先计算 P X PX P X 时:
F ( ( P X ) W V ) W o = 2 h T 2 d ⏟ P X + 2 T h d d v ⏟ ( P X ) W V + 2 T h d v d ⏟ O W o . F_{((PX)W_V)W_o}
=
\underbrace{2hT^2d}_{PX}
+\underbrace{2Thdd_v}_{(PX)W_V}
+\underbrace{2Thd_vd}_{OW_o}. F (( P X ) W V ) W o = P X 2 h T 2 d + ( P X ) W V 2 T h d d v + O W o 2 T h d v d .
预先组合 W V W o W_VW_o W V W o 时:
F P ( X ( W V W o ) ) = 2 h d 2 d v ⏟ W V W o + 2 h T d 2 ⏟ X ( W V W o ) + 2 h T 2 d ⏟ P [ X ( W V W o ) ] . 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)]}. F P ( X ( W V W o )) = W V W o 2 h d 2 d v + X ( W V W o ) 2 h T d 2 + P [ X ( W V W o )] 2 h T 2 d .
后两种顺序都会让 quadratic contraction 落在 d d d 而不是 d v d_v d v 上。由于 d v ≪ d d_v\ll d d v ≪ d ,标准顺序更便宜:2 h T 2 d v 2hT^2d_v 2 h T 2 d v ,而不是 2 h T 2 d 2hT^2d 2 h T 2 d 。
把这些项合在一起,dense MHA prefill FLOPs 是:
F M H A , p r e f i l l = 2 T d h ( 2 d k + d v ) ⏟ X W Q , X W K , X W V + 2 h T 2 ( d k + d v ) ⏟ Q K ⊤ , P V + 2 T h d v d ⏟ O W o . 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}. F MHA , prefill = X W Q , X W K , X W V 2 T d h ( 2 d k + d v ) + Q K ⊤ , P V 2 h T 2 ( d k + d v ) + O W o 2 T h d v d .
MLA。
对于 MLA prefill,首先计算 shared KV latent C K V = X W D K V C^{KV}=XW_{DKV} C K V = X W D K V 。不使用 query compression 时,完整的 score 表达式是:
S = X ( T , d ) W Q ( h , d , d k ) W U K ⊤ ( h , d k , d c ) W D K V ⊤ ( d c , 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}. S = ( T , d ) X ( h , d , d k ) W Q ( h , d k , d c ) W U K ⊤ ( d c , d ) W D K V ⊤ ( d , T ) X ⊤ .
使用 query compression 时,它变成:
S = X ( T , d ) W D Q ( d , d q c ) W U Q ( h , d q c , d k ) W U K ⊤ ( h , d k , d c ) W D K V ⊤ ( d c , 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}. S = ( T , d ) X ( d , d q c ) W D Q ( h , d q c , d k ) W U Q ( h , d k , d c ) W U K ⊤ ( d c , d ) W D K V ⊤ ( d , T ) X ⊤ .
无论以哪种顺序计算 score,output shape 都是 ( h , T , T ) (h,T,T) ( h , T , T ) ,所以 h T 2 hT^2 h T 2 是固定因子;我们只需要比较 contracted tensor dimension。因此应该让这个 dimension 是 d k d_k d k (在 DeepSeek-V2 中,d k = 128 < d c = 512 < d = 5120 d_k=128<d_c=512<d=5120 d k = 128 < d c = 512 < d = 5120 )。于是可以把 formula 分成两部分:X W D Q W U Q XW_{DQ}W_{UQ} X W D Q W U Q 和 X W D K V W U K XW_{DKV}W_{UK} X W D K V W U K 。
Q compression 的 projection order:使用 ( X W D Q ) W U Q (XW_{DQ})W_{UQ} ( X W D Q ) W U Q ,而不是 X ( W D Q W U Q ) X(W_{DQ}W_{UQ}) X ( W D Q W U Q ) 。
不使用 Q compression 时,Q = X W Q Q=XW_Q Q = X W Q 与 MHA 相同(FLOPs:2 T h d d k 2Thdd_k 2 T h d d k )。它只有一次 contraction,因此 query side 没有计算顺序可选。下面的比较只在把 W Q W_Q W Q factorize 成 W D Q W U Q W_{DQ}W_{UQ} W D Q W U Q 后才会出现。
使用 Q compression 时,factorized order 先把每个 token 映射到 C Q C^Q C Q ,再通过 W U Q W_{UQ} W U Q 展开 head axis:
X ( T , d ) W D Q ( d , d q c ) → C Q ( T , d q c ) , F L O P s = 2 T d d q c , C Q ( T , d q c ) W U Q ( h , d q c , d k ) → Q ( h , T , d k ) , F L O P s = 2 T h d q c d k . \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} ( T , d ) X ( d , d q c ) W D Q ( T , d q c ) C Q ( h , d q c , d k ) W U Q → ( T , d q c ) C Q , → ( h , T , d k ) Q , FLOPs FLOPs = 2 T d d q c , = 2 T h d q c d k .
如果先合并 weights,两次 contraction 是
W D Q ( d , d q c ) W U Q ( h , d q c , d k ) → W Q e f f ( h , d , d k ) , F L O P s = 2 h d d q c d k , X ( T , d ) W Q e f f ( h , d , d k ) → Q ( h , T , d k ) , F L O P s = 2 T h d d k . \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} ( d , d q c ) W D Q ( h , d q c , d k ) W U Q ( T , d ) X ( h , d , d k ) W Q eff → ( h , d , d k ) W Q eff , → ( h , T , d k ) Q , FLOPs FLOPs = 2 h d d q c d k , = 2 T h d d k .
因此,两种顺序的 activation-side cost 分别是
F ( X W D Q ) W U Q = 2 T d d q c + 2 T h d q c d k = 66.1 M T , F X ( W D Q W U Q ) = 2 T h d d k = 167.8 M T (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} F ( X W D Q ) W U Q F X ( W D Q W U Q ) = 2 T d d q c + 2 T h d q c d k = 66.1 M T , = 2 T h d d k = 167.8 M T ( DeepSeek-V2 ) .
其中第二个公式没有计算额外的 weight-precomposition cost 2 h d d q c d k 2hdd_{qc}d_k 2 h d d q c d k 。当
d q c ( d + h d k ) < h d d k d_{qc}(d+hd_k)<hdd_k d q c ( d + h d k ) < h d d k
时,factorized order 更便宜。收益来自在两次 contraction 之间引入 C Q : ( T , d q c ) C^Q:(T,d_{qc}) C Q : ( T , d q c ) ,使 per-head contraction 沿 d q c d_{qc} d q c 而不是 d d d 进行。
同理,对于 key and score order,先计算 C K V = X W D K V C^{KV}=XW_{DKV} C K V = X W D K V ,重构 K = C K V W U K K=C^{KV}W_{UK} K = C K V W U K ,再计算 S = Q K ⊤ S=QK^\top S = Q K ⊤ 。
代入 DeepSeek-V2 的 dimensions,key-projection order ( X W D K V ) W U K (XW_{DKV})W_{UK} ( X W D K V ) W U K 的 cost 是 22.0 M T 22.0\text{M}T 22.0 M T ,而 X ( W D K V W U K ) X(W_{DKV}W_{UK}) X ( W D K V W U K ) 的 cost 是 167.8 M T 167.8\text{M}T 167.8 M T ;后者没有计算一次性的 weight-composition cost。
Value and output order:使用 ( P ( C K V W U V ) ) W o (P(C^{KV}W_{UV}))W_o ( P ( C K V W U V )) W o ,而不是 ( ( P C K V ) W U V ) W o ((PC^{KV})W_{UV})W_o (( P C K V ) W U V ) W o 或 P ( C K V ( W U V W o ) ) P(C^{KV}(W_{UV}W_o)) P ( C K V ( W U V W o )) 。
标准顺序在计算 P V PV P V 之前重构 V V V :
F ( P ( C K V W U V ) ) W o = 2 T h d c d v ⏟ C K V W U V + 2 h T 2 d v ⏟ P V + 2 T h d v d ⏟ O W o . 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}. F ( P ( C K V W U V )) W o = C K V W U V 2 T h d c d v + P V 2 h T 2 d v + O W o 2 T h d v d .
如果先计算 P C K V PC^{KV} P C K V ,quadratic term 会变成 2 h T 2 d c 2hT^2d_c 2 h T 2 d c 。如果先合并 W U V W o W_{UV}W_o W U V W o ,它会变成 2 h T 2 d 2hT^2d 2 h T 2 d 。由于 d v < d c < d d_v<d_c<d d v < d c < d ,先重构 V V V 可以让 quadratic contraction 保持在最小的 tensor dimension d v d_v d v 上。
不使用 query compression 时,dense MLA prefill FLOPs 是:
F M L A , p r e f i l l , n o Q = 2 T h d d k ⏟ Q = X W Q + 2 T d d c ⏟ C K V = X W D K V + 2 T h d c ( d k + d v ) ⏟ K = C K V W U K , V = C K V W U V + 2 h T 2 ( d k + d v ) ⏟ S = Q K ⊤ , O = P V + 2 T h d v d ⏟ Y = O W o . \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} F MLA , prefill , noQ = Q = X W Q 2 T h d d k + C K V = X W D K V 2 T d d c + K = C K V W U K , V = C K V W U V 2 T h d c ( d k + d v ) + S = Q K ⊤ , O = P V 2 h T 2 ( d k + d v ) + Y = O W o 2 T h d v d .
使用 query compression 时:
F M L A , p r e f i l l , Q c o m p = 2 T d d q c ⏟ C Q = X W D Q + 2 T h d q c d k ⏟ Q = C Q W U Q + 2 T d d c ⏟ C K V = X W D K V + 2 T h d c ( d k + d v ) ⏟ K = C K V W U K , V = C K V W U V + 2 h T 2 ( d k + d v ) ⏟ S = Q K ⊤ , O = P V + 2 T h d v d ⏟ Y = O W o . \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} F MLA , prefill , Qcomp = C Q = X W D Q 2 T d d q c + Q = C Q W U Q 2 T h d q c d k + C K V = X W D K V 2 T d d c + K = C K V W U K , V = C K V W U V 2 T h d c ( d k + d v ) + S = Q K ⊤ , O = P V 2 h T 2 ( d k + d v ) + Y = O W o 2 T h d v d .
对比。
Prefill FLOPs 可以分成三部分:linear projections 、core-attention computation 和 final output projection 。下面的表格汇总推导结果并代入 DeepSeek-V2 dimensions;之后的图展示除 RoPE branch 之外的 FLOPs 随 sequence length T T T 的变化。
DeepSeek-V2 attention configuration:d = 5120 , h = 128 , d k = d v = 128 , d c = 512 , d q c = 1536 d=5120,\ h=128,\ d_k=d_v=128,\ d_c=512,\ d_{qc}=1536 d = 5120 , h = 128 , d k = d v = 128 , d c = 512 , d q c = 1536 。
4.2 Decode:Latent Cache 与矩阵吸收
Takeaway. 单步 decode 中,current-query length 是 1 1 1 ,cache length 是 t t t 。更省的顺序就是把这个 1 1 1 留在外面,尽早 contraction 掉其他 dimensions,最后再做 expansion。这样无需为全部 t t t 个 cached tokens 重构完整的 K , V K,V K , V :MLA FLOPs 随 t t t 增长得更快,但读取的 cache data 更少。
MHA。
在 decode 当前位置 t t t 之前,历史 K < t , V < t K_{<t},V_{<t} K < t , V < t 已经在此前的 decode steps 中计算完成。当前 step 只计算新的 q t , k t , v t q_t,k_t,v_t q t , k t , v t ;append k t , v t k_t,v_t k t , v t 之后,这一步读取的 cache 是 K ≤ t , V ≤ t K_{\le t},V_{\le t} K ≤ t , V ≤ t 。跨所有 heads,current-token projection cost 是 2 d h ( 2 d k + d v ) 2dh(2d_k+d_v) 2 d h ( 2 d k + d v ) 。
与 prefill 不同,这里只有一个新的 query,因此 quadratic T 2 T^2 T 2 interaction 消失了。现在增长的因子是 t t t :当前 query 扫过 t t t 个 cached keys,再用得到的 probabilities 合并 t t t 个 cached values。
Score、value 和 output contractions 是:
q ( h , 1 , d k ) K ≤ t ⊤ ( h , d k , t ) → s c o r e ( h , 1 , t ) , F L O P s = 2 h t d k , p ( h , 1 , t ) V ≤ t ( h , t , d v ) → o ( h , 1 , d v ) , F L O P s = 2 h t d v , o ( h , 1 , d v ) W o ( h , d v , d ) → y ( 1 , d ) , F L O P s = 2 h d v d . \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} ( h , 1 , d k ) q ( h , d k , t ) K ≤ t ⊤ ( h , 1 , t ) p ( h , t , d v ) V ≤ t ( h , 1 , d v ) o ( h , d v , d ) W o → ( h , 1 , t ) score , → ( h , 1 , d v ) o , → ( 1 , d ) y , FLOPs FLOPs FLOPs = 2 h t d k , = 2 h t d v , = 2 h d v d .
两个 chain 遵循同一个原则:从 singleton current-token state 出发,选择让每个 intermediate 尽可能小的 contraction order。
把所有 dimensions 显式写出后,完整的 value-output chain 是
y ( 1 , d ) = p ( h , 1 , t ) V ≤ t ( h , t , d v ) W o ( h , d v , 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}. ( 1 , d ) y = ( h , 1 , t ) p ( h , t , d v ) V ≤ t ( h , d v , d ) W o .
两种 association 的 cost 不同:
F ( p V ≤ t ) W o = 2 h t d v ⏟ p V ≤ t : ( h , 1 , d v ) + 2 h d v d ⏟ ( p V ≤ t ) W o : ( 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)}, F ( p V ≤ t ) W o = p V ≤ t : ( h , 1 , d v ) 2 h t d v + ( p V ≤ t ) W o : ( 1 , d ) 2 h d v d ,
F p ( V ≤ t W o ) = 2 h t d v d ⏟ V ≤ t W o : ( h , t , d ) + 2 h t d ⏟ p ( V ≤ t W o ) : ( 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)}. F p ( V ≤ t W o ) = V ≤ t W o : ( h , t , d ) 2 h t d v d + p ( V ≤ t W o ) : ( 1 , d ) 2 h t d .
从左到右计算,会让长度为 1 1 1 的 query axis 始终保持为 output axis,并在 output projection 之前先消去 cache-length axis。另一种顺序则要先把 W o W_o W o 应用于全部 t t t 个 cached values,使 intermediate 同时带着 t t t 和 d d d ,因此计算代价更高。
Score path 也是同一个原则。把 dimensions 显式写出来:
x t ( 1 , d ) W Q ( h , d , d k ) K ≤ t ⊤ ( h , d k , t ) → s c o r e t ( 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}. ( 1 , d ) x t ( h , d , d k ) W Q ( h , d k , t ) K ≤ t ⊤ → ( h , 1 , t ) score t .
把 singleton query axis 1 1 1 保持在外层,在它周围依次 contraction 其他 feature 和 cache axes,而不是先把它们展开。
因此 cached MHA decode FLOPs 是:
F M H A , d e c o d e , c a c h e d = 2 d h ( 2 d k + d v ) ⏟ x t W Q , x t W K , x t W V + 2 h t ( d k + d v ) ⏟ q K ≤ t ⊤ , p V ≤ t + 2 h d v d ⏟ o W o . 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}. F MHA , decode , cached = x t W Q , x t W K , x t W V 2 d h ( 2 d k + d v ) + q K ≤ t ⊤ , p V ≤ t 2 h t ( d k + d v ) + o W o 2 h d v d .
MLA。
对于 MLA decode,cache 存储 shared latent,而不是 full per-head K , V K,V K , V 。在这个单独的 decode step 中,新的 latent 是
x t ( 1 , d ) W D K V ( d , d c ) → c t K V ( 1 , d c ) , F L O P s = 2 d d c , \underset{(1,d)}{x_t}
\underset{(d,d_c)}{W_{DKV}}
\rightarrow
\underset{(1,d_c)}{c_t^{KV}},
\qquad
\mathrm{FLOPs}=2dd_c, ( 1 , d ) x t ( d , d c ) W D K V → ( 1 , d c ) c t K V , FLOPs = 2 d d c ,
append 之后的 cache 是 C ≤ t K V = [ C < t K V ; c t K V ] C^{KV}_{\le t}=[C^{KV}_{<t};c_t^{KV}] C ≤ t K V = [ C < t K V ; c t K V ] ,shape 为 ( t , d c ) (t,d_c) ( t , d c ) 。下面先把完整 contraction chains 写出来,不预先假定 association。
Score chain。 不失一般性,只考虑带 Q compression 的情况。完整 contraction 是
s c o r e t ( h , 1 , t ) = x t ( 1 , d ) W D Q ( d , d q c ) W U Q ( h , d q c , d k ) W U K ⊤ ( h , d k , d c ) ( C ≤ t K V ) ⊤ ( d c , 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}}. ( h , 1 , t ) score t = ( 1 , d ) x t ( d , d q c ) W D Q ( h , d q c , d k ) W U Q ( h , d k , d c ) W U K ⊤ ( d c , t ) ( C ≤ t K V ) ⊤ .
先看 t t t 的系数。Cache 是唯一带有 t t t axis 的 factor。如果先把它左边的 factors 全部 contraction 掉,左侧 intermediate 的 shape 是 ( h , 1 , d c ) (h,1,d_c) ( h , 1 , d c ) ,最后一次 contraction 的 cost 是
2 h t d c . 2htd_c. 2 h t d c .
这已经是 length-t t t term 能达到的最小值:h h h 、t t t 和 d c d_c d c 无法消掉,剩下的 query axis 最小就是 1 1 1 。如果更早使用 C ≤ t K V C^{KV}_{\le t} C ≤ t K V ,length-t t t intermediate 还会保留额外的 feature axis。因此,cache 应该最后 contraction:
( x t ( 1 , d ) W D Q ( d , d q c ) W U Q ( h , d q c , d k ) W U K ⊤ ( h , d k , d c ) ) ( C ≤ t K V ) ⊤ ( d c , 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}}. ( ( 1 , d ) x t ( d , d q c ) W D Q ( h , d q c , d k ) W U Q ( h , d k , d c ) W U K ⊤ ) ( d c , t ) ( C ≤ t K V ) ⊤ .
接下来只需要决定下面这个 four-factor prefix 的 association:
x t ( 1 , d ) W D Q ( d , d q c ) W U Q ( h , d q c , d k ) W U K ⊤ ( h , d k , d c ) . \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}}. ( 1 , d ) x t ( d , d q c ) W D Q ( h , d q c , d k ) W U Q ( h , d k , d c ) W U K ⊤ .
下表先列出五种 online binary association,再列出相邻两个或三个 weights 的所有不同 offline precomposition。去掉公共的 final cache contraction 2 h t d c 2htd_c 2 h t d c 后,每个 decode step 的 cost 如下;一次性的 offline composition cost 不计入表格。
代入 DeepSeek-V2 dimensions,即使允许所有 offline precomposition,第一行仍然最小。它的直觉也最直接:先用 singleton x t x_t x t 消去 d d d ,之后 length-1 1 1 axis 会依次经过 d q c d_{qc} d q c 、d k d_k d k 和 d c d_c d c 。因此选择从左到右的顺序,总 cost 是
F s c o r e , Q c o m p = 2 d d q c ⏟ x t W D Q + 2 h d q c d k ⏟ c t Q W U Q + 2 h d k d c ⏟ q t W U K ⊤ + 2 h t d c ⏟ q a b s ( C ≤ t K V ) ⊤ . 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}}. F score , Qcomp = x t W D Q 2 d d q c + c t Q W U Q 2 h d q c d k + q t W U K ⊤ 2 h d k d c + q abs ( C ≤ t K V ) ⊤ 2 h t d c .
不使用 Q compression 时,只需要删掉 W D Q , W U Q W_{DQ},W_{UQ} W D Q , W U Q ,换成 W Q W_Q W Q 。Factors 更少,因此相同的 association argument 留给读者;最终的最小值是
F s c o r e , n o Q = 2 h d d k ⏟ x t W Q + 2 h d k d c ⏟ q t W U K ⊤ + 2 h t d c ⏟ q a b s ( C ≤ t K V ) ⊤ . 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}}. F score , noQ = x t W Q 2 h d d k + q t W U K ⊤ 2 h d k d c + q abs ( C ≤ t K V ) ⊤ 2 h t d c .
Value 和 output chain。 完整的 value-side chain 是
y t ( 1 , d ) = p t ( h , 1 , t ) C ≤ t K V ( t , d c ) W U V ( h , d c , d v ) W o ( h , d v , 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}. ( 1 , d ) y t = ( h , 1 , t ) p t ( t , d c ) C ≤ t K V ( h , d c , d v ) W U V ( h , d v , d ) W o .
先看 t t t 的系数。用 p t p_t p t contraction C ≤ t K V C^{KV}_{\le t} C ≤ t K V ,会通过 singleton query axis 立刻消掉 cache-length axis,cost 是
2 h t d c . 2htd_c. 2 h t d c .
这是 length-t t t term 能达到的最小值。如果先用 W U V W_{UV} W U V 展开 C ≤ t K V C^{KV}_{\le t} C ≤ t K V ,仅这一步就需要 2 h t d c d v 2htd_cd_v 2 h t d c d v 。因此应该先计算 p t C ≤ t K V p_tC^{KV}_{\le t} p t C ≤ t K V ,得到 o c : ( h , 1 , d c ) o_c:(h,1,d_c) o c : ( h , 1 , d c ) 。
接下来只需要决定下面这个 three-factor chain 的 association:
o c ( h , 1 , d c ) W U V ( h , d c , d v ) W o ( h , d v , d ) . \underset{(h,1,d_c)}{o_c}
\underset{(h,d_c,d_v)}{W_{UV}}
\underset{(h,d_v,d)}{W_o}. ( h , 1 , d c ) o c ( h , d c , d v ) W U V ( h , d v , d ) W o .
去掉公共的第一步 contraction 2 h t d c 2htd_c 2 h t d c 后,下表包含两种 online association 和 offline-precomposed variant;后者一次性的 composition cost 不计入表格:
第一行仍然最小。这里仍然是把 length-1 1 1 axis 留在外面,从左到右计算更便宜。完整的 value-side cost 是
F v a l u e = 2 h t d c ⏟ p t C ≤ t K V + 2 h d c d v ⏟ o c W U V + 2 h d v d ⏟ o W o . F_{\mathrm{value}}
=\underbrace{2htd_c}_{p_tC^{KV}_{\le t}}
+\underbrace{2hd_cd_v}_{o_cW_{UV}}
+\underbrace{2hd_vd}_{oW_o}. F value = p t C ≤ t K V 2 h t d c + o c W U V 2 h d c d v + o W o 2 h d v d .
合并这些 association 后,完整的 cached-decode cost 是
F M L A , d e c o d e , c a c h e d , n o Q = 2 d d c + 2 h d d k ⏟ current-token projections + 2 h d k d c + 4 h t d c ⏟ score absorption 和 latent reductions + 2 h d c d v + 2 h d v d ⏟ value 和 output projections , F M L A , d e c o d e , c a c h e d , Q c o m p = 2 d d c + 2 d d q c + 2 h d q c d k ⏟ current-token projections + 2 h d k d c + 4 h t d c ⏟ score absorption 和 latent reductions + 2 h d c d v + 2 h d v d ⏟ value 和 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} F MLA , decode , cached , noQ F MLA , decode , cached , Qcomp = current-token projections 2 d d c + 2 h d d k + score absorption 和 latent reductions 2 h d k d c + 4 h t d c + value 和 output projections 2 h d c d v + 2 h d v d , = current-token projections 2 d d c + 2 d d q c + 2 h d q c d k + score absorption 和 latent reductions 2 h d k d c + 4 h t d c + value 和 output projections 2 h d c d v + 2 h d v d .
对比。
DeepSeek-V2 attention configuration:d = 5120 , h = 128 , d k = d v = 128 , d c = 512 , d q c = 1536 d=5120,\ h=128,\ d_k=d_v=128,\ d_c=512,\ d_{qc}=1536 d = 5120 , h = 128 , d k = d v = 128 , d c = 512 , d q c = 1536 。FLOPs 和 cache values 都不包括 RoPE branch。
在这个不考虑 RoPE 的对比中,MLA decode FLOPs 随 t t t 的增长速度是 MHA 的 4 × 4\times 4 × ,而 cache 只有 MHA 的 1 / 64 1/64 1/64 ,cached elements 减少了 98.4 % 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. 当 d q c ( d + h d k ) < h d d k d_{qc}(d+hd_k)<hdd_k d q c ( d + h d k ) < h d d k 时,factorized projection 更便宜,这一点已经体现在第 4.1 节的对比中。
Decode. 只有当前 token 的 query projection 发生变化;latent cache 和随 t t t 增长的 attention 项保持不变。因此,与 KV compression 相比,Q compression 对 decode 的贡献很少。
5. 如何让 MLA 支持 RoPE
5.1 为什么 RoPE 与 MLA Compression 不兼容
Takeaway. 没有 RoPE 时,W U K W_{UK} W U K 可以吸收到 query 一侧,得到一个对所有 key positions 复用的 latent query。加入 RoPE 后,W U K W_{UK} W U K 前面多了随 key position 变化的 R u R_u R u ;同一个 absorbed query 因而无法复用于所有 keys。
先直接写出带 RoPE 的 Q K ⊤ QK^\top Q K ⊤ 。固定一个 head s s s ,令 q s , t C : ( 1 , d k ) q_{s,t}^C:(1,d_k) q s , t C : ( 1 , d k ) 是 position t t t 的 content query,c u K V : ( 1 , d c ) c_u^{KV}:(1,d_c) c u K V : ( 1 , d c ) 是 position u u u 的 KV latent,并令 W U K , s : ( d c , d k ) W_{UK,s}:(d_c,d_k) W U K , s : ( d c , d k ) 。重构出来的 content key 是 k s , u C = c u K V W U K , s k_{s,u}^C=c_u^{KV}W_{UK,s} k s , u C = c u K V W U K , s 。按照第 3 节的 row-vector RoPE convention,query 和 key 分别右乘 R t ⊤ R_t^\top R t ⊤ 与 R u ⊤ R_u^\top R u ⊤ ,因此这一对 positions 的 score 是
S s , t , u R o P E = ( q s , t C R t ⊤ ) ( ( c u K V W U K , s ) R u ⊤ ) ⊤ = q s , t C R t ⊤ R u W U K , s ⊤ ( c u K V ) ⊤ . \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} S s , t , u RoPE = ( q s , t C R t ⊤ ) ( ( c u K V W U K , s ) R u ⊤ ) ⊤ = q s , t C R t ⊤ R u W U K , s ⊤ ( c u K V ) ⊤ .
问题只在于这条矩阵乘法链应该怎样加括号。先去掉 RoPE,看普通 MLA 为什么可以做 key absorption。用 T T T 表示 query positions 的数量、U U U 表示 key positions 的数量,完整的 score chain 是
S C ( h , T , U ) = Q C ( h , T , d k ) W U K ⊤ ( h , d k , d c ) ( C K V ) ⊤ ( d c , 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}. ( h , T , U ) S C = ( h , T , d k ) Q C ( h , d k , d c ) W U K ⊤ ( d c , U ) ( C K V ) ⊤ .
一种顺序是先从 latent 重构所有 keys,再计算 attention:
Q C ( h , T , d k ) ( W U K ⊤ ( h , d k , d c ) ( C K V ) ⊤ ( d c , U ) ) = Q C ( h , T , d k ) ( K C ) ⊤ ( h , d k , 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}. ( h , T , d k ) Q C ( ( h , d k , d c ) W U K ⊤ ( d c , U ) ( C K V ) ⊤ ) = ( h , T , d k ) Q C ( h , d k , U ) ( K C ) ⊤ .
另一种顺序是先把 W U K W_{UK} W U K 吸收到 query 一侧:
( Q C ( h , T , d k ) W U K ⊤ ( h , d k , d c ) ) ( C K V ) ⊤ ( d c , U ) = Q ^ C ( h , T , d c ) ( C K V ) ⊤ ( d c , 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}. ( ( h , T , d k ) Q C ( h , d k , d c ) W U K ⊤ ) ( d c , U ) ( C K V ) ⊤ = ( h , T , d c ) Q C ( d c , U ) ( C K V ) ⊤ .
W U K W_{UK} W U K 是固定 weight,所以 absorbed query Q ^ C = Q C W U K ⊤ \widehat Q^C=Q^CW_{UK}^\top Q C = Q C W U K ⊤ 与 key position 无关。每个 query 只需要计算一次这个 ( h , T , d c ) (h,T,d_c) ( h , T , d c ) intermediate,随后就能与所有 cached latents 相乘。这正是 key absorption。
现在把 RoPE 放回去。上面的完整 score 可以按“先重构并旋转 key”的顺序计算:
( q s , t C R t ⊤ ) ( R u ( W U K , s ⊤ ( c u K V ) ⊤ ) ) . \left(q_{s,t}^C R_t^\top\right)
\left(
R_u
\left(W_{UK,s}^\top(c_u^{KV})^\top\right)
\right). ( q s , t C R t ⊤ ) ( R u ( W U K , s ⊤ ( c u K V ) ⊤ ) ) .
也可以尝试像刚才一样,把 key reconstruction 之前的 factors 全部放进左边的括号:
( q s , t C R t ⊤ R u W U K , s ⊤ ) ⏟ 希望吸收到 query 一侧的部分 ( c u K V ) ⊤ . \underbrace{\left(q_{s,t}^C R_t^\top R_u W_{UK,s}^\top\right)}_{\text{希望吸收到 query 一侧的部分}}
(c_u^{KV})^\top. 希望吸收到 query 一侧的部分 ( q s , t C R t ⊤ R u W U K , s ⊤ ) ( c u K V ) ⊤ .
区别现在一眼就能看到:这个括号里含有 R u R_u R u 。它会随 key position u u u 改变,所以不能再得到一个对所有 keys 复用的 Q ^ C \widehat Q^C Q C 。如果把所有 query 和 key positions 一次展开,这个括号产生的 intermediate shape 不是 ( h , T , d c ) (h,T,d_c) ( h , T , d c ) ,而是
( h , T , U , d c ) . (h,T,U,d_c). ( h , T , U , d c ) .
其中 U U U 是 key-position axis:prefill 时 U = T U=T U = T ;一步 decode 时 T = 1 T=1 T = 1 、U = t U=t U = t 。换句话说,RoPE 没有让矩阵乘法失效,但它让所谓的“absorbed query”变成每个 ( t , u ) (t,u) ( t , u ) position pair 各自一个。这样既没有消掉 key-length axis,也失去了 absorption 的意义。若 cache 只保存 C K V C^{KV} C K V ,decode 时就只能重新构造并旋转历史 keys,或显式生成这个 ( h , 1 , t , d c ) (h,1,t,d_c) ( h , 1 , t , d c ) intermediate。
因此,普通 RoPE 和 position-independent key absorption 不能放在同一条 feature path 上。第 3 节的 elementwise RoPE 与这里的 rotation-matrix form 完全等价;改写成 elementwise computation 并不会消除对 key position u u u 的依赖,因此同样无法实现 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 保存 C K V C^{KV} C K V 和一份 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.
令 d r d_r d r 表示 decoupled RoPE branch 的 per-head dimension,对应 DeepSeek-V2 技术报告中的 d h R d_h^R d h R 。这里沿用第 3 节的 rotation tensor,只把 tensor dimension 换成 d r d_r d r :
R ( r ) ( T , d r , d r ) , R t , : , : ( r ) = ( R t ( r ) ) ⊤ . \underset{(T,d_r,d_r)}{\mathcal R^{(r)}},
\qquad
\mathcal R^{(r)}_{t,:,:}=(R_t^{(r)})^\top. ( T , d r , d r ) R ( r ) , R t , : , : ( r ) = ( R t ( r ) ) ⊤ .
Query path 是
C Q ( T , d q c ) = X ( T , d ) W D Q ( d , d q c ) , \underset{(T,d_{qc})}{C^Q}
=
\underset{(T,d)}{X}
\underset{(d,d_{qc})}{W_{DQ}}, ( T , d q c ) C Q = ( T , d ) X ( d , d q c ) W D Q ,
Q C ( h , T , d k ) = C Q ( T , d q c ) W U Q ( h , d q c , d k ) , Q ‾ R ( h , T , d r ) = C Q ( T , d q c ) W Q R ( h , d q c , d r ) , \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}}, ( h , T , d k ) Q C = ( T , d q c ) C Q ( h , d q c , d k ) W U Q , ( h , T , d r ) Q R = ( T , d q c ) C Q ( h , d q c , d r ) W QR ,
Q R ( h , T , d r ) = Q ‾ R ( h , T , d r ) R ( r ) ( T , d r , d r ) . \underset{(h,T,d_r)}{Q^R}
=
\underset{(h,T,d_r)}{\overline Q^R}
\underset{(T,d_r,d_r)}{\mathcal R^{(r)}}. ( h , T , d r ) Q R = ( h , T , d r ) Q R ( T , d r , d r ) R ( r ) .
Query-side RoPE operation 可以保持不变。但 key side 不同:如果 RoPE key 以 compressed form 存储,每个 decode step 都必须重新展开并旋转全部 t t t 个历史 keys,产生大量重复计算。因此,完整 MLA 直接缓存 materialized K R K^R K R ,而不是 compressed positional-key latent。
Key-value path 是
C K V ( T , d c ) = X ( T , d ) W D K V ( d , d c ) , \underset{(T,d_c)}{C^{KV}}
=
\underset{(T,d)}{X}
\underset{(d,d_c)}{W_{DKV}}, ( T , d c ) C K V = ( T , d ) X ( d , d c ) W D K V ,
K C ( h , T , d k ) = C K V ( T , d c ) W U K ( h , d c , d k ) , V ( h , T , d v ) = C K V ( T , d c ) W U V ( h , d c , d v ) , \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}}, ( h , T , d k ) K C = ( T , d c ) C K V ( h , d c , d k ) W U K , ( h , T , d v ) V = ( T , d c ) C K V ( h , d c , d v ) W U V ,
K ‾ R ( T , d r ) = X ( T , d ) W K R ( d , d r ) , K R ( T , d r ) = K ‾ R ( T , d r ) R ( r ) ( T , d r , d r ) . \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)}}. ( T , d r ) K R = ( T , d ) X ( d , d r ) W K R , ( T , d r ) K R = ( T , d r ) K R ( T , d r , d r ) R ( r ) .
如果每个 head 都保存一份 positional key,每个 token 会额外增加 h d r hd_r h d r 个 cache elements。完整 MLA 改为让所有 h h h 个 query heads 共享同一个 K R : ( T , d r ) K^R:(T,d_r) K R : ( T , d r ) 。因此 RoPE branch 采用的是 MQA-style 结构:positional query 仍然是 per-head 的 Q R : ( h , T , d r ) Q^R:(h,T,d_r) Q R : ( h , T , d r ) ,但 positional key 只有一个 shared head。这样额外 cache 保持为每个 token d r d_r d r 个 elements。
完整的 per-head query 和 key,是沿最后一个 tensor dimension 的 concat:
Q ( h , T , d k + d r ) = [ Q C ( h , T , d k ) ; Q R ( h , T , d r ) ] , \underset{(h,T,d_k+d_r)}{Q}
=
\left[
\underset{(h,T,d_k)}{Q^C};
\underset{(h,T,d_r)}{Q^R}
\right], ( h , T , d k + d r ) Q = [ ( h , T , d k ) Q C ; ( h , T , d r ) Q R ] ,
K ( h , T , d k + d r ) = [ K C ( h , T , d k ) ; K R ( T , d r ) 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]. ( h , T , d k + d r ) K = [ ( h , T , d k ) K C ; ( T , d r ) K R broadcast over h ] .
为什么 score 会变成两个部分相加?Concat 会把最后一个 tensor dimension 划分成两个互不重叠的 blocks。对于 head s s s 、query position t t t 和 key position u u u ,完整 dot product 是
S s , t , u = ∑ ℓ = 1 d k + d r Q s , t , ℓ K s , u , ℓ = ∑ a = 1 d k Q s , t , a C K s , u , a C + ∑ b = 1 d r Q s , t , b R K u , b R = S s , t , u C + S s , t , u R . \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} S s , t , u = ℓ = 1 ∑ d k + d r Q s , t , ℓ K s , u , ℓ = a = 1 ∑ d k Q s , t , a C K s , u , a C + b = 1 ∑ d r Q s , t , b R K u , b R = S s , t , u C + S s , t , u R .
同一个恒等式也可以写成 block tensor multiplication:
S ( h , T , T ) = [ Q C ( h , T , d k ) Q R ( h , T , d r ) ] [ ( K C ) ⊤ ( h , d k , T ) ( K R ) ⊤ ( d r , T ) broadcast over h ] = S C ( h , T , T ) + S R ( 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}. ( h , T , T ) S = [ ( h , T , d k ) Q C ( h , T , d r ) Q R ] ( h , d k , T ) ( K C ) ⊤ ( d r , T ) ( K R ) ⊤ broadcast over h = ( h , T , T ) S C + ( h , T , T ) S R .
这里不存在 Q C ( K R ) ⊤ Q^C(K^R)^\top Q C ( K R ) ⊤ 或 Q R ( K C ) ⊤ Q^R(K^C)^\top Q R ( K C ) ⊤ 这样的 cross terms:content 和 RoPE coordinates 位于不同的 feature blocks,dot product 只会配对同一个 block 内的 coordinates。因此,两条 score channel 可以分别计算:
S C ( h , T , T ) = Q C ( h , T , d k ) ( K C ) ⊤ ( h , d k , T ) , S R ( h , T , T ) = Q R ( h , T , d r ) ( K R ) ⊤ ( d r , 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}, ( h , T , T ) S C = ( h , T , d k ) Q C ( h , d k , T ) ( K C ) ⊤ , ( h , T , T ) S R = ( h , T , d r ) Q R ( d r , T ) ( K R ) ⊤ ,
P ( h , T , T ) = softmax ( S C + S R d k + d r ) , \underset{(h,T,T)}{P}
=
\operatorname{softmax}\!\left(
\frac{S^C+S^R}{\sqrt{d_k+d_r}}
\right), ( h , T , T ) P = softmax ( d k + d r S C + S R ) ,
O ( h , T , d v ) = P ( h , T , T ) V ( h , T , d v ) , Y ( T , d ) = O ( h , T , d v ) W o ( h , d v , 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}. ( h , T , d v ) O = ( h , T , T ) P ( h , T , d v ) V , ( T , d ) Y = ( h , T , d v ) O ( h , d v , d ) W o .
对于 cached decode,content absorption 和 positional branch 给出
s c o r e ( h , 1 , t ) = ( q C ( h , 1 , d k ) W U K ⊤ ( h , d k , d c ) ) ( C ≤ t K V ) ⊤ ( d c , t ) + q R ( h , 1 , d r ) ( K ≤ t R ) ⊤ ( d r , 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}, ( h , 1 , t ) score = ( ( h , 1 , d k ) q C ( h , d k , d c ) W U K ⊤ ) ( d c , t ) ( C ≤ t K V ) ⊤ + ( h , 1 , d r ) q R ( d r , t ) ( K ≤ t R ) ⊤ ,
o ( h , 1 , d v ) = ( p ( h , 1 , t ) C ≤ t K V ( t , d c ) ) W U V ( h , d c , d v ) . \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}}. ( h , 1 , d v ) o = ( ( h , 1 , t ) p ( t , d c ) C ≤ t K V ) ( h , d c , d v ) W U V .
因此 content key 和 value 仍然留在 C K V C^{KV} C K V 后面,只有 shared rotated key K R K^R K R 被显式缓存。
Cache 补表。
对于 DeepSeek-V2,d c = 512 d_c=512 d c = 512 ,d r = 64 d_r=64 d r = 64 。因此完整 cache 相对 MHA 的缩减是 32768 / 576 ≈ 56.9 × 32768/576\approx56.9\times 32768/576 ≈ 56.9 × ,而不是只计算 RoPE-free latent 时的 64 × 64\times 64 × 。
FLOPs 补表。
这个补表沿用第 4 节的 accounting,统计 decoupled branch 新增的 tensor contractions。这里也计算显式的 S C + S R S^C+S^R S C + S R :每个 score element 做一次 addition,因此 prefill 是 h T 2 hT^2 h T 2 FLOPs,cached decode 是 h t ht h t FLOPs。RoPE rotation 本身的 structured elementwise operations 不进入这个对比。
对于 cached decode,t t t 的 coefficient 现在是 4 h d c + 2 h d r + h = 278,656 4hd_c+2hd_r+h=278{,}656 4 h d c + 2 h d 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 = 1 D=1 D = 1 ),预测一个额外 token。正常 inference 时可以移除这个 module;用于 speculative decoding 时,它可以 draft 一个 token。报告给出的第二个 token acceptance rate 为 85 % 85\% 85% -90 % 90\% 90% ,并报告了 speculative decoding 的 1.8 × 1.8\times 1.8 × TPS。
在 MHA 中,多个 verification query 会复用同一份 KV-cache 数据。由于 decode 是强 memory-bound 的,它们带来的额外 attention 计算可以被 cache 读取的时间掩盖:memory bandwidth 已经成为 bottleneck,而 compute units 仍有余量。MLA 改变了这个平衡。它用 C K V C^{KV} C K V 和 decoupled RoPE key 替代完整的 per-head KV cache,同时 absorption 又增加了围绕 C K V C^{KV} C K V 的 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} } ,
}