feat: Checkpoint 支持 extra 通用扩展数据,用户通过函数自定义保存/恢复优化器等状态

- serialization.py: Checkpoint 新增 extra: dict 字段,
  save() 写入 extra.pt,load() 自动恢复
- train_callback.py: CheckpointCallback 新增 save_extra_fn
  参数,用户传入 (context) -> dict 决定保存哪些额外状态
- train_context.py: TrainContextBuilder 新增 load_extra_fn
  参数,用户传入 (extra, context) 从 checkpoint 恢复状态
This commit is contained in:
2026-05-09 15:50:38 +08:00
parent db99d8b254
commit ca4e6b907c
3 changed files with 28 additions and 4 deletions
+7 -1
View File
@@ -121,11 +121,13 @@ class CheckpointCallback(TrainCallback):
interval: int,
weight_only: bool = False,
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
):
self.save_dir = save_dir
self.interval = interval
self.weight_only = weight_only
self.state_dict_fn = state_dict_fn
self.save_extra_fn = save_extra_fn
self.last_ckpt_iter = 0
@only_on_rank(0)
@@ -139,8 +141,12 @@ class CheckpointCallback(TrainCallback):
else context.model.state_dict()
)
extra = self.save_extra_fn(context) if self.save_extra_fn else None
context.checkpoint = Checkpoint(
state_dict=state_dict, epoch=context.epoch, iteration=context.iteration
state_dict=state_dict,
epoch=context.epoch,
iteration=context.iteration,
extra=extra,
)
context.checkpoint.save(save_path)