skip to content

Why does renaming a submodule attribute in an nn.Module break old checkpoints?

level: principalimportance: should knowfreq 32%

answer

  1. keys are attribute paths, not identities
  2. one rename rewrites every key beneath it
  3. strict=False returns what it could not match
  4. DDP adds a module. prefix
  5. CI should load a pinned checkpoint

basics

~20 s

state_dict keys are built from the attribute path down the module tree, so self.fc produces fc.weight. Rename the attribute and every key under it changes, making a saved checkpoint's keys unrecognisable to the new class.

solid answer

~40 s

An `nn.Module`'s `state_dict()` keys are literally the dotted attribute paths of the module tree: `self.encoder.layers[2].fc` yields `encoder.layers.2.fc.weight`. Nothing else identifies a tensor — not its shape, not its position, not a stable id. So renaming `self.fc` to `self.head`, moving a block one level deeper, or switching an `nn.ModuleList` to an `nn.ModuleDict` rewrites the key strings, and `load_state_dict(strict=True)` then raises with missing and unexpected keys. That is the good outcome: with `strict=False` it returns an `_IncompatibleKeys` result and quietly leaves those weights at their initialization, so you serve a partly random model. The consequence for a team is that model attribute names are a **public contract**, not an implementation detail. Treat renames as versioned migrations: keep a key-remapping function alongside the model, and assert on load that nothing was silently missing.

code

python · 18 lines
python
import torch.nn as nn

class V1(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(4, 4)

class V2(nn.Module):
    def __init__(self):
        super().__init__()
        self.head = nn.Linear(4, 4)   # renamed attribute only

sd = V1().state_dict()
print(list(sd))                        # ['fc.weight', 'fc.bias']

result = V2().load_state_dict(sd, strict=False)
print(result.missing_keys)             # ['head.weight', 'head.bias']
print(result.unexpected_keys)          # ['fc.weight', 'fc.bias']

go deeper

for a junior

Know that state_dict keys come from attribute names, so a checkpoint only loads into a class whose structure matches, and that load_state_dict defaults to strict=True.

for a middle

Explain the key derivation through the tree, including ModuleList indices and the module. prefix added by DDP wrapping, and read a missing/unexpected key error correctly.

for a senior

Handle strict=False responsibly: capture the returned _IncompatibleKeys, assert the missing set is exactly what you intended to reinitialize, and recognise a silent partial load behind an unexplained metric drop.

for a principal

Own the contract. Declare model attribute names as public API, require a versioned remap with any rename, and put a pinned-checkpoint strict-load test in CI so cross-team breakage fails on the causing PR rather than in the registry.

## Where the keys come from `state_dict()` walks the module tree recursively, prefixing each level with the name under which the child was registered. Those names come straight from attribute assignment: `self.encoder = Encoder()` registers `encoder`; inside it, `self.layers = nn.ModuleList(...)` registers `layers` and each child under its index. The result is a flat dict whose keys look like `encoder.layers.2.fc.weight`. That is the entire identity scheme. There is no stable id, no shape-based matching, no fuzzy resolution. A checkpoint is a `{str: Tensor}` mapping and loading is a string lookup. ## What counts as a rename More refactors than people expect change keys: - renaming an attribute (`self.fc` → `self.head`); - extracting a group of layers into a new submodule, which *adds* a prefix level to everything inside; - inlining a submodule, which removes one; - reordering entries in an `nn.ModuleList`, since keys are positional indices; - converting a `ModuleList` to a `ModuleDict`, which swaps indices for string keys; - wrapping the model — `DistributedDataParallel` prefixes everything with `module.`, which is why so many loading scripts contain a line that strips exactly that prefix. None of these change the model's mathematics. All of them change the checkpoint contract. ## The two failure modes **Loud.** `load_state_dict(sd)` defaults to `strict=True` and raises a `RuntimeError` listing missing keys (in the model but not the checkpoint) and unexpected keys (in the checkpoint but not the model). Annoying, and exactly right. **Silent.** `load_state_dict(sd, strict=False)` returns an `_IncompatibleKeys(missing_keys, unexpected_keys)` namedtuple and loads whatever matched. If you ignore the return value — and most code does — the renamed submodule keeps its freshly initialized random weights. The model runs, produces plausible-looking output, and is quietly broken. Evaluation metrics drop in a way that looks like a training problem rather than a loading problem. This is the version that costs a week. Note that `strict=False` is legitimately useful — loading a pretrained backbone into a model with a new head is exactly the case it exists for — which is why the discipline is not "never use it" but "always inspect what it returned". ## How to manage renames deliberately 1. **Assert on load.** Capture the result and check it against an explicit allowlist: `assert set(result.missing_keys) <= EXPECTED_NEW_KEYS`. Anything unexpected fails the job immediately rather than at eval time. 2. **Ship a remap function with the version bump.** A dict of old-prefix → new-prefix, applied to the loaded `state_dict` before `load_state_dict`, converts old checkpoints instead of abandoning them. Keep it next to the model class, versioned with it. 3. **Version the checkpoint.** Save your own metadata alongside the weights — a model-code version string — so the loader can select the right migration rather than guessing from key shapes. 4. **Consider load hooks.** `nn.Module` supports `_register_load_state_dict_pre_hook`, which lets a module rewrite incoming keys as part of its own load. It is a private-ish API, so most teams prefer an explicit remap function they can read, but it is the mechanism the library itself uses for this problem. 5. **Freeze the public names.** In a shared model library, treat the top-level attribute names of any exported model as API. Rename freely inside a block whose weights nobody has saved; never rename `encoder`, `head`, `embed` casually. ## The organisational point The technical fix is small. The reason this is a lead-level question is that the cost lands somewhere other than where the change was made: a refactor in the model repo invalidates checkpoints in the registry, artifacts in an experiment tracker, deployed weights in a serving image, and every colleague's local runs. Nobody sees a compile error, and the CI that tests the model code passes because it builds fresh models. So the guardrail belongs in CI: a test that loads a pinned reference checkpoint into the current model class with `strict=True`. It costs seconds, it fails on exactly the refactors that matter, and it converts an invisible cross-team breakage into a red build on the PR that caused it. ## What is *not* the answer Shape-based automatic matching is tempting and wrong — two layers of equal shape are indistinguishable, and a mismatch would load weights into the wrong place, which is worse than failing. Nor does saving the whole module object with `torch.save(model)` solve it: that pickles the class path, so the artifact breaks when you move or rename the *class*, and it is a security and portability problem besides. Saving `state_dict()` and managing key names deliberately remains the recommended path.

  • What exactly does load_state_dict return when strict=False, and why is ignoring it dangerous?
    It returns `_IncompatibleKeys(missing_keys, unexpected_keys)` — names the model expected but the checkpoint lacked, and names the checkpoint carried that the model has no slot for. Ignore it and any renamed submodule silently retains its random initialization. The model runs and produces plausible output, so the failure surfaces as a mysterious quality regression rather than an error.
  • Why do so many loading scripts strip a leading "module." prefix from checkpoint keys?
    Because the checkpoint was saved from a `DistributedDataParallel`-wrapped model. The wrapper holds the real model as its `module` attribute, so every key gains that prefix. Loading into an unwrapped model then mismatches on all of them. The cleaner habit is to save `ddp_model.module.state_dict()` so the artifact is wrapper-agnostic in the first place.
  • How would you let a team rename model attributes without abandoning older checkpoints?
    Ship a versioned remap function with the change: save a model-code version in the checkpoint metadata, and on load apply an old-prefix to new-prefix mapping before calling `load_state_dict`. Pair it with a strict load afterwards so anything the remap missed fails loudly. The migration lives next to the model class and is reviewed with it.
  • What CI guardrail catches this class of breakage?
    A test that loads one pinned reference checkpoint into the current model class with `strict=True`. It runs in seconds, needs no GPU, and fails precisely on the refactors that rewrite key names — turning an invisible break across the registry, the tracker and the serving image into a red build on the PR that caused it.

saying these in an interview costs you the question

  • Assumes weights are matched by shape or position
  • Uses strict=False and ignores the returned keys
  • Thinks torch.save(model) avoids the naming problem
  • Treats submodule names as private implementation detail
  • Blames a quality regression on training, not loading

context