Lesson 7

The training loop

Introduction

A training loop repeats five actions for every batch: forward pass, loss, gradient reset, backward pass, and parameter update. The loss falls from 0.75 to 0.00 in three epochs on a small dataset.

Learning goal

Run the five actions of one training step, then evaluate a model with predictions, softmax probabilities, and an accuracy function.

Before you start

Tensors, autograd, modules, cross-entropy loss, and a basic Python loop.

Lesson plan

  1. Follow one training step through forward, loss, zero_grad, backward, and step.
  2. Read the printed loss across three epochs and explain every line of the loop.
  3. Turn logits into labels and measure accuracy with a reusable function.

The five steps, and a memory aid

One training step has five actions. The order matters. A small change of order can break the update. Use the phrase Fast Loss = Zero Belly Size to keep the order in mind.

WordLetterStep
FastFforward
LossLloss
ZeroZzero_grad
BellyBbackward
SizeSstep

Read the table from top to bottom. Each word gives one letter, and each letter names one function in the loop.

All the parts together

Earlier lessons built the tensor library, autograd, the module API, and the data loaders. This lesson combines the parts. The model then trains on a small dataset.

The example uses a model named NeuralNetwork with two inputs and two outputs. The model and the train_loader come from the earlier lessons.

The complete loop

Read the loop from top to bottom. Each of the five steps carries a mnemonic label in a comment.

import torch.nn.functional as F


# --- setup ---
torch.manual_seed(123)
model = NeuralNetwork(num_inputs=2, num_outputs=2)
optimizer = torch.optim.SGD(model.parameters(), lr=0.5)

num_epochs = 3

for epoch in range(num_epochs):                 # repeat for every epoch

    model.train()                                # set the training mode
    for batch_idx, (features, labels) in enumerate(train_loader):   # repeat for every batch

        # Fast Loss = Zero Belly Size  ->  F L Z B S
        logits = model(features)                    # 1 · F · forward
        loss = F.cross_entropy(logits, labels)  # 2 · L · loss

        optimizer.zero_grad()                       # 3 · Z · zero_grad
        loss.backward()                            # 4 · B · backward
        optimizer.step()                            # 5 · S · step

        ### LOGGING
        print(f"Epoch: {epoch+1:03d}/{num_epochs:03d}"
              f" | Batch {batch_idx:03d}/{len(train_loader):03d}"
              f" | Train/Val Loss: {loss:.2f}")

    model.eval()                                 # set the evaluation mode
    # Optional model evaluation

The prints produce this output.

Epoch: 001/003 | Batch 000/002 | Train/Val Loss: 0.75
Epoch: 001/003 | Batch 001/002 | Train/Val Loss: 0.65
Epoch: 002/003 | Batch 000/002 | Train/Val Loss: 0.44
Epoch: 002/003 | Batch 001/002 | Train/Val Loss: 0.13
Epoch: 003/003 | Batch 000/002 | Train/Val Loss: 0.03
Epoch: 003/003 | Batch 001/002 | Train/Val Loss: 0.00

The loss falls from 0.75 to 0.00 across three epochs. The fall is a sign that the model converges on the training set.

Details of the loop

Reset the gradients first. Call optimizer.zero_grad() in every step. Without the call, the new gradients add to the old gradients. PyTorch does not reset them for you. The behavior stays in PyTorch because some users want gradient accumulation.

Hyperparameters and the validation set

A learning rate is a hyperparameter. Tune the rate and watch the loss. Choose a rate that makes the loss fall after some epochs. The number of epochs is another hyperparameter. A validation dataset helps you find good settings.

The validation set is like the test set. Use the test set one time only, to avoid a biased evaluation. Use the validation set many times to tune the settings.

Predictions after training

Call model.eval() and disable gradient recording. The model then returns one raw score per class for each input.

model.eval()

with torch.no_grad():
    outputs = model(X_train)

print(outputs)
tensor([[ 2.8569, -4.1618],
        [ 2.5382, -3.7548],
        [ 2.0944, -3.1820],
        [-1.4814,  1.4816],
        [-1.7176,  1.7342]])

The five rows are the five training examples. Each row holds two raw scores, called logits. A larger score means more support for that class.

Softmax probabilities

Softmax turns the logits into probabilities that sum to one.

torch.set_printoptions(sci_mode=False)
probas = torch.softmax(outputs, dim=1)
print(probas)
tensor([[    0.9991,     0.0009],
        [    0.9982,     0.0018],
        [    0.9949,     0.0051],
        [    0.0491,     0.9509],
        [    0.0307,     0.9693]])

Look at the first row. The example has a 0.9991 probability of class 0 and a 0.0009 probability of class 1. set_printoptions makes the numbers easy to read.

Labels with argmax

argmax returns the index of the highest value. With dim=1, it returns the highest value in each row. With dim=0, it returns the highest value in each column.

predictions = torch.argmax(probas, dim=1)
print(predictions)   # tensor([0, 0, 0, 1, 1])

You do not need the softmax values for the labels. Apply argmax to the logits directly. The result is the same.

predictions = torch.argmax(outputs, dim=1)
print(predictions)         # tensor([0, 0, 0, 1, 1])
print(predictions == y_train)  # all True
torch.sum(predictions == y_train)  # 5

The dataset has 5 training examples. All 5 predictions are correct. 5 of 5 is 100 percent prediction accuracy.

An accuracy function

The function below works for a dataset of any size. It reads one batch at a time and counts the correct predictions.

def compute_accuracy(model, dataloader):

    model.eval()
    correct = 0.0
    total_examples = 0

    for idx, (features, labels) in enumerate(dataloader):

        with torch.no_grad():
            logits = model(features)

        predictions = torch.argmax(logits, dim=1)
        compare = labels == predictions
        correct += torch.sum(compare)
        total_examples += len(compare)

    return (correct / total_examples).item()

The function uses the same chunk size as the training batch. For a large dataset, memory limits the number of examples at one time. This design keeps the function usable for any dataset size.

compute_accuracy(model, train_loader)   # 1.0
compute_accuracy(model, test_loader)    # 1.0

Both values are 1.0. The model predicts every training example and every test example correctly.

Common pitfalls

  • Call zero_grad() in every step. Without the call, old gradients add to new gradients.
  • Call backward() before step(). The optimizer needs new gradients.
  • Use model.eval() and torch.no_grad() for evaluation. This choice saves memory and gives stable results.
  • A falling training loss does not prove good generalization. Also measure a validation set.
  • Very large learning rates can make the loss jump or become nan.

Try it

A fresh model needs one training step. Write the five lines in the correct order. Add the mnemonic label to each line. Then say which line changes the model parameters.

Reveal the worked answer
logits = model(features)                    # 1 · F · forward
loss = F.cross_entropy(logits, labels)      # 2 · L · loss
optimizer.zero_grad()                       # 3 · Z · zero_grad
loss.backward()                             # 4 · B · backward
optimizer.step()                            # 5 · S · step

The order is forward, loss, zero_grad, backward, step. The phrase “Fast Loss = Zero Belly Size” gives the same order. Only optimizer.step() changes the model parameters.

Recap

One training step runs five actions: forward, loss, zero_grad, backward, and step. The phrase “Fast Loss = Zero Belly Size” keeps the order. Reset the gradients in every step. Use model.eval() and torch.no_grad() for evaluation.

After training, softmax turns logits into probabilities. argmax turns logits into labels. An accuracy function then measures the trained model on each dataset.

Reference: PyTorch optim documentation.