产品NVIDIA Technical Blog·原文 2026年9月15日

NVIDIA用Transformer Engine把JAX下的无丢弃MoE训练提速 10.4 倍

在DeepSeek-V3 671B上,NVIDIA通过JAX配合Transformer Engine的分组GEMM、专家并行等优化,把GB200上基线 103 TFLOPS/GPU提升到 1,068 TFLOPS/GPU,并在 1,024 张GPU上维持 97% 扩展效率。

AI解读:无丢弃MoE的好处是每个token都能被路由器选中的专家处理,模型质量不受影响,代价是每个专家收到的token数量不固定,张量形状参差不齐,常规GPU内核很难高效跑起来。NVIDIA这次做的就是把这条难走的路在JAX里铺平。

具体动作是给JAX版Transformer Engine补上三块关键能力:面向分组的MXFP8量化、跑专家矩阵乘的分组GEMM、以及优化后的专家并行dispatch/combine通信。官方给出的收益是DeepSeek-V3 671B端到端吞吐提升约 10 倍,GB200单卡算力从基线的 103 TFLOPS升到 1,068 TFLOPS。

真正受影响的不是普通用户,而是用JAX在NVIDIA多卡集群上训练MoE的工程团队。之前为了让计算规整,容量型MoE会扔掉超载token或给空位补零,前者牺牲数据完整性、后者浪费算力;分组GEMM让每个专家按真实token数计算,等于把这两笔账都省了。

限制也要说清楚:这些数字来自DeepSeek-V3 671B这一套特定配置,不同模型需要重新调参。代码随NGC MaxText容器以Transformer Engine内置的形式发布,官方建议先在小型MoE上验证正确性,再逐步放量,并跟踪grouped GEMM和dispatch/combine的延迟。

NVIDIA技术博客作者Tanya Lenz介绍了在JAX中通过NVIDIA Transformer Engine加速无丢弃(dropless)MoE训练的优化方案。

NVIDIA表示,在NVIDIA GB200上训练DeepSeek-V3时,未经优化的基线只有 103 TFLOPS/GPU,其中GPU间通信占累计内核时间的 84%;加入JAX与Transformer Engine的针对性内核优化后,该数字升至 1,068 TFLOPS/GPU,提升 10.4 倍。

在DeepSeek-V3 671B上,NVIDIA观察到端到端吞吐约 10 倍的提升;论文中的扩展数据显示,在NVIDIA GB300 NVL72机架硬件上、GPU数量从 128 扩到 2,048 的过程中,系统在 1,024 张GPU时仍保持 97% 的扩展效率。

NVIDIA把这些优化打包进内置Transformer Engine的NVIDIA NGC MaxText容器,建议从参考配置开始,先在小模型上验证正确性,再逐步放量,并跟踪step time、TFLOPS/GPU、MFU、grouped GEMM延迟和MoE dispatch/combine延迟。

为什么MoE训练难做

NVIDIA把MoE列为大规模AI模型训练的架构趋势之一,并列举DeepSeek、Qwen、Mixtral作为以更低训练算力匹配或超过稠密模型性能的例子。MoE通过条件计算提升效率:不再用一个所有token共享的稠密前馈网络(FFN),而是用多个较小的专家网络加一个学习得到的路由器,由路由器决定激活哪Top-K个专家。

生产级MoE训练会引入稠密模型没有的瓶颈:token路由、专家分发与回收、all-to-all通信、以及形状不规则的专家GEMM。问题还会叠加,因为路由器是学习出来的:训练过程中分布可能严重倾斜,路由器逐渐偏好某些专家;没有两个批次产生相同的专家负载,同一批次内一个专家也可能比另一个收到多得多的token。

每个专家收到的token数不同,就没有整齐的矩形GEMM可以批处理分发,结果形成不规则张量。NVIDIA指出,多数库针对的是形状统一、矩形的张量运算。在专家并行(EP)下,token必须被分发出去、输出再合并回原始顺序;如果分发与合并路径没优化,通信就会成为主导,GPU利用不足。

无丢弃MoE与容量型MoE的区别

NVIDIA对比了两种处理路由到专家的方式。无丢弃MoE中,无论负载多不均衡,每个token都交给它选中的专家处理;这对模型质量有吸引力,但对系统要求高。MegaBlocks的工作把专家计算重写为块稀疏矩阵乘法,让每个专家处理不同数量的token而不丢弃或填充,这需要新的块稀疏GPU内核、优化的分组GEMM,以及专门为可变token数设计的dispatch和combine原语。

标准容量型MoE训练框架则通过约束动态路由来绕开复杂性:给每个专家固定token预算,溢出部分要么裁剪要么填充。这保持计算规整、对硬件友好,但强迫在模型质量与效率之间做直接取舍:丢弃溢出token,模型就在不完整数据上训练;为避免丢弃而填充,则要在浪费的算力和内存上付出代价。

Transformer Engine在JAX中提供的三块能力

NVIDIA表示,选择无丢弃MoE意味着训练栈不能再依赖固定的专家形状,每个涉及专家计算的内核都必须高效处理可变token数;而且专家token数是数据相关的,内核不仅要接受动态形状,还要在形状对CPU不可见时正常工作,以支持CUDA graphs并避免重新编译。Transformer Engine提供的构建块包括:面向分组的MXFP8量化、用于专家矩阵乘的MXFP8分组GEMM、以及为dispatch和combine优化的EP操作。

分组GEMM的思路是把所有专家矩阵乘放进一次内核调用,每个按实际token数计算,只算有有效token的区域。NVIDIA指出,此前做法包括循环调用GEMM内核和批处理GEMM:循环需要把token数从设备拷到主机,这一步在关键路径上,带来拷贝延迟并破坏CUDA graphs;批处理GEMM则按最坏情况的token容量计算,即使实际token更少也要填充,浪费算力。Transformer Engine的grouped_gemm / ragged_dot由cuBLAS和cuBLASLt支撑,在Blackwell GPU上还为专家矩阵乘打开了MXFP8块缩放。

EP负责专家周围的一切。Dispatch阶段把token置换并跨GPU发往分配的专家,结合本地重排与多GPU通信;Combine阶段把处理后的token送回原GPU并累加各专家结果。朴素实现把这两个阶段串成一系列独立操作,GPU在步骤之间停顿,数据多次读写内存,通信与计算基本互相闲置。Transformer Engine的EP实现把Dispatch和Combine融合成紧凑的内核路径,由NCCL EP驱动;NCCL EP针对专家并行路由产生的不规则、不均衡流量模式调优,并采用token去重机制:当一个token被发往同一rank上的多个专家、或远程IB节点上的多个rank时,它只在网络上走一次,在接收节点复制,从而节省网络带宽。

  • 分组GEMM处理每个专家内部的计算,EP处理围绕专家的一切。
  • NCCL EP的token去重使一个token跨网络只传输一次,在接收节点复制。

补充优化与性能数据

NVIDIA还提到两项补充优化。JAX host offloading:中间激活不必在整个前向过程中留在设备上,JAX提供重物化API把激活卸载到主机内存;在DeepSeek-V3训练中,为省内存把query和value投影结果卸载到主机。XLA多流集合通信:EP由Transformer Engine NCCL EP驱动,优化后的FSDP则由XLA原生处理;XLA默认在单流上跑通信,本可并行执行的集合操作被串行化,一些暴露在关键路径上。多流集合通信让编译器在独立CUDA流上并发调度独立集合操作,把跨节点InfiniBand传输与节点内NVIDIA NVLink通信重叠;Latency Hiding Scheduler通过分析副本组并检查死锁风险来决定哪些集合操作可以安全重叠,无需手动标注。

如何复现DeepSeek-V3结果

NVIDIA表示这些优化随内置Transformer Engine的NVIDIA NGC MaxText容器发布。启用MaxText中的TE MoEBlock需要添加标志: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。容器需使用 2026 年 9 月 9 日或更新的版本(ghcr.io/nvidia/jax:maxtext-2026-09-09)。

要精确复现博客中的DeepSeek-V3 671B结果,需在上述基础配置上追加MaxText配置标志、XLA标志和环境变量,包括模型参数model_name: "deepseek3-671b"、max_target_length: 4096、hardware: "gpu_multiprocess",训练设置per_device_batch_size: 6、steps: 15,以及weight_dtype: "bfloat16"、capacity_factor: 1.0、megablox: false等。NVIDIA提醒该配置针对DeepSeek-V3,不同模型需要不同调参。

128 GPU的并行配置为总FSDP=16(ICI 8 × DCN 2)× EP=8,nodes: 32。XLA标志包括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等。环境变量为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。

NVIDIA计划后续加入NVFP4、量化与GEMM融合、以及A2A重叠。博客致谢了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上的贡献。

信息来源