In PyTorch, what breaks if you store layers in a plain Python list inside nn.Module?
answer
- registration happens on attribute assignment
- a list is not an nn.Module
- parameters(), to(), state_dict() all miss them
- ModuleList stores, Sequential chains
- suspiciously small parameter count
basics
~10 sSubmodules in a plain list are never registered, so model.parameters() misses them, the optimizer never updates them, and .to(device) leaves them on the CPU. Use nn.ModuleList or nn.ModuleDict instead.
solid answer
~40 s`nn.Module` registers submodules through `__setattr__`: assigning an `nn.Module` to an attribute puts it in the module's `_modules` dict, which is what the recursive walks use. A Python `list` is not a module, so assigning `self.layers = [nn.Linear(4, 4) for _ in range(3)]` registers the list itself and nothing inside it. The consequences are all silent: `model.parameters()` yields nothing for those layers so the optimizer never updates them, `model.to(device)` never moves their weights, and `state_dict()` omits them so checkpoints save an incomplete model. The fix is `nn.ModuleList` for sequences and `nn.ModuleDict` for name-keyed collections — both are containers that register their children while still supporting indexing and iteration. `nn.ParameterList` and `nn.ParameterDict` do the same for bare `nn.Parameter` collections.
code
python · 20 linesimport torch.nn as nn
class Broken(nn.Module):
def __init__(self):
super().__init__()
self.layers = [nn.Linear(4, 4) for _ in range(3)]
class Fixed(nn.Module):
def __init__(self):
super().__init__()
self.layers = nn.ModuleList(nn.Linear(4, 4) for _ in range(3))
def forward(self, x):
for layer in self.layers:
x = layer(x)
return x
print(len(list(Broken().parameters()))) # 0
print(len(list(Fixed().parameters()))) # 6
print(list(Fixed().state_dict().keys())[:2]) # ['layers.0.weight', 'layers.0.bias']go deeper
Remember the rule: layers you want PyTorch to see go in nn.ModuleList or nn.ModuleDict, never in a plain Python list, tuple or dict.
Explain the mechanism — nn.Module registers submodules on attribute assignment into _modules, and parameters(), to(), state_dict() and eval() all walk that dict. Name ModuleList versus Sequential correctly.
Treat it as a silent-training-failure class. Show how you catch it: parameter counts, state_dict keys, and the device-mismatch error that is the loudest of the four symptoms.
Own the naming consequence. Container choice fixes checkpoint key names across the team's model zoo, so ModuleList indices versus ModuleDict keys is a compatibility decision, not a style preference.
## How registration actually happens `nn.Module` overrides attribute assignment. When you write `self.fc = nn.Linear(4, 4)`, the override notices the value is an `nn.Module` and stores it in `self._modules["fc"]`. When you write `self.w = nn.Parameter(...)`, it stores it in `self._parameters["w"]`. Every recursive operation on a model — `parameters()`, `named_parameters()`, `buffers()`, `state_dict()`, `to()`, `train()`, `eval()`, `apply()`, `modules()` — walks those dicts. A Python `list` is neither a module nor a parameter, so it lands in the plain instance `__dict__`. The `nn.Linear` objects inside it are perfectly functional — you can call them, they hold weights — but the parent module has no idea they exist. ## The four symptoms 1. **No optimization.** `model.parameters()` is empty for those layers. `optim.Adam(model.parameters())` builds a parameter group that omits them. The forward pass still runs, loss still decreases (other layers compensate), and the buried layers stay at their initialization forever. 2. **No device movement.** `model.cuda()` moves what it can find. The listed layers stay on the CPU, and the first forward pass raises a device-mismatch `RuntimeError` — this is usually how people discover the bug, and it is the friendliest of the four symptoms because it is loud. 3. **Incomplete checkpoints.** `state_dict()` has no keys for them. You save, restore, and get a model whose middle is randomly initialized. `load_state_dict(strict=True)` does not complain because the missing weights are missing from both sides. 4. **Mode flags do not propagate.** `model.eval()` walks `_modules`, so listed submodules stay in training mode. Any dropout or normalisation inside them keeps behaving as if training. ## The containers PyTorch provides four registering containers: - **`nn.ModuleList`** — an ordered, indexable, iterable list of modules. It registers each child under its index (`layers.0`, `layers.1`, …). It has **no** `forward()`: you iterate it yourself in your own `forward`. That is the point — it is storage, not composition. - **`nn.ModuleDict`** — the same for string keys, preserving insertion order. Useful for named branches (`self.heads = nn.ModuleDict({"cls": ..., "reg": ...})`). - **`nn.ParameterList` / `nn.ParameterDict`** — the equivalents for collections of bare `nn.Parameter` objects, which have the identical registration problem when kept in a list. ## ModuleList versus Sequential They are easy to confuse. `nn.Sequential` **is** callable: it chains its children, feeding each output into the next, and you never write a `forward` for it. `nn.ModuleList` is not callable — calling it raises `NotImplementedError`. Choose `Sequential` when the data flow is a straight chain; choose `ModuleList` when your `forward` needs control: skip connections, a loop that reuses layer outputs, conditional branches, or collecting every layer's output. ``` for layer in self.layers: # ModuleList: you control the flow x = x + layer(x) # e.g. a residual chain ``` ## Naming and checkpoint keys Registration also decides the `state_dict` key names. A `ModuleList` named `blocks` produces `blocks.0.weight`, `blocks.1.weight`, and so on. A `ModuleDict` named `heads` with key `cls` produces `heads.cls.weight`. Those strings are the checkpoint contract, so converting a `ModuleList` to a `ModuleDict` — or reordering entries — renames keys and breaks loading old checkpoints. ## Diagnosing it in one line ``` print(sum(p.numel() for p in model.parameters())) ``` If the parameter count is implausibly small for the architecture you think you built, something is not registered. Follow up with `list(model.state_dict().keys())` and look for the layers you expected. A zero-parameter model that still trains is the clearest possible signal that a container is a plain list. ## A related trap Storing a module inside a tuple, a set, a dataclass field, or as a value in a plain `dict` fails identically. Only assignment of an `nn.Module` directly to an attribute, or membership in one of the registering containers, gets you into `_modules`. If you truly need a reference to a module you do **not** want registered — for example to avoid duplicate parameters when tying weights — that is one of the rare, deliberate cases for keeping it outside the module tree, and it should carry a comment saying so.
- What is the difference between nn.ModuleList and nn.Sequential?`nn.Sequential` defines a `forward` that chains its children, so you call it directly and the data flows straight through. `nn.ModuleList` has no `forward` — calling it raises `NotImplementedError`. It is registration-aware storage that you iterate inside your own `forward`, which is what you need for skip connections, loops that collect intermediate outputs, or conditional branches.
- How do you spot this bug quickly on a model you did not write?Print `sum(p.numel() for p in model.parameters())` and compare it with what the architecture should hold, then inspect `list(model.state_dict().keys())` for the layers you expect to see. Missing keys mean unregistered submodules. A model that trains while reporting far too few parameters is the giveaway.
- Does the same problem apply to bare nn.Parameter objects kept in a list?Yes, identically — a list of `nn.Parameter` is registered nowhere, so none of them are optimized or moved. Use `nn.ParameterList` or `nn.ParameterDict`, which register each entry under an index or key and appear normally in `parameters()` and `state_dict()`.
- How does using a ModuleList affect checkpoint key names?Children are registered under their index, so a `ModuleList` attribute named `blocks` produces keys like `blocks.0.weight`. Those strings are the checkpoint contract: reordering entries, or switching to a `ModuleDict` with string keys, renames every key and makes old checkpoints fail to load.
saying these in an interview costs you the question
- Says a Python list works because forward still runs
- Thinks nn.ModuleList can be called like Sequential
- Believes .to(device) walks every attribute
- Expects load_state_dict to error on missing layers
- Uses a plain dict for named branches instead of ModuleDict