Accelerating Dropless MoE Training in JAX with NVIDIA Transformer Engine
Mirrored from NVIDIA Developer Blog for archival readability. Support the source by reading on the original site.
Accelerating Dropless MoE Training in JAX with NVIDIA Transformer Engine
AI-Generated Summary
- NVIDIA Transformer Engine with JAX delivers a 10.4x throughput improvement for Mixture of Experts training on NVIDIA GB200, raising DeepSeek-V3 performance from 103 to 1,068 TFLOPS/GPU.
- Dropless MoE preserves model quality by processing every token without dropping or padding, and Transformer Engine enables this through grouped GEMM kernels that handle variable expert token counts in a single call.
- Expert parallelism operations are accelerated by NCCL EP, which fuses dispatch and combine stages and deduplicates tokens to reduce network traffic.
- Additional optimizations including JAX host offloading and XLA multistreaming collectives further reduce memory bottlenecks and overlap communication across NVLink and InfiniBand fabrics.
- The full stack sustains 97% scaling efficiency at 1,024 GPUs on NVIDIA GB300 NVL72 hardware when training DeepSeek-V3 671B.
Next Steps
- Try the NVIDIA NGC MaxText container with Transformer Engine enabled to reproduce the optimized JAX MoE path.
- Read the MaxText MoE Configuration guide for detailed setup instructions.
- Review the Transformer Engine documentation to understand the library's capabilities.
Mixture of experts (MoE) has become one of the defining architectural trends in large-scale AI model training. DeepSeek, Qwen, and Mixtral are examples of MoE models that match or exceed the performance of dense model counterparts at a fraction of the training compute.
MoE models provide efficient training through conditional computation. Instead of one dense feed-forward network (FFN) shared by all tokens, MoE replaces it with many smaller expert networks and a learned router that decides which Top-K experts to activate.
However, making MoE training efficient at scale is challenging. In DeepSeek-V3 training on NVIDIA GB200, an unoptimized baseline achieved just 103 TFLOPS/GPU with inter-GPU communication consuming 84% of accumulated kernel time. With the JAX Python library and NVIDIA Transformer Engine targeted kernel optimizations, that number rose to 1,068 TFLOPS/GPU, a 10.4x improvement. This post discusses how Transformer Engine, a library for accelerating Transformer models on NVIDIA GPUs, with JAX leads to significant performance improvement in MoE model operations.
What are the challenges involved in MoE training?
Production-scale MoE training introduces bottlenecks that don’t exist with dense models: token routing, expert dispatch and gather, all-to-all communication, and ragged expert GEMMs.
The problem compounds because the router is learned. Throughout training, the distribution can become heavily skewed as the router develops preferences for certain experts. No two batches produce the same expert loads, and within a single batch, one expert might receive many more tokens than another. Each expert receives a different number of tokens, so there is no clean rectangular GEMM to batch and dispatch. This results in ragged tensors.
In MoE, tokens are routed dynamically to different experts. This means that the number of tokens assigned to each expert varies unpredictably, which results in ragged tensors (Figure 1). This is a challenge because most libraries are highly optimized for tensor operations that expect uniform, rectangular data structures.
With expert parallelism (EP), tokens must be dispatched and outputs must be combined and restored to original token order. If the dispatch and combine path is not optimized, communication dominates and GPUs are underutilized. A poorly optimized all-to-all forces GPUs to stall and wait for data before doing any useful work.
Solving this requires specialized kernels that can natively handle ragged layouts. This is precisely the problem that Transformer Engine MoE optimizations are designed to solve.
How is dropless MoE different from capacity-based MoE?
Dropless and capacity-based MoE are two different ways to handle token routing to experts.
In dropless MoE, every token is processed by its selected expert no matter how uneven the load. This is attractive for model quality but demanding on the system. MegaBlocks: Efficient Sparse Training with Mixture-of-Experts addressed this by reformulating expert computation as block-sparse matrix multiplication, allowing each expert to operate on a different number of tokens without dropping or padding. This requires new block-sparse GPU kernels, optimized grouped GEMM, and dispatch and combine primitives all designed specifically for variable token counts.
In comparison, standard capacity-based MoE training frameworks sidestep the complexity of dynamic routing by constraining it. Each expert is assigned a fixed token budget, and any overflow is either trimmed or padded to fit. This keeps computation regular and hardware-friendly, but it forces a direct tradeoff between model quality and efficiency: drop the overflow tokens and the model trains on incomplete data, or pad to avoid dropping and pay the cost in wasted compute and memory.
What specialized optimizations are required for dropless MoE?
Committing to dropless MoE means the training stack can no longer rely on fixed expert shapes. Every kernel that touches expert computation has to handle variable token counts efficiently. Additionally, it means that each expert’s token count is variable and data-dependent, so the kernels must not only accept dynamic shapes but also work when those shapes are inaccessible on the CPU to enable CUDA graphs and avoid recompilation.
Transformer Engine provides the following building blocks that make this approach practical in JAX:
- A group-aware MXFP8 quantization
- An MXFP8 grouped GEMM on expert matmuls
- Optimized EP operations for dispatch and combine
Figure 3 shows an expert-parallel MoE layer across two GPUs. The router assigns each token to an expert, dispatch moves tokens to their expert’s GPU. The grouped MLP runs two grouped GEMMs on those variable-length groups, and combine reverses the exchange to restore the original token order.
Optimization 1: Grouped GEMM
In a dense FFN, every token passes through the same weight matrix. In MoE, the router distributes tokens unevenly so each expert receives a different number of tokens per step, breaking the regular GEMM shape that typical kernels are optimized for.
Previous approaches included a loop of GEMM kernels and batched GEMMs. The loop required Device-to-Host copies of token counts. This is on the critical path, which incurs the latency of the Device-to-Host transfer and breaks CUDA graphs. The batched GEMM computed the worst-case token capacity even if fewer tokens were used because they are padded to force fixed expert computation, leading to extra compute.
A grouped GEMM solves this by handling all expert matmuls in a single kernel call, each with its actual token count. It computes only the regions with valid tokens and is more performant as a result.
Transformer Engine grouped_gemm /ragged_dot backs this with cuBLAS and cuBLASLt, mapping directly onto the best-performing NVIDIA GEMM libraries to deliver full Tensor Core utilization even with irregular expert shapes. On NVIDIA Blackwell GPUs, this path also opens up MXFP8 block scaling for expert matmuls utilizing the Transformer Engine grouped quantization kernels.
Optimization 2: Expert parallelism to integrate Dispatch and Combine
After the fused router kernels assign each token to its experts, the model must physically move those tokens to the correct devices, process them, and bring the results back.
This process breaks into two distinct stages: Dispatch and Combine.
- Dispatch: Where the token movement occurs: tokens are permuted and sent across GPUs to their assigned experts, a step that involves both local reordering and multi-GPU communication.
- Combine: Where processed tokens are routed back to their original GPUs and their per-expert results are accumulated.
In a naive implementation, these stages run as a serial chain of separate operations, with the GPU stalling between steps, data getting read and written to memory multiple times, and communication sitting mostly idle while compute runs and vice versa.
The Transformer Engine EP implementation integrates the Dispatch and Combine stages into a tightly fused kernel path. This integration is powered by NCCL EP, a communication backend tuned specifically for the irregular, imbalanced traffic patterns that expert-parallel routing produces.
NCCL EP also employs a token deduplication mechanism: when a token is dispatched to multiple experts on the same rank or to multiple ranks on a remote IB node, it traverses the network only once and is replicated on the receiving node, conserving network bandwidth. EP is the counterpart to grouped GEMM: grouped GEMM handles what happens inside each expert; EP handles everything around it.
Additional optimizations
Additional optimizations include JAX host offloading and XLA multistreaming collectives.
JAX host offloading
Intermediate activations don’t have to be saved on device for the entire forward pass. JAX provides rematerialization APIs for offloading activations to host memory. To save memory in DSv3 training, offload the query and value projection results to host. To learn more, see Reducing High-Bandwidth Memory Bottlenecks in JAX-Based LLM Training with Host Offloading.
XLA multistreaming collectives
While EP is driven by Transformer Engine NCCL EP, optimized FSDP is handled natively in XLA. By default, XLA runs communication on a single stream, so collectives that could execute in parallel are serialized and some end up exposed on the critical path. Multi-stream collectives let the compiler schedule independent collectives concurrently across separate CUDA streams, overlapping cross-node InfiniBand transfers with intra-node NVIDIA NVLink communication to draw on both fabrics at once rather than waiting on one serialized stream.
The Latency Hiding Scheduler (LHS) decides which collectives are safe to overlap by analyzing their replica groups and checking for deadlock risk, so the memory-bandwidth gains are automatic and require no manual annotation. This reduces the percentage of exposed collectives in DSv3 training significantly.
What is the training performance impact of MoE in JAX with Transformer Engine?
We observed a 10x end-to-end throughput gain on DeepSeek-V3 671B through MoE in JAX with Transformer Engine optimizations.
Recall that the baseline JAX training stack was leaving most of the hardware potential on the table. Tackling the stack at each layer, we added cuBLAS GroupedGEMM, XLA multistream collectives, MXFP8 GroupQuant, host activation offloading, and finally an optimized EP implementation.
We plan to add NVFP4, quantization fused with GEMM, and A2A overlap. To learn more about future kernel fusions that will be supported in Transformer Engine JAX bindings, see Boosting MoE Training Throughput with Advanced Fusion Kernels.
High multirack scaling performance with JAX
Training large models at scale demands aggressive optimization. At production scale, this amounts to trillions of tokens and massive batch size inefficiencies. While these are negligible on a single node, they can compound quickly across thousands of GPUs, making every bottleneck in compute, memory, and communication critical to address.
Multirack scaling is where most systems struggle, as communication overhead tends to scale faster than compute. With JAX MoE and Transformer Engine stack applied, this degradation stays remarkably in check. The system sustains 97% efficiency at 1,024 GPUs, a result that speaks directly to the effectiveness of the underlying communication optimizations in preserving throughput as the cluster grows.
How to get started with dropless MoE training
The optimizations ship in the NVIDIA NGC MaxText container with Transformer Engine built in, so you can reproduce and build on them directly. To get started, try the optimized JAX MoE path using the NVIDIA NGC MaxText container with Transformer Engine enabled.
Start with the reference configuration, validate correctness on a small MoE model, then scale up while tracking step time, TFLOPS/GPU, MFU, grouped GEMM latency, and MoE dispatch/combine latency.
Basic usage configuration: TE MoEBlock with MaxText
To enable the TE MoEBlock in MaxText, add the following flags to your MaxText YAML config or pass them as command-line arguments to the training script.
Container
Use the container from September 9, 2026 (ghcr.io/nvidia/jax:maxtext-2026-09-09) or newer. For more details, refer to the container images section of the NVIDIA/JAX-Toolbox GitHub repo.
MaxText configuration (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
Performance reproduction for DeepSeek V3
To exactly reproduce the DeepSeek-V3 671B results presented in this post, extend the basic usage configuration with the following MaxText config flags, XLA flags, and environment variables. Note that this configuration is specific to DeepSeek-V3; different models will require different tuning. It is not required to use the TE MoEBlock itself.
MaxText configuration
The MaxText configuration is provided below. For parameter details, refer to the MaxText MoE Configuration guide.
# Model parameters model_name: "deepseek3-671b" max_target_length: 4096 hardware: "gpu_multiprocess" # Training settings per_device_batch_size: 6 gradient_accumulation_steps: 1 steps: 15 attention: "cudnn_flash_te" remat_policy: "custom" # Transformer Engine MoEBlock with MXFP8 grouped GEMMs 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" # Features 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 routing parameters n_routing_groups: -1 topk_routing_group: -1 capacity_factor: 1.0 megablox: false # 128 GPUs: total 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 flag tuning
For guidance on XLA flag tuning, refer to the XLA GPU flags guide and JAX Toolbox GPU performance guide.
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
Environment variables
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
Learn more
Dropless MoE training preserves model quality, while Transformer Engine grouped GEMM and EP kernels make it efficient at scale. This approach drives a ~10x throughput improvement and 97% scaling efficiency to 1,024 GPUs on DeepSeek-V3 671B. These optimizations ship in the NVIDIA NGC MaxText container with Transformer Engine built in, so you can reproduce and build on them directly.
For information on using the Transformer Engine MoE block in MaxText, refer to the MaxText MoE Configuration guide. For more information on Transformer Engine, refer to the Transformer Engine documentation.
Acknowledgments
Special thanks to Abhinav Goel, MD Fahim Faysal Khan, Jane Liu, Terry Sun, Tj Xu, Ming Huang, Chase Roberts, and Oleg Goncharov for their contributions to MoE enablement and optimization in JAX, XLA, and Transformer Engine. Thanks to Artem Polyakov, Ke Wen, and Subhadeep Bhattacharya for their contributions to NCCL EP and to Igor Safanov for cuBLASLt contributions.
Tags
About the Authors
Seonghee Lee is an engineer on the AI platform software team at NVIDIA, focusing on AI Inference-related products. Seonghee holds a master’s in computer science from Stanford University and a bachelor’s in science from Cornell University, specializing in AI. Before joining NVIDIA, she worked at Microsoft Research on developing real-time AI agent interactions.
Jeremy Berchtold is a senior software performance engineer at NVIDIA developing Transformer Engine, a library for accelerating large-scale transformer model training and inference. His work focuses on Transformer Engine JAX integrations and optimized implementations for Mixture of Experts, quantization, attention, and other performance-critical model components. Prior to Transformer Engine, he contributed to the autonomous vehicle stack at NVIDIA, focusing on complex intersection and junction handling.
Phuong Nguyen is a senior DL performance engineer at NVIDIA, working on developing the Transformer Engine. Her work spans mixed-precision training and compute-communication overlap for both dense and MoE models across JAX and PyTorch. She is also a primary contributor to NCCL EP–a communication backend purpose-built for MoE token routing–and integrates it into Transformer Engine. Before NVIDIA, she worked in high-performance computing, optimizing kernels for sparse linear algebra and scientific computing.
Teddy Do is a performance engineer on the Transformer Engine team within the NVIDIA DLFW organization. She focuses on accelerating large-scale Transformer training and inference for core frameworks like Megatron-LM, as well as external customers’ models. She joined NVIDIA in 2022 after earning her BS/MS in Computer Engineering from Drexel University, previously working on NVIDIA data center product diagnostics where she designed maximum-stress workloads for single-GPU and rack-scale hardware validation.
Tejash Shah is a principal product manager within the AI Platform Software group at NVIDIA, responsible for managing JAX and MLX frameworks. Before NVIDIA, Tejash held software engineering roles at semiconductor companies. He holds five patents in a wide range of technological domains. He earned a master's degree in Computer Science from The University of Texas at Dallas and a bachelor’s degree in Information Technology from Gujarat University.
Discussion (0)
Sign in to join the discussion. Free account, 30 seconds — email code or GitHub.
Sign in →No comments yet. Sign in and be the first to say something.