PyTorch Training Loop

Unlike high-level frameworks that hide training behind fit(), PyTorch usually has you write the training loop yourself: fetch a batch, run the model forward, compute the loss, backpropagate, and step the optimizer, repeated for many epochs with periodic validation. It's a few dozen lines, and writing it explicitly gives you full control over mixed precision, gradient accumulation, custom losses, and logging.

Most training bugs hide in this loop: forgetting model.eval() during validation, not zeroing gradients, data on the wrong device, or a learning rate schedule stepped at the wrong time. This page walks through a production-shaped loop and the techniques that make training faster and more stable.

TL;DR

Quick Example

Core Concepts

nn.Module

Models subclass nn.Module: layers assigned as attributes in __init__ are registered, so their parameters appear in model.parameters() and move with model.to(device), and forward defines computation. Modules nest. state_dict() and load_state_dict() save and restore weights. Use nn.ModuleList or nn.ModuleDict for dynamic collections of layers, since plain Python lists aren't registered.

Loss Functions and Optimizers

Optimizers: AdamW (the default for most deep learning, with decoupled weight decay), SGD with momentum (common in vision), and newer ones like Lion or Muon for specific setups. Exclude biases and normalization weights from weight decay via parameter groups for best results.

Train vs Eval Mode

model.train() enables dropout, and uses batch statistics in batch norm. model.eval() disables dropout and uses running statistics. Forgetting eval() during validation or inference gives noisy, lower-quality results. Mode is independent of gradient tracking, so also use torch.no_grad() or inference_mode() for evaluation.

Learning Rate Schedules

Learning rate is the most important hyperparameter. Common schedules: linear warmup (a few hundred to a few thousand steps) followed by cosine or linear decay, OneCycleLR, or ReduceLROnPlateau (stepped with the validation metric). Note whether a scheduler expects step() per batch or per epoch.

Mixed Precision

torch.autocast runs eligible operations in lower precision (bf16 or fp16) while keeping sensitive ones in fp32, typically giving 1.5–3× speedups and lower memory use on modern GPUs.

Stability and Scale Techniques

For multi-GPU training, see distributed training.

Tracking and Reproducibility

Best Practices

Overfit a Single Batch First

Before full training, confirm the model can drive the loss near zero on one small batch. If it can't, there's a bug in the model, loss, or data pipeline, and finding it now saves hours.

Validate Regularly on Held-Out Data

Track validation metrics every epoch (or every N steps for long runs), keep the best checkpoint, and stop early when validation stops improving. Training loss alone hides overfitting. See model evaluation.

Keep the GPU Busy

Use pin_memory=True, non_blocking=True transfers, enough DataLoader workers, and larger batches where memory allows. Profile if GPU utilization is low. The data pipeline is often the bottleneck.

Save state_dicts, Not Whole Models

torch.save(model.state_dict()) is portable across code changes; pickling entire models ties checkpoints to exact class definitions and file paths. Load with weights_only=True, which is safer given pickle's risks. See insecure deserialization.

Common Mistakes

Forgetting model.eval()

Validation with dropout active and batch norm in training mode gives worse, noisy metrics, and at inference time it gives inconsistent predictions.

Softmax Before CrossEntropyLoss

Stepping the Scheduler Incorrectly

Calling a per-step scheduler once per epoch (or vice versa) produces a completely different learning rate curve than intended. Log the LR to verify the schedule.

FAQ

Why doesn't PyTorch have a fit() method?

PyTorch favors explicit, flexible code: custom training loops make it easy to implement new techniques, multiple losses, or unusual update rules. If you want fit()-style convenience, libraries like PyTorch Lightning, Keras 3 (with a PyTorch backend), and Hugging Face's Trainer provide it on top of PyTorch.

What learning rate should I use?

For AdamW on many deep learning tasks, 1e-4 to 1e-3 is a common starting range, and fine-tuning pretrained models often uses 1e-5 to 5e-5. Use warmup plus decay, run a quick learning rate range test, and tune based on validation results.

Should I use bf16 or fp16 mixed precision?

bf16 when your hardware supports it (NVIDIA Ampere and newer, TPUs, recent CPUs). It's simpler (no loss scaling) and more stable. Use fp16 with GradScaler on older GPUs that lack efficient bf16.

Does torch.compile always speed things up?

Often, but not always. Gains depend on model architecture, hardware, and dynamic shapes. Compilation adds startup time and can hit unsupported patterns (graph breaks). Benchmark with and without it, and use mode="max-autotune" for longer runs where compile time is amortized.

Related Topics

References