Skip to main content

PyTorch Lightning 是用于 ML 研究人员的轻量级 PyTorch 包装器。缩放您的模型。少写样板。

项目描述

用于高性能 AI 研究的轻量级 PyTorch 包装器。缩放您的模型,而不是样板。


网站主要功能如何使用文档示例社区Grid AI许可证

PyPI - Python 版本 PyPI 状态 PyPI 状态 康达 码头工人中心 编解码器

阅读文档 松弛 执照

*Codecov > 90%+,但构建延迟可能会更少

PyTorch Lightning 只是组织了 PyTorch

Lightning 解开 PyTorch 代码,将科学与工程分离。 PT 到 PL


闪电设计理念

Lightning 使用以下原则构建 PyTorch 代码:

Lightning 将以下结构强制到您的代码中,使其可重用和可共享:

  • 研究代码(LightningModule)。
  • 工程代码(您删除,并由培训师处理)。
  • 非必要的研究代码(日志记录等......这在回调中)。
  • 数据(使用 PyTorch DataLoaders 或将它们组织到 LightningDataModule 中)。

完成此操作后,您可以在多个 GPU、TPU、CPU、IPU、HPU 甚至 16 位精度上进行训练,而无需更改代码!

只需 15 分钟即可开始


持续集成

Lightning 在多个 CPU、GPU、TPU、IPU 和 HPU 上以及针对主要 Python 和 PyTorch 版本进行了严格测试。

当前构建状态 <中心>
系统/PyTorch 版本。 1.9 1.10 1.12(最新)
Linux py3.7 [GPUs**] - - -
Linux py3.7 [TPUs***] 圈子CI - -
Linux py3.8 [IPU] 构建状态 - -
Linux py3.8 [HPU] - 构建状态 -
Linux py3.8(带康达) 测试 测试 -
Linux py3.9(带康达) - - 测试
Linux py3.{7,9} - - 测试
OSX py3.{7,9} - - 测试
Windows py3.{7,9} - - 测试
  • ** 测试在两个 NVIDIA P100 上运行
  • *** 测试在 Google GKE TPUv2/3 上运行。TPU py3.7 意味着我们支持 Colab 和 Kaggle 环境。
</center>

如何使用

步骤 0:安装

从 PyPI 简单安装

pip install pytorch-lightning

第 1 步:添加这些导入

import os
import torch
from torch import nn
import torch.nn.functional as F
from torchvision.datasets import MNIST
from torch.utils.data import DataLoader, random_split
from torchvision import transforms
import pytorch_lightning as pl

第 2 步:定义 LightningModule(nn.Module 子类)

LightningModule 定义了一个完整的系统(即:GAN、自动编码器、BERT 或简单的图像分类器)。

class LitAutoEncoder(pl.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, y = 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

注意: Training_step 定义了训练循环。Forward 定义 LightningModule 在推理/预测期间的行为方式。

第 3 步:训练!

dataset = MNIST(os.getcwd(), download=True, transform=transforms.ToTensor())
train, val = random_split(dataset, [55000, 5000])

autoencoder = LitAutoEncoder()
trainer = pl.Trainer()
trainer.fit(autoencoder, DataLoader(train), DataLoader(val))

高级功能

Lightning 拥有40 多项高级功能,专为大规模专业 AI 研究而设计。

这里有些例子:

突出显示的功能代码片段
# 8 GPUs
# no code changes needed
trainer = Trainer(max_epochs=1, accelerator="gpu", devices=8)

# 256 GPUs
trainer = Trainer(max_epochs=1, accelerator="gpu", devices=8, num_nodes=32)
Train on TPUs without code changes
# no code changes needed
trainer = Trainer(accelerator="tpu", devices=8)
16-bit precision
# no code changes needed
trainer = Trainer(precision=16)
Experiment managers
from pytorch_lightning import loggers

# 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())

# neptune
trainer = Trainer(logger=loggers.NeptuneLogger())

# ... and dozens more
EarlyStopping
es = EarlyStopping(monitor="val_loss")
trainer = Trainer(callbacks=[es])
Checkpointing
checkpointing = ModelCheckpoint(monitor="val_loss")
trainer = Trainer(callbacks=[checkpointing])
Export to torchscript (JIT) (production use)
# torchscript
autoencoder = LitAutoEncoder()
torch.jit.save(autoencoder.to_torchscript(), "model.pt")
Export to ONNX (production use)
# 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)

训练循环的专业级控制(高级用户)

对于复杂/专业级别的工作,您可以选择完全控制训练循环和优化器。

class LitAutoEncoder(pl.LightningModule):
    def __init__(self):
        super().__init__()
        self.automatic_optimization = False

    def training_step(self, batch, batch_idx):
        # access your optimizers with use_pl_optimizer=False. Default is True
        opt_a, opt_b = self.optimizers(use_pl_optimizer=True)

        loss_a = ...
        self.manual_backward(loss_a, opt_a)
        opt_a.step()
        opt_a.zero_grad()

        loss_b = ...
        self.manual_backward(loss_b, opt_b, retain_graph=True)
        self.manual_backward(loss_b, opt_b)
        opt_b.step()
        opt_b.zero_grad()

相对于非结构化 PyTorch 的优势

  • 模型变得与硬件无关
  • 代码清晰易读,因为工程代码被抽象掉了
  • 更容易复制
  • 少犯错误,因为闪电处理了棘手的工程
  • 保持所有灵活性(LightningModules 仍然是 PyTorch 模块),但删除了大量样板文件
  • Lightning 与流行的机器学习工具有数十种集成。
  • 每个新 PR 都经过严格测试。我们测试了 PyTorch 和 Python 支持版本、每个操作系统、多 GPU 甚至 TPU 的每个组合。
  • 最小的运行速度开销(与纯 PyTorch 相比,每个 epoch 大约 300 毫秒)。

闪电精简版

在 PyTorch Lightning 1.5 版本中,LightningLite 现在使您能够利用 PyTorch 闪电加速器的所有功能,而无需对您的训练循环进行任何重构。查看 文和 文档了解更多信息。


例子

你好世界
对比学习
自然语言处理
强化学习
想象
经典机器学习

社区

PyTorch Lightning 社区由

  • 10 多名核心贡献者,他们都是专业工程师、研究科学家和博士的混合体。来自顶级 AI 实验室的学生。
  • 680 多个活跃的社区贡献者。

想帮助我们为成千上万的研究人员构建 Lightning 并减少样板文件吗?在此处了解如何做出您的第一个贡献

PyTorch Lightning 也是PyTorch 生态系统的一部分,它要求项目具有可靠的测试、文档和支持。

寻求帮助

如果您有任何问题,请:

  1. 阅读文档
  2. 搜索现有讨论,或添加新问题
  3. 加入我们的 Slack 社区

项目详情