混合专家(MoE)已成为大规模AI模型训练中的标志性架构趋势之一。DeepSeek、Qwen和Mixtral等MoE模型,在训练计算量仅为密集模型一小部分的情况下,性能却能与后者持平甚至超越。
MoE模型通过条件计算实现高效训练。它不再使用所有token共享的一个密集前馈网络(FFN),而是将其替换为许多较小的专家网络,以及一个学习型路由器,由路由器决定激活哪些* *Top-K专家。
然而,要让MoE训练在大规模下保持高效颇具挑战。在NVIDIA GB200上训练DeepSeek-V3时,未经优化的基线仅达到103 TFLOPS/GPU,而GPU间通信占用了累计内核时间的84%。借助JAX Python库和NVIDIA Transformer Engine针对性的内核优化,这一数字提升至1,068 TFLOPS/GPU,实现了10.4倍的改进。本文将讨论Transformer Engine——一个在NVIDIA GPU上加速Transformer模型的库——如何与JAX结合,为MoE模型操作带来显著的性能提升。
MoE训练涉及哪些挑战?
生产规模的MoE训练引入了密集模型所没有的瓶颈:token路由、专家分发与收集、全对全通信,以及不规则专家GEMM。
问题会进一步加剧,因为路由器是学习得到的。在整个训练过程中,随着路由器对某些专家形成偏好,分布可能变得严重倾斜。没有两个批次会产生相同的专家负载,而在单个批次内,一个专家可能接收到的token远多于另一个。每个专家接收的token数量不同,因此不存在干净的矩形GEMM可供批处理和分发。这导致了不规则张量。
在MoE中,token被动态路由到不同的专家。这意味着分配给每个专家的token数量会不可预测地变化,从而产生不规则张量(图1)。这是一个挑战,因为大多数库都针对期望统一、矩形数据结构的张量操作进行了高度优化。

图1. MoE训练中的不规则张量
在专家并行(EP)下,token必须被分发,输出必须被合并并恢复为原始token顺序。如果分发与合并路径未优化,通信将占据主导,GPU利用率不足。优化不佳的全对全通信会迫使GPU停顿,等待数据后才能进行任何有效工作。
解决这一问题需要能够原生处理不规则布局的专用内核。这正是Transformer Engine MoE优化旨在解决的问题。
无丢弃MoE与基于容量的MoE有何不同?
无丢弃MoE和基于容量的MoE是处理token路由到专家的两种不同方式。
在无丢弃 MoE 中,无论负载多么不均衡,每个 token 都会由其选定的专家处理。这对模型质量很有吸引力,但对系统要求很高。MegaBlocks: Efficient Sparse Training with Mixture-of-Experts 通过将专家计算重新表述为块稀疏矩阵乘法解决了这个问题,允许每个专家处理不同数量的 token,而无需丢弃或填充。这需要新的块稀疏 GPU 内核、优化的分组 GEMM,以及专门为可变 token 数量设计的调度和合并原语。
相比之下,标准的基于容量的 MoE 训练框架通过约束动态路由来规避其复杂性。每个专家被分配固定的 token 预算,任何溢出要么被裁剪,要么被填充以适应。这保持了计算的规律性和硬件友好性,但迫使模型质量与效率之间进行直接权衡:丢弃溢出的 token,模型就在不完整的数据上训练;或者填充以避免丢弃,但付出浪费计算和内存的代价。

图 2. 基于容量的 MoE 与无丢弃 MoE 的对比
无丢弃 MoE 需要哪些专门的优化?
致力于无丢弃 MoE 意味着训练栈不能再依赖固定的专家形状。每个涉及专家计算的内核都必须高效处理可变 token 数量。此外,这意味着每个专家的 token 数量是可变的且依赖于数据,因此内核不仅必须接受动态形状,还必须在这些形状在 CPU 上不可访问时也能工作,以启用 CUDA 图并避免重新编译。
Transformer Engine 提供了以下构建块,使这种方法在 JAX 中变得实用:
- 组感知的 MXFP8 量化
- 在专家矩阵乘法上的 MXFP8 分组 GEMM
- 用于调度和合并的优化 EP 操作
图 3 显示了跨两个 GPU 的专家并行 MoE 层。路由器将每个 token 分配给一个专家,调度将 token 移动到其专家的 GPU。分组 MLP 在这些可变长度组上运行两个分组 GEMM,合并反转交换以恢复原始 token 顺序。

图 3. 跨两个 GPU 的专家并行 MoE 层
优化 1:分组 GEMM
在密集 FFN 中,每个 token 都通过相同的权重矩阵。在 MoE 中,路由器不均匀地分配 token,因此每个专家在每一步接收到不同数量的 token,打破了典型内核所优化的规则 GEMM 形状。
先前的方法包括一个 GEMM 内核循环和批量 GEMM。该循环需要将令牌计数从设备复制到主机。这位于关键路径上,会引入设备到主机传输的延迟并破坏 CUDA 图。批量 GEMM 即使使用的令牌更少,也会计算最坏情况下的令牌容量,因为它们被填充以强制固定的专家计算,导致额外的计算。
分组 GEMM 通过在一次内核调用中处理所有专家矩阵乘法来解决这个问题,每个专家矩阵乘法使用其实际的令牌计数。它只计算具有有效令牌的区域,因此性能更高。
Transformer Engine grouped_gemm /ragged_dot
由 cuBLAS 和 cuBLASLt 提供支持,直接映射到性能最佳的 NVIDIA GEMM 库,即使专家形状不规则也能实现完整的 Tensor Core 利用率。在 NVIDIA Blackwell GPU 上,此路径还利用 Transformer Engine 分组量化内核为专家矩阵乘法开启了 MXFP8 块缩放。
优化 2:专家并行以集成 Dispatch 和 Combine
在融合路由器内核将每个令牌分配给其专家后,模型必须将这些令牌物理移动到正确的设备,处理它们,并将结果带回。
此过程分为两个不同的阶段:Dispatch 和 Combine。
Dispatch:令牌移动发生的地方:令牌被置换并跨 GPU 发送到其分配的专家,这一步涉及本地重排序和多 GPU 通信。Combine:处理后的令牌被路由回其原始 GPU,并累积每个专家的结果。
在朴素实现中,这些阶段作为独立操作的串行链运行,GPU 在步骤之间停顿,数据多次读写内存,通信在计算运行时大多空闲,反之亦然。
Transformer Engine EP 实现将 Dispatch 和 Combine 阶段集成到一个紧密融合的内核路径中。这种集成由 NCCL EP 提供支持,这是一个专门针对专家并行路由产生的不规则、不平衡流量模式调优的通信后端。
NCCL EP 还采用令牌去重机制:当一个令牌被分派到同一 rank 上的多个专家或远程 IB 节点上的多个 rank 时,它只遍历网络一次,并在接收节点上复制,从而节省网络带宽。EP 是分组 GEMM 的对应物:分组 GEMM 处理每个专家内部发生的事情;EP 处理其周围的一切。
额外优化
额外优化包括 JAX 主机卸载和 XLA 多流集合。
JAX 主机卸载
中间激活不必在整个前向传递过程中保存在设备上。JAX 提供了重新物化 API,用于将激活卸载到主机内存。为了在 DSv3 训练中节省内存,将查询和值投影结果卸载到主机。要了解更多信息,请参阅使用主机卸载减少基于 JAX 的 LLM 训练中的高带宽内存瓶颈。
XLA 多流集合
EP 由 Transformer Engine NCCL EP 驱动,而优化后的 FSDP 则在 XLA 中原生处理。默认情况下,XLA 在单个流上运行通信,因此本可并行执行的集合通信操作会被串行化,其中一些最终暴露在关键路径上。多流集合通信让编译器能够在独立的 CUDA 流上并发调度相互独立的集合通信操作,将跨节点的 InfiniBand 传输与节点内的 NVIDIA NVLink 通信重叠起来,从而同时利用两种互连结构,而不是等待单个串行流。
延迟隐藏调度器(LHS)通过分析集合通信操作的副本组并检查死锁风险,来决定哪些集合通信操作可以安全地重叠,因此内存带宽方面的收益是自动获得的,无需手动标注。这显著降低了 DSv3 训练中暴露的集合通信操作所占的比例。
在 JAX 中使用 Transformer Engine 的 MoE 对训练性能有何影响?
我们观察到,通过结合 Transformer Engine 优化的 JAX MoE,DeepSeek-V3 671B 实现了 10 倍的端到端吞吐量提升。
回想一下,基线 JAX 训练栈此前让大部分硬件潜力白白闲置。我们逐层攻克整个栈,加入了 cuBLAS GroupedGEMM、XLA 多流集合通信、MXFP8 GroupQuant、主机激活卸载,最后是优化后的 EP 实现。
我们计划加入 NVFP4、与 GEMM 融合的量化,以及 A2A 重叠。要了解更多关于 Transformer Engine JAX 绑定将支持的未来内核融合,请参阅利用高级融合内核提升 MoE 训练吞吐量。

图 4. 在 MaxText 上对 DeepSeek-V3 671B 进行的端到端训练性能,以 TFLOPs/GPU/秒和 tokens/GPU/秒衡量,展示了在 NVIDIA GB300 NVL72 机架硬件上四个渐进式优化阶段的表现
使用 JAX 实现的高多机架扩展性能
大规模训练大型模型需要激进的优化。在生产规模下,这意味着数万亿的 token 和巨大的批次大小低效问题。虽然这些问题在单个节点上微不足道,但在数千个 GPU 上会迅速累积,使得计算、内存和通信中的每一个瓶颈都变得至关重要。
多机架扩展是大多数系统面临困难的地方,因为通信开销的增长速度往往快于计算。应用 JAX MoE 和 Transformer Engine 栈后,这种性能退化得到了显著控制。系统在 1,024 个 GPU 上保持了 97% 的效率,这一结果直接说明了底层通信优化在集群扩展时保持吞吐量的有效性。

图 5. DeepSeek-V3 671B 在 NVIDIA GB300 NVL72 上的多机架扩展性能。TFLOPs/GPU/秒(绿色柱,左轴)以及相对于 128-GPU 基线的扩展效率(黑线,右轴),覆盖从 128 到 2,048 的五个 GPU 数量
如何开始无丢弃 MoE 训练
这些优化已随 NVIDIA NGC MaxText 容器 发布,内置 Transformer Engine,因此你可以直接复现并在此基础上构建。要开始使用,请尝试使用 启用 Transformer Engine 的 NVIDIA NGC MaxText 容器 的优化 JAX MoE 路径。
从参考配置开始,在小规模 MoE 模型上验证正确性,然后逐步扩大规模,同时跟踪步进时间、TFLOPS/GPU、MFU、分组 GEMM 延迟以及 MoE 分发/合并延迟。
基本使用配置:MaxText 中的 TE MoEBlock
要在 MaxText 中启用 TE MoEBlock,请将以下标志添加到你的 MaxText YAML 配置中,或将它们作为命令行参数传递给训练脚本。
容器
使用 2026 年 9 月 9 日的容器(ghcr.io/nvidia/jax:maxtext-2026-09-09)或更新版本。更多详情,请参阅 NVIDIA/JAX-Toolbox 的容器镜像部分 GitHub 仓库。
MaxText 配置(MaxText moe_configuration.md):
te_moe_block: true
te_gmm_quantization: "te_mxfp8"
ragged_buffer_factor: 2.0
te_ep_overflow_check_every_n_steps: 20
sparse_matmul: true
prefuse_moe_weights: true
DeepSeek V3 的性能复现
要精确复现本文中展示的 DeepSeek-V3 671B 结果,请使用以下 MaxText 配置标志、XLA 标志和环境变量扩展基本用法配置。请注意,此配置特定于 DeepSeek-V3;不同的模型需要不同的调优。并非必须使用 TE MoEBlock 本身。
MaxText 配置
MaxText 配置如下所示。有关参数详情,请参阅 MaxText MoE 配置指南。
# 模型参数
model_name: "deepseek3-671b"
max_target_length: 4096
hardware: "gpu_multiprocess"
# 训练设置
per_device_batch_size: 6
gradient_accumulation_steps: 1
steps: 15
attention: "cudnn_flash_te"
remat_policy: "custom"
# 使用 MXFP8 分组 GEMM 的 Transformer Engine MoEBlock
quantization: "te_fp8_currentscaling"
te_moe_block: true
te_gmm_quantization: "te_mxfp8"
ragged_buffer_factor: 2.0
te_ep_overflow_check_every_n_steps: 20
prefuse_moe_weights: true
weight_dtype: "bfloat16"
mu_dtype: "bfloat16"
# 特性
pgle: true
profiler: "xplane"
scan_layers: true
zero_one: false
shardy: true
use_segment: false
skip_first_n_steps_for_profiler: 4
custom_remat_enabled: true
logits_dot_in_fp32: false
use_iota_embed: false
custom_remat_config:
mlpwi: device
mlpwi_0: device
mlpwi_1: device
mlpwo: device
moe_mlpwi_0: offload #remat
moe_mlpwi_1: offload #remat
moe_mlpwo: device
query_proj: remat #offload
key_proj: remat
value_proj: remat #offload
query_wa_proj: device
kv_wa_proj: device
out_proj: device
context: device
# MoE 路由参数
n_routing_groups: -1
topk_routing_group: -1
capacity_factor: 1.0
megablox: false
# 128 个 GPU:总计 FSDP=16(ICI 8 × DCN 2)× EP=8。
nodes: 32
ici_data_parallelism: 1
ici_fsdp_parallelism: 8
ici_tensor_parallelism: 1
ici_expert_parallelism: 8
dcn_data_parallelism: 1
dcn_fsdp_parallelism: 2
dcn_tensor_parallelism: 1
dcn_expert_parallelism: 1
shard_optimizer_over_data: false
shard_exp_on_fsdp: false
XLA 标志调优
有关 XLA 标志调优的指导,请参阅 XLA GPU 标志指南和 JAX Toolbox GPU 性能指南。
xla_gpu_all_reduce_combine_threshold_bytes: 33554432
xla_gpu_all_gather_combine_threshold_bytes: 6442450944
xla_gpu_reduce_scatter_combine_threshold_bytes: 201326592
xla_gpu_experimental_enable_nccl_symmetric_buffers: false
xla_gpu_enable_command_buffer: "'FUSION,CUBLAS,CUDNN,DYNAMIC_SLICE_FUSION'"
xla_gpu_experimental_max_unroll_factor: 8
xla_gpu_memory_limit_slop_factor: 99
环境变量
XLA_PYTHON_CLIENT_MEM_FRACTION: 0.88
CUDA_DEVICE_MAX_CONNECTIONS: 16
XLA_PJRT_GPU_HOST_MEMORY_PREALLOCATE: false
XLA_PJRT_GPU_HOST_MEMORY_LIMIT_GB: 180
了解更多
无丢弃 MoE 训练在保持模型质量的同时,借助 Transformer Engine 的分组 GEMM 和 EP 内核,在大规模下依然高效。该方法在 DeepSeek-V3 671B 上实现了约 10 倍的吞吐量提升,并在 1,024 块 GPU 上达到 97% 的扩展效率。这些优化已随内置 Transformer Engine 的 NVIDIA NGC MaxText 容器 一同发布,因此你可以直接复现并在此基础上继续构建。
有关在 MaxText 中使用 Transformer Engine MoE 块的信息,请参阅 MaxText MoE 配置指南。有关 Transformer Engine 的更多信息,请参阅 Transformer Engine 文档。
致谢
特别感谢 Abhinav Goel、MD Fahim Faysal Khan、Jane Liu、Terry Sun、Tj Xu、Ming Huang、Chase Roberts 和 Oleg Goncharov 为 JAX、XLA 和 Transformer Engine 中 MoE 的启用与优化所做的贡献。感谢 Artem Polyakov、Ke Wen 和 Subhadeep Bhattacharya 为 NCCL EP 所做的贡献,以及 Igor Safanov 为 cuBLASLt 所做的贡献。
