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.
- 9 min read
- 3 reading levels
- Published
Read these first
On this page 6
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 codeA 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
- Resuming training exactly where it stopped — everything else a real checkpoint must hold.
- Fixing missing and unexpected keys — what to do when the names do not line up.
- Buffers vs parameters — the distinction that decides what gets saved.
Developer — Code and libraries.
Setup
pip install torchEverything here runs on CPU in under a second. All outputs are real.
Looking inside one
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())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
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)))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:
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
- Resuming training exactly where it stopped — everything else a real checkpoint must hold.
- Fixing missing and unexpected keys — what to do when the names do not line up.
- Buffers vs parameters — the distinction that decides what gets saved.
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_dictand the metadata versioning hooks. - PyTorch documentation,
torch.load—weights_onlysemantics and the safe-globals allowlist.
What to learn next
- Resuming training exactly where it stopped — everything else a real checkpoint must hold.
- Fixing missing and unexpected keys — what to do when the names do not line up.
- Buffers vs parameters — the distinction that decides what gets saved.