Lesson 8
Saving and loading models
Introduction
A trained model loses its learned weights when the program stops. Save the weights to a file. Then load them into a new model with the same architecture.
- Learning goal
Save model parameters with torch.save, rebuild the same architecture, and restore the parameters with torch.load and load_state_dict.
- Before you start
Model classes, nn.Module, model.parameters(), and a completed training loop.
Lesson plan
- Save the state dict of a trained model with torch.save.
- Read the state dict and confirm that it is a Python dictionary.
- Rebuild the same architecture, load the parameters, and confirm that all keys match.
Save the model parameters
Save a trained model to use it later. The recommended way is state_dict.
torch.save(model.state_dict(), "model.pth")
torch.save writes the values to disk. It does not write the model class or the Python code. The save operation stores only data.
The state_dict is a Python dictionary
The state_dict is a Python dictionary. Each key names a layer. Each value is a tensor of trainable parameters, such as weights and biases.
for name, tensor in model.state_dict().items():
print(name, tuple(tensor.shape))
# Expected for a model with one nn.Linear(2, 2):
# layer.weight (2, 2)
# layer.bias (2,)
A linear layer with 2 inputs and 2 outputs holds 4 weights and 2 biases. The shapes (2, 2) and (2,) count those values. The dictionary keeps the order of the model layers.
"model.pth" is a name for the file on disk. You can use any name. The endings .pth and .pt are common. The file ending does not change the contents.
Load the parameters
Restore the model from disk. First create a new model with the same architecture. Then load the file and apply the parameters.
model = NeuralNetwork(2, 2) # must match the original model
model.load_state_dict(torch.load("model.pth", weights_only=True))
torch.load("model.pth", weights_only=True) reads the file and rebuilds the Python dictionary. model.load_state_dict() applies the parameters to the model. This step restores the learned state from the time you saved the model. A successful load prints one line.
<All keys matched successfully>
weights_only=True tells torch.load to build only tensors and simple Python values. It does not run arbitrary code from the file. Keep this setting for a state dict.
The architecture must match
The line model = NeuralNetwork(2, 2) is necessary when the model is not in memory. The architecture must match the saved model exactly.
The state dict stores its parameters under named keys. A model with different layer names or different layer sizes does not match. PyTorch then reports a missing key, an unexpected key, or a size mismatch. Repair the model definition, then load again.
Load the parameters before you evaluate the model or continue training. Otherwise the new model uses random weights and gives meaningless results.
Reference
| Task | Code |
|---|---|
| Make a tensor | torch.tensor([[1, 2], [3, 4]]) |
| Shape · dtype · device | x.shape · x.dtype · x.device |
| Reshape / transpose | x.view(3, 2) · x.T |
| Matrix multiply | a @ b or a.matmul(b) |
| Track gradients | x.requires_grad_(True) |
| Compute gradients | loss.backward() |
| Read a gradient | w.grad |
| Skip gradient tracking | with torch.no_grad(): ... |
| Define a model | class Net(nn.Module): ... |
| List parameters | model.parameters() |
| Common loss | F.cross_entropy(logits, y) |
| Common optimizer | torch.optim.SGD(params, lr=0.5) |
| Training step | forward → loss → zero_grad → backward → step (Fast Loss = Zero Belly Size) |
| Mode switch | model.train() · model.eval() |
| Move to device | model.to(device) · x.to(device) |
| Save / load | torch.save(sd, "m.pth"); load_state_dict(...) |
| Multi-GPU model | DDP(model, device_ids=[rank]) |
| Launch DDP | torchrun --nproc_per_node=2 script.py |
Common pitfalls
- Save
model.state_dict(), not the model object. A saved model object depends on the original class. - Rebuild the same architecture before you call
load_state_dict. - Keep
weights_only=Truewhen you load a state dict. - A file ending does not change the contents.
.pthand.ptare only conventions. - Load the parameters before you evaluate the model or continue training.
Try it
Build a small two-input, two-output model. Save its parameters. Build a second model and load the parameters. Confirm that the second model holds the saved values.
Reveal the worked answer
import torch
import torch.nn as nn
class NeuralNetwork(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.layer = nn.Linear(in_features, out_features)
def forward(self, x):
return self.layer(x)
model = NeuralNetwork(2, 2)
with torch.no_grad():
model.layer.weight.copy_(torch.tensor([[1.0, 2.0], [3.0, 4.0]]))
model.layer.bias.copy_(torch.tensor([0.5, -0.5]))
torch.save(model.state_dict(), "model.pth")
fresh = NeuralNetwork(2, 2) # must match the original model
print(fresh.load_state_dict(torch.load("model.pth", weights_only=True)))
print(fresh.layer.weight.tolist())
# Expected:
# <All keys matched successfully>
# [[1.0, 2.0], [3.0, 4.0]]
The second model has the same layer names and shapes, so every key matches. The loaded weights equal the saved weights.
Recap
Save model.state_dict(), not the whole model. The state dict is a Python dictionary of parameter tensors. Rebuild the same architecture, then call load_state_dict(torch.load(..., weights_only=True)). A successful load prints <All keys matched successfully>.
Reference: PyTorch save and load tutorial.