PyTorch Lightning 训练框架
FreeGuideOnline
9阅读
2026-07-11
bash pip install pytorch-lightning torch torchvision
安装完成后,让我们从一个最简示例开始:用 Lightning 包装一个全连接网络训练 MNIST。
---
### 核心概念:LightningModule 与 LightningDataModule
PyTorch Lightning 将深度学习项目拆分为两个核心抽象:
#### LightningModule — 组织所有模型逻辑
它继承自 `torch.nn.Module`,但你不用直接写 `forward` 的训练循环。你需要定义以下方法:
- `__init__`:初始化模型层与超参数。
- `forward`:定义推理时的前向传播。
- `training_step`:单个训练批次的损失计算,返回损失张量。
- `validation_step`:单个验证批次的逻辑,可计算指标。
- `configure_optimizers`:返回优化器(和学习率调度器)。
- (可选)`test_step`、`predict_step` 等。
所有训练循环的细节(反向传播、梯度清零、设备传输)都由 Lightning 在内部自动完成。
#### LightningDataModule — 封装数据加载流程
它将数据集拆分为可复用的模块,方便在不同项目间共享:
- `prepare_data`:下载、预处理数据(单进程调用)。
- `setup`:在每张 GPU 上执行,用于划分数据集、做变换。
- `train_dataloader`、`val_dataloader`、`test_dataloader`:返回对应的 DataLoader 对象。
数据模块让你的数据处理逻辑变得干净且可插拔。
---
### 第一个 Lightning 项目:训练 MNIST 分类器
我们将逐步构建一个完整的 MNIST 训练管线。所有代码均可直接运行。
#### 步骤 1:导入依赖并定义 LightningDataModule
```python
import torch
from torch import nn
from torch.utils.data import DataLoader, random_split
from torchvision import transforms, datasets
import pytorch_lightning as pl
class MNISTDataModule(pl.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.transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
def prepare_data(self):
# 下载数据集,仅单进程调用
datasets.MNIST(self.data_dir, train=True, download=True)
datasets.MNIST(self.data_dir, train=False, download=True)
def setup(self, stage: str = None):
# 在每张 GPU 上划分训练/验证集
if stage == "fit" or stage is None:
mnist_full = datasets.MNIST(self.data_dir, train=True, transform=self.transform)
self.mnist_train, self.mnist_val = random_split(
mnist_full, [55000, 5000], generator=torch.Generator().manual_seed(42)
)
if stage == "test" or stage is None:
self.mnist_test = datasets.MNIST(self.data_dir, train=False, transform=self.transform)
def train_dataloader(self):
return DataLoader(self.mnist_train, batch_size=self.batch_size, shuffle=True)
def val_dataloader(self):
return DataLoader(self.mnist_val, batch_size=self.batch_size)
def test_dataloader(self):
return DataLoader(self.mnist_test, batch_size=self.batch_size)
步骤 2:定义 LightningModule
import torch.nn.functional as F
from torch.optim import Adam
from torchmetrics import Accuracy
class MNISTClassifier(pl.LightningModule):
def __init__(self, lr=1e-3):
super().__init__()
self.save_hyperparameters() # 自动保存超参数到 self.hparams
self.lr = lr
self.model = nn.Sequential(
nn.Flatten(),
nn.Linear(28 * 28, 128),
nn.ReLU(),
nn.Linear(128, 10)
)
self.train_acc = Accuracy(task="multiclass", num_classes=10)
self.val_acc = Accuracy(task="multiclass", num_classes=10)
def forward(self, x):
return self.model(x)
def training_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = F.cross_entropy(logits, y)
# 记录训练损失和准确率(自动显示在进度条/日志)
self.log("train_loss", loss, prog_bar=True)
self.train_acc(logits, y)
self.log("train_acc", self.train_acc, prog_bar=True, on_step=False, on_epoch=True)
return loss
def validation_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = F.cross_entropy(logits, y)
self.val_acc(logits, y)
self.log("val_loss", loss, prog_bar=True)
self.log("val_acc", self.val_acc, prog_bar=True, on_step=False, on_epoch=True)
def configure_optimizers(self):
return Adam(self.parameters(), lr=self.lr)
步骤 3:训练与测试
if __name__ == "__main__":
# 初始化数据与模型
dm = MNISTDataModule()
model = MNISTClassifier()
# 配置 Trainer —— 实现自动 GPU 支持、日志、早停等
trainer = pl.Trainer(
max_epochs=5,
accelerator="auto", # 自动检测 GPU/TPU
devices=1, # 使用 1 张 GPU(无 GPU 则用 CPU)
log_every_n_steps=10,
enable_progress_bar=True
)
# 训练
trainer.fit(model, dm)
# 测试
trainer.test(model, datamodule=dm)
运行这段代码,你会看到清晰的进度条、实时的损失和准确率,无需编写任何 model.train()、optimizer.zero_grad() 或 loss.backward()。Lightning 替你完成了所有“脚手架”工作。
深入 Trainer:让训练流程如臂使指
pl.Trainer 是 Lightning 的总控枢纽,提供了数十个开箱即用的功能:
自动化硬件加速
accelerator="auto":自动选择当前环境可用的加速器(GPU、TPU、IPU 等)。devices=4:指定使用 4 张 GPU(如accelerator="gpu")。Lightning 自动处理DistributedDataParallel,无需修改代码。precision=16:开启混合精度训练,大幅减少显存占用并提升速度。
训练控制与回调
Lightning 通过“回调”(Callbacks)在训练生命周期的特定时刻插入逻辑。最常用的回调包括:
- ModelCheckpoint:自动保存最佳模型。
- EarlyStopping:验证指标无改善时提前停止。
- LearningRateMonitor:记录学习率变化到日志。
- Timer:限制训练时间。
示例:配置早停与检查点
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint
callbacks = [
EarlyStopping(monitor="val_loss", patience=3, mode="min"),
ModelCheckpoint(monitor="val_acc", mode="max", save_top_k=1, filename="best-{epoch}-{val_acc:.2f}")
]
trainer = pl.Trainer(callbacks=callbacks, max_epochs=50)
日志与可视化
Lightning 支持多种日志记录器(Logger),默认使用 TensorBoardLogger。只需在 Trainer 中指定:
from pytorch_lightning.loggers import TensorBoardLogger
logger = TensorBoardLogger("logs/", name="mnist_experiment")
trainer = pl.Trainer(logger=logger)
然后通过 tensorboard --logdir logs/ 即可实时查看训练曲线。你也可以通过一行配置切换为 WandB、MLflow 等。
进阶技巧:提升你的 Lightning 工程
1. 自动超参数记录与可视化
在 LightningModule 的 __init__ 中调用 self.save_hyperparameters() 会将所有 __init__ 的参数自动存储到 self.hparams 属性中,并自动被日志系统捕获。你还可以通过 self.hparams.lr 直接访问,非常适合实验追踪。
2. 多数据加载器与复杂验证
如果你的任务需要多个验证集(或同时验证多个数据集),只需定义 val_dataloader 返回一个列表:
def val_dataloader(self):
return [self.dl1, self.dl2]