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
- Define models as
nn.Modulesubclasses; parameters are registered automatically. - Standard step:
zero_grad→ forward → loss →backward→ (clip) →optimizer.step(). - Switch modes with
model.train()andmodel.eval()(they affect dropout and batch norm), and validate undertorch.no_grad(). - AdamW plus a scheduler (warmup then cosine or linear decay) is a strong default.
- Use mixed precision (
torch.autocast, bf16 or fp16) for speed and memory, and gradient clipping for stability. - Checkpoint model, optimizer, scheduler, and step, and track validation metrics for early stopping. Try
torch.compilefor extra speed.
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.
- bf16 has fp32's range, needs no loss scaling, and is preferred on Ampere-and-newer GPUs and TPUs.
- fp16 needs a
torch.amp.GradScalerto avoid gradient underflow.
Stability and Scale Techniques
- Gradient clipping (
clip_grad_norm_) prevents exploding gradients, which matters especially for transformers and RNNs. - Gradient accumulation: call
backward()on several micro-batches (dividing the loss by the number of micro-batches) before oneoptimizer.step(), to simulate larger batches on limited memory. - Activation checkpointing trades compute for memory in deep models.
torch.compile(PyTorch 2.x) traces and optimizes the model into fused kernels, often giving significant speedups with one line.
For multi-GPU training, see distributed training.
Tracking and Reproducibility
- Log training and validation loss, metrics, learning rate, gradient norms, and throughput to TensorBoard, Weights & Biases, or MLflow.
- Seed randomness (
torch.manual_seed, NumPy, Pythonrandom, DataLoader workers) for reproducible experiments, and considertorch.use_deterministic_algorithms(True)when exact reproducibility matters. - Checkpoint everything needed to resume: model, optimizer, scheduler, scaler, epoch and step, and RNG states.
- Higher-level libraries (PyTorch Lightning, Hugging Face Accelerate, and the Trainer) package these patterns. See MLOps.
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
- PyTorch — The framework overview
- PyTorch Tensors & Autograd — What backward() does
- PyTorch Datasets & DataLoaders — Feeding the loop efficiently
- PyTorch Distributed Training — Scaling to many GPUs
- Model Evaluation — Validation metrics and methodology
- Fine-Tuning — Training pretrained models further