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]