From 2a70a473c31b37ab27f86aa9525f1a0d374bd5bc Mon Sep 17 00:00:00 2001 From: lilei Date: Tue, 23 Sep 2025 09:20:37 +0000 Subject: [PATCH 01/15] [fix] fix ppo value_loss problem and add kl_coef to reward --- slime/backends/megatron_utils/loss.py | 1 + 1 file changed, 1 insertion(+) diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index 50b1ffad3e..e592308eab 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -161,6 +161,7 @@ def compute_advantages_and_returns(args, rollout_data): # TODO: optimize this old_rewards = rewards rewards = [] + kl *= args.kl_coef for reward, k in zip(old_rewards, kl): k *= -args.kl_coef k[-1] += reward From e63c00b227b5b3af3164169d7c62044fc5a93899 Mon Sep 17 00:00:00 2001 From: lilei Date: Wed, 24 Sep 2025 05:20:25 +0000 Subject: [PATCH 02/15] [fix] fix ppo value_loss problem and add kl_coef to reward --- slime/backends/megatron_utils/loss.py | 1 - 1 file changed, 1 deletion(-) diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index e592308eab..50b1ffad3e 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -161,7 +161,6 @@ def compute_advantages_and_returns(args, rollout_data): # TODO: optimize this old_rewards = rewards rewards = [] - kl *= args.kl_coef for reward, k in zip(old_rewards, kl): k *= -args.kl_coef k[-1] += reward From 0e9c14c5cb6e1089e77c8d849436ea6d87afd131 Mon Sep 17 00:00:00 2001 From: lilei Date: Thu, 25 Sep 2025 16:40:33 +0000 Subject: [PATCH 03/15] [fix] ppo update_weights form distribute --- .../megatron_utils/update_weight_utils.py | 91 ++++++++++++++++++- slime/utils/arguments.py | 3 +- 2 files changed, 91 insertions(+), 3 deletions(-) diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index dedbea7512..7b7ff5b0f9 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -304,9 +304,16 @@ def __init__(self, args, model, weights, *, model_name, quantization_config, voc self.quantization_config = quantization_config self.param_info_buckets = get_param_info_buckets(self.args, self.model) self.weight_version = 0 + self.use_distribute = self.args.use_critic # TODO def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self.rollout_engines = rollout_engines + if self.use_distribute: + colocate_engine_nums = ( + self.args.actor_num_nodes * self.args.actor_num_gpus_per_node // self.args.rollout_num_gpus_per_engine + ) # TODO + self.connect_rollout_engines_distribute(rollout_engines[colocate_engine_nums:], rollout_engine_lock) + self.rollout_engines = rollout_engines[:colocate_engine_nums] # Here we assume the gpu id of rollout engines and train actors are the same. for i, engine in enumerate(self.rollout_engines): @@ -322,6 +329,46 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self._ipc_gather_group = new_group self._ipc_engine = engine + def connect_rollout_engines_distribute(self, rollout_engines, rollout_engine_lock): + self.rollout_engines_distribute = rollout_engines + self.rollout_engine_lock = rollout_engine_lock + + self._is_distribute_src_rank = ( + mpu.get_data_parallel_rank(with_context_parallel=True) == 0 + and mpu.get_tensor_model_parallel_rank() == 0 + and mpu.get_pipeline_model_parallel_rank() == 0 + ) + + if self._is_distribute_src_rank: + self._group_name = "slime_distribute" + + if self._is_distribute_src_rank: + master_address = ray._private.services.get_node_ip_address() + with socket.socket() as sock: + sock.bind(("", 0)) + master_port = sock.getsockname()[1] + world_size = len(rollout_engines) * self.args.rollout_num_gpus_per_engine + 1 + + refs = [ + engine.init_weights_update_group.remote( + master_address, + master_port, + i * self.args.rollout_num_gpus_per_engine + 1, + world_size, + self._group_name, + backend="nccl", + ) + for i, engine in enumerate(self.rollout_engines_distribute) + ] + self._model_update_groups = init_process_group( + backend="nccl", + init_method=f"tcp://{master_address}:{master_port}", + world_size=world_size, + rank=0, + group_name=self._group_name, + ) + ray.get(refs) + @torch.no_grad() def update_weights(self): self.weight_version += 1 @@ -400,7 +447,22 @@ def _update_bucket_weights_from_tensor(self, param_infos): converted_named_tensors.extend( convert_to_hf(self.args, self.model_name, info.name, param, self.quantization_config) ) - self._update_converted_params_from_tensor(converted_named_tensors) + + refs = [] + if dist.get_rank() == self._ipc_gather_src: + refs.extend(self._update_converted_params_from_tensor(converted_named_tensors)) + else: + self._update_converted_params_from_tensor(converted_named_tensors) + + if self.use_distribute: + if self._is_distribute_src_rank: + refs.extend(self._update_bucket_weights_from_distributed(converted_named_tensors)) + ray.get(refs) + + if self.use_distribute: + if self._is_distribute_src_rank: + converted_named_tensors.clear() + ray.get(self.rollout_engine_lock.release.remote()) def _update_converted_params_from_tensor(self, converted_named_tensors): if use_flattened_tensor_bucket: @@ -451,7 +513,32 @@ def _update_converted_params_from_tensor(self, converted_named_tensors): "weight_version": str(self.weight_version), } refs.append(self._ipc_engine.update_weights_from_tensor.remote(**kwargs)) - ray.get(refs) + return refs + return None + + def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar=None): + # lock the rollout engines to prevent dead lock on broadcast. + while not ray.get(self.rollout_engine_lock.acquire.remote()): + time.sleep(0.1) + + refs = [ + engine.update_weights_from_distributed.remote( + names=[name for name, _ in converted_named_tensors], + dtypes=[param.dtype for _, param in converted_named_tensors], + shapes=[param.shape for _, param in converted_named_tensors], + group_name=self._group_name, + weight_version=str(self.weight_version), + ) + for engine in self.rollout_engines_distribute + ] + + handles = [] + for _, param in converted_named_tensors: + handles.append(dist.broadcast(param.data, 0, group=self._model_update_groups, async_op=True)) + for handle in handles: + handle.wait() + + return refs class UpdateWeightFromDistributed: diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index aeafedd3cc..f570799be4 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1081,7 +1081,8 @@ def slime_validate_args(args): f"rollout_num_gpus {args.rollout_num_gpus} != actor_num_gpus_per_node {args.actor_num_gpus_per_node} " f"* actor_num_nodes {args.actor_num_nodes}, overriding rollout_num_gpus to match actor_num_gpus_per_node * actor_num_nodes." ) - args.rollout_num_gpus = args.actor_num_gpus_per_node * args.actor_num_nodes + if not args.use_critic: + args.rollout_num_gpus = args.actor_num_gpus_per_node * args.actor_num_nodes if args.eval_function_path is None: args.eval_function_path = args.rollout_function_path From 48d3292fbba40de234635529f456a9dc8bfb473d Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 03:04:04 +0000 Subject: [PATCH 04/15] [fix] ppo cp bugs --- slime/backends/megatron_utils/loss.py | 3 ++- slime/backends/megatron_utils/update_weight_utils.py | 6 +++--- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index 50b1ffad3e..0e8bd6da81 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -163,7 +163,8 @@ def compute_advantages_and_returns(args, rollout_data): rewards = [] for reward, k in zip(old_rewards, kl): k *= -args.kl_coef - k[-1] += reward + if k.numel() > 0: + k[-1] += reward rewards.append(k) advantages, returns = list( zip( diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index 7b7ff5b0f9..c2e6c378e7 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -304,14 +304,14 @@ def __init__(self, args, model, weights, *, model_name, quantization_config, voc self.quantization_config = quantization_config self.param_info_buckets = get_param_info_buckets(self.args, self.model) self.weight_version = 0 - self.use_distribute = self.args.use_critic # TODO + self.use_distribute = self.args.use_critic def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self.rollout_engines = rollout_engines if self.use_distribute: colocate_engine_nums = ( self.args.actor_num_nodes * self.args.actor_num_gpus_per_node // self.args.rollout_num_gpus_per_engine - ) # TODO + ) self.connect_rollout_engines_distribute(rollout_engines[colocate_engine_nums:], rollout_engine_lock) self.rollout_engines = rollout_engines[:colocate_engine_nums] @@ -340,7 +340,7 @@ def connect_rollout_engines_distribute(self, rollout_engines, rollout_engine_loc ) if self._is_distribute_src_rank: - self._group_name = "slime_distribute" + self._group_name = "slime_ppo_distribute" if self._is_distribute_src_rank: master_address = ray._private.services.get_node_ip_address() From f1dc47cf0ff53c12ff078eba47e543081096bcfd Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 03:31:54 +0000 Subject: [PATCH 05/15] [Fix] ppo pp bugs --- slime/backends/megatron_utils/data.py | 18 +++++++++++++----- slime/backends/megatron_utils/loss.py | 2 +- 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index ffcfb70429..30fd2d07b0 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -390,22 +390,30 @@ def log_perf_data(rollout_id, args): def sync_actor_critic_data( args, - values: Optional[list[torch.Tensor]] = None, - log_probs: Optional[list[torch.Tensor]] = None, - ref_log_probs: Optional[list[torch.Tensor]] = None, + values: Optional[dict[str, list[torch.Tensor]]], + log_probs: Optional[dict[str, list[torch.Tensor]]], + ref_log_probs: Optional[dict[str, list[torch.Tensor]]], group: Optional[dist.ProcessGroup] = None, ): + # return None when not pp last stage + if not values and not log_probs: + return None, None, None handles = [] - if values is None: + if not values: values = [torch.empty_like(log_prob) for log_prob in log_probs] + else: + values = values["values"] for value in values: handles.append(dist.broadcast(value, src=1, group=group, async_op=True)) if args.kl_coef != 0 or args.use_kl_loss: - if log_probs is None: + if not log_probs: ref_log_probs = [torch.empty_like(value) for value in values] log_probs = [torch.empty_like(value) for value in values] + else: + ref_log_probs = ref_log_probs["ref_log_probs"] + log_probs = log_probs["log_probs"] for ref_log_prob, log_prob in zip(ref_log_probs, log_probs): handles.append(dist.broadcast(ref_log_prob, src=0, group=group, async_op=True)) handles.append(dist.broadcast(log_prob, src=0, group=group, async_op=True)) diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index 0e8bd6da81..0df49ba8ec 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -364,7 +364,7 @@ def value_loss_function(args, batch, logits, sum_of_sample_mean): total_lengths=batch["total_lengths"], response_lengths=batch["response_lengths"], ) - values = torch.cat([value.squeeze(-1) for value in values["values"]], dim=0) + values = torch.cat([value.flatten() for value in values["values"]], dim=0) returns = torch.cat(batch["returns"], dim=0) From 5210d9948fa27fdc0108df0972179834d075f5bc Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 03:32:52 +0000 Subject: [PATCH 06/15] [Fix] ppo pp bugs --- slime/backends/megatron_utils/actor.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index 1439fa6216..50a9ee75ae 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -267,7 +267,7 @@ def train_critic(self, rollout_id, rollout_data): self.model, data_iterator, num_microbatches, - )["values"] + ) if rollout_id < self.args.num_critic_only_steps: # we will only use the shape of log_probs in this situation @@ -330,8 +330,8 @@ def train_actor(self, rollout_id, rollout_data): values, log_probs, ref_log_probs = sync_actor_critic_data( self.args, None, - log_probs["log_probs"], - ref_log_probs["ref_log_probs"] if (self.args.kl_coef != 0 or self.args.use_kl_loss) else None, + log_probs, + ref_log_probs if (self.args.kl_coef != 0 or self.args.use_kl_loss) else None, self._actor_critic_groups, ) rollout_data.update({"values": values}) From d4ad1e7f99e4096ec8358a628f0cde1c7c0934b9 Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 03:34:06 +0000 Subject: [PATCH 07/15] [Fix] ppo pp bugs --- slime/backends/megatron_utils/data.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 30fd2d07b0..b435bdbb88 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -390,9 +390,9 @@ def log_perf_data(rollout_id, args): def sync_actor_critic_data( args, - values: Optional[dict[str, list[torch.Tensor]]], - log_probs: Optional[dict[str, list[torch.Tensor]]], - ref_log_probs: Optional[dict[str, list[torch.Tensor]]], + values: Optional[dict[str, list[torch.Tensor]]] = None, + log_probs: Optional[dict[str, list[torch.Tensor]]] = None, + ref_log_probs: Optional[dict[str, list[torch.Tensor]]] = None, group: Optional[dist.ProcessGroup] = None, ): # return None when not pp last stage From 5ea89e95dc88d42a76997ca7c0833badc46ba788 Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 03:44:59 +0000 Subject: [PATCH 08/15] [Fix] ppo pp bugs --- slime/backends/megatron_utils/data.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index b435bdbb88..49b48492ae 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -401,7 +401,7 @@ def sync_actor_critic_data( handles = [] if not values: - values = [torch.empty_like(log_prob) for log_prob in log_probs] + values = [torch.empty_like(log_prob) for log_prob in log_probs["log_probs"]] else: values = values["values"] for value in values: From 0586d6579eb687be29896c07281de1a1407a8c7c Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 03:48:55 +0000 Subject: [PATCH 09/15] [Fix] ppo pp bugs --- slime/backends/megatron_utils/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/slime/backends/megatron_utils/model.py b/slime/backends/megatron_utils/model.py index f09f52dc88..82bd98335a 100644 --- a/slime/backends/megatron_utils/model.py +++ b/slime/backends/megatron_utils/model.py @@ -477,7 +477,7 @@ def train(rollout_id, model, optimizer, opt_param_scheduler, data_iterator, num_ log_dict[f"train/{role_tag}lr-pg_{param_group_id}"] = opt_param_scheduler.get_lr(param_group) if args.use_wandb: - log_dict["train/step"] = accumulated_step_id + log_dict[f"train/{role_tag}step"] = accumulated_step_id wandb.log(log_dict) if args.ci_test: From 06f8a212b04b3f1c2f17989a1acff1c63663d8c3 Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 08:30:26 +0000 Subject: [PATCH 10/15] [Fix] ppo cp pp bugs --- slime/backends/megatron_utils/actor.py | 28 +--- slime/backends/megatron_utils/data.py | 20 +-- slime/backends/megatron_utils/loss.py | 11 +- .../megatron_utils/update_weight_utils.py | 17 +- slime/utils/arguments.py | 5 +- slime/utils/ppo_utils.py | 150 +++++++++++------- 6 files changed, 130 insertions(+), 101 deletions(-) diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index 50a9ee75ae..f787ee97f5 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -269,22 +269,11 @@ def train_critic(self, rollout_id, rollout_data): num_microbatches, ) - if rollout_id < self.args.num_critic_only_steps: - # we will only use the shape of log_probs in this situation - log_probs = values - ref_log_probs = values - else: - values, log_probs, ref_log_probs = sync_actor_critic_data( - self.args, values, None, None, self._actor_critic_groups - ) + if rollout_id >= self.args.num_critic_only_steps: + sync_data = sync_actor_critic_data(self.args, values, self._actor_critic_groups) + rollout_data.update(sync_data) - rollout_data.update( - { - "values": values, - "log_probs": log_probs, - "ref_log_probs": ref_log_probs, - } - ) + rollout_data.update(values) compute_advantages_and_returns(self.args, rollout_data) @@ -325,16 +314,15 @@ def train_actor(self, rollout_id, rollout_data): store_prefix="", ) rollout_data.update(log_probs) - if self.args.use_critic: - values, log_probs, ref_log_probs = sync_actor_critic_data( + if self.args.kl_coef != 0 or self.args.use_kl_loss: + log_probs.update(ref_log_probs) + sync_data = sync_actor_critic_data( self.args, - None, log_probs, - ref_log_probs if (self.args.kl_coef != 0 or self.args.use_kl_loss) else None, self._actor_critic_groups, ) - rollout_data.update({"values": values}) + rollout_data.update(sync_data) # when there is old actor, we need to update the model params to actor manually if "old_actor" in self.weights: diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 49b48492ae..5ab1c88939 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -390,20 +390,20 @@ def log_perf_data(rollout_id, args): def sync_actor_critic_data( args, - values: Optional[dict[str, list[torch.Tensor]]] = None, - log_probs: Optional[dict[str, list[torch.Tensor]]] = None, - ref_log_probs: Optional[dict[str, list[torch.Tensor]]] = None, + data: Optional[dict[str, list[torch.Tensor]]] = None, group: Optional[dist.ProcessGroup] = None, ): + values, log_probs, ref_log_probs = map(data.get, ("values", "log_probs", "ref_log_probs")) + # return None when not pp last stage if not values and not log_probs: - return None, None, None + return {} + handles = [] if not values: - values = [torch.empty_like(log_prob) for log_prob in log_probs["log_probs"]] - else: - values = values["values"] + values = [torch.empty_like(log_prob) for log_prob in log_probs] + for value in values: handles.append(dist.broadcast(value, src=1, group=group, async_op=True)) @@ -412,12 +412,12 @@ def sync_actor_critic_data( ref_log_probs = [torch.empty_like(value) for value in values] log_probs = [torch.empty_like(value) for value in values] else: - ref_log_probs = ref_log_probs["ref_log_probs"] - log_probs = log_probs["log_probs"] + ref_log_probs = ref_log_probs + log_probs = log_probs for ref_log_prob, log_prob in zip(ref_log_probs, log_probs): handles.append(dist.broadcast(ref_log_prob, src=0, group=group, async_op=True)) handles.append(dist.broadcast(log_prob, src=0, group=group, async_op=True)) for handle in handles: handle.wait() - return values, log_probs, ref_log_probs + return {"values": values, "log_probs": log_probs, "ref_log_probs": ref_log_probs} diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index 0df49ba8ec..fc492be7a7 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -137,7 +137,7 @@ def compute_advantages_and_returns(args, rollout_data): if log_probs is None and values is None: return - if args.kl_coef == 0: + if args.kl_coef == 0 or not log_probs: # when kl_coef is 0, we won't compute ref_log_prob xs = log_probs if log_probs is not None else values kl = [torch.zeros_like(x, dtype=torch.float32, device=x.device) for x in xs] @@ -163,14 +163,17 @@ def compute_advantages_and_returns(args, rollout_data): rewards = [] for reward, k in zip(old_rewards, kl): k *= -args.kl_coef - if k.numel() > 0: + cp_rank = mpu.get_context_parallel_rank() + if cp_rank == 0: k[-1] += reward rewards.append(k) advantages, returns = list( zip( *[ - get_advantages_and_returns(value, reward, args.gamma, args.lambd) - for value, reward in zip(values, rewards) + get_advantages_and_returns(total_length, response_length, value, reward, args.gamma, args.lambd) + for total_length, response_length, value, reward in zip( + total_lengths, response_lengths, values, rewards + ) ] ) ) diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index c2e6c378e7..c1e8cbfd66 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -304,14 +304,15 @@ def __init__(self, args, model, weights, *, model_name, quantization_config, voc self.quantization_config = quantization_config self.param_info_buckets = get_param_info_buckets(self.args, self.model) self.weight_version = 0 - self.use_distribute = self.args.use_critic def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self.rollout_engines = rollout_engines + colocate_engine_nums = ( + self.args.actor_num_nodes * self.args.actor_num_gpus_per_node // self.args.rollout_num_gpus_per_engine + ) + self.use_distribute = len(rollout_engines) > colocate_engine_nums + if self.use_distribute: - colocate_engine_nums = ( - self.args.actor_num_nodes * self.args.actor_num_gpus_per_node // self.args.rollout_num_gpus_per_engine - ) self.connect_rollout_engines_distribute(rollout_engines[colocate_engine_nums:], rollout_engine_lock) self.rollout_engines = rollout_engines[:colocate_engine_nums] @@ -448,11 +449,7 @@ def _update_bucket_weights_from_tensor(self, param_infos): convert_to_hf(self.args, self.model_name, info.name, param, self.quantization_config) ) - refs = [] - if dist.get_rank() == self._ipc_gather_src: - refs.extend(self._update_converted_params_from_tensor(converted_named_tensors)) - else: - self._update_converted_params_from_tensor(converted_named_tensors) + refs = self._update_converted_params_from_tensor(converted_named_tensors) if self.use_distribute: if self._is_distribute_src_rank: @@ -514,7 +511,7 @@ def _update_converted_params_from_tensor(self, converted_named_tensors): } refs.append(self._ipc_engine.update_weights_from_tensor.remote(**kwargs)) return refs - return None + return [] def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar=None): # lock the rollout engines to prevent dead lock on broadcast. diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index f570799be4..1f739dd3e4 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1081,8 +1081,9 @@ def slime_validate_args(args): f"rollout_num_gpus {args.rollout_num_gpus} != actor_num_gpus_per_node {args.actor_num_gpus_per_node} " f"* actor_num_nodes {args.actor_num_nodes}, overriding rollout_num_gpus to match actor_num_gpus_per_node * actor_num_nodes." ) - if not args.use_critic: - args.rollout_num_gpus = args.actor_num_gpus_per_node * args.actor_num_nodes + args.rollout_num_gpus = args.actor_num_gpus_per_node * args.actor_num_nodes + if args.use_critic: + args.rollout_num_gpus += args.critic_num_gpus_per_node * args.critic_num_nodes if args.eval_function_path is None: args.eval_function_path = args.rollout_function_path diff --git a/slime/utils/ppo_utils.py b/slime/utils/ppo_utils.py index 1236f5c4ae..b3d7060cab 100644 --- a/slime/utils/ppo_utils.py +++ b/slime/utils/ppo_utils.py @@ -169,39 +169,8 @@ def get_reinforce_plus_plus_returns( prompt_len = total_len - response_len if cp_size > 1: - from slime.backends.megatron_utils.cp_utils import get_logits_and_tokens_offset_with_cp - - # Step 1: Gather all KL chunks and token_offsets from all ranks - _, _, _, token_offsets = get_logits_and_tokens_offset_with_cp(total_len, response_len) - - object_to_gather = {"kl_chunk": local_kl_chunk.cpu(), "offsets": token_offsets} - gathered_objects = [None] * cp_size - dist.all_gather_object(gathered_objects, object_to_gather, group=mpu.get_context_parallel_group()) - - # Step 2: Reconstruct the full response tensor by splitting and placing each part. - full_kl_response = torch.zeros(response_len, device=device, dtype=dtype) - for obj in gathered_objects: - kl_chunk = obj["kl_chunk"].to(device) - global_offsets = obj["offsets"] - - # Calculate the lengths of part_0 and part_1 for this specific chunk. - s0, e0 = global_offsets[0] - s1, e1 = global_offsets[1] - res_s0, res_e0 = max(0, s0 - prompt_len), max(0, e0 - prompt_len) - res_s1, res_e1 = max(0, s1 - prompt_len), max(0, e1 - prompt_len) - len0 = res_e0 - res_s0 - len1 = res_e1 - res_s1 - - if kl_chunk.numel() > 0: - # Split the received contiguous chunk back into its zigzag parts. - kl_part_0, kl_part_1 = torch.split(kl_chunk, [len0, len1]) - - # Place each part in its own correct location. - if kl_part_0.numel() > 0: - full_kl_response[res_s0:res_e0] = kl_part_0 - if kl_part_1.numel() > 0: - full_kl_response[res_s1:res_e1] = kl_part_1 - + # Step 1,2:Gather all chunks and token_offsets from all ranks and reconstruct the full response tensor by splitting and placing each part + full_kl_response = all_gather_with_cp(local_kl_chunk, total_len, response_len, prompt_len, device, dtype) else: full_kl_response = local_kl_chunk @@ -221,23 +190,9 @@ def get_reinforce_plus_plus_returns( # Step 4: Pick up the results corresponding to our local chunk's parts. if cp_size > 1: - local_returns_chunk_parts = [] - local_s0, local_e0 = token_offsets[0] - local_s1, local_e1 = token_offsets[1] - local_res_s0, local_res_e0 = max(0, local_s0 - prompt_len), max(0, local_e0 - prompt_len) - local_res_s1, local_res_e1 = max(0, local_s1 - prompt_len), max(0, local_e1 - prompt_len) - - if local_res_e0 > local_res_s0: - local_returns_chunk_parts.append(returns_for_seq[local_res_s0:local_res_e0]) - if local_res_e1 > local_res_s1: - local_returns_chunk_parts.append(returns_for_seq[local_res_s1:local_res_e1]) - - local_returns_chunk = ( - torch.cat(local_returns_chunk_parts) - if local_returns_chunk_parts - else torch.tensor([], device=device, dtype=dtype) + local_returns_chunk = get_local_chunk_with_cp( + returns_for_seq, total_len, response_len, prompt_len, device, dtype ) - else: local_returns_chunk = returns_for_seq @@ -276,6 +231,8 @@ def get_reinforce_plus_plus_baseline_advantages( def get_advantages_and_returns( + total_len: int, + response_len: int, values: torch.Tensor, rewards: torch.Tensor, gamma: float, @@ -301,17 +258,36 @@ def get_advantages_and_returns( - advantages: Tensor of shape (response_size,) - returns: Tensor of shape (response_size,) """ + from megatron.core import mpu + + cp_size = mpu.get_context_parallel_world_size() + if cp_size > 0: + device, dtype = rewards.device, rewards.dtype + prompt_len = total_len - response_len + full_rewards = all_gather_with_cp(rewards, total_len, response_len, prompt_len, device, dtype) + full_values = all_gather_with_cp(values, total_len, response_len, prompt_len, device, dtype) + else: + full_rewards = rewards + full_values = values + lastgaelam = 0 advantages_reversed = [] - response_length = rewards.size(0) - for t in reversed(range(response_length)): - nextvalues = values[t + 1] if t < response_length - 1 else 0.0 - delta = rewards[t] + gamma * nextvalues - values[t] + for t in reversed(range(response_len)): + nextvalues = full_values[t + 1] if t < response_len - 1 else 0.0 + delta = full_rewards[t] + gamma * nextvalues - full_values[t] lastgaelam = delta + gamma * lambd * lastgaelam advantages_reversed.append(lastgaelam) - advantages = torch.tensor(advantages_reversed[::-1], dtype=values.dtype, device=values.device) - returns = advantages + values + full_advantages = torch.tensor(advantages_reversed[::-1], dtype=full_values.dtype, device=full_values.device) + full_returns = full_advantages + full_values + + if cp_size > 0: + advantages = get_local_chunk_with_cp(full_advantages, total_len, response_len, prompt_len, device, dtype) + returns = get_local_chunk_with_cp(full_returns, total_len, response_len, prompt_len, device, dtype) + else: + advantages = full_advantages + returns = full_returns + return advantages.detach(), returns @@ -333,3 +309,67 @@ def calculate_log_probs_and_entropy(logits, tokens, tp_group, with_entropy: bool else: entropy = None return log_prob, entropy + + +def all_gather_with_cp(local_chunk, total_len, response_len, prompt_len, device, dtype): + from megatron.core import mpu + + from slime.backends.megatron_utils.cp_utils import get_logits_and_tokens_offset_with_cp + + cp_size = mpu.get_context_parallel_world_size() + + # Step 1: Gather all chunks and token_offsets from all ranks + _, _, _, token_offsets = get_logits_and_tokens_offset_with_cp(total_len, response_len) + + object_to_gather = {"chunk": local_chunk.cpu(), "offsets": token_offsets} + gathered_objects = [None] * cp_size + dist.all_gather_object(gathered_objects, object_to_gather, group=mpu.get_context_parallel_group()) + + # Step 2: Reconstruct the full response tensor by splitting and placing each part. + full_response = torch.zeros(response_len, device=device, dtype=dtype) + for obj in gathered_objects: + chunk = obj["chunk"].to(device) + global_offsets = obj["offsets"] + + # Calculate the lengths of part_0 and part_1 for this specific chunk. + s0, e0 = global_offsets[0] + s1, e1 = global_offsets[1] + res_s0, res_e0 = max(0, s0 - prompt_len), max(0, e0 - prompt_len) + res_s1, res_e1 = max(0, s1 - prompt_len), max(0, e1 - prompt_len) + len0 = res_e0 - res_s0 + len1 = res_e1 - res_s1 + if chunk.numel() > 0: + # Split the received contiguous chunk back into its zigzag parts. + part_0, part_1 = torch.split(chunk, [len0, len1]) + + # Place each part in its own correct location. + if part_0.numel() > 0: + full_response[res_s0:res_e0] = part_0 + if part_1.numel() > 0: + full_response[res_s1:res_e1] = part_1 + + return full_response + + +def get_local_chunk_with_cp(full_response, total_len, response_len, prompt_len, device, dtype): + from slime.backends.megatron_utils.cp_utils import get_logits_and_tokens_offset_with_cp + + _, _, _, token_offsets = get_logits_and_tokens_offset_with_cp(total_len, response_len) + + local_returns_chunk_parts = [] + local_s0, local_e0 = token_offsets[0] + local_s1, local_e1 = token_offsets[1] + local_res_s0, local_res_e0 = max(0, local_s0 - prompt_len), max(0, local_e0 - prompt_len) + local_res_s1, local_res_e1 = max(0, local_s1 - prompt_len), max(0, local_e1 - prompt_len) + + if local_res_e0 > local_res_s0: + local_returns_chunk_parts.append(full_response[local_res_s0:local_res_e0]) + if local_res_e1 > local_res_s1: + local_returns_chunk_parts.append(full_response[local_res_s1:local_res_e1]) + + local_chunk = ( + torch.cat(local_returns_chunk_parts) + if local_returns_chunk_parts + else torch.tensor([], device=device, dtype=dtype) + ) + return local_chunk From d54b1dc4c7dbc02148139736ae931f00dc946454 Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 09:37:28 +0000 Subject: [PATCH 11/15] [Fix] ppo update_from_distribute --- .../megatron_utils/update_weight_utils.py | 111 ++++++------------ 1 file changed, 33 insertions(+), 78 deletions(-) diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index c1e8cbfd66..d44b4d6a41 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -313,7 +313,17 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self.use_distribute = len(rollout_engines) > colocate_engine_nums if self.use_distribute: - self.connect_rollout_engines_distribute(rollout_engines[colocate_engine_nums:], rollout_engine_lock) + self.distributed_weight_updator = UpdateWeightFromDistributed( + args=self.args, + model=self.model, + weights=self.weights, + model_name=self.model_name, + quantization_config=self.quantization_config, + vocab_size=self.vocab_size, + ) + self.distributed_weight_updator.connect_rollout_engines( + rollout_engines[colocate_engine_nums:], rollout_engine_lock, True + ) self.rollout_engines = rollout_engines[:colocate_engine_nums] # Here we assume the gpu id of rollout engines and train actors are the same. @@ -330,46 +340,6 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self._ipc_gather_group = new_group self._ipc_engine = engine - def connect_rollout_engines_distribute(self, rollout_engines, rollout_engine_lock): - self.rollout_engines_distribute = rollout_engines - self.rollout_engine_lock = rollout_engine_lock - - self._is_distribute_src_rank = ( - mpu.get_data_parallel_rank(with_context_parallel=True) == 0 - and mpu.get_tensor_model_parallel_rank() == 0 - and mpu.get_pipeline_model_parallel_rank() == 0 - ) - - if self._is_distribute_src_rank: - self._group_name = "slime_ppo_distribute" - - if self._is_distribute_src_rank: - master_address = ray._private.services.get_node_ip_address() - with socket.socket() as sock: - sock.bind(("", 0)) - master_port = sock.getsockname()[1] - world_size = len(rollout_engines) * self.args.rollout_num_gpus_per_engine + 1 - - refs = [ - engine.init_weights_update_group.remote( - master_address, - master_port, - i * self.args.rollout_num_gpus_per_engine + 1, - world_size, - self._group_name, - backend="nccl", - ) - for i, engine in enumerate(self.rollout_engines_distribute) - ] - self._model_update_groups = init_process_group( - backend="nccl", - init_method=f"tcp://{master_address}:{master_port}", - world_size=world_size, - rank=0, - group_name=self._group_name, - ) - ray.get(refs) - @torch.no_grad() def update_weights(self): self.weight_version += 1 @@ -452,14 +422,16 @@ def _update_bucket_weights_from_tensor(self, param_infos): refs = self._update_converted_params_from_tensor(converted_named_tensors) if self.use_distribute: - if self._is_distribute_src_rank: - refs.extend(self._update_bucket_weights_from_distributed(converted_named_tensors)) + if self.distributed_weight_updator._is_pp_src_rank: + refs.extend( + self.distributed_weight_updator._update_bucket_weights_from_distributed(converted_named_tensors) + ) ray.get(refs) if self.use_distribute: - if self._is_distribute_src_rank: + if self.distributed_weight_updator._is_pp_src_rank: converted_named_tensors.clear() - ray.get(self.rollout_engine_lock.release.remote()) + ray.get(self.distributed_weight_updator.rollout_engine_lock.release.remote()) def _update_converted_params_from_tensor(self, converted_named_tensors): if use_flattened_tensor_bucket: @@ -513,30 +485,6 @@ def _update_converted_params_from_tensor(self, converted_named_tensors): return refs return [] - def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar=None): - # lock the rollout engines to prevent dead lock on broadcast. - while not ray.get(self.rollout_engine_lock.acquire.remote()): - time.sleep(0.1) - - refs = [ - engine.update_weights_from_distributed.remote( - names=[name for name, _ in converted_named_tensors], - dtypes=[param.dtype for _, param in converted_named_tensors], - shapes=[param.shape for _, param in converted_named_tensors], - group_name=self._group_name, - weight_version=str(self.weight_version), - ) - for engine in self.rollout_engines_distribute - ] - - handles = [] - for _, param in converted_named_tensors: - handles.append(dist.broadcast(param.data, 0, group=self._model_update_groups, async_op=True)) - for handle in handles: - handle.wait() - - return refs - class UpdateWeightFromDistributed: def __init__(self, args, model, weights, *, model_name, quantization_config, vocab_size): @@ -547,7 +495,7 @@ def __init__(self, args, model, weights, *, model_name, quantization_config, voc self.quantization_config = quantization_config self.weight_version = 0 - def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): + def connect_rollout_engines(self, rollout_engines, rollout_engine_lock, only_pp_0: bool = False): self.rollout_engines = rollout_engines self.rollout_engine_lock = rollout_engine_lock @@ -557,6 +505,12 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self._is_pp_src_rank = ( mpu.get_data_parallel_rank(with_context_parallel=True) == 0 and mpu.get_tensor_model_parallel_rank() == 0 ) + if only_pp_0: + self._is_pp_src_rank = ( + mpu.get_data_parallel_rank(with_context_parallel=True) == 0 + and mpu.get_tensor_model_parallel_rank() == 0 + and mpu.get_pipeline_model_parallel_rank() == 0 + ) pp_rank = mpu.get_pipeline_model_parallel_rank() if self._is_pp_src_rank: self._group_name = f"slime-pp_{pp_rank}" @@ -566,8 +520,7 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): with socket.socket() as sock: sock.bind(("", 0)) master_port = sock.getsockname()[1] - world_size = self.args.rollout_num_gpus + 1 - + world_size = len(rollout_engines) * self.args.rollout_num_gpus_per_engine + 1 refs = [ engine.init_weights_update_group.remote( master_address, @@ -690,9 +643,14 @@ def _update_expert_bucket_weights_from_distributed(self, named_tensors, pbar=Non converted_hf_tensors = [] for name, param in all_gathered_params: converted_hf_tensors += convert_to_hf(self.args, self.model_name, name, param, self.quantization_config) - self._update_bucket_weights_from_distributed(converted_hf_tensors, pbar=pbar) - def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar=None): + refs = self._update_bucket_weights_from_distributed(converted_hf_tensors) + ray.get(refs) + converted_hf_tensors.clear() + ray.get(self.rollout_engine_lock.release.remote()) + pbar.update(1) + + def _update_bucket_weights_from_distributed(self, converted_named_tensors): # lock the rollout engines to prevent dead lock on broadcast. while not ray.get(self.rollout_engine_lock.acquire.remote()): time.sleep(0.1) @@ -714,7 +672,4 @@ def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar= for handle in handles: handle.wait() - ray.get(refs) - converted_named_tensors.clear() - ray.get(self.rollout_engine_lock.release.remote()) - pbar.update(1) + return refs From 52d6c86a1d21abd0725ea1b35df3a1c67f0dd2c4 Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 09:39:27 +0000 Subject: [PATCH 12/15] [Fix] ppo update_from_distribute --- slime/backends/megatron_utils/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/slime/backends/megatron_utils/model.py b/slime/backends/megatron_utils/model.py index 82bd98335a..f09f52dc88 100644 --- a/slime/backends/megatron_utils/model.py +++ b/slime/backends/megatron_utils/model.py @@ -477,7 +477,7 @@ def train(rollout_id, model, optimizer, opt_param_scheduler, data_iterator, num_ log_dict[f"train/{role_tag}lr-pg_{param_group_id}"] = opt_param_scheduler.get_lr(param_group) if args.use_wandb: - log_dict[f"train/{role_tag}step"] = accumulated_step_id + log_dict["train/step"] = accumulated_step_id wandb.log(log_dict) if args.ci_test: From b946ef7d923bbe2a4d7c6b3e49907e2d827ce63e Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 11:26:56 +0000 Subject: [PATCH 13/15] [Fix] ppo update_from_distribute --- slime/backends/megatron_utils/actor.py | 53 +++--- slime/backends/megatron_utils/cp_utils.py | 10 +- slime/backends/megatron_utils/data.py | 11 +- .../megatron_utils/update_weight_utils.py | 162 ++++++++++-------- slime/utils/ppo_utils.py | 92 ++-------- 5 files changed, 143 insertions(+), 185 deletions(-) diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index f787ee97f5..07d26c08b8 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -261,19 +261,18 @@ def train(self, rollout_id, rollout_data_ref): def train_critic(self, rollout_id, rollout_data): # Create data iterator for log_probs and train. data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) - values = forward_only( - get_values, - self.args, - self.model, - data_iterator, - num_microbatches, + rollout_data.update( + forward_only( + get_values, + self.args, + self.model, + data_iterator, + num_microbatches, + ) ) if rollout_id >= self.args.num_critic_only_steps: - sync_data = sync_actor_critic_data(self.args, values, self._actor_critic_groups) - rollout_data.update(sync_data) - - rollout_data.update(values) + sync_actor_critic_data(self.args, rollout_data, self._actor_critic_groups) compute_advantages_and_returns(self.args, rollout_data) @@ -297,32 +296,32 @@ def train_actor(self, rollout_id, rollout_data): if "ref" in self.weights: if self.args.use_routing_replay: os.environ["ROUTING_REPLAY_STAGE"] = "fallthrough" - ref_log_probs = self.compute_log_prob( - "ref", - data_iterator, - num_microbatches, - store_prefix="ref_", + rollout_data.update( + self.compute_log_prob( + "ref", + data_iterator, + num_microbatches, + store_prefix="ref_", + ) ) - rollout_data.update(ref_log_probs) if self.args.use_routing_replay: os.environ["ROUTING_REPLAY_STAGE"] = "record" - log_probs = self.compute_log_prob( - "old_actor" if self.args.keep_old_actor else "actor", - data_iterator, - num_microbatches, - store_prefix="", + rollout_data.update( + self.compute_log_prob( + "old_actor" if self.args.keep_old_actor else "actor", + data_iterator, + num_microbatches, + store_prefix="", + ) ) - rollout_data.update(log_probs) + if self.args.use_critic: - if self.args.kl_coef != 0 or self.args.use_kl_loss: - log_probs.update(ref_log_probs) - sync_data = sync_actor_critic_data( + sync_actor_critic_data( self.args, - log_probs, + rollout_data, self._actor_critic_groups, ) - rollout_data.update(sync_data) # when there is old actor, we need to update the model params to actor manually if "old_actor" in self.weights: diff --git a/slime/backends/megatron_utils/cp_utils.py b/slime/backends/megatron_utils/cp_utils.py index 6ee9193963..3ba6adfdf2 100644 --- a/slime/backends/megatron_utils/cp_utils.py +++ b/slime/backends/megatron_utils/cp_utils.py @@ -1,3 +1,5 @@ +from typing import Union + import torch import torch.distributed as dist import torch.nn.functional as F @@ -170,7 +172,7 @@ def slice_with_cp(tokens: torch.Tensor, pad_value): return torch.cat([tokens[start_1:end_1], tokens[start_2:end_2]]) -def slice_log_prob_with_cp(log_prob: list[float], total_length: int, response_length: int): +def slice_log_prob_with_cp(log_prob: Union[list[float], torch.Tensor], total_length: int, response_length: int): assert len(log_prob) == response_length cp_size = mpu.get_context_parallel_world_size() @@ -183,4 +185,8 @@ def slice_log_prob_with_cp(log_prob: list[float], total_length: int, response_le chunk_1 = log_prob[logits_offset[0][0] - (prompt_length - 1) : logits_offset[0][1] - (prompt_length - 1)] chunk_2 = log_prob[logits_offset[1][0] - (prompt_length - 1) : logits_offset[1][1] - (prompt_length - 1)] - return chunk_1 + chunk_2 + + if isinstance(log_prob, list): + return chunk_1 + chunk_2 + else: + return torch.cat([chunk_1, chunk_2], dim=0) diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 5ab1c88939..9c759d0600 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -390,10 +390,10 @@ def log_perf_data(rollout_id, args): def sync_actor_critic_data( args, - data: Optional[dict[str, list[torch.Tensor]]] = None, + rollout_data: Optional[dict[str, list[torch.Tensor]]] = None, group: Optional[dist.ProcessGroup] = None, ): - values, log_probs, ref_log_probs = map(data.get, ("values", "log_probs", "ref_log_probs")) + values, log_probs, ref_log_probs = map(rollout_data.get, ("values", "log_probs", "ref_log_probs")) # return None when not pp last stage if not values and not log_probs: @@ -403,7 +403,6 @@ def sync_actor_critic_data( if not values: values = [torch.empty_like(log_prob) for log_prob in log_probs] - for value in values: handles.append(dist.broadcast(value, src=1, group=group, async_op=True)) @@ -411,13 +410,11 @@ def sync_actor_critic_data( if not log_probs: ref_log_probs = [torch.empty_like(value) for value in values] log_probs = [torch.empty_like(value) for value in values] - else: - ref_log_probs = ref_log_probs - log_probs = log_probs for ref_log_prob, log_prob in zip(ref_log_probs, log_probs): handles.append(dist.broadcast(ref_log_prob, src=0, group=group, async_op=True)) handles.append(dist.broadcast(log_prob, src=0, group=group, async_op=True)) for handle in handles: handle.wait() - return {"values": values, "log_probs": log_probs, "ref_log_probs": ref_log_probs} + + rollout_data.update({"values": values, "log_probs": log_probs, "ref_log_probs": ref_log_probs}) diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index d44b4d6a41..668dba6cd1 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -313,18 +313,18 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self.use_distribute = len(rollout_engines) > colocate_engine_nums if self.use_distribute: - self.distributed_weight_updator = UpdateWeightFromDistributed( - args=self.args, - model=self.model, - weights=self.weights, - model_name=self.model_name, - quantization_config=self.quantization_config, - vocab_size=self.vocab_size, - ) - self.distributed_weight_updator.connect_rollout_engines( - rollout_engines[colocate_engine_nums:], rollout_engine_lock, True - ) self.rollout_engines = rollout_engines[:colocate_engine_nums] + self.distributed_rollout_engines = rollout_engines[colocate_engine_nums:] + self._is_distributed_src_rank = ( + mpu.get_data_parallel_rank(with_context_parallel=True) == 0 + and mpu.get_tensor_model_parallel_rank() == 0 + and mpu.get_pipeline_model_parallel_rank() == 0 + ) + self._group_name = "slime" + if self._is_distributed_src_rank: + self._model_update_groups = connect_rollout_engines_from_distributed( + self.args, self._group_name, self.distributed_rollout_engines + ) # Here we assume the gpu id of rollout engines and train actors are the same. for i, engine in enumerate(self.rollout_engines): @@ -421,17 +421,19 @@ def _update_bucket_weights_from_tensor(self, param_infos): refs = self._update_converted_params_from_tensor(converted_named_tensors) - if self.use_distribute: - if self.distributed_weight_updator._is_pp_src_rank: - refs.extend( - self.distributed_weight_updator._update_bucket_weights_from_distributed(converted_named_tensors) + if self.use_distribute and self._is_distributed_src_rank: + refs.extend( + update_weights_from_distributed( + self.args, + self._group_name, + self._model_update_groups, + self.weight_version, + self.distributee_rollout_engines, + converted_named_tensors, ) - ray.get(refs) + ) - if self.use_distribute: - if self.distributed_weight_updator._is_pp_src_rank: - converted_named_tensors.clear() - ray.get(self.distributed_weight_updator.rollout_engine_lock.release.remote()) + ray.get(refs) def _update_converted_params_from_tensor(self, converted_named_tensors): if use_flattened_tensor_bucket: @@ -495,7 +497,7 @@ def __init__(self, args, model, weights, *, model_name, quantization_config, voc self.quantization_config = quantization_config self.weight_version = 0 - def connect_rollout_engines(self, rollout_engines, rollout_engine_lock, only_pp_0: bool = False): + def connect_rollout_engines(self, rollout_engines, rollout_engine_lock): self.rollout_engines = rollout_engines self.rollout_engine_lock = rollout_engine_lock @@ -505,41 +507,14 @@ def connect_rollout_engines(self, rollout_engines, rollout_engine_lock, only_pp_ self._is_pp_src_rank = ( mpu.get_data_parallel_rank(with_context_parallel=True) == 0 and mpu.get_tensor_model_parallel_rank() == 0 ) - if only_pp_0: - self._is_pp_src_rank = ( - mpu.get_data_parallel_rank(with_context_parallel=True) == 0 - and mpu.get_tensor_model_parallel_rank() == 0 - and mpu.get_pipeline_model_parallel_rank() == 0 - ) pp_rank = mpu.get_pipeline_model_parallel_rank() if self._is_pp_src_rank: self._group_name = f"slime-pp_{pp_rank}" if self._is_pp_src_rank: - master_address = ray._private.services.get_node_ip_address() - with socket.socket() as sock: - sock.bind(("", 0)) - master_port = sock.getsockname()[1] - world_size = len(rollout_engines) * self.args.rollout_num_gpus_per_engine + 1 - refs = [ - engine.init_weights_update_group.remote( - master_address, - master_port, - i * self.args.rollout_num_gpus_per_engine + 1, - world_size, - self._group_name, - backend="nccl", - ) - for i, engine in enumerate(self.rollout_engines) - ] - self._model_update_groups = init_process_group( - backend="nccl", - init_method=f"tcp://{master_address}:{master_port}", - world_size=world_size, - rank=0, - group_name=self._group_name, + self._model_update_groups = connect_rollout_engines_from_distributed( + self.args, self._group_name, rollout_engines ) - ray.get(refs) @torch.no_grad() def update_weights(self): @@ -644,32 +619,75 @@ def _update_expert_bucket_weights_from_distributed(self, named_tensors, pbar=Non for name, param in all_gathered_params: converted_hf_tensors += convert_to_hf(self.args, self.model_name, name, param, self.quantization_config) - refs = self._update_bucket_weights_from_distributed(converted_hf_tensors) - ray.get(refs) - converted_hf_tensors.clear() - ray.get(self.rollout_engine_lock.release.remote()) - pbar.update(1) + self._update_bucket_weights_from_distributed(converted_hf_tensors) - def _update_bucket_weights_from_distributed(self, converted_named_tensors): + def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar=None): # lock the rollout engines to prevent dead lock on broadcast. while not ray.get(self.rollout_engine_lock.acquire.remote()): time.sleep(0.1) - refs = [ - engine.update_weights_from_distributed.remote( - names=[name for name, _ in converted_named_tensors], - dtypes=[param.dtype for _, param in converted_named_tensors], - shapes=[param.shape for _, param in converted_named_tensors], - group_name=self._group_name, - weight_version=str(self.weight_version), - ) - for engine in self.rollout_engines - ] + refs = update_weights_from_distributed( + self.args, + self._group_name, + self._model_update_groups, + self.weight_version, + self.rollout_engines, + converted_named_tensors, + ) - handles = [] - for _, param in converted_named_tensors: - handles.append(dist.broadcast(param.data, 0, group=self._model_update_groups, async_op=True)) - for handle in handles: - handle.wait() + ray.get(refs) + converted_named_tensors.clear() + ray.get(self.rollout_engine_lock.release.remote()) + pbar.update(1) + + +def connect_rollout_engines_from_distributed(args, group_name: str, rollout_engines): + master_address = ray._private.services.get_node_ip_address() + with socket.socket() as sock: + sock.bind(("", 0)) + master_port = sock.getsockname()[1] + world_size = len(rollout_engines) * args.rollout_num_gpus_per_engine + 1 + + refs = [ + engine.init_weights_update_group.remote( + master_address, + master_port, + i * args.rollout_num_gpus_per_engine + 1, + world_size, + group_name, + backend="nccl", + ) + for i, engine in enumerate(rollout_engines) + ] + model_update_groups = init_process_group( + backend="nccl", + init_method=f"tcp://{master_address}:{master_port}", + world_size=world_size, + rank=0, + group_name=group_name, + ) + ray.get(refs) + return model_update_groups + + +def update_weights_from_distributed( + args, group_name: str, group, weight_version, rollout_engines, converted_named_tensors +): + refs = [ + engine.update_weights_from_distributed.remote( + names=[name for name, _ in converted_named_tensors], + dtypes=[param.dtype for _, param in converted_named_tensors], + shapes=[param.shape for _, param in converted_named_tensors], + group_name=group_name, + weight_version=str(weight_version), + ) + for engine in rollout_engines + ] + + handles = [] + for _, param in converted_named_tensors: + handles.append(dist.broadcast(param.data, 0, group=group, async_op=True)) + for handle in handles: + handle.wait() - return refs + return refs diff --git a/slime/utils/ppo_utils.py b/slime/utils/ppo_utils.py index b3d7060cab..60472c794a 100644 --- a/slime/utils/ppo_utils.py +++ b/slime/utils/ppo_utils.py @@ -164,13 +164,13 @@ def get_reinforce_plus_plus_returns( final_returns_chunks = [] for i in range(len(rewards)): local_kl_chunk = kl[i] - device, dtype = local_kl_chunk.device, local_kl_chunk.dtype total_len, response_len = total_lengths[i], response_lengths[i] - prompt_len = total_len - response_len if cp_size > 1: # Step 1,2:Gather all chunks and token_offsets from all ranks and reconstruct the full response tensor by splitting and placing each part - full_kl_response = all_gather_with_cp(local_kl_chunk, total_len, response_len, prompt_len, device, dtype) + from slime.backends.megatron_utils.cp_utils import all_gather_with_cp + + full_kl_response = all_gather_with_cp(local_kl_chunk, total_len, response_len) else: full_kl_response = local_kl_chunk @@ -190,9 +190,9 @@ def get_reinforce_plus_plus_returns( # Step 4: Pick up the results corresponding to our local chunk's parts. if cp_size > 1: - local_returns_chunk = get_local_chunk_with_cp( - returns_for_seq, total_len, response_len, prompt_len, device, dtype - ) + from slime.backends.megatron_utils.cp_utils import slice_log_prob_with_cp + + local_returns_chunk = slice_log_prob_with_cp(returns_for_seq, total_len, response_len) else: local_returns_chunk = returns_for_seq @@ -261,11 +261,11 @@ def get_advantages_and_returns( from megatron.core import mpu cp_size = mpu.get_context_parallel_world_size() - if cp_size > 0: - device, dtype = rewards.device, rewards.dtype - prompt_len = total_len - response_len - full_rewards = all_gather_with_cp(rewards, total_len, response_len, prompt_len, device, dtype) - full_values = all_gather_with_cp(values, total_len, response_len, prompt_len, device, dtype) + if cp_size > 1: + from slime.backends.megatron_utils.cp_utils import all_gather_with_cp + + full_rewards = all_gather_with_cp(rewards, total_len, response_len) + full_values = all_gather_with_cp(values, total_len, response_len) else: full_rewards = rewards full_values = values @@ -282,8 +282,10 @@ def get_advantages_and_returns( full_returns = full_advantages + full_values if cp_size > 0: - advantages = get_local_chunk_with_cp(full_advantages, total_len, response_len, prompt_len, device, dtype) - returns = get_local_chunk_with_cp(full_returns, total_len, response_len, prompt_len, device, dtype) + from slime.backends.megatron_utils.cp_utils import slice_log_prob_with_cp + + advantages = slice_log_prob_with_cp(full_advantages, total_len, response_len) + returns = slice_log_prob_with_cp(full_returns, total_len, response_len) else: advantages = full_advantages returns = full_returns @@ -309,67 +311,3 @@ def calculate_log_probs_and_entropy(logits, tokens, tp_group, with_entropy: bool else: entropy = None return log_prob, entropy - - -def all_gather_with_cp(local_chunk, total_len, response_len, prompt_len, device, dtype): - from megatron.core import mpu - - from slime.backends.megatron_utils.cp_utils import get_logits_and_tokens_offset_with_cp - - cp_size = mpu.get_context_parallel_world_size() - - # Step 1: Gather all chunks and token_offsets from all ranks - _, _, _, token_offsets = get_logits_and_tokens_offset_with_cp(total_len, response_len) - - object_to_gather = {"chunk": local_chunk.cpu(), "offsets": token_offsets} - gathered_objects = [None] * cp_size - dist.all_gather_object(gathered_objects, object_to_gather, group=mpu.get_context_parallel_group()) - - # Step 2: Reconstruct the full response tensor by splitting and placing each part. - full_response = torch.zeros(response_len, device=device, dtype=dtype) - for obj in gathered_objects: - chunk = obj["chunk"].to(device) - global_offsets = obj["offsets"] - - # Calculate the lengths of part_0 and part_1 for this specific chunk. - s0, e0 = global_offsets[0] - s1, e1 = global_offsets[1] - res_s0, res_e0 = max(0, s0 - prompt_len), max(0, e0 - prompt_len) - res_s1, res_e1 = max(0, s1 - prompt_len), max(0, e1 - prompt_len) - len0 = res_e0 - res_s0 - len1 = res_e1 - res_s1 - if chunk.numel() > 0: - # Split the received contiguous chunk back into its zigzag parts. - part_0, part_1 = torch.split(chunk, [len0, len1]) - - # Place each part in its own correct location. - if part_0.numel() > 0: - full_response[res_s0:res_e0] = part_0 - if part_1.numel() > 0: - full_response[res_s1:res_e1] = part_1 - - return full_response - - -def get_local_chunk_with_cp(full_response, total_len, response_len, prompt_len, device, dtype): - from slime.backends.megatron_utils.cp_utils import get_logits_and_tokens_offset_with_cp - - _, _, _, token_offsets = get_logits_and_tokens_offset_with_cp(total_len, response_len) - - local_returns_chunk_parts = [] - local_s0, local_e0 = token_offsets[0] - local_s1, local_e1 = token_offsets[1] - local_res_s0, local_res_e0 = max(0, local_s0 - prompt_len), max(0, local_e0 - prompt_len) - local_res_s1, local_res_e1 = max(0, local_s1 - prompt_len), max(0, local_e1 - prompt_len) - - if local_res_e0 > local_res_s0: - local_returns_chunk_parts.append(full_response[local_res_s0:local_res_e0]) - if local_res_e1 > local_res_s1: - local_returns_chunk_parts.append(full_response[local_res_s1:local_res_e1]) - - local_chunk = ( - torch.cat(local_returns_chunk_parts) - if local_returns_chunk_parts - else torch.tensor([], device=device, dtype=dtype) - ) - return local_chunk From ce6b693cc244b40e35c1fb6487cf850a88299b2c Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 11:38:27 +0000 Subject: [PATCH 14/15] [Fix] ppo update_from_distribute --- slime/backends/megatron_utils/data.py | 4 ++-- slime/backends/megatron_utils/update_weight_utils.py | 8 +++----- 2 files changed, 5 insertions(+), 7 deletions(-) diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 9c759d0600..d19d902d35 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -395,9 +395,9 @@ def sync_actor_critic_data( ): values, log_probs, ref_log_probs = map(rollout_data.get, ("values", "log_probs", "ref_log_probs")) - # return None when not pp last stage + # return when not the pp last stage if not values and not log_probs: - return {} + return handles = [] diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index 668dba6cd1..9e8195c992 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -428,7 +428,7 @@ def _update_bucket_weights_from_tensor(self, param_infos): self._group_name, self._model_update_groups, self.weight_version, - self.distributee_rollout_engines, + self.distributed_rollout_engines, converted_named_tensors, ) ) @@ -641,7 +641,7 @@ def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar= pbar.update(1) -def connect_rollout_engines_from_distributed(args, group_name: str, rollout_engines): +def connect_rollout_engines_from_distributed(args, group_name, rollout_engines): master_address = ray._private.services.get_node_ip_address() with socket.socket() as sock: sock.bind(("", 0)) @@ -670,9 +670,7 @@ def connect_rollout_engines_from_distributed(args, group_name: str, rollout_engi return model_update_groups -def update_weights_from_distributed( - args, group_name: str, group, weight_version, rollout_engines, converted_named_tensors -): +def update_weights_from_distributed(args, group_name, group, weight_version, rollout_engines, converted_named_tensors): refs = [ engine.update_weights_from_distributed.remote( names=[name for name, _ in converted_named_tensors], From 10fd8122dfab41c941d8cba0a0eb0771456f54c4 Mon Sep 17 00:00:00 2001 From: lilei Date: Sun, 28 Sep 2025 11:41:22 +0000 Subject: [PATCH 15/15] [Fix] ppo update_from_distribute --- slime/backends/megatron_utils/update_weight_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/slime/backends/megatron_utils/update_weight_utils.py b/slime/backends/megatron_utils/update_weight_utils.py index 9e8195c992..654a8364bb 100644 --- a/slime/backends/megatron_utils/update_weight_utils.py +++ b/slime/backends/megatron_utils/update_weight_utils.py @@ -619,7 +619,7 @@ def _update_expert_bucket_weights_from_distributed(self, named_tensors, pbar=Non for name, param in all_gathered_params: converted_hf_tensors += convert_to_hf(self.args, self.model_name, name, param, self.quantization_config) - self._update_bucket_weights_from_distributed(converted_hf_tensors) + self._update_bucket_weights_from_distributed(converted_hf_tensors, pbar) def _update_bucket_weights_from_distributed(self, converted_named_tensors, pbar=None): # lock the rollout engines to prevent dead lock on broadcast.