专栏名称: AI思想会
连接人工智能技术人才和产业人才的交流平台
今天看啥Skill
TodayRss-海外RSS稳定源
目录
今天看啥  ›  专栏  ›  AI思想会

JAX 训练加速指南:8 个让 TPU 满跑的工程实战习惯

AI思想会  · 公众号  · AI  · 2025-12-04 18:59
    

主要观点总结

本文总结了使用JAX在TPU上进行深度学习训练时需要关注的八个关键点,包括Shape的稳定性、算子的融合度、数据管道的效率等。通过遵循这些指导原则,可以优化训练过程,提高性能。

关键观点总结

关键观点1: 尽早锁定Shape和全局Batch大小

在TPU训练中,Shape的稳定性非常重要。动态Shape会触发重新编译,导致性能下降。因此,应尽早锁定Shape和全局Batch大小,避免在训练过程中发生变化。

关键观点2: 使用bfloat16进行优化

在TPU上,bfloat16是一个好选择,可以兼顾速度、内存和数值稳定性。激活值和梯度可以存储为bfloat16,而优化器状态中的权重则需要保留一份FP32的“主副本”。

关键观点3: 明确切分,避免猜测

在JAX中使用pjit进行GSPMD时,应通过PartitionSpecs明确告诉它想要的切分方式。数据并行时,使用'data'模式;如果模型太大需要张量并行,则添加'model'轴。

关键观点4: 使用jit, vmap, scan三件套

JAX喜欢大块头的Kernel,讨厌成千上万个细碎的小算子。训练Step和任何中大型计算逻辑都必须用jit包起来。遇到Python循环,如果是时间步逻辑就换成lax.scan,如果是批次并行就用vmap。

关键观点5: 避免输入管道拖后腿

Host到Device的数据传输一旦停顿,吞吐量就会下降。应使用高效的数据预取策略,如tf.data或NumPy loader配合prefetch,并做双重缓冲。

关键观点6: PRNG要Fold进Step和DeviceID

JAX的PRNG是无状态的,需要Fold进Step和DeviceID来保证独立性。数据增强/Dropout的Key和参数初始化的Key要分开管理。

关键观点7: 使用Remat和梯度累积优化内存使用

对于深层网络,可以直接使用Activation Checkpointing(jax.checkpoint或nn.remat)来节省显存。对于大Batch但显存不足的情况,可以使用梯度累积(Gradient Accumulation)来切成小的micro-step。

关键观点8: 一定要跑Profiler

使用Profiler来识别性能瓶颈。重点关注Host Waits、Recompiles和未融合好的细碎算子。稳态运行时,关注Tokens/sec或Images/sec以及硬件利用率。


免责声明:本文内容摘要由平台算法生成,仅为信息导航参考,不代表原文立场或观点。 原文内容版权归原作者所有,如您为原作者并希望删除该摘要或链接,请通过 【版权申诉通道】联系我们处理。

原文地址: 访问原文地址
总结与预览地址:访问文章预览/总结
文章地址: 访问文章快照