返回Skills库

PyTorch Lightning

MIT
📊 数据知识
K-Dense-AI深度学习

使用 PyTorch Lightning 组织深度学习代码:LightningModule 结构化、Trainer 多卡训练、数据管道、回调与日志及分布式训练,适用于可扩展的神经网络训练。

PyTorch Lightning

概述

PyTorch Lightning是一个深度学习框架,它组织PyTorch代码以消除样板代码,同时保持完全的灵活性。自动化训练工作流、多设备编排,并实现神经网络训练和跨多个GPU/TPU扩展的最佳实践。

安装

# 基础安装
uv pip install lightning

# 或使用 pytorch-lightning(已弃用但仍可用)
uv pip install pytorch-lightning

可选依赖项:

# 深度学习后端
uv pip install deepspeed  # DeepSpeed 支持
uv pip install fairscale  # FSDP 支持

# 日志记录器
uv pip install wandb     # Weights & Biases
uv pip install mlflow    # MLflow
uv pip install comet-ml  # Comet

# 监控
uv pip install tensorboard

使用场景

当您需要以下操作时使用此技能:

  • 使用PyTorch Lightning构建、训练或部署神经网络
  • 将PyTorch代码组织为LightningModules
  • 配置Trainer用于多GPU/TPU训练
  • 使用LightningDataModules实现数据管道
  • 使用回调、日志记录和分布式训练策略(DDP、FSDP、DeepSpeed)
  • 专业地构建深度学习项目

核心功能

1. LightningModule - 模型定义

将PyTorch模型组织为六个逻辑部分:

  1. 初始化 - __init__()setup()
  2. 训练循环 - training_step(batch, batch_idx)
  3. 验证循环 - validation_step(batch, batch_idx)
  4. 测试循环 - test_step(batch, batch_idx)
  5. 预测 - predict_step(batch, batch_idx)
  6. 优化器配置 - configure_optimizers()

快速模板参考: 请参见scripts/template_lightning_module.py获取完整的样板代码。

详细文档: 阅读references/lightning_module.md获取综合方法文档、钩子、属性和最佳实践。

2. Trainer - 训练自动化

Trainer自动化训练循环、设备管理、梯度操作和回调。关键功能:

  • 多GPU/TPU支持,带有策略选择(DDP、FSDP、DeepSpeed)
  • 自动混合精度训练
  • 梯度累积和裁剪
  • 检查点和早停
  • 进度条和日志记录

快速设置参考: 请参见scripts/quick_trainer_setup.py获取常见的Trainer配置。

详细文档: 阅读references/trainer.md获取所有参数、方法和配置选项。

3. LightningDataModule - 数据管道组织

在可重用的类中封装所有数据处理步骤:

  1. prepare_data() - 下载和处理数据(单进程)
  2. setup() - 创建数据集并应用转换(每个GPU)
  3. train_dataloader() - 返回训练DataLoader
  4. val_dataloader() - 返回验证DataLoader
  5. test_dataloader() - 返回测试DataLoader

快速模板参考: 请参见scripts/template_datamodule.py获取完整的样板代码。

详细文档: 阅读references/data_module.md获取方法详情和使用模式。

4. Callbacks - 可扩展训练逻辑

在特定的训练钩子处添加自定义功能,而无需修改LightningModule。内置回调包括:

  • ModelCheckpoint - 保存最佳/最新模型
  • EarlyStopping - 当指标平稳时停止
  • LearningRateMonitor - 跟踪LR调度器变化
  • BatchSizeFinder - 自动确定最佳批量大小

详细文档: 阅读references/callbacks.md获取内置回调和自定义回调创建。

5. Logging - 实验跟踪

与多个日志平台集成:

  • TensorBoard(默认)
  • Weights & Biases(WandbLogger)
  • MLflow(MLFlowLogger)
  • Neptune(NeptuneLogger)
  • Comet(CometLogger)
  • CSV(CSVLogger)

在任何LightningModule方法中使用self.log("metric_name", value)记录指标。

详细文档: 阅读references/logging.md获取日志器设置和配置。

6. Distributed Training - 扩展到多个设备

根据模型大小选择合适的策略:

  • DDP - 适用于<500M参数的模型(ResNet、较小的Transformer)
  • FSDP - 适用于500M+参数的模型(大型Transformer,推荐给Lightning用户)
  • DeepSpeed - 适用于前沿特性和细粒度控制

配置:Trainer(strategy="ddp", accelerator="gpu", devices=4)

详细文档: 阅读references/distributed_training.md获取策略比较和配置。

7. 最佳实践

  • 设备无关代码 - 使用self.device而不是.cuda()
  • 超参数保存 - 在__init__()中使用self.save_hyperparameters()
  • 指标记录 - 使用self.log()进行跨设备自动聚合
  • 可重现性 - 使用seed_everything()Trainer(deterministic=True)
  • 调试 - 使用Trainer(fast_dev_run=True)测试1个批次

详细文档: 阅读references/best_practices.md获取常见模式和陷阱。

快速工作流

  1. 定义模型:
   class MyModel(L.LightningModule):
       def __init__(self):
           super().__init__()
           self.save_hyperparameters()
           self.model = YourNetwork()

       def training_step(self, batch, batch_idx):
           x, y = batch
           loss = F.cross_entropy(self.model(x), y)
           self.log("train_loss", loss)
           return loss

       def configure_optimizers(self):
           return torch.optim.Adam(self.parameters())
  1. 准备数据:
   # 选项1:直接使用DataLoaders
   train_loader = DataLoader(train_dataset, batch_size=32)

   # 选项2:LightningDataModule(推荐用于可重用性)
   dm = MyDataModule(batch_size=32)
  1. 训练:
   trainer = L.Trainer(max_epochs=10, accelerator="gpu", devices=2)
   trainer.fit(model, train_loader)  # 或 trainer.fit(model, datamodule=dm)

资源

scripts/

用于常见PyTorch Lightning模式的可执行Python模板:

  • template_lightning_module.py - 完整的LightningModule样板代码
  • template_datamodule.py - 完整的LightningDataModule样板代码
  • quick_trainer_setup.py - 常见Trainer配置示例

references/

每个PyTorch Lightning组件的详细文档:

  • lightning_module.md - 综合LightningModule指南(方法、钩子、属性)
  • trainer.md - Trainer配置和参数
  • data_module.md - LightningDataModule模式和方法
  • callbacks.md - 内置和自定义回调
  • logging.md - 日志器集成和使用
  • distributed_training.md - DDP、FSDP、DeepSpeed比较和设置
  • best_practices.md - 常见模式、提示和陷阱

兼容工具

Claude CodeOpenClawHermes Agent

数据来源:claude-scientific-skillsMIT 许可) | 查看上游来源

上游项目:K-Dense-AI/scientific-agent-skills / claude-scientific-skills | 收录时间:2026-08-20 | 更新:2026-08-20

本页面内容基于上游开源许可项目整理,仅供学习参考。AI铺子不对第三方内容承担责任, 详情请参阅免责声明