QuiddityML

1 October 2026 · 9 min read

How to save and resume training in PyTorch (checkpoints)

A checkpoint is a file that stores a training run so it can be continued later. This post explains what to put in one, how to save and load it in PyTorch, why saving only the model weights is not enough to resume, and why most runs keep a latest and a best checkpoint.

A checkpoint is a file written during training that holds what is needed to continue the run later from the same point. Training can take hours or days, and a crash, a closed laptop or a job time limit can end it early. Without a checkpoint the run starts again from the first step. With one, it picks up at the last save.

What training state is

Training changes more than the model's weights. A typical PyTorch run has four pieces of state that move as training goes on.

A complete checkpoint stores all four in one dictionary and writes it with torch.save.

A diagram titled "What a complete checkpoint holds". Four boxes, model.state_dict() with weights and biases, optimizer.state_dict() with Adam moments m and v, scheduler.state_dict() with the position in the schedule, and epoch plus best_val_loss, all point into one file called checkpoint_latest.pt. Below, a check mark says all four saved and the resume continues smoothly, and a cross says weights only and the optimizer restarts cold.

Saving and resuming, in one script

The script below trains a small network on made-up data three ways. Run A trains for 20 epochs with no interruption. Run B trains for 10 epochs, saves a checkpoint, throws its objects away, builds new ones, loads the checkpoint, and trains 10 more. Run C loads the same file but restores only the model weights. The learning rate starts at 0.01 and a StepLR schedule halves it every 10 epochs.

1import os
2import tempfile
3import torch
4import torch.nn as nn
5 
6torch.manual_seed(0)
7X = torch.randn(600, 10)
8y = torch.sin(X[:, :1]) + 0.5 * X[:, 1:2] * X[:, 2:3] + 0.3 * torch.randn(600, 1)
9X_train, y_train, X_val, y_val = X[:500], y[:500], X[500:], y[500:]
10loss_fn = nn.MSELoss()
11 
12def build(seed):
13    torch.manual_seed(seed)
14    model = nn.Sequential(nn.Linear(10, 64), nn.ReLU(), nn.Linear(64, 1))
15    optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
16    scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.5)
17    return model, optimizer, scheduler
18 
19def run_epoch(model, optimizer, scheduler):
20    model.train()
21    for i in range(0, 500, 50):
22        loss = loss_fn(model(X_train[i:i + 50]), y_train[i:i + 50])
23        optimizer.zero_grad()
24        loss.backward()
25        optimizer.step()
26    scheduler.step()
27    model.eval()
28    with torch.no_grad():
29        return loss_fn(model(X_val), y_val).item()
30 
31path = os.path.join(tempfile.mkdtemp(), "checkpoint_latest.pt")
32 
33# Run A: 20 epochs with no interruption
34model, optimizer, scheduler = build(seed=1)
35for epoch in range(20):
36    val_a = run_epoch(model, optimizer, scheduler)
37    if epoch == 10:
38        val_a_11 = val_a
39 
40# Run B: 10 epochs, save, stop
41model, optimizer, scheduler = build(seed=1)
42for epoch in range(10):
43    val_loss = run_epoch(model, optimizer, scheduler)
44torch.save({
45    "epoch": epoch,
46    "model_state_dict": model.state_dict(),
47    "optimizer_state_dict": optimizer.state_dict(),
48    "scheduler_state_dict": scheduler.state_dict(),
49    "val_loss": val_loss,
50}, path)
51print(f"saved after epoch {epoch + 1}, val loss {val_loss:.4f}, lr {scheduler.get_last_lr()[0]}")
52 
53# Run B, resumed: new objects with a different seed, then load everything
54model, optimizer, scheduler = build(seed=999)
55checkpoint = torch.load(path, weights_only=True)
56model.load_state_dict(checkpoint["model_state_dict"])
57optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
58scheduler.load_state_dict(checkpoint["scheduler_state_dict"])
59start_epoch = checkpoint["epoch"] + 1
60print(f"resuming at epoch {start_epoch + 1}, lr {scheduler.get_last_lr()[0]}")
61for epoch in range(start_epoch, 20):
62    val_b = run_epoch(model, optimizer, scheduler)
63    if epoch == 10:
64        val_b_11 = val_b
65 
66# Run C, resumed with the weights only: optimizer and scheduler start fresh
67model, optimizer, scheduler = build(seed=999)
68model.load_state_dict(checkpoint["model_state_dict"])
69print(f"weights-only resume, lr {scheduler.get_last_lr()[0]}")
70for epoch in range(start_epoch, 20):
71    val_c = run_epoch(model, optimizer, scheduler)
72    if epoch == 10:
73        val_c_11 = val_c
74 
75print(f"epoch 11 val loss  uninterrupted {val_a_11:.4f} | full resume {val_b_11:.4f} | weights only {val_c_11:.4f}")
76print(f"epoch 20 val loss  uninterrupted {val_a:.4f} | full resume {val_b:.4f} | weights only {val_c:.4f}")
77print("full resume matches uninterrupted run:", val_a == val_b)
1saved after epoch 10, val loss 0.1553, lr 0.005
2resuming at epoch 11, lr 0.005
3weights-only resume, lr 0.01
4epoch 11 val loss  uninterrupted 0.1529 | full resume 0.1529 | weights only 0.1799
5epoch 20 val loss  uninterrupted 0.1466 | full resume 0.1466 | weights only 0.1538
6full resume matches uninterrupted run: True

The fully restored run ends on the same validation loss as the run that was not interrupted, 0.1466, even though its model, optimizer and scheduler were rebuilt from a different random seed before loading. The file carried everything the next step depended on.

Three details in the loading code are worth copying. The model, optimizer and scheduler are built first and then filled in with load_state_dict, because a state dict holds values and has no code in it. start_epoch is the saved epoch plus one, since the saved epoch already finished. weights_only=True tells torch.load to read tensors and plain Python values and refuse to run other code stored in the file, which is the safer setting for a file from another machine or person.

What goes wrong with weights only

Run C starts from the same weights as run B, and its first epoch back is worse: 0.1799 against 0.1529. Two things were lost.

The first is Adam's running averages. A fresh Adam starts both averages at zero, so for its first steps the update sizes are not scaled by the gradient history the run had built up, and the weights move further than they would have.

The second is the position in the schedule. The saved run had already halved its learning rate to 0.005. The fresh scheduler starts again at 0.01, which the output shows on the "weights-only resume" line.

Run C recovers part of the way and still ends at 0.1538 instead of 0.1466. The gap is small on this toy problem. It tends to be larger on long runs that use warmup or a decayed learning rate, where restarting the schedule sends a nearly trained model back to the largest learning rate.

A sketch of training loss against training step, titled "Resuming without the optimizer state". The loss falls from 0.9 to 0.5 by step 2000, where a dashed line marks the resume. A red curve labeled "weights only" jumps up before coming back down, and a teal curve labeled "full restore" keeps falling smoothly. A box at the bottom says load model, optimizer, scheduler.

The picture is a sketch of the shape, and the size of the jump depends on the optimizer, the learning rate and the model.

Keeping a latest and a best checkpoint

A run usually keeps two checkpoint files, because there are two different reasons to save.

1import os
2import tempfile
3import torch
4import torch.nn as nn
5 
6torch.manual_seed(0)
7X = torch.randn(260, 10)
8y = torch.sin(X[:, :1]) + 0.5 * X[:, 1:2] * X[:, 2:3] + 0.3 * torch.randn(260, 1)
9X_train, y_train, X_val, y_val = X[:200], y[:200], X[200:], y[200:]
10 
11model = nn.Sequential(nn.Linear(10, 64), nn.ReLU(), nn.Linear(64, 1))
12optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
13loss_fn = nn.MSELoss()
14folder = tempfile.mkdtemp()
15 
16def save_checkpoint(state, name):
17    tmp_path = os.path.join(folder, name + ".tmp")
18    torch.save(state, tmp_path)
19    os.replace(tmp_path, os.path.join(folder, name))  # swap in only after the write finished
20 
21best_val_loss = float("inf")
22for epoch in range(12):
23    model.train()
24    for i in range(0, 200, 20):
25        loss = loss_fn(model(X_train[i:i + 20]), y_train[i:i + 20])
26        optimizer.zero_grad()
27        loss.backward()
28        optimizer.step()
29 
30    model.eval()
31    with torch.no_grad():
32        val_loss = loss_fn(model(X_val), y_val).item()
33 
34    improved = val_loss < best_val_loss
35    best_val_loss = min(best_val_loss, val_loss)
36    state = {
37        "epoch": epoch,
38        "model_state_dict": model.state_dict(),
39        "optimizer_state_dict": optimizer.state_dict(),
40        "best_val_loss": best_val_loss,
41        "val_loss": val_loss,
42    }
43    save_checkpoint(state, "checkpoint_latest.pt")      # every epoch
44    if improved:
45        save_checkpoint(state, "checkpoint_best.pt")    # only on a new best
46    print(f"epoch {epoch + 1:2d}  val loss {val_loss:.4f}  latest saved" + ("  best saved" if improved else ""))
47 
48best = torch.load(os.path.join(folder, "checkpoint_best.pt"), weights_only=True)
49latest = torch.load(os.path.join(folder, "checkpoint_latest.pt"), weights_only=True)
50print(f"best checkpoint:   epoch {best['epoch'] + 1}, val loss {best['val_loss']:.4f}")
51print(f"latest checkpoint: epoch {latest['epoch'] + 1}, val loss {latest['val_loss']:.4f}")
1epoch  1  val loss 0.4969  latest saved  best saved
2epoch  2  val loss 0.2874  latest saved  best saved
3epoch  3  val loss 0.2204  latest saved  best saved
4epoch  4  val loss 0.2019  latest saved  best saved
5epoch  5  val loss 0.1839  latest saved  best saved
6epoch  6  val loss 0.1543  latest saved  best saved
7epoch  7  val loss 0.1437  latest saved  best saved
8epoch  8  val loss 0.1454  latest saved
9epoch  9  val loss 0.1451  latest saved
10epoch 10  val loss 0.1479  latest saved
11epoch 11  val loss 0.1470  latest saved
12epoch 12  val loss 0.1462  latest saved
13best checkpoint:   epoch 7, val loss 0.1437
14latest checkpoint: epoch 12, val loss 0.1462

The validation loss stops improving after epoch 7, so the best file stays at epoch 7 while the latest file moves on to epoch 12. best_val_loss is stored in the checkpoint too. A resumed run that forgets it starts from infinity and overwrites the best file on its first epoch back, even when that epoch is worse.

A diagram titled "Keep two checkpoints". A validation loss line over eight epochs reads 0.90, 0.74, 0.62, 0.55, 0.61, 0.48, 0.41, 0.47. A row labeled latest checkpoint has a save icon at all eight epochs, with the note saved every epoch, for crash recovery, lose at most one epoch. A row labeled best checkpoint has a save icon at epochs 1, 2, 3, 4, 6 and 7 and none at epochs 5 and 8, where the loss went up, with the note saved only when validation improves, this is the model you deploy. A code box shows if val_loss < best_val_loss: torch.save(checkpoint, checkpoint_best.pt).

The picture shows the same rule on a run where the validation loss goes back up twice: epochs 5 and 8 write the latest file and leave the best file alone.

save_checkpoint writes to a temporary file and then renames it with os.replace. A crash in the middle of torch.save would otherwise leave a half-written latest checkpoint, which is the one file crash recovery depends on. When a resume fails with an error about missing or unexpected keys, compare the model you built against the one that was saved, since the layer names and shapes have to match.

Loading a checkpoint for evaluation

Using the best checkpoint to make predictions needs the model weights only. Build the model, call load_state_dict with checkpoint["model_state_dict"], then call model.eval() so layers such as dropout switch to their evaluation behaviour. If the file was saved on a GPU and is being loaded on a machine without one, pass map_location="cpu" to torch.load.

Common mistakes

Adam explains the two running averages that make the optimizer state worth saving. Learning rate schedules covers warmup, cosine and step decay, which are what a restarted scheduler repeats from the beginning. Early stopping, covered under regularization, ends training when the validation loss stops improving and returns the best checkpoint. The train, validation and test split explains why the best checkpoint is picked on validation data and scored once on test data.

QuiddityML teaches checkpointing as a concept in Unit 3 of the ML Foundation track, where the exercises include ordering the lines of save_checkpoint, spotting the bug in a resume that restarts at the saved epoch instead of the next one, and writing the save, load and training loop from scratch (quiddityml.com).