[論文レビュー] Exploring the limits of Concurrency in ML Training on Google TPUs
この論文では、モデル並列性、通信最適化、分散評価を用いて、4,096 TPU-v3チップに深層学習の学習をスケーリングする技術を提示している。4つの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学習のベストプラクティスを確立すること。
提案手法
- BERT や Transformers のようなモデルでは、データ並列性がバッチサイズによって制限されるため、大規模なレイヤーを複数のTPUチップに分散して配置するモデル並列性を採用した。
- 4,096チップのTPUメッシュ全体にわたる最適化されたアラルーレッド通信プリミティブを実装し、通信オーバーヘッドを低減した。BERTでは、スケールアップ時に合計デバイス時間の27.3%を占めていた。
- 重み更新のシャーディングとミックス精度学習を組み合わせたSPMDパーティショニングを用いて、モデル並列学習の効率を向上させた。
- ホスト側の入力パイプラインと学習メトリクスの分散評価を最適化し、ホスト側のボトルネックを低減し、エンドツーエンドのスループットを向上させた。
- 特に小バッチまたは頻繁な更新が発生するシナリオにおいて有益な、JAXのマルチクライアント実行モデルを活用して、コンパイルおよび起動のオーバーヘッドを削減した。
- einsumへのgather/scatter操作の置き換えやハイパーパramータのチューニングなど、モデル固有の最適化を適用し、収束時間とスケーラビリティを向上させた。
実験結果
リサーチクエスチョン
- RQ1データ並列性のバッチサイズ制限を克服するために、4,096 TPU-v3チップにモデル並列性を効果的にスケーリングする方法は何か?
- RQ24,096ノード規模で高いスケーリング効率を維持するために、通信およびシステムレベルの最適化はどのようなものが必要か?
- RQ3異なるディープラーニングフレームワーク(TensorFlow 対 JAX)は、大規模な学習ワークロードにおいてどのように性能を発揮するか?それぞれの強みは何か?
- RQ4大規模モデル学習における主なパフォーマンスボトルネックは何か?それらはどのように緩和できるか?
- RQ5システム、コンパイラー、フレームワークレベルの最適化を連携することで、エンドツーエンドの学習時間はどの程度短縮できるか?
主な発見
- 4,096チップのTPU-v3 Multipodは、4つのMLPerfモデルで記録的な16〜28秒の学習時間を達成し、MLPerf-v0.7コンテストで新たなベンチマークを樹立した。
- BERTは16〜4,096チップの範囲で高いスケーラビリティを示し、スケールアップ時に通信オーバーヘッド(アラルーレッド)が合計デバイスステップ時間の27.3%を占めた。
- モデル並列性により、SSD、MaskRCNN、Transformerモデルで顕著な高速化が達成され、Transformerモデルでは4つのTPU-v3コアで2.3倍の高速化が観測された。
- 積極的なコンパイラー最適化、効率的な勾配合計、分散評価の組み合わせにより、システムレベルのボトルネックが軽減され、全体のスループットが向上した。
- JAXは、マルチクライアント実行モデルによるコンパイルおよび起動オーバーヘッドが低いため、2つのMLPerf-v0.7ベンチマークでTensorFlowを上回った。
- 本研究では、通信オーバーヘッドと非効率なパーティショニング(例:パーティショニング後の小さな空間次元)が、モデル並列学習における主なスケーラビリティボトルネックであることが確認された。
より良い研究を、今すぐ始めましょう
論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。
クレジットカード登録不要
このレビューはAIが作成し、人間の編集者が確認しました。