Training Loops
A training loop repeats forward pass, loss computation, backward pass, and optimizer step across batches and epochs. Schedulers adjust the learning rate; checkpoints preserve progress for resume and deployment.
Search across all documentation pages
A training loop repeats forward pass, loss computation, backward pass, and optimizer step across batches and epochs. Schedulers adjust the learning rate; checkpoints preserve progress for resume and deployment.
Quick-reference recipe card - copy-paste ready.
for epoch in range(num_epochs):
model.train()
for x, y in train_loader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad(set_to_none=True)
loss = loss_fn(model(x), y)
loss.backward()
optimizer.step()
scheduler.step()When to reach for this:
"""training_loops.py - complete train/val loop with checkpointing."""
from __future__ import annotations
import torch
import torch.nn as nn
from torch.optim.lr_scheduler import CosineAnnealingLR
from torchvision import datasets, transforms
from torch.utils.data import DataLoader, random_split
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])
dataset = datasets.FashionMNIST(root="./data", train=True, download=True, transform=transform)
train_ds, val_ds = random_split(dataset, [50_000, 10_000])
train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=2, pin_memory=True)
val_loader = DataLoader(val_ds, batch_size=256, shuffle=False, num_workers=2)
model = nn.Sequential(nn.Flatten(), nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 10)).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
scheduler = CosineAnnealingLR(optimizer, T_max=10)
loss_fn = nn.CrossEntropyLoss()
best_val_loss = float("inf")
for epoch in range(10):
model.train()
train_loss = 0.0
for x, y in train_loader:
x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)
optimizer.zero_grad(set_to_none=True)
loss = loss_fn(model(x), y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
train_loss += loss.item() * x.size(0)
scheduler.step()
model.eval()
val_loss = 0.0
correct = 0
with torch.inference_mode():
for x, y in val_loader:
x, y = x.to(device), y.to(device)
logits = model(x)
val_loss += loss_fn(logits, y).item() * x.size(0)
correct += (logits.argmax(1) == y).sum().item()
val_loss /= len(val_ds)
acc = correct / len(val_ds)
print(f"epoch {epoch+1}: val_loss={val_loss:.4f} acc={acc:.3f} lr={scheduler.get_last_lr()[0]:.6f}")
if val_loss < best_val_loss:
best_val_loss = val_loss
torch.save({"epoch": epoch, "model": model.state_dict(), "optimizer": optimizer.state_dict()}, "best.pt")What this demonstrates:
AdamW optimizer with weight decay and cosine LR schedule.inference_mode without gradient tracking.| Optimizer | Use When | Notes |
|---|---|---|
| AdamW | Default for most models | Decoupled weight decay |
| SGD + momentum | CNNs with careful tuning | Often best with LR schedule |
| Adam | Quick prototyping | Weight decay interacts with adaptive LR |
| Lion | Research/experimentation | Memory-efficient alternative |
# Gradient accumulation for large effective batch size
accum_steps = 4
for i, (x, y) in enumerate(train_loader):
loss = loss_fn(model(x), y) / accum_steps
loss.backward()
if (i + 1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad(set_to_none=True)model.eval() before validation; model.train() before training.loss.item() * batch_size and divide by dataset size.step() per epoch for CosineAnnealingLR).x.to(device, non_blocking=True) in the loop.| Alternative | Use When | Don't Use When |
|---|---|---|
| Manual loop | Learning, full control | Repetitive boilerplate across projects |
| PyTorch Lightning | Production training structure | Debugging autograd fundamentals |
| HuggingFace Trainer | Transformer fine-tuning | Custom CNN architectures |
torch.compile | Speed up the forward/backward | Debugging training instability |
CrossEntropyLoss for multi-class (includes softmax).BCEWithLogitsLoss for multi-label binary.weight_decay=0.01 is a common starting point.ckpt = torch.load("best.pt", weights_only=True)
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
start_epoch = ckpt["epoch"] + 1CosineAnnealingLR - smooth decay, popular default.OneCycleLR - fast convergence with warmup.ReduceLROnPlateau - reduce when val metric stalls.shuffle=True in DataLoader).torch.utils.tensorboard), W&B, or MLflow for production.OneCycleLR or custom schedulers.Stack versions: This page was written for Python 3.14.0 (stable 3.14, maintenance 3.13), FastAPI 0.115+, Django 5.2, Flask 3.1, Pydantic 2, PyTorch 2.6+, pandas 2.2+, Polars 1.x, ruff 0.9+, and uv 0.6+.
Reviewed by Chris St. John·Last updated Jul 16, 2026