nn.Sequential is wonderful for about a week. Then you need a skip connection, or two heads sharing a backbone, or a branch that only runs at inference, and a straight pipe of layers can't express any of it. That's when you write your own nn.Module — and it turns out the base class you've been leaning on this whole time has a few sharp edges worth knowing about before you cut yourself.
This part is nn.Module in full: what it tracks, what it doesn't, and the handful of bugs that account for most lost afternoons.
Anatomy of a Module
A Module is two methods. __init__ declares the pieces — layers, parameters, constants — and assigns them to self. forward describes how a tensor flows through those pieces. You never call forward directly; you call the module (model(x)), which runs hooks and then forward. Here's a residual block, the pattern at the heart of every ResNet and Transformer:
import torch.nn as nnclass ResidualBlock(nn.Module):def __init__(self, dim):super().__init__() # never skip this lineself.fc1 = nn.Linear(dim, dim)self.fc2 = nn.Linear(dim, dim)self.norm = nn.LayerNorm(dim)self.act = nn.GELU()def forward(self, x):h = self.act(self.fc1(x))h = self.fc2(h)return self.norm(x + h) # the '+ x' is the skip connection
Parameters, buffers, and the state dict
A Module holds three kinds of tensor, and knowing which is which explains a lot of otherwise-baffling behaviour.
| Kind | Trained by the optimizer? | In state_dict? | Example |
|---|---|---|---|
nn.Parameter | Yes | Yes | a Linear layer's weight |
| Registered buffer | No | Yes | BatchNorm's running_mean |
| Plain attribute tensor | No | No | a constant you forgot to register |
model = ResidualBlock(64)for name, p in model.named_parameters():print(name, tuple(p.shape), p.requires_grad)# fc1.weight (64, 64) True# fc1.bias (64,) True# ...print(model.state_dict().keys())# odict_keys(['fc1.weight', 'fc1.bias', 'fc2.weight', 'fc2.bias',# 'norm.weight', 'norm.bias'])
Composing modules: the three containers
nn.Sequential— a fixed chain, output of each feeds the next. Use it for genuinely linear sub-parts.nn.ModuleList— a list that does register its contents. Use it when depth is a hyperparameter and you loop over layers inforward.nn.ModuleDict— the same, keyed by name, for branches you select at runtime.
import torch.nn as nnclass DeepMLP(nn.Module):def __init__(self, dim, depth):super().__init__()# A plain [ResidualBlock(dim) for _ in range(depth)] would NOT register.self.blocks = nn.ModuleList(ResidualBlock(dim) for _ in range(depth))self.head = nn.Linear(dim, 10)def forward(self, x):for block in self.blocks:x = block(x)return self.head(x)
Because it's residual, you can stack these deep without the gradient vanishing. Parameter count grows linearly with depth and quadratically with width — worth having a feel for when you're choosing a model size against a memory budget:
params ≈ 795·h + 10
Doubling the hidden width roughly doubles the parameter count here. In a Transformer, where width appears in several matrices at once, the same doubling costs closer to 4×.
| width | trainable parameters |
|---|---|
| h = 32 | 25450 |
| h = 64 | 50890 |
| h = 128 | 101770 |
| h = 256 | 203530 |
| h = 512 | 407050 |
Weight initialization is not a detail
nn.Linear initialises its weights with a reasonable default (Kaiming uniform). For most work that's fine. When it isn't — a custom layer, a paper that specifies an init, a training run that won't get moving — you override it with model.apply(), which walks every submodule:
import torch.nn as nndef init_weights(m):if isinstance(m, nn.Linear):nn.init.kaiming_normal_(m.weight, nonlinearity="relu")if m.bias is not None:nn.init.zeros_(m.bias)model = DeepMLP(64, depth=6)model.apply(init_weights) # applied recursively to every submodule
train() vs eval(): the mode switch that bites
Some layers behave differently during training and inference. Dropout is active in train() mode and a no-op in eval(). BatchNorm uses the current batch's statistics in train() and its accumulated running statistics in eval(). The flag is a single boolean on every submodule, flipped recursively by model.train() / model.eval().
Saving and loading — the state dict, not the object
Save model.state_dict() — an ordered dictionary of tensors — not the model object itself. Pickling the object ties the file to your exact class definition and directory layout, and it breaks the first time you refactor. The state dict is just data.
import torchtorch.save(model.state_dict(), "model.pt")# Later — reconstruct the architecture first, then load the weights.model = DeepMLP(64, depth=6)model.load_state_dict(torch.load("model.pt", map_location="cpu", weights_only=True))model.eval()
__init__declares,forwardconnects, and you call the module — neverforwarddirectly.- Layers must be attributes (or live in a
ModuleList/ModuleDict) to be registered and trained. - Parameters are learned; buffers persist but aren't; register buffers explicitly.
model.train()/model.eval()change what Dropout and BatchNorm do — set them deliberately.- Save the
state_dict, rebuild the architecture, thenload_state_dict.
References
- [1]nn.Module — API reference · PyTorch documentation
- [2]Modules — design and behaviour notes · PyTorch documentationRegistration, hooks, train/eval, and buffers explained in depth.
- [3]Saving and Loading Models · PyTorch tutorials
- [4]torch.nn.init — initialization functions · PyTorch documentation
- [5]Delving Deep into Rectifiers (Kaiming initialization) · He et al., ICCV 2015
- [6]Deep Residual Learning for Image Recognition · He et al., CVPR 2016Where the skip connection in the block above comes from.


