Contents
Map

02 · Prog Langs

Training Loop From Scratch

View as:

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.

Learning objectives 35 min
By the end of this page you will be able to:
  • 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:

Layertrain() modeeval() mode
nn.DropoutRandomly zeroes activations at rate pPasses all activations through unchanged (identity)
nn.BatchNorm*Uses current batch mean/variance; updates running statisticsUses 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 and model.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() changes Dropout (becomes identity) and BatchNorm (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

Check yourself
0 / 6 answered
  1. Which two layer types change behaviour between model.train() and model.eval()?
  2. Training loss keeps falling while validation loss rises after epoch 5. What is the most likely explanation?
  3. What's the single most common bug caused by forgetting model.eval()?
  4. Why is optimizer.zero_grad() called before backward() and not after step()?
  5. Why wrap the validation loop in torch.no_grad() even though you're not calling backward() there?
  6. How do you compute an epoch-level average loss correctly when the last batch is smaller than the others?

Exercises

Exercise - Add early stopping and a scheduler

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.

Exercise - Find the bug

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

Last reviewed: 2026-09

⚡AI-assisted content - always verify, always explore multiple perspectives·