A raw PyTorch training script starts at twenty lines and grows into hundreds: device placement, mixed precision context managers, gradient scaling, accumulation, distributed process groups, sampler sharding, checkpoint save and resume, rank-zero logging. Each of those is easy alone and fragile together, and the bugs they cause are silent: a duplicated validation sample, a learning-rate schedule stepping per micro-batch, a checkpoint that resumes the model but not the optimizer.

PyTorch Lightning (now shipped as the lightning package, version 2.6 at the time of writing) separates the research code you care about from that engineering. You write a LightningModule that says what one step computes; the Trainer decides how it runs. This article explains the split from first principles, walks through a complete module, the hook order, the distributed strategies, a memory worked example for a 7B fine-tune on eight GPUs, checkpointing, when to drop down to Fabric, and the failure modes that still get through.

The split: what you write and what the Trainer owns

Who owns what in a Lightning runYou writeLightningModule: model, step,loss, optimizer configLightningDataModule: dataTrainerfit / validate / test loopshook order, accumulation,clipping, logging, ckpthooksAcceleratorcuda, cpu, mps, tpuStrategyddp, fsdp, deepspeedPrecision pluginbf16-mixed, 16-mixedCallbackscheckpoint, LR monitorLoggersTensorBoard, W and B, CSVThe module never calls .cuda(), .backward() or optimizer.step() directly in automatic optimisation:the Trainer and its plugins do, which is why the same module runs on 1 GPU or 64.
The LightningModule declares computation; the Trainer drives loops and delegates hardware, parallelism and numeric precision to pluggable components.

The design rule is that anything which depends on where or how the model runs belongs to the Trainer, and anything that depends on what the model is belongs to the module. Hardware is the Accelerator. How work is spread across processes, and how parameters are replicated or sharded, is the Strategy. Autocast, gradient scaling and parameter dtype are the precision plugin. Cross-cutting behaviour such as checkpointing and early stopping are Callback objects, and metric output goes to loggers.

Because the module never hard-codes any of those, the same file trains on a laptop CPU for debugging and on several nodes for the real run; only Trainer arguments change. That is the main practical value, more than the reduction in lines.

A complete LightningModule

Here is a complete module for causal language-model fine-tuning. It is deliberately plain: the model, one step, logging and optimizer configuration.

import torch
import lightning as L
from transformers import AutoModelForCausalLM, get_cosine_schedule_with_warmup

class CausalLMFineTune(L.LightningModule):
    def __init__(self, model_name: str, lr: float = 2e-5, warmup: int = 100):
        super().__init__()
        self.save_hyperparameters()          # stored in every checkpoint
        self.model = AutoModelForCausalLM.from_pretrained(model_name)

    def training_step(self, batch, batch_idx):
        out = self.model(**batch)            # labels in batch -> out.loss
        self.log("train/loss", out.loss, prog_bar=True)
        return out.loss                      # Trainer does backward and step

    def validation_step(self, batch, batch_idx):
        loss = self.model(**batch).loss
        self.log("val/loss", loss, sync_dist=True)   # mean across ranks

    def configure_optimizers(self):
        opt = torch.optim.AdamW(self.parameters(), lr=self.hparams.lr,
                                weight_decay=0.1)
        sched = get_cosine_schedule_with_warmup(
            opt, self.hparams.warmup, self.trainer.estimated_stepping_batches)
        return {"optimizer": opt,
                "lr_scheduler": {"scheduler": sched, "interval": "step"}}

Three details matter. interval: "step" makes the scheduler advance per optimizer step; the default is per epoch, which is wrong for warmup-plus-cosine LLM schedules. estimated_stepping_batches gives the total number of optimizer steps after accounting for devices and gradient accumulation. And sync_dist=True on the validation metric averages it across ranks, so the number you checkpoint on is global rather than whatever rank zero happened to see.

The Trainer that runs it:

from lightning.pytorch.callbacks import ModelCheckpoint, LearningRateMonitor

trainer = L.Trainer(
    accelerator="gpu", devices=8, num_nodes=1,
    strategy="fsdp",
    precision="bf16-mixed",
    max_steps=3000,                # set explicitly; see failure modes
    accumulate_grad_batches=8,
    val_check_interval=500,
    log_every_n_steps=10,
    callbacks=[ModelCheckpoint(monitor="val/loss", save_last=True, save_top_k=2),
               LearningRateMonitor(logging_interval="step")],
)
L.seed_everything(1234, workers=True)
trainer.fit(CausalLMFineTune("my-org/base-7b"), datamodule=dm, ckpt_path="last")

Data with a LightningDataModule

Data loading is where distributed runs most often go quietly wrong, so Lightning gives it its own class. A LightningDataModule groups download, tokenisation, splitting and loader construction, and splits them by where they must run:

class TokenizedText(L.LightningDataModule):
    def __init__(self, path, tokenizer, batch_size=4, seq_len=2048):
        super().__init__()
        self.save_hyperparameters(ignore=["tokenizer"])
        self.tok = tokenizer

    def prepare_data(self):            # once per node: download, cache to disk
        build_token_cache(self.hparams.path, self.tok, self.hparams.seq_len)

    def setup(self, stage):            # on every rank: load cached shards
        self.train_ds, self.val_ds = load_token_cache(self.hparams.path)

    def train_dataloader(self):
        return DataLoader(self.train_ds, batch_size=self.hparams.batch_size,
                          shuffle=True, num_workers=8, pin_memory=True)

    def val_dataloader(self):
        return DataLoader(self.val_ds, batch_size=self.hparams.batch_size,
                          num_workers=4)

prepare_data runs on one process per node (by default), so it is the place for anything that writes to shared disk; doing that in setup makes eight processes race on the same files. Do not assign state in prepare_data, because other ranks never see it. Note that the loaders contain no DistributedSampler: the Trainer inserts one per rank and calls set_epoch on it each epoch, so shuffling differs between epochs but stays consistent across ranks. Write the module once and the same data path serves one GPU or many.

The training loop, hook by hook

Understanding when Lightning calls your code removes most surprises. A simplified view of one training epoch under automatic optimisation:

  1. on_train_epoch_start, then for each batch the DataLoader yields (already sharded per rank by a distributed sampler unless you set use_distributed_sampler=False).
  2. The batch is moved to the device, then on_train_batch_start runs, then training_step runs inside the precision plugin's autocast context.
  3. The returned loss is divided by accumulate_grad_batches and backward runs. On non-boundary micro-batches under DDP, gradient synchronisation is skipped.
  4. On an accumulation boundary: optional gradient clipping, optimizer.step(), zero_grad, and step-interval schedulers advance. trainer.global_step counts these optimizer steps, not micro-batches.
  5. At val_check_interval the validation loop runs on all ranks, metrics logged with sync_dist are reduced, and callbacks such as ModelCheckpoint act on the result.

If you need GAN-style alternating optimizers or custom clipping, set self.automatic_optimization = False and call self.manual_backward(loss) and the optimizer yourself; the Trainer still owns devices, precision and processes.

Scaling out with strategies

The strategy string is where scaling decisions live:

StrategyWhat is replicatedUse when
ddpFull model, grads and optimizer state on every GPUModel plus Adam state fits on one GPU
fsdpParameters, grads and optimizer state sharded across ranksModel state does not fit on one GPU
deepspeed_stage_2 / _3ZeRO partitioning of optimizer state, grads, then paramsYou already use DeepSpeed configs or offload
ModelParallelStrategyFSDP2 plus tensor parallel, user-definedVery large models; experimental, needs PyTorch 2.4+

For FSDP, pass an FSDPStrategy object rather than the string when you need control. auto_wrap_policy and activation_checkpointing_policy both accept a set of layer classes, typically your transformer block class; sharding_strategy defaults to FULL_SHARD; and state_dict_type is "full" by default or "sharded".

from lightning.pytorch.strategies import FSDPStrategy
from transformers.models.llama.modeling_llama import LlamaDecoderLayer

strategy = FSDPStrategy(
    auto_wrap_policy={LlamaDecoderLayer},
    activation_checkpointing_policy={LlamaDecoderLayer},
    sharding_strategy="FULL_SHARD",
    state_dict_type="sharded",     # each rank writes its own shard
)

Worked example: a 7B fine-tune on eight GPUs

Fine-tune a 7B-parameter model with AdamW on one node of eight 80 GB GPUs, sequence length 2048. With precision="bf16-mixed" the parameters are held in fp32 and autocast computes in bf16, so persistent state per parameter is 4 bytes of weight, 4 of gradient and 8 for Adam's two moments: 16 bytes. That assumes the weights are fp32 when FSDP wraps them; if you load them in bf16, check what dtype your optimizer state ends up in before trusting the table.

ItemDDP (per GPU)FSDP FULL_SHARD (per GPU)
Weights fp32, 7B x 4 B28 GB3.5 GB
Gradients fp3228 GB3.5 GB
Adam moments, 7B x 8 B56 GB7 GB
Persistent total112 GB: does not fit14 GB plus gathered layer

DDP is ruled out before activations are even counted. Under FSDP each rank holds 14 GB of sharded state and temporarily gathers one wrapped block at a time, which is why the auto-wrap policy must name the decoder layer; without it the whole model is one unit and the gather is the full model. That leaves the rest of the 80 GB for activations. With activation checkpointing on each decoder layer, a per-device batch of 4 sequences is a reasonable starting point; measure peak memory with a short run before committing.

The global batch is per-device batch times devices times nodes times accumulation: 4 x 8 x 1 x 8 = 256 sequences, or about 524,000 tokens per optimizer step. Change any factor and the learning rate that worked may stop working, so record the global batch with every run. With 3,000 steps the run sees about 1.6 billion tokens.

Checkpoints and resume

A Lightning checkpoint contains model state, optimizer and scheduler state, loop progress, callback state and the hyperparameters saved by save_hyperparameters. trainer.fit(..., ckpt_path="last") resumes all of it, including the global step and the position in the schedule, which is the difference between a true resume and a warm start.

With FSDP and state_dict_type="full", the full state is gathered to be written as one file; for large models that is slow and can exhaust CPU memory on rank zero. "sharded" writes a directory with one shard per rank, which is fast and scales, but you must convert it before loading the weights outside Lightning. Test the resume path on day one by killing a run deliberately; a checkpoint you have never restored is a hope, not a backup. See checkpointing in depth for the general trade-offs.

Fabric: when the Trainer is too much

Lightning also ships Fabric, which gives you the strategy, precision and launcher machinery without the Trainer's loop. You keep your own for loop and call fabric.backward(loss) instead of loss.backward():

import lightning as L

fabric = L.Fabric(accelerator="cuda", devices=8, strategy="fsdp", precision="bf16-mixed")
fabric.launch()
model, optimizer = fabric.setup(model, optimizer)
loader = fabric.setup_dataloaders(loader)

for step, batch in enumerate(loader):
    loss = model(**batch).loss
    fabric.backward(loss)
    optimizer.step(); optimizer.zero_grad()
    if step % 500 == 0:
        fabric.save("ckpt/step.ckpt", {"model": model, "optimizer": optimizer, "step": step})

Choose Fabric when the training loop itself is the research (RL rollouts, unusual accumulation, interleaved generation) or when porting an existing script with minimal edits. Choose the Trainer when the loop is standard and you want callbacks, resume semantics and logging handled for you.

Debugging before you scale

Most Lightning bugs are cheaper to find on a laptop than on a node of eight GPUs, and the Trainer has flags for exactly that. Use them in this order before a long run:

  • fast_dev_run=True runs one training and one validation batch through every hook, with logging and checkpointing disabled. It catches shape errors and typos in seconds.
  • overfit_batches=2 trains on two batches repeatedly. If loss does not fall towards zero, the bug is in the model, labels or optimizer, not in the data volume.
  • limit_train_batches and limit_val_batches shorten epochs so you can test checkpoint, resume and schedule behaviour in minutes.
  • profiler="simple" reports time per hook, which shows whether the data loader or the step is the bottleneck; "pytorch" gives a full kernel trace.
  • detect_anomaly=True finds the operation that first produced a NaN, at a large speed cost, so enable it only while chasing one.

The sanity check that runs a couple of validation batches before training (controlled by num_sanity_val_steps) is worth keeping: it fails a misconfigured validation loop before you have spent an hour of GPU time.

Failure modes

  • Rank-dependent control flow hangs. Code that runs a collective, including self.log(..., sync_dist=True), on some ranks but not others deadlocks until the NCCL timeout. Never put such calls behind if self.global_rank == 0.
  • Unbounded training. If neither max_epochs nor max_steps is set, Lightning defaults to max_epochs = 1000. Always set one explicitly.
  • Clipping under FSDP. With the Trainer's FSDP precision plugin, gradient_clip_algorithm="norm" (the default algorithm when you set gradient_clip_val) raises a configuration error in current releases, because per-shard norm clipping is wrong under FSDP. Check the docs for your version before relying on clipping there.
  • Scheduler per epoch. Returning a scheduler without "interval": "step" steps it once per epoch, so warmup never ends on a one-epoch fine-tune.
  • Duplicated validation samples. The distributed sampler pads the last batch so every rank has equal work, which repeats a few samples. For exact metrics on small eval sets, evaluate on one device or de-duplicate.
  • fp16 overflow. 16-mixed uses a gradient scaler and can still produce NaNs on some LLMs; prefer bf16-mixed on GPUs that support bf16. Background in mixed precision training.

Trade-offs

ChoiceGainCost
TrainerLoops, resume, callbacks, logging done for youHook order to learn; harder to bend for unusual loops
FabricYour own loop with distributed and precision handledYou own checkpoint and resume semantics
FSDPModel state sharded; large models fitCommunication per layer; checkpoint conversion
DeepSpeed strategyOffload and ZeRO configsSecond config system; see DeepSpeed
Activation checkpointingLarge memory savingRoughly one extra forward pass of compute

What to do next

  1. Move model, step and optimizer code into a LightningModule and data into a LightningDataModule; keep device and dtype calls out of both.
  2. Run on CPU with fast_dev_run=True to exercise every hook once before touching a GPU.
  3. Do the 16-bytes-per-parameter arithmetic for your model and pick DDP or FSDP from it, naming your decoder block in both FSDP policies; see FSDP in depth.
  4. Set max_steps, a step-interval scheduler and an explicit global batch, and log all three.
  5. Kill a run on purpose and resume it with ckpt_path="last"; confirm the loss curve and learning rate continue where they stopped.
  6. Try torch.compile on the inner model only after the baseline is correct, and compare step time and memory.
Key takeaway: Lightning works because it separates what a step computes from how and where it runs. Write the step once, then choose accelerator, strategy and precision as Trainer arguments. Do the memory arithmetic before picking DDP or FSDP, set max_steps and a per-step scheduler explicitly, keep collectives off rank-only code paths, and prove resume works before the long run starts.