ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

PyTorch Lightning 实验管理进阶指南:跟踪超参数、模型拓扑与多实验管理器

PyTorch Lightning 实验管理进阶指南:跟踪超参数、模型拓扑与多实验管理器 人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】pytorch-lightningPretrain, 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.log和self.log_dict记录了标量指标如 loss、accuracy并可以在终端进度条prog_barTrue或 TensorBoard 浏览器中查看这些指标随 epoch 的变化曲线。但真实的研究与工程场景往往不止于此你需要记录生成样本的图像、训练过程的直方图、模型的拓扑计算图甚至需要把超参数与模型权重关联起来进行对比实验。这些非标量内容无法通过self.log直接表达这正是本篇进阶指南要解决的问题先选择一个支持这些能力的实验管理器Logger再直接调用该管理器自身的 API。Lightning 的统一设计是Trainer持有 logger 对象LightningModule内通过self.logger或self.loggers访问它进而拿到experiment句柄调用实验管理器特有的方法。从源码看logger与loggers是 LightningModule 的属性分别返回Trainer持有的单个 logger 与 logger 列表所有内置 logger 都继承自 Logger 基类并统一从 loggers 包 导出可通过from lightning.pytorch import loggers as pl_loggers或from lightning.pytorch.loggers import XXXLogger导入。跟踪音频、图像与其他工件Artifacts要记录直方图、模型拓扑图、图像、音频等高级内容第一步是从 Lightning 支持的多个实验管理器中任选一个将其注入Trainerfrom lightning.pytorch import loggers as pl_loggers tensorboard pl_loggers.TensorBoardLogger(save_dir) trainer Trainer(loggertensorboard)第二步是绕过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_image、add_histogram、add_figure、add_audio等SummaryWriter的全部方法。注意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。LitLoggerLitLogger 是 Lightning 官方维护的云端实验管理方案。安装与使用pip install litloggerfrom lightning.pytorch.loggers import LitLogger lit_logger LitLogger(save_dirlogs/) trainer Trainer(loggerlit_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.mlComet 提供实验跟踪、模型注册与数据集管理能力。安装与配置pip install comet-mlfrom lightning.pytorch.loggers import CometLogger comet_logger CometLogger(api_keyYOUR_COMET_API_KEY) trainer Trainer(loggercomet_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)MLflowMLflow 是开源 MLOps 平台支持实验跟踪、模型打包与模型注册。安装与配置pip install mlflowfrom lightning.pytorch.loggers import MLFlowLogger mlf_logger MLFlowLogger(experiment_namelightning_logs, tracking_urifile:./ml-runs) trainer Trainer(loggermlf_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)TensorBoardTensorBoard 是 Lightning 的默认实验管理器依赖可用时自动启用可直接安装pip install tensorboardfrom lightning.pytorch.loggers import TensorBoardLogger logger TensorBoardLogger() trainer Trainer(loggerlogger)在模块内记录图像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必填日志根目录namelightning_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 successlogger 会把self.hparams以 YAML 形式写入hparams.yamltensorboard.py 的 save/finalize 实现配合 TensorBoard 的 HPARAMS 面板即可做超参数对比与平行坐标图。Weights and BiaseswandbWB 提供强大的实验仪表盘与超参数搜索能力。安装与配置pip install wandbfrom lightning.pytorch.loggers import WandbLogger wandb_logger WandbLogger(projectMNIST, log_modelall) trainer Trainer(loggerwandb_logger) # log gradients and model topology wandb_logger.watch(model)WB 的watch方法可自动跟踪梯度与模型拓扑logall额外记录参数直方图log_freq500改变记录频率默认 100 步log_graphFalse可关闭计算图记录详见 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(keygenerated_images, images[fake_images])同时使用多个实验管理器你完全可以在一次训练中同时把指标写入多个平台把 logger 列表传给Trainer的logger参数即可。from lightning.pytorch.loggers import TensorBoardLogger, WandbLogger logger1 TensorBoardLogger() logger2 WandbLogger() trainer Trainer(logger[logger1, logger2])此时在模块内通过self.loggers复数按索引访问各自的experimentclass 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_rate、another_parameter等入参抓取并保存到self.hparams属性中实现见 hparams_mixin.py。只要你的实验管理器支持跟踪超参数这些参数就会自动出现在其仪表盘上。save_hyperparameters的完整用法还包括显式指定参数名self.save_hyperparameters(arg1, arg3)只保存列出的参数传入单个对象self.save_hyperparameters(params)其中params可以是dict、argparse.Namespace或OmegaConf对象忽略某些参数self.save_hyperparameters(ignorearg2)忽略单个或一组参数如不想记录数据路径、不可序列化的对象控制是否发送给 loggerloggerTrue默认设为False时超参数只保存在self.hparams而不会上报到实验管理器。不同管理器的超参数展示方式略有差异例如 TensorBoard 会额外生成hparams.yaml见上文WB 则支持在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(modelself, input_arrayprototype_array)源码层面的工作流程见 tensorboard.py 的 log_graph 实现值得注意输入数组可省略log_graph(modelself)会回退使用模型上的model.example_input_array属性因此你可以在LightningModule中定义self.example_input_array torch.randn(32, 1, 28, 27)来替代每次显式传input_array类型校验input_array必须是Tensor或tupletuple 表示传给forward()的位置参数否则 TensorBoard 无法追踪logger 会发出警告并跳过传输钩子输入会先经过_on_before_batch_transfer和_apply_batch_transfer_handler即模型定义的 batch transfer hooks再传给experiment.add_graph(model, input_array)完成图写入前置条件TensorBoardLogger(log_graphTrue)只有在tensorboard包可用时才会真正记录计算图构造函数中有显式的可用性检查tensorboard.py。WB 用户则可通过前文提到的wandb_logger.watch(model)同步获得模型拓扑与梯度信息。常见问题与排错思路self.logger为None只在LightningModule.__init__中访问self.logger会出现该问题因为此时Trainer尚未构建。请把对self.logger.experiment的调用移到setup、training_step等训练期方法中。TensorBoard 记录log_graph无输出检查是否安装了tensorboard包、是否设置了log_graphTrue、是否提供了input_array或example_input_array以及输入是否为 Tensor/tuple 类型。TensorBoard 不显示超参数TensorBoard 对是否包含超参数的日志格式敏感混用不同格式的旧日志会导致超参数面板失效需要删除或迁移之前保存的日志目录后重新训练。多 logger 索引错位self.loggers.experiment[i]的索引严格对应Trainer(logger[...])的传入顺序请保持两者一致。关于更基础的指标记录self.log、self.log_dict、reduce_fx归约、default_root_dir目录配置等可回看 日志基础指南 与 实验管理器总览各 Logger 类的完整 API 可继续查阅仓库源码 loggers 目录 下的tensorboard.py、wandb.py、comet.py、mlflow.py、litlogger.py等文件。赞分享人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】pytorch-lightningPretrain, 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 实验跟踪与可视化进阶指南PyTorch Lightning 实验跟踪与可视化进阶指南 前言 在深度学习项目开发过程中实验跟踪和可视化是至关重要的环节。PyTorch Lightnin人工智能深度学习机器学习预训练分布式训练微调如何用Agent Lightning与MLflow实现AI智能体训练的完整实验跟踪与模型管理如何用Agent Lightning与MLflow实现AI智能体训练的完整实验跟踪与模型管理 Agent Lightning作为一款强大的AI智能体训练框架为人工智能强化学习AI Agent大模型PyTorch图像模型实验管理MLflow跟踪完整指南PyTorch图像模型实验管理MLflow跟踪完整指南 PyTorch Image Models timm 是一个强大的深度学习模型库提供了数百种预训练的计人工智能计算机视觉深度学习预训练创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表