Your PyTorch job resumed from a checkpoint and the loss jumped — what did it omit?
answer
- weights alone are a serving artifact
- Adam remembers more than the weights
- the schedule has a counter too
- load on CPU, then place
- the load default changed in 2.6
basics
~20 sAlmost always the optimizer state. Saving only model.state_dict() discards Adam's moment estimates and momentum buffers, so the first steps after resume behave like a freshly initialized optimizer. A resumable checkpoint also needs the scheduler, the AMP scaler and the step counter.
solid answer
~40 sA checkpoint that holds only weights is a *deployment* artifact, not a *resume* artifact. `torch.save` a dictionary containing `model.state_dict()`, `optimizer.state_dict()`, `scheduler.state_dict()`, `scaler.state_dict()` if you use AMP, and the epoch or global step. The optimizer is the usual culprit for a loss spike: Adam's `exp_avg` and `exp_avg_sq` buffers and its step counter live in `optimizer.state_dict()`, and without them the first updates after resume use uncorrected, badly-scaled moments. The second culprit is the scheduler — restart without it and the learning rate reverts to its initial value, which for a warmup-then-decay schedule means jumping back to a rate the model outgrew thousands of steps ago. On load, `torch.load(path, map_location="cpu")` then `load_state_dict` on each object in turn; since PyTorch 2.6 `weights_only` defaults to `True`, which loads plain state-dict payloads fine but rejects arbitrary pickled objects.
code
python · 20 linesimport torch
from torch import nn
model = nn.Linear(4, 2)
opt = torch.optim.AdamW(model.parameters(), lr=1e-3)
sched = torch.optim.lr_scheduler.StepLR(opt, step_size=10, gamma=0.5)
torch.save({
"model": model.state_dict(),
"optimizer": opt.state_dict(),
"scheduler": sched.state_dict(),
"global_step": 1234,
}, "ckpt.pt")
ckpt = torch.load("ckpt.pt", map_location="cpu", weights_only=True)
model.load_state_dict(ckpt["model"])
opt.load_state_dict(ckpt["optimizer"])
sched.load_state_dict(ckpt["scheduler"])
step = ckpt["global_step"]
print(step, sched.get_last_lr())go deeper
Know that torch.save and torch.load work on state_dicts rather than whole models, and that a checkpoint you intend to train from must include the optimizer, not just the weights.
List the contents of a resumable checkpoint and explain what each one restores — moment estimates, the schedule counter, the loss scale, the step number.
Diagnose a loss spike at the resume point by elimination: optimizer state first, scheduler second, then data order. Know the device and weights_only pitfalls of torch.load in 2.6 and later.
Set the policy: checkpoint cadence against restart cost, retention and naming, the separation between resume artifacts and serving artifacts, and a restore drill that is actually exercised rather than assumed.
## Two different artifacts "Save the model" means two incompatible things. A **serving** checkpoint needs only the weights: `torch.save(model.state_dict(), path)`. A **resume** checkpoint needs everything required to make step N+1 identical to what it would have been had the job never stopped. Conflating them is why resumed runs show a loss spike, a visible discontinuity in the curve at the restart point, and sometimes permanently worse final quality. ## What a resume checkpoint contains - `model.state_dict()` — parameters **and** buffers. Buffers matter: BatchNorm's `running_mean` and `running_var` are in here and are not owned by the optimizer. - `optimizer.state_dict()` — two parts: `param_groups` (the learning rate, weight decay, betas as they currently stand) and `state` (the per-parameter buffers). - `scheduler.state_dict()` — mostly `last_epoch`, the counter the schedule's formula is evaluated at. - `scaler.state_dict()` — the current AMP loss scale and growth tracker. - `epoch` and `global_step` — plain integers, needed to resume the schedule of anything driven by step count, and to keep logging aligned. - Optionally RNG states from `torch.get_rng_state()`, `torch.cuda.get_rng_state()`, plus the Python and NumPy generators, if you want bitwise reproducibility of augmentation and dropout. ## Why the optimizer state matters most Adam and AdamW keep, per parameter, an exponential moving average of the gradient (`exp_avg`) and of its square (`exp_avg_sq`), plus a step count used for bias correction. Those buffers are the accumulated knowledge of thousands of iterations about the curvature and scale of each coordinate. Drop them and the optimizer restarts with zeroed moments. The bias correction term makes the very first steps after such a restart unusually large, and the per-coordinate scaling is wrong until the moving averages refill — which takes on the order of `1/(1-beta2)` steps, thousands at the usual `beta2 = 0.999`. The result is exactly the reported symptom: the loss jumps at the resume point and takes a while to recover. Plain SGD with momentum has a smaller version of the same problem via its `momentum_buffer`. ## The scheduler and the scaler A scheduler is stateless apart from its counter, so restoring `last_epoch` is all it takes — but forgetting it is expensive. A cosine schedule restarted from zero jumps the learning rate back to its peak; a warmup schedule restarted from zero re-runs warmup on a model that is well past that phase. Both produce a loss spike that looks identical to the optimizer-state one, so check both. The AMP scaler is a smaller effect: without its state the scaler restarts from its default scale and may skip a handful of steps while it re-probes. Not catastrophic, but free to fix. ## Loading order and devices The robust idiom is to load onto CPU and let each object place its own tensors: 1. Construct the model with the same architecture, then `model.load_state_dict(ckpt["model"])`, then move the model to the device. 2. Construct the optimizer **over the moved parameters**, then `optimizer.load_state_dict(ckpt["optimizer"])` — it casts its state tensors to match each parameter's device. 3. Construct the scheduler over that optimizer and load its state. `map_location="cpu"` in `torch.load` avoids the classic failure of a checkpoint saved from `cuda:3` refusing to load on a box with two GPUs. `load_state_dict` is strict by default and raises on any missing or unexpected key; passing `strict=False` returns a named tuple of `missing_keys` and `unexpected_keys` instead, which is the right tool when deliberately loading a backbone into a model with a new head — and the wrong tool for silencing a mismatch you have not understood, because it will happily load nothing at all. ## weights_only in torch 2.6 and later `torch.load` used to unpickle arbitrarily, which meant loading an untrusted checkpoint could execute code. Since PyTorch 2.6 the `weights_only` argument defaults to `True`, restricting deserialization to tensors and a safe set of primitive containers. Dictionaries of state dicts, ints, floats and lists load fine under that default. What breaks is a checkpoint into which someone pickled a custom class, an argparse namespace, or a whole `nn.Module` object — those now raise, and the fixes are to allowlist the type or, for a checkpoint you trust and control, to pass `weights_only=False` explicitly. The durable fix is not to pickle live objects into checkpoints in the first place. ## What you still cannot restore Even a complete checkpoint does not make a resumed run bit-identical unless you also restore the position within the epoch's data order. The usual pragmatic compromise is to checkpoint on epoch boundaries so the sampler restarts cleanly, and to accept that mid-epoch resume replays some samples. Nondeterministic kernels and different GPU counts are further sources of divergence — worth stating explicitly when someone asks for "exact" resume.
- How would you make a mid-epoch resume replay no samples?You have to checkpoint the data position as well as the model: either save the sampler's epoch and the number of batches consumed and skip forward on resume, or use a sampler that derives its order deterministically from (epoch, seed) and can be fast-forwarded. Most teams sidestep it by checkpointing only on epoch boundaries and accepting the replay.
- What does model.load_state_dict(sd, strict=False) return, and when is it the right call?It returns a named tuple with missing_keys and unexpected_keys instead of raising. It is right when the mismatch is intentional — loading a pretrained backbone into a model with a new head, for example — and you then assert that the reported keys are exactly the ones you expected. It is wrong as a way to silence an error you have not diagnosed, because it can silently load almost nothing.
- A checkpoint saved from cuda:3 fails to load on a two-GPU machine. What is the fix?Pass map_location="cpu" (or a device map) to torch.load. Saved tensors record the device they came from, and the default load tries to restore them there. Loading to CPU first, then moving the model and constructing the optimizer over the placed parameters, makes checkpoints portable across machines with different GPU counts.
saying these in an interview costs you the question
- Saving only model.state_dict() and calling it a checkpoint
- Thinking Adam is stateless because the weights carry everything
- Restoring the model but rebuilding the scheduler from scratch
- Using strict=False to make a key mismatch go away
- Assuming torch.load still unpickles arbitrary objects by default