混合专家(MoE)架构已成为大规模AI模型训练中最具代表性的架构趋势之一。DeepSeek、Qwen和Mixtral等模型都是MoE模型的典型案例,它们能以远低于密集模型的训练算力,达到相当甚至更优的性能表现。

MoE模型通过条件计算实现高效训练。它并非让所有Token共享一个密集前馈网络(FFN),而是用多个更小的专家网络加上一个可学习的路由器来替代,由路由器决定激活哪些Top-K专家。

然而,要在大规模场景下让MoE训练保持高效并不容易。在NVIDIA GB200上进行DeepSeek-V3训练时,未经优化的基线方案每个GPU仅能达到103 TFLOPS,且GPU间通信占据了累计内核时间的84%。而借助JAX Python库与NVIDIA Transformer Engine的针对性内核优化,这一数字提升到了每GPU 1068 TFLOPS,实现了10.4倍的提升。本文将探讨Transformer Engine(一个用于在NVIDIA GPU上加速Transformer模型的库)与JAX结合后,如何在MoE模型运算中带来显著的性能提升。

MoE训练面临哪些挑战

生产级规模的MoE训练会引入密集模型中不存在的瓶颈:Token路由、专家分发与聚合、all-to-all通信,以及不规则的专家矩阵乘法(GEMM)。

由于路由器本身是可学习的,这一问题会进一步加剧。随着训练的推进,路由器会对某些专家形成偏好,导致分布出现严重倾斜。没有任何两个批次会产生相同的专家负载,即便在同一批次内,某个专家收到的Token数量也可能远超另一个专家。每个专家接收到的Token数量各不相同,因此无法形成规整的矩形GEMM来进行批处理和分发,由此产生了不规则张量。

在MoE中,Token会被动态路由到不同的专家。这意味着分配给每个专家的Token数量是不可预测的,从而形成不规则张量(如图1所示)。这是一个挑战,因为大多数库都是针对期望统一矩形数据结构的张量运算进行高度优化的。

在采用专家并行(EP)时,必须对Token进行分发,并将输出合并、还原为原始Token顺序。如果分发与合并路径未经优化,通信将占据主导地位,导致GPU利用率不足。一个优化不佳的all-to-all通信会迫使GPU停滞,等待数据到位后才能开始有效计算。

要解决这一问题,需要能够原生处理不规则布局的专用内核,而这正是Transformer Engine MoE优化所要解决的核心问题。

无丢弃MoE与基于容量的MoE有何不同

无丢弃MoE和基于容量的MoE是处理Token路由到专家这一问题的两种不同方式。

在无丢弃MoE中,无论负载多么不均衡,每个Token都会由其所选专家处理。这种方式对模型质量有利,但对系统要求较高。论文《MegaBlocks: Efficient Sparse Training with Mixture-of-Experts》通过将专家计算重新表述为块稀疏矩阵乘法来解决这一问题,使每个专家可以处理不同数量的Token,而无需丢弃或填充。这需要专门为可变Token数量设计的全新块稀疏GPU内核、优化的分组GEMM,以及分发与合并原语。

相比之下,标准的基于容量的MoE训练框架通过限制动态路由的方式,回避了这种复杂性。每个专家被分配固定的Token预算,任何溢出部分要么被裁剪,要么被填充以适配预算。这样可以让计算保持规整、对硬件友好,但也迫使模型质量与效率之间做出直接权衡:丢弃溢出的Token会导致模型在不完整数据上训练,而填充则要以浪费计算和内存为代价来避免丢弃。

无丢弃MoE需要哪些专门优化

选择无丢弃MoE,意味着训练框架不能再依赖固定的专家形状。每一个涉及专家计算的内核都必须高效处理可变数量的Token。此外,这也意味着每个专家的Token数量是可变且依赖数据的,因此内核不仅要支持动态形状,还要能在CPU无法获取这些形状信息的情况下正常工作,从而支持CUDA图并避免重新编译。

Transformer Engine为JAX提供了以下构建模块,使这一方案在实践中变得可行:

支持分组感知的MXFP8量化

针对专家矩阵乘法的MXFP8分组GEMM

针对分发与合并优化的EP操作

图3展示了跨两个GPU的专家并行MoE层。路由器将每个Token分配给某个专家,分发操作将Token移动到对应专家所在的GPU。分组MLP在这些可变长度的分组上运行两次分组GEMM,合并操作则反向执行数据交换,将Token还原为原始顺序。

优化一:分组GEMM

在密集FFN中,每个Token都会经过相同的权重矩阵。而在MoE中,路由器会不均匀地分发Token,使得每个专家在每一步接收到的Token数量各不相同,这打破了常规内核所优化的规则GEMM形状。

以往的方法包括GEMM内核循环和批量GEMM。循环方式需要将Token数量从设备端拷贝到主机端,这一操作处于关键路径上,会带来设备到主机传输的延迟,并破坏CUDA图的连续性。批量GEMM则会按最坏情况的Token容量进行计算,即便实际使用的Token更少,因为它们被填充以强制形成固定的专家计算规模,从而带来额外的计算开销。

分组GEMM通过在单次内核调用中处理所有专家的矩阵乘法,并按各自实际的Token数量进行计算,解决了这一问题。它只计算有效Token所在的区域,因此性能表现更优。

Transformer Engine的grouped_gemm/ragged_dot以cuBLAS和cuBLASLt作为底层支撑,直接映射到性能最优的NVIDIA GEMM库,即使在专家形状不规则的情况下,也能实现Tensor Core的充分利用。在NVIDIA Blackwell GPU上,这一路径还借助Transformer Engine的分组量化内核,为专家矩阵乘法开启了MXFP8块缩放能力。

优化二:整合专家并行的分发与合并操作

在融合路由器内核为每个Token分配好专家之后,模型必须将这些Token实际移动到对应的设备上,完成处理,并将结果传回。

这一过程可分为两个独立阶段:分发(Dispatch)与合并(Combine)。

分发:Token移动发生的阶段,Token被重新排列并跨GPU发送到各自分配的专家,这一步既涉及本地重排,也涉及多GPU通信。

合并:处理完成的Token被路由回其原始GPU,各专家的结果在此阶段被累加汇总。

在简单实现中,这两个阶段会作为一连串独立操作依次串行执行,GPU在各步骤之间会出现停滞,数据需要多次读写内存,且在计算运行时通信大多处于空闲状态,反之亦然。

Transformer Engine的EP实现将分发与合并阶段整合为一条紧密融合的内核路径。这一整合依托于NCCL EP实现,这是一种专门针对专家并行路由所产生的不规则、不均衡流量模式而调优的通信后端。

NCCL EP还采用了一种Token去重机制:当某个Token被分发给同一计算节点上的多个专家,或分发到远程InfiniBand节点上的多个计算节点时,它只会在网络中传输一次,并在接收节点上进行复制,从而节省网络带宽。EP与分组GEMM是相辅相成的:分组GEMM负责处理每个专家内部的计算,而EP则负责处理其周边的一切事务。

其他优化

其他优化还包括JAX主机卸载和XLA多流集合通信。

JAX主机卸载

中间激活值并不需要在整个前向传播过程中都保存在设备上。JAX提供了重计算(rematerialization)API,可将激活值卸载至主机内存。为了在DSv3训练中节省内存,可将查询(query)和值(value)投影结果卸载到主机端。相关内容可参阅《Reducing High-Bandwidth Memory Bottlenecks in JAX-Based LLM Training with Host Offloading》一文。

XLA多流集合通信

尽管EP由Transformer Engine的NCCL EP驱动,但优化后的FSDP则由XLA原生处理。默认情况下,XLA在单一流上执行通信,因此原本可以并行执行的集合通信操作会被串行化,其中一部分甚至会暴露在关键路径上。多流集合通信允许编译器在独立的CUDA流上并发调度互不依赖的集合通信操作,使跨节点的InfiniBand传输与节点内的NVIDIA NVLink通信能够同时进行,从而同时利用两种互连结构,而不是等待单一串行流完成。

延迟隐藏调度器(Latency Hiding Scheduler,LHS)通过分析各集合通信操作的副本组并检查死锁风险,来判断哪些操作可以安全重叠执行,因此带宽收益是自动获得的,无需人工标注。这大幅降低了DSv3训练中暴露在关键路径上的集合通信比例。

JAX结合Transformer Engine的MoE训练性能提升如何

我们观察到,在DeepSeek-V3 671B模型上,借助JAX结合Transformer Engine的MoE优化,端到端吞吐量实现了10倍提升。

回顾此前的情况,基线JAX训练框架远未发挥出硬件的全部潜力。通过在每一层逐一优化框架,我们相继加入了cuBLAS分组GEMM、XLA多流集合通信、MXFP8分组量化、主机激活卸载,以及最终优化的EP实现。

我们计划进一步加入NVFP4、与GEMM融合的量化,以及A2A重叠等优化。关于Transformer Engine JAX绑定未来将支持的内核融合方案,可参阅《Boosting MoE Training Throughput with Advanced Fusion Kernels》一文。

JAX在多机架规模下的扩展性能表现如何

大规模训练模型需要极致优化。在生产规模下,这意味着数万亿Token以及巨大的批次规模所带来的各种低效问题。这些问题在单节点上可以忽略不计,但在数千个GPU上会迅速累积放大,使计算、内存和通信中的每一个瓶颈都变得至关重要。

多机架扩展是大多数系统面临瓶颈的环节,因为通信开销往往比计算增长得更快。而在应用了JAX MoE与Transformer Engine技术栈后,这种性能衰减得到了显著控制。系统在1024个GPU规模下仍能保持97%的效率,这一结果充分印证了底层通信优化在集群规模扩大时保持吞吐量的有效性。

如何开始无丢弃MoE训练

这些优化已内置于集成了Transformer Engine的NVIDIA NGC MaxText容器中,用户可以直接复现并在此基础上进行开发。要开始使用,可尝试通过启用Transformer Engine的NVIDIA NGC MaxText容器来体验优化后的JAX MoE路径。

建议先从参考配置入手,在小型MoE模型上验证正确性,然后逐步扩大规模,同时跟踪单步耗时、每GPU的TFLOPS、MFU、分组GEMM延迟以及MoE分发/合并延迟等指标。

(此处省略具体配置代码及参数细节,详见原始技术文档中的MaxText配置说明、XLA标志调优建议及环境变量设置。)

延伸阅读

无丢弃MoE训练能够保持模型质量,而Transformer Engine的分组GEMM与EP内核则使其能够在大规模场景下高效运行。这一方案在DeepSeek-V3 671B模型上实现了约10倍的吞吐量提升,并在1024个GPU规模下达到97%的扩展效率。这些优化已内置于集成Transformer Engine的NVIDIA NGC MaxText容器中,用户可直接复现并在此基础上继续开发。

关于在MaxText中使用Transformer Engine MoE模块的更多信息,请参阅MaxText MoE配置指南;关于Transformer Engine的更多信息,请参阅Transformer Engine相关文档。

Q&A

Q1:什么是无丢弃MoE训练?

A:无丢弃MoE训练是指无论专家间负载多么不均衡,每个Token都会被其所选择的专家完整处理,不会被丢弃或截断,这样能更好地保证模型质量,但对训练系统的要求也更高。

Q2:使用Transformer Engine优化MoE训练能带来多大提升?

A:在DeepSeek-V3 671B模型上,借助JAX与Transformer Engine的优化,端到端训练吞吐量实现了约10倍提升,并且在1024个GPU的大规模集群中依然能保持97%的扩展效率。

Q3:普通开发者如何开始尝试这种优化的MoE训练方案?

A:可以直接使用集成了Transformer Engine的NVIDIA NGC MaxText容器,先在小型MoE模型上验证配置的正确性,再逐步扩大规模并跟踪相关性能指标。

NVIDIA