PyTorch Lightning 回调状态持久化:用 state_dict、load_state_dict 与 state_key 让自定义 Callback 可断点续训

PyTorch Lightning 回调状态持久化:用 state_dict、load_state_dict 与 state_key 让自定义 Callback 可断点续训 PyTorch Lightning 回调状态持久化用 state_dict、load_state_dict 与 state_key 让自定义 Callback 可断点续训【免费下载链接】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 中Callback回调常被用来承载训练过程中的运行时状态例如早停的等待轮数、最优分数、SWA 平均权重或计时器剩余时间。如果这些状态只存在内存里一旦训练中断并从 checkpoint 恢复回调就会失忆导致早停从零开始、计时器归零等行为异常。本指南以 docs/source-pytorch/extensions/callbacks_state.rst 为核心系统讲解如何通过实现Callback.state_dict()、Callback.load_state_dict()两个钩子以及为有状态回调定义唯一的state_key把回调内部状态作为模型 checkpoint 的一部分持久化使训练中断后能精确恢复。读完本文你将能够编写任意可恢复的自定义回调并理解 Lightning 底层如何按state_key归集与还原这些状态。一、为什么回调需要保存状态大多数 Lightning 回调是无状态的——它们只是监听训练事件并做出反应例如打印日志、调整学习率。但另一些回调内部维护着跨步骤、跨 epoch的累积信息例如EarlyStopping记录wait_count已连续多少个 epoch 未改善、best_score当前最优指标Timer记录已用时间用于限制训练总时长ModelCheckpoint记录best_k_models与kth_best_model_path决定何时该覆盖旧 checkpoint随机权重平均StochasticWeightAveraging、WeightAveraging记录平均模型权重。这类回调如果不保存状态那么使用Trainer.fit(ckpt_path...)断点续训时Lightning 只会恢复模型权重、优化器与学习率调度器回调内部状态却会回到初始值。其结果可能是已经等待了 20 个 epoch 的早停重新倒计时、训练计时器重新计时最终改变整个训练的行为。解决方案就是让回调自身实现两个钩子把状态交接给 checkpoint 机制state_dict()返回一个可被 pickle 序列化的字典表示回调当前状态load_state_dict(state_dict)把 checkpoint 中保存的字典还原回回调内部属性。这两个钩子的默认实现位于 callback.pystate_dict()默认返回空字典{}load_state_dict()默认什么都不做。因此只有你主动覆写它们回调状态才会被保存。注意state_dict()返回的必须是可 pickle 的对象Note that the returned state must be able to be pickled。checkpoint 最终会被torch.save写入磁盘任何不可序列化的对象如打开的 file handle、未绑定的 CUDA 设备引用都会在保存时报错。二、最小实现单实例有状态回调如果一个有状态回调在 Trainer 中只会以单实例形式使用那么实现state_dict()与load_state_dict()两个钩子就足够了。以文档中的Counter回调为例它统计训练完成了多少个 epoch 或 batchfrom lightning.pytorch.callbacks import Callback class Counter(Callback): def __init__(self, whatepochs, verboseTrue): self.what what self.verbose verbose self.state {epochs: 0, batches: 0} property def state_key(self) - str: # note: we do not include verbose here on purpose return fCounter[what{self.what}] def on_train_epoch_end(self, *args, **kwargs): if self.what epochs: self.state[epochs] 1 def on_train_batch_end(self, *args, **kwargs): if self.what batches: self.state[batches] 1 def load_state_dict(self, state_dict): self.state.update(state_dict) def state_dict(self): return self.state.copy()关键实现细节state_dict()返回的是self.state.copy()而非原始引用避免把内部可变对象直接暴露给 checkpoint 序列化流程防止状态在保存过程中被意外改动load_state_dict()用self.state.update(state_dict)合并字典比整体赋值更稳健——即使 checkpoint 中缺少某个键也不会把self.state里的其他键清空保存与恢复是一对逆操作保存什么结构恢复时就按什么结构读取。三、多实例场景为什么必须定义 state_key文档明确指出如果 Trainer 支持传入同一个回调类型的多个实例那么仅仅实现上面两个钩子是不够的还必须覆写state_key属性否则 Lightning 无法在加载时区分不同实例各自的状态。先看默认实现。在 callback.py 中基类的state_key默认返回property def state_key(self) - str: return self.__class__.__qualname__即默认 key 只是类名。这意味着如果有两个Counter实例它们的state_key都是Counter保存时后一个实例会覆盖前一个实例的状态加载时两个实例也会拿到同一份状态——状态彻底混淆。因此文档中的Counter把what该实例统计的是 epoch 还是 batch编码进 keyproperty def state_key(self) - str: # note: we do not include verbose here on purpose return fCounter[what{self.what}]这样两个实例分别获得Counter[whatepochs]与Counter[whatbatches]两个互不冲突的 key。注意注释中的刻意设计只把会影响状态语义的参数放进 key而把verbose这类纯显示参数排除在外。如果某个参数不影响状态结构却写进了state_key那么仅仅因为显示设置不同就会导致同一个逻辑回调在断点续训时匹配不上。然后像文档中这样把两个实例同时交给 Trainerfrom lightning.pytorch import Trainer # two callbacks of the same type are being used trainer Trainer(callbacks[Counter(whatepochs), Counter(whatbatches)])此时训练产生的 Lightning checkpoint 中回调状态会被归集到callbacks字段下结构与文档展示的一致{ state_dict: ..., callbacks: { Counter{what: batches}: {batches: 32, epochs: 0}, Counter{what: epochs}: {batches: 0, epochs: 2}, ... } }可以看到两个实例的状态分别挂在各自的state_key之下、互不干扰。文档同时提醒如果缺少state_key覆写默认 key 只有类名Counter两个实例的状态将无法区分——这正是多实例有状态回调最容易被忽略的坑。四、源码级原理状态如何进出 checkpoint理解了怎么写再来看 Lightning 底层怎么存、怎么取。相关实现集中在 trainer/call.py 与 checkpoint_connector.py。4.1 保存阶段Trainer 组装 checkpoint 时CheckpointConnector.dump_checkpoint()会构造包含callbacks字段的字典checkpoint_connector.py其中回调状态由_call_callbacks_state_dict()收集def _call_callbacks_state_dict(trainer: pl.Trainer) - dict[str, dict]: Called when saving a model checkpoint, calls and returns every callbacks state_dict, keyed by Callback.state_key. callback_state_dicts {} for callback in trainer.callbacks: state_dict callback.state_dict() if state_dict: callback_state_dicts[callback.state_key] state_dict return callback_state_dicts三点值得注意保存顺序是遍历trainer.callbacks逐个调用callback.state_dict()空状态不保存只有state_dict()返回了非空字典时该回调才会进入callbacks字段因此无状态回调不会污染 checkpointkey 就是callback.state_key所有实例按各自的state_key归集这与前一节的多实例机制一一对应。此外dump_checkpoint的callbacks字段仅在非weights_only模式下写入checkpoint_connector.py 的注释明确标注了callbacks: callback specific state[] # if not weights_only。如果只是为了导出权重做推理可以不携带回调状态。4.2 加载阶段加载时restore_callbacks()会依次调用_call_callbacks_on_load_checkpoint()与_call_callbacks_load_state_dict()checkpoint_connector.py。其中真正把状态写回回调实例的是后者def _call_callbacks_load_state_dict(trainer: pl.Trainer, checkpoint: dict[str, Any]) - None: Called when loading a model checkpoint, calls every callbacks load_state_dict. callback_states: Optional[dict[Union[type, str], dict]] checkpoint.get(callbacks) if callback_states is None: return for callback in trainer.callbacks: state callback_states.get(callback.state_key, callback_states.get(callback._legacy_state_key)) if state: state deepcopy(state) callback.load_state_dict(state)实现要点每个回调先按state_key精确查找自己的状态找不到时回退到_legacy_state_key用于兼容 1.5.0 之前的老 checkpoint见下文找到的状态会先deepcopy一份再交给load_state_dict避免回调在恢复过程中直接持有 checkpoint 字典内部对象的引用防止后续操作污染已加载的数据。4.3 加载时的缺失告警如果 checkpoint 中存在某个state_key但当前 Trainer 里没有对应的回调实例_call_callbacks_on_load_checkpoint()会打印告警call.pyBe aware that when using ckpt_path, callbacks used to create the checkpoint need to be provided during Trainer instantiation. Please add the following callbacks: [...]这提醒你生成 checkpoint 时所用的回调在断点续训时必须原样传入 Trainer否则对应状态无法被任何实例接收。仓库测试 test_callbacks.py 中test_resume_incomplete_callbacks_list_warning专门验证了这一行为保存时用了两个ModelCheckpoint分别监控epoch与global_step恢复时只传入其中一个就会触发Please add the following callbacks: [...]告警。4.4 旧版 checkpoint 兼容1.5.0 之前在 1.5.0 之前回调状态按类型class而不是按state_key保存。为了兼容旧 checkpoint基类保留了_legacy_state_key属性返回回调的类本身callback.py。加载时_call_callbacks_on_load_checkpoint()会读取 checkpoint 里的pytorch-lightning_version若早于1.5.0dev则按_legacy_state_key匹配call.py。测试 test_callbacks.py 中的test_resume_callback_state_saved_by_type_stateful演示了这条兼容路径一个state_key返回类本身的老式回调保存后再用新 Trainer 加载callback.state 111被正确恢复。这意味着你可以放心地让新代码继续读取旧版本产出的 checkpoint无需迁移脚本。五、仓库内的权威实践EarlyStopping 与 ModelCheckpoint 怎么写 state_key与其自己摸索不如直接参考 Lightning 官方回调的写法——它们是state_key设计的范本。5.1 EarlyStoppingEarlyStopping覆写了state_key并用基类提供的_generate_state_key()帮助方法生成字符串early_stopping.pyproperty override def state_key(self) - str: return self._generate_state_key(monitorself.monitor, modeself.mode)_generate_state_key()的实现callback.py是把一组键值对格式化成ClassName{...}形式的字符串def _generate_state_key(self, **kwargs: Any) - str: return f{self.__class__.__qualname__}{repr(kwargs)}对EarlyStopping(monitorval_loss, modemin)生成的 key 就是EarlyStopping{monitor: val_loss, mode: min}测试 test_early_stopping.py 直接断言了这一结果。而它的state_dict()保存了wait_count、stopped_epoch、best_score、patience、stopping_reason等恢复早停判定所必需的全部字段early_stopping.py——这也回答了到底该保存什么凡是影响后续决策的内部变量都要进state_dict。5.2 ModelCheckpointModelCheckpoint的state_key更进一步把监控指标、模式与保存触发条件全部编码进去model_checkpoint.pyproperty override def state_key(self) - str: return self._generate_state_key( monitorself.monitor, modeself.mode, every_n_train_stepsself._every_n_train_steps, every_n_epochsself._every_n_epochs, train_time_intervalself._train_time_interval, )原因很直观训练中完全可以同时存在多个ModelCheckpoint例如一个按val_loss每 1 个 epoch 保存、另一个按global_step每 1000 步保存它们的内部状态各自维护的best_k_models必须隔离。把决定这个实例管什么的构造参数全部纳入 key就保证了同一组参数配置的实例在断点续训时能精确对接。5.3 其他官方实现仓库中还有WeightAveragingweight_averaging.py、StochasticWeightAveragingstochastic_weight_avg.py、Timertimer.py、Finetuningfinetuning.py等官方回调都实现了state_dict/load_state_dict对是研究哪些状态值得持久化的一手素材。六、实践清单让自定义回调可断点续训结合文档与源码编写可恢复的自定义回调时建议遵循以下检查清单识别有状态回调回调内部是否有跨 epoch/batch 累积或决策用的属性如果有就需要持久化。实现state_dict()返回可 pickle的状态字典优先返回拷贝如self.state.copy()只暴露真正需要保存的字段。实现load_state_dict()与state_dict()严格镜像把字典写回内部属性用dict.update或带get的容错读取可提高对新旧版本的兼容性。评估是否多实例该回调是否可能以同一类型多个实例同时传给 Trainer是则必须覆写state_key。设计state_key只编码影响状态语义的构造参数参考EarlyStopping用_generate_state_key(monitor..., mode...)排除verbose这类纯显示参数。断点续训时原样传入回调加载ckpt_path的 Trainer 必须包含保存时使用的全部有状态回调否则 Lightning 会发出 Please add the following callbacks 告警且对应状态无人认领。兼容旧 checkpoint如需要基类的_legacy_state_key会自动处理 1.5.0 之前按类型保存的旧 checkpoint无需额外代码。七、进一步阅读回调基类定义与所有钩子src/lightning/pytorch/callbacks/callback.py回调状态收集/还原的底层实现src/lightning/pytorch/trainer/call.pycheckpoint 组装与回调状态写入src/lightning/pytorch/trainer/connectors/checkpoint_connector.py官方回调实现范本EarlyStoppingearly_stopping.py、ModelCheckpointmodel_checkpoint.py对应测试用例tests/tests_pytorch/callbacks/test_callbacks.py、tests/tests_pytorch/callbacks/test_early_stopping.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),仅供参考