PyTorch Lightning 基础训练指南:从零定义模型、数据集到 Trainer.fit 全流程实战
PyTorch Lightning 基础训练指南从零定义模型、数据集到 Trainer.fit 全流程实战【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000 GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning本文是 PyTorch Lightning 训练入门系列的基础篇面向希望训练模型但又不想手写训练循环的开发者。文章以仓库文档 docs/source-pytorch/model/train_model_basic.rst 为骨架结合仓库内 examples/pytorch/basics/autoencoder.py 真实示例与 src/lightning/pytorch 源码实现完整演示从编写nn.Module、定义LightningModule、构建DataLoader到调用Trainer.fit()完成一次模型训练的每一步。读完后你将掌握 Lightning 的基本编程模型并理解Trainer在底层替你执行的训练循环究竟做了什么——这也是后续学习验证/测试切分、学习率调度、分布式训练等高级特性的起点。快速了解本文的四个核心概念在开始写代码前先建立 Lightning 的最小心智模型。整个框架围绕四个角色运转角色对应类/对象职责网络结构torch.nn.Module定义前向计算纯 PyTorch 代码不感知 Lightning训练配方LightningModule定义训练/验证/测试步骤、优化器、学习率调度器数据供给torch.utils.data.DataLoader或LightningDataModule提供批次数据训练驱动器Trainer执行训练循环、设备调度、日志、断点续训等所有工程逻辑其中LightningModule是本文的核心它是描述你的nn.Module如何被训练的完整配方recipe。在仓库源码 src/lightning/pytorch/core/module.py 中可以看到LightningModule继承自_DeviceDtypeModuleMixin、HyperparametersMixin、ModelHooks、DataHooks、CheckpointHooks以及torch.nn.Module因此它本身就是一个nn.Module同时额外获得了设备/数据类型管理、超参数记录、生命周期钩子、检查点读写等能力。第一步添加导入在文件顶部添加本教程所需的全部导入。除了标准的 PyTorch 与 torchvision 工具外最关键的一行是import lightning as LTrainer、LightningModule等核心符号都从该命名空间导出见仓库 src/lightning/init.py 中导出的Trainer、LightningModule、LightningDataModule、Callback、seed_everything、Fabricimport os import torch from torch import nn import torch.nn.functional as F from torchvision import transforms from torchvision.datasets import MNIST from torch.utils.data import DataLoader import lightning as L第二步定义 PyTorch nn.ModuleLightning 并不要求你用任何特殊方式编写网络结构——它就是普通的nn.Module。下面定义一个小型自编码器Autoencoder的两个部件Encoder把 28×28 的 MNIST 图像压缩为 3 维潜在表示Decoder把 3 维潜在表示还原为 28×28 的图像。class Encoder(nn.Module): def __init__(self): super().__init__() self.l1 nn.Sequential(nn.Linear(28 * 28, 64), nn.ReLU(), nn.Linear(64, 3)) def forward(self, x): return self.l1(x) class Decoder(nn.Module): def __init__(self): super().__init__() self.l1 nn.Sequential(nn.Linear(3, 64), nn.ReLU(), nn.Linear(64, 28 * 28)) def forward(self, x): return self.l1(x)这里没有任何 Lightning 相关代码nn.Module只负责前向计算。仓库中的 examples/pytorch/basics/autoencoder.py 给出了一个更完整的等价实现将 Encoder/Decoder 合并在一个LitAutoEncoder中并用save_hyperparameters()自动记录hidden_dim与learning_rate等超参数。第三步定义 LightningModule——把训练逻辑写成配方LightningModule是整个训练流程的配方它把所有工程上容易出错、但逻辑上相似的部分训练步骤、优化器配置组织成标准钩子交给Trainer统一驱动。在本例中只需实现两个钩子training_step(batch, batch_idx)定义nn.Module之间如何交互——接收一个批次前向计算并返回损失configure_optimizers()定义模型使用的优化器还可以返回学习率调度器。class LitAutoEncoder(L.LightningModule): def __init__(self, encoder, decoder): super().__init__() self.encoder encoder self.decoder decoder def training_step(self, batch, batch_idx): # training_step 定义了训练循环的单步行为 x, _ batch x x.view(x.size(0), -1) z self.encoder(x) x_hat self.decoder(z) loss F.mse_loss(x_hat, x) return loss def configure_optimizers(self): optimizer torch.optim.Adam(self.parameters(), lr1e-3) return optimizer对自动优化默认开启模式training_step返回的损失值会作为反向传播的输入。从源码 src/lightning/pytorch/loops/optimization/automatic.py 可以看出返回值既可以是单个 Tensor也可以是包含loss键的字典此时字典中除loss外的其他键会被作为额外输出收集。仓库测试模型 tests/tests_pytorch/helpers/simple_models.py 就演示了后者返回{loss: loss}同时用self.log(train_loss, loss, prog_barTrue)把指标记录到进度条与日志。提示在training_step里用self.log(...)记录指标是 Lightning 的标准做法指标会按步骤/epoch 自动聚合并写入你配置的 Logger如 TensorBoard、CSV。而configure_optimizers还可以返回更复杂的结构例如{optimizer: optim, lr_scheduler: sched, monitor: val_loss}或优化器列表多优化器场景参见 src/lightning/pytorch/core/optimizer.py 中_init_optimizers_and_lr_schedulers的解析逻辑。第四步定义训练数据集数据部分仍然使用纯 PyTorch API构造torch.utils.data.DataLoader其中包含你的训练数据集。dataset MNIST(os.getcwd(), downloadTrue, transformtransforms.ToTensor()) train_loader DataLoader(dataset)MNIST(os.getcwd(), ...)会在当前工作目录下载 MNIST 数据集并缓存transforms.ToTensor()将图像转为[0, 1]区间的张量。这里DataLoader未显式指定batch_size默认值为 1实际训练中建议根据显存设置 batch size如 32、64并可用num_workers开启多进程加载。仓库中的 examples/pytorch/basics/autoencoder.py 展示了更工程化的做法用LightningDataModule封装数据通过random_split把训练集切成 55000/5000 的 train/val 两份并分别提供train_dataloader()、val_dataloader()、test_dataloader()与predict_dataloader()——这样训练、验证、测试三阶段的数据就统一在一个模块里管理了。第五步用 Trainer.fit 训练模型一切就绪后实例化模型与Trainer调用fit()即可开始训练# model autoencoder LitAutoEncoder(Encoder(), Decoder()) # train model trainer L.Trainer() trainer.fit(modelautoencoder, train_dataloaderstrain_loader)Trainer负责处理全部工程细节把规模化所需的复杂度都抽象掉。Trainer()不带任何参数时会自动选择可用的硬件acceleratorauto、devicesauto并启用默认的回调模型摘要、进度条、断点保存与默认的 CSV 日志。Trainer.fit的完整签名见仓库 src/lightning/pytorch/trainer/trainer.py支持以下关键入参参数含义model要训练的LightningModule必填train_dataloaders训练数据加载器或加载器列表val_dataloaders验证数据加载器可省略datamodule若使用LightningDataModule用它替代前两个参数ckpt_path从指定检查点恢复训练断点续训Trainer构造函数中与基础训练最相关的常用参数还包括trainer L.Trainer( acceleratorauto, # auto | cpu | gpu | tpu | mps ... devicesauto, # 使用的设备数量如 1、[0, 1] 或 auto strategyauto, # 分布式策略如 ddp、fsdp、deepspeed precisionNone, # 混合精度16-mixed、bf16-mixed 等 max_epochsNone, # 训练总轮数 fast_dev_runFalse, # 快速试跑仅运行 1 个 batch验证代码可运行 log_every_n_steps50, # 每 N 步写一次日志 enable_checkpointingTrue, # 是否自动保存检查点 enable_progress_barTrue, # 是否显示进度条 accumulate_grad_batches1, # 梯度累积步数 gradient_clip_valNone, # 梯度裁剪阈值 )关于max_epochsTrainer 默认训练 1 个 epoch如需更多轮次请显式设置。初次接触时可以用fast_dev_runTrue快速验证整条链路数据、模型、钩子、优化器能否跑通。第六步消除训练循环——Trainer 在底层替你做了什么上面的trainer.fit(...)一行代码背后Lightning 替你执行了如下训练循环autoencoder LitAutoEncoder(Encoder(), Decoder()) optimizer autoencoder.configure_optimizers() for batch_idx, batch in enumerate(train_loader): loss autoencoder.training_step(batch, batch_idx) loss.backward() optimizer.step() optimizer.zero_grad()这个朴素版循环对应 Lightning 自动优化模式的三个基本动作前向计算训练损失 → 反向传播 → 优化器更新参数。在真实源码中这三个动作被封装为一个Closure闭包。查看 src/lightning/pytorch/loops/optimization/automatic.py 可以看到Closure类把training_step、backward、zero_grad三个子闭包合并为一个然后交给optimizer.step(closure)调用class Closure(AbstractClosure[ClosureResult]): ... combines three elementary closures into one: training_step, backward and zero_grad. ... def closure(self, *args, **kwargs): step_output self._step_fn() # 1. 调用 training_step if self._zero_grad_fn is not None: self._zero_grad_fn() # 2. 梯度清零 if self._backward_fn is not None and step_output.closure_loss is not None: self._backward_fn(step_output.closure_loss) # 3. 反向传播 return step_output而负责按 epoch 遍历 DataLoader、执行上述优化、按间隔运行验证的驱动器则是训练 epoch 循环类_TrainingEpochLoop见 src/lightning/pytorch/loops/training_epoch_loop.py。它的 docstring 明确说明训练 epoch 循环负责调用*_epoch_{start,end}钩子、按请求间隔运行验证验证由独立的_EvaluationLoop负责并在需要时支持多优化器自动优化——这正是文档中朴素循环之上 Lightning 增加的全部价值。Closure相较于手写循环还处理了许多边界情况支持梯度累积accumulate_grad_batches 1时 loss 会被normalize归一化、混合精度AMP 对损失进行缩放、梯度裁剪、分布式梯度同步等。也就是说你在training_step里写的那几行代码经过 Trainer 编排后可以无缝接入这些高级特性而无需改动循环本身。为什么说循环越复杂Lightning 的价值越大文档中有一句非常关键的话Lightning 的强大之处在于当训练循环变得复杂时——当你加入验证/测试切分、学习率调度器、分布式训练以及各种最新 SOTA 技术时。以手写循环为例你每增加一个特性就要在for循环里再叠加一段样板代码例如每 N 个 epoch 跑一次验证并在验证时切换model.eval()/torch.no_grad()为每个 optimizer 绑定一个lr_scheduler并在 epoch 结束时scheduler.step()在 8 张 GPU 上做 DDP 同步处理DistributedSampler与梯度 all-reduce训练中断后从断点恢复 epoch、step、优化器状态保存/加载模型权重与超参数。而在 Lightning 中这些能力要么是Trainer的内置参数如val_check_interval、accumulate_grad_batches、strategyddp、ckpt_path...要么是LightningModule的对应钩子如validation_step、lr_scheduler_step、on_save_checkpoint。你可以把各种技术任意组合而不需要每次重写一个新的训练循环——这就是配方式编程模型带来的核心收益。仓库中 examples/pytorch/basics/autoencoder.py 展示了这种组合能力的实际形态在同一个LitAutoEncoder上通过LightningCLI一行配置即可同时获得fit、test、predict三个阶段的完整流程含ckpt_pathbest的断点恢复与推理。从基础走向进阶后续学习路径掌握本文的基础流程后可以沿着以下方向继续深入均在当前仓库文档 docs/source-pytorch/model 目录内加入验证与测试在LightningModule中实现validation_step/test_step并在fit时传入val_dataloaders参考 build_model.rst手动控制优化循环当你需要 GAN 双优化器、累积梯度自定义等精细控制时参考 manual_optimization.rst 与 own_your_loop.rst构建更完整的模型save_hyperparameters、forward与推理、模型导出等参考 build_model_advanced.rst 与 build_model_intermediate.rst使用 Trainer 的全部能力包括断点续训、回调、日志、分布式训练等可查阅 common/trainer.rst 与 common/lightning_module.rst。小结本文基于官方基础教程 train_model_basic.rst 的完整流程走通了 Lightning 训练模型的最小闭环普通nn.Module定义结构 →LightningModule定义训练配方 →DataLoader供给数据 →Trainer.fit()驱动训练。同时通过源码揭示了 Trainer 的底层循环本质training_step前向、backward、optimizer.step()、zero_grad()的组合automatic.py以及 epoch 级的循环编排training_epoch_loop.py。理解了这一层你就能明白为什么零代码改动地扩展规模是可行的——因为工程逻辑全部收敛在 Trainer 与标准钩子契约中你的模型代码只需要描述配方剩下的交给框架。【免费下载链接】pytorch-lightningPretrain, finetune ANY AI model of ANY size on 1 or 10,000 GPUs with zero code changes.项目地址: https://gitcode.com/gh_mirrors/py/pytorch-lightning创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →