02.08 · UNIT 03 · PyTorch in One Hour: saving and scale · Lesson
Saving and loading models
Save the state dict to a file. Rebuild the same model, then load the parameters.
PLAIN-LANGUAGE INTRODUCTION
What is this?
Save the state dict to a file. Rebuild the same model, then load the parameters.
One simple example
A linear layer with 2 inputs and 2 outputs stores layer.weight of shape (2,2) and layer.bias of shape (2,).
What goes in?
A trained model and a file name.
What comes out?
A file of parameter tensors, then an identical reloaded model.
Why does it matter?
Training is expensive, so the learned state must survive the program.
What is it not?
A saved state dict does not include the model class or the code.
WORK THROUGH THE IDEA
See the idea in more detail
torch.save(model.state_dict(), 'model.pth')writes the parameter tensors to disk.- The state dict is a Python dictionary. Each key names a layer and each value is a tensor.
- Rebuild the same architecture, then call
model.load_state_dict(torch.load('model.pth', weights_only=True)). - A successful load prints
<All keys matched successfully>. - Common mistake: a model with different layer names or sizes reports a missing key or a size mismatch.