MLA 的 Decode Context Parallelism:用 Q 的通信买 KV 的显存

长上下文 decode 的显存瓶颈几乎全在 KV cache。MLA 已经把每 token 每层的缓存压到 576 维 latent(对比 MHA 约 43× 压缩),但 1M 上下文下 Kimi K3 仍要 27.6 GB——如果每张卡都存一份完整副本,8 卡一组就是无谓的 8 倍冗余。Decode Context Parallelism(DCP)的解法一句话可以说完:沿序列维度把 KV 切成条带,每张卡只缓存自己那一段;KV 全程不通信,通信只发生在 Q 侧,而 Q 的通信量正比于 batch、与上下文长度无关。本文以 SGLang 的 DCP 实现为线索,从 MLA 的两个吸收恒等式讲起,拆到 decode 一步的八阶段流水,最后算清显存和通信两笔账(维度取 Kimi K3 实测值)。


1. MLA 复习:为什么 KV 只剩 576 维

先把符号表摆出来(Kimi K3):

符号 含义 K3 值
dmodeld_{model} hidden 7168
HH head 数 96
dcd_c KV latent(kv_lora_rank) 512
dc′d_c' Q latent(q_lora_rank) 1536
dnd_n 每 head nope 维 128
dRd_R pe 槽位维 64(NoPE:不旋转)
LMLAL_{MLA} MLA 层数 24

完整形式(训练/prefill 视角)分三步。降维(每 token 一次,head 无关):

ctK=WDKVht∈R512,ktR=WKRht∈R64,ctQ=WDQht∈R1536c^K_t = W^{DKV} h_t \in \mathbb{R}^{512}, \quad k^R_t = W^{KR} h_t \in \mathbb{R}^{64}, \quad c^Q_t = W^{DQ} h_t \in \mathbb{R}^{1536}

按 head 升维(96 份):kt,hC=WhUKcˉtKk^C_{t,h} = W^{UK}_h \bar{c}^K_t、vt,hC=WhUVcˉtKv^C_{t,h} = W^{UV}_h \bar{c}^K_t 等。每头注意力照常。关键在于 KV cache 只需要存 [cˉtK∥ktR]=576[\bar{c}^K_t \Vert k^R_t] = 576 维——所有 head 共享一行。对比 MHA 的 2⋅H⋅dv=2×96×128=245762 \cdot H \cdot d_v = 2 \times 96 \times 128 = 24576/token:约 43× 压缩(K3 头多,压缩比比 48B 家族的 14× 更大)。K3 相对 DeepSeek 原版的改动是剥掉 RoPE(NoPE)+ 输出门控(Gated MLA)。

2. 吸收形式:decode 实际执行的东西

decode 时按完整形式算要给每个 token、每个 head 升维出 96 份 K/V——开销随上下文长度线性涨。两个恒等式把它消掉:

恒等式①(K 升维搬进 query 侧):

qC⋅(WhUKcˉK)=(qC⋅WhUK)⏟q^hC∈R512⋅cˉKq^C \cdot (W^{UK}_h \bar{c}^K) = \underbrace{(q^C \cdot W^{UK}_h)}_{\hat{q}^C_h \in \mathbb{R}^{512}} \cdot \bar{c}^K

恒等式②(V 升维推迟到注意力之后):

∑τατ(WhUVcˉτK)=WhUV⋅(∑τατcˉτK)⏟u^h∈R512\sum_{\tau} \alpha_{\tau} (W^{UV}_h \bar{c}^K_{\tau}) = W^{UV}_h \cdot \underbrace{\left( \sum_{\tau} \alpha_{\tau} \bar{c}^K_{\tau} \right)}_{\hat{u}_h \in \mathbb{R}^{512}}

于是 decode 变成 576 维单 KV head 的 MQA:query 每 head 是 [q^hC∥qhR]∈R576[\hat{q}^C_h \Vert q^R_h] \in \mathbb{R}^{576}(每步现算),KV 缓存就是那 576 维 latent(head 共享)。代价是 Q 侧 GEMM 变大(128→512 吸收 + 1536→96×192 升维),但 decode 每 step 只有 B 个 token——划算;收益是升维计算量与上下文长度无关。

这个"单 KV head"的形态正是 DCP 的地基:KV 没有按 head 的结构,才可能干净地按位置切条带。

3. 归并恒等式:LSE 两遍法

CP 把 key 集合切成不相交的 T=∪rTrT = \cup_r T_r。每 rank 对自己的条带算局部 softmax 注意力:

Zr=∑τ∈Tresτ,o(r)=∑τ∈TresτvτZrZ_r = \sum_{\tau \in T_r} e^{s_{\tau}}, \quad o^{(r)} = \frac{\sum_{\tau \in T_r} e^{s_{\tau}} v_{\tau}}{Z_r}

全局输出可以用局部量精确重建:

o=∑rZrZ⋅o(r),ZrZ=exp⁡ ⁣(lser−log⁡∑relser)o = \sum_{r} \frac{Z_r}{Z} \cdot o^{(r)}, \quad \frac{Z_r}{Z} = \exp\!\left( \mathrm{lse}_r - \log \sum_{r} e^{\mathrm{lse}_r} \right)

这就是 FlashAttention 式两遍法的分布式版本(m=max⁡rlserm = \max_r \mathrm{lse}_r,wr=exp2(lser−m)w_r = \mathrm{exp2}(\mathrm{lse}_r - m),out=∑rwro(r)/∑rwr\mathrm{out} = \sum_r w_r o^{(r)} / \sum_r w_r)——精确归并,无近似。SGLang 的 attention backend(FlashMLA/cutedsl)开 return_lse=True 吐出每 rank 的 lse,归并 kernel 就是这个公式的 Triton 实现。

4. Decode 一步的八阶段流水(TP=8 × CP=8,B=64)

一组 8 卡,每个 rank 身兼两职:TP 维上是 12/96 个 head 的 owner,CP 维上是 1/8 条带 KV 的 owner(owner = slot mod 8)。hidden states 每 rank 各有一份完整副本(下一步的 MoE 需要全量)。

步 计算 形状(每 rank) 通信
① 下投影 fused qkv_a_proj:h→cK∥kR∥cQh \to c^K \Vert k^R \Vert c^Q [B,7168] → [B,512]+[B,64]+[B,1536] 无(冗余计算)
② 写 KV 池 set_mla_kv_buffer(dcp_sharded),owner 过滤 每层每 token 一行 576 无(owner 只落盘自己的槽位)
③ Q 升维 q_b_proj:只算自己 12 个 head [B,12,576] —
④ All-gather Q 12 heads → 96 heads [B,96,576] ≈ 7.1 MB ≈ 6.2 MB/rank
⑤ 条带注意力 对本 rank 的 1/8 KV 做 576-MQA,得 (o(r),lser)(o^{(r)}, \mathrm{lse}_r) o(r)o^{(r)}[B,96,512] 无
⑥ All-to-all 把 (o(r),lser)(o^{(r)}, \mathrm{lse}_r) 送到 head owner 收 12 heads × 8 份 partial ≈ 5.5 MB/rank
⑦ LSE 归并 + V 升维 两遍法合并 → u^h\hat{u}_h @ wvcw_{vc} [B,12,128] 无
⑧ 输出投影 o_proj 行并行 + All-reduce [B,7168] ≈ 1.6 MB/rank

注意几个设计点:

  • KV 全程零通信。步骤⑤每个 rank 拿全量 Q 对本地 1/8 条带算注意力——Q 吸收 K 上投影已在 Q 侧完成(恒等式①),所以 KV 不需要任何变换就能直接参与点积。这是"吸收形式"给 DCP 的礼物:如果按完整形式算,每 rank 得先把自己条带的 latent 升维成 96 份 K/V,通信和计算都回来。
  • head 与条带的双重身份在步骤⑥解耦:注意力输出按条带分散在各 rank、按 head 切碎,A2A 把 8 份 partial 送到各自 head 的 owner 处归并——partial 和的合并恰好就是 LSE 恒等式。
  • 步骤①是冗余计算换零通信:下投影每 rank 都算全量(B 个 token 的 GEMM 很便宜),只有 owner 落盘。物理 C 行 → 逻辑 8C(8 个 rank 的物理行拼成逻辑池)。

5. 两笔账

显存(KV cache,bf16):

项目 算式 结果
每 token 每层 576 × 2 B ≈ 1.125 KB
全模型 · 1M token(24 层) × 24 ≈ 27.6 GB
不切分 — 每 rank 27.6 GB
CP=8 切分后 ÷ 8 每 rank ≈ 3.45 GB

通信(每层每 step,B=64,bf16):

通信 算式 每 rank 发/收
④ AG Q 64×96×576×2B ≈ 7.1 MB,发 7/8 ≈ 6.2 MB
⑥ A2A partial 64×96×(512+4)×2B ≈ 6.3 MB,发 7/8 ≈ 5.5 MB
⑧ AR 部分和 64×7168×2B ≈ 0.92 MB,环 AR × 2×(7/8) ≈ 1.6 MB

× 24 层:每 rank ≈ 320 MB / decode step;8 卡总线总流量 ≈ 2.5 GB/步。

核心交换由此显形:Q 的通信量 ∝B\propto B(decode 每 step 只有 64 个 token,一份 [64,96,576] bf16 才 7 MB,复制 8 份也才 ~50 MB/层/步),KV 的显存 ∝\propto 上下文长度(1M 时 27 GB)。用 Q 的冗余通信买 KV 的显存切分,且交换比随上下文增长越来越划算——1M 上下文时切分比不切每 rank 省 ~24 GB。

6. 与组件的对照速查

公式/步骤 SGLang 代码
WDQ/WDKVW^{DQ}/W^{DKV} 一次 GEMM(步骤①) fused_qkv_a_proj_with_mqa
两个 RMSNorm q_a_layernorm / kv_a_layernorm
WUQ/WQRW^{UQ}/W^{QR}(步骤③) q_b_proj
恒等式①(步骤⑤前) q_nope @ w_kc
576-MQA + lse(步骤⑤) FlashMLA/cutedsl backend,return_lse=True
恒等式②(步骤⑦) attn_output @ w_vc
WOW^O 行并行 + AR(步骤⑧) o_proj + all-reduce
owner 过滤写池(步骤②) set_mla_kv_buffer_dcp_sharded

7. 总结

把 DCP 放回更大的图景里:

  • MLA 的贡献是把 KV 压小(43×,压到 head 无关的 576 维 latent);
  • 吸收形式的贡献是让升维开销与序列长度无关——这同时是 decode 效率和解法可行性的来源;
  • DCP 的贡献是把压小的 KV 再沿位置轴切开,靠 LSE 恒等式精确归并,零近似。

三个东西是同一条设计线上的三节:压缩(MLA)→ 计算重排(吸收)→ 空间切分(DCP)。所以 DCP 不是"又一个并行策略",而是 MLA 形态的自然推论——KV 一旦没有按 head 的结构,按位置切就是最干净的下一步。反过来说,MHA 想做同样的事就得在条带边界处理 head 间的交互,代价完全不同。

参考

  • SGLang DCP 实现(fused_qkv_a_proj_with_mqa / set_mla_kv_buffer_dcp_sharded / FlashMLA backend)
  • Kimi K3 模型配置(kv_lora_rank=512, q_lora_rank=1536, 96 heads, NoPE)
  • FlashAttention 两遍法与 LSE 归并恒等式