diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 97122a35031..7e13253f23a 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -39,7 +39,7 @@ def init( role: str, wandb_run_id: str, with_ref: bool = False, - ) -> Optional[int]: + ): super().init(args, role, wandb_run_id, with_ref) init(args) @@ -73,7 +73,11 @@ def init( Timer().start("train_wait") return - start_rollout_id = loaded_rollout_id + 1 + expected_start_rollout_id = 0 if loaded_rollout_id == 0 else (loaded_rollout_id + 1) + assert ( + args.start_rollout_id == expected_start_rollout_id + ), f"{args.start_rollout_id=} {expected_start_rollout_id=}" + self.weights = {"actor": {}} self.update_cpu_params_dict(self.weights["actor"]) @@ -132,7 +136,6 @@ def init( self.prof.start() Timer().start("train_wait") - return start_rollout_id @torch.no_grad() def update_cpu_params_dict(self, params_dict: Dict[str, torch.Tensor]) -> None: diff --git a/miles/ray/placement_group.py b/miles/ray/placement_group.py index 006490f99f9..686967b49c3 100644 --- a/miles/ray/placement_group.py +++ b/miles/ray/placement_group.py @@ -151,13 +151,7 @@ def create_training_models(args, pgs, rollout_manager, wandb_run_id): else: critic_model = None - start_rollout_ids = ray.get( - actor_model.async_init(args, role="actor", with_ref=args.kl_coef != 0 or args.use_kl_loss) - ) - - assert len(set(start_rollout_ids)) == 1 - if args.start_rollout_id is None: - args.start_rollout_id = start_rollout_ids[0] + ray.get(actor_model.async_init(args, role="actor", with_ref=args.kl_coef != 0 or args.use_kl_loss)) if args.use_critic: ray.get(critic_init_handle) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 3b922576945..6d1dd80765a 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -8,6 +8,7 @@ from miles.backends.sglang_utils.arguments import add_sglang_arguments from miles.backends.sglang_utils.arguments import validate_args as sglang_validate_args +from miles.utils.checkpoint_utils import get_latest_checkpointed_iteration def reset_arg(parser, name, **kwargs): @@ -1105,12 +1106,8 @@ def miles_validate_args(args): "please make sure it is a valid megatron checkpoint directory." ) - # TODO: During loading, we need to set the start_rollout_id here. - if ( - args.load is None - or not os.path.exists(args.load) - or not os.path.exists(os.path.join(args.load, "latest_checkpointed_iteration.txt")) - ): + load_ckpt_iter = get_latest_checkpointed_iteration(args.load) + if load_ckpt_iter is None: args.no_load_optim = True args.no_load_rng = True args.finetune = True @@ -1118,6 +1115,8 @@ def miles_validate_args(args): if args.ref_ckpt_step is not None: args.ckpt_step = args.ref_ckpt_step args.start_rollout_id = 0 + else: + args.start_rollout_id = load_ckpt_iter + 1 if args.eval_interval is not None: assert args.eval_prompt_data is not None, "eval_prompt_data must be set when eval_interval is set" diff --git a/miles/utils/checkpoint_utils.py b/miles/utils/checkpoint_utils.py new file mode 100644 index 00000000000..f74d87afba7 --- /dev/null +++ b/miles/utils/checkpoint_utils.py @@ -0,0 +1,13 @@ +from pathlib import Path +from typing import Optional + + +def get_latest_checkpointed_iteration(dir_load: Optional[str]) -> Optional[int]: + if dir_load is None: + return None + + path_txt = Path(dir_load) / "latest_checkpointed_iteration.txt" + if not path_txt.exists(): + return None + + return int(path_txt.read_text())