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.
- The model's weights.
model.state_dict()returns a dictionary that maps each layer's name to its tensor of weights or biases. - The optimizer's state. The optimizer is the object that updates the weights from the gradients. Plain SGD stores nothing between steps, but most optimizers do. Adam keeps two running averages for each weight: one of its recent gradients, usually written $m$, and one of its recent squared gradients, usually written $v$. It uses them to set the size of each update.
optimizer.state_dict()returns them. - The scheduler's state. A learning rate schedule changes the learning rate, the number that scales each update, as training goes on.
scheduler.state_dict()records how far along the schedule the run is. - Bookkeeping. The epoch number (an epoch is one pass over the training data) and the lowest validation loss so far, where the validation loss is the loss measured on held-out examples the model does not train on.
A complete checkpoint stores all four in one dictionary and writes it with torch.save.

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: TrueThe 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.

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.
- Latest checkpoint: written at the end of each epoch, or every fixed number of steps, over the previous one. It exists for crash recovery, so a crash costs at most the work since the last save.
- Best checkpoint: written only when the validation loss reaches a new low. Later epochs can make the model worse on held-out data, so the weights at the end of training are often not the ones to keep. This file is the model you evaluate and ship.
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.1462The 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.

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
- Saving only
model.state_dict()and then resuming training. The weights come back, and the optimizer's running averages and the learning rate schedule start over. - Saving the whole model object with
torch.save(model, path). That stores a reference to the class and the file it was defined in, so the file can stop loading once the code is renamed or moved. A state dict holds tensors and loads into any model with matching layers. - Resuming at the saved epoch instead of the next one. The saved epoch is repeated, and a schedule driven by the epoch number drifts by one.
- Overwriting the best checkpoint with a later, worse epoch. Guard the save with the
val_loss < best_val_losscomparison and restorebest_val_losswhen resuming. - Saving once at the end. A crash at 90% of the run then loses 90% of the work.
- Expecting a resumed run to match bit for bit when the data is shuffled. The script above matched because it reads batches in a fixed order. With a shuffling data loader, the random number generator has its own state, and a resumed run draws a different shuffle unless that state is saved as well.
Related concepts
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).