From abbd95eca5d0e4892a27af6dd8ec08bc1a848339 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:37:47 +0800 Subject: [PATCH 01/11] rm return --- miles/backends/megatron_utils/actor.py | 1 - miles/ray/placement_group.py | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 97122a35031..9a0e65bc75c 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -132,7 +132,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..023814343c1 100644 --- a/miles/ray/placement_group.py +++ b/miles/ray/placement_group.py @@ -151,7 +151,7 @@ def create_training_models(args, pgs, rollout_manager, wandb_run_id): else: critic_model = None - start_rollout_ids = ray.get( + ray.get( actor_model.async_init(args, role="actor", with_ref=args.kl_coef != 0 or args.use_kl_loss) ) From 624d24999faaf9a03cd5a76d439f8142051e4389 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:38:39 +0800 Subject: [PATCH 02/11] more --- miles/utils/arguments.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 3b922576945..78954f4d480 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1118,6 +1118,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 + if TODO: + args.start_rollout_id = TODO 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" From bbcc9d83f4c0e36aa4a75af936d2023a33f2766c Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:39:12 +0800 Subject: [PATCH 03/11] more --- miles/backends/megatron_utils/actor.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 9a0e65bc75c..93c30cb5acd 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -73,7 +73,9 @@ def init( Timer().start("train_wait") return - start_rollout_id = loaded_rollout_id + 1 + expected_start_rollout_id = 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"]) From 15b32b9a011a2252a6ac56087870da2bd754ebaa Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:39:34 +0800 Subject: [PATCH 04/11] fmt --- miles/backends/megatron_utils/actor.py | 4 +++- miles/ray/placement_group.py | 4 +--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 93c30cb5acd..26ef4ac88da 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -74,7 +74,9 @@ def init( return expected_start_rollout_id = loaded_rollout_id + 1 - assert args.start_rollout_id == expected_start_rollout_id, f"{args.start_rollout_id=} {expected_start_rollout_id=}" + 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"]) diff --git a/miles/ray/placement_group.py b/miles/ray/placement_group.py index 023814343c1..2761b84a8a2 100644 --- a/miles/ray/placement_group.py +++ b/miles/ray/placement_group.py @@ -151,9 +151,7 @@ def create_training_models(args, pgs, rollout_manager, wandb_run_id): else: critic_model = None - ray.get( - actor_model.async_init(args, role="actor", with_ref=args.kl_coef != 0 or args.use_kl_loss) - ) + 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: From 5af266f2d3e0a550ffe0d41c09a58c7b93bf81d3 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:39:48 +0800 Subject: [PATCH 05/11] more --- miles/backends/megatron_utils/actor.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 26ef4ac88da..b632fd45177 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) From 010324ebba2689c5bc7e5d964c30626535905246 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:42:14 +0800 Subject: [PATCH 06/11] more --- miles/utils/checkpoint_utils.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) create mode 100644 miles/utils/checkpoint_utils.py 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()) From 15b483bbe0b4b79d316c2a1d8975bc63ee7430a4 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:43:01 +0800 Subject: [PATCH 07/11] more --- miles/utils/arguments.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index 78954f4d480..c294f389a70 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 ckpt_iter is None: args.no_load_optim = True args.no_load_rng = True args.finetune = True @@ -1118,7 +1115,7 @@ 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 - if TODO: + else: args.start_rollout_id = TODO if args.eval_interval is not None: From 405e6c6ee299db2e7d2542bf6386eec39c8be779 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:43:14 +0800 Subject: [PATCH 08/11] more --- miles/utils/arguments.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/miles/utils/arguments.py b/miles/utils/arguments.py index c294f389a70..6d1dd80765a 100644 --- a/miles/utils/arguments.py +++ b/miles/utils/arguments.py @@ -1107,7 +1107,7 @@ def miles_validate_args(args): ) load_ckpt_iter = get_latest_checkpointed_iteration(args.load) - if ckpt_iter is None: + if load_ckpt_iter is None: args.no_load_optim = True args.no_load_rng = True args.finetune = True @@ -1116,7 +1116,7 @@ def miles_validate_args(args): args.ckpt_step = args.ref_ckpt_step args.start_rollout_id = 0 else: - args.start_rollout_id = TODO + 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" From 190b5de8b6add61039b1608fe17a681ef88c6f98 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:44:06 +0800 Subject: [PATCH 09/11] more --- miles/backends/megatron_utils/actor.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index b632fd45177..3b4ef3c6360 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -73,10 +73,12 @@ def init( Timer().start("train_wait") return - expected_start_rollout_id = loaded_rollout_id + 1 + expected_start_rollout_ids = {loaded_rollout_id + 1} + if loaded_rollout_id == 0: + expected_start_rollout_ids |= {0} assert ( - args.start_rollout_id == expected_start_rollout_id - ), f"{args.start_rollout_id=} {expected_start_rollout_id=}" + args.start_rollout_id in expected_start_rollout_ids + ), f"{args.start_rollout_id=} {expected_start_rollout_ids=}" self.weights = {"actor": {}} self.update_cpu_params_dict(self.weights["actor"]) From c6a79050de3855823c02e9535026e9f59bd0d860 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:45:09 +0800 Subject: [PATCH 10/11] more --- miles/backends/megatron_utils/actor.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index 3b4ef3c6360..7e13253f23a 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -73,12 +73,10 @@ def init( Timer().start("train_wait") return - expected_start_rollout_ids = {loaded_rollout_id + 1} - if loaded_rollout_id == 0: - expected_start_rollout_ids |= {0} + expected_start_rollout_id = 0 if loaded_rollout_id == 0 else (loaded_rollout_id + 1) assert ( - args.start_rollout_id in expected_start_rollout_ids - ), f"{args.start_rollout_id=} {expected_start_rollout_ids=}" + 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"]) From bd0bfc83e6de938c83807b74256d04d663af0cd5 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 17 Oct 2025 14:56:44 +0800 Subject: [PATCH 11/11] fix --- miles/ray/placement_group.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/miles/ray/placement_group.py b/miles/ray/placement_group.py index 2761b84a8a2..686967b49c3 100644 --- a/miles/ray/placement_group.py +++ b/miles/ray/placement_group.py @@ -153,10 +153,6 @@ def create_training_models(args, pgs, rollout_manager, wandb_run_id): 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] - if args.use_critic: ray.get(critic_init_handle) actor_model.connect(critic_model)