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
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:
on_train_epoch_start, then for each batch the DataLoader yields (already sharded per rank by a distributed sampler unless you setuse_distributed_sampler=False).- The batch is moved to the device, then
on_train_batch_startruns, thentraining_stepruns inside the precision plugin's autocast context. - The returned loss is divided by
accumulate_grad_batchesand backward runs. On non-boundary micro-batches under DDP, gradient synchronisation is skipped. - On an accumulation boundary: optional gradient clipping,
optimizer.step(),zero_grad, and step-interval schedulers advance.trainer.global_stepcounts these optimizer steps, not micro-batches. - At
val_check_intervalthe validation loop runs on all ranks, metrics logged withsync_distare reduced, and callbacks such asModelCheckpointact 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:
| Strategy | What is replicated | Use when |
|---|---|---|
ddp | Full model, grads and optimizer state on every GPU | Model plus Adam state fits on one GPU |
fsdp | Parameters, grads and optimizer state sharded across ranks | Model state does not fit on one GPU |
deepspeed_stage_2 / _3 | ZeRO partitioning of optimizer state, grads, then params | You already use DeepSpeed configs or offload |
ModelParallelStrategy | FSDP2 plus tensor parallel, user-defined | Very 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.
| Item | DDP (per GPU) | FSDP FULL_SHARD (per GPU) |
|---|---|---|
| Weights fp32, 7B x 4 B | 28 GB | 3.5 GB |
| Gradients fp32 | 28 GB | 3.5 GB |
| Adam moments, 7B x 8 B | 56 GB | 7 GB |
| Persistent total | 112 GB: does not fit | 14 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=Trueruns one training and one validation batch through every hook, with logging and checkpointing disabled. It catches shape errors and typos in seconds.overfit_batches=2trains 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_batchesandlimit_val_batchesshorten 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=Truefinds 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 behindif self.global_rank == 0. - Unbounded training. If neither
max_epochsnormax_stepsis set, Lightning defaults tomax_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 setgradient_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-mixeduses a gradient scaler and can still produce NaNs on some LLMs; preferbf16-mixedon GPUs that support bf16. Background in mixed precision training.
Trade-offs
| Choice | Gain | Cost |
|---|---|---|
| Trainer | Loops, resume, callbacks, logging done for you | Hook order to learn; harder to bend for unusual loops |
| Fabric | Your own loop with distributed and precision handled | You own checkpoint and resume semantics |
| FSDP | Model state sharded; large models fit | Communication per layer; checkpoint conversion |
| DeepSpeed strategy | Offload and ZeRO configs | Second config system; see DeepSpeed |
| Activation checkpointing | Large memory saving | Roughly one extra forward pass of compute |
What to do next
- Move model, step and optimizer code into a
LightningModuleand data into aLightningDataModule; keep device and dtype calls out of both. - Run on CPU with
fast_dev_run=Trueto exercise every hook once before touching a GPU. - 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.
- Set
max_steps, a step-interval scheduler and an explicit global batch, and log all three. - Kill a run on purpose and resume it with
ckpt_path="last"; confirm the loss curve and learning rate continue where they stopped. - Try torch.compile on the inner model only after the baseline is correct, and compare step time and memory.