用于预训练和微调 AI 模型的深度学习框架。
正在部署模型? 使用 LitServe 以纯 Python 构建自定义推理服务器。
快速入门 • 示例 • PyTorch Lightning • Fabric • Lightning Cloud • 社区 • 文档
为什么选择 PyTorch Lightning?
在纯 PyTorch 中训练模型需要编写和维护大量重复的工程代码。处理反向传播、混合精度、多 GPU 和分布式训练容易出错,且往往在每个项目中重新实现。PyTorch Lightning 通过组织 PyTorch 代码来自动化这些基础设施,同时保持对模型逻辑的完全控制。你负责编写科学部分,Lightning 负责处理工程部分,并且无需更改核心代码即可从 CPU 扩展到多节点 GPU。PyTorch 专家仍然可以选择 专家级控制。
一个有趣的类比:如果 PyTorch 是 Javascript,那么 PyTorch Lightning 就是 ReactJS 或 NextJS。
正在寻找 GPU?
Lightning Cloud 是运行 PyTorch Lightning 的最简单方式,无需管理基础设施。一条命令即可开始训练,并获得 GPU、自动扩缩容、监控功能以及免费层级。无需配置云环境。
你也可以在自己的硬件或云上运行 PyTorch Lightning。
Lightning 有 2 个核心包
PyTorch Lightning:大规模训练和部署 PyTorch。
Lightning Fabric:专家控制。
Lightning 让你可以精细控制希望在 PyTorch 之上添加多少抽象层。
快速入门
安装 Lightning:
pip install lightning
高级安装选项
安装可选依赖
pip install lightning['extra']
Conda
conda install lightning -c conda-forge
安装稳定版本
从源代码安装未来版本
pip install https://github.com/Lightning-AI/lightning/archive/refs/heads/release/stable.zip -U
安装前沿版本
从源代码安装每日构建版(无保证)
pip install https://github.com/Lightning-AI/lightning/archive/refs/heads/master.zip -U
或来自测试 PyPI
pip install -iU https://test.pypi.org/simple/ pytorch-lightning
PyTorch Lightning 示例
定义训练工作流。以下是一个玩具示例(探索真实示例):
# main.py
# ! pip install torchvision
import torch, torch.nn as nn, torch.utils.data as data, torchvision as tv, torch.nn.functional as F
import lightning as L
# --------------------------------
# Step 1: Define a LightningModule
# --------------------------------
# A LightningModule (nn.Module subclass) defines a full *system*
# (ie: an LLM, diffusion model, autoencoder, or simple image classifier).
class LitAutoEncoder(L.LightningModule):
def __init__(self):
super().__init__()
self.encoder = nn.Sequential(nn.Linear(28 * 28, 128), nn.ReLU(), nn.Linear(128, 3))
self.decoder = nn.Sequential(nn.Linear(3, 128), nn.ReLU(), nn.Linear(128, 28 * 28))
def forward(self, x):
# in lightning, forward defines the prediction/inference actions
embedding = self.encoder(x)
return embedding
def training_step(self, batch, batch_idx):
# training_step defines the train loop. It is independent of forward
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)
self.log("train_loss", loss)
return loss
def configure_optimizers(self):
optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)
return optimizer
# -------------------
# Step 2: Define data
# -------------------
dataset = tv.datasets.MNIST(".", download=True, transform=tv.transforms.ToTensor())
train, val = data.random_split(dataset, [55000, 5000])
# -------------------
# Step 3: Train
# -------------------
autoencoder = LitAutoEncoder()
trainer = L.Trainer()
trainer.fit(autoencoder, data.DataLoader(train), data.DataLoader(val))
在您的终端上运行模型
pip install torchvision
python main.py
从 PyTorch 转换到 PyTorch Lightning
PyTorch Lightning 只是组织有序的 PyTorch - Lightning 将 PyTorch 代码解耦,从而将科学部分与工程部分分离。

示例
探索 PyTorch Lightning 支持的各类训练方式。预训练并微调任意类型的模型,以执行分类、分割、摘要等任意任务:
| 任务 | 描述 | 运行 |
|---|---|---|
| Hello world | 预训练 - Hello world 示例 | |
| 图像分类 | 微调 - 使用 ResNet-34 模型对汽车图像进行分类 | |
| 图像分割 | 微调 - 使用 ResNet-50 模型对图像进行分割 | |
| 目标检测 | 微调 - 使用 Faster R-CNN 模型检测目标 | |
| 文本分类 | 微调 - 文本分类器(BERT 模型) | |
| 文本摘要 | 微调 - 文本摘要(Hugging Face transformer 模型) | |
| 音频生成 | 微调 - 音频生成器(transformer 模型) | |
| LLM 微调 | 微调 - LLM(Meta Llama 3.1 8B) | |
| 图像生成 | 预训练 - 图像生成器(diffusion 模型) | |
| 推荐系统 | 训练 - 推荐系统(分解与嵌入) | |
| 时间序列预测 | 训练 - 使用 LSTM 进行时间序列预测 |
高级功能
Lightning 拥有超过 40+ 高级功能 专为大规模专业 AI 研究而设计。
以下是一些示例:
在数千块 GPU 上训练,无需修改代码
# 8 GPUs
# no code changes needed
trainer = Trainer(accelerator="gpu", devices=8)
# 256 GPUs
trainer = Trainer(accelerator="gpu", devices=8, num_nodes=32)
无需修改代码即可在 TPU 等其他加速器上训练
# no code changes needed
trainer = Trainer(accelerator="tpu", devices=8)
16-bit precision
# no code changes needed
trainer = Trainer(precision=16)
实验经理
from lightning import loggers
# litlogger
trainer = Trainer(logger=LitLogger())
# tensorboard
trainer = Trainer(logger=TensorBoardLogger("logs/"))
# weights and biases
trainer = Trainer(logger=loggers.WandbLogger())
# comet
trainer = Trainer(logger=loggers.CometLogger())
# mlflow
trainer = Trainer(logger=loggers.MLFlowLogger())
# ... and dozens more
早停
es = EarlyStopping(monitor="val_loss")
trainer = Trainer(callbacks=[es])
检查点
checkpointing = ModelCheckpoint(monitor="val_loss")
trainer = Trainer(callbacks=[checkpointing])
导出到 torchscript (JIT)(生产环境使用)
# torchscript
autoencoder = LitAutoEncoder()
torch.jit.save(autoencoder.to_torchscript(), "model.pt")
导出为 ONNX(生产环境使用)
# onnx
with tempfile.NamedTemporaryFile(suffix=".onnx", delete=False) as tmpfile:
autoencoder = LitAutoEncoder()
input_sample = torch.randn((1, 64))
autoencoder.to_onnx(tmpfile.name, input_sample, export_params=True)
os.path.isfile(tmpfile.name)
相较于非结构化 PyTorch 的优势
- 模型变得与硬件无关
- 由于工程代码被抽象化,代码更易于阅读
- 更容易复现
- 因为 Lightning 处理了棘手的工程细节,所以能减少错误
- 保留了所有灵活性(LightningModules 仍然是 PyTorch 模块),但去除了大量样板代码
- Lightning 与众多流行的机器学习工具集成了数十种功能。
- 经过每个新 PR 的严格测试。我们测试所有受支持的 PyTorch 和 Python 版本组合、所有操作系统、多 GPU 甚至 TPU。
- 极小的运行速度开销(与纯 PyTorch 相比,每个 epoch 约 300 毫秒)。
Lightning Fabric:专家级控制
在任意设备、任意规模上运行,并对 PyTorch 训练循环和扩展策略进行专家级控制。你甚至可以编写自己的 Trainer。
Fabric 专为最复杂的模型设计,例如基础模型扩展、LLM、扩散模型、transformers、强化学习、主动学习。无论规模大小。
| 需要修改的内容 | 生成的 Fabric 代码(请复制!) |
|---|---|
|
|
主要特性
轻松从 CPU 切换到 GPU(Apple Silicon、CUDA、…)、TPU、多 GPU 甚至多节点训练
# Use your available hardware
# no code changes needed
fabric = Fabric()
# Run on GPUs (CUDA or MPS)
fabric = Fabric(accelerator="gpu")
# 8 GPUs
fabric = Fabric(accelerator="gpu", devices=8)
# 256 GPUs, multi-node
fabric = Fabric(accelerator="gpu", devices=8, num_nodes=32)
# Run on TPUs
fabric = Fabric(accelerator="tpu")
开箱即用地使用最先进的分布式训练策略(DDP、FSDP、DeepSpeed)和混合精度
# Use state-of-the-art distributed training techniques
fabric = Fabric(strategy="ddp")
fabric = Fabric(strategy="deepspeed")
fabric = Fabric(strategy="fsdp")
# Switch the precision
fabric = Fabric(precision="16-mixed")
fabric = Fabric(precision="64")
所有设备逻辑样板代码均已为您处理
# no more of this!
- model.to(device)
- batch.to(device)
使用 Fabric 原语构建您自己的自定义 Trainer,用于训练、检查点、日志记录等
import lightning as L
class MyCustomTrainer:
def __init__(self, accelerator="auto", strategy="auto", devices="auto", precision="32-true"):
self.fabric = L.Fabric(accelerator=accelerator, strategy=strategy, devices=devices, precision=precision)
def fit(self, model, optimizer, dataloader, max_epochs):
self.fabric.launch()
model, optimizer = self.fabric.setup(model, optimizer)
dataloader = self.fabric.setup_dataloaders(dataloader)
model.train()
for epoch in range(max_epochs):
for batch in dataloader:
input, target = batch
optimizer.zero_grad()
output = model(input)
loss = loss_fn(output, target)
self.fabric.backward(loss)
optimizer.step()
您可以在我们的 examples
示例
自监督学习
卷积架构
强化学习
GANs
经典机器学习
持续集成
Lightning 在多种 CPU、GPU 和 TPU 上,并针对主要 Python 和 PyTorch 版本进行了严格测试。
*Codecov 覆盖率 > 90%+,但构建延迟可能导致显示值较低
当前构建状态
| 系统 / PyTorch 版本 | 1.13 | 2.0 | 2.1 |
|---|---|---|---|
| Linux py3.9 [GPUs] | |||
| Linux (multiple Python versions) | |||
| OSX (multiple Python versions) | |||
| Windows (multiple Python versions) |
社区
The lightning community 由
- 10+ 核心贡献者,他们均为来自顶级 AI 实验室的职业工程师、研究科学家和博士生的混合群体。
- 800+ 社区贡献者。
想要帮助我们构建 Lightning 并为数千名研究人员减少样板代码?在此了解如何进行你的第一次贡献
Lightning 也是 PyTorch 生态系统 的一部分,该生态系统要求项目具备完善的测试、文档和支持。
寻求帮助
如果您有任何问题,请: