尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

LightningCLI 进阶实战:在同一项目中自由混合模型、数据集、优化器与学习率调度器(PyTorch Lightning)

LightningCLI 进阶实战:在同一项目中自由混合模型、数据集、优化器与学习率调度器(PyTorch Lightning) LightningCLI 进阶实战在同一项目中自由混合模型、数据集、优化器与学习率调度器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-lightningLightningCLI是 PyTorch Lightning 提供的命令行配置工具本文讲解其中级进阶能力当项目从单一模型 单一数据集演进为多个模型 多个数据集时如何在不改动代码的前提下从命令行任意组合 model、data、optimizer 与 lr_scheduler。读完本文你将掌握省略model_class/datamodule_class的子类模式、自定义优化器与调度器的运行时切换、从任意包按类名或完整导入路径选择类以及按类名查看特定帮助等完整实战技能。本文基于仓库中的官方文档 lightning_cli_intermediate_2.rst 展开并结合作者源码 cli.py 与测试用例 test_cli.py 进行底层原理印证。为什么需要任意模型 × 任意数据集的组合PyTorch Lightning 项目通常从一个模型、一个数据集起步。随着项目规模增长你会引入越来越多的模型与数据集此时最理想的状态是直接通过命令行把任意模型与任意数据集自由搭配完全不需要修改代码。# 自由组合一切 $ python main.py fit --modelGAN --dataMNIST $ python main.py fit --modelTransformer --dataMNIST如果没有LightningCLI这类配置往往需要写大量样板代码boilerplate常见形态如下# 选择模型 if args.model gan: model GAN(args.feat_dim) elif args.model transformer: model Transformer(args.feat_dim) ... # 选择数据集 if args.data MNIST: datamodule MNIST() elif args.data imagenet: datamodule Imagenet() ... # 混合它们 trainer.fit(model, datamodule)每新增一个模型或数据集就需要往if/elif分支里追加代码组合数量越多维护成本越高。官方文档强烈建议避免这种样板代码直接使用LightningCLI。从源码看LightningCLI的构造参数中model_class与datamodule_class都是可选的默认None并提供了subclass_mode_model、subclass_mode_data两个开关来控制子类模式见 cli.py。这正是实现上述自由组合能力的核心机制。前置准备安装 jsonargparse 并回顾基础用法使用LightningCLI需要额外的 Python 依赖。你可以选择安装 Lightning 的全部额外依赖pip install lightning[pytorch-extra]或者只安装LightningCLI真正依赖的jsonargparse带 signatures 支持用于从函数签名自动生成参数pip install jsonargparse[signatures]在进入本文的进阶主题前建议先阅读基础篇 lightning_cli_intermediate.rst 与 lightning_cli.rst掌握最小可用的LightningCLI(DemoModel, BoringDataModule)写法、fit/validate/test/predict四个子命令以及--model.learning_rate这类点分嵌套参数的传参方式。本文所有示例沿用的DemoModel、BoringDataModule均来自 boring_classes.py它们是 Lightning 自带的演示类。支持多个 LightningModule要支持多个模型实例化LightningCLI时省略model_class参数即可# main.py from lightning.pytorch.cli import LightningCLI from lightning.pytorch.demos.boring_classes import DemoModel, BoringDataModule class Model1(DemoModel): def configure_optimizers(self): print(⚡, using Model1, ⚡) return super().configure_optimizers() class Model2(DemoModel): def configure_optimizers(self): print(⚡, using Model2, ⚡) return super().configure_optimizers() cli LightningCLI(datamodule_classBoringDataModule)现在就能在命令行里任意选择模型# 使用 Model1 python main.py fit --model Model1 # 使用 Model2 python main.py fit --model Model2原理子类模式如何被激活从源码可以看到省略model_class与显式开启子类模式是等价的二者最终会落到同一个内部开关上cli.pyself._model_class model_class or LightningModule self.subclass_mode_model (model_class is None) or subclass_mode_model self._datamodule_class datamodule_class or LightningDataModule self.subclass_mode_data (datamodule_class is None) or subclass_mode_data也就是说只要model_classNonesubclass_mode_model就自动为TrueCLI 会以LightningModule为基类接受任何注册到解析器中的子类数据侧同理。随后add_lightning_class_args会根据该开关选择add_subclass_arguments子类模式还是add_class_arguments普通模式见 cli.py。显式限定基类的子类模式Tip除了省略model_class你还可以传入一个基类并设置subclass_mode_modelTrue。这样 CLI 只接受给定基类的子类避免在命令行中选到不相关的类起到类型约束作用。cli LightningCLI(MyModelBase, BoringDataModule, subclass_mode_modelTrue)测试用例 test_lightning_cli_config_and_subclass_mode 验证了在subclass_mode_modelTrue, subclass_mode_dataTrue时配置文件中可以通过class_path指定任意子类并且该配置会被原样保存回配置文件。支持多个 LightningDataModule与模型侧对称要支持多个数据模块实例化时省略datamodule_class参数# main.py import torch from lightning.pytorch.cli import LightningCLI from lightning.pytorch.demos.boring_classes import DemoModel, BoringDataModule class FakeDataset1(BoringDataModule): def train_dataloader(self): print(⚡, using FakeDataset1, ⚡) return torch.utils.data.DataLoader(self.random_train) class FakeDataset2(BoringDataModule): def train_dataloader(self): print(⚡, using FakeDataset2, ⚡) return torch.utils.data.DataLoader(self.random_train) cli LightningCLI(DemoModel)运行时选择数据集# 使用 FakeDataset1 python main.py fit --data FakeDataset1 # 使用 FakeDataset2 python main.py fit --data FakeDataset2同样地你可以给出基类并设置subclass_mode_dataTrue让 CLI 只接受该基类的子类数据模块cli LightningCLI(DemoModel, MyDataModuleBase, subclass_mode_dataTrue)Tip如果datamodule_classNone子类模式数据参数组不再是必填项。这是因为你可能想直接使用LightningModule内部自带的 dataloader 而不传入数据模块源码在 cli.py 中通过requiredFalse显式体现了这一设计意图。支持多个优化器torch.optim中的标准优化器开箱即用python main.py fit --optimizer AdamW如果所用优化器需要额外参数可以直接通过 CLI 追加无需改动任何代码python main.py fit --optimizer SGD --optimizer.lr0.01更进一步任何torch.optim.Optimizer的自定义子类都可以被用作 CLI 可选的优化器# main.py import torch from lightning.pytorch.cli import LightningCLI from lightning.pytorch.demos.boring_classes import DemoModel, BoringDataModule class LitAdam(torch.optim.Adam): def step(self, closure): print(⚡, using LitAdam, ⚡) super().step(closure) class FancyAdam(torch.optim.Adam): def step(self, closure): print(⚡, using FancyAdam, ⚡) super().step(closure) cli LightningCLI(DemoModel, BoringDataModule)运行时选择任意优化器# 使用 LitAdam python main.py fit --optimizer LitAdam # 使用 FancyAdam python main.py fit --optimizer FancyAdam原理优化器参数组如何注入模型LightningCLI默认auto_configure_optimizersTrue会自动注册优化器与调度器参数组。相关实现位于 cli.pyparser.add_optimizer_args((Optimizer,))以torch.optim.Optimizer为基类注册optimizer参数组用户若在自定义的add_arguments_to_parser中已注册过优化器则不会重复注册。真正把命令行选出的优化器装配到模型上靠的是_add_configure_optimizers_method_to_modelcli.py它读取解析器里link_to AUTOMATIC的优化器/调度器参数组用instantiate_class(self.model.parameters(), optimizer_init)实例化优化器随后把模型原有的configure_optimizers方法覆盖为LightningCLI.configure_optimizers的偏函数。默认实现configure_optimizerscli.py非常简洁无调度器时直接返回优化器有调度器时返回[optimizer], [lr_scheduler]或带monitor的字典。如果你的模型覆盖了configure_optimizers会收到一条警告提示该方法将被LightningCLI覆盖测试 test_cli_configure_optimizers_warning 对该行为做了断言。支持多个学习率调度器torch.optim.lr_scheduler中的标准学习率调度器同样开箱即用python main.py fit --optimizerAdam --lr_scheduler CosineAnnealingLR注意必须同时指定--optimizer--lr_scheduler才会生效——调度器是挂在优化器之上的。如果需要额外参数继续通过 CLI 追加python main.py fit --optimizerAdam --lr_schedulerReduceLROnPlateau --lr_scheduler.monitortrain_loss前提是你的训练流程中确实记录了train_loss这一指标。ReduceLROnPlateau属于按指标触发的调度器需要monitor指定被监控的指标名。值得一提的是LightningCLI 内部定义了一个ReduceLROnPlateau包装类cli.py在 PyTorch 原版基础上增加了monitor属性并在 LRSchedulerTypeTuple 中用它替换了原版从而让 CLI 能够接受ReduceLROnPlateau及其子类。测试 test_cli_reducelronplateau 验证了通过--lr_schedulerReduceLROnPlateau --lr_scheduler.monitorfoo传入后模型configure_optimizers返回的调度器确实带有monitor foo。同样任何torch.optim.lr_scheduler.LRScheduler的自定义子类都可以作为可选的调度器# main.py import torch from lightning.pytorch.cli import LightningCLI from lightning.pytorch.demos.boring_classes import DemoModel, BoringDataModule class LitLRScheduler(torch.optim.lr_scheduler.CosineAnnealingLR): def step(self): print(⚡, using LitLRScheduler, ⚡) super().step() cli LightningCLI(DemoModel, BoringDataModule)运行时选择任意调度器# LitLRScheduler python main.py fit --optimizerAdam --lr_scheduler LitLRScheduler从实现上看add_lr_scheduler_argscli.py与优化器参数组对称以LRSchedulerTypeTuple为基类注册lr_scheduler参数组默认link_toAUTOMATIC并在实例化时跳过optimizer参数由 Lightning 自动把优化器实例注入。从任意包中选择类前面各节中可被选中的自定义类都定义在运行LightningCLI的同一个 Python 文件里。如果你想只用类名选择任意包中的类只需在入口文件中导入对应包即可即便这些导入看起来没有直接被使用from lightning.pytorch.cli import LightningCLI import my_code.models # noqa: F401 import my_code.data_modules # noqa: F401 import my_code.optimizers # noqa: F401 cli LightningCLI()现在可以在命令行中使用其中任何一个类python main.py fit --model Model1 --data FakeDataset1 --optimizer LitAdam --lr_scheduler LitLRScheduler# noqa: F401注释用于避免 linter 报告导入未被使用的警告——这些导入的意义在于把类注册进 Python 的模块命名空间供 CLI 按名字解析。另外未导入的子类也可以通过完整导入路径直接选择python main.py fit --model my_code.models.Model1底层对应的实例化逻辑在instantiate_classcli.py它从形如{class_path: ..., init_args: {...}}的配置中切分出模块路径与类名通过__import__动态导入并实例化。这就是只要给得出完整模块路径类就一定能被 CLI 找到的原因。查看特定类的帮助信息当 CLI 接受多个模型或数据集时主帮助信息中不会列出某个具体类的专属参数因为此时类是运行时才确定的。为了查看某个类的具体参数帮助需要借助额外的帮助参数它们接受类名或其完整导入路径。例如python main.py fit --model.help Model1 python main.py fit --data.help FakeDataset2 python main.py fit --optimizer.help Adagrad python main.py fit --lr_scheduler.help StepLR每条命令都会打印出对应类Model1、FakeDataset2、Adagrad、StepLR的__init__参数清单、类型与默认值方便你在不改代码的情况下确认每个参数名与传参格式。测试用例对此行为有直接验证test_lightning_cli_help 断言在子类模式下fit --help的输出包含--model.help与--data.help并且--data.helpDataDirDataModule会展开出--data.data_dir或--data.init_args.data_dir这样的具体参数test_cli_help_message 断言--optimizer.helptorch.optim.Adam完整路径与--optimizer.helpAdam短类名输出的帮助信息完全一致说明两种写法等价且都可用。组合起来一个任意搭配的完整工作流综合以上各节你可以得到一个高度模块化的训练入口把模型、数据模块、优化器、调度器分别放入不同包入口文件只做导入与LightningCLI()实例化然后全部组合问题都交给命令行# 模型 A 数据集 B 自定义优化器 自定义调度器 python main.py fit --model Model1 --data FakeDataset1 --optimizer LitAdam --lr_scheduler LitLRScheduler # 同一入口换模型换数据集换优化器 python main.py fit --model my_code.models.Transformer --data MNIST --optimizer AdamW这种模式下新增一个模型、数据集、优化器或调度器都不再需要改动main.py或任何训练逻辑只需保证类可被导入同文件定义、导入包、或给出完整导入路径三选一项目的可扩展性与可复现性都得到显著提升。关于配置的保存与复现LightningCLI默认通过SaveConfigCallback在每次训练开始时把完整配置写入日志目录cli.py为多组合实验留档提供了保障。延伸阅读继续深入学习LightningCLI的配置保存、参数链接、回调注册等高级能力可以参考 lightning_cli_advanced.rst、lightning_cli_advanced_2.rst 与 lightning_cli_advanced_3.rstCLI 相关的全部测试集中在 test_cli.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创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表