TPUをトレーニング途中で終了させ、数秒で復旧しました:MaxTextによる弾力的なトレーニング入門
分散AIトレーニングは、単一のマシンの故障が通常、マルチノードジョブ全体をクラッシュさせ、時間のかかるフルワークロードのインフラストラクチャ再起動を余儀なくされるため、非常に不安定であることが知られています。これを解決するために、GoogleのJAXエコシステムはPathwaysを介した弾力的なトレーニングを利用しており、ハードウェア障害をキャッチ可能なPython例外に変換することで、実行中のプロセスが生き残れるようにします。予期せぬ障害が発生した場合、システムは故障したワーカーのみを自動的に置き換え、Cloud Storageから最後に有効なチェックポイントを復元し、トレーニングをその場で再開します。これにより、メインコントローラープロセスを再起動することなく、合計ダウンタイムを2分未満に最小限に抑えます。