ARTICLE DETAIL

资讯详情

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

PyTorch Lightning 1.7 升级指南:Trainer 内部重构、rank_zero 模块整合与 TBPTT 输出格式变更全解析

PyTorch Lightning 1.7 升级指南:Trainer 内部重构、rank_zero 模块整合与 TBPTT 输出格式变更全解析 人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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 1.7 版本发布说明中的 devel 破坏性变更Breaking Changes清单为主体逐一解读每一项变更的来龙去脉、迁移方法和源码依据帮助你在升级到 1.7 时避免踩坑。读完本文你将掌握rank_zero系列工具的正确导入路径、Trainer内部 Mixin 体系的重构方向、分布式后端配置方式的变化、Profiler 与精度插件的 API 迁移以及多优化器 TBPTT 场景下outputs数据格式的交换规则。PyTorch Lightning 1.7 是一次以内部架构整理为主题的版本升级。它移除了大量遗留 API、把分散的rank_zero工具统一收口、让Trainer的职责划分更清晰并调整了多优化器与 TBPTTTruncated Backpropagation Through Time组合场景下的钩子输出格式。本文以仓库内 docs/source-pytorch/upgrade/sections/1_7_devel.rst 的变更表格为主体结合 src/lightning/pytorch 下的实际源码逐条讲解之前怎么用、现在怎么改、为什么这么改。一、变更总览1.7 的破坏性变更分四类从1_7_devel.rst表格看1.7 的变更可以归纳为四个主题主题涉及内容迁移要点rank_zero工具集中化rank_zero_only/rank_zero_debug/rank_zero_info/rank_zero_warn/rank_zero_deprecation/LightningDeprecationWarning统一改从pl.utilities.rank_zero导入Trainer 内部重构Mixin 类合并、run_stage移除、call_hook改为受保护方法使用Trainer.{fit,validate,test,predict}与公开属性分布式与精度插件PL_TORCH_DISTRIBUTED_BACKEND环境变量、PrecisionPlugin检查点钩子改用策略构造参数与load_state_dict钩子输出格式多优化器 TBPTT 下的outputs维度顺序交换维度或临时加new_formatTrue下文按此四条主线展开每条都给出可直接复制的迁移前后代码对照。二、rank_zero 系列工具集中到pl.utilities.rank_zero1.7 中最大的一类变更是把散落在pl.utilities.distributed与pl.utilities.warnings中的 rank 相关工具统一收口到pl.utilities.rank_zero。完整映射如下旧导入路径新导入路径pl.utilities.distributed.rank_zero_onlypl.utilities.rank_zero.rank_zero_onlypl.utilities.distributed.rank_zero_debugpl.utilities.rank_zero.rank_zero_debugpl.utilities.distributed.rank_zero_infopl.utilities.rank_zero.rank_zero_infopl.utilities.warnings.rank_zero_warnpl.utilities.rank_zero.rank_zero_warnpl.utilities.warnings.rank_zero_deprecationpl.utilities.rank_zero.rank_zero_deprecationpl.utilities.warnings.LightningDeprecationWarningpl.utilities.rank_zero.LightningDeprecationWarning迁移示例旧写法from pl.utilities.distributed import rank_zero_only, rank_zero_info from pl.utilities.warnings import rank_zero_warn, LightningDeprecationWarning改为新写法from lightning.pytorch.utilities.rank_zero import ( rank_zero_only, rank_zero_info, rank_zero_warn, LightningDeprecationWarning, )从当前源码看这次收口是物理层面的整合而非简单转发。src/lightning/pytorch/utilities/rank_zero.py 是所有工具的最终出口它从lightning.fabric.utilities.rank_zero重新导出LightningDeprecationWarning、rank_zero_debug、rank_zero_deprecation、rank_zero_info、rank_zero_only、rank_zero_warn等并配置了 PL 自己的日志器# src/lightning/pytorch/utilities/rank_zero.py from lightning.fabric.utilities.rank_zero import ( LightningDeprecationWarning, WarningCache, rank_prefixed_message, rank_zero_debug, rank_zero_deprecation, rank_zero_info, rank_zero_module, rank_zero_only, rank_zero_warn, ) rank_zero_module.log logging.getLogger(__name__)而rank_zero_only的真正实现位于 src/lightning/fabric/utilities/rank_zero.py它通过读取RANK、LOCAL_RANK、SLURM_PROCID、JSM_NAMESPACE_RANK等环境变量确定当前进程 rank其中LOCAL_RANK优先于SLURM_PROCID判断防止 SLURM 托管环境下误判再把rank_zero_only.rank设定为实际 rank保证装饰器在 0 号进程之外静默跳过。LightningDeprecationWarning也定义在同一文件中并被注册为rank_zero_deprecation的默认警告类别class LightningDeprecationWarning(DeprecationWarning): Deprecation warnings raised by Lightning. rank_zero_module.rank_zero_deprecation_category LightningDeprecationWarning也就是说pl.utilities.warnings在 1.7 之后只保留PossibleUserWarning一类用户警告见 src/lightning/pytorch/utilities/warnings.py所有与 rank 相关的打印、调试、警告与弃用提示都应改从pl.utilities.rank_zero导入。仓库的 src/lightning/pytorch/CHANGELOG.md 中 1247-1252 行也明确列出了这六项替换。三、Trainer 内部重构公开 API 收敛Mixin 体系拆除1.7 对Trainer做了一次大规模瘦身把历史遗留的 Mixin 基类合并进Trainer本体并移除了一批长期废弃的方法与属性。3.1Trainer.run_stage退役改用具体方法通用入口Trainer.run_stage被移除必须根据目的调用专用方法Trainer.fit、Trainer.validate、Trainer.test、Trainer.predict。从 src/lightning/pytorch/trainer/trainer.py 可以看到fit内部实际调用的是self._run_stage()即带下划线的内部实现run_stage只是 1.7 之前暴露给外部的一个通用包装如今外部使用者不再需要它直接使用语义清晰的具体方法即可。3.2 三个 Mixin 合并进 Trainer旧版本中Trainer的能力由多个 Mixin 拼装而成1.7 将其拆除旧派生基类新的实现位置TrainerCallbackHookMixin直接使用Trainer基类TrainerOptimizersMixinsrc/lightning/pytorch/core/optimizer.pyTrainerDataLoadingMixinTrainer方法与 src/lightning/pytorch/trainer/connectors/data_connector.py即DataConnector如果你在自定义代码中isinstance(trainer, TrainerCallbackHookMixin)或从这些 Mixin 派生子类需要改为依赖Trainer本身与core/optimizer.py中的LightningOptimizer类。这属于内部实现细节的收敛一般用户代码中较少直接接触。3.3 属性迁移device_ids、root_device、移除项Trainer 上多个属性发生变化Trainer.data_parallel_device_ids→Trainer.device_ids属性改名。当前Trainer.device_ids在 src/lightning/pytorch/trainer/trainer.py 中实现它会优先返回策略暴露的并行设备列表否则回退到[self.strategy.root_device]并据此派生num_devices。Trainer.root_gpu→Trainer.strategy.root_device.indexGPU 场景下需要主 GPU 编号时不再直接读 Trainer 属性而是通过策略层获取trainer.strategy.root_device.index。从 src/lightning/pytorch/strategies/ddp.py 可以看到root_device定义为self.parallel_devices[self.local_rank]即当前进程对应的那台设备。Trainer.should_rank_save_checkpoint直接移除不再有任何替代。Trainer.lightning_optimizers→ 使用Strategy及其属性优化器的查询与维护职责移交给策略对象。3.4Trainer.call_hook变为受保护方法禁止外部调用旧的Trainer.call_hook被拆分并改为内部方法不应在用户代码中使用Trainer._call_callback_hooksTrainer._call_lightning_module_hookTrainer._call_ttp_hookTrainer._call_accelerator_hook当前源码中这些调用统一收口在 src/lightning/pytorch/trainer/call.py例如其中定义了_call_callback_hooks(trainer, hook_name, ...)、_call_lightning_module_hook(trainer, hook_name, ...)以及_call_strategy_hook等src/lightning/pytorch/plugins/precision/precision.py 的pre_backward也通过call._call_callback_hooks(...)/call._call_lightning_module_hook(...)触发钩子。以_开头即表明这些是私有契约未来可能继续变动自定义回调或插件应通过标准钩子接口工作而不是直接调用它们。3.5 AMP 与verbose_evaluate的处置Trainer.use_amp/LightningModule.use_amp两处布尔属性都被移除混合精度完全交给 PyTorch 原生 AMPtorch.autocast/torch.amp。需要判断是否使用 AMP 时应查询当前精度插件trainer.precision_plugin的配置而不是依赖被移除的use_amp。Trainer.verbose_evaluate被移除评估循环的详细输出改由循环构造器控制EvaluationLoop(verbose...)。3.6Trainer.get_deprecated_arg_names()移除这个专用于收集过期构造参数名的遗留方法在 1.7 被直接删除。Trainer构造参数的校验逻辑早已迁入连接器connector体系不再需要这个通用方法。四、分布式后端配置环境变量让位于策略构造参数1.7 移除了通过环境变量PL_TORCH_DISTRIBUTED_BACKEND指定分布式后端的遗留方式改为在策略构造函数中显式传入process_group_backend参数。旧写法1.7 之前靠环境变量export PL_TORCH_DISTRIBUTED_BACKENDnccl python train.py新写法1.7 起在策略构造器里配置from lightning.pytorch import Trainer from lightning.pytorch.strategies import DDPStrategy strategy DDPStrategy(process_group_backendnccl) trainer Trainer(strategystrategy)从源码看process_group_backend已作为一等构造参数被多个策略支持src/lightning/pytorch/strategies/ddp.py 的DDPStrategy.__init__接收process_group_backend: Optional[str] None并保存在self._process_group_backend随后在_get_process_group_backend中回退到_get_default_process_group_backend_for_device(self.root_device)即根据设备类型推断默认后端见 src/lightning/pytorch/strategies/ddp.py。同样的参数也出现在DeepSpeedStrategy、FSDPStrategy、ModelParallelStrategy见 src/lightning/pytorch/strategies/deepspeed.py、src/lightning/pytorch/strategies/fsdp.py、src/lightning/pytorch/strategies/model_parallel.py中。CHANGELOG 中 src/lightning/pytorch/CHANGELOG.md 亦记录了该环境变量方式的移除。这样做的好处是后端选择成为策略对象的显式配置与分布式环境、启动方式解耦配置一目了然且可在同一进程内为不同策略指定不同后端。五、精度插件检查点钩子改为load_state_dictPrecisionPlugin现称Precision的检查点相关钩子发生变更旧钩子新接口PrecisionPlugin.on_load_checkpointPrecisionPlugin.load_state_dict(state_dict)PrecisionPlugin.on_save_checkpointPrecisionPlugin.state_dict()即把加载 / 保存检查点的职责从钩子形式统一为标准的state_dict/load_state_dict协议与 PyTorch 模块的惯例保持一致。当前基类 src/lightning/pytorch/plugins/precision/precision.py 中class Precision(FabricPrecision, CheckpointHooks)而具体插件例如 AMP 插件通过state_dict()返回GradScaler的状态、load_state_dict()恢复之见 src/lightning/pytorch/plugins/precision/amp.pyoverride def state_dict(self) - dict[str, Any]: if self.scaler is not None: return self.scaler.state_dict() return {} override def load_state_dict(self, state_dict: dict[str, Any]) - None: if self.scaler is not None: self.scaler.load_state_dict(state_dict)升级要点如果你自定义了精度插件并覆写过on_load_checkpoint/on_save_checkpoint请把它们重写为load_state_dict/state_dict。六、性能分析器Profiler基类合并profile_iterable移除6.1BaseProfiler→Profiler旧的基类pytorch_lightning.profiler.BaseProfiler被合并为pytorch_lightning.profiler.Profiler。当前 src/lightning/pytorch/profilers/profiler.py 中的Profiler是一个抽象基类定义了start(action_name)、stop(action_name)、summary()等抽象接口并提供profile(action_name)上下文管理器、_prepare_filename、_prepare_streams、setup(stage, local_rank, log_dir)、teardown(stage)、describe()等通用基础设施。自定义分析器应继承Profiler并实现start/stopfrom lightning.pytorch.profilers import Profiler class MyProfiler(Profiler): def start(self, action_name: str) - None: ... def stop(self, action_name: str) - None: ... def summary(self) - str: return MyProfiler report6.2SimpleProfiler.profile_iterable/AdvancedProfiler.profile_iterable移除这两个用于包装可迭代对象并逐项计时的辅助属性被删除。需要给可迭代对象逐项计时时应改用Profiler.profile(action_name)上下文管理器包裹循环体内的工作见 src/lightning/pytorch/profilers/profiler.pyprofiler MyProfiler() with profiler.profile(load training data): # 加载/处理单个 batch 的代码 ...SimpleProfiler与AdvancedProfiler本身继续存在src/lightning/pytorch/profilers/simple.py、src/lightning/pytorch/profilers/advanced.py只是不再提供profile_iterable。七、重点行为变更多优化器 TBPTT 下outputs的维度交换这是 1.7 中最容易在升级后静默出错的行为变更涉及两个训练钩子的outputs参数维度顺序。7.1on_train_batch_end(outputs, ...)2D 列表维度交换旧格式outputs是形状为(n_optimizers, tbptt_steps)的 2D 列表 新格式outputs形状变为(tbptt_steps, n_optimizers)即优化器维度与TBPTT 时间步维度互换。# 旧写法outputs[optimizer_idx][tbptt_step] def on_train_batch_end(self, outputs, batch, batch_idx): loss_opt0_step0 outputs[0][0] # 新写法outputs[tbptt_step][optimizer_idx] def on_train_batch_end(self, outputs, batch, batch_idx): loss_opt0_step0 outputs[0][0]注仅当同时使用多个优化器且启用TBPTT时该格式变更才生效单一优化器或未开启 TBPTT 的场景不受影响。7.2training_epoch_end(outputs)3D 列表维度交换旧格式outputs形状为(n_optimizers, n_batches, tbptt_steps) 新格式形状变为(n_batches, tbptt_steps, n_optimizers)即按batch → tbptt 步 → 优化器排列。# 旧写法outputs[optimizer_idx][batch_idx][tbptt_step] def training_epoch_end(self, outputs): ... # 新写法outputs[batch_idx][tbptt_step][optimizer_idx] def training_epoch_end(self, outputs): ...7.3 过渡期迁移开关new_formatTrue如果暂时不想改钩子内部的索引逻辑可以在钩子签名中追加new_formatTrue参数临时使用新格式def on_train_batch_end(self, outputs, batch, batch_idx, new_formatTrue): # 直接按新维度顺序 (tbptt_steps, n_optimizers) 处理 ... def training_epoch_end(self, outputs, new_formatTrue): # 直接按新维度顺序 (n_batches, tbptt_steps, n_optimizers) 处理 ...注意new_formatTrue只是一个过渡兼容开关最终都应迁移到新的维度顺序即不带该参数、直接按新格式编写。仓库 src/lightning/pytorch/CHANGELOG.md 记录了这两处格式废弃与替换的对应关系PR #12182。八、回调与设备统计DeviceStatsMonitor内部化键名前缀device_stats_monitor.prefix_metric_keys这一公开属性在 1.7 被移除/内部化。从当前源码看键名前缀逻辑已改为模块级私有函数_prefix_metric_keys(metrics_dict, prefix, separator)src/lightning/pytorch/callbacks/device_stats_monitor.py在回调内部把设备统计指标统一加上DeviceStatsMonitor.{hook_name}/前缀后交给 logger。DeviceStatsMonitor的公开 API 现在是cpu_stats与filter_keys两个构造参数src/lightning/pytorch/callbacks/device_stats_monitor.pyfrom lightning.pytorch.callbacks import DeviceStatsMonitor # 只记录 GPU 显存的峰值与当前值 device_stats DeviceStatsMonitor( filter_keys{allocated_bytes.all.current, allocated_bytes.all.peak} ) trainer Trainer(callbacks[device_stats])九、升级自查清单把上述变更整理成一张可直接对照检查的清单导入路径全局搜索utilities.distributed与utilities.warnings中的rank_zero_*与LightningDeprecationWarning统一改为from lightning.pytorch.utilities.rank_zero import ...。入口方法确认代码中没有调用Trainer.run_stage全部改用fit/validate/test/predict。Trainer 属性data_parallel_device_ids→device_idsroot_gpu→strategy.root_device.index删除对should_rank_save_checkpoint、use_amp、verbose_evaluate、lightning_optimizers的引用。私有调用不要从外部调用Trainer.call_hook自定义插件/回调走标准钩子接口。分布式后端删除PL_TORCH_DISTRIBUTED_BACKEND环境变量用法改为DDPStrategy(process_group_backend...)DeepSpeed / FSDP / ModelParallel 同理。精度插件on_load_checkpoint→load_state_dicton_save_checkpoint→state_dict。Profiler继承pytorch_lightning.profiler.Profiler删除对BaseProfiler与profile_iterable的引用改用profile()上下文管理器。TBPTT 格式多优化器 TBPTT 场景下按新维度顺序改写on_train_batch_end与training_epoch_end过渡期可加new_formatTrue。全部落实后你的代码即可平滑升级到 1.7并受益于更清晰的Trainer结构、统一的 rank 工具与更贴近 PyTorch 惯例的插件协议。若需结合版本间的历史迁移路径可继续查阅仓库内的 docs/source-pytorch/upgrade/migration_guide.rst 与 docs/source-pytorch/upgrade/sections 目录下的其他版本小节。赞分享人工智能深度学习机器学习预训练分布式训练微调【免费下载链接】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 1.7 升级指南进阶篇Strategy 重构、Callback 钩子迁移与 Trainer API 统一PyTorch Lightning 1.7 升级指南进阶篇Strategy 重构、Callback 钩子迁移与 Trainer API 统一 导读 本文档人工智能深度学习机器学习预训练分布式训练微调PyTorch Lightning 1.5 升级到 2.0 常规用户迁移指南Trainer 与回调 API 变更全解析PyTorch Lightning 1.5 升级到 2.0 常规用户迁移指南Trainer 与回调 API 变更全解析 导读 本文以官方升级文档 v1.5 常人工智能深度学习机器学习预训练分布式训练微调OSV-Scanner v1 到 v2 迁移指南CLI 变更、命令重构与输出格式升级全解析OSV Scanner v1 到 v2 迁移指南CLI 变更、命令重构与输出格式升级全解析 导读 本文以 OSV Scanner 官方迁移文档 docs/m漏洞扫描供应链安全应用安全CLI开发工具MCP 服务创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表