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
- Follow one training step through forward, loss, zero_grad, backward, and step.
- Read the printed loss across three epochs and explain every line of the loop.
- 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.
| Word | Letter | Step |
|---|---|---|
| Fast | F | forward |
| Loss | L | loss |
| Zero | Z | zero_grad |
| Belly | B | backward |
| Size | S | step |
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
- The model has 2 inputs and 2 outputs. The numbers match the toy dataset: 2 input features and 2 class labels.
- The optimizer is stochastic gradient descent, SGD, with a learning rate of 0.5.
model.train()andmodel.eval()set the mode. Some parts, such as dropout and batch normalization, behave differently in each mode. The toy model has neither part. Keep both calls to avoid surprises when the model changes.- Pass the logits to
cross_entropy. The function applies softmax internally. The step is fast and numerically stable. loss.backward()computes the gradients in the graph that PyTorch built in the background.optimizer.step()uses the gradients to update the model parameters. For SGD, the step multiplies each gradient by the learning rate and subtracts the result from the parameter.
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()beforestep(). The optimizer needs new gradients. - Use
model.eval()andtorch.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.