Checkpoints, Export and Inference
Resuming training exactly where it stopped
Saving the weights alone gets you a model that continues from roughly the right place; saving the optimizer, the scheduler and the step count gets you a run that is indistinguishable from one that never stopped.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
To carry on properly you must save more than the weights — you must save the optimizer's memory too.
Think of pausing a long car journey. Writing down your location is not enough. You also need your current speed, which way you were pointing, and how much fuel you had. Start again from a standstill facing the wrong way and you will get there eventually, but not by the route you were on.
The model's weights are your location. The optimizer's memory is your speed and direction.
Why the optimizer has a memory
Modern optimizers do not react to the latest correction alone. They keep a running sense of which directions have been steady and which have been jumpy. They then move confidently in the steady directions and carefully in the noisy ones.
That accumulated sense takes hundreds of steps to build. Throw it away and the optimizer starts guessing again from nothing, taking a few large, unhelpful steps until it rebuilds.
You will not see an error. The loss will still fall. It will fall along a different path than it would have. The difference is invisible unless you compare against a run that was never interrupted.
How it works
what a real checkpoint holds:
model weights -> where you are
optimizer state -> your speed and direction
scheduler state -> how the step size was supposed to shrink
epoch / step -> how far through the journey you are
random state -> which way the dice were about to fallA real example you have seen
A video that resumes at the right timestamp but restarts the download from zero. It looks like it worked, and then it stutters for the next five minutes while it catches up. Position restored, buffer lost.
Remember this
- Weights alone are not enough to continue a run cleanly.
- Save the optimizer and scheduler state as well, plus the step number.
- Nothing errors when you get this wrong — the run takes a different path.
What to learn next
- Fixing missing and unexpected keys — when the checkpoint and the model disagree.
- Exponential moving average of weights — a second set of weights worth checkpointing.
- Early stopping and keeping the best model — deciding which checkpoint to keep.
Developer — Code and libraries.
Setup
pip install torchRuns on CPU in under a second. The point of the script below is to prove the difference rather than assert it.
Proof, three runs compared
import torch
import torch.nn as nn
def build():
torch.manual_seed(0) # same starting weights every time
model = nn.Sequential(nn.Linear(8, 16), nn.ReLU(), nn.Linear(16, 1))
opt = torch.optim.Adam(model.parameters(), lr=0.05)
sched = torch.optim.lr_scheduler.StepLR(opt, step_size=2, gamma=0.5)
return model, opt, sched
torch.manual_seed(1)
x, y = torch.randn(32, 8), torch.randn(32, 1)
def train(model, opt, sched, steps, start=0):
losses = []
for step in range(start, start + steps):
opt.zero_grad(set_to_none=True)
loss = nn.functional.mse_loss(model(x), y)
loss.backward()
opt.step()
sched.step()
losses.append(round(loss.item(), 4))
return losses
# 1) the reference: six steps, never interrupted
m, o, s = build()
reference = train(m, o, s, 6)
# 2) three steps, save EVERYTHING, then resume in fresh objects
m, o, s = build()
first_half = train(m, o, s, 3)
torch.save({"model": m.state_dict(), "opt": o.state_dict(),
"sched": s.state_dict(), "step": 3}, "full.pt")
m2, o2, s2 = build()
ckpt = torch.load("full.pt", weights_only=True)
m2.load_state_dict(ckpt["model"]); o2.load_state_dict(ckpt["opt"])
s2.load_state_dict(ckpt["sched"])
full_resume = first_half + train(m2, o2, s2, 3, start=ckpt["step"])
# 3) the common mistake: save only the weights
m3, _, _ = build()
m3.load_state_dict(ckpt["model"])
o3 = torch.optim.Adam(m3.parameters(), lr=0.05) # fresh, empty Adam state
s3 = torch.optim.lr_scheduler.StepLR(o3, step_size=2, gamma=0.5)
weights_only_resume = first_half + train(m3, o3, s3, 3)
print("uninterrupted :", reference)
print("full resume :", full_resume)
print("weights only :", weights_only_resume)
print("\nfull resume matches? ", full_resume == reference)
print("weights only matches? ", weights_only_resume == reference)uninterrupted : [1.2773, 1.125, 1.0259, 0.9668, 0.9082, 0.8785] full resume : [1.2773, 1.125, 1.0259, 0.9668, 0.9082, 0.8785] weights only : [1.2773, 1.125, 1.0259, 0.9668, 0.8794, 0.7127] full resume matches? True weights only matches? False
The full resume is identical, digit for digit, to the run that never stopped. That is the standard to aim for, and it is achievable.
Now look at the third row, because it is more interesting than "worse". Steps 5 and 6 read 0.8794 and 0.7127 against the reference's 0.9082 and 0.8785. The weights-only resume produced a lower loss.
That is not a bonus. Adam's bias correction is aggressive on its first few steps, so a freshly reset optimizer takes large steps that happen to cut the training loss here. The trajectory changed. On a real problem this shows up as a small dip followed by worse validation, and no log line anywhere says a thing about it.
What belongs in a production checkpoint
import torch
import torch.nn as nn
model = nn.Linear(4, 1)
opt = torch.optim.AdamW(model.parameters(), lr=1e-3)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=100)
scaler = torch.amp.GradScaler("cpu", enabled=False) # enabled=True on GPU AMP
state = {
"model": model.state_dict(),
"optim": opt.state_dict(),
"sched": sched.state_dict(),
"scaler": scaler.state_dict(), # loss scale, if you use mixed precision
"epoch": 7,
"global_step": 3521,
"best_metric": 0.912,
"torch_rng": torch.get_rng_state(), # so dropout and shuffling continue
"config": {"lr": 1e-3, "batch_size": 64}, # what produced this run
}
torch.save(state, "run.pt")
loaded = torch.load("run.pt", weights_only=True)
print("stored:", sorted(loaded))
print("resuming at epoch", loaded["epoch"], "step", loaded["global_step"])stored: ['best_metric', 'config', 'epoch', 'global_step', 'model', 'optim', 'scaler', 'sched', 'torch_rng'] resuming at epoch 7 step 3521
The config entry is not a PyTorch requirement. It is there because in six months you will find run.pt on a disk and have no idea what learning rate produced it.
Writing the file without corrupting it
import os
import torch
def save_atomic(state, path):
"""Write to a temporary file, then rename. Rename is atomic on one filesystem."""
tmp = path + ".tmp"
torch.save(state, tmp)
os.replace(tmp, path) # either the old file or the new one, never half
save_atomic({"step": 1}, "safe.pt")
print("wrote safe.pt, size", os.path.getsize("safe.pt"), "bytes")wrote safe.pt, size 864 bytes
The byte count depends on your PyTorch version; the pattern is the point. A job killed midway through torch.save leaves a truncated file, and if that file overwrote your only checkpoint the run is gone. Writing to a temporary name and renaming means the old checkpoint stays valid until the new one is complete.
Common mistakes
Saving only model.state_dict(). The output above shows what it costs. Save the optimizer.
Resuming without the scheduler. The learning rate silently restarts at its initial value. On a cosine schedule two-thirds of the way through a run, that is a large, sudden increase.
Loading the optimizer into a differently-ordered parameter list. Optimizer state is keyed by parameter index. Build the model and optimizer identically before loading, or Adam's moments land on the wrong tensors with no error at all.
Forgetting the epoch number. The run restarts at epoch 0, re-runs the warmup, and re-reads data you had finished with.
Ignoring the dataloader position. Restoring the RNG state gets you the same shuffle, but not the same position within an epoch. Simplest honest fix: checkpoint at epoch boundaries. If you need mid-epoch resumption, torchdata's StatefulDataLoader records the loader's position for you.
Saving every step to the same filename. Keep the last few plus the best. A single file is a single point of failure, and a checkpoint written during a crash is often the corrupt one.
Try it yourself
Change Adam to SGD with no momentum in the proof script and rerun. The weights-only row should now match the reference, because plain SGD keeps no state. Then add momentum=0.9 and watch it diverge again.
What to learn next
- Fixing missing and unexpected keys — when the checkpoint and the model disagree.
- Exponential moving average of weights — a second set of weights worth checkpointing.
- Early stopping and keeping the best model — deciding which checkpoint to keep.
Researcher — Mathematics and papers.
What the optimizer is actually carrying
Adam maintains per-parameter first and second moment estimates $m_t$ and $v_t$ with
$$ m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t, \qquad v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2 $$
where $g_t$ is the gradient at step $t$ and $\beta_1, \beta_2$ are the decay rates, defaulting to 0.9 and 0.999. The update divides by $\sqrt{\hat v_t}$, where $\hat v_t = v_t / (1 - \beta_2^t)$ is the bias-corrected second moment and $t$ is the step counter, also stored in the optimizer state.
Discarding that state resets $m_0 = v_0 = 0$ and $t = 1$. At $t=1$ the bias correction divides by $1 - \beta_2 = 10^{-3}$, so the first resumed update has magnitude close to the full learning rate regardless of the true gradient scale — which is precisely the mechanism that produced the anomalous drop in the third row of the output. The effective step size at $t=1$ after a reset is $O(\eta)$ rather than the $O(\eta \cdot |g|/\sqrt{v})$ the run had settled into. With $\beta_2 = 0.999$, $v$ needs on the order of $1/(1-\beta_2) = 1000$ steps to re-equilibrate.
Memory cost follows directly: Adam stores $2P$ floats of state for $P$ parameters, which is the $8P$ bytes in the FSDP memory table and the reason optimizer sharding was the first ZeRO stage.
Exact resumption, and its limits
Bitwise-identical resumption requires the full state of every stochastic component: parameters, optimizer state including step counters, scheduler last_epoch, GradScaler scale and growth tracker, and every RNG stream — torch.get_rng_state(), torch.cuda.get_rng_state_all(), numpy.random.get_state() and random.getstate() if your augmentations use them. DataLoader workers each derive a seed from the base seed and the worker id, so worker RNG is reproducible from the base seed but the consumption point within an epoch is not recoverable without a stateful loader.
Two further limitations are worth stating plainly. Nondeterministic kernels (atomics in scatter-add, cuDNN algorithm selection) break bitwise equality even in an uninterrupted run unless torch.use_deterministic_algorithms(True) is set. And a resumed run at a different world size cannot be bitwise identical, since the gradient reduction order differs — see DDP determinism.
Checkpoint policy at scale
Checkpoint interval is an optimization against expected failure rate. Young's approximation gives an optimal interval of $\sqrt{2 \cdot C \cdot M}$ where $C$ is the cost of writing a checkpoint and $M$ the mean time between failures; the expected wasted fraction is then roughly $\sqrt{2C/M}$. On large clusters where $M$ falls to hours, this pushes intervals to minutes and makes write cost the binding constraint — motivating asynchronous checkpointing, which copies state to pinned host memory and writes from a background thread, and sharded formats that parallelise the write across ranks.
References
- Kingma and Ba (2015), Adam: A Method for Stochastic Optimization — the moment estimates and bias correction that the state carries.
- Young (1974), A first order approximation to the optimum checkpoint interval, CACM 17(9) — the interval formula still used today.
- PyTorch documentation, Saving and Loading a General Checkpoint for Inference and/or Resuming Training — the canonical dictionary layout.
- Mohan et al. (2021), CheckFreq: Frequent, Fine-Grained DNN Checkpointing, FAST — asynchronous checkpointing and its measured overhead.
What to learn next
- Fixing missing and unexpected keys — when the checkpoint and the model disagree.
- Exponential moving average of weights — a second set of weights worth checkpointing.
- Early stopping and keeping the best model — deciding which checkpoint to keep.