如何不造出另一个 Megatron(二):极简并行方案
我在上一篇博客说了一堆流水线并行的不好,这一篇来填坑:不用 PP,参数怎么放,长上下文怎么切,跨机通信又怎么办?
在讨论这个话题之前,我们先做一些假设和限制,以此来收缩我们的设计空间。要不然的话,什么功能都往上加,就又变成 Megatron 了,还不如直接把 Megatron 拿来用。
约束条件
- 支持万亿级别参数量(1T+ params)
- 支持一百万上下文长度
- NVL72 集群
- 不依赖 CPU offload 存放模型状态和激活值。
方案推导
我们先粗略地算一算需要多少显卡才有足够多的显存。首先考虑每个参数所需要的常驻显存:
- FP32 主权重:4 B
- FP32 AdamW 优化器状态:4+4 B
- FP32 梯度:4 B
- BF16 权重:2 B
- 注:如果用 Muon 可以减小优化器状态,但是这里我们为了保守起见,优化器状态取 8 B。
- 注:BF16 权重可以从主权重中重建,因此理论上可以不必常驻在显存里面,但是大部分实现会把它留在显存里面。这里为了保守起见,我们将其计入常驻显存。
- 注:为了保守起见,这里我们也不考虑低精度量化。
假设是万亿参数量,这就是 18 TB 的显存需求。在上面这套精度和状态格式下,分片改变的是状态放在哪里,不改变去重后的总量。一张 GB200 有 200 GB 显存,光是常驻显存就至少得有 90 张卡。也就是一个 NVLink 通信域无法放下,必须有一部分通信走 RDMA。
那我们先来解决一下参数的切分。
因为我们考虑的是万亿级别参数量的模型,而现在这个规模的模型全都是 MoE 架构,并且专家的参数量占总参数量的 98% 以上。所以只要解决了专家的切分,我们基本上就解决了常驻显存的问题。
得益于 NVL72 互联,我们可以直接使用 EP=64 做通信和计算,EP 的 dispatch/combine 在 NVLink 内部,不需要走 RDMA。这样的纯 NVLink 算子更容易实现。之前 DeepSeek-V3 里面为了跨机通信绞尽脑汁想出来的 DeepEP 和 DualPipe 都不需要了。(再次感慨一下 DeepEP 真的是太精巧了,更多故事见前一篇博客文章。)
但是显然如果只有 EP=64,依然放不下所有参数。所以我们用 FSDP 对专家参数再切一刀,比如增加一维 eFSDP=4。这样一来,一共就是 256 张卡。一个专家的参数被切到分布在 4 个机柜里的 4 张卡上。在计算这一层的时候,FSDP 可以预取下一层的参数。经过 FSDP all-gather 之后,每个机柜都有了当前层完整的参数。这样就可以保证 EP=64 的通信依然发生在 NVLink 内部。
好了,搞定了占 98% 的专家参数之后,剩下占 2% 的稠密参数就可以随意解决了,用 FSDP 切一下就行。
解决了参数的放置问题,接下来算算动态显存占用。以 GLM-5.2 架构为例:
- 假设我们在反向传播的时候重算每一层,我们至少也需要把每一层的激活值保存下来。
78 层 * 6144 隐状态 * BF16 = 0.958 MB/token - 另外我们把 DSA Indexer 的 top-k 结果也存下来,因为这一部分跑得比较慢。
21 层 * top 2048 * int32 = 0.172 MB/token - EP 通信缓冲区。每个词元需要
6144 BF16 隐状态 + FP32 概率 * 256 专家 = 13312 B。考虑最坏情况,所有 EP rank 都往当前 rank 发,也就是放大EP倍。13312 B * 64 = 0.852 MB/token - 每层重算加上反向传播大概需要
2.5 MB/token。 - FSDP:
753B 参数 / 78 层 * (2% 稠密参数 + 98% 专家参数 / EP 64) * (BF16 当前层权重 + BF16 下一层预取权重 + FP32 梯度缓冲区) = 2.727 GB。
再次得益于 NVL72 互联,我们可以直接开启 CP=64,把一条序列切到一整个机柜里面。每张卡的动态显存占用是 1048576 * (0.958 + 0.172 + 0.852 + 2.5) / 64 = 73.433 GB。加上参数所占的显存 753 * 18 * (2% / 64 + 98% / 256) = 56.122 GB 和 FSDP 2.727 GB,总共约 133 GB,离 200 GB 还有 67 GB 的浮动空间,可以用来放置其他未建模的临时缓冲区。
如果 FA4 PR 2816 合并了,正反向传播还能各省约 0.5 MB/token。这样一来,缩减到 128 卡也能放下。或者同样是 256 卡,可以把 microbatch 从1增加到2。
到这里,整个方案的骨架其实已经出来了,就是 FSDP x CP x EP 这三种并行方式。专家参数和稠密参数使用不同的切分方式:
- 稠密网格:
(d_rep, fsdp, cp) - 专家网格:
(e_fsdp, ep) ep <= NVL且fsdp * cp <= NVL,保证所有每层都要做一次的关键路径通信留在 NVLink 内部。d_rep * fsdp * cp == e_fsdp * ep == world- 稠密参数在
fsdp x cp内做切分,跨d_rep复制。 - 专家参数分散在
ep内,每个 rank 再沿e_fsdp切分。 cp大小跟ep无关,取决于需要支持的上下文长度和显存占用量。- 总数据并行(Data Parallelism, DP)数量是
d_rep * fsdp。 - 如果要扩展到万卡级别,可以加上
e_rep一维。
上面 256 卡 1M 上下文的例子取
(d_rep, fsdp, cp) = (4, 1, 64)(e_fsdp, ep) = (4, 64)
eFSDP 通信掩盖
上面的推导过程解决了显存方面的顾虑,接下来我们来分析通信。最大的顾虑就是 FSDP 通信能不能被每一层的计算掩盖。还是沿用前面的 GLM-5.2 架构为例做计算。
- 首先估算每个 rank 需要完整持有的一层专家参数大小:
753B 参数 / 78 层 * 98% 专家参数 * BF16 / EP 64 ≈ 300 MB - 这些参数被切成了
e_fsdp份,所以每一个 rank 需要收到(e_fsdp - 1) / e_fsdp的比例。当e_fsdp比较大的时候,这个比例就趋近于 1。我们这里取上限作为最差情况计算。所以每一层前向传播的时候,FSDP all-gather 需要收到 300 MB 数据。 - 前向结束后权重分片被释放,所以在反向传播的时候,又需要一次 FSDP all-gather 300 MB。然后 FP32 梯度需要 FSDP reduce-scatter 600 MB。
- 因此一层总的通信量大约是 1200 MB。
- 每张 GB200 一般配套 400 Gbps 的 RDMA 带宽;GB300 则是 800 Gbps。这里我们取 45 GB/s 做估算。
- 因此每一层的 eFSDP 通信时间约为 27 ms。
注意到这个通信量是跟输入词元数量无关的。所以我们可以把问题“eFSDP 通信能不能被掩盖”换成“一个 microbatch 至少需要多少词元才能有足够多的计算来掩盖 eFSDP 通信”。
那我们接下来再估算一下计算。
- 首先算一下每一层激活的参数量:
753B 参数 / 78 层 * (2% 稠密参数 + 98% 专家参数 * 8/256) ≈ 0.5B。 - 先考虑线性层计算。前向传播需要
2 FLOP/param/token;反向传播中的重计算、dgrad 和 wgrad 各需要2 FLOP/param/token。因此一共是4 GFLOP/token的计算量。 - 因为注意力的计算量比较麻烦,这里不仔细建模,直接简单假定注意力和线性层的时间比例。根据 profile 数据,在 256k token(262,144)的上下文下为 2:1,在 1M 上下文下为 4:1。这里先保守按 2:1 算(低估 1M 上下文的注意力的计算时间)。
- GB200 纸面数据是 2.5 PFLOP/s,假设能达到 80% 的理论性能(高估)。
- 保守估计(高估通信时间、低估计算时间),掩盖 eFSDP 通信需要的词元数量是
(2.5 PFLOP/s * 80%) * 27 ms / (4 GFLOP * (1 + 2)) ≈ 4,500
也就是每张卡一个微批处理只要有 5000 个词元就能盖住 eFSDP 通信。可以说轻轻松松。
其实这段分析也解释了 eFSDP 和 PP 一个很大的区别。PP 在通信和计算上的浪费是结构性的,只能通过调整 stage 数量和 microbatch 数量来减少气泡,或者像 Zero Bubble 那样用更复杂的调度和更多显存去从理论上消除气泡。而 eFSDP 则是计算和通信两者在时间上的比较,只要输入长度超过阈值,通信的额外代价就几乎是零。
当然 eFSDP 也有结构性的浪费。首层前向传播的 all-gather 和末尾反向传播的 reduce-scatter 都是无法掩盖的。不过相对来说这 1/78 的代价就比较好接受,毕竟要让 PP 的气泡降到 1.28% 可不是一件容易的事情。
CP
河神:年轻的樵夫呦,你掉的是这把金斧头,还是这把银斧头,还是这把铁斧头呢?
上下文并行(Context Parallelism, CP)就是在序列维度切一刀,一条完整的序列被拆到多张卡上计算。大部分的计算是针对单个词元的隐状态的,并没有序列的概念,CP 并不影响。只有注意力是跨词元的,需要处理序列。CP 的方案也有好几种——
- Ring Attention:Q 留在本地,KV 分块环传,边计算边合并 LSE。通信方式是点对点,所以有时候也被叫做 P2P CP。代价是需要多轮通信、需要合并状态,算子的效率下降,并且有可能需要修改注意力算子的实现。
- Ulysses:用 all-to-all 把原先在序列维度的切分转置成按注意力头切分,然后使用正常的注意力算子进行计算,算完了再进行一次 all-to-all 转置回去,因此有时候也被叫做 A2A CP。好处是不用改算子,不用合并 LSE。坏处是受到了注意力头数量的限制,Q/K/V/O 都需要通信。
- USP:结合上面两者,Ring 走 RDMA,Ulysses 走 NVLink。有点太复杂了,在 NVL72 下也用不上了。
- All-gather KV:顾名思义,先 all-gather KV,然后本地的 Q 就可以跟全局的 KV 正常计算了。好处就是实现起来特别简单,没有注意力头的数量限制,没有算子实现的限制。坏处就是要把全局的 KV 传一遍,而且需要有显存把这一层的全局 KV 临时存下。
现代架构(GQA、MLA、DSA)都在大幅压缩 KV 的大小,所以 All-gather KV 方案显得又好写又快。MLA 传的是 latent KV。DSA 除了 latent KV,还需要传 Indexer 的 K。
线性注意力(GDN、KDA)因为状态是递推的,所以跟普通的 softmax attention 不一样,没法用 all-gather KV 方案。除了重算和 Ulysses 两个方案以外,还有一个精妙的 KCP。KCP 我还没搞懂,但是跑得很快就是了。
另外值得一提的是,因果掩码(causal mask)让不同 CP rank 的计算量不均匀,越往后的 rank 需要的计算越多,对长上下文尤其明显。可以加入 zig-zag 排列来解决负载不均衡的问题。
TP
我们这个极简并行方案里面只保留了 FSDP x CP x EP,不仅砍掉了 PP,也没加上 TP。其实 TP 倒不会带来太多实现和维护上的负担,这里没加上主要还是因为想要保持极简,以及 TP 并不能保证额外带来很多收益。
首先我们讨论专家的张量并行,也称为 ETP。ETP 和 EP 可以一起使用。就是把专家按照行或列切成 ETP 份,然后正常做 EP,做计算的时候维度缩小了 ETP 倍,算完之后再进行 TP 通信和聚合。
ETP 和 EP 对 NVLink 通信域的切分是正交的,所以如果设置了 ETP=4,另一边就只能降到 EP=16。这样做有什么好处呢?
记 \(A_i\) 为 EP=64 时每一张卡上面的负载,记 \(B_g\) 为 ETP=4 EP=16 时每组的负载。那么有
本质上,就是用 4 张卡负载的平均值,替代各自的负载。定义负载偏斜为
\[\operatorname{skew}(X) \coloneqq \frac{\max_{j} X_j}{\overline{X}}\]总计算量和平均负载不变,而平均值不会超过组内最大值,所以加入了 ETP 之后负载偏斜会降低:
那么代价是什么呢?
- GEMM 会变慢。注意到专家的中间维度本来就小(2048 甚至 1536),ETP 切了之后就更小了。
- EP 通信增多。一个词元要到达一个 ETP 组的全部 4 张卡,EP 通信量要变成不开 ETP 的4倍。
- 激活显存占用增多,跟上面同样的道理。
- 每层多了 ETP 组内的 reduce-scatter 和 all-reduce 通信,又多了两个同步点。
所以负载偏斜变好,并不等于整体跑得更快。
而如果只是想解决负载不均衡的问题,还有别的思路。比如最近 MoonEP 通过动态冗余专家和在线规划实现完美均衡。
另一方面,考虑稠密投影和注意力机制那一边的 TP。
对于稠密投影来说,我们已经使用 CP 在词元维度进行了切分,因此已经让稠密投影的计算并行化了。在这个基础上,理论上来说依然可以加入 TP 在矩阵乘法内部进行进一步的切分。但注意到 CP 和 TP 对 NVLink 通信域的切分是正交的,增加了 TP 就得降低 CP,所以两者对稠密投影的并行加速的贡献是冲突的。
注意力就得一个个讨论了。
- GQA:Qwen3.5 只有 2 个 KV 头,如果 TP 大于 2,就需要复制 KV。反过来,对 CP 来说,KV 头少反而使得 all-gather KV 非常轻松。
- MLA:展开成了 MHA 之后头很多,倒是没什么问题。但是在强化学习后训练里面,为了减少训练框架和推理引擎之间的数值差异,可能会考虑以矩阵吸收后的状态进行计算。对这种情况,好像目前的算子都不能支持 TP。
- DSA:目前几个算子实现基本上会在内部把头的数量填充到 64 的倍数,所以这就限制了 TP 能带来的加速。
- GDN:看起来好像算子是支持的,但是 GDN 的头数也不是很多,有可能加速有限。
总体上来说,TP 和 CP 在争夺并行度的上限,而且 TP 往往伴随着算子效率的下降,不实测的话难以说清到底会变快多少。CP 的实现比起 TP 又容易得多,也能非常方便地扩展到 64 张卡上,所以为了简单起见,我这里就只选择了 CP,删除了 TP。
番外:如果没有 NVL72 呢?
前面我们假设的是 NVL72 集群,但如果只能租到 NVL8 集群,我们这套方案还能不能适用呢?这里我们考虑 B300 而不是 B200,一是因为 B300 的显存要大得多,有 309 GB,二是因为 B300 配套了 800 Gbps 的 RDMA 通信。
先放松一下上下文长度,考虑 256k 上下文(每个序列 262,144 个词元)。把前面的计算搬下来:
(e_fsdp, ep) = (32, 8),cp=8,d_rep * fsdp == 32- 每个 CP 组的 microbatch 放一条 256k 序列,每张卡上词元的数量至少是
262144 * 32 / 256 = 32768。 - 动态显存
- EP 通信缓冲区变小了:
13312 B * 8 = 0.106 MB/token - 其他的动态显存占用不变:激活值检查点
0.958 MB/token,DSA Indexer top-k0.172 MB/token,每层重算加上反向传播1.5 MB/token(假设 FA4 PR 2816 合并了) - FSDP:
753B 参数 / 78 层 * (2% 稠密参数 + 98% 专家参数 / EP 8) * (BF16 当前层权重 + BF16 下一层预取权重 + FP32 梯度缓冲区) = 11.005 GB - 总的动态显存占用是
32768 * (0.106 + 0.958 + 0.172 + 1.5) MB + 11.005 GB = 100.66 GB
- EP 通信缓冲区变小了:
- 参数显存
- 如果
(d_rep, fsdp, cp) = (1, 32, 8),那么所有参数跨 RDMA 切分成 256 份:753B params * 18 B/param / 256 = 52.95 GB - 如果
(d_rep, fsdp, cp) = (32, 1, 8),那么稠密参数只在 NVLink 域内切 8 份:753B param * 18 B/param * (2% / 8 + 98% / 256) = 85.77 GB
- 如果
- 从显存看来,256 张 B300 NVL8 支持 256k 上下文有明显余量。
- 每层通信
- 每个 rank 需要完整持有的一层专家参数:
753B 参数 / 78 层 * 98% * BF16 / EP 8 = 2.365 GB - 每层 eFSDP 的通信量:
2.365 GB * (1 + 1 + 2) = 9.460 GB - 假设稠密参数的 FSDP 也走 RDMA,那么还得加上:
753B 参数 / 78 层 * 2% * (2 + 2 + 4) B = 1.544 GB - 一层总的通信量是 11 GB。
- 按照 90 GB/s 算,这就是 122 ms。
- 每个 rank 需要完整持有的一层专家参数:
- 每层计算
- 每层激活参数量仍然是
0.5B,计算量4 GFLOP/token - 假设注意力时间是线性层的 2 倍,算力按 80% 算
- 要掩盖通信,每卡最少需要的词元数量:
(2.5 PFLOP/s * 80%) * 122 ms / (4 GFLOP * (1 + 2)) ≈ 20,300
- 每层激活参数量仍然是
注意到前面我们列出来了,每张卡上的词元数是 32,768,所以完全可以被盖住。我没想到其实在 256k 上下文的情况下,B300 NVL8 集群完全可以高效地训练。
然而考虑一百万上下文长度的话,这个方案就不行了,光是动态显存就不够。主要原因还是在 NVLink 域内 CP 只能开到 8。
那我们试着考虑一下二维 CP。假设 CP=8*8=64。
- 因为 CP 开到了 64,每张卡的词元数量降到了
1,048,576 / 64 = 16384 - 每层的 latent KV:
(512 kv_lora + 64 rope + 128 indexer-K) * BF16 = 1408 B/token - 一百万上下文就是
1.476 GB - 前向一次 BF16 all-gather,反向一次 BF16 all-gather 和一次 FP32 reduce-scatter。
- all-gather 分两步
- RDMA:每张卡从其他 7 个节点的 7 张卡上各接收 1/64 的 KV:
1.476 GB * 7/64 = 161.5 MB - NVLink:每张卡再从同一节点的 7 张卡上各接收 1/8 的 KV:
1.476 GB * 7/8 = 1.292 GB
- RDMA:每张卡从其他 7 个节点的 7 张卡上各接收 1/64 的 KV:
- reduce-scatter 类似地也是分两步,先做 NVLink 再做 RDMA。
- 比起 NVL72 的 CP64,NVL8 CP64 额外的 RDMA 通信时间是
161.5 MB * (1 + 1 + 2) / 90 GB/s = 7.2 ms - 依然假设一百万上下文的注意力时间是线性层的 4 倍。一层的计算时间是
4 GFLOP/token * 16384 token * (1 + 4) / (2.5 PFLOP/s * 80%) = 164 ms - 注意这
7.2 ms在关键路径上,不能被掩盖,相当于额外付出了 4.4% 的时间。 (2.5 PFLOP/s * 80%) * 122 ms / (4 GFLOP * (1 + 4)) = 12200 < 16384
也就是如果这个二维 CP 实现得好的话,一百万上下文可以在比 NVL72 多付出 4.4% 的时间的情况下搞定。上面还高估了算力,因此高估了通信的额外开销,实际上应该还能更低。
总结
绕了一大圈,最后剩下的其实就三样东西:FSDP、CP、EP。
这倒不是说 PP 和 TP 有什么原罪,也不是说 Megatron 不好。恰恰相反,Megatron 之所以会长成今天这个样子,就是因为它得支持各种模型、各种机器、各种并行组合。每一个功能单独拿出来,大概都能找到一个非加不可的理由。只不过所有这些理由叠在一起,最后就变成了 Megatron。我对能维护这样复杂系统的团队只有深深的敬意。
但如果我们先把问题限制住,情况就不太一样了。
NVL72 给了一个巨大的高速通信域,所以 EP 和 CP 可以留在 NVLink 里面;MoE 把绝大多数参数塞进了专家,所以跨通信域最麻烦的就只剩专家权重的切分;而长上下文本身又提供了足够多的计算,让跨机 eFSDP 的通信可以藏在计算后面。于是 PP 没有了必须存在的理由,TP 也从必选项变成了实测有收益再说。
甚至换到只有 NVL8 的机器上,同样的思路也还能继续往下推:NVLink 不够,那就把 CP 拆成两维,多付一点无法掩盖的 RDMA 通信。RDMA 带宽增加到 800 Gbps 又进一步让这个额外的通信开销变得容易承受。
人有不为也,而后可以有为。
——《孟子·离娄下》
做系统大概也是这样。加一个功能,总能找到理由,而真正困难的是知道哪些东西可以不做,哪些复杂度可以拒绝。
如何不造出下一个 Megatron?也许答案不是学会更多并行方式,而是在学完、算清楚之后,有底气说一句:
这个不做。