Efficient MoE Training for Biological Foundation Models
NVIDIA published a tutorial describing how its Transformer Engine (TE) primitives are used in the BioNeMo MoE recipe to train mixture-of-experts biological foundation models. The post frames MoE as a way to scale model capacity by activating only a subset of expert subnetworks per token, and outlines three implementation challenges the recipe addresses.
The first is fragmented expert computation. The post states that a naive implementation, such as the Hugging Face baseline, iterates over all experts in a Python loop, with each expert triggering separate kernel launches. TE's GroupedLinear instead submits multiple linear transformations in one call, gathering expert weights and input tokens and accepting per-expert token counts (split_sizes). The post notes Hugging Face Transformers also provides grouped_mm, but says TE can additionally fuse GroupedLinear with MXFP8 quantization, activation, routing-weight scaling, and intermediate data movement into a GroupedMLP kernel.
The second is model and activation memory. The BioNeMo recipe supports FP8 and MXFP8 training, representing weights and activations with 8 bits instead of BF16's 16. MXFP8 assigns a scaling factor to each block of 32 consecutive values, and the post states it is hardware-accelerated on NVIDIA Blackwell GPUs.
The third is quantization overhead. Master weights remain in 16 bits, so quantization and dequantization steps convert between formats around the low-precision GEMM. The TE Sequential API detects the GroupedLinear to ScaledSwiGLU to GroupedLinear pattern and replaces it with ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8 for the forward pass and a matching fused backward operation, avoiding materialization of some intermediates.
In a training benchmark on eight NVIDIA B200 Tensor Core GPUs, the post reports the BioNeMo recipe delivered up to 2.21x the throughput of the Hugging Face baseline. The tutorial covers running a two-GPU L0_sanity configuration first, then scaling to a Mixtral-8x7B configuration with expert parallelism (EP=8) and MXFP8 across eight GPUs, selecting BF16 or MXFP8 based on GPU and memory requirements. It lists prerequisites including Python, PyTorch, and distributed training familiarity; a CUDA-enabled environment; at least two GPUs for expert parallelism; and Blackwell GPUs for the fused MXFP8 GroupedMLP kernel.
Based on reporting from the original publisher. Visit the source for full context and later updates.
Publisher excerpt
As language models grow, scaling dense architectures becomes increasingly expensive. In a dense transformer, every token passes through every layer, so adding...