PyTorch Lightning 实验管理进阶指南:跟踪超参数、模型拓扑与多实验管理器
2026/9/19 15:57:10 网站建设 项目流程
  • 人工智能
  • 深度学习
  • 机器学习
  • 预训练
  • 分布式训练
  • 微调

【免费下载链接】pytorch-lightning

Pretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.

项目地址:https://gitcode.com/gh_mirrors/py/pytorch-lightning
点击查看免费下载

导读:本文基于 PyTorch Lightning 官方文档「Track and Visualize Experiments (intermediate)」编写,面向已经掌握self.log基础指标记录的用户,系统讲解如何在 Lightning 训练循环中跟踪音频、图像、直方图等复杂工件(artifacts),如何接入 LitLogger、Comet.ml、MLflow、TensorBoard、Weights and Biases 等第三方实验管理器,如何通过save_hyperparameters自动记录超参数,以及如何利用log_graph可视化模型拓扑结构。读完本文,你将能够在一个或多个实验管理器仪表盘上完整复现一次实验的指标、超参数与计算图。


为什么需要"进阶"日志能力

在 日志基础指南 中,我们用self.logself.log_dict记录了标量指标(如 loss、accuracy),并可以在终端进度条(prog_bar=True)或 TensorBoard 浏览器中查看这些指标随 epoch 的变化曲线。

但真实的研究与工程场景往往不止于此:你需要记录生成样本的图像、训练过程的直方图、模型的拓扑计算图,甚至需要把超参数与模型权重关联起来进行对比实验。这些"非标量"内容无法通过self.log直接表达,这正是本篇进阶指南要解决的问题:先选择一个支持这些能力的实验管理器(Logger),再直接调用该管理器自身的 API

Lightning 的统一设计是:Trainer持有 logger 对象,LightningModule内通过self.logger(或self.loggers)访问它,进而拿到experiment句柄,调用实验管理器特有的方法。从源码看,loggerloggers是 LightningModule 的属性,分别返回Trainer持有的单个 logger 与 logger 列表;所有内置 logger 都继承自 Logger 基类,并统一从 loggers 包 导出,可通过from lightning.pytorch import loggers as pl_loggersfrom lightning.pytorch.loggers import XXXLogger导入。


跟踪音频、图像与其他工件(Artifacts)

要记录直方图、模型拓扑图、图像、音频等高级内容,第一步是从 Lightning 支持的多个实验管理器中任选一个,将其注入Trainer

from lightning.pytorch import loggers as pl_loggers tensorboard = pl_loggers.TensorBoardLogger(save_dir="") trainer = Trainer(logger=tensorboard)

第二步是绕过self.log,直接访问 logger 的底层实验 API。在LightningModule的任意函数或 hook 中:

def training_step(self): tensorboard = self.logger.experiment tensorboard.add_image() tensorboard.add_histogram(...) tensorboard.add_figure(...)

这里self.logger.experiment返回的是该日志后端原生的 experiment 对象。以 TensorBoard 为例,它底层封装了torch.utils.tensorboard.SummaryWriter(或 tensorboardX),因此你可以调用add_imageadd_histogramadd_figureadd_audioSummaryWriter的全部方法。

注意experiment属性在除LightningModule.__init__之外的任何函数中都可以访问——因为__init__执行时Trainer尚未创建,self.logger还不存在(可参考 module.py 中 logger 属性的实现,它在self._trainer is None时返回None)。


支持的实验管理器一览

官方文档的 supported_exp_managers.rst 详细列出了五个开箱即用的实验管理器。它们的接入模式完全一致:安装依赖 → 实例化 Logger → 传入Trainer(logger=...)→ 在模块内通过self.logger.experiment使用其原生 API。

LitLogger

LitLogger 是 Lightning 官方维护的云端实验管理方案。安装与使用:

pip install litlogger
from lightning.pytorch.loggers import LitLogger lit_logger = LitLogger(save_dir="logs/") trainer = Trainer(logger=lit_logger)

在模块内用其 API 跟踪文件类工件:

class LitModel(LightningModule): def any_lightning_module_function_or_hook(self): lit_logger = self.logger.experiment lit_logger.log_file("generated_images.txt")

Comet.ml

Comet 提供实验跟踪、模型注册与数据集管理能力。安装与配置:

pip install comet-ml
from lightning.pytorch.loggers import CometLogger comet_logger = CometLogger(api_key="YOUR_COMET_API_KEY") trainer = Trainer(logger=comet_logger)

在模块内记录图像:

class LitModel(LightningModule): def any_lightning_module_function_or_hook(self): comet = self.logger.experiment fake_images = torch.Tensor(32, 3, 28, 28) comet.add_image("generated_images", fake_images, 0)

MLflow

MLflow 是开源 MLOps 平台,支持实验跟踪、模型打包与模型注册。安装与配置:

pip install mlflow
from lightning.pytorch.loggers import MLFlowLogger mlf_logger = MLFlowLogger(experiment_name="lightning_logs", tracking_uri="file:./ml-runs") trainer = Trainer(logger=mlf_logger)

experiment_name指定实验名称,tracking_uri指定元数据存储位置(此处为本地文件./ml-runs,也可指向远程 MLflow 服务地址)。在模块内:

class LitModel(LightningModule): def any_lightning_module_function_or_hook(self): mlf_logger = self.logger.experiment fake_images = torch.Tensor(32, 3, 28, 28) mlf_logger.add_image("generated_images", fake_images, 0)

TensorBoard

TensorBoard 是 Lightning 的默认实验管理器(依赖可用时自动启用),可直接安装:

pip install tensorboard
from lightning.pytorch.loggers import TensorBoardLogger logger = TensorBoardLogger() trainer = Trainer(logger=logger)

在模块内记录图像:

class LitModel(LightningModule): def any_lightning_module_function_or_hook(self): tensorboard_logger = self.logger.experiment fake_images = torch.Tensor(32, 3, 28, 28) tensorboard_logger.add_image("generated_images", fake_images, 0)

进阶参数TensorBoardLogger的完整签名(见 tensorboard.py 源码)包含:

参数默认值说明
save_dir必填日志根目录
name"lightning_logs"实验名,日志实际保存在save_dir/name/version/下;设为空字符串则不创建按实验名区分的子目录
versionNone实验版本号。不指定时 logger 自动扫描目录并分配下一个可用版本(见_get_next_version实现,tensorboard.py);传入字符串则直接作为子目录名,否则使用version_${version}
log_graphFalse是否把计算图写入 TensorBoard,需要模型定义了example_input_array属性
default_hp_metricTrue当调用log_hyperparams而未提供 metric 时,写入一个占位指标hp_metric(否则无 metric 的超参数记录会被忽略)
prefix""加在指标 key 前的前缀字符串
sub_dirNone子目录,日志保存在save_dir/name/version/sub_dir/
**kwargs透传给SummaryWriter的额外参数,例如max_queue(刷新前待写日志队列大小)、flush_secs(自动刷新间隔秒数)

此外,训练成功结束后(status == "success"),logger 会把self.hparams以 YAML 形式写入hparams.yaml(tensorboard.py 的 save/finalize 实现),配合 TensorBoard 的 HPARAMS 面板即可做超参数对比与平行坐标图。

Weights and Biases(wandb)

W&B 提供强大的实验仪表盘与超参数搜索能力。安装与配置:

pip install wandb
from lightning.pytorch.loggers import WandbLogger wandb_logger = WandbLogger(project="MNIST", log_model="all") trainer = Trainer(logger=wandb_logger) # log gradients and model topology wandb_logger.watch(model)

W&B 的watch方法可自动跟踪梯度与模型拓扑:log="all"额外记录参数直方图,log_freq=500改变记录频率(默认 100 步),log_graph=False可关闭计算图记录(详见 wandb.py 中的 watch 文档);训练结束后可调用wandb_logger.experiment.unwatch(model)移除钩子。在模块内记录图像有两种方式:

class MyModule(LightningModule): def any_lightning_module_function_or_hook(self): wandb_logger = self.logger.experiment fake_images = torch.Tensor(32, 3, 28, 28) # Option 1 wandb_logger.log({"generated_images": [wandb.Image(fake_images, caption="...")]}) # Option 2 for specifically logging images wandb_logger.log_image(key="generated_images", images=[fake_images])

同时使用多个实验管理器

你完全可以在一次训练中同时把指标写入多个平台:把 logger 列表传给Trainerlogger参数即可。

from lightning.pytorch.loggers import TensorBoardLogger, WandbLogger logger1 = TensorBoardLogger() logger2 = WandbLogger() trainer = Trainer(logger=[logger1, logger2])

此时在模块内通过self.loggers(复数)按索引访问各自的experiment

class MyModule(LightningModule): def any_lightning_module_function_or_hook(self): tensorboard_logger = self.loggers.experiment[0] wandb_logger = self.loggers.experiment[1] fake_images = torch.Tensor(32, 3, 28, 28) tensorboard_logger.add_image("generated_images", fake_images, 0) wandb_logger.add_image("generated_images", fake_images, 0)

这也与源码中loggers属性返回 list 的设计一一对应(module.py)。需要说明的是,self.loggers.experiment实际是一个按索引取值的列表,索引顺序与传入Trainer(logger=[...])的顺序一致。


跟踪超参数

要让实验管理器自动记录超参数,只需在LightningModule.__init__中调用一次save_hyperparameters()

class MyLightningModule(LightningModule): def __init__(self, learning_rate, another_parameter, *args, **kwargs): super().__init__() self.save_hyperparameters()

其原理是:save_hyperparameters自动检查调用处所在帧的__init__签名,把learning_rateanother_parameter等入参抓取并保存到self.hparams属性中(实现见 hparams_mixin.py)。只要你的实验管理器支持跟踪超参数,这些参数就会自动出现在其仪表盘上。

save_hyperparameters的完整用法还包括:

  • 显式指定参数名self.save_hyperparameters('arg1', 'arg3')只保存列出的参数;
  • 传入单个对象self.save_hyperparameters(params),其中params可以是dictargparse.NamespaceOmegaConf对象;
  • 忽略某些参数self.save_hyperparameters(ignore='arg2'),忽略单个或一组参数(如不想记录数据路径、不可序列化的对象);
  • 控制是否发送给 loggerlogger=True(默认),设为False时超参数只保存在self.hparams而不会上报到实验管理器。

不同管理器的超参数展示方式略有差异,例如 TensorBoard 会额外生成hparams.yaml(见上文),W&B 则支持在experiment.config中追加自定义配置(参考 wandb.py 的说明)。


跟踪模型拓扑(计算图)

多个实验管理器都支持可视化模型拓扑结构。TensorBoard 的log_graph是其中最常用的方式,示例:

def any_lightning_module_function_or_hook(self): tensorboard_logger = self.logger prototype_array = torch.Tensor(32, 1, 28, 27) tensorboard_logger.log_graph(model=self, input_array=prototype_array)

源码层面的工作流程(见 tensorboard.py 的 log_graph 实现)值得注意:

  1. 输入数组可省略log_graph(model=self)会回退使用模型上的model.example_input_array属性,因此你可以在LightningModule中定义self.example_input_array = torch.randn(32, 1, 28, 27)来替代每次显式传input_array
  2. 类型校验input_array必须是Tensortuple(tuple 表示传给forward()的位置参数),否则 TensorBoard 无法追踪,logger 会发出警告并跳过;
  3. 传输钩子:输入会先经过_on_before_batch_transfer_apply_batch_transfer_handler(即模型定义的 batch transfer hooks),再传给experiment.add_graph(model, input_array)完成图写入;
  4. 前置条件TensorBoardLogger(log_graph=True)只有在tensorboard包可用时才会真正记录计算图(构造函数中有显式的可用性检查,tensorboard.py)。

W&B 用户则可通过前文提到的wandb_logger.watch(model)同步获得模型拓扑与梯度信息。


常见问题与排错思路

  • self.loggerNone:只在LightningModule.__init__中访问self.logger会出现该问题,因为此时Trainer尚未构建。请把对self.logger.experiment的调用移到setuptraining_step等训练期方法中。
  • TensorBoard 记录log_graph无输出:检查是否安装了tensorboard包、是否设置了log_graph=True、是否提供了input_arrayexample_input_array,以及输入是否为 Tensor/tuple 类型。
  • TensorBoard 不显示超参数:TensorBoard 对"是否包含超参数"的日志格式敏感,混用不同格式的旧日志会导致超参数面板失效,需要删除或迁移之前保存的日志目录后重新训练。
  • 多 logger 索引错位self.loggers.experiment[i]的索引严格对应Trainer(logger=[...])的传入顺序,请保持两者一致。

关于更基础的指标记录(self.logself.log_dictreduce_fx归约、default_root_dir目录配置等),可回看 日志基础指南 与 实验管理器总览;各 Logger 类的完整 API 可继续查阅仓库源码 loggers 目录 下的tensorboard.pywandb.pycomet.pymlflow.pylitlogger.py等文件。

  • 人工智能
  • 深度学习
  • 机器学习
  • 预训练
  • 分布式训练
  • 微调

【免费下载链接】pytorch-lightning

Pretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.

项目地址:https://gitcode.com/gh_mirrors/py/pytorch-lightning
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询