PyTorch Lightning 进阶实战:用 YAML 配置文件管理全部超参数(LightningCLI 配置驱动训练全解)
2026/9/19 10:04:16 网站建设 项目流程

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", action=ActionConfigFile, help="Path to a configuration file in json or yaml format." )

也就是说,-c--config是每个 LightningCLI 都内置的参数,它基于jsonargparseActionConfigFile实现,同时支持 JSON 与 YAML 两种格式的配置文件。

如果你还不熟悉 YAML 语法,建议先阅读 什么是 YAML 配置文件(FAQ)。


使用配置文件运行 CLI

使用 YAML 配置文件运行 CLI 非常简单:

python main.py fit --config config.yaml

其中fit是 LightningCLI 注册的子命令之一。从源码看,LightningCLI.subcommands()定义了四个内置子命令:fitvalidatetestpredict,每个子命令都有自己独立的子解析器(见 cli.py),因此配置文件中的内容需要与子命令对应。

命令行参数覆盖配置

命令行中给出的独立参数可以覆盖配置文件中的选项。例如,用配置文件启动训练,但把max_epochs覆盖为 100:

python 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_once(test_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 保存文件,避免多进程竞争写文件;
  • 防止覆盖:默认overwrite=False,如果目标目录已存在同名config.yaml会抛出RuntimeError,提示你删除旧文件、传save_config_callback=None禁用保存,或通过overwrite: True允许覆盖;
  • 跨进程同步:通过trainer.strategy.broadcast()file_existsalready_saved同步到所有 rank,保证各进程状态一致。

禁用配置保存

如果不想保存配置,实例化LightningCLI时传入save_config_callback=None即可:

cli = LightningCLI(DemoModel, BoringDataModule, save_config_callback=None)

修改保存的文件名

想要把保存的配置文件改成其他名字(如name.yaml),使用save_config_kwargs

cli = 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_none=False) # Required for proper reproducibility trainer.logger.log_hyperparams({"config": config}) cli = LightningCLI(..., save_config_callback=LoggerSaveConfigCallback)

注意这里的self.parser.dump(self.config, skip_none=False)是把解析后的完整配置序列化为字符串——保留None值是为了保证可复现性。仓库测试test_lightning_cli_logger_save_config(test_cli.py)完整验证了这一用法:配置被写入 TensorBoard 的 hparams 事件文件中,同时日志目录下不再生成config.yaml

禁用默认的 log_dir 保存行为

如果只想保留自定义的保存逻辑、不向log_dir写文件,有两种方式:

# 方式一:子类中调用 super().__init__(..., save_to_log_dir=False) class MySaveConfigCallback(SaveConfigCallback): def __init__(self, *args, **kwargs): super().__init__(*args, save_to_log_dir=False, **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: 10learning_rate: 0.02正是 DemoModel 构造函数签名中的默认值,说明--print_config输出的是解析器从类签名中提取并填充默认值后的完整配置。同时,打印输出开头会带有版本头# lightning.pytorch==<版本号>(由init_parser中的dump_header注入),测试test_lightning_cli_print_config(test_cli.py)对此有断言验证。

打印指定模型的配置

--print_config也支持与其他命令行参数组合。对于支持多个模型的 CLI(即subclass_mode_model=True的注册模式),默认情况下没有选中任何模型,因此打印出的配置不包含模型设置。要打印某个特定模型的默认配置,需要显式指定:

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_path+init_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

配置项可以是intstr这样的简单 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.PyTorchProfilerprofile_memory参数——它的类型是动态决定的,解析阶段无法获知预期类型。此时配置文件应这样写:

trainer: profiler: class_path: lightning.pytorch.profilers.PyTorchProfiler dict_kwargs: profile_memory: true

仓库测试test_pytorch_profiler_init_args(test_cli.py)验证了这一点:profile_memory会被保留在dict_kwargs中并最终生效,同时record_shapes这类能在类签名中解析的参数会被移动到init_args。类似地,CometLoggerWandbLogger等 logger 的某些参数也通过dict_kwargs传递(见 test_cli.py)。


组合多个配置文件

CLI 可以同时接收多个配置文件,它们会按顺序依次解析。假设有两个包含共同设置的配置文件:

# config_1.yaml trainer: num_epochs: 10 ... # config_2.yaml trainer: num_epochs: 20 ...

多个配置文件一起传入时,最后一个配置文件中的值生效,因此上例中num_epochs = 20

python main.py fit --config config_1.yaml --config config_2.yaml

这种"后者覆盖前者"的语义让配置组合变得非常灵活。测试test_lightning_cli_config_with_subcommandtest_lightning_cli_config_before_subcommand(test_cli.py)还展示了配置文件与子命令的位置关系:--config既可以出现在子命令之前也可以出现在之后,解析结果等价;当多个配置同时存在时,同样遵循"最后解析的生效"原则。

与默认配置文件的配合

除了显式传入--config,你还可以通过parser_kwargs为每个子命令设置default_config_files,让 CLI 在无参数时自动加载指定配置(见测试test_lightning_cli_parse_kwargs_with_subcommands,test_cli.py):

parser_kwargs = { "fit": {"default_config_files": ["fit.yaml"]}, "validate": {"default_config_files": ["validate.yaml"]}, } cli = LightningCLI(DemoModel, BoringDataModule, parser_kwargs=parser_kwargs)

使用选项组:分组配置文件

选项组也可以作为独立的配置文件传入。假设有以下三个独立的配置文件,分别对应trainermodeldata三组选项:

# 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)梳理:

  1. setup_parser:初始化解析器,注册--config/-c参数,并按subcommands()fit/validate/test/predict创建子解析器;
  2. parse_arguments:调用parser.parse_args(),读取命令行与配置文件中的参数(配置文件由ActionConfigFile注入),得到self.config
  3. _set_seed:根据seed_everything配置设置随机种子(True时自动随机选择,见 cli.py);
  4. instantiate_classes:按配置实例化modeldatatrainer,并把SaveConfigCallback追加进 trainer 的callbacks列表(除非fast_dev_run开启或显式传入save_config_callback=None);
  5. _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_callback=Nonesave_config_kwargs调整
自定义保存逻辑继承SaveConfigCallback重写save_config
复杂对象配置class_path+init_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),仅供参考

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

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

立即咨询