AB
AiBoss站
快讯

NVIDIA 介绍 JAX 无丢弃 MoE 训练优化方案

NVIDIA 开发者博客发文介绍在 JAX 中结合 Transformer Engine 加速无丢弃(dropless)MoE 训练的方法,称在 DeepSeek-V3 671B 训练中观察到端到端吞吐提升,并给出分组 GEMM、专家并行通信融合等优化思路。

NVIDIA 博客介绍 JAX 无丢弃 MoE 训练优化

NVIDIA 开发者博客发布技术文章,介绍如何在 JAX 中借助 NVIDIA Transformer Engine 加速无丢弃(dropless)混合专家(MoE)模型训练。文章称,在 NVIDIA GB200 上训练 DeepSeek-V3 时,未优化的基线为 103 TFLOPS/GPU,而经过针对性内核优化后提升至 1,068 TFLOPS/GPU,即 10.4 倍提升;文章还提到在 GB300 NVL72 硬件上训练 DeepSeek-V3 671B 时,整套技术栈在 1,024 块 GPU 上保持了 97% 的扩展效率。上述数字均来自该博客自述,尚未获独立验证。

背景与影响

MoE 通过路由器为每个 token 选择 Top-K 专家,以条件计算降低训练算力需求,DeepSeek、Qwen、Mixtral 等模型均采用该架构。但生产级 MoE 训练面临 token 路由不均、专家负载倾斜、all-to-all 通信占比高等问题。文章指出,在未优化基线上,GPU 间通信占累计内核时间的 84%。

所谓无丢弃 MoE,是指每个 token 都由其被选中的专家处理,不做丢弃或填充,以保持模型质量;代价是专家收到的 token 数量可变,形成不规则张量。文章介绍的三类优化包括:用分组 GEMM 在单次内核调用中处理各专家的实际 token 数;通过 NCCL EP 融合 dispatch 与 combine 阶段,并对发往同一节点的重复 token 去重以减少网络流量;以及 JAX 主机卸载与 XLA 多流集合通信,用于缓解显存瓶颈并让 NVLink 与 InfiniBand 传输重叠。

限制与来源

需要说明的是,本文所述性能数据、硬件配置与优化效果均出自 NVIDIA 官方技术博客,属于厂商自述,未提供第三方复现结果。文章提到的后续计划(如 NVFP4、量化与 GEMM 融合、A2A 重叠)尚未确认落地时间。相关容器、配置指南与文档的可用性、支持范围及具体功能,请以 NVIDIA 官网当前信息为准。

来源:NVIDIA Developer Blog