Skip to content
AI-grafen
EUniversityDeep learning· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

Checkpoints, saving and restarting

Be able to save and load models and optimiser state and resume training.

Prerequisites

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 partWhy
The model's weightsobviously
The optimiser stateAdam's m and v — without them the momentum restarts from zero
The scheduler's stateotherwise the learning rate jumps back
The epoch and the stepto know where you were
The random stateso that the data order and the dropout carry on the same way
The scaler state (AMP)the loss scaling factor
The configurationthe 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 checkpointAn export
The purposeresuming trainingrunning in production
The contenteverything aboveonly the weights and the graph
The format.pt with a state_dictsafetensors, ONNX, TorchScript
The sizelarge (the optimiser state dominates)small
The bindingto the code and the versionstandalone

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 ruleWhy
Write to a temporary file and renamean interruption in the middle of the writing otherwise gives a broken file
Save the latest and the best separatelythe latest for restarting, the best for use
Keep the N most recentthe disk fills up faster than you think
Save on time, not only per epochlong epochs give too sparse save points
Log the git commit in the checkpointotherwise 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

All the sources and licences