ocp.v1.training.pathways module#
Public API for training.pathways package.
Snapshotter#
- class orbax.checkpoint.experimental.v1.training.pathways.Snapshotter(*, replica_axis_index=0)[source][source]#
Manages asynchronous backups of JAX array states to pinned host memory.
- __init__(*, replica_axis_index=0)[source][source]#
Initializes Snapshotter.
- Parameters:
replica_axis_index (
int) – The index of the mesh axis representing data/model replicas. Typically corresponds to physical slices (connected by DCN).
- save(step, state)[source][source]#
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.
- Parameters:
step (
int) – The training step number associated with the state.state (
PyTree) – The PyTree to be saved.
- Return type:
None
- load(abstract_state, *, reset_snapshot_state=True)[source][source]#
Move arrays from workers onto TPU devices.
Uses abstract_state.sharding to properly re-partition onto the new mesh.
- Parameters:
abstract_state (
PyTree) – An abstract representation of the state, used to provide the target shardings for the restored arrays on the TPU devices.reset_snapshot_state (
bool) – If True, clears snapshot history and resets it to contain only the returned restored state (in host-pinned memory).
- Return type:
PyTree- Returns:
The restored array state.
- Raises:
RuntimeError – If no snapshots are available to restore from.
- join()[source][source]#
Blocks until all queued snapshot requests have been processed.
- Return type:
None
- property latest: CheckpointMetadata[None] | None#
Returns the training step of the most recently pinned backup.
- Return type:
Optional[CheckpointMetadata[None],None]