google/orbax

Orbax provides common checkpointing and persistence utilities for JAX users

解決的問題

Orbax 提供了一種標準化的方式來儲存和還原 JAX 模型的狀態(檢查點),並處理模型持久化。它解決了在不同框架和儲存格式之間管理模型權重和優化器狀態的複雜性,特別是在大規模分散式訓練中。

工作原理

Orbax 提供了一個可組合的 API,允許使用者將 JAX pytrees(例如模型權重和優化器狀態)儲存(檢查點)和載入(還原)到儲存裝置中。它支援非同步檢查點以最小化訓練中斷,並在處理自訂類型和各種儲存格式方面提供彈性。

適用對象

專為使用 JAX 進行大規模模型訓練與評估的機器學習實務人員和研究人員設計,包括那些開發基礎模型或高效率序列模型的人。

主要特色

  • 非同步檢查點:透過在不阻斷訓練的情況下儲存狀態,減少開銷。
  • 分散式支援:廣泛應用於 MaxText 和 PaxML 等高性能 JAX 框架中。
  • 彈性儲存:支援多種儲存格式和自訂類型。
  • 廣泛整合:已整合至主要的 JAX 生態系統,包括 Flax、Gemma 和 AXLearn。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案