Phase 23 of 25 · Topic 23.5

Model Training Loop, Loss Functions & Optimizers

1Concept

A standard PyTorch training loop iterates through epochs, computing forward predictions, evaluating the loss function (`nn.CrossEntropyLoss`), running backpropagation (`loss.backward()`), and updating weights (`optimizer.step()`).

2Architecture Diagram

Epoch Loop: [ Forward Pass ] -> [ Compute Loss ] -> [ zero_grad() ] -> [ backward() ] -> [ step() ]

3Code Example

Python 3.12
training_loop_code = '''
import torch
import torch.nn as nn
import torch.optim as optim

model = nn.Linear(4, 2)
criterion = nn.CrossEntropyLoss()
optimizer = optim.AdamW(model.parameters(), lr=0.001)

# Single training step demonstration
inputs = torch.randn(8, 4)
labels = torch.randint(0, 2, (8,))

optimizer.zero_grad()           # 1. Clear previous gradients
outputs = model(inputs)         # 2. Forward pass
loss = criterion(outputs, labels) # 3. Compute loss
loss.backward()                 # 4. Backward pass
optimizer.step()                # 5. Update weights

print(f"Step Loss: {loss.item():.4f}")
'''
print("=== PyTorch Training Step Pipeline ===")
print(training_loop_code.strip())

4Expected Output

=== PyTorch Training Step Pipeline ===
import torch
import torch.nn as nn
import torch.optim as optim

model = nn.Linear(4, 2)
criterion = nn.CrossEntropyLoss()
optimizer = optim.AdamW(model.parameters(), lr=0.001)

# Single training step demonstration
inputs = torch.randn(8, 4)
labels = torch.randint(0, 2, (8,))

optimizer.zero_grad()           # 1. Clear previous gradients
outputs = model(inputs)         # 2. Forward pass
loss = criterion(outputs, labels) # 3. Compute loss
loss.backward()                 # 4. Backward pass
optimizer.step()                # 5. Update weights

print(f"Step Loss: {loss.item():.4f}")

5Key Takeaways

  • AdamW is the industry standard optimizer with decoupled weight decay regularization.
  • Always invoke `optimizer.zero_grad()` before `loss.backward()`.
  • Use learning rate schedulers (`optim.lr_scheduler`) for dynamic learning rate decay.