PyTorch Lightning 回调状态持久化:用 state_dict、load_state_dict 与 state_key 让自定义 Callback 可断点续训
2026/9/19 23:11:25 网站建设 项目流程

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_modelskth_best_model_path,决定何时该覆盖旧 checkpoint;
  • 随机权重平均(StochasticWeightAveragingWeightAveraging)记录平均模型权重。

这类回调如果不保存状态,那么使用Trainer.fit(ckpt_path=...)断点续训时,Lightning 只会恢复模型权重、优化器与学习率调度器,回调内部状态却会回到初始值。其结果可能是:已经等待了 20 个 epoch 的早停重新倒计时、训练计时器重新计时,最终改变整个训练的行为。

解决方案就是让回调自身实现两个钩子,把状态"交接"给 checkpoint 机制:

  • state_dict():返回一个可被 pickle 序列化的字典,表示回调当前状态;
  • load_state_dict(state_dict):把 checkpoint 中保存的字典还原回回调内部属性。

这两个钩子的默认实现位于 callback.py:state_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 或 batch:

from lightning.pytorch.callbacks import Callback class Counter(Callback): def __init__(self, what="epochs", verbose=True): 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 f"Counter[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",保存时后一个实例会覆盖前一个实例的状态,加载时两个实例也会拿到同一份状态——状态彻底混淆。

因此文档中的Counterwhat(该实例统计的是 epoch 还是 batch)编码进 key:

@property def state_key(self) -> str: # note: we do not include `verbose` here on purpose return f"Counter[what={self.what}]"

这样两个实例分别获得"Counter[what=epochs]""Counter[what=batches]"两个互不冲突的 key。注意注释中的刻意设计:只把会影响状态语义的参数放进 key,而把verbose这类纯显示参数排除在外。如果某个参数不影响状态结构,却写进了state_key,那么仅仅因为显示设置不同,就会导致同一个"逻辑回调"在断点续训时匹配不上。

然后像文档中这样,把两个实例同时交给 Trainer:

from lightning.pytorch import Trainer # two callbacks of the same type are being used trainer = Trainer(callbacks=[Counter(what="epochs"), Counter(what="batches")])

此时训练产生的 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 callback's `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

三点值得注意:

  1. 保存顺序是遍历trainer.callbacks,逐个调用callback.state_dict()
  2. 空状态不保存:只有state_dict()返回了非空字典时,该回调才会进入"callbacks"字段,因此无状态回调不会污染 checkpoint;
  3. key 就是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 callback's `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.py):

Be 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(分别监控epochglobal_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 EarlyStopping

EarlyStopping覆写了state_key,并用基类提供的_generate_state_key()帮助方法生成字符串(early_stopping.py):

@property @override def state_key(self) -> str: return self._generate_state_key(monitor=self.monitor, mode=self.mode)

_generate_state_key()的实现(callback.py)是把一组键值对格式化成"ClassName{...}"形式的字符串:

def _generate_state_key(self, **kwargs: Any) -> str: return f"{self.__class__.__qualname__}{repr(kwargs)}"

EarlyStopping(monitor="val_loss", mode="min"),生成的 key 就是:

EarlyStopping{'monitor': 'val_loss', 'mode': 'min'}

测试 test_early_stopping.py 直接断言了这一结果。而它的state_dict()保存了wait_countstopped_epochbest_scorepatiencestopping_reason等恢复早停判定所必需的全部字段(early_stopping.py)——这也回答了"到底该保存什么":凡是影响后续决策的内部变量,都要进state_dict

5.2 ModelCheckpoint

ModelCheckpointstate_key更进一步,把监控指标、模式与保存触发条件全部编码进去(model_checkpoint.py):

@property @override def state_key(self) -> str: return self._generate_state_key( monitor=self.monitor, mode=self.mode, every_n_train_steps=self._every_n_train_steps, every_n_epochs=self._every_n_epochs, train_time_interval=self._train_time_interval, )

原因很直观:训练中完全可以同时存在多个ModelCheckpoint(例如一个按val_loss每 1 个 epoch 保存、另一个按global_step每 1000 步保存),它们的内部状态(各自维护的best_k_models)必须隔离。把决定"这个实例管什么"的构造参数全部纳入 key,就保证了同一组参数配置的实例在断点续训时能精确对接。

5.3 其他官方实现

仓库中还有WeightAveraging(weight_averaging.py)、StochasticWeightAveraging(stochastic_weight_avg.py)、Timer(timer.py)、Finetuning(finetuning.py)等官方回调都实现了state_dict/load_state_dict对,是研究"哪些状态值得持久化"的一手素材。

六、实践清单:让自定义回调可断点续训

结合文档与源码,编写可恢复的自定义回调时,建议遵循以下检查清单:

  1. 识别有状态回调:回调内部是否有跨 epoch/batch 累积或决策用的属性?如果有,就需要持久化。
  2. 实现state_dict():返回可 pickle的状态字典;优先返回拷贝(如self.state.copy()),只暴露真正需要保存的字段。
  3. 实现load_state_dict():与state_dict()严格镜像,把字典写回内部属性;用dict.update或带get的容错读取可提高对新旧版本的兼容性。
  4. 评估是否多实例:该回调是否可能以同一类型多个实例同时传给 Trainer?是,则必须覆写state_key
  5. 设计state_key:只编码影响状态语义的构造参数(参考EarlyStopping_generate_state_key(monitor=..., mode=...)),排除verbose这类纯显示参数。
  6. 断点续训时原样传入回调:加载ckpt_path的 Trainer 必须包含保存时使用的全部有状态回调,否则 Lightning 会发出 "Please add the following callbacks" 告警,且对应状态无人认领。
  7. 兼容旧 checkpoint(如需要):基类的_legacy_state_key会自动处理 1.5.0 之前按类型保存的旧 checkpoint,无需额外代码。

七、进一步阅读

  • 回调基类定义与所有钩子:src/lightning/pytorch/callbacks/callback.py
  • 回调状态收集/还原的底层实现:src/lightning/pytorch/trainer/call.py
  • checkpoint 组装与回调状态写入:src/lightning/pytorch/trainer/connectors/checkpoint_connector.py
  • 官方回调实现范本:EarlyStopping(early_stopping.py)、ModelCheckpoint(model_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),仅供参考

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

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

立即咨询