Building Models with nn.Module
model.train() and model.eval()
One flag switches layers like Dropout and BatchNorm between practice behaviour and exam behaviour — and forgetting to flip it is the most common evaluation bug in PyTorch.
- 7 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
model.train() and model.eval() flip one switch that tells every layer whether it is practising or performing.
A cricketer practises with a weighted bat and a bowling machine set to awkward angles. The handicaps are deliberate — practice is made harder so match day feels easier. On match day, the handicaps come off. Same player, same skills, different behaviour.
Some PyTorch layers work exactly like this. During training they deliberately make life harder, and during evaluation they behave cleanly. The switch that tells them which situation they are in is the mode.
Why it exists
The best-known practice handicap is dropout — a layer that randomly silences a fraction of the network on every training pass. It forces the network not to over-rely on any one internal pathway, which fights overfitting.
But silencing at exam time would be sabotage. Predictions would change on every run of the same input.
Another layer, BatchNorm, keeps notes during training and reads them back at exam time — the next lesson is all about it. Both layers need to know which situation they are in. Neither can guess. So the model carries a flag, and you set it.
How it works
model.train() model.eval()
| |
practice mode exam mode
| |
dropout: ON dropout: OFF
batchnorm: writes batchnorm: reads
| |
outputs vary a bit same input → same outputThe flag does nothing by itself. Each layer reads it and chooses its own behaviour.
A real example you have seen
When a photo app tags your face, that model runs in exam mode — steady and repeatable. The messy, randomised practice behaviour existed only during training, months earlier, on someone's servers.
Remember this
model.train()= practice mode.model.eval()= exam mode.- Dropout and BatchNorm behave differently in each — most other layers ignore the flag.
- Evaluating in the wrong mode gives noisy, wrong-looking results with no error message.
What to learn next
- BatchNorm and its running statistics — the other layer that reads this flag, with higher stakes.
- Buffers vs parameters — where BatchNorm keeps the notes it writes in train mode.
- Overfitting and underfitting — the disease dropout exists to treat.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on CPU. The train-mode numbers below depend on the seed and torch version; the pattern — varying then constant — is the point.
See the switch flip
import torch
from torch import nn
torch.manual_seed(0)
model = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Dropout(p=0.5), nn.Linear(8, 1))
x = torch.ones(1, 4)
model.train() # this is the default state
print("training flag:", model.training)
print("three train-mode passes on the SAME input:")
for _ in range(3):
print(f" {model(x).item():.4f}")
model.eval()
print("training flag:", model.training)
print("three eval-mode passes on the SAME input:")
for _ in range(3):
print(f" {model(x).item():.4f}")training flag: True three train-mode passes on the SAME input: -0.1923 -0.5477 -0.1100 training flag: False three eval-mode passes on the SAME input: -0.3621 -0.3621 -0.3621
Same input, three different answers in train mode, one steady answer in eval mode. That is dropout doing its job, and then standing down.
What the flag actually is
model.train() sets self.training = True on the model and every registered submodule, using the same recursive walk as everything else in the register. model.eval() sets it to False. Nothing else happens. Each layer's forward reads its own self.training and branches.
Layers that care: Dropout (all variants), BatchNorm (all variants), and a few relatives like InstanceNorm with tracked stats. Linear, Conv2d, ReLU, LSTM without dropout — indifferent.
eval() does not stop gradients
This trips almost everyone once. The mode flag and gradient tracking are independent switches:
model.eval()
with torch.no_grad(): # THIS is what stops gradient bookkeeping
preds = model(x)eval() changes layer behaviour. torch.no_grad() stops autograd from recording, which saves memory and time. Proper evaluation uses both. Using only no_grad leaves dropout on; using only eval wastes memory building a graph nobody will use.
The evaluation loop that belongs in your muscle memory:
model.eval()
with torch.no_grad():
for x, y in val_loader:
...
model.train() # back to practice before the next epochCommon mistakes
Forgetting model.eval() before validation. Validation loss looks worse and jumps around, because dropout is still deleting neurons. Hours have been lost tuning learning rates to fix "noisy validation" that was a missing mode switch.
Forgetting model.train() after validation. The sneakier twin. Training continues in eval mode: dropout never fires again and BatchNorm stops updating. The model still trains, but not the way you designed.
Assuming eval() makes runs deterministic overall. It stops dropout randomness. Other sources — shuffled data, random augmentation — are separate.
Calling eval() once at startup and never train(). Then dropout never fires at all, and your regularisation exists only in the architecture diagram.
Try it yourself
Set p=0.9 in the Dropout layer and rerun. Watch train-mode outputs swing much harder while eval-mode outputs stay steady — then reason about what p=0.9 does to the effective width of the network during practice.
What to learn next
- BatchNorm and its running statistics — the other layer that reads this flag, with higher stakes.
- Buffers vs parameters — where BatchNorm keeps the notes it writes in train mode.
- Overfitting and underfitting — the disease dropout exists to treat.
Researcher — Mathematics and papers.
Dropout's two halves
Dropout (Srivastava et al., 2014, Dropout: A Simple Way to Prevent Neural Networks from Overfitting) samples a mask $m \sim \mathrm{Bernoulli}(1-p)$ per element per forward pass. PyTorch implements inverted dropout: in train mode the kept activations are scaled by $\frac{1}{1-p}$, so that
$$ \mathbb{E}[\tilde{h}] = \frac{(1-p) \cdot h}{1-p} = h $$
Where $h$ is the activation, $\tilde h$ its dropped-and-rescaled version, and $p$ the drop probability. Because the expectation is preserved at train time, eval mode is the identity — no rescaling at inference, which is why the switch is behavioural rather than arithmetic.
The classical interpretation: training samples from an ensemble of $2^n$ thinned subnetworks sharing weights, and eval approximates the ensemble's geometric mean with a single pass. The approximation is exact for linear layers, heuristic otherwise.
The three-switch matrix
Mode, gradient recording, and inference optimisation are orthogonal:
| Switch | Controls | Scope |
|---|---|---|
train()/eval() | layer behaviour (dropout, BN) | per-module flag |
torch.no_grad() | autograd graph construction | context manager |
torch.inference_mode() | graph + version-counter bookkeeping | context manager |
inference_mode (stable since torch 1.10) is no_grad plus the promise that produced tensors will never enter autograd, allowing PyTorch to skip view/version tracking. Measurably faster for small-batch serving; tensors created inside cannot later be used in a graph, which is the trade.
Mode as an evaluation-protocol problem
Monte Carlo dropout (Gal and Ghahramani, 2016, Dropout as a Bayesian Approximation) deliberately runs inference in train mode: $T$ stochastic passes yield a predictive distribution whose variance estimates model uncertainty. The mode flag becomes an epistemic instrument — a reminder that "eval mode" encodes a protocol choice, not a law.
Mode also interacts with fine-tuning: freezing a backbone while leaving it in train mode lets BatchNorm running statistics drift toward the new dataset even though no parameter updates — covered in depth in BatchNorm and transfer learning. Getting the flag right per-submodule, not per-model, is the difference between the two standard fine-tuning recipes.
Reading
- Srivastava et al. (2014) — dropout, and the ensemble interpretation.
- Gal and Ghahramani (2016) — MC dropout and uncertainty.
- PyTorch autograd notes (
docs.pytorch.org/docs/stable/notes/autograd.html) — the precise semantics ofno_gradvsinference_mode.
What to learn next
- BatchNorm and its running statistics — the other layer that reads this flag, with higher stakes.
- Buffers vs parameters — where BatchNorm keeps the notes it writes in train mode.
- Overfitting and underfitting — the disease dropout exists to treat.