NVIDIA通过主机卸载优化JAX LLM训练

realtime news Jul 10, 2026 19:02

NVIDIA的主机卸载技术为JAX LLM训练提升了GPU内存效率,使得批处理规模更大,吞吐量更快。

NVIDIA通过主机卸载优化JAX LLM训练

NVIDIA为基于JAX的大型语言模型(LLM)训练引入了一种新的主机卸载技术,解决了限制现代AI工作负载可扩展性的GPU高带宽内存(HBM)瓶颈问题。借助最新的NVIDIA Blackwell架构,该方法通过在前向传播过程中将选定激活移动到CPU内存,并在后向传播过程中将其流式返回,从而实现更大的批处理规模和更快的训练吞吐量。

随着模型规模、序列长度和批处理规模的增长,HBM经常成为LLM训练的限制因素。NVIDIA于2026年7月10日在公司博客中详细介绍的主机卸载解决方案,为激活重计算提供了替代方法。激活重计算是一种常见但计算成本高昂的内存约束处理方法。相较于重新计算激活,该方法将激活暂时存储在CPU内存中,并在需要时检索。

为何NVIDIA的Blackwell架构脱颖而出

Blackwell GPU与NVIDIA的Grace CPU配对,通过NVLink-C2C实现高达900 GB/s的双向带宽。这种高速连接使主机卸载变得可行,可以在GPU和CPU内存之间实现快速数据传输。在NVIDIA即将推出的Vera和Rubin平台上,这一带宽翻倍至1.8 TB/s,进一步提高了对内存密集型工作负载进行卸载的可行性。

除了硬件外,NVIDIA将JAX加速线性代数(XLA)编译器集成到其系统中,支持流水线数据传输与GPU计算重叠,从而最大化吞吐量。软硬件的紧密结合确保数据移动不会阻塞训练流水线,这是普通集群中常见的问题。

大型模型的性能提升

使用基于JAX的MaxText框架进行的测试展示了主机卸载对两种高要求LLM工作负载的影响:密集型Llama 3.1(4050亿参数)和稀疏型DeepSeek-V3(6710亿参数)。对于DeepSeek-V3,采用流水线传输的主机卸载实现了908.2 TFLOPs/s/设备的性能——比激活重计算提升了57%,比非流水线卸载提升了67.7%。这些优化还支持更大的批处理配置,将GPU内存利用率提升至165.2 GiB,同时保持高吞吐量。

即使在内存需求较低的场景中,例如Llama 3.1,卸载同样证明了其价值。通过LHS启用的QKV卸载,吞吐量提升了2.9%,表明即使是较小的改进也能在大规模训练中积少成多。

将JAX定位为可扩展AI的关键

JAX是由Google和NVIDIA支持的开源机器学习库,已成为扩展LLM的关键框架。其生态系统包括用于分布式训练的工具,如用于优化的Optax和用于检查点的Orbax。包括主机卸载在内的最新创新进一步巩固了JAX在处理大规模工作负载及优化内存效率方面的声誉。

行业对内存优化的关注并非新鲜事。Google最近在2026年4月10日详细介绍了针对TPU训练的类似卸载技术,反映了利用CPU资源克服GPU内存限制的更广泛趋势。然而,NVIDIA的方法针对其专有的互连和硬件进行了定制,为运行在其系统上的JAX用户提供了无与伦比的整合能力。

对AI开发者的意义

主机卸载对GPU内存成为限制因素的工作负载最为有利,例如具有高参数数量、长上下文长度或大批处理规模的模型训练。开发者可以通过更新其JAX环境并启用特定的XLA标志(包括延迟隐藏调度器和流水线卸载)来实现这一功能。

随着AI模型的持续增长,像主机卸载这样的内存优化技术将对保持效率和成本效益至关重要。NVIDIA对软硬件紧密集成的重视提供了竞争优势,特别是在公司准备推出具有更高互连性能的Rubin平台之际。

对于希望在NVIDIA GPU上尝试JAX的开发者,NVIDIA提供了一系列工具,包括NVIDIA JAX-Toolbox和LLM训练的预构建容器。随着GPU硬件的不断发展,这些进步可能会塑造可扩展AI开发的未来。

Image source: Shutterstock