latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
tensorplay.distributed.checkpoint
tensorplay.distributed.checkpoint saves and restores the state of a
distributed training job — one or more state_dict structures split across
the ranks of a process group — and makes the checkpoint look like a single
artifact on the storage medium. Unlike a naive per-rank save, the package
coordinates the write across ranks: the individual pieces are gathered into a
global plan, every rank writes its share of the data, and a metadata file is
committed once all ranks have finished, so a partially written checkpoint can
be detected by the missing or incomplete metadata.
A checkpoint on the file system is a directory with one or more per-rank data
files plus a .metadata file describing the layout — which tensors live where
and how the pieces map back to the original state dictionary. Loading reads
that metadata and restores each value into the (already allocated) state
dictionary it was loaded from, so the shape and dtype of the objects you load
into come from the caller, not from the checkpoint.
import tensorplay as tp
import tensorplay.distributed as dist
from tensorplay.distributed.checkpoint import save, load
dist.init_process_group("nccl")
model = tp.nn.Linear(512, 512).cuda(dist.get_rank())
model = tp.nn.parallel.DistributedDataParallel(model, device_ids=[dist.get_rank()])
state_dict = {"model": model.state_dict()}
save(state_dict, checkpoint_id="file:///checkpoints/run-1")
# later, on every rank:
model2 = tp.nn.Linear(512, 512).cuda(dist.get_rank())
load({"model": model2.state_dict()}, checkpoint_id="file:///checkpoints/run-1")
Saving and loading
Save a state dictionary with coordinated metadata commit. |
|
Stage the input and execute the checkpoint write asynchronously. |
|
Load checkpoint values into an existing state dictionary. |
|
save()snapshots the input state dictionary (detaching and cloning tensors so the checkpoint never aliases live training state), flattens it into a write plan, coordinates the plan across ranks, writes the data through a storage writer, and commits the metadata once every rank has written its share. Passcheckpoint_id(a storage URI such asfile:///path) or an explicitstorage_writer.async_save()returns control to the caller immediately: the input is staged (copied to pin/shared memory where configured) and the write runs on a background thread or process, chosen byAsyncCheckpointerType. It returns a future (or anAsyncSaveResponsewhen a stager is used) that you can wait on at a safe point, typically after the next training step.save_state_dict()/load_state_dict()are the lower-level entry points that take an explicit storage writer/reader and acoordinator_rank, used when embedding the checkpoint logic inside a custom pipeline.load()reads the metadata, builds a load plan, and restores values in place into the existing state dictionary. On failure it rolls back the state dictionary to the snapshot taken before the load, so a failed load never leaves the model half-modified.
State dictionary helpers
For sharded training (FSDP, DTensor, or a plain distributed model), the state
dict you would normally assemble yourself differs across ranks. The
get_model_state_dict() /
get_optimizer_state_dict() family
walks the model and optimizer for you, producing the per-rank view that
checkpoint.save expects, and the matching set_* functions apply a saved
state dict back onto the live model and optimizer. StateDictOptions
controls the details: full_state_dict collapses shards into an
all-parameters view, cpu_offload moves CPU tensors to pinned memory before
saving, ignore_frozen_params skips parameters that do not require
gradients, and broadcast_from_rank0 lets rank 0 distribute the full state
dict instead of each rank loading its own shard. The optimizer state
dictionaries produced and consumed by the *_optimizer_state_dict functions
are OptimizerStateType, a plain dict[str, Any] keyed by parameter-qualified
names.
load_sharded_optimizer_state_dict() is
a companion helper for loading a sharded optimizer state dict (the format
FSDP produces for its optimizer) against the planner, so a checkpoint saved by
FSDP can be resumed by a job with a different world size.
Storages
Read checkpoint transactions written by |
|
|
|
Writes MEGA-format shards ( |
|
Reads MEGA checkpoints produced by |
|
The storage layer decides where the bytes land. FileSystemWriter
writes one file per rank into the checkpoint directory by default, optionally
with metadata after every rank finishes; FileSystemReader
reads such a directory back. SerializationFormat
chooses the on-disk tensor encoding (torch_save, the default, or
safetensors). The HuggingFace storages write checkpoints in the layout used
by Hugging Face model repositories (config + sharded weight files), letting
you save a model that can be loaded by the Hugging Face ecosystem directly.
StorageWriter /
StorageReader are the abstract
interfaces a custom backend implements.
Planners
The planner turns a state dictionary into a list of concrete data items and
back. On save it walks the state dict, assigns each tensor a WriteItem that
records where the tensor lives (TensorWriteData) and how its shards map into
the checkpoint (ChunkStorageMetadata), then assembles a global SavePlan
from the per-rank local plans. On load it turns the metadata back into
ReadItems and a LoadPlan. The default planners handle tensors, sharded
tensors, and stateful objects out of the box; implementing the abstract
SavePlanner /
LoadPlanner interfaces lets you
customize which values are stored and how.
Metadata and supporting types
Metadata is the description of a
whole checkpoint: for each key in the state dictionary the
TensorStorageMetadata (or
BytesStorageMetadata for
non-tensor bytes) says what is stored, at which MetadataIndex, and how the
stored chunks reassemble into the full tensor. StorageMeta
carries the storage-level properties, and TensorProperties
the dtype/layout/device of the tensor being stored.
The Stateful protocol marks objects
that know how to serialize themselves: a value with state_dict() and
load_state_dict() methods is saved through them instead of being pickled
directly. CheckpointableTensor is
the interface a tensor-like type implements to participate in the checkpoint
(being shardable and reloadable by shape/dtype). Errors raised while saving or
loading surface as CheckpointException.
Where to go next
FSDP — the sharding strategy whose state dicts
get_model_state_dict/get_optimizer_state_dictare designed to collect.distributed tensors — the sharded tensor type that checkpoint stores and restores transparently.
the distributed package — process groups, init, and collectives that the checkpoint machinery uses to coordinate ranks.
Help improve this page
Found an error, an unclear step, or a missing example?
tensorplay.distributed.autograd
tensorplay.distributed.autograd extends automatic differentiation across worker boundaries. When a forward pass spans several workers — because rpc_sync() / rpc_async calls ran parts of the computation remotely, or a rem
tensorplay.distributed.device_mesh
A device mesh is the execution context for distributed tensors . It is an n-dimensional array whose entries are global ranks: the value at coordinates (i, j, ...) is the rank of the process holding that position of the m

