Skip to content

Commit faa2a2d

Browse files
fix: respect continue_from_checkpoint in multi-stage runs (#619)
1 parent 9a7e26e commit faa2a2d

1 file changed

Lines changed: 5 additions & 1 deletion

File tree

‎trinity/cli/launcher.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -387,7 +387,11 @@ def run(
387387
from trinity.trainer import get_latest_hf_checkpoint_path
388388

389389
state_manager = StateManager(path=cfg.get_checkpoint_job_dir())
390-
latest_stage = state_manager.load_stage().get("latest_stage", 0)
390+
latest_stage = (
391+
state_manager.load_stage().get("latest_stage", 0)
392+
if cfg.continue_from_checkpoint
393+
else 0
394+
)
391395
prev_stage_checkpoint = None
392396
for i, stage_config in enumerate(cfg):
393397
if i < latest_stage:

0 commit comments

Comments
 (0)