A training checkpoint is the job's only memory across failures. On a large cluster failures are routine: a GPU throws an uncorrectable memory error, a link flaps, a host reboots, a spot node is reclaimed. Each one kills the job, and everything since the last good checkpoint is recomputed. So checkpointing is a trade between two costs you can measure: time spent saving, and work lost when something fails.
Other pages on this site cover the neighbouring questions: what a resumable checkpoint must contain and how to continue training from one, and how failure rates scale with fleet size in fleet operations. This page is about the save path itself: how many bytes a checkpoint is, where those bytes travel, which part of the trip stalls the GPUs, how asynchronous and sharded saving remove most of the stall, how to commit a checkpoint so a crash mid-save never leaves a bad one marked as latest, and how to verify a restore. The code uses PyTorch's torch.distributed.checkpoint (DCP); signatures were checked against the PyTorch documentation in October 2026, and this module has changed between releases, so check the version you run.
What a checkpoint weighs
Start from first principles: what must be saved, per parameter. For mixed-precision training with Adam, the usual layout keeps an FP32 master copy of each weight (4 bytes), two FP32 optimizer moments (4 bytes each) and the BF16 working copy the forward pass uses (2 bytes). The BF16 copy can be re-derived from the master, so many pipelines skip it. Gradients are not saved; they are recomputed on the next step.
| State | Bytes per parameter | 70B-parameter model |
|---|---|---|
| FP32 master weights | 4 | 280 GB |
| Adam first moment | 4 | 280 GB |
| Adam second moment | 4 | 280 GB |
| BF16 weights (optional) | 2 | 140 GB |
| Total, without BF16 copy | 12 | 840 GB |
Add a little for the learning-rate scheduler, the step counter, RNG states and the data loader's position. Those are tiny but not optional: without the loader position, resuming replays or skips data, and the model you get is not the model you would have trained. A checkpoint of weights alone is an export, not a checkpoint. Optimizers with different state change the arithmetic: SGD with momentum keeps one moment, and some memory-saving optimizers keep factored or 8-bit state. Measure your actual state dict size once rather than trusting the table.
The save path and sharding
Those bytes start in GPU memory and end on storage that survives the node. In between they cross the PCIe or chip-to-chip link to host memory, the host's memory system, possibly a local NVMe drive, and the network to shared storage. The slowest hop sets the save time, and for a whole cluster the slowest hop is usually aggregate write bandwidth on shared storage, because every node writes at once.
Sharding is what makes the save parallel. With FSDP or ZeRO-style training each rank holds a slice of the parameters and optimizer state, and each rank writes its own slice. That spreads 840 GB across every GPU in the job. The opposite pattern, gathering the full state on rank 0 and writing one file, serialises the whole checkpoint through one host's memory and network link and is the most common reason a checkpoint that took minutes at small scale takes an hour at large scale. DCP writes one or more files per rank plus a metadata file that maps logical tensors to shards, and it can reshard on load, so a checkpoint saved on 1,024 GPUs can be loaded on 512.
Worked example: 70B on 1,024 GPUs
Take the 70B model trained on 1,024 GPUs, 128 nodes of 8, with the 12-byte layout above. Assume shared storage sustains 50 GB/s of aggregate write bandwidth for this job, and that the job as a whole fails on average every 4 hours. These are assumptions to make the arithmetic concrete, not measurements; substitute your own.
- Per GPU shard: 840 GB / 1,024 = about 0.82 GB. Per node: about 6.6 GB of host staging memory. Small.
- Synchronous save time: 840 GB / 50 GB/s = about 17 seconds, during which all 1,024 GPUs are idle.
- Asynchronous save: the blocking part is the device-to-host copy of 0.82 GB per GPU, well under a second at PCIe rates even with eight GPUs sharing a host. Call it 1 second. The 17-second write runs in the background.
The classic first-order rule for the interval, from Young and later refined by Daly, is that the optimal interval is about the square root of two times the blocking checkpoint cost times the mean time between failures. Expected waste is then roughly the save cost divided by the interval, plus half the interval divided by the MTBF, plus restart time per failure, which is the same in both cases and left out here.
| Mode | Blocking cost C | Interval | Save overhead C/T | Expected rework T/2M | Total |
|---|---|---|---|---|---|
| Synchronous | 17 s | sqrt(2 x 17 x 14,400) = 700 s | 2.4% | 2.4% | 4.9% |
| Asynchronous | 1 s | sqrt(2 x 1 x 14,400) = 170 s | 0.6% | 0.6% | 1.3% (incl. commit lag) |
Async saving adds one cost the formula misses: a checkpoint is not usable until its background write has finished and been committed, so the restore point lags the save by the write time. That adds about the write time over the MTBF, 17 / 14,400 or roughly 0.1%, which is included above. The 170-second interval is ten times the 17-second write, so saves do not overlap; if contention on shared storage stretches writes toward the interval, lengthen the interval rather than letting saves queue. At 3.6 percentage points of a 1,024-GPU job, async saving is worth roughly 37 GPUs of throughput. The same arithmetic at 64 GPUs looks different: each GPU now holds 13 GB, each node stages 105 GB of host memory, and host RAM, not storage, can become the constraint.
Async sharded saving with PyTorch DCP
DCP's entry points are dcp.save, dcp.load and dcp.async_save, each taking a state dict and a checkpoint_id (a directory path for the file-system backend). async_save returns a future and accepts an async_checkpointer_type (a thread by default, or a separate process). By default it stages tensors into CPU memory before returning, which is the blocking copy from the worked example, and writes in the background. The helpers get_state_dict and set_state_dict in torch.distributed.checkpoint.state_dict produce and consume sharded model and optimizer state for FSDP and other parallel wrappers, and an object implementing the Stateful protocol is called during save and load so that everything is captured together.
import json, os
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict
from torch.distributed.checkpoint.stateful import Stateful
class TrainState(Stateful):
def __init__(self, model, optim, sched, loader, step=0):
self.model, self.optim, self.sched, self.loader, self.step = model, optim, sched, loader, step
def state_dict(self):
msd, osd = get_state_dict(self.model, self.optim)
return {"model": msd, "optim": osd, "sched": self.sched.state_dict(),
"loader": self.loader.state_dict(), "step": self.step}
def load_state_dict(self, sd):
set_state_dict(self.model, self.optim, model_state_dict=sd["model"], optim_state_dict=sd["optim"])
self.sched.load_state_dict(sd["sched"]); self.loader.load_state_dict(sd["loader"]); self.step = sd["step"]
def commit(root, ckpt_dir, step):
dist.barrier() # every rank's future has resolved
if dist.get_rank() == 0:
tmp = os.path.join(root, "latest.json.tmp")
with open(tmp, "w") as f:
json.dump({"step": step, "dir": ckpt_dir}, f)
f.flush(); os.fsync(f.fileno())
os.replace(tmp, os.path.join(root, "latest.json")) # atomic publish
# COMMIT_LAG * step_time must exceed the expected background write time; COMMIT_LAG < SAVE_EVERY.
state, pending = TrainState(model, optim, sched, loader), None
for batch in loader:
train_step(batch); state.step += 1
if pending is not None and state.step == pending[2] + COMMIT_LAG:
fut, d, s = pending; fut.result(); commit(ROOT, d, s); pending = None
if state.step % SAVE_EVERY == 0:
d = os.path.join(ROOT, f"step_{state.step:09d}")
pending = (dcp.async_save({"train": state}, checkpoint_id=d), d, state.step)Every rank reaches the commit branch on the same step, so the barrier inside commit cannot deadlock; deciding per rank with fut.done() could. The loader object needs a state_dict method; plain PyTorch DataLoader does not have one, so use a stateful loader or track the sample offset yourself. Restore is the mirror image: read latest.json, build the model and optimizer as for training, call dcp.load({"train": state}, checkpoint_id=d), and continue from state.step. Wait on the pending future before exiting, or the final checkpoint is lost.
Committing a checkpoint
A checkpoint is only useful if you can tell a complete one from a partial one. A job that dies halfway through a save leaves some ranks' files written and others missing. If the restart logic picks the newest directory, it loads garbage or fails, and if it fails on every restart, the job is wedged. The rule: never infer completeness from a directory's existence. Publish completeness explicitly, after every rank has finished, with an atomic operation. On a POSIX file system that is write-to-temp, fsync, then rename, as in the commit function above. On object storage it is writing a small marker object last, and having restore require it.
Retention follows from the commit rule: delete old checkpoints only after a newer one has committed, keep at least two, and keep some older ones on a longer schedule, because the failure you need to roll back from is sometimes a loss spike or a data problem discovered hours later, not a crash. The file-system writer's sync_files option, on by default, fsyncs shard files; leave it on unless you have measured that your storage's durability makes it redundant.
Failure modes and verification
| Failure | What you see | Fix |
|---|---|---|
| Gather-to-rank-0 save | Save time grows with cluster size; rank 0 host runs out of memory | Sharded save, one writer per rank |
| Partial checkpoint marked latest | Restart loops failing on load | Explicit commit after barrier; atomic rename or marker object |
| Overlapping async saves | Save time creeps up, storage saturated | At most one save in flight; interval above write time |
| Host memory exhaustion from staging | Node OOM-killed during save | Size staging per node; fewer GPUs per host or fewer, larger saves |
| Missing loader or RNG state | Loss curve jumps or repeats data after resume | Stateful object that saves loader, scheduler, RNG |
| Storage contention from many jobs | Everyone's checkpoints slow at the top of the hour | Stagger schedules; local NVMe tier then drain |
| Checkpoint saved after divergence | Restored model has NaN or a loss spike | Check loss and finiteness before committing; keep older checkpoints |
The last row matters because a checkpoint is not automatically good. Commit only after a cheap health check: loss finite and within a band of its recent average, and gradient norm not exploding. A restore drill is the real verification: periodically load the latest checkpoint into a separate small job, possibly with a different world size, run a few steps, and compare the loss with the training job's log at that step.
Trade-offs
Async saving trades host memory and complexity for GPU time, and on most large jobs it is the right trade. A local NVMe tier, writing to the node's own drive first and draining to shared storage in the background, cuts the time a checkpoint is exposed to shared-storage contention, but a node that dies before draining loses its shard, so the newest checkpoint is not durable until the drain finishes; commit only after it does. Fast storage paths such as GPUDirect Storage can shorten the device-to-storage path. Saving less, by skipping the BF16 copy or keeping optimizer state in lower precision, shrinks every hop, at the price of a format others must understand. On preemptible capacity, interval choice interacts with the provider's notice period; preemptible and Spot VMs works through that case.
What to do next
- Measure your real state dict size per rank and in total; do not trust the 12-bytes rule for your optimizer.
- Measure blocking save time and background write time separately, at full scale, with other jobs running.
- Switch any gather-to-rank-0 save to a sharded save.
- Move to async saving with at most one save in flight, and size host staging memory per node.
- Put model, optimizer, scheduler, loader position and RNG state behind one Stateful object.
- Add an explicit commit step: barrier, then atomic publish of a latest pointer or marker; make restore require it.
- Gate the commit on a loss and finiteness check; keep at least two committed checkpoints plus a longer-lived tier.
- Recompute the interval from measured blocking cost and the job's observed failure rate each time scale changes.
- Schedule a restore drill that loads on a different world size and compares loss.