英伟达发布JAX大语言模型训练中高带宽内存瓶颈削减研究,采用主机内存卸载技术。
大语言模型训练负载日益面临GPU内存瓶颈,在计算资源充分利用之前就触及上限。模型权重、梯度、优化器状态、通信缓冲区和中间激活都争夺GPU高带宽内存(HBM)容量。随着模型规模、序列长度和批处理大小的增加,HBM容量往往成为扩展的主要制约因素。
开源JAX库中的主机卸载技术通过在前向传播时将特定激活转移到固定主机内存,在后向传播时按需流回的方式来缓解HBM压力。这是激活重新计算的替代方案——后者通过重新计算激活而非从主机内存重新加载来节省空间。
主机卸载在NVIDIA Grace Blackwell系统上特别有效。在这些系统中,NVIDIA Grace CPU与NVIDIA Blackwell GPU通过NVLink-C2C连接,提供900 GB/s的双向带宽,使固定主机内存成为理想的激活中转站。Vera CPU与Rubin GPU进一步提升性能,将双向带宽翻倍至1.8 TB/s的相干传输速率。然而,仅有高带宽的CPU-GPU连接还不够;要实现性能提升,激活传输必须与GPU计算工作充分重叠。
实验采用MaxText框架进行——这是一个JAX大语言模型训练框架,利用加速线性代数(XLA)编译器在NVIDIA GPU上进行大规模训练。所有测试在NVIDIA GB200 NVL72系统上进行,跨128个GPU、两种工作负载:Llama 3.1 405B(密集仅解码器变换器,用于研究固定批大小下的查询、键、值(QKV)激活卸载)和DeepSeek-V3 671B(稀疏混合专家模型,配备多头隐式注意力(MLA),用于研究吞吐量和内存容量效应)。
DeepSeek-V3 671B包含61个解码器层:前3层采用密集多层感知机(MLP)块,其余层采用混合专家(MoE)块。针对占主导地位的重复MoE解码器层,激活卸载策略选择性地卸载MLA查询和键/值投影中间结果、以及MoE上投影中间结果。这些激活规模足够大,直接影响是否能容纳更大批处理配置。
启用卸载、延迟隐藏调度器(LHS)和流水线传输后,DeepSeek-V3 671B达到908.2 TFLOPs/s/设备——比相同批配置下的激活重新计算快57%,比不用LHS或流水线的卸载方案快67.7%。与密集型Llama工作负载的情况不同(仅需LHS即可隐藏延迟),DeepSeek-V3的大规模MoE和MLA激活足迹使流水线传输对总吞吐量产生了明显的正面影响。这种性能优势体现了NVIDIA软硬件的紧密协同设计——在Blackwell系统中,XLA自定义调度标志与专用复制流通道协同配合,确保数据异步传输。这种集成使平台能解锁对缺乏编译器-互连深度集成的架构而言难以实现的大规模批配置。
在容量对比中,在设备上保存选定激活允许微批处理大小2、全局批处理大小256的配置,而优化的主机卸载启用了微批处理大小8、全局批处理大小1024。不使用卸载或重新计算时,该配置下设备出现内存溢出错误。主机卸载通过将选定激活存储移出GPU内存来解决此问题,为模型状态、通信缓冲区、运行时工作区和活跃计算留下更多HBM。启用LHS和流水线传输后,卸载配置使用165.2 GiB GPU内存,而不启用这些优化时为145.6 GiB。增长反映了在GPU内存中保留更多复制缓冲区和预取激活以实现传输与计算重叠的权衡,用部分内存容量换取更好的重叠效率和更高吞吐量。
Llama 3.1 405B实验在合成数据上运行10步,批大小2、序列长度8,192、全分片数据并行度128、bfloat16激活和NVFP4 4比特权重量化。带LHS的QKV激活卸载将吞吐量从2,669提升至2,746 TFLOPs/s/设备——相比无卸载基线提升2.9%。禁用LHS后QKV卸载吞吐量降至2,569 TFLOPs/s/设备,验证了主机卸载依赖与其他GPU工作的有效重叠。在此配置下,LHS单独在无流水线情况下实现最佳吞吐量2,746 TFLOPs/s/设备,相比之下有流水线时为2,718 TFLOPs/s/设备,因为LHS已将大部分传输延迟隐藏于计算和通信之后。
70.9 GiB主机内存值代表全部126层的总QKV激活存储,而非单一时刻的GPU内存节省量。在批大小2、序列长度8,192条件下,单层bfloat16 QKV激活需约576 MiB:查询512 MiB,键和值各32 MiB。启用层级扫描循环后,后向传播一次处理一层,GPU上同时仅需保留一层的QKV激活。在此工作负载中,QKV卸载主要通过用与计算和通信重叠的传输替代后向通过的QKV重新计算来优化性能,GPU峰值内存仍由模型状态、通信缓冲区和运行时工作区主导。密集的Llama 3.1 405B模型的收益小于稀疏的DeepSeek-V3 671B,但揭示了相同机制:有针对性的QKV卸载用可与计算重叠的传输替代后向通过重新计算。