Checkpoints, Export and Inference

What a state_dict really is

A state_dict is an ordinary Python dictionary mapping names to tensors — every learned parameter and every remembered buffer, and nothing at all about the shape of your model.

On this page 6
  1. Why it works that way
  2. What goes in the bag
  3. How it works
  4. A real example you have seen
  5. Remember this
  6. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A state_dict is a labelled box of numbers: every name in your model paired with the numbers it holds.

Think of a recipe and a shopping bag. The recipe says "flour, sugar, cardamom" and the order to use them. The bag holds the actual flour and sugar.

Your model class is the recipe. The state_dict is the bag. Saving a model saves the bag, not the recipe — which is why you need the class definition again before you can load anything.

Why it works that way

PyTorch could have saved the whole model as one object. It can, and it is a bad idea.

Saving the object stores a reference to your class by name and file location. Rename the file, move the class, upgrade PyTorch, and the file will not open. Worse, opening it can run code hidden in the file — the subject of its own lesson.

Saving only the numbers avoids all of that. The file becomes a plain list of names and numbers, readable by any version, with no code inside.

What goes in the bag

Two kinds of numbers live there.

Parameters are what the model learns — the weights it adjusts during training.

Buffers are numbers the model remembers but does not learn. The running average that a normalisation layer keeps is the common example. They are saved too, and forgetting them gives you a model that behaves differently after loading.

How it works

your model class          the state_dict
-------------------       ----------------------------
fc  = Linear(3, 2)  ->    "fc.weight"  : 2x3 numbers
                          "fc.bias"    : 2 numbers
bn  = BatchNorm1d(2) ->   "bn.weight"  : 2 numbers   (learned)
                          "bn.bias"    : 2 numbers   (learned)
                          "bn.running_mean" : 2      (remembered)
                          "bn.running_var"  : 2      (remembered)

NOT in the dict: the class, the layer sizes, the optimizer, your code

A real example you have seen

Your phone's contacts backup. It saves the names and numbers, not the phone. Restore it on a new phone and the contacts come back — as long as you have a phone to put them in.

Remember this

  • A state_dict is a plain dictionary: names to tensors.
  • It holds learned parameters and remembered buffers, nothing else.
  • You always need the model class in code before you can load one.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Everything here runs on CPU in under a second. All outputs are real.

Looking inside one

inspect_state_dict.py
import torch
import torch.nn as nn

class TinyNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(3, 2)
        self.bn = nn.BatchNorm1d(2)
        self.register_buffer("seen_batches", torch.zeros(1))  # not learned, still saved

    def forward(self, x):
        self.seen_batches += 1
        return self.bn(self.fc(x))

model = TinyNet()
sd = model.state_dict()

for name, tensor in sd.items():
    print(f"{name:28} {str(tuple(tensor.shape)):8} {tensor.dtype}")

print("\ntype of state_dict:", type(sd).__name__)
print("is fc.weight the same object as the model's?",
      sd["fc.weight"].data_ptr() == model.fc.weight.data_ptr())
Output
seen_batches                 (1,)     torch.float32
fc.weight                    (2, 3)   torch.float32
fc.bias                      (2,)     torch.float32
bn.weight                    (2,)     torch.float32
bn.bias                      (2,)     torch.float32
bn.running_mean              (2,)     torch.float32
bn.running_var               (2,)     torch.float32
bn.num_batches_tracked       ()       torch.int64

type of state_dict: OrderedDict
is fc.weight the same object as the model's? True

Four things that surprise people, all visible above.

The keys are attribute paths. fc.weight is the weight attribute of the fc attribute. Rename self.fc to self.linear1 in your class and every old checkpoint stops loading.

Buffers are in there. bn.running_mean, bn.running_var and my hand-registered seen_batches are not learned by the optimizer. They are saved anyway, because eval mode depends on them.

num_batches_tracked is an integer. Not every entry is a float, and not every entry has a shape — that one is a scalar with shape ().

The tensors are shared, not copied. state_dict() hands you views of the live parameters. Keep the dictionary, train for ten steps, and your saved copy has changed too. torch.save serialises immediately so files are safe; a state_dict held in a Python variable is not. Use copy.deepcopy(model.state_dict()) when you want a snapshot in memory.

Saving and loading

save_load.py
import torch
import torch.nn as nn

torch.manual_seed(0)
model = nn.Sequential(nn.Linear(4, 3), nn.ReLU(), nn.Linear(3, 1))

torch.save(model.state_dict(), "tiny.pt")          # save the numbers, not the class

fresh = nn.Sequential(nn.Linear(4, 3), nn.ReLU(), nn.Linear(3, 1))
x = torch.randn(2, 4)
print("before loading, same output? ", torch.allclose(model(x), fresh(x)))

fresh.load_state_dict(torch.load("tiny.pt", weights_only=True))
print("after loading, same output?  ", torch.allclose(model(x), fresh(x)))

print("\nkeys in the file:", list(torch.load("tiny.pt", weights_only=True)))
Output
before loading, same output?  False
after loading, same output?   True

keys in the file: ['0.weight', '0.bias', '2.weight', '2.bias']

Note the keys: 0 and 2, not 1. In an nn.Sequential the key is the position, and position 1 is the ReLU, which has nothing to save. Insert a layer at the front later and every key shifts by one, which is a silent way to break every checkpoint you own.

weights_only=True restricts loading to plain tensors and refuses to execute code. It is the default from PyTorch 2.6 onward, and worth passing explicitly so your code reads the same on older versions.

What is not in there

A state_dict does not contain the model architecture, the optimizer's state, the epoch number, the learning-rate schedule, or your random number generator. A real training checkpoint is a dictionary of several state_dicts:

python
torch.save({
    "model": model.state_dict(),
    "optim": optimizer.state_dict(),     # Adam's moments live here
    "sched": scheduler.state_dict(),
    "epoch": epoch,
}, "checkpoint.pt")

Optimizers have a state_dict too, with a twist: it is keyed by parameter index, not by name. Change the order in which parameters are registered and the optimizer state maps onto the wrong tensors, quietly. The resume lesson shows what that costs.

Common mistakes

torch.save(model) instead of torch.save(model.state_dict()). It works today on your machine, and fails when the class moves or PyTorch upgrades. It also produces a file that runs code on load.

Renaming layers or reordering Sequential. The dictionary key is the attribute path. Renaming is a checkpoint-breaking change, and the mismatch lesson covers the repair.

Assuming state_dict() copies. It does not. Hold a snapshot with copy.deepcopy, or write it to disk.

Forgetting buffers exist. Custom saving code that filters to model.parameters() drops BatchNorm statistics. The model loads without complaint and evaluates differently.

Saving from a GPU and loading on a CPU-only box. torch.load tries to restore each tensor to the device it came from. Pass map_location="cpu" to avoid a RuntimeError about no CUDA device.

Try it yourself

Add self.dropout = nn.Dropout(0.5) to TinyNet and print the state_dict again. Work out why nothing new appears. Then swap BatchNorm1d for LayerNorm and count how many keys vanish.

What to learn next

Researcher — Mathematics and papers.

How it is assembled

nn.Module.state_dict() walks the module tree depth-first. At each module it calls _save_to_state_dict, which writes self._parameters then self._buffers into the destination dictionary under the current prefix, and then recurses into self._modules with the prefix extended by the child's name. That ordering explains the output above: seen_batches is a buffer of the root module and therefore precedes every child's entries.

Three registries determine what is captured. _parameters holds nn.Parameter instances, assigned through __setattr__ interception. _buffers holds tensors registered by register_buffer, with persistent=False excluding an entry from the state_dict while keeping the device-movement behaviour. _modules holds child modules. A plain tensor assigned as an attribute lands in none of them: it is not saved, and it is not moved by .to(device) — the most common cause of a device mismatch inside a custom module.

Modules may override _save_to_state_dict and _load_from_state_dict to change their serialised form. nn.utils.parametrize and weight normalisation use this to store the underlying parametrisation rather than the materialised weight, and quantised modules use it to store packed integer representations plus scales. The version-tracking hook (_version in _metadata) exists so a module can migrate an older layout on load — BatchNorm used it when num_batches_tracked was introduced.

The optimizer's state_dict

Optimizer.state_dict() returns {"state": {...}, "param_groups": [...]} where state is keyed by the integer index of each parameter in the flattened param_groups ordering, not by name. load_state_dict therefore requires that the new optimizer was constructed over parameters in the same order and with the same group structure. Nothing checks names, so a reordering loads Adam's first-moment estimate for layer 3 onto layer 5 and training continues, degraded, without an error. Any code that builds parameter groups from a dictionary iteration order, or that adds parameters conditionally, is exposed to this.

Storage and file layout

torch.save writes a ZIP container: a pickle describing the object graph, plus one binary record per storage. Tensors sharing a storage — views, slices, transposes — serialise as one record with offsets, so saving model.state_dict() after torch.nn.utils.prune or any view-heavy manipulation can produce a file far larger than the parameter count suggests, since the whole base storage is written. tensor.clone() before saving materialises the view and removes the surplus.

weights_only=True swaps the standard Unpickler for _weights_only_unpickler, which permits only a fixed allowlist of globals (tensor types, storages, collections.OrderedDict, primitives) and raises on anything else. It became the default in PyTorch 2.6. torch.serialization.add_safe_globals extends the allowlist for cases such as a dataclass stored alongside the weights.

References

  • PyTorch documentation, Saving and Loading Models — the recommended state_dict workflow and its rationale.
  • PyTorch source, torch/nn/modules/module.py — _save_to_state_dict, _load_from_state_dict and the metadata versioning hooks.
  • PyTorch documentation, torch.load — weights_only semantics and the safe-globals allowlist.

What to learn next