PyTorch Lightning(整理版)
原始资料: Learn PyTorch Lightning - ReCoDE Deep Learning Best Practices created: 2026-07-02 22:12 整理说明: 本版本结合原始笔记和可读取的原始教程重排、补全和翻译,保留常用英文术语。
内容简要概括
PyTorch Lightning 是对 PyTorch 训练流程的轻量封装,目标是把研究代码和工程训练循环分开,让模型更容易维护、复现实验和扩展到 GPU/多卡训练。使用时通常把模型逻辑写进 LightningModule,把数据交给 DataLoader 或 LightningDataModule,把训练控制交给 Trainer。新项目优先使用 lightning 新入口;旧项目或旧教程中仍可能看到 pytorch_lightning。
PyTorch Lightning、LightningModule、Trainer、LightningDataModule、training_step、validation_step、configure_optimizers、DataLoader、ModelCheckpoint、mixed precision、GPU 训练、训练循环
目录
- 1. 新旧导入入口
- 2. Lightning 的职责拆分
- 3. 最小 LightningModule 示例
- 4. 数据准备与 LightningDataModule
- 5. Trainer、验证和回调
- 6. 自定义数据集时怎么做
1. 新旧导入入口
原笔记里的判断可以保留:
pytorch_lightning = 老入口,专指 PyTorch Lightning
lightning = 新入口,官方现在更推荐,包含 PyTorch Lightning + Fabric
新教程和新项目中优先写:
import lightning as L
如果某个老项目明确依赖旧包名,再使用:
import pytorch_lightning as pl
注意:原始教程示例使用 pytorch_lightning as pl,这是旧入口风格。整理笔记时保留旧示例,实际新项目可把 pl.LightningModule 改成 L.LightningModule,把 pl.Trainer 改成 L.Trainer。
2. Lightning 的职责拆分
一句话记忆:
model 负责定义“学什么”
train_loader 负责提供“用什么数据学”
trainer 负责控制“怎么训练”
更完整的分工:
| 部分 | 放在哪里 | 负责什么 |
|---|---|---|
| 模型计算 | LightningModule.forward |
定义前向传播。 |
| 训练步骤 | LightningModule.training_step |
取出 batch、计算 logits、计算 loss、返回 loss。 |
| 验证/测试步骤 | validation_step、test_step |
计算验证或测试指标,并用 self.log 记录。 |
| 优化器 | configure_optimizers |
返回 optimizer,必要时返回 scheduler。 |
| 数据加载 | DataLoader 或 LightningDataModule |
准备训练、验证、测试数据。 |
| 训练控制 | Trainer |
控制 epoch、设备、精度、callbacks、logger。 |
3. 最小 LightningModule 示例
旧入口版本:
import pytorch_lightning as pl
import torch
from torch import nn
from torch.nn import functional as F
class LitModel(pl.LightningModule):
def __init__(self):
super().__init__()
self.layer = nn.Linear(28 * 28, 10)
def forward(self, x):
return self.layer(x.view(x.size(0), -1))
def training_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = F.cross_entropy(logits, y)
return loss
def validation_step(self, batch, batch_idx):
x, y = batch
logits = self(x)
loss = F.cross_entropy(logits, y)
self.log("val_loss", loss)
def configure_optimizers(self):
return torch.optim.Adam(self.parameters(), lr=0.001)
model = LitModel()
字段和函数说明:
| 名称 | 说明 |
|---|---|
__init__ |
定义网络层和需要保存的超参数。 |
forward |
推理时的前向计算,也可在 training_step 中调用。 |
training_step |
每个训练 batch 执行一次,通常返回 loss。 |
batch_idx |
当前 batch 的序号,有些任务不用也可以保留参数。 |
validation_step |
每个验证 batch 执行一次,用于计算验证 loss 或指标。 |
self.log("val_loss", loss) |
把指标交给 logger、进度条或 checkpoint callback 使用。 |
configure_optimizers |
统一声明 optimizer,避免手写训练循环。 |
4. 数据准备与 LightningDataModule
最简单时可以直接使用 PyTorch DataLoader:
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
train_loader = DataLoader(
datasets.MNIST("", train=True, download=True, transform=transforms.ToTensor()),
batch_size=32,
shuffle=True,
)
x, y = next(iter(train_loader))
print(x.shape, y.shape)
参数说明:
| 参数 | 作用 |
|---|---|
train=True |
读取训练集;如果是 False 通常读取测试集。 |
download=True |
本地没有数据时自动下载。 |
transform=transforms.ToTensor() |
把图片转成 tensor。 |
batch_size=32 |
每个 batch 包含 32 个样本。 |
shuffle=True |
训练时打乱样本顺序。 |
LightningDataModule 适合把数据下载、切分、预处理、DataLoader 创建集中管理。常见结构:
import lightning as L
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
class MNISTDataModule(L.LightningDataModule):
def __init__(self, data_dir="./data", batch_size=32):
super().__init__()
self.data_dir = data_dir
self.batch_size = batch_size
self.transform = transforms.ToTensor()
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=None):
full = datasets.MNIST(self.data_dir, train=True, transform=self.transform)
self.train_set, self.val_set = random_split(full, [55000, 5000])
self.test_set = datasets.MNIST(self.data_dir, train=False, transform=self.transform)
def train_dataloader(self):
return DataLoader(self.train_set, batch_size=self.batch_size, shuffle=True)
def val_dataloader(self):
return DataLoader(self.val_set, batch_size=self.batch_size)
def test_dataloader(self):
return DataLoader(self.test_set, batch_size=self.batch_size)
5. Trainer、验证和回调
最小训练:
model = LitModel()
trainer = pl.Trainer(max_epochs=3)
trainer.fit(model, train_loader)
参数说明:
| 参数 | 作用 |
|---|---|
max_epochs=3 |
完整遍历训练集 3 次。 |
trainer.fit(model, train_loader) |
启动训练,把模型和训练数据交给 Trainer。 |
常见 checkpoint callback:
from pytorch_lightning.callbacks import ModelCheckpoint
checkpoint_callback = ModelCheckpoint(
monitor="val_loss",
dirpath="./my_model",
filename="sample-mnist-{epoch:02d}-{val_loss:.2f}",
)
trainer = pl.Trainer(max_epochs=3, callbacks=[checkpoint_callback])
参数说明:
| 参数 | 作用 |
|---|---|
monitor="val_loss" |
根据 self.log("val_loss", loss) 记录的指标保存模型。 |
dirpath="./my_model" |
checkpoint 保存目录。 |
filename=... |
checkpoint 文件名模板,可包含 epoch 和指标值。 |
callbacks=[checkpoint_callback] |
把回调交给 Trainer 执行。 |
高级训练能力通常也由 Trainer 控制,例如 GPU、分布式训练、mixed precision。新版本 Lightning 中建议查看当前官方 API,因为旧示例里的 gpus=2、precision=16 在新版本中可能已被 accelerator、devices、precision="16-mixed" 等参数替代。
6. 自定义数据集时怎么做
自定义数据集可以按这个顺序做:
- 写一个继承
torch.utils.data.Dataset的类,实现__len__和__getitem__。 - 在
__getitem__中返回(x, y),或返回任务需要的字典结构。 - 用
DataLoader(dataset, batch_size=..., shuffle=...)包装数据集。 - 如果项目变复杂,把这些逻辑放进
LightningDataModule。 - 在
Trainer.fit(model, datamodule=dm)中传入数据模块。
典型调用:
dm = MNISTDataModule(data_dir="./data", batch_size=64)
model = LitModel()
trainer = L.Trainer(max_epochs=5)
trainer.fit(model, datamodule=dm)
这样模型文件只关心训练逻辑,数据文件只关心数据准备,训练脚本只负责把它们组装起来。