A distributed checkpoint is complete only when every required shard and the shared training position belong to the same generation and are durably visible.
Distributed checkpoint manifests and coherent resume
Name one update boundary
Stop at an optimizer-step boundary and assign a generation ID. Record model parameters and buffers, optimizer moments, scheduler and scaler state, completed update number, sampler epoch and next data position. If a save begins between accumulated microbatches, the checkpoint also needs partial gradients and accepted-example counts, or it cannot resume the same update. Most systems choose a clean step boundary instead. Single-process resume state provides the base inventory.
Stage shards before advertisement
Each worker writes a new shard under a generation-specific staging name, syncs it and reports a checksum and size. A coordinator writes a manifest only after every expected shard is present and verified, then atomically advertises that generation as the latest complete one. A crash during shard write should leave the previous complete generation discoverable. A filename that exists is not sufficient evidence; it can contain a partial write or belong to a different run.
Tie the manifest to data order
Store the dataset manifest hash, number of ranks, sampler seed and epoch with the checkpoint. If a resumed job changes rank count, sharding and sampler order may change; the model weights can load while the next examples differ. Some checkpoint formats can reshard model and optimizer state, but that does not reconstruct an identical data stream automatically. Define whether the goal is exact update replay or a statistically valid continuation and test the appropriate claim.
Coordinate every rank
The coordinator should not publish a pointer while another rank is still writing. A barrier or explicit acknowledgments must cover data durability as well as process arrival. On restart, all workers read the same manifest and reject missing or mismatched shards before any rank enters training collectives. If one rank loads generation 83 while another loads 84, synchronized gradients cannot repair their inconsistent optimizer moments. Sampler accounting must agree too.
Test failure injection
Kill one worker after it creates half a shard, kill the coordinator before pointer publication, then kill it immediately after publication. In the first two cases, readers should select the previous complete generation; in the third, they should select the new one only when all checksums pass. Compare the first resumed fixed-batch update with an uninterrupted reference at the same rank count. The code below illustrates the coordinator-side commit; production workers must supply durable shard acknowledgments.
Implementation
import hashlib
import json
import os
from pathlib import Path
from tempfile import TemporaryDirectory
def publish_generation(checkpoint_root: Path, generation: int,
shard_payloads: dict[int, bytes], expected_ranks: int) -> None:
if set(shard_payloads) != set(range(expected_ranks)):
raise ValueError("incomplete worker generation")
generation_dir = checkpoint_root / f"generation-{generation}"
generation_dir.mkdir()
shard_manifest = []
for worker_rank, payload in sorted(shard_payloads.items()):
shard_name = f"rank-{worker_rank}.bin"
shard_path = generation_dir / shard_name
with shard_path.open("wb") as shard_file:
shard_file.write(payload)
shard_file.flush()
os.fsync(shard_file.fileno())
shard_manifest.append({"rank": worker_rank, "name": shard_name,
"sha256": hashlib.sha256(payload).hexdigest()})
manifest = {"generation": generation, "rank_count": expected_ranks,
"shards": shard_manifest}
staged_pointer = checkpoint_root / "latest.pending"
with staged_pointer.open("w", encoding="utf-8") as pointer_file:
json.dump(manifest, pointer_file)
pointer_file.flush()
os.fsync(pointer_file.fileno())
os.replace(staged_pointer, checkpoint_root / "latest.json")
with TemporaryDirectory() as scratch_directory:
checkpoint_root = Path(scratch_directory)
publish_generation(checkpoint_root, 83, {0: b"rank-zero", 1: b"rank-one"}, 2)
restored_manifest = json.loads((checkpoint_root / "latest.json").read_text())
assert len(restored_manifest["shards"]) == 2Performance and operating cost
A checkpoint stores parameters, buffers and optimizer state; adaptive optimizers can require several times the parameter bytes. Sharded writing distributes I/O but adds coordination and checksum work proportional to shard size. The example uses fsync on files and an atomic pointer replacement; a real durable filesystem may also require directory syncing and an object-store-specific commit protocol. Profile save stalls and retain at least one previous complete generation.
Common Mistakes
- Do not advertise a generation before every shard is durable and checked.
- Do not call weights-only loading an exact optimizer resume.
- Do not assume changing rank count preserves the same sampler order.
