A large training run is a long computation on hardware that fails often. Meta's Llama 3 paper reports 466 job interruptions during a 54-day snapshot of pre-training on 16,384 H100 GPUs, 419 of them unexpected and roughly 78% attributed to confirmed or suspected hardware problems; the team still kept effective training time above 90%. That number is not luck. It is the product of checkpoint policy: how often you save, what you save, how fast you can come back, and whether the run that comes back is the same run that died.
The mechanics of writing bytes, sharded and asynchronous saves and atomic commit markers, are covered in the GPU checkpointing deep dive. This article is about the decisions on top of that machinery. You will derive the optimal checkpoint interval from first principles, price a real failure rate, list every piece of state an exact resume needs, handle restarting on a different number of GPUs, design tiered retention, and build the test that proves a resume is faithful.
The cost model: saving versus rework
Every checkpoint policy trades two kinds of loss. Saving costs time: while the job stalls to copy state, GPUs do no useful work. Not saving costs time too: when a failure arrives, everything since the last durable checkpoint is recomputed, and the restart itself, rescheduling, loading weights, rebuilding communicators and warming up, is pure overhead.
Write C for the stall per save, T for the compute between saves, R for the restart cost and M for the mean time between failures of the whole job. Failures land uniformly within an interval, so on average you lose T/2 of work plus R per failure. Per unit of wall-clock time you pay C/T for saving and (T/2 + R)/M for failures. The waste fraction is W(T) = C/T + (T/2 + R)/M.
Note what M is: the job's failure rate, not a GPU's. If each component fails independently, the job's failure rate is the sum of its parts, so doubling the GPU count roughly halves M. A checkpoint interval that was generous at 512 GPUs can be wasteful at 16,000.
The optimal interval
Differentiate W with respect to T and set it to zero: -C/T^2 + 1/(2M) = 0, so T = sqrt(2CM). This is Young's 1974 result. Daly (2006) refined it with higher-order terms that matter when C is not small next to M; when it is small, the refinement is close to sqrt(2CM) - C. At the optimum the saving term and the rework term are equal, which gives a quick audit: if your job spends far more time saving than it loses to rework, save less often, and vice versa.
The restart term R does not depend on T at all. No interval can recover it; only faster restarts can. That is why mature stacks invest as much in hot spares, cached container images, pre-built communicators and fast restore tiers as in faster saves.
Worked example: pricing a real failure rate
Take the Llama 3 snapshot: 54 days is 77,760 minutes, and 419 unexpected interruptions give M of about 186 minutes, one failure every three hours or so. Assume a restart costs R = 10 minutes. The table computes the optimal interval and the resulting waste for three save stalls, compares it with a naive hourly schedule, and shows what happens if restarts drop to 2 minutes. These are model outputs from the formula above, not measurements of Meta's job.
| Stall per save C | Optimal T | Waste at optimal T | Waste at T = 60 min | Waste at optimal T, R = 2 min |
|---|---|---|---|---|
| 0.25 min | 9.6 min | 10.6% | 22.0% | 6.3% |
| 1 min | 19.3 min | 15.8% | 23.2% | 11.5% |
| 3 min | 33.4 min | 23.4% | 26.6% | 19.1% |
Three lessons fall out. First, with a one-minute stall you should save about every 19 minutes, far more often than the hourly habit many teams start with. Second, cutting the stall from one minute to fifteen seconds, the kind of cut asynchronous saving can buy, is worth a few points of goodput. Third, the restart term is at least as large as either of the others, so profile restart before save bandwidth.
One correction for asynchronous saves: the stall is small, but the checkpoint is not durable until its upload commits. If an upload takes U minutes, a failure during that window falls back to the previous checkpoint, so the expected rework is closer to T/2 + U. Measure C as the stall and add U to the rework term, rather than pretending the save was free.
What an exact resume must restore
An exact resume reproduces the run that would have happened without the failure. That requires more than weights. The table lists what a pre-training job must persist and what breaks when an item is missing.
| State | Where it lives | If you forget it |
|---|---|---|
| Model parameters | Sharded; saved through DCP | Any parallel layout via resharding |
| Optimizer state (moments, step counts) | Usually two fp32 tensors per parameter for Adam | Missing it resets Adam; loss spikes |
| LR scheduler | scheduler.state_dict() | Warmup restarts silently if lost |
| Gradient scaler (fp16 only) | scaler.state_dict() | Scale resets; early overflow skips |
| Step and samples-seen counters | Plain integers | Drive schedule, logging and data position |
| Data loader position | Per rank; StatefulDataLoader | Repeated or skipped data |
| RNG states | Python, NumPy, torch CPU, CUDA per rank | Different dropout and sampling |
| Code version and config | Git hash, resolved config | Silent behaviour change on resume |
The sharded state goes through PyTorch Distributed Checkpoint. get_state_dict and set_state_dict from torch.distributed.checkpoint.state_dict produce and consume state dicts with canonical fully qualified names that are portable across FSDP, DDP and tensor-parallel layouts, and an object implementing the Stateful protocol lets dcp.save and dcp.load call its state_dict and load_state_dict for you. Per-rank state, RNG and data loader position, is different on every rank, so this sketch keeps it in a small file per rank rather than in the shared state dict, where one key would collide across ranks. torchdata's StatefulDataLoader is a drop-in DataLoader with state_dict and load_state_dict; its docs note that it aggregates state across its worker processes but not across ranks.
A training loop with an adaptive interval
The loop below saves on a step boundary every rank agrees on, waits for the previous asynchronous save before committing it, measures the real stall and re-derives the interval from it. dcp.async_save returns a future whose result is the checkpoint metadata. train_step and latest_committed are yours.
import math, os, random, time
import numpy as np, torch
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, optimizer, scheduler):
self.model, self.optimizer, self.scheduler = model, optimizer, scheduler
self.step, self.samples_seen = 0, 0
def state_dict(self):
model_sd, optim_sd = get_state_dict(self.model, self.optimizer)
return {"model": model_sd, "optim": optim_sd, "sched": self.scheduler.state_dict(),
"step": self.step, "samples_seen": self.samples_seen}
def load_state_dict(self, sd):
set_state_dict(self.model, self.optimizer,
model_state_dict=sd["model"], optim_state_dict=sd["optim"])
self.scheduler.load_state_dict(sd["sched"])
self.step, self.samples_seen = sd["step"], sd["samples_seen"]
def rank_extras(loader):
return {"py": random.getstate(), "np": np.random.get_state(),
"torch": torch.get_rng_state(), "cuda": torch.cuda.get_rng_state(),
"loader": loader.state_dict(), "world": dist.get_world_size()}
def steps_between_saves(stall_s, mtbf_s, step_s):
# Young's interval, converted to steps and agreed by every rank.
obj = [max(1, round(math.sqrt(2 * stall_s * mtbf_s) / step_s))]
dist.broadcast_object_list(obj, src=0)
return obj[0]
def train(state, loader, root, mtbf_s, step_s, global_batch, gloo_pg):
pending, pending_dir = None, None
every = steps_between_saves(30.0, mtbf_s, step_s) # first guess for the stall
for batch in loader:
train_step(state, batch)
state.step += 1
state.samples_seen += global_batch
if state.step % every:
continue
if pending is not None: # previous save must be durable first
pending.result()
dist.barrier()
if dist.get_rank() == 0:
open(f"{pending_dir}/COMMITTED", "w").close()
ckpt_dir = f"{root}/step_{state.step:09d}"
os.makedirs(ckpt_dir, exist_ok=True)
torch.save(rank_extras(loader), f"{ckpt_dir}/extras_rank{dist.get_rank()}.pt")
t0 = time.monotonic() # host-side stall is what training pays
pending = dcp.async_save({"train": state}, checkpoint_id=ckpt_dir,
process_group=gloo_pg) # CPU-capable group
pending_dir = ckpt_dir
every = steps_between_saves(time.monotonic() - t0, mtbf_s, step_s)Two details matter. The interval is computed on rank 0 and broadcast, because wall clocks and timings differ across ranks, and a job where rank 17 decides to save one step later than everyone else deadlocks inside a collective. The PyTorch async checkpoint recipe runs on a gloo process group, so pass a gloo group created with dist.new_group if your default group is NCCL. And torch.load defaults to weights_only=True in recent PyTorch releases, which refuses NumPy RNG state, so the per-rank file is loaded with it explicitly off; only do that for files your own job wrote.
Restarting on a different number of GPUs
Restarts often land on a different number of GPUs: a node is drained, spares run out, or you deliberately shrink to keep going. DCP supports load-time resharding, so weights and optimizer state saved under one parallel layout load into another. The hard parts are everything around the tensors.
def resume(state, loader, root, base_seed):
latest = latest_committed(root) # newest step_* dir holding COMMITTED
if latest is None:
return
dcp.load({"train": state}, checkpoint_id=latest) # reshards if the layout changed
path = f"{latest}/extras_rank{dist.get_rank()}.pt"
ex = torch.load(path, weights_only=False) if os.path.exists(path) else None
if ex is not None and ex["world"] == dist.get_world_size():
random.setstate(ex["py"]); np.random.set_state(ex["np"])
torch.set_rng_state(ex["torch"]); torch.cuda.set_rng_state(ex["cuda"])
loader.load_state_dict(ex["loader"])
else: # topology changed: rebuild, do not guess
seed = hash((base_seed, state.step, dist.get_rank())) % 2**31
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
loader.sampler.set_start(state.samples_seen) # your deterministic global sampler- Data position. Per-rank loader state is meaningless when the data-parallel size changes. Keep a global samples-seen counter in the shared state and build the sampler so that it can start from any global offset deterministically.
- Global batch. Keep the global batch constant by changing gradient accumulation steps, or the optimizer sees a different noise scale and the LR schedule no longer means what it did.
- RNG. Exact equality with the uninterrupted run is impossible after a topology change. Re-seed deterministically from the step and rank, and record that this segment is not bitwise comparable.
Tiered storage and retention
Not every checkpoint needs the same durability. A copy in host memory, or mirrored to a peer node's memory, restores in seconds and covers the common case of a process or single GPU failing. Local or shared NVMe survives process crashes and most node reboots. An object store survives losing the cluster. Write to the fast tier at the Young interval and promote every few checkpoints to the slow tier, then restore from the fastest tier holding a committed copy.
Retention needs rules written down before the disk fills: keep the last few committed checkpoints on fast tiers, keep a sparser set on the durable tier, never delete the newest committed one, and pin milestones such as end of warmup, data-mix changes and any step you might want to branch from. Garbage-collect only directories that carry a commit marker and are older than the newest one; an uncommitted directory may still be uploading.
Proving a resume is faithful
A checkpoint you have never restored is a hypothesis. The cheapest strong test is a resume equivalence run: train a small configuration for 2N steps straight, then train it for N steps, save, kill the process, resume and train N more. With deterministic kernels enabled the two loss curves should match bit for bit; without them they should match to within the noise you see between two identical uninterrupted runs. A jump at the resume step points at missing optimizer or scheduler state; a slow drift points at data order or RNG.
Run that test in CI on every training-loop change, and periodically restore the real job's latest durable checkpoint on spare capacity and compare its validation loss with the logged value.
Failure modes
- Resuming from an uncommitted checkpoint. The newest directory exists but its upload never finished. Always select by commit marker.
- Ranks disagree on when to save. Time-based triggers evaluated per rank deadlock collectives. Decide on one rank, broadcast, save on a step boundary.
- Learning-rate warmup reruns. The scheduler was rebuilt rather than restored; the loss spikes right after every restart.
- Data seen twice. The loader restarted at the epoch start; evaluation looks better than it should because the model saw those samples again.
- Corrupt state saved faithfully. A silent data corruption or a NaN made it into a checkpoint; restarting reloads the poison. Validate loss and parameter norms before committing, and keep older checkpoints to roll back past it. See GPU hardware faults for where such errors come from.
- Interval tuned at the wrong scale. M came from a smaller cluster.
Trade-offs
Asynchronous saving cuts the stall but holds an extra copy of state in host memory and delays durability. In-memory tiers restart fast but vanish with the node. Exact resume needs determinism that can cost kernel speed, so many teams check bitwise only on a small CI configuration. Sharding optimizer state with ZeRO or FSDP shrinks each rank's share of a save, lowering C and the optimal interval.
What to do next
- Pull your job's interruption log and compute M from it, not from a vendor sheet.
- Measure the real stall C and the restart cost R end to end, including scheduling.
- Set the interval to sqrt(2CM), broadcast it as a step count, and re-derive it after each save.
- Persist every row of the state table, including scheduler, scaler, counters, loader and RNG.
- Commit with a marker and resume only from committed checkpoints.
- Make the sampler resumable from a global sample offset so you can restart on fewer GPUs.
- Add a fast restore tier and write retention rules before the disk fills.
- Put a 2N versus N-plus-N resume test in CI and run a restore drill on the real job.