[论文解读] Exploring the limits of Concurrency in ML Training on Google TPUs
本论文提出了一套技术,通过模型并行、通信优化和分布式评估,将深度学习训练扩展至4,096个TPU-v3芯片,实现了在四个MLPerf模型上16–28秒的创纪录训练时间。该方法通过解决通信、数据流水线和优化器分片中的瓶颈,在Google TPU Multipod上实现了近乎完美的可扩展性。
Recent results in language understanding using neural networks have required training hardware of unprecedentedscale, with thousands of chips cooperating on a single training run. This paper presents techniques to scaleML models on the Google TPU Multipod, a mesh with 4096 TPU-v3 chips. We discuss model parallelism toovercome scaling limitations from the fixed batch size in data parallelism, communication/collective optimizations,distributed evaluation of training metrics, and host input processing scaling optimizations. These techniques aredemonstrated in both the TensorFlow and JAX programming frameworks. We also present performance resultsfrom the recent Google submission to the MLPerf-v0.7 benchmark contest, achieving record training times from16 to 28 seconds in four MLPerf models on the Google TPU-v3 Multipod machine.
研究动机与目标
- 将深度学习模型扩展至完整的4,096芯片Google TPU-v3 Multipod,以实现最大训练吞吐量。
- 通过在大型模型(如BERT、SSD和Transformers)中采用模型并行,克服数据并行的固定小批量大小限制。
- 在大规模下优化通信、系统级协调和输入流水线性能,以最小化延迟并最大化硬件利用率。
- 在TensorFlow和JAX框架中均实现高性能训练,重点关注跨栈系统与编译器优化。
- 通过分析可扩展性瓶颈并评估框架特异性优势,建立大规模ML训练的最佳实践。
提出的方法
- 采用模型并行将大型层分布到多个TPU芯片上,尤其适用于BERT和Transformers等受小批量大小限制的数据并行模型。
- 在4,096芯片的TPU网格上实现优化的all-reduce通信原语,以减少通信开销,该开销在大规模BERT训练中占总设备时间的27.3%。
- 使用SPMD分区配合权重更新分片和混合精度训练,以提升模型并行训练的效率。
- 优化主机输入流水线和训练指标的分布式评估,以减少主机端瓶颈并提升端到端吞吐量。
- 利用JAX的多客户端执行模型,降低编译和启动开销,尤其适用于小批量或频繁更新的场景。
- 应用模型特定的优化,如用einsum替换gather/scatter操作,并调整超参数以提升收敛时间和可扩展性。
实验结果
研究问题
- RQ1如何有效将模型并行扩展至4,096个TPU-v3芯片,以克服数据并行的批量大小限制?
- RQ2在4,096节点规模下,需要哪些通信和系统级优化以保持高可扩展效率?
- RQ3在大规模训练工作负载下,不同深度学习框架(TensorFlow与JAX)的表现如何?各自的优点是什么?
- RQ4大规模模型训练中的主要性能瓶颈是什么?如何缓解?
- RQ5通过协调系统、编译器和框架级别的优化,端到端训练时间可缩短到何种程度?
主要发现
- 配备4,096个芯片的TPU-v3 Multipod在四个MLPerf模型上实现了16至28秒的创纪录训练时间,成为MLPerf-v0.7竞赛的新基准。
- BERT在16至4,096个芯片上表现出高可扩展性,大规模下通信开销(all-reduce)占总设备步时间的27.3%。
- 模型并行在SSD、MaskRCNN和Transformer模型中实现了显著加速,Transformer模型在四个TPU-v3核心上观察到2.3倍的加速。
- 结合激进的编译器优化、高效的梯度求和与分布式评估,减少了系统级瓶颈并提升了整体吞吐量。
- 由于多客户端执行模型中启动和编译开销更低,JAX在两个MLPerf-v0.7基准测试中优于TensorFlow。
- 研究证实,通信开销和低效的分区(如分区后空间维度过小)是模型并行训练中的主要可扩展性瓶颈。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。