Checkpoints, Export and Inference
Fixing missing and unexpected keys
When a checkpoint and a model disagree, PyTorch tells you exactly which names are missing and which are unwanted — and strict=False is the switch that turns that useful error into a silent, untrained model.
- 10 min read
- 3 reading levels
- Published
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.
Loading a checkpoint is matching labels, and PyTorch will name every label that does not match.
Picture unpacking a delivery of labelled jars into a labelled shelf. Some shelf slots have no jar — those are missing. Some jars have no slot — those are unexpected. And sometimes a jar is the right name but the wrong size, so it will not sit in the slot.
PyTorch reports all three. The error message is genuinely one of the most helpful in the library. Read it.
Why the labels drift apart
Almost always one of four reasons.
You trained on several GPUs. The multi-GPU wrapper adds module. to the front of every label. The checkpoint says module.fc.weight; your plain model expects fc.weight.
You renamed something. self.fc became self.classifier and every label changed with it.
You are borrowing a pretrained model. Someone trained on 1000 categories, you have 3. The body fits and the final layer cannot.
The library changed. A new version added or renamed an internal buffer.
The dangerous shortcut
Search the internet for this error and the first answer is strict=False. That switch tells PyTorch to stop complaining about mismatched labels.
It does not fix anything. It loads whatever matched and leaves everything else at its random starting values, without a word. If nothing matched, you now have a completely untrained model that loaded "successfully".
strict=False is safe only when you print what it skipped and agree with the list.
How it works
checkpoint says model expects verdict
--------------- ------------- -------
module.fc.weight fc.weight name mismatch: fix the prefix
fc.weight (1000x8) fc.weight (3x8) size mismatch: new head, expected
backbone.conv1.weight backbone.conv1.weight match
(nothing) head.bias missing: will stay randomA real example you have seen
Restoring a phone backup onto a different brand of phone. Contacts and photos come across. The app layout does not, because the labels do not exist on the new phone. You want to be told which parts did not transfer.
Remember this
- The error names missing keys and unexpected keys. Read both lists.
strict=Falsesilences the complaint; it does not load anything extra.- Whenever you use
strict=False, print the report and check it is what you meant.
What to learn next
- Why torch.load can run code, and safetensors — the other thing that can go wrong at load time.
- Transfer learning in PyTorch — the workflow that produces most of these mismatches.
- Saving and logging from one rank only — how to avoid writing the
module.prefix in the first place.
Developer — Code and libraries.
Setup
pip install torchRuns on CPU in a second. Every error message below is a real captured one.
The module. prefix, and the trap under it
import torch
import torch.nn as nn
from torch.nn.modules.utils import consume_prefix_in_state_dict_if_present
torch.manual_seed(0)
trained = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 3))
# what a DDP run actually writes to disk: every key gains a "module." prefix
ddp_file = {"module." + k: v for k, v in trained.state_dict().items()}
target = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 3))
try:
target.load_state_dict(ddp_file)
except RuntimeError as e:
print(e)
report = target.load_state_dict(ddp_file, strict=False) # loads nothing, quietly
print("\nstrict=False -> missing:", report.missing_keys[:2], "...")
consume_prefix_in_state_dict_if_present(ddp_file, "module.")
print("after stripping:", target.load_state_dict(ddp_file))Error(s) in loading state_dict for Sequential:
Missing key(s) in state_dict: "0.weight", "0.bias", "2.weight", "2.bias".
Unexpected key(s) in state_dict: "module.0.weight", "module.0.bias", "module.2.weight", "module.2.bias".
strict=False -> missing: ['0.weight', '0.bias'] ...
after stripping: <All keys matched successfully>Read the middle line again. With strict=False, every single parameter failed to load and the call returned without raising. Your model is exactly as random as it was before. Train it and you will spend a day wondering why a pretrained checkpoint performs like noise.
The right fix is one line. consume_prefix_in_state_dict_if_present edits the dictionary in place, and does nothing if the prefix is absent, so it is safe to call unconditionally.
Better still, avoid creating the problem: save ddp_model.module.state_dict(), not ddp_model.state_dict(). See saving from one rank.
Size mismatch: the transfer-learning case
import torch
import torch.nn as nn
torch.manual_seed(0)
pretrained = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 1000)) # 1000 classes
mine = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 3)) # my 3 classes
try:
mine.load_state_dict(pretrained.state_dict())
except RuntimeError as e:
print(str(e).strip())
# keep the body, drop the head, and say so out loud
donor = {k: v for k, v in pretrained.state_dict().items() if not k.startswith("2.")}
report = mine.load_state_dict(donor, strict=False)
print("\nreused :", [k for k in donor])
print("left fresh:", report.missing_keys)
assert report.unexpected_keys == [], report.unexpected_keys # nothing silently ignoredError(s) in loading state_dict for Sequential:
size mismatch for 2.weight: copying a param with shape torch.Size([1000, 8]) from checkpoint, the shape in current model is torch.Size([3, 8]).
size mismatch for 2.bias: copying a param with shape torch.Size([1000]) from checkpoint, the shape in current model is torch.Size([3]).
reused : ['0.weight', '0.bias']
left fresh: ['2.weight', '2.bias']This is strict=False used correctly. The head is dropped deliberately, the report is printed, and the assert guarantees that nothing in the file was ignored by accident.
That assert is the habit worth stealing. unexpected_keys being non-empty means the checkpoint contained something your model has no home for — usually a sign that you are loading the wrong file, or that a name changed.
A loader that refuses to fail silently
import torch
import torch.nn as nn
def load_reporting(model, state, allow_missing=()):
"""Load leniently, but shout about anything not explicitly permitted."""
report = model.load_state_dict(state, strict=False)
surprises = [k for k in report.missing_keys
if not any(k.startswith(p) for p in allow_missing)]
if surprises or report.unexpected_keys:
raise RuntimeError(
f"unplanned mismatch\n missing: {surprises}\n"
f" unexpected: {report.unexpected_keys}")
print(f"loaded ok; deliberately fresh: {report.missing_keys}")
torch.manual_seed(0)
donor = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 1000)).state_dict()
mine = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 3))
body = {k: v for k, v in donor.items() if not k.startswith("2.")}
load_reporting(mine, body, allow_missing=("2.",)) # the new head is expected
try:
load_reporting(mine, {"nonsense.weight": torch.zeros(1)})
except RuntimeError as e:
print("\n" + str(e))loaded ok; deliberately fresh: ['2.weight', '2.bias'] unplanned mismatch missing: ['0.weight', '0.bias', '2.weight', '2.bias'] unexpected: ['nonsense.weight']
Thirty seconds of work, and a category of silent bug is gone from your project.
A repair table
| symptom | cause | fix |
|---|---|---|
every key has module. | saved from DDP or DataParallel | consume_prefix_in_state_dict_if_present(sd, "module.") |
every key needs module. | loading a plain file into a DDP model | load into ddp.module, before wrapping |
every key has _orig_mod. | model was compiled with torch.compile | strip that prefix, or save model._orig_mod.state_dict() |
| only the last layer mismatches on size | different class count | drop those keys, strict=False, print the report |
| one renamed layer | you renamed an attribute | rebuild the dict with a rename map |
num_batches_tracked unexpected | old checkpoint, newer PyTorch | harmless; allow it explicitly |
Renaming by hand is a dictionary comprehension, nothing more:
renamed = {k.replace("fc.", "classifier."): v for k, v in old.items()}Common mistakes
Reaching for strict=False first. It answers the error and not the question. Print missing_keys and unexpected_keys before you decide anything.
Never printing the report. load_state_dict returns a NamedTuple. Ignoring the return value is throwing away the only evidence you have.
Assuming a matching key was loaded correctly. Names can match while meanings do not, if two models share a layer name for different purposes. Nothing catches that but a sanity check on the outputs.
Loading into a torch.compiled model. Compilation wraps the module and adds _orig_mod. to every key. Save and load the original, then compile.
Ignoring num_batches_tracked. Usually harmless, but if the checkpoint is missing BatchNorm statistics rather than the counter, eval mode will behave very differently from training mode.
Try it yourself
Take the new_head.py script and change the head filter to k.startswith("0.") by mistake. Look at the report and work out what the model now consists of, before you run it.
What to learn next
- Why torch.load can run code, and safetensors — the other thing that can go wrong at load time.
- Transfer learning in PyTorch — the workflow that produces most of these mismatches.
- Saving and logging from one rank only — how to avoid writing the
module.prefix in the first place.
Researcher — Mathematics and papers.
What load_state_dict actually does
The method walks the module tree calling each module's _load_from_state_dict with a prefix, accumulating four lists: missing_keys, unexpected_keys, error_msgs, and a set of successfully loaded keys. Per-parameter copying happens through param.copy_(input_param) under torch.no_grad(), so it preserves the destination tensor's identity, device and dtype — the checkpoint's dtype is cast to the parameter's, not the other way round. That is why loading a float16 checkpoint into a float32 model succeeds silently and gives you float32 weights with float16 precision.
Shape checking is exact: input_param.shape != param.shape produces an error_msgs entry, and unlike missing or unexpected keys these are raised regardless of strict. strict=False suppresses only name mismatches. A common misreading is that strict=False tolerates shape differences; it does not, which is why the transfer-learning pattern must delete the offending keys rather than rely on the flag.
The return value is _IncompatibleKeys(missing_keys, unexpected_keys), whose __repr__ renders as <All keys matched successfully> when both lists are empty — a deliberately reassuring string that appears in the output above.
assign=True, added in PyTorch 2.1, replaces the parameter object with the checkpoint tensor instead of copying into it. That preserves the checkpoint's dtype and device, and is the mechanism behind meta-device loading: construct the model on the meta device with zero allocation, then load_state_dict(sd, assign=True) to materialise it, halving peak memory when loading a large model.
Prefix conventions and their origin
Both DataParallel and DistributedDataParallel store the wrapped model under the attribute name module, so the recursive prefixing produces module. on every key. torch.compile returns an OptimizedModule holding the original under _orig_mod, producing _orig_mod. for the same reason. Neither is special-cased in the serialisation format, because both are ordinary module attributes; consume_prefix_in_state_dict_if_present exists only as a convenience over the general rename.
Hooks provide the principled alternative to post-hoc renaming. register_load_state_dict_pre_hook runs before matching and can rewrite keys, which is how libraries migrate checkpoints across their own version changes without breaking users. _register_state_dict_hook does the mirror operation on save. Modules that change their internal parameterisation between releases are expected to use these, and the _metadata version field carried in the state_dict is what such a hook branches on.
Practical guarantees this buys you
A load that reports empty missing_keys and empty unexpected_keys and raises no shape error guarantees that every parameter and persistent buffer in the model received a value from the file. It guarantees nothing about semantics: two models can agree on every name and shape while meaning different things by them. For pipelines where a wrong-but-loadable checkpoint is a real risk, the cheap defence is to store a hash of the architecture — repr(model) or a sorted list of (name, shape) — alongside the weights, and compare on load.
References
- PyTorch documentation,
torch.nn.Module.load_state_dict—strict,assign, and the returned_IncompatibleKeys. - PyTorch source,
torch/nn/modules/module.py—_load_from_state_dict, the error accumulation, and the pre-hook interface. - PyTorch documentation, Meta device and
assign=True— low-memory loading of large models.
What to learn next
- Why torch.load can run code, and safetensors — the other thing that can go wrong at load time.
- Transfer learning in PyTorch — the workflow that produces most of these mismatches.
- Saving and logging from one rank only — how to avoid writing the
module.prefix in the first place.