ggaaooppeenngg

为什么计算机科学是无限的但生命是有限的

MiMo-V2.5 推理优化解读:Hybrid SWA 的工程落地

小米 MiMo 团队的《Full-Pipeline Inference Optimization for MiMo-V2.5 Series》(arXiv:2607.13095)核心论点很直白:Hybrid SWA 理论上能把 KVCache 和 attention 计算量压到 Full Attention 的 1/7,但如果 KVCache 系统不跟着改造,实际反而会变慢

本文不写综述,重点拆解几个关键工程问题:SWA 的 KV cache 怎么存、怎么 offload、共享前缀命中时的重算问题、numa_balancing 为什么 +10%、以及 Prefill 为什么也要开 MTP


1. 论文速览

MiMo-V2.5-Pro 的架构组合拳:

1
2
Hybrid SWA + 稀疏 MoE + 多模态编码器
70 层 = 10 Full Attention + 60 SWA(窗口 W=128)

理论上 attention FLOPs 和 KVCache 存储都降到 1/7,但论文的核心观点是:这些理论红利不会自动兑现。Hybrid SWA 在 KVCache 管理、前缀匹配、Full 和 SWA 层语义一致性上都引入了新问题。

论文涉及的优化面很广:KVCache 分层(L1 GPU / L2 Host / L3 GCache)、调度(LLM-Router)、Prefill/Decode pipeline、多模态等。其中 LLM-Router 和 GCache 属于常见的分布式缓存和调度实现,本文不展开。重点聚焦在 SWA 相关 KVCache 这条主线上。


2. SWA 的 KV Cache 怎么存

核心设计:物理分离 + 逻辑统一

物理层:两个独立池子

每层存储(L1 GPU / L2 Host / L3 GCache)都并排维护两个 KV Pool:

1
2
3
4
5
6
7
8
9
┌─────────────────────────────────────┐
│ Full KV Pool │
│ ├─ 大小: O(seq_len × 10 层) │
│ └─ 淘汰: 按整个序列做 LRU │
├─────────────────────────────────────┤
│ SWA KV Pool │
│ ├─ 大小: O(W × 60 层) ← 严格 O(W) │
│ └─ 淘汰: 独立按 window 做 eviction │
└─────────────────────────────────────┘

关键点:SWA pool 物理上就只有 W 大小,写入新 token 时按窗口 evict 旧的。不是「存完整历史然后只用尾部 128 个」,而是存进来就已经是窗口截断后的样子

直观对比(128K token 序列):

Full Attention Hybrid SWA
60 个 SWA 层总量 60 × 128K = 7.68M 60 × 128 = 7.68K
10 个 Full 层总量 10 × 128K = 1.28M 10 × 128K = 1.28M
合计 8.96M ≈ 1.29M(省 7×)

Hybrid SWA 存储量几乎等于只算 Full 层的量,SWA 层几乎不占空间。

逻辑层:单序列视图 + 双索引前缀树

上层(前缀树、调度器)只看到一个逻辑序列,底下用 Full→SWA mapping 做透明分层。每个前缀树节点存两套元数据:

1
2
3
4
5
PrefixTreeNode {
tokens: [t0, t1, ..., tN]
full_seg_idx: [Full KV pool 位置] # Full 层复用
swa_seg_map: [SWA KV pool 位置 or 空] # 判断 window 安全性
}

用途

  1. 命中判定:token 相等 + tail W 个 token 的 swa_seg_map 都非空 → 才算真命中
  2. 独立淘汰:window 外的 SWA 段可以先扔,Full 段还留着

一句话总结:SWA 存的就是「每 SWA 层永远只存 W 个 token 的 KV」,SWA pool 和 Full pool 物理分离,上层用双索引前缀树抽象成单序列视图。


3. SWA Cache 怎么 Offload 到 CPU 内存

存什么:只存「窗口内有效 slot」

对一条 100K token 的序列:

1
2
Full 层 (10 层):  每层 100K 个 slot  ← 全部下沉到 CPU
SWA 层 (60 层): 每层只有 W=128 个 slot ← 只存窗口内

SWA 下沉到 CPU 的就是窗口截断后的形态,不是完整历史。

怎么存:Host 侧独立的 SWA Pool

L2 结构完全镜像 L1:

1
2
3
Host DRAM (L2):
Full KV Pool: [10 层, seq_len, head_dim], paged + pinned memory
SWA KV Pool: [60 层, W=128, head_dim], 独立 eviction

工程要点:

  • Pin memory:Host 池子用固定内存,D2H/H2D 用异步 DMA
  • Paged 分配:SWA pool 拆成 block,滑出窗口的 block 归还池子
  • 和 Full pool 完全独立:分开的地址空间、eviction 策略

怎么搬:按 SWA mask 传输

论文 §3.1.1 的关键一句:

“Cross-tier transfers are performed based solely on the SWA mask, ensuring only valid window data is moved.”

D2H / H2D 时不搬窗口外的数据

1
2
3
D2H writeback: 只把 mask=1 的 slot 打包 DMA 到 host
H2D prefetch: 只拉窗口内的 W 个 token 的 KV
→ 数据量小,layerwise prefetch 能 overlap 计算

这里省的是带宽。按 mask 走,60 层每层就 128 个 token,总量小到可以忽略。

D2H 时机

  • Eviction 前的备份:GPU SWA pool 满了,先 D2H 到 host 再释放
  • 一致性修复(§3.1.4):前缀树节点合并 / prefill chunk 完成时,检查 device 和 host 的 SWA 占用差 → host 补齐槽位 → 异步 D2H
  • 请求结束 / 会话切换:高价值会话主动 D2H 到 host,甚至下沉到 L3

一句话总结:SWA 下沉到 CPU 的就是窗口截断后的 128 个 slot,Host 侧用独立 pool + pinned memory 存,搬运只按 SWA mask 走有效位置。


4. 共享前缀命中时,SWA 要不要重算

这是 Hybrid SWA 落地生产最难受的一个点

问题本质

典型场景:共享 system prompt(2000 tokens),下面挂了很多用户会话:

1
2
3
4
PrefixTree:
system_prompt (2000 tokens) ← 非叶子节点
├─ user_A_session (往下延展了 50K tokens)
└─ user_B_session (往下延展了 30K tokens)

关键观察:user_A 会话延展了 50K tokens,当前窗口在 [52000-128, 52000] 附近。system_prompt 段(位置 0-2000)早就滑出窗口了,SWA KV 被 evict 掉了。

这就是「伪命中」

user_D 带着同一个 system prompt 进来:

1
2
3
4
5
命中判定:
1. Token equality 检查 → 匹配到 2000 tokens ✅
2. Full KV 检查 → 都在 → Full 命中 2000 tokens ✅
3. Window-safe 检查 → tail 128 个 SWA slot 在不在?
└─► ❌ 已被 evict

按「window-safe length」规则,SWA 匹配长度被 clip。Full 层能复用 2000 tokens,但 SWA 层不行——SWA 层要看窗口内 KV,位置 2001 的 attention 需要 [1873, 2001] 的 KV,这段没了 → 只能重算 SWA 层的 prefill。

代价

  • Full 层:省了(Full pool 还留着)
  • SWA 层:60 层全部重跑一次 prefill(占 6/7)
  • 60 层重算,几乎等于白干

小米的应对:SWA 保留策略(§3.1.4 第 4 条)

论文直接点名了这个场景:

“Medium/short sequence SWA retention strategy. Based on user request patterns, we retain relatively dense SWA KV Cache at fixed length positions for medium/short sequences… particularly beneficial for long agent sessions, multi-user shared system prompts, and repeated tool calls to the same codebase.

翻译:对高频复用的前缀(比如 system prompt),SWA pool 不严格 O(W),而是在关键位置留几份 SWA 快照

推测的实现方式:

1
2
3
4
5
6
7
普通序列的 SWA pool:
只留 tail W 个 slot ← 严格 O(W)

高频前缀(如 system prompt):
在几个 anchor 位置多留 W 个 slot
比如: [0, W], [500-W, 500], [1000-W, 1000], [2000-W, 2000]
← 相当于"检查点",每个都能作为 SWA 计算的起点

这样 user_D 命中时:匹配到 system_prompt 末尾 tail W=128 的 SWA slot 确实存在(是保留下来的 anchor),Full + SWA 都能完整复用 → 不需要重算 SWA 层

Trade-off 账

严格 O(W) 高频前缀密集保留
SWA 存储占比 极低(1/7) 略高(几倍 W)
单节点并发能力 稍低
共享前缀命中率 差(伪命中) 好(真命中)

核心思路用一点点 SWA 存储换取显著提升的共享前缀命中率。因为 system prompt 被成千上万用户共享,每次省掉 60 层 SWA prefill 的收益,远远大于多存几份 128-slot 快照的代价。

一句话总结:非叶子共享前缀命中时,默认严格 O(W) 实现会因为 tail 被 evict 导致 SWA 层必须重算。小米的解法是对高频共享前缀主动保留多个 SWA 快照,让 Full + SWA 都能完整复用,真正兑现 Hybrid SWA 在多用户共享场景下的效率红利。


5. numa_balancing 为什么影响 +10%

论文里就一句话:「关掉 numa_balancing,端到端 +10%」。背后是操作系统调度和 GPU 推理的碰撞

numa_balancing 是干嘛的

现代服务器都是 NUMA 架构,CPU 访问本 NUMA 节点内存 << 访问远端节点内存(延迟差 1.5-2×)。kernel.numa_balancing=1 是 Linux 的「自动 NUMA 均衡」特性:

1
2
3
4
5
6
7
内核周期性做的事:
1. 扫描进程页表 → 把页面标记为 PROT_NONE
2. 进程下次访问 → 触发缺页异常
3. 内核在 fault handler 里统计 NUMA locality
4. 如果 CPU 和 page 不在同一 node:
- 迁移页面到 CPU 所在 node
- 或者迁移线程到 page 所在 node

但在 GPU 推理场景下变成灾难

SGLang 有自己的 --numa-node 配置,会主动把 rank 绑到指定 NUMA node(PIN 死)。冲突来了

  • SGLang 说:这个 rank 就吃 Node 0,别乱跑
  • 内核 numa_balancing 说:我来「优化」一下

内核会做几件破坏性的事:

  1. 周期性把页面标 PROT_NONE:即使你已经手动 pin 好了,下一次访问必触发 page fault
  2. Page fault 引起大 stall:GPU 计算 → 触发 page fault → 内核处理 → GPU 继续(延迟数百 μs 到 ms)
  3. 随机页面迁移:pinned memory 被迁移 → GPU 之前记录的 DMA 地址失效

为什么这个 bug 特别难缠

论文里这句话说到点子上了:

“In multi-node multi-GPU deployments, these gaps appear at random positions across ranks, and each inter-rank synchronization is bottlenecked by the slowest rank.”

关键组合拳

  • 随机性:numa_balancing 是周期性异步扫描,不同 rank、不同时刻会随机中招
  • 全局同步放大伤害:LLM 推理每层 attention 之后都要 all_reduce最慢那个 rank 拖累所有人
1
2
3
所有 rank 计算 ──► 集合通信同步 ──► 下一层

最慢那个 rank 拖累所有人

如果 8 个 rank 里有 1 个刚好被 page fault 命中,那 8 个 rank 全都得等它——尾延迟灾难。对一次前向 60+ 层,只要偶尔一层有一个 rank 中招,整个 latency 就翻倍。

为啥关掉直接 +10%

  1. SGLang 已经手动 pin 好了 NUMA,调度器安排得明明白白
  2. numa_balancing 是「帮倒忙」,它假设进程没做优化,主动帮你搬,反而破坏了原有布局
  3. 代价是 hard cost:page fault、TLB shootdown、页面迁移都是不可避免的 CPU 侧开销

关掉后:内核不再主动折腾 → 页表稳定 → 没有意外 page fault → GPU kernel 之间没有随机 gap → 集合通信不再被拖尾 → 直接 +10%。

一句话总结:SGLang 已经手动做了 NUMA pin,numa_balancing 在这基础上做二次干扰,通过 page fault + 页面迁移引入随机 stall,被 GPU 集合通信同步放大成尾延迟灾难,关掉就直接 +10%。


6. Prefill 为什么也要开 MTP

默认直觉是「MTP 是加速 decode 的,跟 prefill 有啥关系」,但正是这个直觉害惨了 agentic 场景。

MTP 是什么

MTP (Multi-Token Prediction) = 多 token 预测。MiMo-V2.5 系列原生支持 3 层 MTP

1
2
3
4
主模型: 预测第 t+1 个 token
MTP Layer 1: 同时预测第 t+2 个 token
MTP Layer 2: 同时预测第 t+3 个 token
MTP Layer 3: 同时预测第 t+4 个 token

核心思想:一次 forward 出 4 个 token,然后主模型验证,接受的部分直接输出。理想情况下 decode 速度接近 4×。

关键前提:MTP 层本身有参数、有 KV cache,需要**跟着上下文一起「预热」**才能预测准。

原始问题:Prefill 不跑 MTP 的后果

默认实现:Prefill 阶段只跑主模型,跳过 MTP 层。

问题:MTP 层需要看历史 context 才能预测。如果 prefill 跳过 MTP:

1
2
3
4
5
6
7
Prefill 结束时:
主模型 KV cache: [完整 prompt 的 KV] ✅
MTP 层 KV cache: [空 / invalid] ❌

Decode 第 1 个 token:
主模型: ✅
MTP 层: context 是空的,瞎猜 → 接受率极低 ❌

MTP 需要多久「热身」:论文说是 128 个 token。因为 MTP 层需要逐个 decode 的 token 慢慢累积 KV。

这段时间 MTP 基本白搭:主模型每步生成 1 个 token,MTP 提议的 3 个 token 因为 context 不足几乎全被拒,有效加速 ≈ 1×。

为什么 agentic 场景特别惨

论文这句话点破了关键:

Since agentic scenarios involve mostly short output sequences, this limitation significantly limited MTP’s effective speedup.”

Agentic 场景的输出特征:

1
2
3
一次工具调用响应: ~20-50 tokens
一次 function call: ~30-80 tokens
一次思考步骤: ~50-150 tokens

几乎所有输出都在 128 token 之内——这正好是 MTP 需要热身的窗口。结果:

1
2
理论加速比: 3× (3层 MTP)
Agentic 场景实际加速比: ~1× (刚热身完就结束了)

等于 MTP 层白装了

解法:Prefill 阶段也跑 MTP

小米的改动很直接:prefill 时也让 MTP 层参与前向

1
2
3
4
5
6
7
Prefill 结束时:
主模型 KV cache: [完整 prompt 的 KV] ✅
MTP 层 KV cache: [完整 prompt 的 MTP KV] ✅ ← 新增

Decode 第 1 个 token:
主模型: ✅
MTP 层: 已经用整个 prompt 预热过,接受率立即很高 ✅

效果(论文数据)

  • 0-128 token 加速: 2.3× ← 从近 1× 提升到 2.3×,agentic 场景直接受益
  • 128-256 token 加速: 1.5×

工程适配

论文说:

“By introducing MTP support during prefill with dedicated adaptations and optimizations for HiCache L2/L3…”

关键点:MTP 层有自己的 KV cache,之前 HiCache 那套 offloading 系统只管主模型的 KV。要 prefill 阶段跑 MTP,就得:

  1. MTP KV 也要进 HiCache:L1 (GPU) ↔ L2 (Host) ↔ L3 (GCache)
  2. Prefix cache 得覆盖 MTP KV:前缀树节点要标 MTP 层的状态
  3. 传输和存储成本:MTP 有 3 层,KVCache 总量变多,L2/L3 得扛住

Trade-off 账

维度 不开 prefill MTP 开 prefill MTP
Prefill 计算 70 层 70 + 3 = 73 层(多 4%)
Decode 0-128 tokens 加速 ~1× 2.3×
Agentic 场景总吞吐 显著提升

核心逻辑用 prefill 的少量额外开销,换 decode 从第 1 个 token 就享受 MTP 加速。Agentic 场景短输出多,这笔账非常划算。

具体场景:Agent 调用工具,输入 5K token,输出 60 token

1
2
3
没开 prefill MTP: 100ms (prefill) + 1200ms (decode) = 1300ms
开了 prefill MTP: 104ms (prefill) + 520ms (decode) = 624ms
省 52%

一句话总结:Prefill 开 MTP 是为了给 MTP 层的 KV cache 做「上下文预热」,避免 decode 初期因为 MTP 无历史导致接受率极低。Agentic 场景输出短,正好都落在这个「未热身窗口」里,所以收益特别大——0-128 token 加速直接从 ~1× 提升到 2.3×。


总结

这篇论文的核心贡献不是某个单点优化,而是把 Hybrid SWA + MoE + 多模态这套组合架构从理论拉到生产的完整工程实践。本文聚焦在 SWA 相关 KVCache 这条主线,拆解了几个关键问题:

  1. SWA 存储:物理分离 Full/SWA pool,SWA 严格 O(W),前缀树双索引做透明分层
  2. SWA offload:Host 侧独立 SWA pool + pinned memory,搬运只按 mask 走有效位置
  3. 共享前缀命中:默认严格 O(W) 会伪命中,小米用「SWA 快照保留」策略换真实命中率
  4. numa_balancing:操作系统自动优化 vs 手动 NUMA pin 的冲突,关掉 +10%
  5. Prefill MTP:给 MTP 层 KV cache 做上下文预热,agentic 短输出场景 0-128 token 从 ~1× 提升到 2.3×

可迁移的启示

  • 架构红利不会自动兑现,工程系统必须跟着架构一起改造
  • Hybrid SWA 的价值在多用户共享场景,但需要专门的缓存策略才能兑现
  • 推理系统的 pipeline 各阶段不是孤立的,跨阶段依赖的组件需要协同优化
  • **系统级坑(numa_balancing、THP、CPU governor)**对已经手动优化过的高性能场景是纯干扰,该关就关

参考

  • 论文:MiMo Team, Xiaomi. Full-Pipeline Inference Optimization for MiMo-V2.5 Series: Pushing Hybrid SWA Efficiency to the Limit. arXiv:2607.13095, 2026.
  • 模型:Xiaomi MiMo Team. MiMo-V2.5 / MiMo-V2.5-Pro. HuggingFace, 2026.

本文基于 MiMo-V2.5 论文、个人笔记《MiMo-V2.5 SWA offload 方法》整理。

DeepSeek-V4 的架构图有一个特点:本身就是一份并行性说明书。画成并行分支的模块,基本都是在告诉系统实现者”这里有 overlap 的空间”。本文梳理 V4 中可 overlap 的算子对,对照论文承诺与 SGLang 实现现状。


一、Overlap 的三个前提

  1. 算法独立 — 两条路径无数据依赖(论文画成并行分支)
  2. 资源不冲突 — 不争用同一 SM 或同一 buffer
  3. 工程可兜住 — CUDA stream / multi-kernel launch 的硬件支持

DeepSeek 的特殊性在于:MoE + MHC(Multi-Head Compressor)双重稀疏结构,天然产生多条独立路径


二、Attention 分支内的 Overlap

2.1 wqkv_a 与 Compressor 的 Overlap

数据依赖

1
2
3
4
5
6
7
wqkv_a:  输入 = hidden (x)

输出 = q_a + kv_a + k_rope

compressor: 输入 = hidden (x) + past KV (来自 KV cache)

输出 = summary_K_4 + summary_K_128

两者都读 hidden,但 compressor 不依赖 wqkv_a 的输出,可以完全并行。

资源 pattern 互补

Kernel 性质 瓶颈资源
wqkv_a GEMM (M=4096, N=2112, K=7168) 中 GEMM MFU ~30-40%,CU 大量闲置
compressor flash_c4 / flash_c128 memory-bound HBM 带宽吃满,MFMA 单元闲置

一个 compute-bound 倾向,一个 memory-bound 倾向,硬件资源不冲突。

L2 Cache 共享

hidden 的大小:[4096, 7168] BF16 = 56 MB,接近 MI355X 的 L2 cache(32 MB)。

  • 串行wqkv_a 读完 hiddencompressor 再读——大概率已被 q_b 挤出 L2,又一次从 HBM 读
  • 并行:两个 kernel 同时从 HBM 拉 hidden,第二次的 cache line 是免费的

2.2 Indexer 的 Overlap:只依赖 q_a,不依赖 q_b

这是 V4 论文里 “MLA Co-Design with Indexer” 的核心设计。

Indexer 两条路径的依赖

1
2
3
4
5
6
7
Indexer K-side:  W_K^I · x  →  RoPE  →  summary_K

只需要 hidden (x),不需要 wqkv_a 的输出

Indexer Q-side: W_Q^I · q_lora → Hadamard + RoPE

只需要 q_lora (q_a 的输出),不需要 q_b 的输出

q_b 是 critical path 上最重的 GEMM(929 µs),indexer 不等 q_b 完成就可以启动

为什么 Q-side 用 q_lora 而不是 hidden

V4 论文的算法决策:

  • 实验结论:召回率无差异
  • 系统收益:每层省一次 [T, 7168] → [T, ...] 的大 GEMM
  • 代价:Q-side 多一个 q_lora_ready event 依赖(K-side 完全自由)

代码里精确实现了这个依赖:

1
2
3
4
# K-side:不依赖 q_lora,可以立刻启动
stream_indexer.wait_stream(current_stream)
# Q-side:只等 q_lora_ready,不等 q_b
# (在 indexer 内部用 q_lora_ready 精确控制)

2.3 三者在时间轴上的 Overlap

1
2
3
4
5
6
7
8
9
10
11
12
13
时间轴(prefill 4096 tokens,典型层):

0 µs 77 µs 83 µs 1012 µs
│ wqkv_a │ q_a │ q_b (929 µs) │
│─────────┴─────┴───────────────────────────────┤
│ │ │
│ └──► indexer K-side + Q-side │ ← 不等 q_b
│ (q_lora_ready 触发)
│ │
│ └──► compressor (c4 + c128) │ ← 完全独立
│ (只等 x)

└──► kv_write (等 wqkv_a 输出切片) │

端到端时间 = max(q_b, indexer, compressor, kv_write) = ~1012 µs(由 q_b 主导)

串行时 = wqkv_a + q_a + q_b + indexer + compressor + kv_write = ~2290 µs


三、MoE 分支内的 Overlap

3.1 Shared Expert 与 Routed Expert

两条 expert 路径完全独立,SGLang 用 alt_stream 实现,是已实现的最好案例。

3.2 MoE Wave 模型:计算与通信 Overlap

Wave 分 chunk 乒乓:dispatch → compute → combine 流水,掩盖全量通信延迟。V4 论文 Fig.5 的时序图明确画出了这一点。

3.3 Combine 通信与 Shared Expert 尾部计算

Combine(NVLink all-to-all)不需要 shared expert 的完整输出,可以和 shared expert 的 down_proj 最后一部分 overlap。SGLang 目前串行等待,未利用。


四、为什么 Multi-Stream Overlap 能赚到时间

CUDA stream 是 GPU 上的 FIFO 工作队列;stream 本身只给调度器自由度,性能要靠”互补的资源占用”赚出来。

机制 1:单 Kernel 资源利用率低(最主要)

q_b 把计算压满但 HBM 闲;swa_scatter 把 HBM 压满但计算闲。不同 stream 让它们同时跑,各用各的硬件资源。

机制 2:小 Kernel Grid 填不满 SM

MI355X 有 256 个 CU,但 trace 里很多 kernel 的 grid 很小(rocprim cumsum 只有 1 个 CU,占用率 0.4%)。多 stream 让调度器把多个小 kernel 同时塞进 CU。

机制 3:隐藏 CPU Launch Overhead

每次 hipLaunchKernel 在 host 端要 ~3-5 µs。Trace 里 ~30000 个 kernel × 4 µs ≈ 120 ms 纯 launch 开销。多 stream 让 host 预先 enqueue,GPU 不饥饿。

机制 4:同源输入的 L2 Cache 共享

wqkv_acompressor 都读 hidden,并行时第二次读直接 cache hit,省 ~20% 输入带宽。

反向判断:什么情况下多 Stream 不赚

场景 多 stream 是否有用 原因
两个高 MFU GEMM(都跑 70%+) SM 都被占满,只是 timesharing
两个 memory-bound op 读不同数据 HBM 总带宽是上限
两个 op 有数据依赖(A → B) 必须串行
一个 compute-bound + 一个 memory-bound 经典 case
一个大 GEMM + 一群 µs 级小 kernel ✅✅ 赚得最多

五、工程实现:SGLANG_OPT_USE_MULTI_STREAM_OVERLAP

这个关键优化靠一个环境变量控制,不在 --help 里:

1
2
3
4
5
# 开启
SGLANG_OPT_USE_MULTI_STREAM_OVERLAP=1 python -m sglang.launch_server ...

# 验证
python -c "import os; print(os.environ.get('SGLANG_OPT_USE_MULTI_STREAM_OVERLAP'))"

读取位置在 sglang/srt/layers/deepseek_v4.py_forward_prepare_multi_stream 方法入口,通过 os.environ.get() 判断,默认关闭。

实测效果

在 decode 阶段(batch size 中等,61 层 DeepSeek-V4),开启后整体 forward 有 ~3.5% 的端到端提升

状态 forward latency (decode) 相对提升
关闭(串行) 基准
开启(multi-stream overlap) -3.5%

提升主要来自:

  • q_b GEMM(929 µs)与 indexer + compressor 并行:decode 时 q_b 仍是瓶颈,但 indexer K-side 和 compressor 可以完全隐藏在其执行期间
  • 小 kernel(swa_scatter、cumsum 等)与主体 GEMM 并行:这些 µs 级 kernel 在串行时被 q_b 的 launch gap 放大,overlap 后基本被吸收

注意:3.5% 是 decode 阶段的收益。prefill 阶段因为 q_b 的 GEMM 更大(M=4096),overlap 的相对收益会被稀释,但绝对时间节省更显著(每层 ~1.28 ms,61 层 ~78 ms)。

没开时,整个 _forward_prepare_multi_stream 退化成串行,论文里”sparse attention 模块在 prefill 阶段近乎 free”的 claim 直接失效。


六、对照论文图:承诺 vs 现状

论文图中的结构 论文章节 承诺的并行性 SGLang 状态
Indexer ∥ wqkv_a §3.2.2 双分支并行 ✅ 算法并行,⚠️ 需 SGLANG_OPT_USE_MULTI_STREAM_OVERLAP=1 开启
Compressor c4 ∥ c128 §3.2.1 双尺度评分并行 ❌ 串行
Shared ∥ Routed Expert §3.3 双 expert 路径并行 ✅ 已实现
Wave dispatch/compute/combine §3.3.3 乒乓 overlap ✅ 已实现
KV store ∥ next layer indexer §3.2 层间 pipeline ❌ 未做

七、总结

DeepSeek-V4 论文在算法层面为 overlap 留了很大空间,尤其是 Fig.3(MHC 结构)和 Fig.5(Wave 时序图)。当前 SGLang 实现只吃到了 Shared Expert / Wave 的部分,仍有明显优化空间。

实测数据验证了 overlap 的价值:

  • decode 阶段:开启 SGLANG_OPT_USE_MULTI_STREAM_OVERLAP,整体 forward 提升 ~3.5%
  • prefill 阶段:每层节省 1.28 ms,61 层共 **78 ms**,sparse attention 模块接近”free”的目标

更一般的启示:算法论文画依赖图时多想一步系统实现,系统实现时多对照论文的并行性承诺。


参考:DeepSeek-V4 论文(arXiv:2606.02405)§3.2 Attention Mechanism, §3.3 Expert Mixture, §3.3.3 Dynamic Expert Routing & Wave Scheduling

背景

训练超长序列 LLM 时,单卡显存放不下完整的 KV cache,需要对序列维度做并行(Sequence Parallelism)。目前主流有两种方案:

  • Ulysses(DeepSpeed-Ulysses):两次 All-to-All,按 head 切分
  • Ring Attention:环形传递 KV blocks,分块计算

两者目标相同,但设计假设完全不同。本文从原理、通信量、代码实现到架构选择动机,做完整对比。


一、核心思想对比

Ulysses Attention

1
2
3
Step1: [N/P, h, d] ──All-to-All──→ [N, h/P, d]
Step2: 本地做标准 Attention(每张卡拿到全序列、部分 head)
Step3: 输出 ──All-to-All──→ 回到 [N/P, h, d]

关键:两次 All-to-All 把”序列切”转成”head 切”,每张卡对全序列做部分 head 的 attention。

Ring Attention

1
2
3
4
Round 0: attn(Q_i, K_i, V_i)           ← 本地 attention
Round 1: 收到 K_{i-1}, V_{i-1} → attn(Q_i, K_{i-1}, V_{i-1})
...
Round P-1: 收到所有 KV → online softmax 合并结果

关键:KV 沿环形传递,每张卡只存自己的 KV chunk,计算和通信完全重叠。


二、通信量数学推导

Ulysses:All-to-All 的 $(P-1)/P^2$

Ulysses 输入是 [N/P, h, d](序列已被切,head 完整),输出是 [N, h/P, d](序列完整,head 被切)。

每 GPU 发送 P-1 个 chunk,每个 chunk 大小:

$$
\frac{N}{P} \times \frac{h}{P} \times d = \frac{N \cdot h \cdot d}{P^2}
$$

总发送量:

$$
\text{send per GPU} = (P-1) \times \frac{N \cdot h \cdot d}{P^2} = N \cdot h \cdot d \cdot \frac{P-1}{P^2}
$$

两个 1/P 的来源:

  • 第一个 1/P:序列维度已被切(N/P
  • 第二个 1/P:head 维度再切一次(h/P

Ring:P2P 的 $(P-1)/P$

Ring 只传 KV(Q 不动),每轮传 (N/P) × d_kv

$$
\text{send per GPU} = (P-1) \times \frac{N}{P} \times d_{kv} = N \cdot d_{kv} \cdot \frac{P-1}{P}
$$

只有一个 1/P(序列切分),没有 head 切分。

对比(P=8, h=64, d=128, d_kv=576)

方法 通信量/GPU 比值
Ulysses (MHA) $N × 64 × 128 × 7/64 = N × 896$ 1.8×
Ring (MHA KV) $N × 128×2 × 7/8 = N × 224$
Ring (MLA) $N × 576 × 7/8 = N × 504$

注意:MHA 的 KV 是 h × d × 2 = 16384/token,MLA 的 c_kv 只有 576/token,这是后续分析的关键。


三、MLA 对 Ulysses 的致命问题

MLA 的 KV cache 结构

DeepSeek-V3/V4 使用 MLA(Multi-head Latent Attention),KV cache 不是 multi-head 的:

1
2
3
c_kv [seq, 512]     ← 压缩的共享 latent
k_rope [seq, 64] ← RoPE 位置编码部分
合计:576 dim/token(vs MHA 的 16384 dim/token)

只有 1 个”头”,无法按 head 切分。

DeepSpeed Ulysses 代码验证

1
2
3
4
# deepspeed/sequence/layer.py 核心逻辑
q = _SeqAllToAll.apply(group, query, scatter_idx=2, gather_idx=0)
k = _SeqAllToAll.apply(group, key, scatter_idx=2, gather_idx=0) # K 和 Q 对称处理
v = _SeqAllToAll.apply(group, value, scatter_idx=2, gather_idx=0)

Q、K、V 走完全相同的 All-to-All 路径,没有”KV 走 All-Gather”的分支。当 num_kv_heads=1 时:

1
2
1 KV head / 4 GPUs → [1, 0, 0, 0]
GPU 1-3:没有 KV head → 无法计算 ❌

Ulysses SP 上限 = num_kv_heads

模型 num_kv_heads Ulysses SP 上限 实际可用性
DiT (视觉) 32 (MHA) 32
LLaMA-3 8 (GQA) 8 ⚠️ 受限
DeepSeek-V3/V4 1 (MLA) 1 ❌ 不可用

四、MLA 场景:Ring vs Ulysses 精确对比

通信量(P=8)

Ring Attention(MLA)

  • 只传 c_kv:每轮 (N/8) × 576,共 7 轮
  • 总计:$N × 504$ / GPU

Ulysses(MLA,Q All-to-All + KV All-Gather)

  • Q All-to-All:$7 × (N/8) × 64 × 512 = N × 28,672$
  • KV All-Gather:$7 × (N/8) × 576 = N × 504$
  • Output reverse:$N × 28,672$
  • 总计:$N × 57,848$ / GPU

Ring 比 Ulysses 省 115×。

根本原因:MLA 的非对称性——Q 巨大(32768 dim/token)、KV 极小(576 dim/token)。Ring 只移动小的 KV,Ulysses 被迫移动巨大的 Q。

缩放性对比

P MHA Ulysses MLA Ring MLA 比 MHA 省
4 $N × 6,144$ $N × 432$ 14.2×
8 $N × 3,584$ $N × 504$ 7.1×
16 $N × 1,792$ $N × 540$ 3.3×
64 $N × 428$ $N × 567$ 0.75×(MHA 反超)

交叉点:$P = 4hd/d_{kv} ≈ 57$,但 Ulysses SP 上限 = 64,所以实践中 MLA Ring 几乎总是更优


五、为什么 Ulysses 在视觉模型流行,文本 LLM 不用?

四个结构性原因

1. KV head 数量趋势

1
2
3
4
2020: MHA (GPT-3) num_kv_heads = 96  → Ulysses 随便用
2022: MQA (PaLM) num_kv_heads = 1 → Ulysses 废了
2023: GQA (LLaMA-2) num_kv_heads = 8 → Ulysses 受限
2024: MLA (DeepSeek-V3) num_kv_heads = 1 → Ulysses 废了

文本 LLM 全面转向 GQA/MLA 压缩 KV heads,Ulysses 的前提条件被釜底抽薪。

2. 文本 LLM 的 TP 已经做了同样的事

1
2
Tensor Parallelism: 按 head 切权重 → 本地算 attention → All-Reduce
Ulysses: 按 head 切数据 → 本地算 attention → All-to-All

本质重叠,TP 已经切了 head 之后,Ulysses 没有额外收益。

3. 推理阶段 Decode 占 80% 时间

1
2
Prefill: ~20% 时间(可以序列并行)
Decode: ~80% 时间(Q=1 token,序列并行无用)

Ulysses 对 Decode 完全无用。视觉扩散模型没有 Decode 阶段,全程受益。

4. 视频生成 token 数极高

1
Sora 级别: (64×64) × 120 frames = 491,520 tokens

必须序列并行,且 DiT 用标准 MHA(32 heads),Ulysses 完美适配。

一句话总结

Ulysses 在视觉模型流行 = MHA(head 够多)+ 无 decode + 没被 TP 覆盖 + 序列极长。文本 LLM 四条全占不到。


六、DeepSeek 的实际选择

训练:全程 Ring/CP

DeepSeek-V3 技术报告显示,长上下文扩展(32K→128K)用的是 Ring Attention(Context Parallel),不是 Ulysses:

  • MLA 只有 1 个 KV latent → Ulysses 物理上不可用
  • KV 只有 576 dim → Ring 通信量本来就小
  • Ring 的通信-计算 overlap → 长序列时通信几乎免费

推理:SGLang 的 CP 配置

1
2
3
--enable-nsa-prefill-context-parallel
--attn-cp-size 8
--nsa-prefill-cp-mode round-robin-split

选择 Ring 风格 CP 的原因:

  • MLA 的 KV 只有 1 个共享 latent,切不了 head
  • KV 极小(576 dim),Ring 通信可接受
  • round-robin 切分 tokens 均衡负载

七、LoRA 压缩能让 Ulysses 复活吗?

思路:在压缩态做通信

MLA 的 Q 路径:hidden(7168) → q_lora(1536) → expand → Q[128, 576]

能否在 q_lora 维度(1536)做 All-Gather,而非展开后的 Q(73728)?

通信量更新(LoRA 压缩版)

通信 数据 P=8 总量
q_lora All-Gather $N × 1536 × 7/8$ $N × 1,344$
c_kv All-Gather $N × 576 × 7/8$ $N × 504$
o_lora Reduce-Scatter $N × 1024 × 7/8$ $N × 896$
总计 $N × 2,744$

vs 原始 Ulysses(展开态):$N × 57,848$ → 压缩 21×

但仍然比 Ring 大 5.4×

1
2
Ring:              N × 504
Ulysses (LoRA): N × 2,744 ← Q 和 O 的压缩态还是要传

根本原因:Ring 的 Q 和 O 根本不过网络,留在本地计算。Ulysses 无论怎么压缩,都要把 Q/O(或其压缩态)在网络上搬一次,这是架构级差距,压缩只能缩小、不能逆转。


八、全景决策树

1
2
3
4
5
6
7
8
9
10
num_kv_heads >= SP_size?

├─ YES (DiT, ViT, 视觉 MHA)
│ └─ Ulysses ✅(All-to-All,通信量 1/P²)

└─ NO (GQA=8, MLA=1, 文本 LLM)
├─ 训练:Ring/CP ✅
└─ 推理:
├─ Prefill:Ring/CP
└─ Decode:partial attn + All-Reduce

总结

Ulysses 的核心优势是 1/P² 通信缩放,但前提是 num_kv_heads ≥ SP。MLA 把 KV 压缩到 1 个 latent,直接废掉这个前提。Ring 只传 KV(576 dim),Q 留本地,反而成了 MLA 的最优搭档。

这不是巧合——MLA 和 Ring CP 是刻意协同设计,不是将就。


相关阅读:DeepSeek-V3 技术报告、DeepSpeed-Ulysses 论文(arXiv:2309.14509)、Ring Attention 论文(arXiv:2310.01889)

从残差到 mHC:一条清晰的三代演进路

如果你训练过深层 Transformer,一定对残差连接(Residual Connection)不陌生。它是 ResNet 留下的最重要的遗产之一,也是今天所有大语言模型的标配。

但残差连接有个根本问题:它只有一条信息流

DeepSeek 在 2025 年连发两篇论文,把这个问题彻底讲透了。故事的三步是:

  1. Residual(2015)x_{l+1} = x_l + F(x_l) — 一条流,稳定但表达弱
  2. HC(2024):把残差流拓宽到 n 条并行流,表达能力暴涨,但训练直接崩
  3. mHC(2025):给 HC 加上数学约束,又强又稳

这篇文章把这三代讲清楚,重点放在 mHC 上。


残差连接的瓶颈:一条流不够用

标准 Transformer 里,每一层的计算是这样的:

1
2
x = x + Attention(Norm(x))  # 残差连接
x = x + MoE(Norm(x)) # 残差连接

Layer 0 到 Layer 42 的所有信息,全都挤在同一个向量里做加法。

浅层的语法信息、中层的语义信息、深层的推理信息,互相覆盖、互相干扰。这是残差连接的天花板。

HC(Hyper-Connections) 提出了一个很自然的想法:

为什么不把 1 条流扩成 n 条并行流,让不同深度的信息走不同的”车道”?

HC 的公式长这样:

$$
X_{l+1} = B_l X_l + C_l F(A_l X_l)
$$

  • A_l:把 n 条流压缩成 1 条,送给 Attention/MoE
  • F:正常的层计算
  • C_l:把 F 的输出扩回 n 条流
  • B_l:n×n 矩阵,控制 n 条流之间怎么混合(残差项)

n=4 的时候,FLOPs 几乎没增加(F 只算 1 次),但信息容量翻了 4 倍。

听起来完美,对吧?


HC 的致命缺陷:训练会崩

问题出在 B_l 上。

HC 里的 B_l 是完全可学习的,没有任何约束。当模型叠到 60 层、参数量到万亿级时,这个无约束的矩阵会出问题:

多层复合后,信号被无限放大或衰减。

数学上,HC 跨层递归展开后得到:

$$
x_L = (\prod B_i) x_l + \sum (\prod B_j) C_i^T F(…)
$$

$Π B_i$ 是多层 $B$ 矩阵的复合。因为 $B$ 无约束,这个复合矩阵的谱范数可以远大于 1,也可以接近 0。

实测结果(DeepSeek 27B 实验):
HC 的复合映射增益(Amax Gain Magnitude)达到 ~3000,意味着信号在前向传播中可以被放大 3000 倍。训练到第 12k 步,loss 直接 spike,梯度范数爆炸。

HC 的作者(ByteDance Seed 团队)在 ICLR 2025 发表了这个想法,但没有解决稳定性问题。


mHC:给 HC 加上”安全阀”

DeepSeek 团队在 2025 年提出的 mHC(Manifold-Constrained Hyper-Connections),只做了一件事:

把 HC 中的残差混合矩阵 B_l 约束在双随机矩阵流形上。

什么是双随机矩阵?

一个 n×n 矩阵,满足:

  • 所有元素非负
  • 每行之和等于 1
  • 每列之和等于 1

这样的矩阵,谱范数 永远 ≤ 1。也就是说,它永远不会放大信号。

而且,双随机矩阵在乘法下是封闭的:两个双随机矩阵相乘,结果还是双随机矩阵。这意味着叠 100 层,复合映射仍然稳定。

一句话:mHC = HC + 双随机约束,恢复了残差连接的恒等映射性质,同时保留了多流的表达能力。


怎么把矩阵变成双随机?Sinkhorn-Knopp 算法

mHC 的核心算法是 Sinkhorn-Knopp 迭代,操作很简单:

1
2
3
4
5
6
7
8
9
# 输入: B_raw (4×4, 可能是负数)
M = exp(B_raw) # 先保证非负

# 交替行列归一化 20 次
for _ in range(20):
M = M / M.sum(dim=1, keepdim=True) # 行归一化
M = M / M.sum(dim=0, keepdim=True) # 列归一化

# 输出: B (4×4 双随机矩阵)

为什么第一步用 exp()?因为 B_raw 是神经网络输出的 logit,可能有负数。直接归一化会产生负权重,违反”非负”约束。exp() 把任意实数映射到正数,完美解决。

20 次迭代是实验得出的经验值,足够收敛,又不会太慢。


4 条流到底解决了什么问题?

很多人问:为什么是 4 条流?不是 2 条或 8 条?

三个核心作用:

1. 梯度高速公路

普通残差的梯度路径是 43 个乘法项的连乘,容易消失或爆炸。
4 条流提供了 4 条并行的梯度路径,一条堵了还有其他。类似 DenseNet 的思路,但高效得多(F 只算 1 次)。

2. 信息分离存储

1
2
3
4
stream_0: 可能专注浅层信息(位置、语法)
stream_1: 可能专注中层信息(语义、实体关系)
stream_2: 可能专注深层信息(推理链)
stream_3: 可能做"工作记忆"(当前层临时计算)

Layer 5 学到的特征,通过 B 矩阵的合理混合,可以在 Layer 30 仍然清晰可用。而在 1 条流里,这个特征早就被 25 次加法淹没了。

3. 动态容量分配

B 矩阵是输入依赖的(由当前层的 hidden state 动态生成):

  • 简单 token(”the”, “a”):B ≈ 单位矩阵,4 条流基本不混合,省计算
  • 复杂 token(需要深度推理):B 大幅混合,让更多层的信息参与

这是 1 条流做不到的——普通残差对所有 token 都是 x + F(x)

为什么选 n=4?
论文实测:n=4 是性价比最优的点。n=2 效果不够,n=8 的 Sinkhorn 开销(8×8 矩阵 × 20 次迭代)显著增加,但收益递减。


工程优化:为什么只多 6.7% 开销?

mHC 看起来很重:每行代码都有矩阵运算、Sinkhorn 迭代、4 倍 hidden state……

但 DeepSeek 做了三件事,把开销压到了 仅 +6.7%

1. Kernel Fusion(算子融合)

用 TileLang 把整个 mHC 计算——RMSNorm + 线性投影 + Sigmoid + Sinkhorn——融合成一个 CUDA kernel,减少内存读写。

2. Selective Recomputing(选择性重计算)

前向时只保存每块的第一个层输入,反向时重新计算 mHC 的中间激活。最优块大小由公式给出:

$L_r^* \approx \sqrt{\frac{nL}{n+2}}$

3. Overlapping with DualPipe

把 mHC 的计算和流水线通信重叠。在 DualPipe 调度中,F_post,res kernel 放到高优先级流,避免阻塞 All-to-All 通信。


实验效果:又强又稳

DeepSeek 用 27B MoE 模型做了对比实验:

指标 Baseline HC mHC
训练稳定性 稳定 loss spike @12k步 稳定
最终 loss 下降 +0.021 vs baseline +0.027 vs baseline
复合映射增益 1.0 ~3000 ~1.6
训练开销 0% +?% +6.7%

下游任务(BBH / DROP / GSM8K 等 8 个 benchmark)上,mHC 全面超过 baseline 和 HC。

最重要的结论: mHC 让 1.6T 参数、60+ 层、MoE 模型能够稳定训练——这是普通 HC 做不到的。


一句话总结

mHC = HC + 双随机矩阵约束,是残差连接的最终进化形态,让万亿参数 MoE 能稳定训练。

如果你正在设计超大模型,mHC 是目前最值得考虑的残差连接方式。它不改 FLOPs,不改层计算,只改残差流的拓扑结构——但就是这一点,让深层训练从”可能崩”变成了”一定稳”。


参考资料:

  • Hyper-Connections, Zhu et al., ByteDance Seed, ICLR 2025, arXiv:2409.19606
  • mHC: Manifold-Constrained Hyper-Connections, Xie et al., DeepSeek, 2025, arXiv:2512.24880
  • DeepSeek-V4 Technical Report, 2026

写于 2026-06-02,基于 DeepSeek-V4-Pro 代码和 mHC 原始论文整理。

一句话概括

MegaMoE = 把 EP(Expert Parallelism)中的 All-to-All 通信和 MoE 计算融合到一个 kernel 里,让 NVLink 传输和 Tensor Core 计算时间上完全重叠。

传统 EP 里通信和计算串行,GPU 利用率只有 50-60%;MegaMoE 通过 Symmetric Memory + 细粒度 Scheduler + 单 Kernel 状态机,把利用率推到接近 100%。


传统 EP MoE vs MegaMoE

传统 EP(如 DeepEP)

1
2
3
4
时间 →

[dispatch all-to-all] → 等完 → [GEMM1] → [SwiGLU] → [GEMM2] → 等完 → [combine all-to-all]
NVLink 通信 空闲 计算 计算 计算 空闲 NVLink 通信

问题:通信和计算串行,GPU 要么在算要么在传,利用率约 50-60%。

MegaMoE

1
2
3
4
5
6
时间 →

┌───────────────────── 一个 Mega Kernel ─────────────────────┐
│ NVLink dispatch ←→ GEMM1 ←→ SwiGLU ←→ GEMM2 ←→ NVLink combine │
│ (通信和计算同时进行,流水线式重叠) │
└──────────────────────────────────────────────────────────────┘

通信隐藏在计算背后,GPU 利用率接近 100%。


传统 DeepEP 的 Overlap 能力分析

常见误解:”传统 DeepEP dispatch/combine 是无法和 MoE 计算 overlap 的”

更准确的说法:传统 DeepEP 可以通过 low-latency kernel + 多 stream + two-batch overlap 实现部分重叠,但需要框架手工编排,重叠率有限,对 decode 小 batch 几乎无效。

DeepEP 的两种模式

DeepEP 提供了两套 kernel,overlap 能力完全不同:

DeepEP 模式 能否 overlap 原因
Normal kernel(高吞吐) 基本不能 dispatch/combine 占满 SM 跑带宽,SM 被通信占住,GEMM 抢不到资源
Low-latency kernel(纯 RDMA) 可以 只用少量 SM 发 RDMA,大部分 SM 空着,可以留给 GEMM

你提到的”无法 overlapped”更接近 Normal kernel 的情况。

框架层怎么强行 overlap(传统方案)

SGLang / TRT-LLM 的做法是 micro-batch 切分 + 多 stream

1
2
3
4
时间线 ──────────────────────────────────►
stream A: [dispatch chunk0] [combine chunk0]
stream B: [GEMM chunk0] [dispatch chunk1] [combine chunk1]
stream A: [GEMM chunk1]
  • chunk0 在算的时候,chunk1 在通信
  • 需要 2 条 CUDA stream + event 同步
  • 代价:buffer 翻倍,GEMM tile 变小导致 Tensor Core 利用率下降

这就是 TBO (Two-Batch Overlap),SGLang 里有 --enable-two-batch-overlap 开关。

为什么传统方案”能 overlap 但不彻底”

四个硬限制:

  1. Kernel 边界 = 同步点
    每次 launch 有 5–10μs 开销,小 batch 下被 launch 开销主导

  2. SM 资源争抢
    Normal dispatch 想跑满 NVLink 就要用很多 SM,GEMM 也想要全部 SM,互相挤压

  3. 依赖链串行

    1
    dispatch(x) → GEMM1 → SwiGLU → GEMM2 → combine

    单个 token 的依赖必须串行,overlap 只能靠不同 micro-batch 之间的并行

  4. Decode 阶段 batch 小
    再切一半更糟,GEMM 效率暴跌

重叠率实测

方案 重叠率 适用场景
DeepEP normal 0%(串行) Prefill 大 batch
DeepEP low-latency + TBO 30–60% Prefill 中等 batch
MegaMoE 接近 100% 所有场景(包括 decode 小 batch)

对性能标定(profiler)的意义

如果在 perf 数据库里标注 MoE 路径,至少分三档:

  1. DeepEP normal:通信和计算串行(吞吐高,延迟差)
  2. DeepEP low-latency + TBO:部分 overlap(需要 enable_two_batch_overlap=true
  3. MegaMoE:原生 overlap,不依赖切 batch

这三档在 SLA 模型里会给出显著不同的 ITL/TTFT,混在一起标定会让 profiler 拟合出错。


核心实现机制

1. Symmetric Memory(对称内存)

传统 all-to-all:

1
GPU_0 → cudaMemcpyAsync → GPU_1(需要显式同步,退出 kernel)

Symmetric Memory:

1
2
3
所有 GPU 共享一片 symmetric buffer
GPU_0 直接写入 GPU_1 的 symmetric buffer(RDMA over NVLink)
GPU_1 看到数据就开始算,不需要全局 barrier

关键:不需要等所有 token 传完再开始算。传一批就算一批(streaming)。

实现层级

  • 硬件层:NVSwitch 全连接,任何 GPU pair 都有直达链路,memory controller 识别远端地址自动走 NVLink
  • 驱动/Runtime 层:CUDA VMM 把远端 GPU 物理内存映射到本地虚拟地址(cuMemMap / nvshmem_ptr
  • 应用层:MegaMoE 直接用普通指针读写远端 buffer,一条 PTX load/store 指令完成

注意:只有 symmetric allocation 的那块 buffer 地址一致,不是整个 GPU 显存对称。

2. 单 Kernel 状态机

传统方式多个 kernel 之间有全局同步点,无法实现 fine-grained overlap:

1
kernel_1 (dispatch) → kernel 边界(全局 barrier)→ kernel_2 (GEMM1) → ...

MegaMoE 的 scheduler/mega_moe.cuh 用一个状态机让所有 SM 在同一个 kernel 内自主领任务:

1
2
3
4
5
6
7
8
9
10
enum class BlockPhase { None, Linear1, Linear2 };

while (true) {
auto block = scheduler.get_next_block();
switch (block.phase) {
case Linear1: /* 做 W1/W3 GEMM */ break;
case Linear2: /* 做 W2 GEMM */ break;
case None: return;
}
}

没有 kernel 边界,没有全局 barrier,通信和计算可以交错到 warp 级别。

3. Wave-Based Expert Processing

256 个 expert 太多,无法同时处理(Token Pool 显存有限),分 wave 分批:

1
2
3
4
Wave 0: expert [0..31]   → 处理完 → 释放 pool
Wave 1: expert [32..63] → 处理完 → 释放 pool
...
Wave 7: expert [224..255]

每个 wave 内:

  • 所有 SM 抢 L1 block(gate + up projection)
  • SwiGLU 在 SMEM 中直接做激活
  • 所有 SM 抢 L2 block(down projection)

kNumExpertsPerWave 由 SMEM 容量、BLOCK_M、pool_capacity 共同决定,让 wave 内 expert 数刚好喂饱所有 SM 又不撑爆 Token Pool。

4. Arrival Count 轮询(细粒度 Overlap 的关键)

1
2
3
4
// 轮询等待:不是等所有 token 到齐,而是等够一个 BLOCK_M 就开始算
while (volatile_count < kNumSMs * kNumRanks) {
// spin-wait 直到其他 SM/rank 报告完成
}
  • 不需要全局 barrier(不需要所有 token 都收完才开始)
  • 收到一个 block 的 token 就立刻开始算
  • 这就是 streaming overlap 的本质

5. L1/L2 交错

1
2
Expert_0: [L1 block 0][L1 block 1][L1 block 2] → SwiGLU → [L2 block 0][L2 block 1]
Expert_1: [L1 block 0][L1 block 1] → SwiGLU → [L2 block 0]

L1 结果留在 Token Pool 中,L2 直接从 pool 读(在 L2 cache 中热),一个 wave 结束后才释放 pool 空间。

对应关系:

MegaMoE 术语 权重矩阵 操作 含义
L1 (Linear 1) W1 + W3(拼在一起) gate + up projection hidden → 2×intermediate
SwiGLU 无权重 SiLU(gate) × up 激活函数
L2 (Linear 2) W2 down projection intermediate → hidden

W1 和 W3 拼在一起做一次 GEMM,可以复用 TMA 加载的 activation tile,减少 GMEM → SMEM 搬运。


SM/Warp 级别并行

GPU 硬件层次:

1
2
3
4
GPU (整块芯片)
└── SM (Streaming Multiprocessor) × 132 (H20) / 192 (B200)
└── Warp × 最多 64 resident / SM
└── Thread × 32 / Warp(SIMT 锁步执行)

在 MegaMoE 中,一个 Mega Kernel launch 占满所有 SM,每个 SM 内部:

1
2
Warp Group A(若干 warp):负责通信(polling + NVLink store/load)
Warp Group B(若干 warp):负责计算(Tensor Core MMA)

Warp Group A 和 B 在 SM 内部并行执行——这就是”通信和计算在 warp 级别交错”的含义。


为什么必须 GB200/GB300

MegaMoE 依赖三个 GB200/GB300 独有(或大幅增强)的硬件能力:

硬件特性 GB200/GB300 H20 影响
GPU 间拓扑 72 GPU NVSwitch 全连接 8 GPU NVLink mesh H20 无法 symmetric memory 跨机
Symmetric Memory ✅ kernel 内 load/store 远端 ❌ 需要 NCCL API 无法做 in-kernel 通信
NVSwitch Reduction ✅ 硬件完成 combine ❌ SM 做 reduce combine 占 SM 资源
FP4 Tensor Core ✅ 原生指令(SM100) ❌ 软件模拟(SM90) 计算慢 → 通信占比低 → overlap 收益有限
NVLink 带宽/GPU 900 GB/s(18 ports) 450 GB/s(NV18 mesh) H20 通信带宽是瓶颈

为什么计算足够快才值得 overlap

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
假设: 通信时间 = T_comm, 计算时间 = T_compute

情况 A — 计算快 (FP4 Tensor Core):
T_compute = 2ms, T_comm = 5ms
不 overlap: 7ms
overlap: max(2, 5) = 5ms ← 省 28%

情况 B — 计算慢 (BF16 软件模拟):
T_compute = 20ms, T_comm = 5ms
不 overlap: 25ms
overlap: max(20, 5) = 20ms ← 只省 20%

而且 overlap 不是免费的:
- 需要 SM 专门做通信 warp(减少计算并行度)
- Symmetric Memory 的 load/store 比本地慢 10-20x
- Scheduler 状态机有开销

只有当计算非常快、通信成为瓶颈时,overlap 的净收益才大于开销。

和 DeepEP / flashinfer_mxfp4 的区别

MegaMoE vs DeepEP

DeepEP MegaMoE
通信粒度 整个 batch dispatch 完再算 tile 级别通信和计算交错
overlap 方式 不同 CUDA stream 异步(kernel 间) 同一个 kernel 内 tile 级 pipeline
权重格式 任意 Runner backend 都能接 只支持 DeepGEMM 格式(权重需预转换)
内存模型 普通 GPU buffer NVSHMEM 对称内存 (SymmBuffer)
batch 上限 无(取决于显存) 有(NUM_MAX_TOKENS_PER_RANK

MegaMoE vs flashinfer_mxfp4

flashinfer_mxfp4 MegaMoE
解决的问题 单卡 MoE 计算效率 多卡 EP 通信+计算 overlap
Fuse 范围 GEMM1 + SwiGLU + GEMM2 dispatch + GEMM1 + SwiGLU + GEMM2 + combine
通信 不涉及(单卡或假设已 gather 好) 融合 All-to-All 通信
适用场景 TP 模式(experts 复制在每卡) EP 模式(experts 分布在不同卡)
硬件要求 任何 Hopper/Blackwell 需要 NVLink + Symmetric Memory(GB200+)
状态 已发布,可用 开发中(DeepGEMM PR #304)

两者互补:flashinfer_mxfp4 管单卡 MoE kernel 效率,MegaMoE 管 EP 多卡通信+计算融合。


Python API 调用流程

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
# 1. 多进程初始化 + 分配对称内存
buffer = deep_gemm.get_symm_buffer_for_mega_moe(
group, # 进程组 (torch.distributed)
num_experts=256,
num_max_tokens_per_rank=8192,
num_topk=6,
hidden=4096,
intermediate_hidden=2048
)

# 2. 权重预处理(FP4 packing + 布局变换)
l1_transformed, l2_transformed = deep_gemm.transform_weights_for_mega_moe(
l1_weights, # W1, W3 (gate + up)
l2_weights # W2 (down)
)

# 3. 加载输入到 symmetric buffer
buffer.x[:N].copy_(x_fp8) # FP8 激活
buffer.x_sf[:N].copy_(x_sf) # 激活的 scale factor
buffer.topk_idx[:N].copy_(topk_idx) # 路由结果
buffer.topk_weights[:N].copy_(topk_weights)

# 4. 一次调用,完成 dispatch + GEMM1 + SwiGLU + GEMM2 + combine
y = torch.empty((N, 4096), dtype=torch.bfloat16, device='cuda')
deep_gemm.fp8_fp4_mega_moe(y, l1_transformed, l2_transformed, buffer)
# y 就是最终结果,全程只有这一个 kernel launch

总结:MegaMoE 的关键创新

机制 作用 依赖硬件
Symmetric Memory kernel 内直接 load/store 远端 GPU NVSwitch 全连接
Arrival Count 轮询 收到一个 block 就开始算,不等全部 低延迟 NVLink
Wave-based Scheduler SM 自主领任务,无全局 barrier SM 数量充足(Blackwell 192 SM)
L1/L2 交错 L1 结果留在 pool,L2 直接读 大 L2 cache(Blackwell 100MB+)
FP8×FP4 GEMM 计算足够快,才值得 overlap 通信 原生 FP4 Tensor Core(SM100)

一句话:MegaMoE = Symmetric Memory(消除通信 barrier)+ 细粒度 Scheduler(一个 block 到了就算)+ 单 Kernel 状态机(避免 kernel launch 开销),把 EP MoE 的 GPU 利用率从 ~60% 推到接近 100%。


局限性与适用边界

  • 权重格式锁定:只支持 DeepGEMM 的 FP4 权重布局,需要预转换
  • 激活函数限制:依赖 SwiGLU 无状态特性,L1 结果可直接给 L2 用;若 MoE 层有跨 token 依赖(如全局 norm),pipeline 会断
  • 硬件绑定:Symmetric Memory + NVSwitch 全连接 + 原生 FP4 三个条件缺一不可,H20 及更早硬件无法使用
  • 跨机不支持:Symmetric Memory 只在单机 NVSwitch 域内有效,跨机仍需回退到 DeepEP
  • 负载不均影响fetch_expert_recv_count() 的 spin-wait 在 expert 负载不均时可能导致 SM 空转

参考

背景

DeepSeek-V4 系列模型(V4 / V4-Pro)在 MoE(Mixture of Experts)层大量使用 MXFP4 量化权重,配合 FP8/BF16 激活 实现高效推理。在 SGLang / vLLM 中,flashinfer_mxfp4 是一个专为这种量化格式设计的 MoE Runner Backend。

本文将基于 H20 (SM90 Hopper) 实测数据和源码分析,完整解析 FlashInfer MXFP4 的实现原理,以及与 Marlin MoE 的架构差异。


MXFP4 量化格式

格式定义

MXFP4 是 OCP(Open Compute Project)标准的 4-bit 浮点格式:

1
2
每个权重元素: [1 sign | 2 exponent | 1 mantissa] → 4-bit 浮点数
每 32 个元素: 共享一个 8-bit 指数缩放因子 (E8M0 block scale)

与 INT4 的对比

维度 MXFP4 INT4 (Marlin)
表示方式 浮点 (E2M1) 定点整数
动态范围 大(指数缩放) 小(线性,需精确校准)
缩放粒度 每 32 元素(block scale) 每 128 元素(group quantization)
解量化 value = fp4 × 2^(scale - 127) value = int4 × scale[channel]

优势:浮点表示让 MXFP4 在训练和推理中都有更好的数值稳定性。


FlashInfer MXFP4 架构

调用链(H20 / SM90)

1
2
3
4
5
6
7
8
9
10
11
SGLang --moe-runner-backend flashinfer_mxfp4

Python: trtllm_fp4_block_scale_moe()

C++: get_cutlass_fused_moe_module() → 加载预编译 .so

CUTLASS 3.x: cute::gemm::kernel::DefaultGemmUniversal

├─ TMA (Tensor Memory Accelerator) → 异步搬权重 tile 到 SMEM
├─ WGMMA (Warpgroup MMA) → SM90 用 FP16 TC 模拟 FP4
└─ Warp Specialization → Producer warpgroup 专搬数据,Consumer warpgroup 专计算

核心特性

  1. Expert-First 遍历:按 expert 分组 token,而非按 token 遍历 expert
  2. 全 Fuse:W1 + W3 + SwiGLU + W2 四个步骤 fuse 成一个 kernel
  3. TMA + Warp Specialization:数据搬运和计算 overlap
  4. Grouped GEMM:一次 launch 处理所有 256 个 expert

完整 Pipeline(DeepSeek-V4 为例)

模型配置

1
2
3
4
5
DeepSeek-V4-Flash:
- 256 experts, topk=6
- hidden_dim=4096, moe_intermediate=2048
- 权重: MXFP4 (4-bit), 激活: BF16
- Router: Sigmoid + Group TopK (非 Softmax)

第一阶段:Routing(路由选择)

1
2
3
4
5
6
7
8
9
10
11
12
# 文件: trtllm_fused_moe_routing_deepseek.cu

# 输入: router_logits [num_tokens, 256]
scores = torch.sigmoid(router_logits + bias) # DeepSeek-V4 用 Sigmoid, 非 Softmax

# Group TopK: 256 experts → 8 组 × 32
group_scores = scores.view(num_tokens, 8, 32).max(dim=-1) # 每组最高分
top_groups = group_scores.topk(4).indices # 选 4 组 (128 experts)

# 从 128 个 expert 中选 top-6
topk_indices = ... # [num_tokens, 6]
topk_weights = ... # 归一化后的权重

为什么用 Sigmoid?

  • 独立评分(非 Softmax 的零和游戏)
  • 多 expert 可同时激活
  • 训练更稳定

第二阶段:Gather(Token 重排)

1
2
3
4
# 文件: cutlass_fused_moe_kernels.cuh → expandInputRowsKernel

# 输入: hidden_states [num_tokens, 4096] (按 token 顺序)
# 输出: expanded_input [num_tokens × 6, 4096] (按 expert 连续排列)

目的:把路由到同一个 expert 的 token 收集到一起,形成连续内存,供后续 GEMM 使用。

1
2
3
4
5
6
7
8
9
10
原始 (token 顺序):
[t0 | t1 | t2 | t3]

路由结果:
expert_3: [t0, t1, t3]
expert_7: [t0, t2]

Gather 后 (expert 顺序):
[e3: t0, t1, t3 | e7: t0, t2 | ...]
↑ 连续 ↑ 连续

TMA 优势:连续输入让 TMA 一次加载整个 tile,达到满带宽。


第三阶段:Compute(Fused GEMM)

1
# 文件: cutlass_fused_moe_kernels.cuh → CUTLASS Grouped GEMM

这是最核心的部分,W1 + W3 + SwiGLU + W2 全部 fuse

Step 3a: GEMM1 (Gate + Up Projection)

1
2
3
4
5
6
对每个 expert_i:
M = expert_i 分到的 token 数 (例如 3)
input_i = expanded_input[offset : offset+M] # [M, 4096]

# W13 = [W1; W3] 拼接,一次 GEMM 算出
gate_up = input_i @ W13_expert_i # [M, 4096] → split → [M, 2048] + [M, 2048]

MXFP4 权重存储

1
2
3
// W1_expert_i: stored as [2048, 2048] int8
// 实际逻辑形状: [2048, 4096] FP4 (每 byte 存 2 个 FP4 元素)
// Scale: [2048, 64] float8_e8m0 (每 32 个元素一个 scale)

Step 3b: SwiGLU Activation

1
2
3
# 在 GEMM1 输出还在 SMEM/寄存器时立即执行(fuse 关键)
intermediate = SiLU(gate) * up # [M, 2048]
# DeepSeek-V4 还有 swiglu_limit=10.0 的 clamp

Step 3c: GEMM2 (Down Projection)

1
2
output_i = intermediate @ W2_expert_i  # [M, 4096]
# W2_expert_i: stored as [4096, 1024] int8 → 逻辑 [4096, 2048] FP4

CUTLASS Grouped GEMM 调度

1
2
3
4
5
6
7
8
9
10
11
12
// 一次 kernel launch 处理所有 256 个 expert
cutlass::gemm::kernel::GroupedGemmKernel<...>::run();

// 内部调度:
// SM 空闲 → 从 problem list 取下一个 expert
// expert_i 的 M=0 (没 token) → 跳过
// expert_j 的 M=12 → 分配 SM 算 GEMM

// TMA + Warp Specialization:
// Warpgroup 0 (Producer): TMA_LOAD(weight_tile, desc)
// Warpgroup 1 (Consumer): WGMMA(tile_C, tile_A, tile_B)
// Producer 搬下一个 expert 的权重时,Consumer 正在算当前 expert

Pipeline 示意(H20)

1
2
3
4
Cycle 0-100:  Producer 搬 expert_3 权重,Consumer 算 expert_1 (上一轮)
Cycle 100-200: Consumer 算 expert_3 (权重已到),Producer 搬 expert_7
Cycle 200+: Consumer 算 expert_7,Producer 搬 expert_5
→ 数据搬运和计算完全 overlap!

第四阶段:Scatter(结果写回)

1
2
3
4
# 文件: cutlass_fused_moe_kernels.cuh → finalizeMoeRoutingKernel

# 输入: expert_outputs [num_tokens × 6, 4096] (按 expert 排列)
# 输出: final_output [num_tokens, 4096] (恢复 token 原始顺序)

加权求和

1
2
3
for token_j:
output[j] = Σ (weight_k × expert_output[permute_map[j][k]])
k=1..6

示例

1
2
token_0 → expert_3 (0.6), expert_7 (0.4)
output[0] = 0.6 × out_0_e3 + 0.4 × out_0_e7

为什么叫 Scatter?

  • expert_3 的输出: [out_0_e3, out_1_e3, out_3_e3]
  • 需要写回: output[0], output[1], output[3](不连续!)
  • 这就是 scatter(分散写入),与 gather(聚集读取)相对

性能实测(H20, DeepSeek-V4-Flash, TP=4)

Benchmark 结果

Concurrency FlashInfer MXFP4 Marlin 差距
conc=1 20.2 tok/s/GPU 20.5 tok/s/GPU ≈持平
conc=2 42.1 tok/s/GPU 28.1 tok/s/GPU +50%
conc=4 32.5 tok/s/GPU 31.8 tok/s/GPU ≈持平

阶段耗时分解(估)

阶段 占比 说明
Routing ~2% 纯 element-wise,快
Gather ~3% 内存拷贝
GEMM1 (W13) ~42% 计算密集
Activation ~1% fused 在 GEMM1 后
GEMM2 (W2) ~48% 计算密集
Scatter ~4% 加权求和 + 写回

GEMM 占 90%,所以 TMA + Grouped GEMM 对总体性能影响最大。

为什么 conc=2 差距最大?

  1. Expert-First + Grouped GEMM:一次 launch 处理所有 expert,TMA 预取下一个 tile
  2. 中等 batch:每个 expert 分到 2-3 个 token → GEMM tile 够大,WGMMA 饱和
  3. Warp Specialization:Producer 搬数据,Consumer 计算,完美 overlap

Marlin 的劣势

  • Token-First:每个 token 的 topk expert 不同 → 权重反复换入换出
  • 无 TMA:用 cp.async 加载,线程要参与地址计算 → SM 利用率 < 50%
  • 无法全 fuse:中间结果 (gate, up, mid) 要写回 GMEM

为什么 conc=4 差距消失?

1
2
3
4
5
conc=4: prefill=50k tokens
→ Prefill 时间: ~1500ms (占 85%+)
→ Decode MoE 时间: ~265ms (占 15%-)

即使 MoE 快 50%,总体提升也只有: 265ms × 50% / 1765ms ≈ 7.5%

Prefill 主导后,MoE decode 的优化被稀释。


与 Marlin MoE 的架构对比

核心差异

维度 FlashInfer MXFP4 Marlin (SGLang)
量化格式 MXFP4 (浮点4位) INT4 / MXFP4
遍历顺序 Expert-First Token-First
融合程度 W1+W3+SwiGLU+W2 全 fuse W1+W3 部分 fuse,W2 单独
数据加载 TMA (硬件 DMA) cp.async / LDG (线程参与)
调度方式 Warp Specialization (Producer/Consumer) 单 warp 既搬又算
Kernel 架构 CUTLASS 3.x Grouped GEMM 手写 CUDA,Ampere 设计
TMA 利用 ✅ 异步 tile 预取 ❌ 无 TMA (用 cp.async)
多 expert 并行 Grouped GEMM 一次 launch 逐 expert 串行或小并行

Expert-First vs Token-First

Token-First (Marlin)

1
2
3
4
5
for token_0:
expert_3: W1 → W3 → SwiGLU → W2 ← 写回 GMEM
expert_7: W1 → W3 → SwiGLU → W2 ← 换权重
for token_1:
... # 重复上述,权重反复换入换出

Expert-First (FlashInfer)

1
2
3
4
5
6
7
for expert_3:
# 固定权重 W1/W3/W2,常驻 SMEM
gate_up = input_e3 @ W13_e3 # 连续 token 一起算
mid = SiLU(gate) * up # 在寄存器/SMEM 直接算
output = mid @ W2_e3 # 中间结果不落地
for expert_7:
... # 换一次权重,算所有路由到它的 token

为什么 Expert-First 能全 fuse?

  • 固定权重 → W1/W3/W2 常驻 SMEM,不换出
  • 变 batch → 同一 expert 的多个 token 连续处理
  • 中间结果 (gate, up, mid) 全在寄存器/SMEM,不写 GMEM

TMA + WGMMA 的硬件优势

TMA (Tensor Memory Accelerator)

SM90 (Hopper) 引入的硬件 DMA 引擎

  • 异步搬数据从 HBM 到 SMEM,不占用 CUDA core
  • 一条 tcgen05.1d 指令触发,后台独立运行
  • 支持多维 tensor 描述符(TMA Descriptor)

vs 传统 LDG

1
2
3
4
5
6
7
8
LDG (传统):
线程发 load → 等 HBM (~400 cycles) → 数据到 SMEM → 继续算
→ 线程 stall,SM 有空闲

TMA:
线程发 TMA_LOAD → 立即返回 → 可以去 setup 下一个 tile
→ TMA 后台搬,线程继续算上一轮结果
→ SM 利用率 80%+

WGMMA (Warpgroup MMA)

SM90 引入的 warpgroup 级矩阵乘指令

  • 一个 warpgroup = 4 个 warp = 128 线程
  • 直接读写 TMEM(Tensor Memory,SM90 新增寄存器文件)
  • 比传统 mma.sync 高一个抽象层级

SM100 (Blackwell) 的进化

1
2
SM90: TMA → SMEM → WGMMA → TMEM
SM100: TMA_GATHER4 → TMEM (跳过 SMEM) → UTCMMA (原生 FP4 TC)

Blackwell 的 TMA_GATHER4 还能一次加载 4 个不连续 token(专为 MoE 的 sparse routing 优化)。


总结

FlashInfer MXFP4 的核心优势

  1. Expert-First + 全 Fuse:中间结果不落地 GMEM,省 ~2× hidden_dim × batch 带宽
  2. TMA + Warp Specialization:搬运和计算 overlap,SM 利用率最大化
  3. Grouped GEMM:一次 launch 处理 256 experts,launch overhead 最小化
  4. Blackwell 未来兼容:cute_dsl_fused_moe_nvfp4 已支持 SM100 原生 FP4 TC

适用场景

场景 推荐后端
H20/H100, conc=2~8 ✅ FlashInfer MXFP4
A100, 或 conc=1 调试 Marlin
Blackwell (B200) FlashInfer NVFP4 (cute_dsl)
追求最大吞吐 FlashInfer MXFP4

一句话

FlashInfer MXFP4 在 Hopper 上通过 Expert-First + TMA + Warp Specialization + 全 Fuse,把 MoE 推理的瓶颈从内存带宽转移到了计算,实现了中等并发下 50% 的吞吐提升


参考资料


如果你对 MoE 推理优化感兴趣,可以看看我之前写的 Marlin MoE Kernel 深度分析,对比两种实现的差异。

什么是 WideEP

WideEP(Wide Expert Parallelism)不是 SGLang 里的一个具体模块,而是社区对一种大规模 MoE 部署模式的俗称。官方文档里叫 Large-Scale EP

核心思路只有一句话:

把 Expert Parallelism 撑得很宽(几十到上百张 GPU),配合 DP Attention 消除通信,用 DeepEP 解决 All-to-All 瓶颈。


为什么需要 WideEP

以 DeepSeek-V3/V4 为例:256 个 routed expert,FP8 权重约 25GB,加上 dense 层、KV Cache,单机 8 卡跑高吞吐 decode 场景显存压力极大。

传统方案有三个瓶颈:

方案 瓶颈
TP(Tensor Parallel) 所有卡存全部专家,每卡 ~25GB,显存放不下
普通 EP(8 卡) 每卡 ~32 个专家,batch 小,HBM 带宽利用率低
无 DeepEP 的 EP NCCL All-to-All 延迟高,跨机扩展不了

WideEP 的解法:用更多卡分摊专家 → 每卡只存 4 个专家 → batch 变大 → HBM 利用率高;DP Attention → 省掉 KV Cache 冗余。


WideEP 的核心组成

1. Expert Parallelism(MoE 层)

1
256 experts ÷ 64 张卡 = 每卡 4 个专家

token 经过 gate 计算 TopK 后,通过 DeepEP All-to-All 精确 dispatch 到目标专家所在卡,算出结果再 combine 回来。

2. DP Attention(Attention 层)

传统 TP 下 Attention 需要 AllReduce 汇总,跨 64 卡延迟巨大。WideEP 改用 Data Parallel

  • 每卡存一份完整的 Attention 权重(QKV/O projection)
  • 每卡独立算自己的 batch,KV Cache 独立不共享
  • 零通信,省掉 AllReduce

为什么放得下?DeepSeek-V3/V4 的 dense 层(Attention + Gate + Shared Expert + Norm)总共才 ~4GB,复制 64 份完全没问题。

3. DeepEP All-to-All

DeepEP 是 DeepSeek 开源的 MoE 专用通信库,解决了 NCCL 做 All-to-All 的三大问题:

特性 NCCL All-to-All DeepEP
通信粒度 全量 按 token routing 结果精确发送
SM Overlap 不支持 支持(通信和 GEMM 重叠)
Low Latency 模式 有(decode 专用)
跨节点 一视同仁 NVLink + IB 分层优化

各层并行方式一览

并行方式 通信
Embedding / RMSNorm 每卡完整副本(DP)
Attention QKV / Output 每卡完整副本
Attention compute DP,每卡独立 batch
KV Cache 每卡独立
MoE Gate 每卡完整副本
MoE Experts EP,每卡 4 个 DeepEP All-to-All
Shared Expert / LM Head 每卡完整副本

关键设计取舍:dense 层(4GB)复制 64 份 → 省掉通信;MoE 层(25GB)必须 EP 切分 → 靠 DeepEP 通信。


关于 --tp 64 的误解

很多人看到 --tp 64 以为权重被切了 64 份,其实不是。

在 SGLang 里,tp 只是定义一个包含 N 张卡的通信组,组内的并行方式由其他参数决定:

1
2
3
4
5
6
# parallel_state.py
moe_tp_size = tp_size // moe_ep_size // moe_dp_size
# tp=64, ep=64, dp=1 → moe_tp_size = 64 // 64 // 1 = 1

attn_tp_size = tp_size // attn_dp_size
# tp=64, dp=64 → attn_tp_size = 64 // 64 = 1

两个都是 1,意味着没有任何层做真正的 Tensor Parallel 切分。tp=64 的真实含义是”这 64 张卡组成一个协作集群”,具体怎么分工由 ep_sizedp_size 等参数决定。


SGLang 配置示例

基础 WideEP(单机 8 卡)

1
2
3
4
5
6
7
python -m sglang.launch_server \
--model-path deepseek-ai/DeepSeek-V3 \
--tp 8 \
--ep-size 8 \
--moe-a2a-backend deepep \
--deepep-mode auto \
--moe-runner-backend deep_gemm

完整 WideEP(DP Attention + 低延迟模式)

1
2
3
4
5
6
7
8
9
10
11
python -m sglang.launch_server \
--model-path deepseek-ai/DeepSeek-V3 \
--tp 64 \
--dp-size 64 \
--enable-dp-attention \
--ep-size 64 \
--moe-a2a-backend deepep \
--deepep-mode low_latency \
--enable-two-batch-overlap \
--enable-eplb \
--mem-fraction-static 0.85

多节点部署(2 × 8 卡)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
# Node 0
python -m sglang.launch_server \
--model-path deepseek-ai/DeepSeek-V3 \
--tp 16 --ep-size 16 \
--nnodes 2 --node-rank 0 \
--dist-init-addr <MASTER_IP>:29500 \
--moe-a2a-backend deepep

# Node 1
python -m sglang.launch_server \
--model-path deepseek-ai/DeepSeek-V3 \
--tp 16 --ep-size 16 \
--nnodes 2 --node-rank 1 \
--dist-init-addr <MASTER_IP>:29500 \
--moe-a2a-backend deepep

两个重要的优化特性

TBO(Two-Batch Overlap)

把请求拆成 micro-batch,在 attention 和 dispatch/combine 之间穿插执行:

1
--enable-two-batch-overlap
  • 吞吐量最高提升
  • 零额外显存开销
  • 原理:attention 算完 → yield → dispatch 和下一个 batch 的 attention 并行

EPLB(Expert Parallelism Load Balancer)

运行时收集 expert 激活统计,动态调整 expert 放置/复制策略,解决负载不均:

1
--enable-eplb

配合大 batch size(如 --max-running-requests 128)效果更好。


与之前方案的对比

WideEP vs 无 DeepEP 的 EP(moe-a2a-backend=none

a2a_backend=none WideEP(deepep
token 跨卡 ❌ 不移动 ✅ 精确 dispatch
每卡计算 只算本地 expert + AllReduce 只算本地 expert(无 reduce)
可扩展性 最多 8~16 卡 64~128+ 卡
通信算子 AllReduce DeepEP All-to-All
KV Cache TP 共享(冗余) DP 独立(不冗余)

none 模式下,token 不跨卡,每卡算完自己的 expert 后 AllReduce 汇总——浪费计算(路由到的 expert 不在本卡就白算),跨机扩展不了。


通信 Backend 选择

Backend 描述 约束
none(默认) 用 AllReduce/AllGather 支持 ep < tp(混合 EP+TP)
deepep DeepEP 通信库 必须 ep == tp
mooncake 弹性推理 + RDMA 必须 ep == tp
mori AMD ROCm 优化 必须 ep == tp,仅支持 normal 模式
flashinfer FlashInfer All-to-All 无特殊约束

总结

WideEP = DP Attention + EP MoE(DeepEP All-to-All),本质上就是:

  1. MoE 层:专家分散到 64+ 卡,DeepEP dispatch/combine
  2. Dense 层:每卡完整副本,DP 独立算,零通信
  3. tp=64:只是通信组大小,不代表 TP 切分(moe_tp_size=1

DeepEP 出来之前,大规模 EP 做不了——NCCL All-to-All 延迟太高。DeepEP 解决了这个瓶颈,WideEP 才成为实用的部署模式。


参考文档:

引言

DeepSeek-V4 引入 MXFP4 量化后,MoE 层的计算效率成为推理性能的关键瓶颈。SGLang 的 Marlin Runner Backend 专门针对 INT4/MXFP4 量化权重优化 MoE 的 GEMM 计算。本文深入分析其实现原理、数据流以及设计权衡。

MoE 层的计算流程

一个标准的 MoE 层包含四个阶段:

1
2
3
4
5
6
7
8
9
10
11
12
13
MoE Layer

├── ① Router → topk_ids, topk_weights

├── ② Token Dispatch (All-to-All) ← A2A backend

├── ③ Expert Compute ← Runner Backend (本文主角)
│ ├── W1 GEMM (gate + up 融合)
│ ├── SwiGLU 激活
│ ├── W2 GEMM (down projection)
│ └── Weighted sum reduce

└── ④ Token Combine ← A2A backend

Runner Backend 只负责第 ③ 步,不同 backend 优化的是同一组计算在不同量化格式下的执行效率:

Runner Backend 权重格式 核心 kernel 适用场景
triton FP8/BF16 Triton fused MoE 通用 FP8
deep_gemm FP8 block DeepGemm DeepSeek-V3 FP8
marlin INT4/MXFP4 Marlin WNA16 GPTQ/AWQ/MXFP4 量化
flashinfer_mxfp4 MXFP4 FlashInfer MXFP4 DeepSeek-V4 MXFP4

Marlin MoE 完整数据流

以 DeepSeek-V4 配置为例:hidden_size=7168, intermediate_size=3072, num_experts=384, topk=6, 权重为 MXFP4 格式。

Step 1: Block Size 启发式选择

1
2
3
4
5
6
7
8
9
M, K = hidden_states.shape  # M=tokens, K=7168
E = w1.shape[0] # num_experts=384
N = w2.shape[1] * 16 # 3072 (Marlin 打包因子 16)
topk = topk_ids.shape[1] # 6

# 启发式选择 M 方向分块大小
for block_size_m in [8, 16, 32, 48, 64]:
if M * topk / E / block_size_m < 0.9:
break

设计意图

  • Context 阶段(M 大):用大 block(64),让每个 block 尽量填满
  • Generation 阶段(M 小):用小 block(8),避免最后一个 block 填充率过低

潜在问题:假设 token 在专家间均匀分布,但 Router 可能有热点。热点专家需要多个 block,最后一个 block 可能只填了 10%,浪费计算。

Step 2: Token-Expert 对齐

1
2
3
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
topk_ids, block_size_m, global_num_experts
)

作用

  1. expert_id 对 token 排序
  2. block_size_m 对齐(padding),让每个专家的 token 数是 block 的倍数
  3. 返回排序后的 token 索引和每个 block 对应的 expert_id

目的:让后续 Marlin GEMM kernel 能用 block-sparse 方式高效执行。

Step 3: W1 GEMM (gate + up 融合)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
intermediate_cache1 = moe_wna16_marlin_gemm(
hidden_states, # [M, 7168]
intermediate_cache1, # output buffer
w1, # [E, 7168/pack, 2*3072*pack]
w1_scale, # 量化 scale
sorted_token_ids, # token → expert 映射
expert_ids, # 每个 block 的 expert_id
num_tokens_post_padded, # padding 后的总 token 数
topk_weights, # 路由权重
moe_block_size=block_size_m,
top_k=topk,
mul_topk_weights=False, # ← W1 阶段不乘路由权重
size_m=M, size_n=2*N, size_k=K,
)

关键点

  • mul_topk_weights=False:W1 输出是纯 GEMM 结果,不乘权重
  • W1 权重是 gate 和 up 融合 的:[E, K/pack, 2*N*pack]
  • 输出 [M*topk, 2*N]:前半 N 是 gate,后半 N 是 up

为什么 W1 不乘权重?

假设 token 的 top-2 是 Expert 3 (weight=0.6) 和 Expert 7 (weight=0.4):

1
2
3
4
5
6
7
8
W1 输出(不乘权重):
[Expert 3 的 gate/up, Expert 7 的 gate/up]

经过 SwiGLU:
[Expert 3 的 activated, Expert 7 的 activated]

W2 输出(乘权重):
Expert 3 输出 × 0.6 + Expert 7 输出 × 0.4

如果 W1 就乘了权重,SwiGLU 激活函数作用在”已经被压缩的信号”上,精度会下降。

Step 4: SwiGLU + Clamp

1
2
3
4
5
if clamp_limit is not None:
# DeepSeek-V4: swiglu_limit=10.0
swiglu_limit_func(intermediate_cache2, intermediate_cache1, clamp_limit)
else:
silu_and_mul(intermediate_cache1, intermediate_cache2)

输入 [M*topk, 2*N] 拆成 gate 和 up:

1
2
3
4
5
gate = input[:, :N]      # SwiGLU gate branch
up = input[:, N:] # SwiGLU up branch

# SiLU(clamp(gate)) * clamp(up)
output = F.silu(torch.clamp(gate, max=10)) * torch.clamp(up, -10, 10)

Clamp 的作用:防止激活值爆炸,避免量化误差被放大。DeepSeek-V4 的 swiglu_limit=10.0 是经验值。

Step 5: W2 GEMM (down projection)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
intermediate_cache3 = moe_wna16_marlin_gemm(
intermediate_cache2, # [M*topk, 3072]
intermediate_cache3, # output buffer
w2, # [E, 3072/pack, 7168*pack]
w2_scale,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
topk_weights,
moe_block_size=block_size_m,
top_k=1, # ← 注意这里是 1
mul_topk_weights=True, # ← W2 阶段乘路由权重
size_m=M*topk, size_n=K, size_k=N,
).view(-1, topk, K) # reshape 为 [M, topk, K]

关键变化

  • top_k=1:W2 的输入已经是展开的 M×topk 个 token,每个只对应 1 个专家
  • mul_topk_weights=True:在 GEMM 内部就乘上路由权重,融合减少一次 memory pass

Step 6: 跨专家归约

1
2
3
4
5
6
if is_mxfp4_marlin:
# MXFP4:用 torch.sum(atomic_add 不支持 BF16 on SM < 90)
output = torch.sum(intermediate_cache3, dim=1) # [M, topk, K] → [M, K]
else:
# 普通 INT4:用 CUDA kernel 做 sum + scale
moe_sum_reduce(intermediate_cache3, output, routed_scaling_factor)

routed_scaling_factor=2.5(DeepSeek-V4)在这里乘进去。

MXFP4 特殊处理

DeepSeek-V4 的专家权重是 MXFP4 格式(4-bit 分块量化),Marlin kernel 对此有特殊路径:

1
2
3
4
5
6
7
is_mxfp4_marlin = (
num_bits == 4
and w1_zeros is None # 无 zero point
and w2_zeros is None
and w1_scale.dtype == torch.float8_e8m0 # scale 是 E8M0 格式
and w2_scale.dtype == torch.float8_e8m0
)

MXFP4 + E8M0 scale 要求激活必须是 BF16(不支持 FP16):

1
2
if is_mxfp4_marlin and hidden_states.dtype == torch.float16:
marlin_hidden_states = hidden_states.to(torch.bfloat16) # 强制转 BF16

MXFP4 特殊处理

  • 不用 atomic_add(因为 BF16 的 atomic_add 在 SM < 90 上不支持)
  • 改用 torch.sum 替代归约
  • Scale 是 E8M0 格式(8-bit exponent-only),和普通 INT4 的 per-tensor/per-channel scale 不同

内存布局优化

1
2
3
4
5
6
7
# 两个中间 buffer 共享底层存储
intermediate_cache13 = torch.empty(
(M * topk * max(2*N, K),), # 按 W1 和 W2 输出中较大的分配
device=device, dtype=dtype,
)
intermediate_cache1 = intermediate_cache13[:M*topk*2*N].view(-1, 2*N)
intermediate_cache3 = intermediate_cache13[:M*topk*K].view(-1, K)

节省显存:W1 的输出 [M*topk, 2*N] 在被 SwiGLU 消费后就不需要了,W2 可以写入同一区域。

但有个细节:W2 GEMM 的输入是 intermediate_cache2(SwiGLU 输出,维度 N),不是 intermediate_cache1(维度 2N)。所以实际复用链是:

1
2
3
intermediate_cache1 [M*topk, 2*N] → (SwiGLU) → intermediate_cache2 [M*topk, N]

intermediate_cache3 [M*topk, K] ← (W2 GEMM 输出)

intermediate_cache2 如果能复用 intermediate_cache1 的后半部分(up 分支在 SwiGLU 后就没用了),可以进一步节省显存。

Marlin 在 SGLang 中的注册

1
2
3
@register_fused_func("none", "marlin")
def fused_experts_none_to_marlin(dispatch_output, quant_info, runner_config):
...

“none” 是 A2A backend(无 EP,纯 TP),“marlin” 是 runner backend。

这意味着 Marlin MoE 目前只支持纯 TP 模式(每张卡有所有专家的 1/N 权重),不支持 EP(Expert Parallelism)。

为什么 Marlin 没做 EP?

  1. MXFP4 权重大小不是瓶颈

    • 384 专家 × 33MB/8 = 1.6GB/卡(TP=8)
    • 如果 EP=8,每卡 48 专家 × 33MB = 1.6GB/卡
    • 权重大小一样,EP 的优势(减少单卡显存)消失
  2. 通信量对比

    • TP 的 AllReduce:2 × B × hidden(W1 + W2 各一次)
    • EP 的 All-to-All:topk × B × hidden = 6 × B × hidden(topk=6)
    • EP 通信量是 TP 的 3 倍
  3. Marlin 是量化 kernel,EP 是通信问题

    • 理论上可以接 deepep,但 MXFP4 + 高 topk 场景下收益不大
    • SGLang 团队可能评估后觉得优先级不高

完整数据流图

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
hidden_states [M, 7168] (BF16)

▼ moe_align_block_size(topk_ids)
│ → sorted_token_ids, expert_ids (按专家排序+对齐)

▼ moe_wna16_marlin_gemm (W1, gate+up fused)
│ 每个 block 读对应 expert 的 MXFP4 权重 → 反量化 → GEMM
│ [M, 7168] × [E, 7168/pack, 2*3072*pack] → [M*6, 6144]
│ mul_topk_weights=False (不乘权重)

▼ SwiGLU + clamp (swiglu_limit=10.0)
│ gate = clamp([:, :3072], max=10)
│ up = clamp([:, 3072:], -10, 10)
│ → SiLU(gate) * up → [M*6, 3072]

▼ moe_wna16_marlin_gemm (W2, down)
│ [M*6, 3072] × [E, 3072/pack, 7168*pack] → [M*6, 7168]
│ mul_topk_weights=True (内部乘路由权重)
│ → reshape [M, 6, 7168]

▼ torch.sum(dim=1) 或 moe_sum_reduce
│ [M, 6, 7168] → [M, 7168] (× routed_scaling_factor=2.5)

▼ output [M, 7168] (BF16)

总结

方面 评价 建议
Block size 选择 简单启发式,可能不够鲁棒 考虑按专家负载动态调整
MXFP4 处理 强制转 BF16 有额外开销 统一模型 dtype,避免 forward 时转换
SwiGLU clamp 好的实践,防止激活爆炸 验证量化范围是否覆盖 [-10, 10]
内存复用 做得好 intermediate_cache2 也可复用
mul_topk_weights 设计合理,W2 融合权重乘法
EP 支持 仅 TP,未实现 EP MXFP4 + 高 topk 下收益不大,暂不急需

Marlin MoE 是 MXFP4 量化 MoE 的高效实现,通过 block-sparse GEMM、SwiGLU 融合、内存复用等手段优化计算。在 DeepSeek-V4 的 MXFP4 + 高 topk 场景下,选择纯 TP 而非 EP 是理性的设计决策。

参考资料

DeepSeek-V4 series incorporate several key upgrades: (1) hybrid attention architecture that combines Compressed Sparse Attention (CSA) and Heavily Compressed Attention (HCA); (2) Manifold-Constrained Hyper-Connections (mHC); (3) Muon optimizer.

DeepSeek-V4 彻底改变了 KV Cache 的设计,从 V3 的 MLA(Multi-head Latent Attention)转向 CSA + HCA 混合注意力架构,通过序列维度压缩而非仅靠隐维度压缩来大幅降低 KV Cache 开销。


附录:Python 计算器代码

以下代码可直接运行,计算任意序列长度下的 KV Cache 大小:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
# DeepSeek-V4-Pro KV Cache Size Calculator!

# 参考文献:
# - config.json: https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/config.json
# - model.py: https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/inference/model.py
# - paper: https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/blob/main/DeepSeek_V4.pdf!

# ── 模型参数(来自 config.json) ──────────────────────────────!
NUM_LAYERS = 61
HEAD_DIM = 512 # c: KV entry 维度(RoPE 包含在 512 内)
ROPE_DIM = 64 # r: 最后 64 维携带 RoPE
NOPE_DIM = HEAD_DIM - ROPE_DIM # 448
WINDOW = 128 # w: 滑动窗口大小
INDEX_DIM = 128 # c_I: indexer KV 维度
CSA_RATIO = 4 # CSA 压缩比
HCA_RATIO = 128 # HCA 压缩比
N_CSA = 30 # CSA 层数(2,4,6,...,60)
N_HCA = 31 # HCA 层数(0,1,3,5,...,59)

# ── 精度配置 ──────────────────────────────────────!
def bf16_bytes():
"""BF16 基准:所有元素 2 bytes。"""
s_kv = HEAD_DIM * 2 # 512 × 2 = 1024
s_idx = INDEX_DIM * 2 # 128 × 2 = 256
return s_kv, s_idx

def mixed_bytes():
"""混合精度部署:nope FP8 (1B) + rope BF16 (2B) + indexer FP4 (0.5B)。"""
s_kv = NOPE_DIM * 1 + ROPE_DIM * 2 # 448×1 + 64×2 = 576
s_idx = INDEX_DIM * 0.5 # 128 × 0.5 = 64 (FP4)
return s_kv, s_idx

# ── 核心计算 ──────────────────────────────────────!
def calc_kvcache(seq_len: int, precision="bf16"):
"""
计算 DeepSeek-V4-Pro 的 KV Cache 大小。

Args:
seq_len: 序列长度(token 数)
precision: "bf16" 或 "mixed"(FP8+BF16+FP4)

Returns:
包含各组件和总大小的字典(字节)
"""
s_kv, s_idx = bf16_bytes() if precision == "bf16" else mixed_bytes()

# CSA 层 (ratio=4): Shared-KV + Indexer + SWA
csa_compressed = (seq_len // CSA_RATIO + WINDOW) * s_kv
csa_indexer = (seq_len // CSA_RATIO) * s_idx
csa_swa = WINDOW * s_kv
csa_per_layer = csa_compressed + csa_indexer + csa_swa

# HCA 层 (ratio=128): Shared-KV + SWA(无 Indexer)
hca_compressed = (seq_len // HCA_RATIO + WINDOW) * s_kv
hca_swa = WINDOW * s_kv
hca_per_layer = hca_compressed + hca_swa

total = N_CSA * csa_per_layer + N_HCA * hca_per_layer

return {
"csa_per_layer": csa_per_layer,
"hca_per_layer": hca_per_layer,
"csa_total": N_CSA * csa_per_layer,
"hca_total": N_HCA * hca_per_layer,
"total_bytes": total,
"total_gib": total / (1024**3),
# 组件分解
"csa_shared_kv": N_CSA * (seq_len // CSA_RATIO) * s_kv,
"csa_indexer": N_CSA * (seq_len // CSA_RATIO) * s_idx,
"hca_shared_kv": N_HCA * (seq_len // HCA_RATIO) * s_kv,
"all_windows": (N_CSA + N_HCA) * WINDOW * s_kv,
}

def fmt(b):
"""格式化字节数为可读字符串。"""
if b >= 1024**3:
return f"{b / (1024**3):.2f} GiB"
return f"{b / (1024**2):.2f} MiB"


def report(seq_len: int):
print("=" * 60)
print(f" DeepSeek-V4-Pro KV Cache | seq_len = {seq_len:,}")
print("=" * 60)
print(f" {N_CSA} CSA 层 (ratio={CSA_RATIO}) + {N_HCA} HCA 层 (ratio={HCA_RATIO})")
print(f" head_dim={HEAD_DIM} (nope={NOPE_DIM} + rope={ROPE_DIM})")
print(f" index_dim={INDEX_DIM}, window={WINDOW}")
print()

for prec in ["bf16", "mixed"]:
r = calc_kvcache(seq_len, prec)
s_kv, s_idx = bf16_bytes() if prec == "bf16" else mixed_bytes()
label = "BF16 基准" if prec == "bf16" else "FP8 kv + BF16 rope + FP4 indexer"
print(f" ── {label} (S_kv={s_kv}, S_idx={s_idx}) ──")
print(f" CSA 单层: {fmt(r['csa_per_layer']):>10s} HCA 单层: {fmt(r['hca_per_layer']):>10s}")
print(f" 组件分解:")
for name, val in [
("CSA Shared-KV 压缩", r["csa_shared_kv"]),
("CSA Indexer", r["csa_indexer"]),
("HCA Shared-KV 压缩", r["hca_shared_kv"]),
("滑动窗口 (61 层)", r["all_windows"]),
]:
print(f" {name:30s} {fmt(val):>10s} ({val/r['total_bytes']*100:.1f}%)")
print(f" {'─' * 56}")
print(f" {'Total':30s} {fmt(r['total_bytes']):>10s}")
print()

if __name__ == "__main__":
report(1_048_576) # 1M tokens

运行结果(Python 3.x):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
============================================================
DeepSeek-V4-Pro KV Cache | seq_len = 1,048,576
============================================================
30 CSA 层 (ratio=4) + 31 HCA 层 (ratio=128)
head_dim=512 (nope=448 + rope=64)
index_dim=128, window=128

── BF16 基准 (S_kv=1024, S_idx=256) ──
CSA 单层: 320.12 MiB HCA 单层: 8.12 MiB
组件分解:
CSA Shared-KV 压缩 7,500.00 MiB (77.9%)
CSA Indexer 1,920.00 MiB (19.5%)
HCA Shared-KV 压缩 248.00 MiB (2.5%)
滑动窗口 (61 层) 7.50 MiB (0.1%)
─────────────────────────────────────────────────────────────
Total 9,675.50 MiB (9.62 GiB)

── FP8 kv + BF16 rope + FP4 indexer (S_kv=576, S_idx=64.0) ──
CSA 单层: 160.07 MiB HCA 单层: 4.57 MiB
组件分解:
CSA Shared-KV 压缩 4,218.75 MiB (87.4%)
CSA Indexer 480.00 MiB (9.7%)
HCA Shared-KV 压缩 139.50 MiB (2.8%)
滑动窗口 (61 层) 4.25 MiB (0.1%)
─────────────────────────────────────────────────────────────
Total 4,822.50 MiB (4.83 GiB)

对比 V3.2(61 层,每 token 1152 bytes):

  • V3.2 BF16: 61 × 1,048,576 × 1152 / 1024³ ≈ 83.9 GiB
  • V4-Pro BF16: 9.62 GiB (8.7× 压缩)
  • V4-Pro 混合精度: 4.83 GiB (17.4× 压缩) ✅!

一、核心架构参数!

参数 符号 V4-Pro V4-Flash
总参数量 1600B 285B
总层数 $L_{layers}$ 61 43
CSA 层数 30
HCA 层数 31
压缩后 latent 维度 (c_KV) $c$ 512
index_head_dim $c_I$ 128
index_topk $k$ 1024
sliding_window $w$ 128

数据来源:config.json

compress_ratios 数组 → 层类型

compress_ratios 有 62 个元素(61 层 + 1 层 MTP):

1
2
3
索引: 0  1  2  3  4  5  ... 59 60 61
值: 128 128 4 128 4 128 ... 128 4 0
类型: HCA HCA CSA HCA CSA HCA ... HCA CSA MTP
类型 数量 层索引
HCA (ratio=128) 31 0,1,3,5,7,…,59
CSA (ratio=4) 30 2,4,6,8,…,60
总计 61

二、KV Cache 结构革命!

V3 MLA vs V4 CSA/HCA!

对比项 V3 MLA V4 CSA/HCA
KV 投影 Linear(dim, 576) = 512+64 Linear(dim, 512)
RoPE 位置 独立存储 64 维,拼接在 512 后 包含在 512 维内(后 64 维)
KV cache 每 entry 640 bytes (FP8+BF16) 576 bytes (FP8+BF16 混合)
压缩方式 隐维度压缩 (7168→512) 序列维度压缩 (4:1 或 128:1)

V4 KV entry 构成(512 维)

1
2
3
4
5
6
7
KV entry (512 dims)
┌────────────────────────────┬──────────────┐
│ nope 维度 (448 个元素) │ rope 维度 (64)│
│ FP8: 448 × 1 byte │ BF16: 64 × 2 │
│ = 448 bytes │ = 128 bytes │
└────────────────────────────┴──────────────┘
合计 = 576 bytes/entry (混合精度)

论文原文(Section 2.3.4):

“我们采用混合存储格式:RoPE 维度使用 BF16 精度,其余维度使用 FP8 精度。”

代码证据!

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
# model.py, Attention.__init__
self.head_dim = args.head_dim # 512
self.rope_head_dim = args.rope_head_dim # 64
self.nope_head_dim = self.head_dim - self.rope_head_dim # 448!

# KV cache 分配:512 维,不是 576
self.register_buffer("kv_cache",
torch.zeros(args.max_batch_size, kv_cache_size, self.head_dim),
persistent=False)

# 写入时:前 448 维 FP8,后 64 维 BF16
kv = self.wkv(x) # Linear(7168 → 512)
kv = self.kv_norm(kv)
apply_rotary_emb(kv[..., -64:], freqs_cis) # RoPE 对后 64 维原地施加
act_quant(kv[..., :-64], 64, scale_fmt, scale_dtype, True) # FP8 量化(前 448 维)

结论:V4 的 512 维包含 RoPE(后 64 维),所以存储是 576 bytes/entry(混合精度),相比 V3 的 640 bytes/token 还小。


三、压缩机制详解!

1. HCA(Heavily Compressed Attention)— 重度压缩!

  • 压缩比:128:1(每 128 个 token → 1 个压缩 KV entry)
  • 无 overlap(论文明确:”does not perform overlapped compression”)
  • 无稀疏选择:对所有压缩条目做全量密集注意力(提供全局”低分辨率概览”)
  • 滑动窗口:同样保留最近 128 个 token!

单层 KV Cache 计算(1M tokens):

组件 entries bytes/entry 小计
压缩 KV 1M/128 576 4.5 MB
SWA 128 1088 135 KB
总计 ~4.6 MB/层 (无 Indexer)

2. CSA(Compressed Sparse Attention)— 轻度压缩!

  • 压缩比:4:1(每 4 个 token → 1 个压缩 KV entry)
  • 窗口大小:8 个 token(含 50% overlap,步长=4)
  • 稀疏选择:Lightning Indexer(FP4 精度)选取 top-1024 最相关的压缩块
  • 滑动窗口:保留最近 128 个 token 的原始未压缩 KV!

单层 KV Cache 计算(1M tokens):

组件 entries bytes/entry 小计
压缩 KV 1M/4 576 144 MB
Indexer 1M/4 64 (FP4) 16 MB
SWA 128 1088 135 KB
总计 ~160 MB/层

四、索引器(Indexer)机制(CSA 独有)!

索引器是 CSA 的”导航系统”,用来决定哪些压缩块最值得关注

结构

  • 独立 Compressor:自己的 wkv 和 wgate 权重(head_dim=128,不是 512)
  • FP4 量化:Hadamard 旋转 + fp4_act_quant(全 128 维)
  • 计算流程
    1
    2
    Query token t → indexer query (128 维) → 和所有压缩块的索引器 key 算分
    → Top-k (1024) → 选出最相关的 1024 个压缩块

Indexer Cache 大小(1M tokens)

1
(1M/4) × 128 dims × 0.5 byte (FP4) = 16 MB/层

论文 Section 5.2.1

“索引器中的 QK 激活被缓存、加载,并完全以 FP4 精度计算。”


五、RoPE 的特殊处理!

三处应用 RoPE(Section 2.3.3)

  1. Query 向量 (q_t,i) → last 64 dims
  2. KV entry 向量 (C_Comp) → last 64 dims
  3. Core attention 输出 (o_t,i) → last 64 dims(位置 = -i)

为什么输出也要加 RoPE?!

因为 KV entry 同时作为 K 和 V,压缩后的 entry 带绝对位置信息。如果直接输出:

1
o_t,i = Σ attn_weight_j × C_Comp_j(R(j))  ← 携带绝对位置 R(j)

对输出施加 RoPE(position=-i):

1
R(-i) × o_t,i = Σ attn_weight_j × C_Comp_j(R(j-i))  ← 转为相对位置

论文原文

“注意力输出的贡献将与 query 和 KV entry 之间的距离相关。”


六、精度配置!

两种精度模式!

符号 含义 BF16 基准 混合精度部署
$S_{kv}$ Shared-KV 每元素字节 $512 \times 2 = 1024$ $448 \times 1 + 64 \times 2 = 576$
$S_{idx}$ Indexer 每元素字节 $128 \times 2 = 256$ $128 \times 0.5 = 64$ (FP4)

论文依据!

Section 2.3.4

“RoPE 维度使用 BF16 精度,其余维度使用 FP8 精度。”

Section 5.2.1

“索引器中的 QK 激活完全以 FP4 精度计算。”


七、总 KV Cache 计算公式!

通用公式!

设序列长度 $N$,CSA 层数 $N_{CSA}=30$,HCA 层数 $N_{HCA}=31$:

$$\text{Total KV} = N_{CSA} \times \text{PerCSA} + N_{HCA} \times \text{PerHCA}$$

其中:

$$\text{PerCSA} = \underbrace{(\frac{N}{4} + w) \times S_{kv}}_ {\text{Shared-KV}} + \underbrace{\frac{N}{4} \times S_ {idx}}_ {\text{Indexer}}$$

$$\text{PerHCA} = \underbrace{(\frac{N}{128} + w) \times S_{kv}}_{\text{Shared-KV only}}$$


八、数值结果(N = 1,048,576 = 1M tokens)!

BF16 基准(vLLM 博客算法)

CSA 30 层

组件 计算 结果
Shared-KV 压缩 262,144 × 1024 256.00 MiB
Shared-KV 窗口 128 × 1024 0.125 MiB
Indexer 压缩 262,144 × 256 64.00 MiB
CSA 单层 320.13 MiB
30 层 CSA 320.13 × 30 9,603.75 MiB

HCA 31 层

组件 计算 结果
Shared-KV 压缩 8,192 × 1024 8.00 MiB
Shared-KV 窗口 128 × 1024 0.125 MiB
HCA 单层 8.13 MiB
31 层 HCA 8.13 × 31 251.88 MiB

总计(BF16):9,603.75 + 251.88 = 9,855.63 MiB ≈ 9.62 GiB ✅(和 vLLM 博客完全一致)

混合精度实际部署(FP8 + BF16 + FP4)

组件 bytes/entry 30 层 CSA 31 层 HCA
CSA 压缩 KV 576 4.23 GiB
CSA Indexer 64 (FP4) 0.47 GiB
CSA SWA 576 2.11 MiB
HCA 压缩 KV 576 0.14 GiB
HCA SWA 576 2.18 MiB
总计 ~4.84 GiB

压缩比

  • V3.2 KV Cache (61 层) = 83.9 GiB
  • V4-Pro 混合精度 = 4.84 GiB
  • 压缩比:83.9 / 4.84 ≈ 17.3× ✅!

九、各组件占比(BF16 基准)!

组件 MiB 百分比
CSA Shared-KV 压缩 7,680.00 77.9% ← 绝对主体
CSA Indexer 1,920.00 19.5% ← 不可忽略
HCA Shared-KV 压缩 248.00 2.5% ← 128× 压缩极高效
滑动窗口 (全部 61 层) 7.63 0.1% ← 可忽略
总计 9,855.63 100%

CSA 的 Indexer 占了近 20%,是不能漏算的部分。


十、为什么能这么小?关键洞察!

  1. 序列维度压缩(V3 只压缩隐维度,V4 同时压缩序列长度)

    • CSA: L → L/4
    • HCA: L → L/128
  2. 稀疏注意力(CSA 只访问 top-1024 个压缩块,而非全部 L/4 个)

  3. 分层设计:CSA 负责精准的中程依赖,HCA 负责全局概览,互为补充

  4. 混合精度

    • RoPE 维度:BF16(2 bytes,保证位置精度)
    • 其余维度:FP8(1 byte,节省空间)
    • Indexer:FP4(0.5 byte,最激进)
  5. 磁盘 KV Cache:压缩后的 KV 块可以落盘存储,共享前缀的请求复用缓存*


十一、Forward 计算流程图(V4-Pro)!

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
                Hidden State [batch, seq_len, 7168]

┌───────────────┴───────────────┐
▼ ▼
Q Down-proj KV 通路
[7168 → 1536] [7168 → 512] (wkv)
│ │
▼ ▼
Q Up-proj (1536→128×512) ┌──→ 压缩 (CSA: ratio=4 / HCA: ratio=128)
│ │ ├─→ RoPE (后 64 维, BF16)
Q [batch, seq_len, 128, 512] │ ├─→ FP8 量化 (前 448 维)
│ │ └─→ 写入 KV cache
┌─→ RoPE (后 64 维) │
│ │ │
│ Q [batch, seq_len, 128, 512] KV cache [batch, N/ratio, 512]
│ │ │
│ └──→ Indexer (仅 CSA) ──→ Top-k → 1024 个选中
│ │
└──→ Attention(Q, K=selected_1024, V=selected_1024)

└──→ + SWA (最近 128 个 token)


Output [batch, seq_len, 7168]

十二、和 vLLM 部署对齐!

vLLM 博客明确验证了 V4 的计算:

“using fp4 indexer cache and fp8 attention cache, which further reduces the KV cache size by roughly compared to the bf16 estimate!”

来源 KV Cache (1M tokens) 说明
论文 Figure 1 9.62 GiB (10% of V3.2) BF16 估计
vLLM 博客 ~4.8 GiB FP4+FP8 混合精度
本文计算 4.84 GiB 精确对齐 ✅

十三、总结!

DeepSeek-V4 通过序列维度压缩 + 混合精度存储 + 稀疏选择,将 1M tokens 的 KV Cache 从 V3.2 的 83.9 GiB 压缩到 4.84 GiB(17.3× 压缩)。

核心创新:

  • CSA:4:1 压缩 + overlap + Indexer 稀疏选择
  • HCA:128:1 重度压缩,提供全局概览
  • 混合精度:RoPE (BF16) + 内容 (FP8) + Indexer (FP4)

这套架构让百万 token 上下文从”不可能”变成”日常可用” 🐾!


参考!

  1. DeepSeek-V4 论文 (2026) Section 2.3
  2. vLLM 博客:DeepSeek V4 in vLLM (2026-04-24)
  3. config.json
  4. 代码追踪: model.py, compressor.py, c4.cuh
  5. 论文 Section 2.3.3: Partial Rotary Positional Embedding
  6. 论文 Section 5.2.1: FP4 Quantization-Aware Training

概述

SGLang 的 Pipeline Parallelism (PP) 模式下,HiCache 负责 KV cache 的异步 prefetch(Host→GPU)和 backup(GPU→Host)。本文从 PP 调度器的外层事件循环出发,追踪 Load 和 Write 的完整时序,揭示 write_ack 比 load_ack 多延迟的根本原因。

一、PP 事件循环

每个 iteration 执行以下 4 步:

1
2
3
4
5
iter=N:
① check_hicache_events() ← 查询 HiCache 异步事件(load_ack / write_ack)
② get_next_batch_to_run() ← 选下一批 batch(prefix match → eviction → load)
③ _pp_launch_batch() ← launch forward
④ _pp_process_batch_result() ← 处理上一批 batch 的结果(insert → write)

核心设计process_batch_result 处理的是 mbs[next_mb_id]——上一轮 launch 的 batch。同一个 iter 内,调度器同时在处理两个不同 batch 的生命周期阶段。

二、Load 时序(Host→GPU)

触发时机

iter=N 的 get_next_batch_to_run —— prefill 阶段 prefix match 发现 host_hit,发起 load_back。

完整时序

步骤 发生在 说明
1. load 发起 iter=N get_next_batch_to_run match_prefix → host_hit → load_back → start_loading
2. CUDA copy 启动 iter=N get_next_batch_to_run GPU 从 Host 异步拉取 KV cache,CUDA event 入队
3. forward 逐层等待 iter=N _pp_launch_batch forward 通过 consumer_index 逐层等待对应 layer 的 load 完成
4. load_ack 消费 iter≥N+1 check_hicache_events loading_check → event.query()=True → 消费 ack

关键:load 是 prefetch——在 forward 之前触发,CUDA copy 和 forward 可以重叠(逐层等待、逐层执行)。

时序图

1
2
3
4
5
6
7
8
9
10
11
12
iter=N:
① check_hicache_events()
└─ 消费更早的 load_ack
② get_next_batch_to_run()
└─ match_prefix → host_hit → load_back() → start_loading()
└─ CUDA copy 启动,event 入队
③ _pp_launch_batch()
└─ forward 逐层等待 load 完成

iter=N+1:
① check_hicache_events()
└─ loading_check() → event.query()=True → load_ack ✓

延迟:load_ack 比 load 发起晚 1 iter

三、Write 时序(GPU→Host)

触发时机

iter=N 的 process_batch_result —— 处理上一轮 launch 的 batch 结果,insert 时触发 write_backup。

完整时序

步骤 发生在 说明
1. forward 执行 iter=N-1 _pp_launch_batch Prefill batch 在 GPU 上计算,生成 KV cache
2. write 发起 iter=N process_batch_result 处理 iter=N-1 的 batch → insert → write_backup → start_writing
3. CUDA copy 启动 iter=N process_batch_result GPU 异步写回 Host,CUDA event 入队
4. write_ack 消费 iter≥N+1 check_hicache_events writing_check → event.query()=True → 消费 ack

关键:write 是 post-write——必须在 forward 算完、拿到完整 KV cache 后才能触发。

时序图

1
2
3
4
5
6
7
8
9
10
11
12
13
iter=N-1:
② get_next_batch_to_run() → 选出 Prefill A
③ _pp_launch_batch() → forward(Prefill A)

iter=N:
④ _pp_process_batch_result()
└─ 处理 iter=N-1 的 Prefill A
└─ insert() → write_backup() → start_writing()
└─ CUDA copy 启动,event 入队

iter=N+1:
① check_hicache_events()
└─ writing_check() → event.query()=True → write_ack ✓

延迟:write_ack 比 write 发起晚 1 iter,但 write 本身比 forward 晚 1 iter(因为 process_batch_result 处理上一批 batch)。

四、延迟对比

以同一个 Prefill batch 为基准

追踪一个 Prefill batch 从 launch 到 write_ack 的完整生命周期:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
iter=1:
② get_next_batch_to_run()
└─ 选出 Prefill A
└─ match_prefix → host_hit → load_back() → start_loading()
③ _pp_launch_batch() → forward(Prefill A)

iter=2:
① check_hicache_events()
└─ loading_check() → load_ack ✓
↑ load_ack: 1 iter delay

iter=3:
④ process_batch_result()
└─ 处理 iter=1 的 Prefill A
└─ insert() → write_backup() → start_writing()
① check_hicache_events()
└─ (write_ack 还没完成)

iter=4:
① check_hicache_events()
└─ writing_check() → write_ack ✓
↑ write_ack: 3 iter delay from Prefill launch

对比表

事件 触发时机 ack 入队时机 ack 消费时机 相对于 Prefill launch 的延迟
load iter=1 get_next_batch_to_run iter=1 iter=2 1 iter
write iter=3 process_batch_result(处理 iter=1 的 batch) iter=3 iter=4 3 iter

为什么 write_ack 多延迟?

两个因素叠加:

1. process_batch_result 滞后一轮

它处理 mbs[next_mb_id](上一轮 launch 的 batch),以 PP2 为例偏移 2 iter。所以 write 比 load 晚 2 iter 才触发。

2. CUDA event query 需要等下一轮 check

无论 load 还是 write,ack 入队后最早等下一轮 check_hicache_events 才能消费。两者各加 1 iter。

综合

1
2
Load:  iter=1 发起 → iter=1 ack 入队 → iter=2 消费
Write: iter=1 forward → iter=3 process_result 发起 → iter=3 ack 入队 → iter=4 消费

Load 是 prefetch(forward 之前),Write 是 post-write(forward 之后)。这个架构差异决定了 Write 必然比 Load 多延迟。

五、PP 间同步问题

PP0 和 PP1 共享同一个 HiCache 实例(radix tree + host memory)。由于 output relay 延迟,PP0 和 PP1 的 iter 进度存在 1-2 iter 的偏移,需要三层同步机制保证一致性。

5.1 逻辑时钟保证重放一致性

PP1 的 writing_check/loading_check 不再直接消费 ack,而是将事件通过 Gloo P2P 通道 replay 给 PP0。PP0 作为唯一的事件消费者,按逻辑时钟顺序处理所有 ack,确保 PP0 和 PP1 的 radix tree 操作顺序一致。

问题:早期实现中 PP1 的 writing_check() 绕过 check_hicache_events guard,直接消费 write_ack(即 ack theft),导致 PP0 端 pending event 永远无法完成,radix tree 分叉。

修复:将 PP rank 分支移入 writing_check/loading_check 内部,用 PPHiCacheEventsReq 控制请求替代 dict wrapper,强制 PP1 replay 事件给 PP0。

5.2 Count Sync 保证 CP 一致性

PP0 比 PP1 早 1-2 iter 积累 write ack(output relay 延迟),导致 ack_write_queue 积累差异。PR #22878 通过 piggybacking write-ack consumption counts 在 PP ranks 间同步:

1
2
3
iter=N:
PP0: writing_check() → 消费 3 个 write_ack → count=3
PP1: 等待 PP0 的 count → count=3 → radix tree 操作对齐

Count sync 确保 PP0 和 PP1 对同一个 radix tree node 的 checkpoint(CP)操作一致,避免一个 stage 认为 node 已 backup、另一个 stage 还在等待的情况。

5.3 PP1 同步消费 PP0 的 ack

PP1 不再独立消费 ack,而是通过同步机制确保 PP0 消费 ack 后,PP1 的 radix tree 状态与 PP0 对齐:

1
2
3
4
5
PP0: writing_check() → event.query()=True → 消费 write_ack
→ radix tree: insert() → finalize() → node 状态更新
↓ TP all_reduce sync
PP1: 等待 PP0 完成 → radix tree: 同步执行 finalize()
→ node 状态与 PP0 一致

三层保障

同步层 机制 保证什么
逻辑时钟 Event Replay(Gloo P2P) PP0/PP1 事件处理顺序一致
Count Sync PR #22878 piggyback counts PP0/PP1 checkpoint 状态一致
Ack 同步 PP1 replay → PP0 消费 → TP all_reduce radix tree 节点状态一致

核心目标:无论 PP0 和 PP1 的 iter 偏移多少,radix tree 的结构和节点状态在两个 stage 上始终保持一致。

六、总结

维度 Load Write
触发函数 get_next_batch_to_run process_batch_result
触发契机 prefix match 发现 host_hit insert → _inc_hit_count
与 forward 关系 forward 之前(prefetch) forward 之后(post-write)
CUDA copy 与 forward 可重叠(逐层等待) 串行(forward 算完才触发)
ack 消费延迟 +1 iter +3 iter(含 process_result 偏移)

核心结论

  1. PP 调度器每个 iter 同时做两件事:get_next_batch_to_run 选下一批 batch,process_batch_result 处理上一批 batch 的结果
  2. Load 是 prefetch(forward 之前触发),Write 是 post-write(forward 之后触发)
  3. process_batch_result 处理 mbs[next_mb_id] 的 iter 偏移是 Write 延迟的根本原因
  4. PP 间需要 count sync 补偿 output relay 带来的时序差异

本文基于 SGLang 源码分析,涉及文件:hiradix_cache.pycache_controller.py