Training Loop From Scratch
The One-Line Definition
The canonical PyTorch training loop is a hand-written for epoch / for batch nested loop that runs forward pass → compute loss → backward() → optimizer.step() → zero_grad() for every batch, with an explicit model.eval() pass over held-out data to track generalization.
- Write the canonical training and validation loop and explain each step
- Explain exactly what model.train() and model.eval() change
- Compute correct epoch-level metrics when the last batch is smaller
- Detect overfitting from train and validation curves
There's no hidden "train" button in PyTorch - you write the training loop yourself. That's intentional: it forces you to see exactly what happens on each pass through the data, which makes it much easier to debug when something goes wrong (loss not decreasing, model not learning, etc.). Once you've written it by hand a few times, wrapper libraries like Trainer classes make a lot more sense because you know what they're doing under the hood.
Unlike higher-level frameworks that hide the loop behind a .fit() call, PyTorch's philosophy is "the training loop is just Python" - full control over logging, gradient clipping, custom schedulers, mixed precision, and multi-loss objectives without fighting a framework abstraction. Every production training script (including transformers.Trainer internally) is a variation on the same five-step loop.
flowchart TD
Start(["🚀 Start Epoch"]) --> TM["model.train()"]
TM --> B{"For each\ntraining batch"}
B --> F["➡️ Forward pass"]
F --> L["📉 Compute loss"]
L --> ZG["🧹 zero_grad()"]
ZG --> BW["⬅️ backward()"]
BW --> OS["🔧 optimizer.step()"]
OS --> B
B -->|"epoch done"| EM["model.eval()"]
EM --> V{"For each\nvalidation batch"}
V --> VF["➡️ Forward pass\n(no_grad)"]
VF --> VL["📊 Track val loss/acc"]
VL --> V
V -->|"done"| Next(["🔁 Next epoch"])
style TM fill:#d8dfe8,stroke:#b0bac8
style F fill:#e8e0d4,stroke:#c8b89a
style L fill:#e8e0d4,stroke:#c8b89a
style ZG fill:#dde4dc,stroke:#b0c4b0
style BW fill:#dde4dc,stroke:#b0c4b0
style OS fill:#dde4dc,stroke:#b0c4b0
style EM fill:#ddd8e4,stroke:#b8b0c8
style VF fill:#ddd8e4,stroke:#b8b0c8
style VL fill:#ddd8e4,stroke:#b8b0c8
The Canonical Loop
Every epoch (one full pass over the training data) has two phases: a training phase, where the model sees examples, makes predictions, gets corrected, and updates its internal settings; and a validation phase, where the model is tested on data it hasn't been updated with, to check whether it's actually learning general patterns rather than just memorizing the training examples.
import torch
import torch.nn as nn
import torch.optim as optim
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = MyModel().to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
num_epochs = 10
for epoch in range(num_epochs):
# ---- Training phase ----
model.train() # enables dropout, uses batch stats for BatchNorm
running_loss = 0.0
for batch_x, batch_y in train_loader:
batch_x, batch_y = batch_x.to(device), batch_y.to(device)
optimizer.zero_grad() # clear stale gradients
outputs = model(batch_x) # forward pass
loss = loss_fn(outputs, batch_y)
loss.backward() # backward pass
optimizer.step() # update weights
running_loss += loss.item() * batch_x.size(0)
train_loss = running_loss / len(train_loader.dataset)
# ---- Validation phase ----
model.eval() # disables dropout, uses running stats for BatchNorm
val_loss, correct = 0.0, 0
with torch.no_grad(): # no graph needed - saves memory and compute
for batch_x, batch_y in val_loader:
batch_x, batch_y = batch_x.to(device), batch_y.to(device)
outputs = model(batch_x)
loss = loss_fn(outputs, batch_y)
val_loss += loss.item() * batch_x.size(0)
correct += (outputs.argmax(dim=1) == batch_y).sum().item()
val_loss /= len(val_loader.dataset)
val_acc = correct / len(val_loader.dataset)
print(f"Epoch {epoch+1}/{num_epochs} | train_loss={train_loss:.4f} "
f"val_loss={val_loss:.4f} val_acc={val_acc:.4f}")
model.train() vs model.eval()
Some layers behave differently depending on whether the model is learning or being tested. Dropout randomly "turns off" parts of the network during training to prevent over-reliance on any single piece - but you want the full network active when actually using it. Batch Normalization uses statistics from the current batch during training, but switches to stable, pre-computed averages during evaluation so predictions don't depend on what else happens to be in the same batch. Forgetting to switch modes is one of the most common PyTorch bugs.
model.train() and model.eval() toggle the .training flag on every submodule (recursively), which changes behavior for exactly two layer types that most architectures include:
| Layer | train() mode | eval() mode |
|---|---|---|
nn.Dropout | Randomly zeroes activations at rate p | Passes all activations through unchanged (identity) |
nn.BatchNorm* | Uses current batch mean/variance; updates running statistics | Uses stored running mean/variance from training, ignoring the current batch |
Forgetting model.eval() before validation/inference causes silently wrong results (dropout still randomly zeroing activations, BatchNorm using unstable single-batch statistics) - not a crash, just degraded and non-reproducible metrics, which makes this bug easy to miss. Always pair it with torch.no_grad() for inference to also skip graph construction.
Study Notes
- The loop is always: forward → loss →
zero_grad()→backward()→step(), repeated per batch, nested inside a per-epoch loop - Call
model.train()before the training phase andmodel.eval()before the validation/inference phase, every epoch - Wrap validation/inference in
torch.no_grad()to skip autograd graph construction and save memory model.eval()changesDropout(becomes identity) andBatchNorm(uses running stats instead of batch stats) - the two layers most commonly responsible for train/eval mismatch bugs- Track both train and validation loss per epoch to catch overfitting (train loss keeps falling, val loss starts rising)
Check Yourself
- Which two layer types change behaviour between model.train() and model.eval()?
- Training loss keeps falling while validation loss rises after epoch 5. What is the most likely explanation?
- What's the single most common bug caused by forgetting
model.eval()? - Why is
optimizer.zero_grad()called beforebackward()and not afterstep()? - Why wrap the validation loop in
torch.no_grad()even though you're not callingbackward()there? - How do you compute an epoch-level average loss correctly when the last batch is smaller than the others?
Exercises
Extend the canonical loop with a learning-rate scheduler (cosine or step), gradient clipping, and early stopping that keeps the best checkpoint by validation loss with a patience of 3 epochs.
Solution
Call scheduler.step() once per epoch (or per step for warmup/cosine-per-step schedules), clip with torch.nn.utils.clip_grad_norm_ after backward and before optimizer.step(), and track best_val_loss with a counter that resets on improvement; stop when the counter reaches the patience and reload the best state_dict.
A colleague's loop computes accuracy = correct / len(train_loader) and never calls model.eval(). List every consequence you can predict, then fix both.
Solution
Accuracy is divided by the number of batches, so it is inflated by roughly the batch size; validation runs with dropout on and BatchNorm using batch statistics, so metrics are noisy and pessimistic. Divide by len(loader.dataset) and wrap validation in model.eval() plus torch.no_grad().
References
- PyTorch documentation, torch.nn.Module (train/eval) (2026)
- PyTorch tutorials, Optimizing Model Parameters (2026)
Last reviewed: 2026-09