NVIDIA的Transformer引擎将JAX中的MoE训练加速10倍
NVIDIA发布了对其Transformer引擎的优化,这些优化显著加速了使用JAX训练Mixture-of-Experts (MoE)模型的速度。通过解决诸如token路由和GPU间通信等瓶颈,该公司宣称在DeepSeek-V3(一种拥有6710亿参数的MoE模型)的端到端训练吞吐量上实现了10倍加速。这些改进还使得高达1024个GPU的集群达到了97%的扩展效率。
近年来,MoE架构作为一种高效扩展AI模型的方法,受到了越来越多的关注。与激活所有参数处理每个token的密集模型不同,MoE选择性地将每个token路由到特定的子网络或“专家”子集。这种方法允许模型达到万亿级参数规模,而不需要线性增加计算成本,使其在大规模语言和多模态模型中尤为具有吸引力。NVIDIA的最新工作将这些原则与其高性能硬件和软件堆栈相结合。
突破MoE瓶颈
训练MoE模型面临独特的挑战,尤其是在token路由和处理“稀疏张量”方面,其中数据分布在专家之间不均。NVIDIA的Transformer引擎引入了专门的内核,例如分组GEMM(通用矩阵乘法)和专家并行操作,以解决这些低效问题。例如,分组GEMM内核通过在单次操作中处理不同数量的token,优化了不规则的张量形状,避免了填充或拆分数据的额外开销。
该引擎还利用了专为MoE复杂流量模式设计的NCCL EP通信协议。通过将token分发和合并操作融合到单一内核中,GPU可以保持高利用率,最大限度地减少因数据传输等待而导致的空闲时间。这些优化结合了MXFP8量化和JAX主机卸载等技术,进一步推动了硬件性能的极限。
无丢失MoE:以质量为先的方法
NVIDIA的框架聚焦于“无丢失” MoE,其中每个token都被处理,而不受专家之间负载不平衡的影响。虽然计算成本较高,但这种方法通过避免数据丢失提高了模型质量。为支持这一点,Transformer引擎依赖于块稀疏矩阵操作和动态形状内核,确保即使是分布不均的token也能被高效处理。
这与基于容量的MoE方法形成对比,后者将token裁剪或填充以适应固定预算,从而在硬件简化与模型保真度之间做出妥协。NVIDIA的进展使开发者能够在不牺牲性能的情况下优先考虑模型质量。
与JAX的超大规模扩展
NVIDIA的努力不仅限于单GPU性能。该公司展示了其MoE训练堆栈在拥有1024个GPU的集群上仍能保持97%的效率。达到这种扩展水平对于训练万亿级token数据集至关重要,因为节点间通信开销可能迅速成为瓶颈。诸如XLA多流集合和延迟隐藏调度器(LHS)等创新进一步减少了大规模部署中的低效问题。
为什么这很重要
随着模型规模和复杂度的增长,MoE架构成为AI行业的核心关注点。微软和艾伦人工智能研究所的最新研究强调了MoE在平衡性能和效率方面的潜力,其应用涵盖语言模型到多模态系统。NVIDIA的Transformer引擎优化使其成为这些进步的关键推动者,降低了训练下一代AI模型所需的成本和时间。
对于开发者来说,这些优化已包含在NVIDIA NGC MaxText容器中,该容器包括用于复现NVIDIA在DeepSeek-V3上成果的预配置工具。随着MoE的采用不断增长,这些工具可能会成为企业扩展AI基础设施的关键,且成本可控。
展望未来,NVIDIA计划集成更多功能,如NVFP4量化和高级内核融合,进一步提高性能。目前,其Transformer引擎代表了使MoE训练在大规模上变得切实可行的重大进展。