Building Models with nn.Module
nn.Module is the base class for all PyTorch models. Subclass it to define layers, implement forward(), and register parameters that optimizers update during training.
Search across all documentation pages
nn.Module is the base class for all PyTorch models. Subclass it to define layers, implement forward(), and register parameters that optimizers update during training.
Quick-reference recipe card - copy-paste ready.
import torch.nn as nn
class MLP(nn.Module):
def __init__(self, in_dim: int, hidden: int, out_dim: int):
super().__init__()
self.net = nn.Sequential(
nn.Linear(in_dim, hidden),
nn.ReLU(),
nn.Linear(hidden, out_dim),
)
def forward(self, x):
return self.net(x)When to reach for this:
state_dict()."""building_models.py - custom CNN for image classification."""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
class SmallCNN(nn.Module):
def __init__(self, num_classes: int = 10):
super().__init__()
self.conv1 = nn.Conv2d(1, 32, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
self.pool = nn.MaxPool2d(2)
self.fc1 = nn.Linear(64 * 7 * 7, 128)
self.fc2 = nn.Linear(128, num_classes)
self.dropout = nn.Dropout(0.25)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.pool(F.relu(self.conv1(x))) # 28->14
x = self.pool(F.relu(self.conv2(x))) # 14->7
x = x.flatten(1)
x = self.dropout(F.relu(self.fc1(x)))
return self.fc2(x)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = SmallCNN(num_classes=10).to(device)
x = torch.randn(8, 1, 28, 28, device=device)
logits = model(x)
print("logits shape:", logits.shape) # (8, 10)
total_params = sum(p.numel() for p in model.parameters())
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"params: {total_params:,} trainable: {trainable:,}")What this demonstrates:
__init__ become part of model.parameters().forward defines the computation; call via model(x), not model.forward(x)..to(device) moves all parameters and buffers to GPU.flatten(1) preserves batch dimension while collapsing spatial dims.__init__ registers submodules (nn.Linear, nn.Conv2d) and parameters (nn.Parameter).forward runs the computation graph; hooks can intercept inputs/outputs.nn.Module tracks training vs eval mode for dropout and batch norm behavior.state_dict() returns a dict of parameter tensors for save/load.named_modules() and named_children() traverse the module tree.| Pattern | Use When | Example |
|---|---|---|
nn.Sequential | Linear chain of layers | MLP classifiers |
Subclass nn.Module | Custom forward logic | ResNet skip connections |
nn.ModuleList | Variable number of layers | Dynamic depth networks |
nn.ModuleDict | Named submodules | Multi-head architectures |
# Freeze backbone, train head only
for param in model.conv1.parameters():
param.requires_grad = False
# Different learning rates per group
optimizer = torch.optim.Adam([
{"params": model.conv1.parameters(), "lr": 1e-5},
{"params": model.fc2.parameters(), "lr": 1e-3},
])forward() directly - skips hooks and nn.Module wrappers. Fix: always call model(x).super().__init__() - submodules not registered. Fix: call super().__init__() first in __init__.model.eval() for inference - dropout and batch norm behave incorrectly. Fix: model.eval() before inference; model.train() for training.__init__.model.to(device) and x.to(device).model.eval() or nn.GroupNorm for small batches.| Alternative | Use When | Don't Use When |
|---|---|---|
Subclass nn.Module | Custom architectures | Trivial 3-layer MLP (use Sequential) |
nn.Sequential | Simple feedforward stacks | Need branching or skip connections |
PyTorch Lightning LightningModule | Structured training loops | Learning PyTorch fundamentals |
torch.nn.functional | Stateless operations in forward | Need learnable parameters |
nn.Module classes hold learnable parameters (Linear, Conv2d).nn.functional provides stateless functions (relu, conv2d with explicit weights).__init__; use functional for one-off ops.sum(p.numel() for p in model.parameters() if p.requires_grad)numel() returns total elements in a tensor.state_dict but not trained.self.register_buffer("name", tensor).nn.Module can contain other modules as attributes.model.children() returns direct children; model.modules() is recursive.model.eval() for validation and inference.def init_weights(m):
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
model.apply(init_weights)register_forward_hook inspects intermediate activations.print(model)
# or
from torchinfo import summary
summary(model, input_size=(8, 1, 28, 28))torchinfo shows parameter counts per layer.def forward(self, x: torch.Tensor) -> torch.Tensor.def forward(self, x):
residual = x
out = self.block(x)
return F.relu(out + residual).to(device) - partial moves cause device errors.Stack versions: This page was written for Python 3.14.0 (stable 3.14, maintenance 3.13), FastAPI 0.115+, Django 5.2, Flask 3.1, Pydantic 2, PyTorch 2.6+, pandas 2.2+, Polars 1.x, ruff 0.9+, and uv 0.6+.
Reviewed by Chris St. John·Last updated Jul 19, 2026