Coding for ML

The PyTorch Training Loop

Writing the fundamental PyTorch model training loop: zero_grad, forward pass, loss calculation, backward, and step.

🟡 intermediate5 min readcodingpytorch
The PyTorch Training Loop represents the core execution pattern for training neural networks. Unlike high level abstractions, PyTorch training requires explicitly managing data batch iteration, zeroing gradient buffers, executing model forward passes, calculating loss, invoking backward autograd, and stepping optimizer parameters.

The 5 Essential Steps of PyTorch Training

Unlike high level frameworks that hide training loops inside model.fit(), PyTorch provides explicit control over the training loop.

Every standard PyTorch training step follows 5 Mandatory Operations:

┌─────────────────────────────────────────────────────────────┐
│ 1. optimizer.zero_grad()   Clear previous gradients         │
│ 2. outputs = model(inputs) Forward pass computation         │
│ 3. loss = criterion(...)   Compute scalar loss              │
│ 4. loss.backward()         Autograd backward pass           │
│ 5. optimizer.step()        Update parameters via optimizer  │
└─────────────────────────────────────────────────────────────┘
Batch Data ──► [ 1. zero_grad ] ──► [ 2. Forward Pass ] ──► [ 3. Calculate Loss ] ──► [ 4. loss.backward ] ──► [ 5. optimizer.step ]

Standard PyTorch Training Loop Pattern

import torch
import torch.nn as nn
import torch.optim as optim

# Setup device, model, loss function, and optimizer
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = MyNeuralNetwork().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

num_epochs = 10

for epoch in range(num_epochs):
  # Set model to training mode (enables Dropout and BatchNorm tracking)
  model.train()
  running_loss = 0.0

  for batch_idx, (inputs, targets) in enumerate(train_loader):
    # Move tensors to GPU hardware
    inputs, targets = inputs.to(device), targets.to(device)

    # 1. Zero Parameter Gradients
    optimizer.zero_grad()

    # 2. Forward Pass
    outputs = model(inputs)

    # 3. Calculate Loss
    loss = criterion(outputs, targets)

    # 4. Backward Pass (Compute Gradients via Autograd)
    loss.backward()

    # 5. Optimizer Step (Update Weights)
    optimizer.step()

    running_loss += loss.item()

  # Validation Phase
  model.eval()  # Set model to evaluation mode
  val_loss = 0.0
  with torch.no_grad():  # Disable autograd engine to save memory
    for val_inputs, val_targets in val_loader:
      val_inputs, val_targets = val_inputs.to(device), val_targets.to(device)
      val_outputs = model(val_inputs)
      val_loss += criterion(val_outputs, val_targets).item()

  print(
      f'Epoch {epoch+1}/{num_epochs} | Train Loss:'
      f' {running_loss/len(train_loader):.4f} | Val Loss:'
      f' {val_loss/len(val_loader):.4f}'
  )

Critical PyTorch Rules

  1. model.train() vs model.eval(): model.eval() disables Dropout random dropping and fixes BatchNorm statistics to use running averages.
  2. torch.no_grad() in Validation: Wrapping evaluation loops inside with torch.no_grad(): disables autograd graph building, saving GPU memory and speeding up inference by 50 percent.
  3. Always Clear Gradients: Forgetting optimizer.zero_grad() causes gradients from successive batches to add together, resulting in exploding gradients.

Say this out loud

The PyTorch training loop explicitly manages data iteration, zeroing gradient buffers, forward passes, loss calculation, backward autograd, and optimizer parameter updates. Toggling model.train() and model.eval() switches Dropout and BatchNorm behaviors, while torch.no_grad() saves memory during validation loops.

Followups to expect

  1. What is Gradient Accumulation in PyTorch? Calling loss.backward() across multiple mini-batches before invoking optimizer.step() and zero_grad(), simulating larger effective batch sizes on memory constrained GPUs.
  2. What is Gradient Clipping (nn.utils.clip_grad_norm_)? Clipping parameter gradient norms to a maximum threshold before optimizer.step() to prevent exploding gradients in Recurrent Neural Networks.

Check yourself

Question 1 of 3

Why must optimizer.zero_grad() be called before computing backward pass gradients in PyTorch?

More in Coding for ML

See all →
Implement Linear Regression from Scratch5 minImplement Self-Attention from Scratch5 minNumPy Broadcasting & Vectorization5 min