资讯动态

PyTorch Lightning 进阶实战:用 YAML 配置文件管理全部超参数(LightningCLI 配置驱动训练全解)

发布时间:2026/9/19 10:04:20 来源:尧图企业网站定制
PyTorch Lightning 进阶实战用 YAML 配置文件管理全部超参数LightningCLI 配置驱动训练全解【免费下载链接】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 的LightningCLI为核心系统讲解如何用 YAML 配置文件驱动训练从--config加载配置、命令行覆盖、--print_config生成配置模板到SaveConfigCallback自动保存配置实现实验可复现再到多配置文件组合与选项组用法并结合仓库源码cli.py与测试用例test_cli.py深入剖析底层机制。读完本文你将掌握一套代码与配置分离、实验一键可复现的生产级实验管理方案。前置要求本文假设你已经阅读过 组合模型与数据集中间篇对LightningCLI(DemoModel, BoringDataModule)的基本用法已有了解。为什么需要配置文件随着项目逐渐复杂化可配置的选项数量会变得非常庞大如果全部通过单个命令行参数来控制会很不方便。为此使用 LightningCLI 实现的 CLI天然支持从配置文件读取输入默认的配置文件格式为YAML。从源码看LightningCLI.init_parser()在初始化解析器时就会注册配置文件入口# src/lightning/pytorch/cli.py 中 init_parser 的实现 parser.add_argument( -c, --config, actionActionConfigFile, helpPath to a configuration file in json or yaml format. )也就是说-c与--config是每个 LightningCLI 都内置的参数它基于jsonargparse的ActionConfigFile实现同时支持 JSON 与 YAML 两种格式的配置文件。如果你还不熟悉 YAML 语法建议先阅读 什么是 YAML 配置文件FAQ。使用配置文件运行 CLI使用 YAML 配置文件运行 CLI 非常简单python main.py fit --config config.yaml其中fit是 LightningCLI 注册的子命令之一。从源码看LightningCLI.subcommands()定义了四个内置子命令fit、validate、test、predict每个子命令都有自己独立的子解析器见 cli.py因此配置文件中的内容需要与子命令对应。命令行参数覆盖配置命令行中给出的独立参数可以覆盖配置文件中的选项。例如用配置文件启动训练但把max_epochs覆盖为 100python main.py fit --config config.yaml --trainer.max_epochs 100这里体现了 LightningCLI 的解析优先级设计命令行参数 配置文件 默认值。得益于jsonargparse的点号dotted key语法--trainer.max_epochs会被解析到嵌套的trainer命名空间下无需修改配置文件即可临时调整单个超参数。关于配置中的数值键原文档示例中使用了num_epochs这样的键名。需要说明的是实际生效的键名取决于Trainer的真实参数名如max_epochs这由add_lightning_class_args自动从类签名中提取因此配置文件中应使用类签名中真实存在的参数名。自动保存配置SaveConfigCallback为了简化实验记录并保证可复现性默认情况下LightningCLI会把完整的 YAML 配置自动保存到日志目录。这意味着多次以不同超参数运行fit后每次运行各自的日志目录下都会有一个config.yaml文件lightning_logs/ ├── version_0/ │ └── config.yaml ├── version_1/ │ └── config.yaml └── version_7/ └── config.yaml这些文件可以直接用来复现实验python main.py fit --config lightning_logs/version_7/config.yaml这一机制在仓库测试test_lightning_cli_save_config_only_oncetest_cli.py中有直接验证训练结束后config.yaml文件存在且回调的already_saved标志被置为True因此后续test阶段不会重复保存。SaveConfigCallback 的底层实现自动保存配置由专用回调 SaveConfigCallback 完成它会被自动添加到Trainer中。从源码可以看到其关键行为保存位置setup()阶段通过trainer.log_dir定位日志目录见 cli.py仅全局主进程rank zero执行if trainer.is_global_zero判断保证只有 rank 0 保存文件避免多进程竞争写文件防止覆盖默认overwriteFalse如果目标目录已存在同名config.yaml会抛出RuntimeError提示你删除旧文件、传save_config_callbackNone禁用保存或通过overwrite: True允许覆盖跨进程同步通过trainer.strategy.broadcast()把file_exists与already_saved同步到所有 rank保证各进程状态一致。禁用配置保存如果不想保存配置实例化LightningCLI时传入save_config_callbackNone即可cli LightningCLI(DemoModel, BoringDataModule, save_config_callbackNone)修改保存的文件名想要把保存的配置文件改成其他名字如name.yaml使用save_config_kwargscli LightningCLI(..., save_config_kwargs{config_filename: name.yaml})扩展 SaveConfigCallback把配置写入 Logger你也可以继承SaveConfigCallback实现自定义逻辑例如在保存文件之外把配置额外写入 logger方便在 TensorBoard / WandB 等实验管理平台上查看class LoggerSaveConfigCallback(SaveConfigCallback): def save_config(self, trainer: Trainer, pl_module: LightningModule, stage: str) - None: if isinstance(trainer.logger, Logger): config self.parser.dump(self.config, skip_noneFalse) # Required for proper reproducibility trainer.logger.log_hyperparams({config: config}) cli LightningCLI(..., save_config_callbackLoggerSaveConfigCallback)注意这里的self.parser.dump(self.config, skip_noneFalse)是把解析后的完整配置序列化为字符串——保留None值是为了保证可复现性。仓库测试test_lightning_cli_logger_save_configtest_cli.py完整验证了这一用法配置被写入 TensorBoard 的 hparams 事件文件中同时日志目录下不再生成config.yaml。禁用默认的 log_dir 保存行为如果只想保留自定义的保存逻辑、不向log_dir写文件有两种方式# 方式一子类中调用 super().__init__(..., save_to_log_dirFalse) class MySaveConfigCallback(SaveConfigCallback): def __init__(self, *args, **kwargs): super().__init__(*args, save_to_log_dirFalse, **kwargs) # 方式二通过 save_config_kwargs 直接传入 cli LightningCLI(..., save_config_kwargs{save_to_log_dir: False})需要特别注意的是save_config方法只会在 rank zero 上被调用。这意味着你可以在其中自由实现自定义保存逻辑而无需担心多进程 rank 与竞态条件问题。但也正因为它只在 rank zero 运行任何集合通信collective call都会导致进程挂起等待广播如果你的自定义逻辑需要集合通信应当改为实现setup方法见 SaveConfigCallback.setup 的源码注释。用 --print_config 生成配置文件模板CLI 的--help选项可以帮助你了解有哪些可配置项以及如何使用。但是从零手写一份配置文件既耗时又容易出错。为此LightningCLI 提供了--print_config参数把当前配置打印到标准输出而不实际运行命令。以LightningCLI(DemoModel, BoringDataModule)为例执行python main.py fit --print_config会生成一份包含所有默认值的配置形如seed_everything: null trainer: logger: true ... model: out_dim: 10 learning_rate: 0.02 data: data_dir: ./ ckpt_path: null这里的out_dim: 10与learning_rate: 0.02正是 DemoModel 构造函数签名中的默认值说明--print_config输出的是解析器从类签名中提取并填充默认值后的完整配置。同时打印输出开头会带有版本头# lightning.pytorch版本号由init_parser中的dump_header注入测试test_lightning_cli_print_configtest_cli.py对此有断言验证。打印指定模型的配置--print_config也支持与其他命令行参数组合。对于支持多个模型的 CLI即subclass_mode_modelTrue的注册模式默认情况下没有选中任何模型因此打印出的配置不包含模型设置。要打印某个特定模型的默认配置需要显式指定python main.py fit --model DemoModel --print_config此时生成的配置形如seed_everything: null trainer: ... model: class_path: lightning.pytorch.demos.boring_classes.DemoModel init_args: out_dim: 10 learning_rate: 0.02 ckpt_path: null注意区别在 subclass 模式下模型配置以class_pathinit_args的形式出现。标准实验流程一个推荐的实验标准流程是# 1. 打印一份配置作为参考模板 python main.py fit --print_config config.yaml # 2. 按需修改配置可以删掉所有默认参数只保留要改的项 nano config.yaml # 3. 使用编辑后的配置开始训练 python main.py fit --config config.yaml配置项的三层结构class_path / init_args / dict_kwargs配置项可以是int、str这样的简单 Python 对象也可以是复杂对象。复杂对象由两部分组成class_path类的完整导入路径init_args传递给类构造函数的参数。例如假设模型定义如下# model.py class MyModel(L.LightningModule): def __init__(self, criterion: torch.nn.Module): self.criterion criterion那么对应的配置为model: class_path: model.MyModel init_args: criterion: class_path: torch.nn.CrossEntropyLoss init_args: reduction: mean ...LightningCLI底层使用 jsonargparse 来解析配置文件和自动创建对象因此你不需要手动写任何反序列化逻辑。从源码看instantiate_class会按class_path动态导入模块并调用构造函数见 cli.py而LightningCLI.instantiate_classes则统一完成 model、data、trainer 的实例化。便捷提示Lightning 会自动注册所有LightningModule的子类因此对它们不必须写完整的导入路径直接用类名代替即可。dict_kwargs绕过解析校验的特殊键解析器会尽力推断应该接受的参数名与类型。但总会存在一些尚未支持、或客观上无法支持的情况。为克服这些限制配置中有一个特殊键dict_kwargs其中的参数不会在解析阶段被校验但会被用于类实例化。一个典型例子是lightning.pytorch.profilers.PyTorchProfiler的profile_memory参数——它的类型是动态决定的解析阶段无法获知预期类型。此时配置文件应这样写trainer: profiler: class_path: lightning.pytorch.profilers.PyTorchProfiler dict_kwargs: profile_memory: true仓库测试test_pytorch_profiler_init_argstest_cli.py验证了这一点profile_memory会被保留在dict_kwargs中并最终生效同时record_shapes这类能在类签名中解析的参数会被移动到init_args。类似地CometLogger、WandbLogger等 logger 的某些参数也通过dict_kwargs传递见 test_cli.py。组合多个配置文件CLI 可以同时接收多个配置文件它们会按顺序依次解析。假设有两个包含共同设置的配置文件# config_1.yaml trainer: num_epochs: 10 ... # config_2.yaml trainer: num_epochs: 20 ...多个配置文件一起传入时最后一个配置文件中的值生效因此上例中num_epochs 20python main.py fit --config config_1.yaml --config config_2.yaml这种后者覆盖前者的语义让配置组合变得非常灵活。测试test_lightning_cli_config_with_subcommand、test_lightning_cli_config_before_subcommandtest_cli.py还展示了配置文件与子命令的位置关系--config既可以出现在子命令之前也可以出现在之后解析结果等价当多个配置同时存在时同样遵循最后解析的生效原则。与默认配置文件的配合除了显式传入--config你还可以通过parser_kwargs为每个子命令设置default_config_files让 CLI 在无参数时自动加载指定配置见测试test_lightning_cli_parse_kwargs_with_subcommandstest_cli.pyparser_kwargs { fit: {default_config_files: [fit.yaml]}, validate: {default_config_files: [validate.yaml]}, } cli LightningCLI(DemoModel, BoringDataModule, parser_kwargsparser_kwargs)使用选项组分组配置文件选项组也可以作为独立的配置文件传入。假设有以下三个独立的配置文件分别对应trainer、model、data三组选项# trainer.yaml num_epochs: 10 # model.yaml out_dim: 7 # data.yaml data_dir: ./data那么fit命令可以这样运行python main.py fit --trainer trainer.yaml --model model.yaml --data data.yaml [...]这就是选项组groups of options模式--trainer、--model、--data这些键分别对应add_core_arguments_to_parser注册的嵌套命名空间见 cli.py每个分组都可以单独用一个 YAML 文件提供便于团队中不同角色维护各自关心的配置片段。从源码看完整调用链理解配置驱动的全流程可以顺着LightningCLI.__init__的调用链cli.py梳理setup_parser初始化解析器注册--config/-c参数并按subcommands()为fit/validate/test/predict创建子解析器parse_arguments调用parser.parse_args()读取命令行与配置文件中的参数配置文件由ActionConfigFile注入得到self.config_set_seed根据seed_everything配置设置随机种子True时自动随机选择见 cli.pyinstantiate_classes按配置实例化model、data、trainer并把SaveConfigCallback追加进 trainer 的callbacks列表除非fast_dev_run开启或显式传入save_config_callbackNone_run_subcommand调用对应子命令方法fit/validate/test/predict执行训练流程。整个过程把配置文件 → 解析 → 实例化 → 自动保存配置 → 训练串联成一个闭环从机制上保证了实验记录与复现的一致性。总结本文围绕 LightningCLI 的配置文件能力覆盖了以下核心实践能力关键用法加载配置python main.py fit --config config.yaml命令行覆盖--trainer.max_epochs 100生成配置模板python main.py fit --print_config config.yaml自动保存配置默认保存到log_dir/config.yaml可用save_config_callbackNone或save_config_kwargs调整自定义保存逻辑继承SaveConfigCallback重写save_config复杂对象配置class_pathinit_args特殊参数用dict_kwargs绕过校验组合配置多个--config按顺序解析后者覆盖前者分组配置--trainer trainer.yaml --model model.yaml --data data.yaml这套机制让配置与代码分离、实验可复现真正落地每次实验的完整超参数集合都被固化在日志目录中配合--print_config模板生成与多配置组合能力你可以轻松构建适合专业项目的模块化、可复现的实验工作流。若需要进一步定制复杂项目的 CLI 行为可继续阅读 为复杂项目定制 CLI进阶三 与 扩展 LightningCLI专家篇。【免费下载链接】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),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价