cs.DCMay 21, 2026

Orbax: Distributed Checkpointing with JAX

Authors: Colin GaffneyShutong LiDaniel NgAnastasia PetrushkinaNiket KumarAdam CogdellMridul SahuYaning Liang+8 more

Organizations: 1Google, Mountain View, USA · 3Google DeepMind, London, UK · 2Google DeepMind, Mountain View, USA

Abstract

In a landscape of high-performance distributed ML systems, JAX has emerged as a framework of choice. However, JAX's modular design philosophy leaves it without a standardized checkpointing solution. In this paper, we introduce Orbax, a modular, JAX-native checkpointing library that abstracts the complexities of distributed accelerator systems while also providing flexibility for user-friendly checkpoint manipulations throughout the ML model lifecycle. We demonstrate performance exceeding comparable PyTorch competitors by up to 3.5×\times for saving and 2×\times for loading. The library is available at https://github.com/google/orbax.

Explore similar work

CardsList