Tensor by Tensor

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

  1. torch.save(model.state_dict(), 'model.pth') writes the parameter tensors to disk.
  2. The state dict is a Python dictionary. Each key names a layer and each value is a tensor.
  3. Rebuild the same architecture, then call model.load_state_dict(torch.load('model.pth', weights_only=True)).
  4. A successful load prints <All keys matched successfully>.
  5. Common mistake: a model with different layer names or sizes reports a missing key or a size mismatch.
Open the detailed notes ↗