# Copyright 2026 The Orbax Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Manages asynchronous backups of JAX array states to pinned host memory."""
import logging
import queue
import threading
from typing import Any
from etils import epath
import jax
from orbax.checkpoint.experimental.v1._src.training.metadata import types as training_metadata_types
from orbax.checkpoint.experimental.v1._src.tree import types as tree_types
def is_shardable_array(x: ...) -> bool: # pyrefly: ignore[invalid-annotation]
"""Returns True if x is a concrete shardable array."""
return isinstance(x, jax.Array)
def _is_prng_key(x: Any) -> bool:
"""Returns True if x has a JAX PRNGKeyArray dtype."""
return (
hasattr(x, "dtype")
and hasattr(x, "shape")
and jax.dtypes.issubdtype(x.dtype, jax.dtypes.prng_key)
)
def _unwrap_if_prng_key(x: Any) -> Any:
"""Extracts the underlying key data buffer if x is a PRNGKeyArray."""
return jax.random.key_data(x) if _is_prng_key(x) else x
def _wrap_if_prng_key(x: Any, orig_x: Any) -> Any:
"""Wraps raw array x into a PRNGKeyArray if orig_x was a PRNGKeyArray."""
if _is_prng_key(orig_x):
return jax.random.wrap_key_data(x, dtype=orig_x.dtype)
return x
def _get_pathwaysutils():
"""Imports and returns pathwaysutils experimental modules."""
try:
# pylint: disable=g-import-not-at-top
from pathwaysutils.experimental import concatenate_by_mesh_axis # pyrefly: ignore[missing-import]
from pathwaysutils.experimental import split_by_mesh_axis # pyrefly: ignore[missing-import]
except ImportError as e:
raise ImportError(
"Snapshotter requires pathwaysutils. Please ensure pathwaysutils is"
" installed or linked via"
" //orbax/checkpoint/experimental/v1:pathways_support."
) from e
return concatenate_by_mesh_axis, split_by_mesh_axis
[docs]
class Snapshotter:
"""Manages asynchronous backups of JAX array states to pinned host memory."""
[docs]
def __init__(self, *, replica_axis_index: int = 0):
"""Initializes Snapshotter.
Args:
replica_axis_index: The index of the mesh axis representing data/model
replicas. Typically corresponds to physical slices (connected by DCN).
"""
self._concatenate_by_mesh_axis, self._split_by_mesh_axis = (
_get_pathwaysutils()
)
self._latest_snapshot: tuple[tree_types.PyTree, int] | None = None
self._lock = threading.Lock()
self._queue = queue.Queue(maxsize=1)
self.replica_axis_index = replica_axis_index
self._worker_thread = threading.Thread(target=self._worker, daemon=True)
self._worker_thread.start()
def _worker(self):
"""Processes background snapshot requests from the queue."""
while True:
pinned_state, step = self._queue.get()
try:
unwrapped_state = jax.tree.map(_unwrap_if_prng_key, pinned_state)
jax.block_until_ready(unwrapped_state)
except (jax.errors.JaxRuntimeError, RuntimeError) as e:
logging.exception("Failed to snapshot state at step %d: %s", step, e)
else:
with self._lock:
self._latest_snapshot = (pinned_state, step)
finally:
self._queue.task_done()
[docs]
def save(self, step: int, state: tree_types.PyTree) -> None:
"""Backs up JAX array states to pinned host memory, asynchronously.
If previous snapshotting requests are still in progress, this request may
be skipped.
The saved state contains all replicas of user data. When restoring via
`load_pytree`, snapshotter is able to reconstruct user data even if some
replicas are unavailable, making it resilient to failures of some replicas.
Args:
step: The training step number associated with the state.
state: The PyTree to be saved.
"""
if self._queue.full():
logging.warning("Snapshotter busy. Skipping snapshot for step %d", step)
return
def _pin_leaf(x):
if not is_shardable_array(x):
return x
data = _unwrap_if_prng_key(x)
pinned = jax.device_put(
data, data.sharding.with_memory_kind("pinned_host")
)
return _wrap_if_prng_key(pinned, x)
pinned_state = jax.tree.map(_pin_leaf, state)
self._queue.put((pinned_state, step))
[docs]
def load(
self,
abstract_state: tree_types.PyTree,
*,
reset_snapshot_state: bool = True,
) -> tree_types.PyTree:
"""Move arrays from workers onto TPU devices.
Uses `abstract_state.sharding` to properly re-partition onto the new mesh.
Args:
abstract_state: An abstract representation of the state, used to provide
the target shardings for the restored arrays on the TPU devices.
reset_snapshot_state: If True, clears snapshot history and resets it to
contain only the returned restored state (in host-pinned memory).
Returns:
The restored array state.
Raises:
RuntimeError: If no snapshots are available to restore from.
"""
with self._lock:
if self._latest_snapshot is None:
raise RuntimeError("No snapshots available to restore from.")
pinned_state, step = self._latest_snapshot
def is_replica_active(arr):
try:
data = _unwrap_if_prng_key(arr)
data[...].block_until_ready()
return True
except jax.errors.JaxRuntimeError as _:
return False
def get_active_pytree(x):
mesh_axis_name = x.sharding.mesh.axis_names[self.replica_axis_index]
data = _unwrap_if_prng_key(x)
all_replicas = self._split_by_mesh_axis.split_by_mesh_axis(
data,
mesh_axis_name,
)
active_replicas = [
replica for replica in all_replicas if is_replica_active(replica)
]
if not active_replicas:
raise RuntimeError(
"No active replicas found."
)
reconstructed_state = (
self._concatenate_by_mesh_axis.concatenate_by_mesh_axis(
active_replicas,
mesh_axis_name,
)
)
return _wrap_if_prng_key(reconstructed_state, x)
pinned_state = jax.tree.map(
lambda x: get_active_pytree(x) if is_shardable_array(x) else x,
pinned_state,
)
def _device_put_pinned(x, abs_x):
if is_shardable_array(x):
data = _unwrap_if_prng_key(x)
put_x = jax.device_put(
data, abs_x.sharding.with_memory_kind("pinned_host")
)
return _wrap_if_prng_key(put_x, x)
return x
# Re-shard on host to the target device mesh
host_target_state = jax.tree.map(
_device_put_pinned,
pinned_state,
abstract_state,
)
def _device_put_to_device(x, abs_x):
if is_shardable_array(x):
data = _unwrap_if_prng_key(x)
put_x = jax.device_put(data, abs_x.sharding.with_memory_kind(None))
return _wrap_if_prng_key(put_x, x)
return x
# Move from host back to device (TPU) memory.
restored_state = jax.tree.map(
_device_put_to_device,
host_target_state,
abstract_state,
)
unwrapped_restored = jax.tree.map(_unwrap_if_prng_key, restored_state)
jax.block_until_ready(unwrapped_restored)
if reset_snapshot_state:
with self._lock:
self._latest_snapshot = (host_target_state, step)
return restored_state
[docs]
def join(self) -> None:
"""Blocks until all queued snapshot requests have been processed."""
self._queue.join()
@property
def latest(self) -> training_metadata_types.CheckpointMetadata[None] | None:
"""Returns the training step of the most recently pinned backup."""
with self._lock:
if self._latest_snapshot is None:
return None
_, step = self._latest_snapshot
return training_metadata_types.CheckpointMetadata(
step=step,
path=epath.Path(),
metadata=None,
)