ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

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」系列的高级篇面向希望深度掌控训练过程可视化的开发者你将学会直接使用内置的TQDMProgressBar与RichProgressBar通过子类化定制它们的显示行为从零继承ProgressBar构建完全属于自己的进度条并基于Logger抽象基类把任意实验管理器实验追踪系统无缝集成进 Lightning 的训练流程。文中所有结论均与当前仓库源码src/lightning/pytorch/callbacks/progress/、src/lightning/pytorch/loggers/logger.py及测试用例tests/tests_pytorch/callbacks/progress/一一对应保证可直接复制运行。本文对应的原始文档为 logging_expert.rst属于可视化专题下的进阶内容建议先阅读 logging_basic.rst 与 logging_intermediate.rst 掌握基础用法再进入本专题。更换进度条从内置组件开始PyTorch Lightning 的进度条本质上是一个回调Callback通过Trainer(callbacks...)注入即可生效。Lightning 默认使用TQDMProgressBar其类文档明确注明 This is the default progress bar used by Lightning它负责在终端中渲染训练过程。如果你只是想改变展示方式可以直接换成自带的进度条也可以继承后自定义。使用 TQDMProgressBar将TQDMProgressBar实例传入 Trainer 的callbacks参数即可from lightning.pytorch.callbacks import TQDMProgressBar trainer Trainer(callbacks[TQDMProgressBar()])从 tqdm_progress.py 的源码看TQDMProgressBar提供三个可调参数用于控制渲染行为参数类型默认值作用refresh_rateint1每隔多少个 batch 刷新一次进度条设为0可完全禁用显示process_positionint0进度条向下偏移的行数当你同时有多个进度条需要对齐展示时有用leaveboolFalse每个 epoch 结束后是否把已完成的进度条保留在终端中而不是清除此外源码在初始化时会对refresh_rate做一次_resolve_refresh_rate归一化处理在 Colab检测到COLAB_GPU环境变量上自动把默认的1提高到20以避免刷新过快导致崩溃如果设置了TQDM_MINITERS环境变量则取它与refresh_rate的较大值作为最低刷新频率。默认的TQDMProgressBar最多会渲染四种进度条sanity check验证集预检、训练进度验证开始时暂停、结束时恢复且支持val_check_interval下的多次验证、验证进度、测试进度从源码的on_predict_*系列钩子可以看出预测predict阶段也有独立的进度条。对于IterableDataset这类无限数据集进度条通过convert_inf把inf转为None后由 tqdm 处理为永不结束的滚动条。使用 RichProgressBar彩色与格式化输出RichProgressBar基于 rich 库提供彩色、排版精美的进度条。首先安装依赖pip install rich然后传入 Trainerfrom lightning.pytorch.callbacks import RichProgressBar trainer Trainer(callbacks[RichProgressBar()])从 rich_progress.py 的构造函数可以看到它比 tqdm 版本多了两个参数参数类型默认值作用refresh_rateint100注意与 tqdm 版本不同这里是每秒刷新次数而非每个 batchleaveboolFalseepoch 结束后是否保留已完成进度条themeRichProgressBarTheme默认主题各组件样式console_kwargsdictNone传给 richConsole的额外参数若未安装 rich源码要求rich 10.2.2实例化会直接抛出ModuleNotFoundError并提示安装命令。另外源码注释特别提醒PyCharm 用户需要在 run/debug 配置中开启 emulate terminal 才能看到彩色渲染效果。Rich 进度条还针对IterableDataset等无长度数据集做了专门处理源码自定义了CustomBarColumn、CustomInfiniteTask与CustomProgress当任务总数为inf时以脉冲动画显示并隐藏剩余时间列避免无意义的估算。自定义 Rich 主题RichProgressBarTheme是一个 dataclass可以精确控制进度条每个组件的颜色与格式from lightning.pytorch.callbacks import RichProgressBar from lightning.pytorch.callbacks.progress.rich_progress import RichProgressBarTheme # 创建自己的主题 theme RichProgressBarTheme(descriptiongreen_yellow, progress_bargreen1) # 正常初始化 progress_bar RichProgressBar(themetheme) trainer Trainer(callbacksprogress_bar)主题支持的字段定义于 rich_progress.py及其默认样式如下字段默认值控制内容description描述文字样式如 Epoch 1、Testing 等progress_bar#6206E0进行中的进度条颜色progress_bar_finished#6206E0已完成进度条颜色progress_bar_pulse#6206E0无限数据集IterableDataset时的脉冲动画样式batch_progress批次计数如 10/50样式timedim已用时间与剩余时间样式processing_speeddim underline处理速度it/s样式metricsitalic指标文字样式metrics_text_delimiter 多个指标之间的分隔符metrics_format.3f指标数值的格式化字符串颜色取值遵循 rich 的样式语法颜色名、十六进制色值、样式组合均可各组件最终通过configure_columns方法组装成进度条布局。定制内置进度条子类化并覆写方法无论是TQDMProgressBar还是RichProgressBar官方推荐的做法都是继承后覆写其初始化方法。以 tqdm 版本为例覆写init_validation_tqdm修改验证阶段的描述文字from lightning.pytorch.callbacks import TQDMProgressBar class LitProgressBar(TQDMProgressBar): def init_validation_tqdm(self): bar super().init_validation_tqdm() bar.set_description(running validation...) return barTQDMProgressBar中可覆写的初始化方法有五个见 tqdm_progress.pyinit_sanity_tqdm()验证预检阶段的进度条init_train_tqdm()训练阶段的进度条init_validation_tqdm()验证阶段的进度条init_test_tqdm()测试阶段的进度条init_predict_tqdm()预测阶段的进度条这些方法都会用统一的BAR_FORMAT{l_bar}{bar}| {n_fmt}/{total_fmt} [{elapsed}{remaining}, {rate_noinv_fmt}{postfix}]和内部Tqdm类构造进度条。注意源码中Tqdm类覆写了format_num给浮点数/字符串补零对齐防止进度条数值变化时左右抖动。除了进度条初始化还可以覆写get_metrics控制进度条末尾展示哪些指标。基类 progress_bar.py 的文档给出了一个实用范例——去掉默认显示的版本号v_numdef get_metrics(self, trainer, model): # 不显示版本号 items super().get_metrics(trainer, model) items.pop(v_num, None) return items基类默认的get_metrics逻辑是把 Trainer 收集到的progress_bar_metrics即训练中调用self.log(..., prog_barTrue)记录的指标与get_standard_metrics目前只包含 logger 的版本号形如Epoch 1: 4%|▎ | 40/1095 [00:0301:37, 10.84it/s, v_num10]合并若两者出现同名键会通过rank_zero_warn发出警告提示prog_barTrue的指标会覆盖标准指标。构建你自己的进度条继承 ProgressBar如果内置的两个进度条都无法满足需求可以继承抽象基类ProgressBar它同时是一个Callback会自动接入 Trainer 的钩子体系。官方文档示例from lightning.pytorch.callbacks import ProgressBar class LitProgressBar(ProgressBar): def __init__(self): super().__init__() # 别忘了调用父类初始化 :) self.enable True def disable(self): self.enable False def on_train_batch_end(self, trainer, pl_module, outputs, batch_idx): super().on_train_batch_end(trainer, pl_module, outputs, batch_idx) # 别忘了调用父类实现 :) percent (self.train_batch_idx / self.total_train_batches) * 100 sys.stdout.flush() sys.stdout.write(f{percent:.01f} percent complete \r) bar LitProgressBar() trainer Trainer(callbacks[bar])这段代码在on_train_batch_end钩子中计算当前批次占比并直接写入stdout实现一个极简的百分比进度输出。基类 progress_bar.py 为自定义进度条提供了以下关键基础设施total_train_batches/total_val_batches/total_test_batches_current_dataloader/total_predict_batches_current_dataloader各类阶段的批次总数属性。它们会随 epoch 变化例如max_epochs -1且设置了max_steps时训练总数会取剩余步数并支持多 dataloader 场景无限数据集返回inf。自定义进度条应通过这些属性确定总迭代数而不是自行硬编码。disable()/enable()基类将这两个方法声明为抽象接口默认raise NotImplementedError要求实现类提供禁用/启用能力——例如学习率查找Learning Rate Finder等预训练流程会临时开关进度条。在setup钩子中基类会自动检测trainer.is_global_zero非 rank 0 进程直接调用disable()即多卡训练时进度条只在主进程显示。print(*args, **kwargs)在不破坏进度条布局的前提下打印日志tqdm 版本会转发给当前活跃进度条的write方法。has_dataloader_changed(dataloader_idx)/reset_dataloader_idx_tracker()多 dataloader 切换时的状态跟踪辅助方法。仓库针对该机制有完整的单元测试例如 test_tqdm_progress_bar.py 通过MockTqdm记录了进度条n、total、description的变化序列验证了批次推进、epoch 描述、多验证等行为test_rich_progress_bar.py 则覆盖了 Rich 进度条的主题、刷新与异常场景。编写自定义进度条后可参考这些测试验证自己的实现。集成新的实验管理器继承 LoggerLightning 中所有实验管理器logger都基于Logger抽象基类。要接入一套全新的实验追踪系统只需继承 logger.py 中的Logger并实现必要的抽象接口from lightning.pytorch.loggers import Logger class LitLogger(Logger): property def name(self) - str: return my-experiment property def version(self): return version_0 def log_metrics(self, metrics, stepNone): print(my logged metrics, metrics) def log_hyperparams(self, params, *args, **kwargs): print(my logged hyperparameters, params)Logger 的抽象接口Logger定义于 src/lightning/pytorch/loggers/logger.py继承自 Fabric 的Loggersrc/lightning/fabric/loggers/logger.py其中被声明为抽象方法、必须实现的成员有name属性实验名称。version属性实验版本号可为整数或字符串。log_metrics(metrics, stepNone)记录指标。metrics是以指标名为键、数值为值的字典step为该指标对应的步数训练循环会持续调用它务必实现为收到即写。log_hyperparams(params, *args, **kwargs)记录超参数。params可以是argparse.Namespace或字典*args/**kwargs取决于具体 logger 的扩展需求。除了抽象方法基类还提供了一系列可选覆写的钩子集成时按需实现root_dir/log_dir/save_dir/group_separator目录相关属性。root_dir是所有版本实验的根目录log_dir是当前版本的输出目录save_dir是本地日志保存根目录不落盘则返回Nonegroup_separator是目录分组的默认分隔符/。log_graph(model, input_arrayNone)记录模型计算图。save()/finalize(status)保存与收尾基类的finalize默认调用save()status取值为成功、失败或中断等状态。after_save_checkpoint(checkpoint_callback)PyTorchLogger新增每次ModelCheckpoint保存新检查点后被调用可用于同步记录检查点信息。多卡环境与 rank_zero_experiment在分布式训练中logger 通常通过rank_zero_experiment装饰器暴露experiment属性见 src/lightning/fabric/loggers/logger.pyrank 0 返回真实的实验对象其他 rank 返回_DummyExperiment占位对象其所有方法均为空操作从而保证非主进程不会重复写入日志。集成自研追踪系统时建议沿用这一模式保护底层实验句柄。集成后的完整流程实现LitLogger后即可像内置 logger 一样使用trainer Trainer(loggerLitLogger())训练过程中self.log(loss, value)产生的指标会自动流转到log_metricssave_hyperparameters()记录的超参则会流向log_hyperparams实验名称与版本号用于组织输出目录并显示在进度条的v_num字段中。内置的 CSV、TensorBoard、WB、MLflow 等 logger位于 src/lightning/pytorch/loggers/都是该接口的成熟实现编写自定义 logger 时可对照参考。仓库还提供了DummyLoggersrc/lightning/pytorch/loggers/logger.py用于在禁用用户 logger 的特定功能时保持代码可运行。小结本文覆盖了实验可视化高级路径上的四个层次直接替换TQDMProgressBar/RichProgressBar、局部定制覆写init_*_tqdm、get_metrics等方法、完全重写继承ProgressBar并利用total_*_batches等基础设施以及生态扩展继承Logger接入任意实验管理器。每一层的实现都可在仓库源码中找到对应的类与测试用例进一步探索可继续阅读progress_bar.py进度条基类与get_standard_metricstqdm_progress.py默认 tqdm 进度条实现rich_progress.pyRich 进度条与主题定义logger.py实验管理器抽象基类test_tqdm_progress_bar.py 与 test_rich_progress_bar.py进度条行为测试若需要了解指标记录prog_barTrue、on_step等如何驱动进度条展示可结合 logging.rst 与 logging_advanced.rst 一起阅读。赞分享人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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人工智能深度学习机器学习预训练分布式训练微调PyTorch Lightning N-Bit Precision 专家指南自定义 Precision Plugin 集成新精度技术PyTorch Lightning N Bit Precision 专家指南自定义 Precision Plugin 集成新精度技术 导读 本文面向希望将自人工智能深度学习机器学习预训练分布式训练微调PyTorch Lightning 高级日志与实验可视化指南PyTorch Lightning 高级日志与实验可视化指南 前言 在深度学习模型训练过程中有效的日志记录和可视化对于监控模型性能、调试问题和优化训练流程至关人工智能深度学习机器学习预训练分布式训练微调创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表