google/orbax

Orbax provides common checkpointing and persistence utilities for JAX users

What it solves

Orbax 提供了一种标准化的方式来保存和恢复 JAX 模型的状态(检查点/checkpointing)并处理模型持久化。它解决了在不同框架和存储格式之间管理模型权重和优化器状态的复杂性,特别是针对大规模分布式训练。

How it works

Orbax 提供了一个可组合的 API,允许用户将 JAX pytrees(例如模型权重和优化器状态)保存(检查点)并从存储中加载(恢复)。它支持异步检查点,以最大限度地减少训练中断,并为处理自定义类型和各种存储格式提供了灵活性。

Who it’s for

它专为使用 JAX 进行大规模模型训练和评估的机器学习从业者和研究人员设计,包括那些构建基础模型或高性能序列模型的人员。

Highlights

  • Asynchronous Checkpointing: 减少开销,通过在不阻塞训练的情况下保存状态来实现。
  • Distributed Support: 在 MaxText 和 PaxML 等高性能 JAX 框架中被广泛使用。
  • Flexible Storage: 支持各种存储格式和自定义类型。
  • Broad Integration: 已集成到主要的 JAX 生态系统,包括 Flax, Gemma, 和 AXLearn。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目