This is part one of three. It covers everything you need to train real models with PyTorch Lightning, starting from zero Lightning knowledge. By the end you can turn a plain PyTorch model into a Lightning one, train and validate it, log metrics, save and resume checkpoints, switch from CPU to GPU by changing one argument, and read the errors you are most likely to meet. Mid-level and Senior take the same topics further (multi-GPU, configuration files, production); nothing here is thrown away.
You should already know the basics of PyTorch: what a tensor is, what an nn.Module is, and roughly what a training loop does. You do not need to be good at it. Lightning exists precisely because the training loop is where most people's PyTorch code gets long, repetitive and fragile.
Each section ends with a Try it task. Do them as you go. They take a few minutes each, and the ideas only stick once you have watched your own model train, stop, and resume.
The guide was checked against Lightning 2.6.6, released on 10 September 2026. Where an older tutorial shows something different, a later section tells you which form is current.
What Lightning is, and the problem it solves
PyTorch Lightning is a Python library that sits on top of PyTorch and takes over the engineering part of training a model, so that your code only contains the research part. The research part is the model, the loss and the optimizer. The engineering part is everything around them: looping over epochs and batches, moving tensors to the right device, turning gradients on and off, calling backward, saving checkpoints, writing logs, using several GPUs, and using mixed precision.
The diagram is the whole idea. You write the recipe once. The Trainer cooks it on whatever kitchen you give it.
Why does this matter? Because the engineering code is the part that breaks. A hand-written loop that works on your laptop CPU typically needs rewriting when you move to a GPU, rewriting again for two GPUs, and again when you add mixed precision. Each rewrite is a chance to forget optimizer.zero_grad(), to leave the model in training mode during validation, or to log from every process at once. Lightning's promise is that the same LightningModule runs unchanged on a CPU, one GPU, many GPUs or many machines. Changing hardware means changing Trainer arguments, not model code.
Two things Lightning is not are worth stating early, because beginners often assume otherwise.
It is not a different framework from PyTorch. A LightningModule is a subclass of torch.nn.Module. Your layers, losses, optimizers and tensors are all plain PyTorch. If you know how to write a PyTorch model, you already know how to write the interesting half of a Lightning one.
It is not a service. There is no server, daemon or dashboard to run. Lightning is a library you import. You run a Python script, the script uses the Trainer, and the Trainer writes files (logs and checkpoints) to disk. That simplicity is why it fits so easily into the rest of an MLOps stack: a training job is just a script that a scheduler or CI system can launch.
What people use it for:
Research that scales
Prototype on a laptop CPU, then run the identical model on a GPU server without rewriting the loop.
Reproducible training
A fixed structure for steps, logging and checkpoints means two people's projects look alike and can be reviewed alike.
Safe long runs
Built-in checkpointing and resuming means a crashed job continues instead of starting from zero.
Tracking out of the box
One self.log call sends a metric to TensorBoard, CSV files, MLflow, Weights & Biases and others.
- Think of a PyTorch training script you have written or read. Highlight, line by line, which lines are about the model and which are about running it (devices, loops, saving, printing).
- Count each group. The running lines are what Lightning takes over.
What came before: the hand-written loop
To see what Lightning removes, look at what it replaces. Here is a typical plain-PyTorch training loop for a classifier, with a validation pass and a device move:
import torch
import torch.nn.functional as F
device = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
for epoch in range(5):
model.train()
for x, y in train_loader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
loss = F.cross_entropy(model(x), y)
loss.backward()
optimizer.step()
model.eval()
val_loss = 0.0
with torch.no_grad():
for x, y in val_loader:
x, y = x.to(device), y.to(device)
val_loss += F.cross_entropy(model(x), y).item()
print(f"epoch {epoch} val_loss {val_loss / len(val_loader):.4f}")
torch.save(model.state_dict(), "model.pt")
This is fine for a first experiment. Now list what it does not do. It has no progress bar. It does not log to a file. It does not save the best checkpoint, only the last. It cannot resume, because the optimizer state is not saved. It does not stop early. It uses one device. And every feature you add (mixed precision, gradient clipping, several GPUs) adds lines to the middle of the loop, tangled with your research code.
Before Lightning, people handled this in one of two ways. Some copied a personal "training utilities" file from project to project, which slowly became an undocumented in-house framework. Others adopted a heavyweight framework that hid the loop completely and made customisation painful. Lightning takes a middle path: it owns the loop, but it exposes hooks, which are methods you can override at specific moments, so you keep control where it matters.
The same task in Lightning has no loop in your code at all. You write what happens for one batch, and the Trainer repeats it. The rest of this guide shows how.
lightning package also contains Fabric, a smaller tool for people who want to keep their own loop but still get device and distribution handling. This series teaches the Trainer. If you ever need a fully custom loop, Fabric is the current answer; custom loop objects inside the Trainer were removed long ago.
- In the plain loop above, find the three lines that implement one optimization step:
zero_grad,backwardandstep. - Keep these in mind. In a Lightning
training_stepyou will not write them; the Trainer will.
The mental model: four nouns
Lightning has a large API, but you can navigate nearly all of it with four nouns. Learn these and every page of the documentation becomes easier to place.
LightningModule. This is your model plus its training recipe, in one class. It holds the layers (in __init__), the forward pass (forward), what to do with one training batch (training_step), what to do with one validation batch (validation_step), and which optimizer to use (configure_optimizers). It is the research code.
Trainer. This is the engine. You create it with arguments that describe how to run (how many epochs, which hardware, which precision) and then call trainer.fit(model, ...). It owns the loops, the device placement, the checkpointing and the logging. It is the engineering code.
DataModule or DataLoaders. This is where the data lives. The simplest option is to pass ordinary PyTorch DataLoader objects straight to the Trainer. When your data handling grows (downloading, splitting, transforms), you package it into a LightningDataModule, a class that bundles the data steps so they can be reused and shared.
Callback. This is a small object with hook methods like on_train_start or on_train_batch_end, passed to the Trainer in a list. It holds reusable, non-essential behaviour: saving the best model, stopping early, monitoring the learning rate. Callbacks keep that behaviour out of your module.
layers, steps, optimizer
or plain DataLoaders
epochs, hardware, precision
checkpoint, early stop
lightning_logs/
.ckpt files
The direction of control is the thing to internalise. In plain PyTorch you call the model. In Lightning the Trainer calls your methods, at the right moments, in the right order. This is sometimes called inversion of control, and it explains a common beginner confusion: you never write for batch in loader and you never call training_step yourself.
A few more words you will meet in the first hour:
- Epoch: one full pass over the training data.
- Step or batch: one iteration of the loop, over one batch of data.
- Hook: a method with a fixed name that the Trainer calls at a fixed moment.
training_stepis the most important hook. - Checkpoint: a
.ckptfile holding everything needed to resume training (the weights, the optimizer state, the epoch counter and more). - Logger: the adapter that sends metrics somewhere you can look at them.
- Accelerator: the hardware type (CPU, GPU and so on).
- Without looking back, write one sentence each for the four nouns: LightningModule, Trainer, DataModule and Callback.
- Then decide which of the four owns each of these: the loss function, the number of epochs, the train/validation split, saving the best checkpoint.
Installing and checking the setup
Lightning needs Python 3.10 or newer. Python 3.9 support was removed in release 2.6.1. It also needs PyTorch, which pip installs for you if it is missing. Use a virtual environment so the project's packages stay separate from the rest of your machine.
On Linux or macOS:
python3.12 -m venv .venv
source .venv/bin/activate
python -m pip install lightning
On Windows (PowerShell):
py -3.12 -m venv .venv
.venv\Scripts\activate
pip install lightning
There are two package names on PyPI. lightning is the recommended one; you import it as import lightning as L. pytorch-lightning is the same Trainer code under the older pytorch_lightning namespace. Both are at 2.6.6, and either works, but pick one per project.
pytorch_lightning.LightningModule and you pass it to lightning.Trainer, you get a confusing type error saying the model must be a LightningModule, even though it looks like one. The two packages define separate classes. Use import lightning as L everywhere in a project, including in every tutorial you copy from, and rewrite any pytorch_lightning imports you find.
If you need a GPU build of PyTorch, install PyTorch first using the selector at pytorch.org/get-started, which gives you the right command for your CUDA version, and then install Lightning. If you install Lightning first, pip picks a default PyTorch wheel for your platform, and on some systems that is a CPU-only build. You would then see the GPU "not available" error described in the troubleshooting section.
Apple Silicon Macs use the GPU through the mps accelerator. accelerator="auto" finds it for you. Multi-device training is not available on MPS, which is fine for learning. Some PyTorch operations are not implemented for MPS yet, and PyTorch has an environment variable, PYTORCH_ENABLE_MPS_FALLBACK=1, that makes those operations run on the CPU instead of failing.
Windows works, with two differences. Distributed runs use the gloo backend instead of NCCL, and you must wrap the entry point of your script in if __name__ == "__main__":. That guard is good practice everywhere, and it is mandatory on Windows and macOS whenever DataLoader workers start new processes. WSL2 is a practical option if you want a Linux-like experience.
For the examples in this guide, also install torchvision, which provides the MNIST dataset:
pip install torchvision
An optional extras bundle adds the pretty progress bar, command-line tooling and plotting helpers:
pip install "lightning[extra]"
Now verify the installation. This one-liner prints the Lightning version, the PyTorch version, and whether a CUDA GPU is visible:
python -c "import lightning as L, torch; print(L.__version__, torch.__version__, torch.cuda.is_available())"
You should see something like 2.6.6 2.10.0 False on a CPU-only laptop. False is fine for this guide. On a Mac, check MPS instead:
python -c "import torch; print(torch.backends.mps.is_available())"
Then run a smoke test with a model that Lightning ships for exactly this purpose:
import lightning as L
from lightning.pytorch.demos.boring_classes import BoringModel
L.Trainer(fast_dev_run=True).fit(BoringModel())
fast_dev_run=True runs one training batch and one validation batch and then stops, with logging and checkpointing disabled. If this finishes without an error, your installation works end to end. You will use fast_dev_run constantly; it is the fastest way to find out whether a change broke your code.
lightning==2.6.6 (together with your PyTorch version) in requirements.txt so a colleague, or a CI job next month, installs what you tested. Release 2.6.6 is also the minimum you should run today, because it fixes a vulnerability in how checkpoints are loaded; the checkpoint section explains why.
- Create a virtual environment and install
lightningandtorchvision. - Run the version one-liner and note your output.
- Run the
BoringModelsmoke test and read what the Trainer prints.
Your first project, step by step
We will train a small classifier on MNIST, the dataset of 28 by 28 handwritten digits. The model is deliberately tiny, because the point is the Lightning structure, not the accuracy. We build it in four steps: the module, the data, the Trainer, and then a run.
Step 1: write the LightningModule
import torch
import torch.nn.functional as F
import lightning as L
class LitClassifier(L.LightningModule):
def __init__(self, lr: float = 1e-3):
super().__init__()
self.save_hyperparameters()
self.net = torch.nn.Sequential(
torch.nn.Flatten(),
torch.nn.Linear(28 * 28, 128),
torch.nn.ReLU(),
torch.nn.Linear(128, 10),
)
def forward(self, x):
return self.net(x)
def training_step(self, batch, batch_idx):
x, y = batch
loss = F.cross_entropy(self(x), y)
self.log("train_loss", loss, prog_bar=True)
return loss
def validation_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = F.cross_entropy(logits, y)
acc = (logits.argmax(dim=1) == y).float().mean()
self.log("val_loss", loss, prog_bar=True)
self.log("val_acc", acc, prog_bar=True)
def configure_optimizers(self):
return torch.optim.Adam(self.parameters(), lr=self.hparams.lr)
Read it top to bottom, because every line earns its place.
The class inherits from L.LightningModule, which is itself an nn.Module. That is why self.parameters() works and why you build layers in __init__ exactly as you would in PyTorch. Always call super().__init__() first.
self.save_hyperparameters() records the arguments of __init__ (here, lr) into self.hparams, and also into every checkpoint the Trainer writes. This matters later: it lets you rebuild the model from a checkpoint without remembering how you constructed it. It is a one-line habit with a large payoff.
forward defines how the model turns an input into predictions. By convention it is what you call at inference time. Here it just runs the Sequential stack. Lightning does not require forward, but writing it keeps inference clean.
training_step(batch, batch_idx) receives one batch from the training loader and the batch's index. You compute the loss and return it. Returning the loss is what tells Lightning what to differentiate. Notice what is absent: no optimizer.zero_grad(), no loss.backward(), no optimizer.step(), no .to(device). With the default automatic optimization, the Trainer does all of that.
self.log("train_loss", loss, prog_bar=True) records a metric. It sends the value to your logger and, because of prog_bar=True, shows it on the progress bar. One call, two destinations.
validation_step is the same idea for validation data. The Trainer turns gradients off and puts the model in evaluation mode for you before calling it, and switches back afterwards. You do not write model.eval() or torch.no_grad().
configure_optimizers returns the optimizer. The Trainer calls it once at the start of training. self.hparams.lr reads the learning rate that save_hyperparameters stored.
.cuda() or .to(device) inside the module
The Trainer moves the model and each batch to the right device for you. If you hard-code .cuda(), your code breaks on a CPU machine and fights the Trainer on multi-GPU. When you need to create a new tensor inside the module, use self.device, for example torch.zeros(3, device=self.device).
Step 2: prepare the data
For now we use ordinary PyTorch data loaders. The Trainer accepts them directly.
from torch.utils.data import DataLoader, random_split
from torchvision import transforms
from torchvision.datasets import MNIST
def make_loaders(batch_size: int = 64, data_dir: str = "data"):
tfm = transforms.ToTensor()
full = MNIST(data_dir, train=True, download=True, transform=tfm)
train_set, val_set = random_split(full, [55000, 5000])
train_dl = DataLoader(train_set, batch_size=batch_size, shuffle=True, num_workers=2)
val_dl = DataLoader(val_set, batch_size=batch_size, num_workers=2)
return train_dl, val_dl
The first call downloads MNIST into a data/ folder. We split the 60,000 training images into 55,000 for training and 5,000 for validation. The split matters: validation data is data the model never trains on, so its metrics tell you whether the model generalises rather than memorises. We shuffle the training loader and not the validation one.
Step 3: create the Trainer and run
import lightning as L
from model import LitClassifier
from data import make_loaders
def main():
L.seed_everything(42, workers=True)
train_dl, val_dl = make_loaders()
model = LitClassifier(lr=1e-3)
trainer = L.Trainer(max_epochs=3, accelerator="auto", devices="auto")
trainer.fit(model, train_dataloaders=train_dl, val_dataloaders=val_dl)
if __name__ == "__main__":
main()
L.seed_everything(42, workers=True) seeds Python, NumPy and PyTorch, and with workers=True also the DataLoader workers, so reruns start from the same random state. It does not make everything bit-for-bit identical on a GPU (the Trainer's deterministic=True argument pushes further), but it removes the biggest source of run-to-run variation.
The Trainer has three arguments here. max_epochs=3 stops after three passes over the data. accelerator="auto" picks a GPU if one is available and the CPU otherwise. devices="auto" uses all the devices of that type. This pair is why the same script runs on a laptop and on a GPU server unchanged.
trainer.fit(...) starts training. The if __name__ == "__main__": guard keeps the script safe for worker processes.
Step 4: run it
python train.py
On first run you will see MNIST download, then Lightning's messages, then a progress bar per epoch. Training three epochs on a CPU takes a minute or so; on a GPU it is faster. When it finishes, look at the folder: a new lightning_logs/version_0/ directory exists next to your script. That is where the default logger wrote the metrics and where the default checkpoint was saved. Run the script again and you get version_1, then version_2: each run gets its own folder, so you never overwrite a previous run.
rich package is installed, and the older tqdm bar and plain table otherwise. Screenshots in older tutorials show tqdm. The numbers are the same; only the look differs.
- Create the three files above and run
python train.py. - Find the
lightning_logs/version_0/folder and list what is inside. - Change
max_epochsto 1 and run again. Confirm that aversion_1folder appears.
Reading what the Trainer prints
Beginners often scroll past the Trainer's output. Reading it is a skill, because it tells you what Lightning decided on your behalf. A typical start looks like this (details vary by version and machine):
GPU available: False, used: False
TPU available: False, using: 0 TPU cores
| Name | Type | Params | Mode
---------------------------------------------
0 | net | Sequential | 101 K | train
---------------------------------------------
101 K Trainable params
0 Non-trainable params
101 K Total params
0.407 Total estimated model params size (MB)
Sanity Checking: | | 0/? [00:00<?, ?it/s]
Epoch 0: 100%|██████████| 860/860 [00:08<00:00, 100.2it/s, v_num=0, train_loss=0.21, val_loss=0.18, val_acc=0.95]
`Trainer.fit` stopped: `max_epochs=3` reached.
Take it line by line.
The hardware lines tell you which accelerators Lightning found and which it is using. If you expected a GPU and see GPU available: False, stop and fix that now, before waiting hours on a CPU run.
The model summary table lists each top-level submodule with its type and parameter count, and then the totals. It is the quickest sanity check that your model has the size you intended. A model with two million parameters when you expected two thousand is a bug you can spot here.
Sanity Checking is a short validation run on two batches before training begins. Its purpose is to catch bugs in validation_step immediately, instead of at the end of your first epoch. If your validation code is broken, you find out in seconds. If the sanity check bothers you, num_sanity_val_steps=0 disables it, but leave it on while you learn.
The epoch line shows the batch count (860 is 55,000 divided by 64, rounded up), the speed, and the values you marked prog_bar=True. v_num=0 is the run's version number, matching lightning_logs/version_0.
The final line states why training stopped. Here it is max_epochs=3 reached. If you later set up early stopping, this line will say so.
The default epoch cap is worth knowing. If you set neither max_epochs nor max_steps, the Trainer defaults to 1000 epochs, which is almost never what you want while experimenting. Always set one explicitly.
- Run your script again and check which hardware line the Trainer prints.
- Find the parameter count in the summary and verify it by hand: 784 times 128 plus 128, plus 128 times 10 plus 10.
- Set
num_sanity_val_steps=0and notice the Sanity Checking bar disappear.
The LightningModule in depth
You have seen the four methods that matter most. Here is how each behaves in more detail, so you can predict what the Trainer does.
What is required
To call trainer.fit, a module needs training_step and configure_optimizers, plus a training dataloader (from the module, a DataModule, or passed to fit). Everything else is optional. If you misspell one, you get a MisconfigurationException that says No training_step() method defined, which is covered in the errors section.
training_step return values
You can return the loss tensor directly, or a dictionary that contains a "loss" key. Under automatic optimization (the default), the Trainer then calls zero_grad, backward and the optimizer step. If you return None, that batch is skipped, which is occasionally useful.
forward versus training_step
forward is the model's behaviour; training_step is the training recipe. A habit worth building: put the prediction logic in forward and call self(x) from the step methods, as we did. Then inference code reuses exactly the same path as training, and you avoid subtle train-versus-serve differences.
configure_optimizers return shapes
The simplest return is a single optimizer. For a learning-rate scheduler you return a dictionary:
def configure_optimizers(self):
optimizer = torch.optim.Adam(self.parameters(), lr=self.hparams.lr)
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=0.9)
return {
"optimizer": optimizer,
"lr_scheduler": {"scheduler": scheduler, "interval": "epoch", "frequency": 1},
}
interval says whether the scheduler steps each "epoch" or each "step". For ReduceLROnPlateau you must also add "monitor": "val_loss", naming the logged metric the scheduler should watch. Returning several optimizers, as a GAN needs, requires switching to manual optimization, which is an advanced topic covered at the next level.
- Add the scheduler version of
configure_optimizersto your module and run one epoch. - Use the
LearningRateMonitorcallback later to watch the rate change, or simply printself.trainer.optimizers[0].param_groups[0]["lr"]fromon_train_epoch_end.
Logging metrics with self.log
Everything you want to track goes through self.log(name, value). It is the single most-used call after training_step, so it is worth understanding its defaults.
self.log("train_loss", loss, prog_bar=True)
self.log("val_acc", acc, prog_bar=True)
self.log_dict({"val_loss": val_loss, "val_acc": acc})
What self.log does with the value. It sends it to the logger (so it appears in TensorBoard or the CSV file) and records it in trainer.callback_metrics, a dictionary of the latest logged values. That dictionary is what callbacks such as ModelCheckpoint and EarlyStopping read when you tell them to monitor a name. This is why the name string matters so much: it must match exactly everywhere.
Step versus epoch. The defaults depend on where you call it. In training_step, the value is logged every step (on_step=True) and not aggregated per epoch. In validation_step and test_step, it is logged once per epoch (on_epoch=True), averaged over the batches. You can override either: self.log("train_loss", loss, on_step=False, on_epoch=True) gives you a smooth per-epoch curve for training too.
How often it writes. By default the Trainer writes training-step metrics every 50 steps (log_every_n_steps=50). If your dataset is tiny and you have fewer than 50 batches per epoch, you will see a warning that the interval is larger than the number of batches, and few points on your curves. Lower log_every_n_steps in the Trainer.
Logging several values. Use self.log_dict({...}) for a dictionary of metrics. Passing a dictionary to plain self.log is not supported.
Values must be single numbers. self.log expects a scalar. If you pass a tensor with more than one element you get an error suggesting value.mean().
self.log in __init__, or after calling the model by hand outside fit, raises a "not managed by the Trainer control flow" error. Logging is part of the loop's bookkeeping, so it only works while the loop is running. Logging inside predict hooks is not supported either.
Seeing the logs
The default log location is lightning_logs/version_N/. What is inside depends on what is installed. If TensorBoard is installed, the default logger is a TensorBoardLogger, and you can view curves with:
pip install tensorboard
tensorboard --logdir lightning_logs/
Open the address it prints (usually http://localhost:6006) and you get live loss and accuracy curves for every version side by side. If neither TensorBoard nor tensorboardX is installed, Lightning falls back to a CSVLogger and prints a warning saying so. The CSV file, metrics.csv, is perfectly good for a first look and loads straight into pandas.
You can choose a logger explicitly and use several at once:
from lightning.pytorch.loggers import CSVLogger, TensorBoardLogger
trainer = L.Trainer(
max_epochs=3,
logger=[TensorBoardLogger("logs/", name="mnist"), CSVLogger("logs/", name="mnist_csv")],
)
Other built-in loggers include MLFlowLogger, WandbLogger and CometLogger, which send the same self.log values to those platforms; see the MLflow and Weights & Biases guides in this series (MLflow, Weights & Biases). To turn logging off completely, pass logger=False.
- Install TensorBoard, run your training, then start
tensorboard --logdir lightning_logs/and find theval_acccurve. - Change the train loss to
on_step=False, on_epoch=Trueand compare the curve shape. - Open the CSV file from a CSV-logged run in a spreadsheet and match its columns to your
self.lognames.
Organising data: DataLoaders and the LightningDataModule
Passing two DataLoader objects to fit is fine for a single script. As soon as you want the same data setup in training, testing and a teammate's project, the setup code needs a home. That home is the LightningDataModule: a class that packages download, splitting, transforms and loaders so that data handling is as reusable as the model.
import lightning as L
from torch.utils.data import DataLoader, random_split
from torchvision import transforms
from torchvision.datasets import MNIST
class MNISTDataModule(L.LightningDataModule):
def __init__(self, data_dir: str = "data", batch_size: int = 64):
super().__init__()
self.data_dir = data_dir
self.batch_size = batch_size
self.tfm = transforms.ToTensor()
def prepare_data(self):
MNIST(self.data_dir, train=True, download=True)
MNIST(self.data_dir, train=False, download=True)
def setup(self, stage: str):
if stage == "fit":
full = MNIST(self.data_dir, train=True, transform=self.tfm)
self.train_set, self.val_set = random_split(full, [55000, 5000])
if stage == "test":
self.test_set = MNIST(self.data_dir, train=False, transform=self.tfm)
def train_dataloader(self):
return DataLoader(self.train_set, batch_size=self.batch_size, shuffle=True, num_workers=2)
def val_dataloader(self):
return DataLoader(self.val_set, batch_size=self.batch_size, num_workers=2)
def test_dataloader(self):
return DataLoader(self.test_set, batch_size=self.batch_size, num_workers=2)
You then train with trainer.fit(model, datamodule=dm). The Trainer calls the DataModule's methods in a fixed order, and the split between the first two is the most important thing to understand.
prepare_data() is for work that must happen once: downloading files, writing tokenised data to disk. When you later train on many GPUs, the Trainer runs prepare_data on only one process per machine, so several processes do not download the same file into the same place at once. Because it runs on one process only, do not assign state here (self.something = ...). Other processes would never see it.
setup(stage) is for work that every process needs: building the datasets, splitting them, applying transforms. It receives a stage string, one of "fit", "validate", "test" or "predict", so you can build only what the current phase needs. Assign self.train_set and similar here.
The *_dataloader() methods return the loaders. Each is called when its phase needs it.
Good: split the work
- Download in
prepare_data(runs once per machine) - Split and assign datasets in
setup(runs everywhere) - Build loaders in
*_dataloader
Bad: everything in __init__
- Downloads and splits happen when the object is created, even for a quick import
- State set in
prepare_datais missing on other processes - Hard to run only the test phase
A few DataLoader settings affect speed more than people expect. num_workers is how many background processes prepare batches; with zero, the main process loads data and the GPU waits. Lightning even warns you when it looks too low: a PossibleUserWarning saying the dataloader does not have many workers which may be a bottleneck. A sensible start is 2 to 4 workers, tuned later. pin_memory=True speeds up host-to-GPU copies, and persistent_workers=True keeps workers alive between epochs so you do not pay their start-up cost each time. If you set persistent_workers=True, num_workers must be above zero.
On Windows and macOS, workers start as new processes, so your script must be guarded by if __name__ == "__main__":. Without it you get an error about starting a new process before the current one has finished its bootstrap.
- Move your data code into
MNISTDataModuleand train withdatamodule=dm. - Add a
printtoprepare_dataand one tosetup, and run once. Note the order in which they appear. - Set
num_workers=0and then4, and compare the epoch time.
Validating, testing and predicting
The Trainer has four entry points, one per phase of a model's life.
| Method | Purpose | Calls your hook | Data it uses |
|---|---|---|---|
trainer.fit(...) |
Train, validating along the way | training_step, validation_step |
train and validation loaders |
trainer.validate(...) |
Run validation on its own | validation_step |
validation loader |
trainer.test(...) |
Final, one-off evaluation | test_step |
test loader |
trainer.predict(...) |
Produce outputs for new data | predict_step |
predict loader |
The distinction between validation and test is about discipline, not mechanics. You look at validation metrics many times while you tune the model, so you gradually overfit your choices to that set. The test set is the one you touch rarely, ideally once, to get an honest estimate of how the final model performs. That is why it has its own method and its own hook.
To add test support, write a test_step that mirrors the validation step:
def test_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
self.log("test_loss", F.cross_entropy(logits, y))
self.log("test_acc", (logits.argmax(dim=1) == y).float().mean())
Then evaluate:
trainer.test(model, datamodule=dm)
The Trainer prints a table of the logged test metrics and also returns them as a list of dictionaries. For prediction, define predict_step:
def predict_step(self, batch, batch_idx):
x, _ = batch
return self(x).argmax(dim=1)
and call preds = trainer.predict(model, dataloaders=some_loader). You get a list with one result per batch. Predict hooks cannot log.
A habit worth adopting: after training, evaluate with the best checkpoint rather than the last weights. Pass ckpt_path="best" to test, which works when you have a ModelCheckpoint that monitors a metric, as the next section sets up. If you call validate or test with no model and no checkpoint, the Trainer defaults to "best" and warns you.
validate to check a loaded model
After loading a checkpoint, trainer.validate(model, dataloaders=val_dl) is a fast way to confirm it still scores what it did when you saved it. It catches silent problems such as a mismatched transform.
- Add
test_stepto your module and runtrainer.test(model, datamodule=dm)after fitting. - Add
predict_step, predict on one batch, and compare the predicted digits to the labels.
Checkpoints, early stopping and resuming
A long training run you cannot resume is a liability. Lightning writes a checkpoint by default, but the default only keeps the most recent state. Almost every real project configures two callbacks: ModelCheckpoint to save good models, and EarlyStopping to stop when progress stalls.
from lightning.pytorch.callbacks import EarlyStopping, ModelCheckpoint
ckpt = ModelCheckpoint(
dirpath="checkpoints/",
filename="{epoch}-{val_loss:.3f}",
monitor="val_loss",
mode="min",
save_top_k=3,
save_last=True,
)
early = EarlyStopping(monitor="val_loss", patience=3, mode="min")
trainer = L.Trainer(max_epochs=30, callbacks=[ckpt, early])
trainer.fit(model, datamodule=dm)
print(ckpt.best_model_path)
What a checkpoint holds. A .ckpt file is not just weights. It contains the model's state_dict, the optimizer and scheduler states, the epoch and global step counters, the state of callbacks, the hyperparameters recorded by save_hyperparameters, and the Lightning version. That completeness is what makes true resuming possible.
The ModelCheckpoint arguments. monitor="val_loss" names a logged metric to rank checkpoints by. mode="min" means lower is better (use "max" for accuracy). save_top_k=3 keeps the three best and deletes the rest (-1 keeps all). save_last=True also writes a last.ckpt, the most recent state, which is what you resume from after a crash. The filename template can include logged metric names in braces, so files are named like epoch=4-val_loss=0.083.ckpt.
EarlyStopping. It watches a metric and stops training when it has not improved for patience checks. The word checks matters: patience counts validation runs, not epochs. With default settings validation runs once per epoch, so the two coincide, but they differ if you validate more or less often.
Loading a model back
model = LitClassifier.load_from_checkpoint(ckpt.best_model_path)
model.eval()
Two points trip people up. First, load_from_checkpoint is a class method that returns a new model. Calling it on an existing instance, model.load_from_checkpoint(path), builds a new model and discards it, leaving your instance unchanged. Second, you must call model.eval() yourself before inference, since you are outside the Trainer's control. Because we used save_hyperparameters(), you did not have to pass lr; it was restored from the file. You can override saved values by passing keyword arguments.
Resuming training
To continue an interrupted run with the optimizer state and counters intact, pass the checkpoint to fit:
trainer.fit(model, datamodule=dm, ckpt_path="last")
"last" and "best" are special values for checkpoints created earlier in the same session; you can also pass a file path. The old Trainer argument resume_from_checkpoint no longer exists; ckpt_path on fit replaced it. To save a checkpoint by hand at any moment, use trainer.save_checkpoint("manual.ckpt").
A security point you should know on day one
Checkpoint files are built on Python's pickle format, and loading an untrusted pickle can run arbitrary code. Treat a .ckpt file from a stranger like an executable. Since PyTorch 2.6, torch.load defaults to weights_only=True, which blocks most of that. Lightning 2.6.6, released on 10 September 2026, closed a further hole in load_from_checkpoint that could be exploited through crafted checkpoint metadata, and 2.6.5 was still affected. So: run 2.6.6 or newer, and never load checkpoints you do not trust.
You may see UnpicklingError: Weights only load failed when loading your own checkpoints. It happens when your saved hyperparameters contain custom Python objects. The cleanest fix is to keep hyperparameters to plain numbers, strings and lists. For files you fully trust, you can pass weights_only=False to load_from_checkpoint or to fit.
ModelCheckpoint(monitor="val_loss") only works if some hook calls self.log("val_loss", ...) with exactly that string, and if validation actually runs. A typo, or no validation loader, gives a MisconfigurationException saying the monitored key could not be found in the returned metrics.
- Add both callbacks, train for a few epochs, and list the
checkpoints/folder. - Load the best checkpoint with
load_from_checkpoint, calleval(), and runtrainer.validateon it. - Interrupt a run with Ctrl+C after two epochs and resume it with
ckpt_path="last". Check that the epoch counter continues.
Hardware and precision: the Trainer flags you will change most
Because the Trainer owns the hardware, switching machines is a matter of arguments. You have already used accelerator="auto" and devices="auto". Here are the values behind them.
The accelerator is the type of hardware: "cpu", "gpu", "cuda", "mps", "tpu" or "auto". "gpu" resolves to CUDA on NVIDIA machines and MPS on Apple Silicon. devices is how many: an integer, a list of indices like [0, 1], "auto", or -1 for all.
L.Trainer(accelerator="cpu") # force the CPU
L.Trainer(accelerator="gpu", devices=1) # one GPU
L.Trainer(accelerator="gpu", devices=[1]) # specifically GPU number 1
L.Trainer(accelerator="auto", devices="auto") # whatever is available
Using more than one GPU switches on a strategy, which is how the work is distributed. The default strategy="auto" picks sensibly: ordinary data-parallel training (DDP) when you ask for several devices. Multi-GPU work has extra rules (the script runs once per process, and metrics need synchronising), and it is covered properly at the mid level. For now, know that the old arguments gpus=2, tpu_cores and num_processes were removed in version 2.0. The current spelling is accelerator="gpu", devices=2.
Precision controls the number format. The default is full 32-bit floats. Mixed precision runs most of the maths in 16 bits to save memory and time:
L.Trainer(precision="bf16-mixed") # recent NVIDIA GPUs (Ampere or newer)
L.Trainer(precision="16-mixed") # older GPUs; uses a gradient scaler
Teach yourself the string forms. Older material writes precision=16; the integer still maps to "16-mixed", but the strings say what is happening. Mixed precision rarely changes accuracy noticeably and often lets you use a larger batch or finish sooner. Start with the default while you debug, and add mixed precision once the model trains correctly.
accelerator="gpu" on a machine without a working CUDA build of PyTorch fails with a MisconfigurationException saying the CUDAAccelerator cannot run because the accelerator is not available. The usual cause is a CPU-only PyTorch wheel, so check torch.cuda.is_available() first. Using accelerator="auto" avoids the error by falling back to the CPU.
A note for readers in the Gulf and Egypt who rent GPUs from a cloud provider: GPU instance availability and pricing differ by region, and some employers require data to stay in-country. The nice property of Lightning here is that none of this touches your model code. You develop on a CPU or a Mac, run on whatever region's GPU instance you are allowed to use, and change only accelerator, devices and precision.
- Run your script once with
accelerator="cpu"and once with"auto". If you have a GPU, compare epoch times. - If you have a recent NVIDIA GPU, add
precision="bf16-mixed"and watch the speed and memory use.
Debugging flags that save hours
Training runs are slow, and a bug found at epoch 40 is expensive. The Trainer has flags that shrink a run so you can find bugs quickly. They are the most underused feature among beginners.
fast_dev_run=True runs one batch of training and one of validation, and stops. It disables loggers and checkpointing, so it leaves no files behind. Pass an integer, fast_dev_run=5, for five batches. Run it every time you change the module, before launching a real run.
limit_train_batches and limit_val_batches use only a fraction (a float such as 0.1) or a fixed number (an int such as 20) of the data. This is a full epoch loop on a tenth of the data, a good middle ground between a smoke test and a real run.
overfit_batches=10 trains repeatedly on just ten batches. A healthy model should drive the training loss close to zero on ten batches. If it cannot, something is wrong with the model, the loss or the labels, and no amount of extra data will help. This is one of the best sanity checks in deep learning.
detect_anomaly=True makes PyTorch report the operation that produced a NaN or infinite gradient. It slows training a lot, so use it only while hunting such a problem.
profiler="simple" prints a table of how long each hook took, which tells you whether time goes to data loading or computation.
# a quick, cheap run while developing
trainer = L.Trainer(max_epochs=2, limit_train_batches=0.1, limit_val_batches=0.1)
# can the model learn at all?
trainer = L.Trainer(max_epochs=50, overfit_batches=10)
fast_dev_run=True (does the code run?), then overfit_batches=10 (can the model learn?), then limit_train_batches=0.1 (does it behave on a slice?), then the full run. Each step is cheaper than the next and rules out a whole class of bugs.
- Introduce a deliberate bug (misspell a variable in
validation_step) and see how quicklyfast_dev_run=Truereveals it compared with a full run. - Fix it, then run with
overfit_batches=10and watch the training loss fall.
Configuration: the Trainer arguments worth knowing
The Trainer has many arguments, and all of them are keyword-only. You will use a dozen in your first month. Here they are, grouped by what you are trying to do, with the defaults that matter.
Run length. max_epochs (set it; the default is 1000 if neither it nor max_steps is set), max_steps (stop after a number of optimizer steps), and max_time (a limit such as "00:02:00:00", meaning two hours, in days:hours:minutes:seconds form). Setting max_epochs=-1 means train until something else stops you, such as early stopping.
A realistic beginner configuration pulls several of these together:
trainer = L.Trainer(
max_epochs=20,
accelerator="auto",
devices="auto",
precision="32-true",
gradient_clip_val=1.0,
log_every_n_steps=10,
default_root_dir="runs/mnist",
callbacks=[ckpt, early],
)
You can also let a configuration file and command line drive these values using LightningCLI, which turns your module and DataModule into a command-line program with YAML configs. That is the subject of the Mid-level guide, so skip it for now.
- Set
default_root_dir="runs/mnist"and confirm the logs move there. - Add
gradient_clip_val=1.0andval_check_interval=0.5, then confirm validation now runs twice per epoch.
Common errors and how to read them
Lightning's error messages are unusually informative. They usually name both the problem and the fix. The skill is reading them calmly. Here are the ones beginners meet most.
"No training_step() method defined." The Trainer found no training_step on your class. The cause is almost always a misspelling (training_stap) or indentation that put the method outside the class. The same message appears for a missing configure_optimizers. Fix the name.
"ModelCheckpoint(monitor='val_loss') could not find the monitored key in the returned metrics." The checkpoint callback wants a metric you never logged under that exact name. Check three things: that you call self.log("val_loss", ...), that the spelling matches character for character, and that validation actually ran (you passed a validation loader and did not set limit_val_batches=0). EarlyStopping raises a closely related error with the words "conditioned on metric which is not available", with the same causes.
"CUDAAccelerator can not run on your system since the accelerator is not available." You requested a GPU that PyTorch cannot see. Run python -c "import torch; print(torch.cuda.is_available())". If it prints False, you probably have a CPU-only PyTorch build. Reinstall PyTorch from pytorch.org with the right CUDA version, or use accelerator="auto".
"self.log(name, value) was called, but the tensor must have a single element." You logged a tensor with several values. Reduce it with .mean() or log elements separately.
"You are trying to self.log() but it is not managed by the Trainer control flow." You logged outside a Trainer-driven hook, for example in __init__. Move the call into a step or epoch hook.
"Trainer(strategy='ddp') is not compatible with an interactive environment." You tried multi-GPU in a Jupyter or Colab notebook. Either run the code as a script, or use strategy="ddp_notebook", which is notebook-compatible (and not available on Windows).
"strategies from the DDP family are not supported on the MPS accelerator." On a Mac you asked for several devices. Use one device on MPS.
"Weights only load failed" (UnpicklingError). Covered in the checkpoint section: your checkpoint contains custom Python objects. Keep hyperparameters simple, or for trusted files pass weights_only=False.
"The class ... requested by the checkpoint does not resolve to an already imported subclass." This is a safety check added in 2.6.6. Import the module that defines your model class before calling load_from_checkpoint.
CUDA out of memory. The batch does not fit on the GPU. Lower the batch size first. If you still need the larger effective batch, use accumulate_grad_batches, then precision="bf16-mixed".
- Rename
training_steptotraining_stapand run. Read the message, then fix it. - Set
monitor="val_losss"inModelCheckpointand read the error. Notice that it lists the metric names that are available.
Old tutorials: what changed, and what to write instead
Most Lightning content online was written for version 1.x. Copying from it is the most common source of beginner errors. Here is a translation table for the things you will meet. Lightning 2.0 removed them, and the current spelling is on the right.
| Old (do not use) | Current |
|---|---|
Trainer(gpus=2) |
Trainer(accelerator="gpu", devices=2) |
Trainer(resume_from_checkpoint=path) |
trainer.fit(model, ckpt_path=path) |
training_epoch_end(self, outputs) |
on_train_epoch_end(self), keeping your own list |
validation_epoch_end(self, outputs) |
on_validation_epoch_end(self) |
Trainer(precision=16) |
Trainer(precision="16-mixed") |
trainer.tune(model) |
Tuner(trainer).lr_find(model) |
import pytorch_lightning as pl |
import lightning as L (or stay with one namespace) |
self.log({"a": a}) |
self.log_dict({"a": a}) |
Two more things changed more recently. Release 2.6.1 deprecated LightningModule.to_torchscript, because TorchScript itself is being phased out of PyTorch; for exporting a model, use torch.export.export() or to_onnx instead. And Python 3.9, which older machines may still run, is no longer supported.
When a snippet from the internet fails, check its date and its imports before anything else. If it imports pytorch_lightning and passes gpus=, it predates version 2.0, and the fix is the table above.
- Find a Lightning tutorial online, locate any lines from the "old" column, and rewrite them in the current form.
- Run the rewritten code with
fast_dev_run=Trueto check it.
Putting it all together
Here is one small project that uses everything from this guide: a DataModule, a module with metrics, checkpointing, early stopping, logging, resuming and a final test. It is the template to copy for your next experiment. Save it as mnist_full.py and put MNISTDataModule from earlier in datamodule.py.
import torch
import torch.nn.functional as F
import torchmetrics
import lightning as L
from lightning.pytorch.callbacks import EarlyStopping, ModelCheckpoint
from lightning.pytorch.loggers import CSVLogger
from datamodule import MNISTDataModule
class LitMNIST(L.LightningModule):
def __init__(self, lr: float = 1e-3, hidden: int = 128):
super().__init__()
self.save_hyperparameters()
self.net = torch.nn.Sequential(
torch.nn.Flatten(),
torch.nn.Linear(28 * 28, hidden),
torch.nn.ReLU(),
torch.nn.Linear(hidden, 10),
)
self.val_acc = torchmetrics.classification.MulticlassAccuracy(num_classes=10)
self.test_acc = torchmetrics.classification.MulticlassAccuracy(num_classes=10)
def forward(self, x):
return self.net(x)
def training_step(self, batch, batch_idx):
x, y = batch
loss = F.cross_entropy(self(x), y)
self.log("train_loss", loss, on_step=False, on_epoch=True, prog_bar=True)
return loss
def validation_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
self.log("val_loss", F.cross_entropy(logits, y), prog_bar=True)
self.val_acc(logits, y)
self.log("val_acc", self.val_acc, prog_bar=True)
def test_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
self.log("test_loss", F.cross_entropy(logits, y))
self.test_acc(logits, y)
self.log("test_acc", self.test_acc)
def configure_optimizers(self):
return torch.optim.Adam(self.parameters(), lr=self.hparams.lr)
def main():
L.seed_everything(42, workers=True)
dm = MNISTDataModule(batch_size=64)
model = LitMNIST(lr=1e-3, hidden=128)
ckpt = ModelCheckpoint(
dirpath="checkpoints/",
filename="{epoch}-{val_loss:.3f}",
monitor="val_loss",
mode="min",
save_top_k=2,
save_last=True,
)
early = EarlyStopping(monitor="val_loss", patience=3, mode="min")
trainer = L.Trainer(
max_epochs=20,
accelerator="auto",
devices="auto",
logger=CSVLogger("logs/", name="mnist"),
callbacks=[ckpt, early],
)
trainer.fit(model, datamodule=dm)
trainer.test(model, datamodule=dm, ckpt_path="best")
print("best checkpoint:", ckpt.best_model_path)
if __name__ == "__main__":
main()
Notice the new piece: torchmetrics. Lightning installs it as a dependency. You create a metric object as an attribute of the module (self.val_acc), call it with predictions and targets, and pass the object to self.log. Lightning then handles accumulating across batches, averaging at epoch end, resetting between epochs, and (later, on many GPUs) combining across processes. That removes a lot of hand-written bookkeeping and a classic bug where accuracy is averaged wrongly across batches of different sizes.
Run it:
python mnist_full.py
Then walk through what happened, as a checklist for any Lightning run:
- The sanity check ran two validation batches.
- Each epoch trained, validated, logged
val_lossandval_acc, and saved a checkpoint ifval_losswas in the top two. - Training either reached 20 epochs or
EarlyStoppingended it, and the last line of Trainer output says which. trainer.test(..., ckpt_path="best")loaded the best checkpoint and printed test metrics.logs/mnist/version_0/metrics.csvholds the numbers andcheckpoints/holds the models.
To reuse the trained model later without the Trainer:
from mnist_full import LitMNIST
model = LitMNIST.load_from_checkpoint("checkpoints/last.ckpt")
model.eval()
with torch.no_grad():
logits = model(batch_of_images)
Importing the class first is not optional in 2.6.6: the loader checks that the checkpoint's class is already imported, as the error section explained.
- Run the full script and read the final test accuracy.
- Change
hiddento 256 and run again; compareval_accin the two CSV files. - Load
checkpoints/last.ckptin a separate script and classify one image.
What you can now do, and what comes next
You started this guide knowing PyTorch but not Lightning. You can now write a module, feed it data, log, checkpoint, resume, switch hardware, debug cheaply and read the common errors.
What comes next. The Mid-level guide builds on this with multi-GPU training with DDP, mixed precision in depth, LightningCLI and YAML configuration, more loggers and callbacks, the Tuner, and wiring Lightning into CI and containers. The Senior guide covers large-model strategies such as FSDP and DeepSpeed, fault tolerance on schedulers, security of the checkpoint supply chain, upgrades, and where Lightning stops being the right tool.
Neighbouring guides in this series help you around Lightning. For experiment tracking beyond a CSV file, read MLflow or Weights & Biases. To package a training job so it runs the same everywhere, read Docker. For scaling out to a cluster, see Kubernetes and Ray. For large-model training libraries that work with PyTorch, see DeepSpeed, and for fine-tuning language models, see PEFT and TRL.
The best next step is not more reading. Take a small dataset of your own, write a module and DataModule for it using the template above, and make it run with fast_dev_run=True first. When you can do that without looking at this guide, you are ready for the Mid-level one.