Checkpoints, saving and restarting
Be able to save and load models and optimiser state and resume training.
Prerequisites
- DGit — version controlrequired
- DTrain a neural network in PyTorchrequired
Intuition
Saving only the model weights is not enough to be able to carry on training. A checkpoint has to contain everything the training depends on:
| The part | Why |
|---|---|
| The model's weights | obviously |
| The optimiser state | Adam's m and v — without them the momentum restarts from zero |
| The scheduler's state | otherwise the learning rate jumps back |
| The epoch and the step | to know where you were |
| The random state | so that the data order and the dropout carry on the same way |
| The scaler state (AMP) | the loss scaling factor |
| The configuration | the hyperparameters, so that the checkpoint is self-describing |
The most common mistake is to save only the model's state_dict() and believe that the training can be resumed. It can carry on, but not identically — and a restarted Adam often gives a visible hiccup in the loss curve.
Formal
Two completely different purposes, two formats:
| A checkpoint | An export | |
|---|---|---|
| The purpose | resuming training | running in production |
| The content | everything above | only the weights and the graph |
| The format | .pt with a state_dict | safetensors, ONNX, TorchScript |
| The size | large (the optimiser state dominates) | small |
| The binding | to the code and the version | standalone |
torch.load runs pickle and can therefore execute code when reading. Never load a checkpoint from an unknown source. Use weights_only=True when you only need the weights, and safetensors for everything that is shared — the format cannot by construction contain code.
A saving routine that holds in practice:
| The rule | Why |
|---|---|
| Write to a temporary file and rename | an interruption in the middle of the writing otherwise gives a broken file |
| Save the latest and the best separately | the latest for restarting, the best for use |
| Keep the N most recent | the disk fills up faster than you think |
| Save on time, not only per epoch | long epochs give too sparse save points |
| Log the git commit in the checkpoint | otherwise nobody knows which code created it |
strict=False when loading is convenient and dangerous: it allows layers to be missing or superfluous, which is exactly what you want when fine-tuning a subset — but also what hides the fact that you have loaded the wrong model. Always check the return value:
missing, unexpected = model.load_state_dict(sd, strict=False)
If the lists are non-empty when they should not be, you have just silently loaded a model with random weights in some layers.
Code
import os, random, subprocess, tempfile
from pathlib import Path
import numpy as np, torch
def git_commit():
try:
return subprocess.check_output(["git", "rev-parse", "--short", "HEAD"], text=True).strip()
except Exception:
return "unknown"
def save(path: Path, model, opt, sched, scaler, epoch, step, best, config):
ck = {
"model": model.state_dict(),
"optimiser": opt.state_dict(),
"scheduler": sched.state_dict() if sched else None,
"scaler": scaler.state_dict() if scaler else None,
"epoch": epoch, "step": step, "best": best,
"config": config, "git_commit": git_commit(),
"rng": {"python": random.getstate(),
"numpy": np.random.get_state(),
"torch": torch.get_rng_state(),
"cuda": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None},
}
path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp = tempfile.mkstemp(dir=path.parent, suffix=".tmp")
os.close(fd)
torch.save(ck, tmp)
os.replace(tmp, path) # atomic — never a half-written checkpoint
def load(path: Path, model, opt=None, sched=None, scaler=None):
ck = torch.load(path, map_location="cpu", weights_only=False)
missing, unexpected = model.load_state_dict(ck["model"], strict=True)
if opt:
opt.load_state_dict(ck["optimiser"])
if sched and ck.get("scheduler"):
sched.load_state_dict(ck["scheduler"])
if scaler and ck.get("scaler"):
scaler.load_state_dict(ck["scaler"])
r = ck["rng"]
random.setstate(r["python"]); np.random.set_state(r["numpy"])
torch.set_rng_state(r["torch"])
if r["cuda"] and torch.cuda.is_available():
torch.cuda.set_rng_state_all(r["cuda"])
return ck["epoch"], ck["step"], ck["best"], ck["git_commit"]
def tidy(directory: Path, keep=3):
files = sorted(directory.glob("step-*.pt"), key=lambda p: p.stat().st_mtime)
for f in files[:-keep]:
f.unlink()
# The export for production — safetensors cannot contain code
from safetensors.torch import save_file, load_file
save_file({k: v.contiguous() for k, v in model.state_dict().items()}, "model.safetensors")
model.load_state_dict(load_file("model.safetensors"))
Test the restart before you need it. Save at step 100, load, run to step 110, and compare the loss with an unbroken run. If they differ you have forgotten something in the checkpoint — and you would rather notice that now than after a day's training has been interrupted.
Mastery means
- Saves and loads the model and the optimiser state
- Resumes training exactly
- Knows the difference between a checkpoint and an export format
Sign in to do the exercises and build your mastery up.
Sources
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- safetensors (Apache-2.0) — Apache-2.0
- Hugging Face — dokumentation (Apache-2.0) — Apache-2.0