☰
从PyTorch到Lightning:重构深度学习训练流程的实战解析
2026/10/5 8:28:50 网站建设 项目流程

从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(), lr=1e-3) model = MyModel() trainer = pl.Trainer(max_epochs=10) 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_epoch=True, prog_bar=True) 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_size=64, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=64) 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(f"epoch {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(dim=1) correct += pred.eq(y).sum().item() total += y.size(0) print(f"val 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_size=64, learning_rate=2e-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(task="multiclass", num_classes=10) 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_bar=True) self.log("val_acc", self.val_acc, prog_bar=True) def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lr=self.hparams.learning_rate)

注意上面代码里我用了torchmetrics.Accuracy,这是Lightning生态里配套的指标库,你不需要自己手写acc的计算逻辑。所有指标的最后聚合、多卡同步它都帮你处理好,比自己记一个total/correct再跨进程同步要省事得多。

3.3 第二步:用DataLoader或DataModule管理数据

在原生脚本里,数据加载就是两个DataLoader变量。在Lightning里你可以直接在LightningModule里实现train_dataloader和val_dataloader:

def train_dataloader(self): return DataLoader(self.train_dataset, batch_size=64, shuffle=True, num_workers=4) def val_dataloader(self): return DataLoader(self.val_dataset, batch_size=64, num_workers=4)

但更推荐的做法是单独建一个LightningDataModule。它能把"下载/清洗/划分/加载"这个过程完整沉淀下来,复用性更强:

class MNISTDataModule(pl.LightningDataModule): def __init__(self, data_dir="./data", batch_size=64): super().__init__() self.data_dir = data_dir self.batch_size = batch_size def setup(self, stage=None): dataset = MNIST(self.data_dir, train=True, download=True, transform=transforms.ToTensor()) self.train_dataset, self.val_dataset = random_split(dataset, [55000, 5000]) def train_dataloader(self): return DataLoader(self.train_dataset, batch_size=self.batch_size, num_workers=4) def val_dataloader(self): return DataLoader(self.val_dataset, batch_size=self.batch_size, num_workers=4)

3.4 第三步:用Trainer把整个流程跑起来

重头戏来了。以前需要自己手写的一大段循环和配置,现在全部收敛为一行:

dm = MNISTDataModule() model = LitMNIST(hidden_size=128, learning_rate=1e-3) trainer = pl.Trainer(max_epochs=10, accelerator="auto", devices="auto") trainer.fit(model, dm)

accelerator="auto"和devices="auto"会自动检测当前环境是CPU还是GPU,如果是GPU就自动用CUDA训练。相比原生代码里写model.cuda()再一个个搬tensor,Lightning连to(device)这一步都给你省了,因为训练时tensor自动会被放到正确设备上。这一点新上手的朋友最容易搞懵,记住一个原则:在Lightning里,你永远不需要自己在代码里调用.to(device),除非你在做非常特殊的raw tensor操作。

3.5 重构后多出来的功能:你现在白拿了什么

和原来的脚本对比,这次重构并不是单纯的"代码搬家",你白拿了至少四样东西:

  • 断点续训:Trainer的checkpoint_callback默认会保存last.ckpt,训练中断后可以直接从断点恢复继续跑。
  • EarlyStopping和ModelCheckpoint:监控验证指标,效果不好自动停,效果好的权重自动保存,全是指定式配置,不用自己写逻辑。
  • TensorBoard日志:所有self.log()的记录都会自动写进日志目录,跑完tensorboard --logdir=lightning_logs就能看曲线。
  • 多卡训练:把devices变成4,strategy="ddp",跑分布式训练几乎不需要改模型代码。

这些功能如果用原生PyTorch实现,每一样都是一大段需要反复测试的代码,而现在都集成在了一个你本来就绕不开的流程环节里。

4. 高频功能实测:checkpoint、日志、EarlyStopping和梯度裁剪的配置细节

接下来聊实际用得最多的四个功能配置。我不打算罗列所有参数,只分享我测试下来最实用的组合和踩过的坑。

4.1 ModelCheckpoint:不只是保存最好模型

ModelCheckpoint可能是Lightning里配置最繁琐但收益最高的组件。它的核心参数有monitor、mode、save_top_k、filename。我最常用的配置是:

from pytorch_lightning.callbacks import ModelCheckpoint checkpoint_callback = ModelCheckpoint( monitor="val_loss", mode="min", save_top_k=3, filename="epoch={epoch}-val_loss={val_loss:.4f}", auto_insert_metric_name=False, )

这里有两个容易被忽略的点。第一,filename里的大括号占位符必须和monitor的指标名一致,否则保存时会报错或者显示missing。第二,auto_insert_metric_name=False可以让文件名里不自动加指标名,我一般都会关掉,因为文件名太长在管理模型版本时很不方便。

如果不配ModelCheckpoint,Lightning默认只保存last.ckpt。也就是说官方模板里"自动保存最佳模型"这件事,是靠这个callback实现的,而不是Trainer参数天然带的功能,这一点文档里写得很隐晦,容易误会。

4.2 EarlyStopping:监控指标怎么选才不坑

EarlyStopping的配置本身很简单:

from pytorch_lightning.callbacks import EarlyStopping early_stop = EarlyStopping(monitor="val_loss", patience=3, mode="min")

但我想告诉大家一个真实感受:监控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_epochs=10, precision="bf16-mixed", # 或者 "16-mixed" gradient_clip_val=1.0, gradient_clip_algorithm="norm", )

precision="bf16-mixed"在Ampere及以上的NVIDIA GPU上非常实用,显存占用能降三分之一以上,训练速度通常还有提升。但注意,bf16和fp16是有区别的:bf16的范围和fp32一样,所以训练稳定性更好,多数情况下不会出现fp16那种loss突然变NaN的问题;但bf16在少数不支持它的硬件上会报错,所以老卡用户得用"16-mixed"加accumulate_grad_batches来缓一下。

梯度裁剪我用的是algorithm="norm",它按全局梯度范数裁剪,比按值裁剪(value)更常用也更安全。当年用原生PyTorch时我都是手写在backward之后step之前,现在一行配置搞定,出错概率反而更低了。

4.4 日志系统:logger是抽象层,不是某个特定后端

新手最容易卡住的地方是日志怎么配。Lightning里Trainer的logger参数可以接一个或多个Logger对象,比如TensorBoardLogger、CSVLogger、WandbLogger:

from pytorch_lightning.loggers import TensorBoardLogger, CSVLogger tb_logger = TensorBoardLogger("logs", name="mnist_experiment") csv_logger = CSVLogger("logs", name="csv_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_epochs=20, accelerator="gpu", devices=8, strategy="ddp", )

你甚至不需要写if torch.cuda.device_count() > 1 else这种丑陋的分支判断,Lightning在内部处理好了rank分配、通信初始化、模型复制和梯度同步。

5.2 多卡训练下batch_size的含义变化

这个点我必须专门强调,因为它坑的人太多了。在多卡训练中,每张卡上跑的是独立的一个batch,如果你设置了Trainer(devices=8),但DataLoader的batch_size=32,那么每张卡每步处理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, datamodule=dm)。这样Lightning对数据的掌控是完整的,分布式采样器的设置也是自动的。

5.4 DDP之外的选择:什么时候需要DeepSpeed或FSDP

Lightning不止支持DDP,还支持strategy="deepspeed_stage_2"、strategy="fsdp"等。我自己用下来,如果模型超过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 + mask

6.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了,也依然能反哺写代码的方式。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询