Orbax: Distributed Checkpointing with JAX
Fuente:
arXiv
Enregistré dans:
| Auteurs principaux: | , , , , , , , , , , , , , , , |
|---|---|
| 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 |