Orbax: Distributed Checkpointing with JAX

Fuente: arXiv
Enregistré dans:
Détails bibliographiques
Auteurs principaux: Gaffney, Colin, Li, Shutong, Ng, Daniel, Petrushkina, Anastasia, Kumar, Niket, Cogdell, Adam, Sahu, Mridul, Liang, Yaning, Bansal, Nikhil, Pan, Justin, Mau, Angel, Agrawal, Abhishek, Berlot, Marco, Sang, Ruoxin, Sodhia, Kiranbir, Iyer, Rakesh
Format: Preprint
Publié: 2026
Sujets:
Accès en ligne:
Tags: Ajouter un tag
Pas de tags, Soyez le premier à ajouter un tag!
_version_ 1866917537702215680
author Gaffney, Colin
Li, Shutong
Ng, Daniel
Petrushkina, Anastasia
Kumar, Niket
Cogdell, Adam
Sahu, Mridul
Liang, Yaning
Bansal, Nikhil
Pan, Justin
Mau, Angel
Agrawal, Abhishek
Berlot, Marco
Sang, Ruoxin
Sodhia, Kiranbir
Iyer, Rakesh
author_facet Gaffney, Colin
Li, Shutong
Ng, Daniel
Petrushkina, Anastasia
Kumar, Niket
Cogdell, Adam
Sahu, Mridul
Liang, Yaning
Bansal, Nikhil
Pan, Justin
Mau, Angel
Agrawal, Abhishek
Berlot, Marco
Sang, Ruoxin
Sodhia, Kiranbir
Iyer, Rakesh
contents 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.
format Preprint
id arxiv_https___arxiv_org_abs_2605_23066
institution arXiv
publishDate 2026
record_format arxiv
spellingShingle Orbax: Distributed Checkpointing with JAX
Gaffney, Colin
Li, Shutong
Ng, Daniel
Petrushkina, Anastasia
Kumar, Niket
Cogdell, Adam
Sahu, Mridul
Liang, Yaning
Bansal, Nikhil
Pan, Justin
Mau, Angel
Agrawal, Abhishek
Berlot, Marco
Sang, Ruoxin
Sodhia, Kiranbir
Iyer, Rakesh
Distributed, Parallel, and Cluster Computing
Machine Learning
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.
title Orbax: Distributed Checkpointing with JAX
topic Distributed, Parallel, and Cluster Computing
Machine Learning
url https://arxiv.org/abs/2605.23066