Salesforce封顶MoE四类显存尖峰,百亿模型训到百万上下文

Flattening Every Memory Peak in Long-Context Mixture-of-Experts Training

Shrey Pandit, Xuan-Phi Nguyen, Yiran Zhao, Shafiq Joty

cs.DC

2026-09-13

MoE长上下文训练会被dispatch、词表投影、checkpoint、优化器四类尖峰轮流打死。Salesforce用四个精确调度在启动时封顶工作集,120B–667B训到1M上下文,达FSDP2的8–32倍。

这篇在解决什么

训 MoE 时,平均显存好看没有用,任何一个组件的峰值超过 HBM,这一步就死。常用并行会封住权重分片,仍留下四块随工作负载涨、而且涨法不一样的活集合:专家 dispatch 跟路由矩阵走,词表投影跟 token×词表走,checkpoint 边界跟深度×长度走,AdamW 状态跟参数量走。砍掉当前最大的一块,下一块就会露出来。Salesforce AI Research 的目标是四个一起封顶,并且每个都能单独开关。

方法

四个算子只改计算和搬运的次序与粒度,模型、精度、优化器和损失不动,所以仍是精确的全参数 BF16。

四者活集合不相交,嵌进 MoP 的 rank 布局(密集权重 ZeRO-3、注意力序列并行、专家并行),得到一项启动时可检查的 per-rank 预算。

结果

单点对照在 8×H200、输入固定时测。

算子指标结果对照
PipelinedLLEPdispatch 峰值-56.9% 到 -59.3%LLEP,65K token/rank、top-8
PipelinedLLEP速度1.01–1.10×同 LLEP
Ring-DTP词表投影峰值-86.6%时间增加不到 5%
SCO最大 batch+17.7%吞吐变化不到 2%
OffloadStreamAdamW优化器步1.93 s(2.05×)CPU Adam 3.95 s

步长分块相对连续分块快 1.03–1.35×,最多再省 1.37 GiB。组合后在 120B / 241B / 667B、16/32/64 张 H200 上,对同一 GPU 数下调过的 FSDP2 最优配置:三条规模都能训 1M 上下文,可达长度 8–32×。FSDP2 在 120B 过 128K、241B 过 32K、667B 过 64K 就 OOM。在 FSDP2 还能跑的最长点上,120B@128K 吞吐 7.6×,667B@64K 吞吐 10.4×。最大全局 batch 为 150 万 / 180 万 / 300 万 distinct token,分别是基线的 12× / 7× / 3×。每 GPU 算力从 128K 的 91–110 TFLOP/s 升到 1M 的 213–233,因为注意力随长度变重;FSDP2 全程低于 40。附录称训练质量不变。

为什么重要

这是给已经在用 FSDP、专家并行、checkpoint、优化器卸载的团队看的调度论文,不是新的 MoE 结构。价值在于四个峰值的相对高度会随模型、上下文、卡数换位,必须同时有界,而且界限在启动时就能算出来,避免「打掉一个 OOM、下一个再炸」。损失保持精确,数字可以直接和标准全参 BF16 比。1M 上下文、百亿到千亿 MoE,如果卡数加不起,这套组合把显存换成可预测的流式时间。

局限与存疑

作者列了边界。all-to-all 和环通信默认快互联,文中数字来自节点内 NVLink,更慢的织物上多块的代价会变。c 越小接收界越紧、块数越多,前向时间只在一段区间对块数不敏感,所以 c 是按实测曲线选的,不是越小越好。checkpoint 和优化器流式消耗主机内存和链路,实验节点有 2 TB 主机内存,更小的机器可能先卡在 CPU。如何按拓扑自动选 (D, Ep, P, c, β) 仍开放。端到端比的是 FSDP2 扫描后的最好点,不是每一家内部的 MoE 栈;「训练质量不变」放在附录,正文没有给损失曲线。

术语

原文与代码

相关论文

全部论文解读