google/paxml

Pax is a Jax-based machine learning framework for training large scale models. Pax allows for advanced and fully configurable experimentation and parallelization, and has demonstrated industry leading model flop utilization rates.

何を解決するか

Paxml(またはPax)は、Jaxの上に構築された大規模機械学習実験の設定および実行を可能にするフレームワークです。特にCloud TPU VMやTPU Podスライスなどの分散ハードウェア上で動作することを想定して設計されており、複雑なモデル設定の管理と展開を簡素化します。

動作方法

PaxはJaxによる計算処理を活用し、設定ファイルを通じて実験を構造的に定義します。SPMD(Single Program, Multiple Data)をpjitpmapで実装するほか、パイプライン並列処理もサポートしています。マルチホストでのデータインフィードを処理し、トレーニング中に重複バッチが発生しないように、データを異なるJaxプロセスにシャーディングします。また、SeqIO、Lingvo、カスタム実装など、さまざまなデータ入力タイプと統合され、モデルへのトレーニング、評価、デコード用のデータ供給が可能になります。

対象ユーザー

数十億から数兆パラメータに及ぶ非常に大きなモデルをトレーニングする機械学習研究者やエンジニア向けに設計されており、TPUまたはGPUインフラ上でこれらの実験を管理するための堅牢なシステムを必要とする方々に適しています。

特徴

  • TPU最適化: Cloud TPU VMおよびPodスライスと深く統合され、高性能なトレーニングを実現します。
  • スケーラビリティ: 大規模言語モデル向けの弱スケーリングをサポートし、使用するチップ数に比例してモデルサイズを拡大可能。
  • 柔軟な並列処理: pjit(SPMD)およびpmapによる分散実行をサポート。
  • データ管理: マルチホストインフィードとSeqIOを介した評価データの自動シャーディング/パディングを内蔵。
  • パフォーマンスメトリクス: トータルトレーニング効率を測定するためのModel FLOPs Utilization(MFU)を活用。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト