JAX中加速Dropless MoE训练:NVIDIA Transformer Engine解析

本文介绍如何在 JAX 框架下结合 NVIDIA Transformer Engine 实现高效的 Dropless MoE 大模型训练。
混合专家模型(MoE)通过稀疏激活机制,在不成比例增加计算量的前提下大幅扩展模型参数规模,已成为 DeepSeek、Qwen、Mixtral 等前沿模型的核心架构。然而传统 MoE 训练存在 token 丢弃问题——路由器负载不均衡时,超出专家容量的 token 会被丢弃,导致信息损失和训练不稳定。Dropless MoE 旨在消除这一丢弃行为,但对底层算子提出了更高要求。本文梳理了在 JAX 生态中借助 NVIDIA Transformer Engine 加速 Dropless MoE 训练的技术路径:Transformer Engine 通过分组 GEMM(Grouped GEMM)支持不规则形状的专家计算、通过 FP8 低精度训练降低显存和带宽压力,并与 JAX 的 `shard_map` 等分布式并行原语协同,适配专家并行(Expert Parallelism)策略。三者结合,在保障模型训练质量的同时实现接近硬件极限的计算效率。
混合专家模型(Mixture of Experts, MoE)已经成为大规模AI模型训练中最具标志性的架构趋势之一。DeepSeek、Qwen、Mixtral 等一系列前沿模型都采用了 MoE 结构,其核心思路是通过稀疏激活的专家网络,在不成比例增加计算成本的前提下大幅扩展模型参数规模。而如何高效地训练这类模型,尤其是在 JAX 生态中实现所谓的"Dropless MoE",正成为工程实践中的关键议题。
本文基于 NVIDIA 开发者博客的分享,梳理在 JAX 框架下借助 NVIDIA Transformer Engine 加速 Dropless MoE 训练的技术路径与工程价值。
MoE 为什么成为大模型的主流架构
MoE 的基本理念是用多个"专家"子网络替代传统 Transformer 中稠密的前馈层(FFN),并通过一个路由器(router / gating network)为每个 token 动态选择少量专家进行计算。这样一来,模型的总参数量可以做得很大,但每个 token 实际参与计算的参数只占很小一部分,从而在推理和训练时保持相对可控的算力开销。

DeepSeek、Qwen 和 Mixtral 都是这一趋势的代表。它们通过 MoE 结构在保持每 token 计算量不变的情况下,把模型容量推向了数千亿甚至更高的参数级别。可以说,MoE 已经从一个学术性的探索方向,转变为工业级大模型训练的默认选项之一。
Dropless MoE:解决 token 丢弃问题
传统 MoE 训练中一个棘手的问题是负载不均衡。路由器把 token 分配给专家时,往往会出现某些专家"过载"、某些专家"闲置"的情况。为了保证计算图的形状固定、便于硬件并行,工程实现通常会为每个专家设定一个容量上限(capacity factor)。一旦某个专家收到的 token 超过容量,多余的 token 就会被丢弃(dropped),无法得到该专家的处理。
这种 token 丢弃会带来两个负面影响:一是信息损失,被丢弃的 token 得不到应有的专家计算,影响模型质量;二是训练效果的不确定性,因为丢弃行为依赖于批次内的数据分布。
Dropless MoE 的目标正是消除这种丢弃行为——让所有 token 都能路由到对应的专家并完成计算。实现这一点的关键在于支持可变长度、不规则形状的计算(ragged / grouped computation),而这恰恰对底层算子和硬件调度提出了更高的要求。
Transformer Engine 在 JAX 中的角色
NVIDIA Transformer Engine 是专为加速 Transformer 类模型训练与推理而设计的库,其核心能力包括对 FP8 等低精度计算的支持,以及针对 NVIDIA GPU 高度优化的算子实现。在 JAX 生态中,Transformer Engine 提供了与 JAX/Flax 集成的接口,使开发者能够在保持 JAX 函数式编程与自动微分优势的同时,享受到硬件级别的性能加速。
对于 Dropless MoE 训练而言,Transformer Engine 的价值主要体现在几个方面:
高效的分组矩阵运算
Dropless MoE 需要对每个专家处理数量不等的 token,本质上是一系列不同规模的矩阵乘法(grouped GEMM)。Transformer Engine 针对这类不规则计算做了专门优化,避免了为对齐形状而填充(padding)带来的算力浪费。
低精度训练支持
借助 FP8 混合精度,MoE 中大量的专家前馈计算可以在显著降低显存占用和带宽压力的同时,保持训练稳定性。这对于动辄上百个专家的大规模 MoE 模型尤为重要。
与 JAX 并行策略的协同
JAX 天然支持 pmap、shard_map 等分布式并行原语,Transformer Engine 的算子能够与专家并行(expert parallelism)等切分策略配合,将不同专家分布到不同设备上,充分利用多 GPU 集群的算力。
工程意义与实践启示
把 Dropless MoE、JAX 和 Transformer Engine 三者结合起来,本质上是在追求"模型质量"与"训练效率"之间的更优平衡。消除 token 丢弃提升了模型的训练质量与可复现性,而 Transformer Engine 的硬件优化则确保这种"不丢弃"的代价不会转化为无法承受的计算开销。
对于正在或计划训练大规模 MoE 模型的团队来说,这一技术组合提供了一条相对成熟的工程路线:在 JAX 这样兼具灵活性与可扩展性的框架上,通过 Transformer Engine 获取接近硬件极限的性能,同时借助 Dropless 策略避免传统 MoE 的质量损失。
随着 DeepSeek、Qwen、Mixtral 等 MoE 模型持续验证这一架构的有效性,围绕 MoE 训练效率的底层工具链竞争也会愈发激烈。NVIDIA 通过 Transformer Engine 在 JAX 生态的布局,正是这一趋势的直接体现。
小结
MoE 架构解决了大模型"参数规模"与"计算成本"之间的矛盾,而 Dropless MoE 进一步解决了训练过程中的 token 丢弃问题。在 JAX 框架下,NVIDIA Transformer Engine 通过分组 GEMM、FP8 低精度支持以及与分布式并行策略的协同,为高效训练这类模型提供了坚实的底层支撑。对于关注前沿大模型训练技术的工程师而言,这是一条值得深入研究的实践路径。
相关推荐

iOS 27、iPadOS 27与macOS 27:一次讨论背后的信息缺口
Hacker News上关于iOS 27、iPadOS 27与macOS 27的讨论引发关注。本文梳理该话题背景、苹果版本命名策略的可能转向,并说明在信息有限时如何理性看待此类系统更新传闻。

ComfyUI Prompt Studio:从参考图到可用提示词的工作流
ComfyUI Prompt Studio 是一款开源工作流工具,能从参考图和简单创意自动生成可投产的图像提示词、多模型定制提示和 MiniMax 视频脚本,支持忠实重建与自由创作两种模式。

RSI短期不会发生?新论文用NeurIPS实验给出否定答案
一篇新论文让AI智能体重做未发表的NeurIPS论文并由原作者评分,结果显示当前智能体无法胜任开放式ML研究,据此论证递归自我改进(RSI)短期内不会发生。本文解析其实验设计、核心论证与局限。