Original Note

PyTorch Lightning(整理版) - Read

PyTorch Lightning(整理版)

原始资料: Learn PyTorch Lightning - ReCoDE Deep Learning Best Practices created: 2026-07-02 22:12 整理说明: 本版本结合原始笔记和可读取的原始教程重排、补全和翻译,保留常用英文术语。

内容简要概括

PyTorch Lightning 是对 PyTorch 训练流程的轻量封装,目标是把研究代码和工程训练循环分开,让模型更容易维护、复现实验和扩展到 GPU/多卡训练。使用时通常把模型逻辑写进 LightningModule,把数据交给 DataLoaderLightningDataModule,把训练控制交给 Trainer。新项目优先使用 lightning 新入口;旧项目或旧教程中仍可能看到 pytorch_lightning

PyTorch LightningLightningModuleTrainerLightningDataModuletraining_stepvalidation_stepconfigure_optimizersDataLoaderModelCheckpointmixed precision、GPU 训练、训练循环

目录


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_steptest_step 计算验证或测试指标,并用 self.log 记录。
优化器 configure_optimizers 返回 optimizer,必要时返回 scheduler。
数据加载 DataLoaderLightningDataModule 准备训练、验证、测试数据。
训练控制 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=2precision=16 在新版本中可能已被 acceleratordevicesprecision="16-mixed" 等参数替代。

6. 自定义数据集时怎么做

自定义数据集可以按这个顺序做:

  1. 写一个继承 torch.utils.data.Dataset 的类,实现 __len____getitem__
  2. __getitem__ 中返回 (x, y),或返回任务需要的字典结构。
  3. DataLoader(dataset, batch_size=..., shuffle=...) 包装数据集。
  4. 如果项目变复杂,把这些逻辑放进 LightningDataModule
  5. 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)

这样模型文件只关心训练逻辑,数据文件只关心数据准备,训练脚本只负责把它们组装起来。