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.

On this page 6
  1. Why the labels drift apart
  2. The dangerous shortcut
  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.

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 random

A 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=False silences 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

Developer — Code and libraries.

Setup

bash
pip install torch

Runs on CPU in a second. Every error message below is a real captured one.

The module. prefix, and the trap under it

fix_prefix.py
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))
Output
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

new_head.py
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 ignored
Output
Error(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

safe_load.py
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))
Output
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

symptomcausefix
every key has module.saved from DDP or DataParallelconsume_prefix_in_state_dict_if_present(sd, "module.")
every key needs module.loading a plain file into a DDP modelload into ddp.module, before wrapping
every key has _orig_mod.model was compiled with torch.compilestrip that prefix, or save model._orig_mod.state_dict()
only the last layer mismatches on sizedifferent class countdrop those keys, strict=False, print the report
one renamed layeryou renamed an attributerebuild the dict with a rename map
num_batches_tracked unexpectedold checkpoint, newer PyTorchharmless; allow it explicitly

Renaming by hand is a dictionary comprehension, nothing more:

python
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

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