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 值 |
|---|---|---|
| hidden | 7168 | |
| head 数 | 96 | |
| KV latent(kv_lora_rank) | 512 | |
| Q latent(q_lora_rank) | 1536 | |
| 每 head nope 维 | 128 | |
| pe 槽位维 | 64(NoPE:不旋转) | |
| MLA 层数 | 24 |
完整形式(训练/prefill 视角)分三步。降维(每 token 一次,head 无关):
按 head 升维(96 份):、 等。每头注意力照常。关键在于 KV cache 只需要存 维——所有 head 共享一行。对比 MHA 的 /token:约 43× 压缩(K3 头多,压缩比比 48B 家族的 14× 更大)。K3 相对 DeepSeek 原版的改动是剥掉 RoPE(NoPE)+ 输出门控(Gated MLA)。
2. 吸收形式:decode 实际执行的东西
decode 时按完整形式算要给每个 token、每个 head 升维出 96 份 K/V——开销随上下文长度线性涨。两个恒等式把它消掉:
恒等式①(K 升维搬进 query 侧):
恒等式②(V 升维推迟到注意力之后):
于是 decode 变成 576 维单 KV head 的 MQA:query 每 head 是 (每步现算),KV 缓存就是那 576 维 latent(head 共享)。代价是 Q 侧 GEMM 变大(128→512 吸收 + 1536→96×192 升维),但 decode 每 step 只有 B 个 token——划算;收益是升维计算量与上下文长度无关。
这个"单 KV head"的形态正是 DCP 的地基:KV 没有按 head 的结构,才可能干净地按位置切条带。
3. 归并恒等式:LSE 两遍法
CP 把 key 集合切成不相交的 。每 rank 对自己的条带算局部 softmax 注意力:
全局输出可以用局部量精确重建:
这就是 FlashAttention 式两遍法的分布式版本(,,)——精确归并,无近似。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: | [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,得 | [B,96,512] | 无 |
| ⑥ All-to-all | 把 送到 head owner | 收 12 heads × 8 份 partial | ≈ 5.5 MB/rank |
| ⑦ LSE 归并 + V 升维 | 两遍法合并 → @ | [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 的通信量 (decode 每 step 只有 64 个 token,一份 [64,96,576] bf16 才 7 MB,复制 8 份也才 ~50 MB/层/步),KV 的显存 上下文长度(1M 时 27 GB)。用 Q 的冗余通信买 KV 的显存切分,且交换比随上下文增长越来越划算——1M 上下文时切分比不切每 rank 省 ~24 GB。
6. 与组件的对照速查
| 公式/步骤 | SGLang 代码 |
|---|---|
| 一次 GEMM(步骤①) | fused_qkv_a_proj_with_mqa |
| 两个 RMSNorm | q_a_layernorm / kv_a_layernorm |
| (步骤③) | q_b_proj |
| 恒等式①(步骤⑤前) | q_nope @ w_kc |
| 576-MQA + lse(步骤⑤) | FlashMLA/cutedsl backend,return_lse=True |
| 恒等式②(步骤⑦) | attn_output @ w_vc |
| 行并行 + 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 归并恒等式