Skip to content
Back to student guides
PyTorch LightningLLMsFine-tuning & training3 levels97 sectionsCovers Lightning 2.6

The Complete PyTorch Lightning Guide

Organise PyTorch training code and scale it across GPUs and nodes with Lightning. Taught at three levels — Beginner, Mid-level and Senior — each with an in-depth guide, interview prep, and practical tips.

Official docs AI-drafted · community review in progressHelp review it
19sections
31examples

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.

YOUR CODEmodel, loss, optimizer
→
LIGHTNINGMODULEthe recipe
→
TRAINERthe engine
→
ANY HARDWARECPU, GPU, many GPUs

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.

Try it
  1. 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).
  2. 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:

train_plain.py
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 and Fabric are different things The 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.
Try it
  1. In the plain loop above, find the three lines that implement one optimization step: zero_grad, backward and step.
  2. Keep these in mind. In a Lightning training_step you 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.

You write
LightningModule
layers, steps, optimizer
DataModule
or plain DataLoaders
You configure
Trainer
epochs, hardware, precision
Callbacks
checkpoint, early stop
Lightning produces
Logs
lightning_logs/
Checkpoints
.ckpt files
The Trainer calls your module's methods; your module never calls the Trainer.

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_step is the most important hook.
  • Checkpoint: a .ckpt file 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).
Try it
  1. Without looking back, write one sentence each for the four nouns: LightningModule, Trainer, DataModule and Callback.
  2. 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:

BASH
python3.12 -m venv .venv
source .venv/bin/activate
python -m pip install lightning

On Windows (PowerShell):

BASH
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.

Do not mix the two import styles If your module subclasses 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:

BASH
pip install torchvision

An optional extras bundle adds the pretty progress bar, command-line tooling and plotting helpers:

BASH
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:

BASH
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:

BASH
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:

smoke_test.py
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.

Pin the version in a requirements file Lightning does not follow semantic versioning: a minor release can remove deprecated features. Write 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.
Try it
  1. Create a virtual environment and install lightning and torchvision.
  2. Run the version one-liner and note your output.
  3. Run the BoringModel smoke 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

model.py
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.

Never call .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.

data.py
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

train.py
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

BASH
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.

The progress bar looks different on your machine Since release 2.6.0 the Trainer uses a Rich progress bar and a Rich model summary if the 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.
Try it
  1. Create the three files above and run python train.py.
  2. Find the lightning_logs/version_0/ folder and list what is inside.
  3. Change max_epochs to 1 and run again. Confirm that a version_1 folder 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):

TEXT
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.

Try it
  1. Run your script again and check which hardware line the Trainer prints.
  2. Find the parameter count in the summary and verify it by hand: 784 times 128 plus 128, plus 128 times 10 plus 10.
  3. Set num_sanity_val_steps=0 and 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:

PYTHON
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.

Try it
  1. Add the scheduler version of configure_optimizers to your module and run one epoch.
  2. Use the LearningRateMonitor callback later to watch the rate change, or simply print self.trainer.optimizers[0].param_groups[0]["lr"] from on_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.

PYTHON
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().

You can only log from inside a Trainer-driven hook Calling 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:

BASH
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:

PYTHON
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.

Try it
  1. Install TensorBoard, run your training, then start tensorboard --logdir lightning_logs/ and find the val_acc curve.
  2. Change the train loss to on_step=False, on_epoch=True and compare the curve shape.
  3. Open the CSV file from a CSV-logged run in a spreadsheet and match its columns to your self.log names.

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.

datamodule.py
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_data is 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.

Try it
  1. Move your data code into MNISTDataModule and train with datamodule=dm.
  2. Add a print to prepare_data and one to setup, and run once. Note the order in which they appear.
  3. Set num_workers=0 and then 4, 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:

PYTHON
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:

PYTHON
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:

PYTHON
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.

Use 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.
Try it
  1. Add test_step to your module and run trainer.test(model, datamodule=dm) after fitting.
  2. 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.

PYTHON
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

PYTHON
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:

PYTHON
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.

The monitored name must match what you logged 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.
Try it
  1. Add both callbacks, train for a few epochs, and list the checkpoints/ folder.
  2. Load the best checkpoint with load_from_checkpoint, call eval(), and run trainer.validate on it.
  3. 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.

PYTHON
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:

PYTHON
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.

Asking for a GPU you do not have 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.

Try it
  1. Run your script once with accelerator="cpu" and once with "auto". If you have a GPU, compare epoch times.
  2. 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.

PYTHON
# 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)
A debugging routine When something is wrong, work through these in order: 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.
Try it
  1. Introduce a deliberate bug (misspell a variable in validation_step) and see how quickly fast_dev_run=True reveals it compared with a full run.
  2. Fix it, then run with overfit_batches=10 and 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:

PYTHON
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.

Try it
  1. Set default_root_dir="runs/mnist" and confirm the logs move there.
  2. Add gradient_clip_val=1.0 and val_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".

How to read a Lightning traceback Scroll to the bottom first: the last line is the exception type and its message, and Lightning's messages usually include a hint in plain English. Then read upward until you reach a line in your file. That line, not the Trainer internals above it, is where to start fixing.
Try it
  1. Rename training_step to training_stap and run. Read the message, then fix it.
  2. Set monitor="val_losss" in ModelCheckpoint and 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.

Try it
  1. Find a Lightning tutorial online, locate any lines from the "old" column, and rewrite them in the current form.
  2. Run the rewritten code with fast_dev_run=True to 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.

mnist_full.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:

BASH
python mnist_full.py

Then walk through what happened, as a checklist for any Lightning run:

  1. The sanity check ran two validation batches.
  2. Each epoch trained, validated, logged val_loss and val_acc, and saved a checkpoint if val_loss was in the top two.
  3. Training either reached 20 epochs or EarlyStopping ended it, and the last line of Trainer output says which.
  4. trainer.test(..., ckpt_path="best") loaded the best checkpoint and printed test metrics.
  5. logs/mnist/version_0/metrics.csv holds the numbers and checkpoints/ holds the models.

To reuse the trained model later without the Trainer:

PYTHON
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.

Try it
  1. Run the full script and read the final test accuracy.
  2. Change hidden to 256 and run again; compare val_acc in the two CSV files.
  3. Load checkpoints/last.ckpt in 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.

Sources