ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

从PyTorch到Lightning:重构深度学习训练流程的实战解析

从PyTorch到Lightning:重构深度学习训练流程的实战解析 从PyTorch到Lightning不止是少写几十行样板代码我先说个真实经历。去年接了一个视觉模型的项目模型本身是一个Transformer变体真正让人头疼的不是网络结构怎么搭而是训练脚本里的那一大坨断点续训、梯度裁剪、学习率调度、验证指标汇总、多卡同步、日志落盘……每换一个数据集就要把这套流程从头再抄一遍。后来我同事甩了一句你用PyTorch Lightning试试我一开始是拒绝的——总觉得这种封装框架会把我写代码的自由给抢走。结果用了三个星期之后我的想法发生了彻底改变甚至把之前项目里两套自研的训练模板全部替换成了Lightning。这篇文章不想写成官方文档的翻译而是从一个实际用者的角度把PyTorch Lightning到底解决了什么问题、它的核心机制是怎么回事、怎么从原生PyTorch平滑迁过去以及那些文档里没写清楚的坑一次性说透。无论你是刚看完PyTorch基础教程的新手还是已经被训练样板代码烦透的老手看完这篇应该都能自己动手把项目迁过来并且知道迁的时候要躲开哪些雷。1. 先搞清楚Lightning到底替你干了哪些活很多朋友第一次接触Lightning最容易产生的困惑是它和我直接写一个train.py有什么区别我帮你把这个问题掰开揉碎。1.1 它没替你写模型它替你写的是流程原生PyTorch训练一个模型核心流程不外乎这几步构造模型、遍历DataLoader拿batch、前向算loss、反向更新梯度、算验证指标、存checkpoint、写TensorBoard。这套流程本身不复杂麻烦在于它会把你真正要研究的东西——模型结构、损失函数设计、实验对比——给淹没掉。Lightning做的事情特别简单它把训练循环、验证循环、测试循环、预测循环全部包装成了一个标准的Trainer而你只需要告诉它模型怎么做前向、loss怎么算、指标怎么log。换句话说模型怎么算由你定什么时候算、怎么调度、怎么并行、怎么保存由Trainer说了算。我第一次把代码迁到Lightning后最直观的感受是项目里的train.py从三百多行缩到了八十多行。剩下来的几乎全是模型定义本身和数据处理逻辑。这叫把研究代码和工程代码分离听着抽象实际体验就是——你再也不用在每次实验前祈祷这个脚本别出幺蛾子。1.2 你可能不需要学一堆新概念Lightning的API设计其实很克制核心只需要理解两个东西LightningModule和Trainer。LightningModule是nn.Module的子类模型结构、优化器、loss计算都放在这里面但它比普通nn.Module多了一组钩子方法比如training_step、validation_step、configure_optimizers。Trainer则是一个调度器负责调用这些钩子方法、管理设备、控制日志和保存。注意你完全可以把LightningModule当普通的nn.Module用甚至单独拿出去做推理或者嵌入到其他框架里它不依赖任何全局状态。这一点对渐进式迁移特别友好后面我会详细讲。1.3 一个最小的Lightning训练流程长什么样我先给一个最极简的例子让没接触过的朋友对整体长什么样有个直觉import pytorch_lightning as pl import torch from torch import nn class MyModel(pl.LightningModule): def __init__(self): super().__init__() self.fc nn.Linear(784, 10) def forward(self, x): return self.fc(x) def training_step(self, batch, batch_idx): x, y batch logits self(x) loss nn.functional.cross_entropy(logits, y) self.log(train_loss, loss) return loss def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lr1e-3) model MyModel() trainer pl.Trainer(max_epochs10) trainer.fit(model) # 数据呢这里我省略了DataLoader真实使用会传train_dataloaders参数你发现没有整个脚本里最复杂的训练逻辑就只剩了training_step这一个函数。梯度怎么累积、什么时候调优化器、什么时候调scheduler、反向传播怎么执行——全是Trainer在幕后替你干了。2. LightningModule与Trainer的分工为什么这样设计是聪明的现在问题来了这种拆分到底聪明在哪很多吐槽Lightning的人会说它不过是把我的代码藏起来了出了问题更难查。这个说法一半对一半错。错的那一半在于它没有藏代码它只是把稳定不变的流程代码标准化了而你自己的核心逻辑依然完完整整留在你的文件里。2.1 研究代码与工程代码的物理隔离我特别喜欢Lightning的一点是它用类的方法边界把思想和管线切得干干净净。以前写训练脚本时我经常发现自己在同一个文件里同时干着两件不相干的事调模型结构和调分布式训练参数。Lightning的做法是模型类里你只写思想所有的分布式、精度、日志、checkpoint配置都扔给Trainer。这样带来的直接好处是代码的可读性和可复用性。同一个LightningModule我在单卡上能训练扔到8卡上能训练换到TPU上也能训练——只要改Trainer的参数就行模型的代码一行都不用动。这在我之前用原生PyTorch写DDP的时候是难以想象的DDP的初始化、进程分组、数据sampler切分每一步都必须小心伺候。2.2 钩子方法的设计哲学写你想写的剩下的交给我LightningModule里那几个_step方法谈不上什么黑科技但它给了你一套标准化的叙事结构。比如validation_step你只需要返回一个指标值Lightning会自动把所有进程上的结果汇总好再算平均值def validation_step(self, batch, batch_idx): x, y batch logits self(x) loss nn.functional.cross_entropy(logits, y) self.log(val_loss, loss, on_epochTrue, prog_barTrue) return loss你可以同时写on_validation_epoch_end来做更复杂的汇总比如收集所有batch的预测结果然后统一算指标。这套设计几乎完整覆盖了我能想到的所有训练流程变体。2.3 代价是什么一套额外需要学习的命名约定Lightning并不是零成本的。它引入了一套钩子方法的命名约定比如training_step、validation_step、test_step、predict_step、on_train_epoch_end等。这套命名是扁平的没有复杂的继承关系但你必须熟悉它们各自的调用时机。刚开始用的时候我是对着官方文档翻了好几次心里总犯嘀咕这个hook到底是在每个batch后调用还是在每个epoch的最后调用我的建议是不要死记硬背抓一个核心规律。training_step/validation_step/test_step是每个batch粒度的钩子on_*_epoch_start/end是每个epoch粒度的钩子on_*_batch_start/end是batch前后的通用钩子。把这三个粒度记牢90%的调用时机问题就能解决。剩下的10%看一遍官方文档里的hooks表格即可。3. 手把手把原生PyTorch训练脚本改写成Lightning工程理论讲再多不如直接动手。这里我用一个经典MNIST分类任务演示怎么把一个原生PyTorch脚本平滑重构成Lightning工程。我故意从真实项目中常见的乱糟糟脚本出发这样更有参考价值。3.1 改造前的原生PyTorch脚本先看看痛点假设你现在有一个标准的原生训练脚本大概长这样model Net() optimizer optim.Adam(model.parameters()) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) val_loader DataLoader(val_dataset, batch_size64) for epoch in range(10): model.train() for batch_idx, (x, y) in enumerate(train_loader): optimizer.zero_grad() out model(x) loss F.nll_loss(out, y) loss.backward() optimizer.step() # 每隔几步打印一下loss if batch_idx % 100 0: print(fepoch {epoch} batch {batch_idx} loss {loss.item()}) model.eval() total, correct 0, 0 with torch.no_grad(): for x, y in val_loader: out model(x) pred out.argmax(dim1) correct pred.eq(y).sum().item() total y.size(0) print(fval acc: {correct / total})这段代码看着还行但它有一个隐藏的致命伤训练逻辑和验证逻辑和打印逻辑全是在一个for循环里手搓的。当你把它扩展到多卡训练断点续训指标写入TensorBoard的时候每一个扩展都会改动这段代码的核心逻辑改动的次数多了各种bug就来了。3.2 第一步把网络和训练逻辑包进LightningModule迁移的第一步是把模型的定义和训练逻辑整合到LightningModule里。注意这里的关键不是把模型定义搬过来就完了还要把loss怎么算用什么优化器日志记什么这些决策也一起放进去import pytorch_lightning as pl import torch from torch import nn from torch.nn import functional as F from torch.utils.data import DataLoader, random_split from torchvision.datasets import MNIST from torchvision import transforms class LitMNIST(pl.LightningModule): def __init__(self, hidden_size64, learning_rate2e-4): super().__init__() self.save_hyperparameters() self.net nn.Sequential( nn.Flatten(), nn.Linear(28 * 28, hidden_size), nn.ReLU(), nn.Linear(hidden_size, hidden_size), nn.ReLU(), nn.Linear(hidden_size, 10) ) self.val_acc torchmetrics.Accuracy(taskmulticlass, num_classes10) def forward(self, x): return self.net(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) 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_barTrue) self.log(val_acc, self.val_acc, prog_barTrue) def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lrself.hparams.learning_rate)注意上面代码里我用了torchmetrics.Accuracy这是Lightning生态里配套的指标库你不需要自己手写acc的计算逻辑。所有指标的最后聚合、多卡同步它都帮你处理好比自己记一个total/correct再跨进程同步要省事得多。3.3 第二步用DataLoader或DataModule管理数据在原生脚本里数据加载就是两个DataLoader变量。在Lightning里你可以直接在LightningModule里实现train_dataloader和val_dataloaderdef train_dataloader(self): return DataLoader(self.train_dataset, batch_size64, shuffleTrue, num_workers4) def val_dataloader(self): return DataLoader(self.val_dataset, batch_size64, num_workers4)但更推荐的做法是单独建一个LightningDataModule。它能把下载/清洗/划分/加载这个过程完整沉淀下来复用性更强class MNISTDataModule(pl.LightningDataModule): def __init__(self, data_dir./data, batch_size64): super().__init__() self.data_dir data_dir self.batch_size batch_size def setup(self, stageNone): dataset MNIST(self.data_dir, trainTrue, downloadTrue, transformtransforms.ToTensor()) self.train_dataset, self.val_dataset random_split(dataset, [55000, 5000]) def train_dataloader(self): return DataLoader(self.train_dataset, batch_sizeself.batch_size, num_workers4) def val_dataloader(self): return DataLoader(self.val_dataset, batch_sizeself.batch_size, num_workers4)3.4 第三步用Trainer把整个流程跑起来重头戏来了。以前需要自己手写的一大段循环和配置现在全部收敛为一行dm MNISTDataModule() model LitMNIST(hidden_size128, learning_rate1e-3) trainer pl.Trainer(max_epochs10, acceleratorauto, devicesauto) trainer.fit(model, dm)acceleratorauto和devicesauto会自动检测当前环境是CPU还是GPU如果是GPU就自动用CUDA训练。相比原生代码里写model.cuda()再一个个搬tensorLightning连to(device)这一步都给你省了因为训练时tensor自动会被放到正确设备上。这一点新上手的朋友最容易搞懵记住一个原则在Lightning里你永远不需要自己在代码里调用.to(device)除非你在做非常特殊的raw tensor操作。3.5 重构后多出来的功能你现在白拿了什么和原来的脚本对比这次重构并不是单纯的代码搬家你白拿了至少四样东西断点续训Trainer的checkpoint_callback默认会保存last.ckpt训练中断后可以直接从断点恢复继续跑。EarlyStopping和ModelCheckpoint监控验证指标效果不好自动停效果好的权重自动保存全是指定式配置不用自己写逻辑。TensorBoard日志所有self.log()的记录都会自动写进日志目录跑完tensorboard --logdirlightning_logs就能看曲线。多卡训练把devices变成4strategyddp跑分布式训练几乎不需要改模型代码。这些功能如果用原生PyTorch实现每一样都是一大段需要反复测试的代码而现在都集成在了一个你本来就绕不开的流程环节里。4. 高频功能实测checkpoint、日志、EarlyStopping和梯度裁剪的配置细节接下来聊实际用得最多的四个功能配置。我不打算罗列所有参数只分享我测试下来最实用的组合和踩过的坑。4.1 ModelCheckpoint不只是保存最好模型ModelCheckpoint可能是Lightning里配置最繁琐但收益最高的组件。它的核心参数有monitor、mode、save_top_k、filename。我最常用的配置是from pytorch_lightning.callbacks import ModelCheckpoint checkpoint_callback ModelCheckpoint( monitorval_loss, modemin, save_top_k3, filenameepoch{epoch}-val_loss{val_loss:.4f}, auto_insert_metric_nameFalse, )这里有两个容易被忽略的点。第一filename里的大括号占位符必须和monitor的指标名一致否则保存时会报错或者显示missing。第二auto_insert_metric_nameFalse可以让文件名里不自动加指标名我一般都会关掉因为文件名太长在管理模型版本时很不方便。如果不配ModelCheckpointLightning默认只保存last.ckpt。也就是说官方模板里自动保存最佳模型这件事是靠这个callback实现的而不是Trainer参数天然带的功能这一点文档里写得很隐晦容易误会。4.2 EarlyStopping监控指标怎么选才不坑EarlyStopping的配置本身很简单from pytorch_lightning.callbacks import EarlyStopping early_stop EarlyStopping(monitorval_loss, patience3, modemin)但我想告诉大家一个真实感受监控val_loss还是监控val_acc结果差异极大。我自己的经验是分类任务里监控val_loss通常比监控val_acc更稳定因为acc是离散值步进不平滑容易出现原地不动几轮然后突然跳变的情况而val_loss是连续值对模型的整体拟合程度更敏感。除非你有明确的业务目标比如必须到达某个准确率门槛否则优先用loss做监控指标。另外一个容易被忽视的细节是EarlyStopping应该关注是否需要恢复最佳权重。Lightning默认训练结束时模型权重就是训练结束那一刻的权重并不自动回滚到validation最好的那一版。如果你想要早停后自动拿最优权重需要配合ModelCheckpoint去加载它或者在fit之后手动load_from_checkpoint。嫌麻烦的话可以看EarlyStopping的restore_best_weights参数设为True可以自动恢复但前提是你也得配置一个指向同一monitor指标的ModelCheckpoint。4.3 梯度裁剪和混合精度一键开启带来的性能变化Lightning把混合精度和梯度裁剪做成了Trainer的开关参数这个设计有时候会让人低估它的重要性trainer pl.Trainer( max_epochs10, precisionbf16-mixed, # 或者 16-mixed gradient_clip_val1.0, gradient_clip_algorithmnorm, )precisionbf16-mixed在Ampere及以上的NVIDIA GPU上非常实用显存占用能降三分之一以上训练速度通常还有提升。但注意bf16和fp16是有区别的bf16的范围和fp32一样所以训练稳定性更好多数情况下不会出现fp16那种loss突然变NaN的问题但bf16在少数不支持它的硬件上会报错所以老卡用户得用16-mixed加accumulate_grad_batches来缓一下。梯度裁剪我用的是algorithmnorm它按全局梯度范数裁剪比按值裁剪value更常用也更安全。当年用原生PyTorch时我都是手写在backward之后step之前现在一行配置搞定出错概率反而更低了。4.4 日志系统logger是抽象层不是某个特定后端新手最容易卡住的地方是日志怎么配。Lightning里Trainer的logger参数可以接一个或多个Logger对象比如TensorBoardLogger、CSVLogger、WandbLoggerfrom pytorch_lightning.loggers import TensorBoardLogger, CSVLogger tb_logger TensorBoardLogger(logs, namemnist_experiment) csv_logger CSVLogger(logs, namecsv_metrics) trainer pl.Trainer(logger[tb_logger, csv_logger])我推荐至少接一个CSVLogger因为它会输出一个纯文本的指标变化表哪怕TensorBoard因为环境问题打不开你也能用pandas直接读这个CSV做分析。TensorBoard和Wandb适合给人类看曲线CSV适合给代码做后处理两者结合最踏实。5. 多卡训练与分布式Lightning把DDP的复杂性藏在了哪里多卡训练是最让人觉得Lightning真值的场景。你要是用原生PyTorch写过DDP一定记得那些折磨人的细节环境变量初始化、DistributedSampler、barrier同步、rank判断、模型广播……在Lightning里这些几乎全部消失。5.1 从单卡到多卡你要改的只有Trainer参数下面这段代码在1张卡上跑和8张卡上跑唯一要改的地方是devices:trainer pl.Trainer( max_epochs20, acceleratorgpu, devices8, strategyddp, )你甚至不需要写if torch.cuda.device_count() 1 else这种丑陋的分支判断Lightning在内部处理好了rank分配、通信初始化、模型复制和梯度同步。5.2 多卡训练下batch_size的含义变化这个点我必须专门强调因为它坑的人太多了。在多卡训练中每张卡上跑的是独立的一个batch如果你设置了Trainer(devices8)但DataLoader的batch_size32那么每张卡每步处理32个样本一次梯度更新对应的总样本数是32 * 8 256。也就是说从单卡迁到8卡等效batch size变大了8倍学习率不变的话往往需要相应调大或者缩小batch size保持变量。Lightning里对等效batch size的追踪其实是交给你自己的它并不会隐式调整学习率。要控制这种行为你可以在DataLoader里直接用batch_size也可以在Trainer上设置num_nodes和devices但更常见的是在实验设计层面统一规划。5.3 多卡下最容易翻车的部分采样器与数据重复原生DDP里每个进程用一个DistributedSampler保证每个进程看到的样本不重叠。Lightning会自动帮你处理这个但前提是你要把它封装成LightningDataModule并在fit时传入。如果你图省事直接在LightningModule里写return DataLoader(...)Lightning也能用只不过它需要靠一些魔法来判断怎么为每个进程分配数据容易出现数据重复或者漏样本。实测下来最稳妥的写法永远是把数据准备的逻辑放进LightningDataModule然后trainer.fit(model, datamoduledm)。这样Lightning对数据的掌控是完整的分布式采样器的设置也是自动的。5.4 DDP之外的选择什么时候需要DeepSpeed或FSDPLightning不止支持DDP还支持strategydeepspeed_stage_2、strategyfsdp等。我自己用下来如果模型超过10亿参数DDP的显存复用效率会迅速下降这时候deepspeed_stage_3或者FSDP是更合适的选择。Lightning把它们封装成了和DDP一样配置即可用的strategy迁移成本很低。但这里我不打算展开太多因为大模型训练本身就是另一套方法论普通实验用DDP就非常够用了。6. 那些文档不会主动告诉你的踩坑记录用任何框架踩坑都是难免的。下面这些坑是我在实际项目中真实碰到过的每一条都卡过我不少时间写出来给各位排雷。6.1save_hyperparameters与超参搜索的联动LightningModule.__init__里写self.save_hyperparameters()是一个很自然的习惯。它的作用是自动存下__init__的所有参数存到self.hparams里然后load_from_checkpoint时能用这些参数重建模型。听起来很美好但有个隐蔽的坑如果你的__init__参数里包含一个不能序列化的对象比如一个自定义的tokenizer实例save_hyperparameters会直接报错。解决办法有两个要么只存可序列化的参数要么在__init__里手动指定self.save_hyperparameters(ignore[tokenizer])。这个ignore参数是我后来整理代码时才发现的遇到过的朋友应该能理解我此刻的激动。6.2 数据预处理千万别写在__init__里第二个坑和LightningDataModule的生命周期有关。很多人习惯在DataModule.__init__里直接下载数据、做预处理这是错的。原因特别简单__init__是在实际使用数据之前调用的但Lightning会在构建Trainer、做分布式环境初始化的时候也实例化你的DataModule。如果你在__init__里做大量数据处理多卡训练时会发现每个进程都重复做一遍预处理既慢又浪费资源。正确的做法是__init__只存配置setup()里做数据下载、划分和预处理train_dataloader/val_dataloader/test_dataloader只负责组装DataLoader。6.3 把to(device)写进代码是反模式这是Lightning的兼容性问题里最容易被新手触发的一个。因为Lightning在训练时已经自动管理了设备你在自定义的forward里手动写x x.cuda()在单卡上能跑但一旦切到多卡或者CPU就直接炸穿。更合理的做法是要么彻底不写设备相关代码要么在极个别必须处理raw tensor的场景用self.device这个由Lightning提供的属性def forward(self, x): mask torch.ones_like(x).to(self.device) return x mask6.4 和原生PyTorch混用时注意torch.no_grad()和.eval()的管理Lightning在验证和测试阶段会自动设置model.eval()并进入no_grad上下文所以你在validation_step里写with torch.no_grad():其实是多余的。但反过来也意味着你不能依赖validation_step里的显式eval()状态去做一些需要梯度的事。如果你想在验证阶段仍保留某些操作需要梯度比如计算某些依赖梯度的指标需要手动用torch.enable_grad()重新开启并做好释放显存的准备。6.5 一个关于import的玄学坑pytorch_lightning和lightning的区别现在Lightning的包名经历了迭代早期是import pytorch_lightning as pl后面官方推出了新命名import lightning.pytorch as pl。网上很多教程混着用导致初学者经常遇到ModuleNotFoundError。我的建议是这俩本质同一个库但如果你用官方最新的2.x版本import lightning.pytorch是新写法旧写法也兼容。如果你在某个老项目里看到pytorch_lightning不用慌直接照着项目原样写就行如果你开新项目直接用lightning.pytorch更长远。唯一的坑是别在同一个项目里两种import混着用Lightning内部的状态管理和插件注册会因为双模块路径而出现奇怪问题。6.6 实测定次多时的训练速度变化这是我个人很关心的一个指标。有朋友会问用了Lightning训练速度会不会变慢我的实测感受是封装带来的性能损失几乎可以忽略。在同样的模型、数据、GPU配置下Lightning的训练速度和原生PyTorch脚本基本持平因为底层已经被torch.compile、DDP通信库、AMP优化等工具优化过了Lightning只是协调它们的调用没有引入额外的重计算瓶颈。真正会拖慢速度的反而是那些看不见的选择比如num_workers设太小、pin_memory没开、persistent_workers没设。这些和Lightning无关但它经常被归咎于Lightning慢其实属于数据加载的老问题。迁移到Lightning后同样的踩坑模式可能会因为封装而更难感知所以我还是建议在Trainer之外单独设置好DataLoader的参数。7. 要不要迁移的最终判断文章写到这里最后回到最初的问题PyTorch Lightning到底解决了什么问题适合哪些人不适合哪些人7.1 我的个人使用标准我用Lightning用了大约一年现在几乎所有个人项目都跑了Lightning但我也保留着一旦写C推理代码、部署代码或者研究V100以外非常特殊的硬件时就直接用原生PyTorch的习惯。给你一个我自己的参考标准该用Lightning做研究实验、深度学习课程作业、Kaggle比赛快速验证、企业内部模型训练流程标准化、需要频繁切换卡数或精度的情况下。不该用Lightning做底层推理优化、写框架级别的训练系统、研究编译器相关的黑科技、调试极其冷门的硬件性能问题。说白了Lightning的适用场景有一个共同特征你更关心做什么而不是怎么把流程跑起来。7.2 如果你打算迁移我建议的路径如果是老项目我不建议一次性把所有代码全部推翻重写风险太大。更稳妥的路径是先把一个最小的模型封装成LightningModule配好DataModule在单卡上跑通再逐步增加ModelCheckpoint、EarlyStopping、多卡、混合精度这些附加功能。每加一个功能就验证一次结果是否合理直到整个流程都稳定下来。这样每一步都有回退空间不会出现改了一晚上最后整个模型都不work的挫败感。7.3 关于生态的一点点展望Lightning现在不只是一个训练封装层了。它配套的torchmetrics、LightningDataModule、LightningFlash之类的东西正在往深度学习全流程套件方向走。我不知道未来它会不会变成类似Keras之于TensorFlow那样的事实标准但至少从目前社区活跃度、维护频率、企业采用率来看它已经是我见过的最有可能承担这个角色的PyTorch上层框架。如果现在你还在犹豫要不要学我的建议是把它当作PyTorch之后的第二工具——学它不会亏因为你学的其实是架构良好的训练流程应该怎么组织这件事这个思路即便哪一天不用Lightning了也依然能反哺写代码的方式。
返回列表