MoE 模型的长上下文训练,很少死于「平均显存不够」,更多死于某一个组件的峰值先爆。Salesforce AI Research 的四位研究者(Shrey Pandit、Xuan-Phi Nguyen、Yiran Zhao、Shafiq Joty)9 月 13 日放出的预印本,把这个问题拆到了组件级:并行训练方案里通常有四路显存峰值不受约束,每一路的增长曲线都不一样,先爆哪一路取决于模型规模、上下文长度和设备数量——你把最大的一路压下去,暴露出来的就是下一路。
四路峰值,各有各的增长曲线
论文点名的四路是:专家分发(expert dispatch)随路由矩阵增长;词表投影(vocabulary projection)随 token 数乘以词表大小增长;梯度检查点边界(checkpoint boundaries)随深度乘以序列长度增长;优化器状态随参数量增长。这四路在常用并行方案里都没有上界,而任何一路超出设备显存,训练就挂。所以论文的核心主张是:目标是同时压住每一个峰值,而不是压平均占用。
四个调度组件:GPU 工作集在启动时固定
针对四路峰值,论文给出四个调度方案,共同点是「GPU 工作集在启动时固定」:
- PipelinedLLEP:扩展 least-loaded expert parallelism,对每个数据源贡献给一个分发块的 token 数设上限,并以分块方式让通信与计算重叠;
- Ring-DTP:在词表投影处让激活或权重分片沿环形拓扑流转,把每块 logits 折叠进一个在线 log-sum-exp,避免一次性物化整个投影;
- SCO(Selective checkpoint offload):把每个检查点边界处唯一的长生命周期张量留在 CPU 内存;
- OffloadStreamAdamW:把优化器卸载后串行的 CPU Adam 更新改造成桶式流水线。
关键设计约束是:四个组件只改变计算和数据搬运的顺序与粒度,不改变数学——损失和梯度保持精确(exact),不是近似。
数字:峰值砍掉六成到八成半
匹配组件测试(matched component tests)的结果:MoE 分发峰值最高削减 59.3% 且吞吐不降;词表投影峰值削减 86.6%;被卸载的优化器步骤加速 2.05 倍。四组件组合在 120B 到 667B 参数的 MoE 模型上,以 1M 上下文长度训练——相对调优过的 FSDP2 基线,上下文可达长度是它的 8 到 32 倍,吞吐最高 10.4 倍。
泼点冷水
三点值得注意:其一,这是 9 月 13 日刚挂出的预印本(v1),外部复现尚未出现;其二,论文页没有附带代码仓库链接,四个组件的工程实现细节目前只能从 PDF 里读;其三,8 到 32 倍可达长度是与「调优过的 FSDP2 基线」的对比,这个倍数是相对值,不是绝对能力。对要上生产的人来说,先等代码或第三方复现再说。
所以呢
这篇论文的价值不在某个单一数字,而在它把「长上下文 MoE 训练爆显存」从一个笼统的工程抱怨,拆成了四路可分别治理的峰值,并证明调度层的改动可以做到数学精确。长上下文竞赛里,注意力机制的创新(稀疏、线性、混合)拿走了大部分头条,但训练基础设施这种「无聊」的工作,往往才是决定谁能真正把 1M 上下文跑起来的那块木板。