cs.DC · 2605.23066 Copy arXiv ID · May 21, 2026 Save Orbax: Distributed Checkpointing with JAX Authors: Colin Gaffney , Shutong Li , Daniel Ng , Anastasia Petrushkina , Niket Kumar , Adam Cogdell , Mridul Sahu , Yaning 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 May 18, 2026 · Shujie Han, Feng Jiang, Patrick P. C. Lee +5 Fault Tolerance
Jun 14, 2026 · Qiyue Liang, Steven Ingram, George Vanica +4 Pytorch Single-Agent Baselines
Nov 16, 2023 · Alexander Rutherford, Benjamin Ellis, Matteo Gallici +18 Multi-Agent Reinforcement Learning Deep Reinforcement Learning
May 18, 2026 · cs.DC J/K move · Enter open · S save
Shujie Han, Feng Jiang, Patrick P. C. Lee, Xiao Zhang +4
Northwestern Polytechnical University · The Chinese University of Hong Kong · National University of Defense Technology
Large Language Model (LLM) training is frequently interrupted by a heterogeneous spectrum of failures, from common GPU crashes to catastrophic cluster-wide outages. Existing checkpointing systems rely on monolithic, single-tier storage backend, forcing a trade-off between state-saving overhead and recovery speed. We propose TierCheck, a cluster-aware tiered checkpointing system that aligns storage placement with failure heterogeneity. TierCheck adopts a three-tier design that maintains lightweight differential checkpoints in local and peer memory for fast localized recovery, while asynchronously migrating heavyweight base checkpoints to remote persistent storage. It also ensures strict global consistency across tiers without stalling training, and achieves fast cluster-aware checkpoint restoration during recovery. Evaluations on models up to 40 billion parameters show that TierCheck achieves low training overhead, reduces end-to-end checkpointing time to under 10s, and supports high-frequency checkpointing, ultimately striking an optimal balance between low-overhead persistence and fast recovery.