用pytorch lightning简化模型训练流程!
引言
前面几篇文章,只要一涉及模型训练,总少不了要写一段重复的循环代码——加载数据、前向传播、计算梯度、更新参数。说真的,每次都是那几行,写久了确实有点审美疲劳。

你可以自己封装一个train方法,或者干脆换一条路——用PyTorch Lightning把这些琐事丢给框架去处理,自己专心折腾模型的结构和逻辑。后者的好处是,代码更清爽,维护起来也省心。
什么是PyTorch Lightning?
PyTorch Lightning是一个高级库,可以看作是在PyTorch之上又做了一层封装。它把那些流程化的、重复的部分(比如训练循环、日志记录、检查点)统一打包,让项目结构更清晰,也更容易复用。
在原生PyTorch里,训练代码往往散落在各个文件里,维护起来挺头疼。Lightning把这些功能组织成标准的接口,还内置了分布式训练的支持——多GPU、TPU、甚至多节点并行,都不用自己手动处理数据划分和梯度聚合。
另外,它集成了TensorBoard等可视化工具,训练过程中的指标监控和结果查看都变得很方便。
不用PyTorch Lightning时,如何训练模型?
如果不借助Lightning,训练模型就得直接跟PyTorch的API打交道。定义模型、损失函数、优化器,再手写训练和验证循环。下面是一个典型的例子,注意那个for循环部分——代码虽然能跑,但结构上总感觉有点“糙”。
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
# 假设我们有一个自定义的数据集和模型
dataset = ... # 你的数据集
model = ... # 你的模型,继承自 nn.Module
loss_function = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# 数据加载
dataloader = DataLoader(dataset, batch_size=64, shuffle=True)
# 训练循环
for epoch in range(num_epochs):
for inputs, targets in dataloader:
optimizer.zero_grad() # 清零梯度
outputs = model(inputs) # 前向传播
loss = loss_function(outputs, targets) # 计算损失
loss.backward() # 反向传播
optimizer.step() # 更新模型参数
print(f"Epoch {epoch}, Loss: {loss.item()}")
PyTorch Lightning如何简化并结构化训练代码?
Lightning引入了几层新的抽象,把训练流程拆成几个模块化的部分:
- :定义模型的结构,以及训练、验证、测试的逻辑。
LightningModule
- :管理整个训练过程——数据加载、优化器、训练循环、检查点等。
Trainer
- :在训练的不同阶段插入自定义操作,比如日志记录或模型保存。
Callbacks
- :专门处理数据加载和预处理,让数据和模型解耦。
LightningDataModule
这样一来,模型、数据、训练逻辑被分到不同的模块里,代码结构变得清晰、模块化,而且很多复杂的任务(比如分布式训练、日志记录、可视化)都内置好了,不用自己从头造轮子。
还是用刚才那个PyTorch的例子,换成Lightning后,代码会优雅很多。核心逻辑保持不变,但整体框架一下子规整了。
import pytorch_lightning as pl
如果你有多张GPU卡,或者需要多节点分布式训练,直接在Trainer里指定设备ID或分布式环境变量就行,模型代码根本不用改。
# 单节点多GPU训练
trainer = Trainer(gpus=4)
# 多节点分布式训练
trainer = Trainer(accelerator="ddp")
# 执行训练前,需要配置训练节点,设置主节点和专用端口
# 启动训练时指定主节点,例如:
python -m torch.distributed.launch --nproc_per_node=NUM_GPUS_PER_NODE
--nnodes=NUM_NODES --node_rank=NODE_RANK
--master_addr=MASTER_ADDR --master_port=MASTER_PORT train.py
注意,启动多节点训练时,需要搭配torch.distributed.launch并正确配置节点参数。这些细节Lightning已经帮你封装好,你只需要在Trainer里声明一句accelerator="ddp"就行。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名