google/orbax
Orbax provides common checkpointing and persistence utilities for JAX users
何を解決するか
Orbaxは、JAXモデルの状態(チェックポイント化)を保存および復元する標準化された方法を提供し、モデルの永続化を扱います。大規模な分散学習において、異なるフレームワークやストレージ形式間でモデルの重みやオプティマイザの状態を管理する複雑さに対処します。
仕組み
Orbaxは、ユーザーがJAXのpytree(モデルの重みやオプティマイザの状態など)をストレージに保存(チェックポイント化)および読み込み(復元)できる、合成可能なAPIを提供します。非同期のチェックポイント化をサポートしており、トレーニングの中断を最小限に抑え、カスタム型やさまざまなストレージ形式の扱いに柔軟性を提供します。
対象ユーザー
JAXを用いて大規模なモデルのトレーニングや評価を行う機械学習の実践者や研究者、特に基礎モデルや高性能なシーケンスモデルを構築している人向けに設計されています。
特徴
- 非同期チェックポイント化: トレーニングをブロッキングせずに状態を保存することで、オーバーヘッドを低減します。
- 分散処理対応: MaxTextやPaxMLなど、高性能なJAXフレームワークで広く使用されています。
- 柔軟なストレージ: 複数のストレージ形式やカスタム型をサポートします。
- 広範な統合: Flax、Gemma、AXLearnなど、主要なJAXエコシステムに統合されています。
関連
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト