Optimisers, Schedulers and the Training Loop
PyTorch Lightning
Lightning splits training into the part only you can write — model, loss, optimiser — and the machinery everyone rewrites badly, which its Trainer runs for you.
- 7 min read
- 3 reading levels
- Published
Read these first
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
PyTorch Lightning is a library that runs the repetitive machinery of training, so you write only the parts unique to your model.
Think of a tiffin kitchen you can rent. The gas, the washing station, the serving line, the cleanup crew — all provided and reliable. You walk in with the one thing nobody can bring for you: your recipe.
Plain PyTorch is cooking at home: total freedom, and you also scrub every pot. Lightning is the rented kitchen: your recipe, their machinery.
Why it exists
Every training script contains the same dozen chores — moving data to the right device, looping over epochs, calling backward, saving checkpoints, validating, logging. Each chore is short. Each is also a place to plant a quiet bug. The previous lessons in this section met several: the forgotten eval(), the un-copied best model, the mis-averaged metric.
Lightning's bet: those chores are identical across projects, so write them once, test them hard, and let everyone reuse them. What remains yours is the recipe — which model, which loss, which optimiser.
How it works
you write: Lightning runs:
┌──────────────────────┐ ┌────────────────────────────┐
│ the model │ │ epoch and batch loops │
│ what one step means │ ──→ │ device moves, backward, step│
│ which optimiser │ │ checkpoints, logging │
└──────────────────────┘ │ early stopping, multi-GPU │
└────────────────────────────┘Your part goes in a class; Lightning's part is called the Trainer. One fit call connects them.
Where you have seen this
This split is how most mature tools work: you bring the unique part, machinery handles the rest. Website builders run the servers while you design pages. Lightning brought that deal to model training, and a large share of research code on GitHub now uses it.
Remember this
- Lightning separates your recipe from repeated machinery.
- The machinery bugs from earlier lessons are handled for you, tested.
- Underneath, it is ordinary PyTorch — nothing is hidden or replaced.
What to learn next
- Experiment tracking — where
self.logcalls become comparable runs. - DataLoader workers and speed — the input pipeline Lightning keeps warning you about.
- Reading a shape error — the debugging section, for when any framework's stack trace lands in your lap.
Developer — Code and libraries.
Setup
pip install torch pytorch-lightningHonest size note: the pytorch-lightning wheel is small (a few MB), but it pulls helper dependencies (fsspec, tqdm, PyYAML, aiohttp, lightning-utilities); expect a few tens of MB of new packages on a machine that already has torch. Outputs captured with pytorch-lightning 2.6.5, torch 2.5.1, CPU.
A complete training run
import torch
import torch.nn as nn
import pytorch_lightning as pl
from torch.utils.data import DataLoader, TensorDataset
torch.manual_seed(0)
X = torch.randn(256, 10)
y = (2 * X[:, 0] - X[:, 3] > 0).long()
loader = DataLoader(TensorDataset(X, y), batch_size=32)
class Classifier(pl.LightningModule):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Linear(10, 32), nn.ReLU(), nn.Linear(32, 2))
self.loss_fn = nn.CrossEntropyLoss()
def training_step(self, batch, batch_idx):
xb, yb = batch # no .to(device): Lightning moved it already
loss = self.loss_fn(self.net(xb), yb)
self.log("train_loss", loss, prog_bar=True)
return loss # Lightning calls backward and step for you
def configure_optimizers(self):
return torch.optim.AdamW(self.parameters(), lr=1e-3)
trainer = pl.Trainer(max_epochs=3, accelerator="cpu", logger=False,
enable_checkpointing=False, enable_progress_bar=False)
model = Classifier()
trainer.fit(model, loader)
print(f"final train_loss: {trainer.callback_metrics['train_loss'].item():.4f}")GPU available: True (cuda), used: False TPU available: False, using: 0 TPU cores ┌───┬─────────┬──────────────────┬────────┬───────┬───────┐ │ │ Name │ Type │ Params │ Mode │ FLOPs │ ├───┼─────────┼──────────────────┼────────┼───────┼───────┤ │ 0 │ net │ Sequential │ 418 │ train │ 0 │ │ 1 │ loss_fn │ CrossEntropyLoss │ 0 │ train │ 0 │ └───┴─────────┴──────────────────┴────────┴───────┴───────┘ Trainable params: 418 Non-trainable params: 0 Total params: 418 Total estimated model params size (MB): 0.002 Modules in train mode: 5 Modules in eval mode: 0 Total FLOPs: 0 `Trainer.fit` stopped: `max_epochs=3` reached. final train_loss: 0.6070
The hardware lines at the top reflect the capture machine and will read differently on yours. Lightning also prints advisory hints (DataLoader worker counts, logger suggestions) that are trimmed here, and the summary-table format shifts between Lightning versions. The parameter count and the final loss reproduce.
Walkthrough
LightningModule — a nn.Module with opinions. You fill in named methods: training_step defines what one batch means, configure_optimizers returns the optimiser. Optional siblings — validation_step, test_step, predict_step — follow the same shape.
training_step returns the loss, and that is all. No zero_grad, no backward, no step — the Trainer calls them in the right order around your method. Writing them yourself here is a bug, not extra safety.
No .to(device) anywhere. The Trainer owns placement. The same file runs on CPU, one GPU (accelerator="gpu"), or several (devices=4) without edits — the entire class of device mismatch errors is designed away.
self.log(...) — streams metrics to whatever logger is attached (TensorBoard by default; disabled above for a clean demo). With prog_bar=True it also lives on the progress bar.
The chores as callbacks — the previous lesson in two lines:
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint
trainer = pl.Trainer(max_epochs=100,
callbacks=[EarlyStopping(monitor="val_loss", patience=5),
ModelCheckpoint(monitor="val_loss", save_top_k=1)])These are the tested versions of the deepcopy-and-patience code you wrote by hand — including the restore-best behaviour.
Common mistakes
Calling backward() or optimizer.step() inside training_step. Double-stepping, broken gradient accumulation, and mysterious loss curves follow. If you genuinely need manual control, set self.automatic_optimization = False and take over explicitly — half-taking-over is the bug.
Sprinkling .to(device) or .cuda() in the module. It works on your machine and breaks on multi-GPU, where each process owns a different device. Create tensors with device=self.device when you must make one inside the module.
Expecting pl.seed_everything from import lightning tutorials to match. Two package names ship the same code: pytorch_lightning (used here) and the newer lightning (imported as import lightning as L). Pick one per project; mixing both in one process causes subtle isinstance failures with callbacks.
Treating the Trainer as magic during debugging. Trainer(fast_dev_run=True) runs a single batch through train and validation — the Lightning-native version of overfit a single batch first; there is also overfit_batches=1 for the full ritual.
Try it yourself
Add a validation_step (same shape as training_step, logging "val_loss"), pass a second DataLoader to fit, and attach the EarlyStopping callback. You have now rebuilt lessons 6 and 8 of this section in about ten lines.
What to learn next
- Experiment tracking — where
self.logcalls become comparable runs. - DataLoader workers and speed — the input pipeline Lightning keeps warning you about.
- Reading a shape error — the debugging section, for when any framework's stack trace lands in your lap.
Researcher — Mathematics and papers.
What the abstraction buys, precisely
The Trainer inserts well-defined hooks around your methods; the run loop is roughly:
for epoch: for batch:
batch = transfer_to_device(batch)
loss = module.training_step(batch)
precision_plugin.backward(loss) # fp32 / fp16+GradScaler / bf16
clip_gradients(); optimizer.step(); zero_grad()
[scheduler.step() per its interval]Precision policy (precision="16-mixed" etc.), gradient accumulation (accumulate_grad_batches), clipping (gradient_clip_val), and distribution (strategy="ddp", FSDP, DeepSpeed) are orthogonal flags on this loop rather than code edits — the practical argument for the framework is that these features compose correctly, which hand-rolled loops historically get wrong in the interactions (e.g. clipping unscaled fp16 gradients; see gradient clipping).
Overhead
The hook dispatch adds per-step Python overhead measured in tens of microseconds — noise for real models, measurable for micro-benchmarks with sub-millisecond steps. For serious throughput analysis, profile with the trainer's profiler (profiler="advanced") rather than assuming; data loading dominates far more often, as covered in is the GPU waiting for data.
Lineage and alternatives
- Falcon et al. (2019–), PyTorch Lightning — the framework itself; the design paper equivalent is the documentation's "Lightning philosophy".
- Hugging Face Accelerate — a thinner deal: keep your own loop, let the library own device placement and distribution. Preferred when the loop itself is the research.
- Lightning Fabric — the same trade at a lower level from the Lightning team;
fabric.setup(model, optimizer)and your loop survives. - torchtnt (Meta) — a loop framework without the module conventions.
The honest decision rule: exotic training dynamics — meta-learning with higher-order gradients, GAN-style alternating updates beyond the built-in support, RL interaction loops — fight the abstraction; measurement-heavy but structurally standard training benefits most from it.
Reproducibility surface
pl.seed_everything(42, workers=True) seeds Python, NumPy and torch, and propagates seeds into DataLoader workers — a packaged version of the checklist in random seeds and reproducibility. Checkpoints serialise model, optimiser, scheduler, loop counters and RNG states, making Trainer(..., ckpt_path=...) resumes bit-honest in a way ad-hoc torch.save(model.state_dict()) resumes are not.
What to learn next
- Experiment tracking — where
self.logcalls become comparable runs. - DataLoader workers and speed — the input pipeline Lightning keeps warning you about.
- Reading a shape error — the debugging section, for when any framework's stack trace lands in your lap.