fix: 断点续训恢复优化器/调度器状态及采样器剩余长度

- 使用Checkpoint.load()替代手动加载model.safetensors,恢复optimizer/scheduler状态
- TrainContextBuilder从checkpoint.extra恢复优化器和调度器state_dict
- ResumableDistributedSampler.__len__返回剩余样本数而非总数
- 训练前对state_dict置空避免mp.spawn pickle 7GB大对象
This commit is contained in:
2026-05-26 13:50:25 +08:00
parent dd1b39f435
commit a548d4553e
3 changed files with 18 additions and 13 deletions
+2 -1
View File
@@ -71,7 +71,8 @@ class TrainContextBuilder:
if self._checkpoint is not None:
context.epoch = max(self._checkpoint.epoch, cfg.start_epoch)
context.iteration = max(self._checkpoint.iteration, cfg.start_batch)
context.model.load_state_dict(self._checkpoint.state_dict)
if self._checkpoint.state_dict:
context.model.load_state_dict(self._checkpoint.state_dict)
context.checkpoint = self._checkpoint
else:
context.checkpoint = Checkpoint(