A distributed checkpoint is written by many processes at once, each holding only part of the model and optimizer state. The difficult part is not writing in parallel; it is reading back into a different layout. You trained on 256 GPUs with tensor parallelism of 8 and want to fine-tune on 64 with tensor parallelism of 4, or export a single file for inference. If the checkpoint recorded state as "rank 37's tensors", none of that works without a conversion script per layout pair.
PyTorch Distributed Checkpoint (DCP, torch.distributed.checkpoint) solves this by recording every tensor in global coordinates: its full logical shape plus the offsets and sizes of each chunk that was written. Loading becomes a geometry problem. Each new rank states which region of each global tensor it needs, and the planner intersects those regions with the saved chunks. This article explains that mechanism: the on-disk format, the planner protocol, resharding worked by hand in one and two dimensions, deduplication, partial loads and format conversion. Checkpoint sizing, asynchronous saving and commit protocols are in GPU Checkpointing Deep Dive; save intervals and exact resume are in LLM Training Checkpointing. APIs here are described as they appear in current PyTorch; check signatures against the release you run.
What is on disk
Using FileSystemWriter, a DCP checkpoint is a directory with two kinds of file. Data files are named __{rank}_{n}.distcp: rank-prefixed, numbered, holding the raw bytes of whatever that rank wrote. A single .metadata file describes everything. It holds a Metadata object whose state_dict_metadata maps each fully-qualified key (model.layers.0.attn.wq.weight, optim.state.model.layers.0.attn.wq.weight.exp_avg) to either a BytesStorageMetadata for pickled non-tensor objects or a TensorStorageMetadata with three fields: properties (dtype, layout and similar), size (the global shape) and chunks, a list of ChunkStorageMetadata entries each carrying offsets and sizes. Storage-specific data maps each chunk to the file and byte range that holds it.
You can read the index without loading any tensor data, which makes it the first tool for debugging a checkpoint:
from torch.distributed.checkpoint import FileSystemReader
from torch.distributed.checkpoint.metadata import TensorStorageMetadata
md = FileSystemReader("/ckpt/step_12000").read_metadata()
for key, entry in sorted(md.state_dict_metadata.items()):
if isinstance(entry, TensorStorageMetadata):
chunks = [(tuple(ch.offsets), tuple(ch.sizes)) for ch in entry.chunks]
print(key, tuple(entry.size), entry.properties.dtype, len(chunks), chunks[:2])
else:
print(key, "bytes")Two checks belong in every pipeline that consumes checkpoints. First, for each tensor the chunks should tile the global shape exactly: no gaps, no overlaps. Second, the set of keys should match what the loading code expects; a renamed module shows up here as a key that exists on one side only.
The planner protocol
Both save and load are a short conversation between every rank and a coordinator (rank 0 by default), driven by a planner object. For a save:
- Set up. Each rank calls
set_up_plannerwith its local state dict. Distributed tensors (DTensor, or the older ShardedTensor) are expanded into their local shards, each tagged with its global offsets. - Local plan.
create_local_planlists the write items this rank could contribute: one per tensor shard, one per non-tensor object. - Global plan. The coordinator gathers every local plan and runs
create_global_plan. This is where replicated items are deduplicated, where files are assigned to ranks, and where theMetadataindex is assembled. - Finish. Each rank receives its slice of the global plan, may adjust it in
finish_plan, and the storage writer writes its items. - Commit. Write results flow back to the coordinator, which writes
.metadata. A directory without it is an incomplete checkpoint and must be treated as one.
Loading mirrors this. Each rank's local plan lists read items: for every tensor in the state dict it was handed, the intersection of the region it holds with each saved chunk listed in the metadata. That is why loading is always in place: you construct the model in the target layout first, take its state dict, and DCP fills those tensors. The planners are pluggable. DefaultSavePlanner always deduplicates and exposes dedup_save_to_lowest_rank (its older dedup_replicated_tensors argument is deprecated and no longer has any effect); DefaultLoadPlanner exposes allow_partial_load. Custom planners are how frameworks rename keys, transform tensors on the way in, or skip state.
Worked example: resharding by intersection
Make it concrete with a tensor of global shape [12, 4], an embedding table of twelve rows. Two ranks saved it sharded by rows, so the metadata holds chunk A at offsets [0, 0] sizes [6, 4] and chunk B at offsets [6, 0] sizes [6, 4]. Three ranks now load it, also sharded by rows: rank 0 holds rows 0 to 3, rank 1 rows 4 to 7, rank 2 rows 8 to 11.
The planner's work for each pair of (requested region, saved chunk) is a box intersection, done one dimension at a time:
def intersect(req_off, req_size, ch_off, ch_size):
"""Return (offset in chunk, offset in request, lengths), or None if disjoint."""
in_chunk, in_req, lengths = [], [], []
for ro, rs, co, cs in zip(req_off, req_size, ch_off, ch_size):
lo, hi = max(ro, co), min(ro + rs, co + cs)
if lo >= hi:
return None
in_chunk.append(lo - co)
in_req.append(lo - ro)
lengths.append(hi - lo)
return in_chunk, in_req, lengths
chunks = {"A": ([0, 0], [6, 4]), "B": ([6, 0], [6, 4])}
for rank, start in enumerate([0, 4, 8]):
for name, (co, cs) in chunks.items():
hit = intersect([start, 0], [4, 4], co, cs)
if hit:
print(f"rank {rank} <- chunk {name}: {hit}")
# rank 0 <- chunk A: ([0, 0], [0, 0], [4, 4])
# rank 1 <- chunk A: ([4, 0], [0, 0], [2, 4])
# rank 1 <- chunk B: ([0, 0], [2, 0], [2, 4])
# rank 2 <- chunk B: ([2, 0], [0, 0], [4, 4])Read the second line: rank 1 takes rows 4 and 5 of the global tensor from offset 4 inside chunk A, and writes them at offset 0 of its own shard. The third line takes the first two rows of chunk B, which are global rows 6 and 7, into offset 2. Nothing in this calculation knows how many ranks saved the checkpoint, which is the whole point.
Two dimensions work identically. Suppose an [8, 8] weight was saved under 2-way FSDP sharding on dimension 0 and 2-way tensor parallelism on dimension 1: four chunks of [4, 4] at offsets [0,0], [0,4], [4,0] and [4,4]. Load it with 4-way tensor parallelism on dimension 1 and no FSDP, and each rank wants an [8, 2] column strip. Rank 0's strip at offsets [0, 0] intersects the two left-hand chunks, so it issues two reads of [4, 2] each. Every rank does the same against its own pair of chunks: eight reads in total, each a strided slice of a saved chunk.
Loading into a different layout
The state-dict helpers in torch.distributed.checkpoint.state_dict handle the awkward part: producing keys that do not depend on how the model is wrapped, and mapping optimizer state from parameter ids to parameter names so it can be resharded like the weights. A save and a load into a different parallel layout look like this:
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import (
get_state_dict, set_state_dict, StateDictOptions)
# --- save, in the training job (any layout) ---
model_sd, optim_sd = get_state_dict(model, optimizer)
dcp.save({"model": model_sd, "optim": optim_sd, "step": step},
checkpoint_id=f"/ckpt/step_{step}")
# --- load, in a job built with a different mesh ---
model = build_model(new_mesh) # parameters already sharded for the new layout
optimizer = build_optimizer(model)
model_sd, optim_sd = get_state_dict(model, optimizer)
state = {"model": model_sd, "optim": optim_sd, "step": 0}
dcp.load(state, checkpoint_id="/ckpt/step_12000") # fills tensors in place
set_state_dict(model, optimizer,
model_state_dict=state["model"], optim_state_dict=state["optim"])Three details matter. The optimizer must exist, with its state allocated, before loading, because DCP can only fill tensors that already exist; get_state_dict takes care of this for standard optimizers. Non-tensor values such as step are stored as bytes and returned whole. And StateDictOptions controls the shape of what you get: full_state_dict=True gathers unsharded tensors, with cpu_offload=True to keep them off the GPU, which is what an export path wants, while strict decides whether missing keys are fatal.
For fine-tuning, where the new model has heads the checkpoint lacks, pass a DefaultLoadPlanner(allow_partial_load=True) so keys absent from the checkpoint are left at their initial values instead of failing the load. Keep that option out of resume paths: there, a missing key is a bug you want to hear about.
To leave the DCP world, torch.distributed.checkpoint.format_utils provides dcp_to_torch_save (a DCP directory to one torch.save file) and torch_save_to_dcp for the reverse. Conversion runs in a single process and materialises full tensors, so size the machine for the whole model in host memory.
Replicated tensors and deduplication
When several ranks hold the same tensor (a replicated LayerNorm weight under tensor parallelism, or every parameter under plain data parallelism), each rank's local plan offers it. The default global plan keeps one copy, so a data-parallel checkpoint is not N times the model size. dedup_save_to_lowest_rank changes which rank does the writing; by default each duplicated item goes to whichever candidate rank has the least data assigned so far, which balances bytes written across ranks.
The same rule creates the most damaging DCP bug. A plain tensor saved under one key from several ranks is assumed to be identical on all of them. Per-rank state that genuinely differs, such as RNG states, data-loader positions or per-rank counters, will be deduplicated to a single rank's value, and every rank resumes with it. Give per-rank state per-rank keys (rng/rank_17), or store it outside DCP in small per-rank files.
Failure modes
| Symptom | Cause | Fix |
|---|---|---|
| Load fails with missing or unexpected keys | Module renamed, or a wrapper prefix such as module. or _orig_mod. leaked into keys | Build state dicts with get_state_dict; diff key sets from .metadata in CI |
| Shape mismatch on an embedding or output layer | Vocabulary padded to a multiple tied to the tensor-parallel degree, so the global shape changed | Fix the padding rule, or load with a custom planner that slices |
| Every rank resumes with identical data order | Per-rank state deduplicated | Per-rank keys |
Checkpoint directory has data but no .metadata | Job died before commit | Treat as absent; resume from the previous checkpoint |
| Load is slow on object storage | Many small reads from many ranks | Fewer, larger files; co-locate the load with storage, or stage to local NVMe |
| Silent dtype change | Loading bf16 chunks into fp32 tensors, or the reverse | Assert dtypes against the metadata before loading |
Storage throughput is the other operational limit. Data transfer paths, including reading from NVMe straight into GPU memory, are covered in the GPUDirect Storage article.
Trade-offs
| Approach | Strength | Weakness |
|---|---|---|
Gather to rank 0, one torch.save file | Simple, portable | Rank 0 memory and bandwidth bound; impractical for large models |
Per-rank torch.save files | Fast, parallel | Bound to the saving layout; resharding needs bespoke scripts |
| DCP with global-offset metadata | Parallel writes, resharding on load, dedup | More moving parts; per-rank state needs care |
| Framework formats (Megatron, Orbax) | Integrated with that framework's parallelism | Conversion needed to move between ecosystems |
Pick DCP when the same checkpoint must outlive the layout that wrote it, which for anything trained with FSDP or tensor parallelism is almost always. Keep a converted single-file export for consumers that do not run PyTorch distributed.
What to do next
- Print the
.metadataindex of your latest checkpoint and check that each tensor's chunks tile its global shape. - Write a CI test that saves a small model on 2 ranks and loads it on 3, then compares every tensor with the original.
- Audit your state dict for per-rank values stored under shared keys.
- Confirm the resume path refuses a directory without
.metadata. - Decide where
allow_partial_loadis allowed (fine-tuning) and where it is forbidden (resume). - Time a
dcp_to_torch_saveexport once, so you know what the export host needs in memory and how long it takes.