diff --git a/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_grpo_npu.py b/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_grpo_npu.py index 257edd0766..da47268341 100644 --- a/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_grpo_npu.py +++ b/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_grpo_npu.py @@ -1,6 +1,5 @@ import os -import slime.utils.misc as U from slime.utils.external_utils.command_utils import execute_train_npu MODEL_NAME = os.environ.get("SLIME_SCRIPT_MODEL_NAME", "Qwen3-VL-2B-Instruct") @@ -78,7 +77,6 @@ def execute(): "--weight-decay 0.1 " "--adam-beta1 0.9 " "--adam-beta2 0.98 " - "--optimizer-cpu-offload " "--overlap-cpu-optimizer-d2h-h2d " "--use-precision-aware-optimizer " @@ -120,7 +118,9 @@ def execute(): ) misc_args = ( - "--actor-num-nodes 1 " f"--actor-num-gpus-per-node 8 " f"--rollout-num-gpus 8 " + "--actor-num-nodes 1 " + "--actor-num-gpus-per-node 8 " + "--rollout-num-gpus 8 " "--no-gradient-accumulation-fusion " "--use-flash-attn " ) diff --git a/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_ppo_npu.py b/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_ppo_npu.py index ae3559f51e..b14ceadf00 100644 --- a/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_ppo_npu.py +++ b/examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_ppo_npu.py @@ -1,6 +1,6 @@ import os +import tempfile -import slime.utils.misc as U from slime.utils.external_utils.command_utils import execute_train_npu MODEL_NAME = os.environ.get("SLIME_SCRIPT_MODEL_NAME", "Qwen3-VL-2B-Instruct") @@ -29,6 +29,10 @@ def get_megatron_model_type(model_name: str) -> str: def execute(): + critic_config = tempfile.NamedTemporaryFile("w", suffix=".yaml", delete=False) + critic_config.write("lr: 1e-5\n") + critic_config.close() + ckpt_args = f"--hf-checkpoint /path/to/model/checkpoints/{MODEL_NAME} " wandb_args = ( @@ -69,7 +73,6 @@ def execute(): "--eps-clip 4e-4 " "--num-critic-only-steps 1 " "--normalize-advantages " - "--critic-lr 1e-5 " ) optimizer_args = ( @@ -79,7 +82,6 @@ def execute(): "--weight-decay 0.1 " "--adam-beta1 0.9 " "--adam-beta2 0.98 " - "--optimizer-cpu-offload " "--overlap-cpu-optimizer-d2h-h2d " "--use-precision-aware-optimizer " @@ -122,9 +124,7 @@ def execute(): misc_args = ( "--actor-num-nodes 1 " - "--actor-num-gpus-per-node 4 " - "--critic-num-nodes 1 " - "--critic-num-gpus-per-node 4 " + "--actor-num-gpus-per-node 8 " "--rollout-num-gpus 8 " "--no-gradient-accumulation-fusion " "--use-flash-attn " @@ -138,6 +138,7 @@ def execute(): exit() train_args = ( + f"--critic-config-path {critic_config.name} " f"{ckpt_args} " f"{rollout_args} " f"{optimizer_args} " diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index f68d665537..adf9ad3a33 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -1,7 +1,6 @@ import logging import os import random -import socket from argparse import Namespace from contextlib import nullcontext @@ -10,14 +9,13 @@ import torch import torch.distributed as dist from megatron.core import mpu -from ray.actor import ActorHandle from torch_memory_saver import torch_memory_saver from transformers import AutoConfig, AutoTokenizer from slime.ray.train_actor import TrainRayActor from slime.utils import train_dump_utils from slime.utils.data import process_rollout_data -from slime.utils.distributed_utils import get_gloo_group, init_process_group +from slime.utils.distributed_utils import get_gloo_group from slime.utils.logging_utils import init_tracking from slime.utils.memory_utils import clear_memory, print_memory from slime.utils.misc import Box @@ -30,7 +28,7 @@ from ...utils.tensor_backper import TensorBackuper from .checkpoint import load_checkpoint from .cp_utils import slice_log_prob_with_cp, slice_with_cp -from .data import DataIterator, get_data_iterator, log_perf_data, log_rollout_data, sync_actor_critic_data +from .data import DataIterator, get_data_iterator, log_perf_data, log_rollout_data from .initialize import init, is_megatron_main_rank from .loss import compute_advantages_and_returns, get_log_probs_and_entropy, get_values from .model import forward_only, initialize_model_and_optimizer, save, train @@ -83,12 +81,6 @@ def init( logger.info(f"Set torch_memory_saver.memory_margin_bytes to {x}") torch_memory_saver.memory_margin_bytes = x - if role == "critic": - self.args.load = self.args.critic_load - self.args.save = self.args.critic_save - self.args.lr = self.args.critic_lr - self.args.lr_warmup_iters = self.args.critic_lr_warmup_iters - (self.model, self.optimizer, self.opt_param_scheduler, loaded_rollout_id) = initialize_model_and_optimizer( args, role ) @@ -360,9 +352,9 @@ def compute_log_prob( store_prefix=store_prefix, ) - def train(self, rollout_id: int, rollout_data_ref: Box) -> None: + def train(self, rollout_id: int, rollout_data_ref: Box, external_data=None): if self.args.debug_rollout_only: - return + return None if self.args.offload_train: self.wake_up() @@ -371,25 +363,22 @@ def train(self, rollout_id: int, rollout_data_ref: Box) -> None: rollout_data = self._get_rollout_data(rollout_data_ref) if self.role == "critic": - return self.train_critic(rollout_id, rollout_data) + result = self.train_critic(rollout_id, rollout_data) else: - return self.train_actor(rollout_id, rollout_data) + self.train_actor(rollout_id, rollout_data, external_data=external_data) + result = None - def train_critic(self, rollout_id: int, rollout_data: RolloutBatch) -> None: - # Create data iterator for log_probs and train. + if self.args.offload_train: + self.sleep() + + return result + + def train_critic(self, rollout_id: int, rollout_data: RolloutBatch): + """Train critic and return CPU values (used as old-values for the next actor train).""" data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) - rollout_data.update( - forward_only( - get_values, - self.args, - self.model, - data_iterator, - num_microbatches, - ) - ) - if rollout_id >= self.args.num_critic_only_steps and not self.args.critic_train_only: - sync_actor_critic_data(self.args, rollout_data, self._actor_critic_groups) + # Compute current critic values (used as old_values for value loss and for actor advantages). + rollout_data.update(forward_only(get_values, self.args, self.model, data_iterator, num_microbatches)) compute_advantages_and_returns(self.args, rollout_data) @@ -403,7 +392,13 @@ def train_critic(self, rollout_id: int, rollout_data: RolloutBatch) -> None: num_microbatches, ) - def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None: + if mpu.is_pipeline_last_stage() and "values" in rollout_data: + from slime.backends.megatron_utils.data import tensors_to_cpu + + return {"values": tensors_to_cpu(rollout_data["values"])} + return {} + + def train_actor(self, rollout_id: int, rollout_data: RolloutBatch, external_data=None) -> None: # Create data iterator for log_probs and train. data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) @@ -455,11 +450,12 @@ def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None: RoutingReplay.clear_all_forward() if self.args.use_critic: - sync_actor_critic_data( - self.args, - rollout_data, - self._actor_critic_groups, - ) + if external_data is not None and mpu.is_pipeline_last_stage(): + values = external_data.get("values") + if values is not None: + from slime.backends.megatron_utils.data import tensors_to_gpu + + rollout_data["values"] = tensors_to_gpu(values) if self._active_model_tag != "actor": self._switch_model("actor") @@ -623,26 +619,3 @@ def load_other_checkpoint(self, model_tag: str, path: str) -> None: self.weights_backuper.backup(model_tag) self._active_model_tag = model_tag - - def connect_actor_critic( - self, - actor_handle: ActorHandle | None = None, - master_address: str | None = None, - master_port: int | None = None, - ) -> None: - if self.role == "actor": - master_address = ray.util.get_node_ip_address() - with socket.socket() as sock: - sock.bind(("", 0)) - master_port = sock.getsockname()[1] - actor_handle.connect_actor_critic.remote(master_address=master_address, master_port=master_port) - - group_name = "actor_critic" - world_size = 2 - self._actor_critic_groups = init_process_group( - backend="nccl", - init_method=f"tcp://{master_address}:{master_port}", - world_size=world_size, - rank=0 if self.role == "actor" else 1, - group_name=group_name, - ) diff --git a/slime/backends/megatron_utils/arguments.py b/slime/backends/megatron_utils/arguments.py index 9444fe88a0..c504652385 100644 --- a/slime/backends/megatron_utils/arguments.py +++ b/slime/backends/megatron_utils/arguments.py @@ -11,6 +11,7 @@ def validate_args(args): """Run megatron's own validate_args plus slime-specific megatron validations.""" + _megatron_validate_args(args) # always use varlen @@ -116,9 +117,6 @@ def megatron_parse_args(extra_args_provider, skip_hf_validate=False): _hf_validate_args(args, hf_config) args.rank = 0 - if args.critic_train_only: - args.world_size = args.critic_num_nodes * args.critic_num_gpus_per_node - else: - args.world_size = args.actor_num_nodes * args.actor_num_gpus_per_node + args.world_size = args.actor_num_nodes * args.actor_num_gpus_per_node args = _set_default_megatron_args(args) return args diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 8a7f768b3b..19db1f475a 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -610,53 +610,32 @@ def log_perf_data(rollout_id: int, args: Namespace) -> None: ) -def sync_actor_critic_data( - args: Namespace, - rollout_data: RolloutBatch | None = None, - group: dist.ProcessGroup | None = None, -) -> None: +def tensors_to_cpu(tensor_list): + """Move a list of GPU tensors to CPU for Ray object store transfer. + + Args: + tensor_list: List of GPU tensors, or None. + + Returns: + List of CPU tensors (detached), or None if input is None. """ - Broadcast `values` (from critic) and optionally `log_probs`/`ref_log_probs` - (from actor) across PP ranks to align data dependencies. + if tensor_list is None: + return None + return [t.detach().cpu() for t in tensor_list] + - - Values are broadcast from src=1. - - Log-probs and ref-log-probs are broadcast from src=0 when KL is used. - Updates `rollout_data` in place with the synchronized tensors. +def tensors_to_gpu(tensor_list, device=None): + """Move a list of CPU tensors back to GPU. + + Args: + tensor_list: List of CPU tensors, or None. + device: Target CUDA device. If None, uses current device. + + Returns: + List of GPU tensors, or None if input is None. """ - log_probs_key = "log_probs" if not args.use_rollout_logprobs else "rollout_log_probs" - values, log_probs, ref_log_probs = map(rollout_data.get, ("values", log_probs_key, "ref_log_probs")) - - # return when not the pp last stage - if not values and not log_probs: - return - - handles = [] - - 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)) - - if args.kl_coef != 0 or args.use_kl_loss: - if not log_probs: - log_probs = [torch.empty_like(value) for value in values] - if not ref_log_probs: - ref_log_probs = [torch.empty_like(value) for value in values] - for ref_log_prob, log_prob in zip(ref_log_probs, log_probs, strict=False): - handles.append(dist.broadcast(log_prob, src=0, group=group, async_op=True)) - handles.append(dist.broadcast(ref_log_prob, src=0, group=group, async_op=True)) - - for handle in handles: - handle.wait() - - rollout_data.update( - { - k: v - for k, v in { - "values": values, - log_probs_key: log_probs, - "ref_log_probs": ref_log_probs, - }.items() - if v is not None - } - ) + if tensor_list is None: + return None + if device is None: + device = torch.cuda.current_device() + return [t.to(device=device, dtype=torch.float32) for t in tensor_list] diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index 5707c57cca..c7ad36839a 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -581,6 +581,10 @@ def compute_advantages_and_returns(args: Namespace, rollout_data: RolloutBatch) Early returns if both `log_probs` and `values` are None (intermediate pipeline stages). + If ``args.custom_advantage_function_path`` is set, it is called after KL computation + and must populate ``rollout_data["advantages"]`` and + ``rollout_data["returns"]``. + Args: args: Configuration specifying estimator type, KL coefficient, normalization settings, and other hyperparameters. @@ -615,8 +619,14 @@ def compute_advantages_and_returns(args: Namespace, rollout_data: RolloutBatch) ) for i in range(len(log_probs)) ] + rollout_data["kl"] = kl + + if args.custom_advantage_function_path is not None: + custom_adv_fn = load_function(args.custom_advantage_function_path) + custom_adv_fn(args, rollout_data) + advantages, returns = rollout_data["advantages"], rollout_data["returns"] - if args.advantage_estimator in ["grpo", "gspo"]: + elif args.advantage_estimator in ["grpo", "gspo"]: rewards = torch.tensor(rewards, dtype=torch.float32, device=kl[0].device) returns = get_grpo_returns(rewards, kl) # TODO: is the copy necessary? diff --git a/slime/ray/actor_group.py b/slime/ray/actor_group.py index 451c48396d..c9ce215558 100644 --- a/slime/ray/actor_group.py +++ b/slime/ray/actor_group.py @@ -108,9 +108,25 @@ def async_init(self, args, role, with_ref=False, with_opd_teacher=False): for actor in self._actor_handlers ] - def async_train(self, rollout_id, rollout_data_ref): - """Do one rollout training""" - return [actor.train.remote(rollout_id, rollout_data_ref) for actor in self._actor_handlers] + def async_train(self, rollout_id, rollout_data_ref, external_data=None): + """Do one rollout training. Returns a list of Ray refs (one per worker). + + For critics, each ref resolves to ``{"values": [cpu tensors...]}`` (or ``{}`` + for non-last-PP-stage workers). Actor refs resolve to ``None``. + + ``external_data`` may be a list (one item per worker) or a single dict + broadcast to all workers. + """ + if isinstance(external_data, list): + assert len(external_data) == len(self._actor_handlers) + return [ + actor.train.remote(rollout_id, rollout_data_ref, external_data=ed) + for actor, ed in zip(self._actor_handlers, external_data, strict=False) + ] + return [ + actor.train.remote(rollout_id, rollout_data_ref, external_data=external_data) + for actor in self._actor_handlers + ] def save_model(self, rollout_id, force_sync=False): """Save actor model""" @@ -129,13 +145,5 @@ def offload(self): def clear_memory(self): return ray.get([actor.clear_memory.remote() for actor in self._actor_handlers]) - def connect(self, critic_group): - return ray.get( - [ - actor.connect_actor_critic.remote(critic) - for actor, critic in zip(self._actor_handlers, critic_group._actor_handlers, strict=False) - ] - ) - def set_rollout_manager(self, rollout_manager): return ray.get([actor.set_rollout_manager.remote(rollout_manager) for actor in self._actor_handlers]) diff --git a/slime/ray/placement_group.py b/slime/ray/placement_group.py index 1487ff17c5..7279a366d3 100644 --- a/slime/ray/placement_group.py +++ b/slime/ray/placement_group.py @@ -77,47 +77,36 @@ def _create_placement_group(num_gpus): def create_placement_groups(args): - """Create placement groups for actor and rollout engines.""" + """Create placement groups for actor, critic, and rollout engines.""" num_gpus = 0 if args.debug_train_only: num_gpus = args.actor_num_nodes * args.actor_num_gpus_per_node rollout_offset = 0 - if args.use_critic: - num_gpus += args.critic_num_nodes * args.critic_num_gpus_per_node - critic_offset = args.actor_num_nodes * args.actor_num_gpus_per_node elif args.debug_rollout_only: num_gpus = args.rollout_num_gpus rollout_offset = 0 elif args.colocate: num_gpus = args.actor_num_nodes * args.actor_num_gpus_per_node rollout_offset = 0 - if args.use_critic: - num_gpus += args.critic_num_nodes * args.critic_num_gpus_per_node - critic_offset = args.actor_num_nodes * args.actor_num_gpus_per_node else: num_gpus = args.actor_num_nodes * args.actor_num_gpus_per_node + args.rollout_num_gpus rollout_offset = args.actor_num_nodes * args.actor_num_gpus_per_node - if args.use_critic: - num_gpus += args.critic_num_nodes * args.critic_num_gpus_per_node - critic_offset = args.actor_num_nodes * args.actor_num_gpus_per_node - rollout_offset += args.critic_num_nodes * args.critic_num_gpus_per_node logger.info(f"Creating placement group with {num_gpus} GPUs...") pg, actor_pg_reordered_bundle_indices, actor_pg_reordered_gpu_ids = _create_placement_group(num_gpus) - rollout_pg_reordered_bundle_indices = actor_pg_reordered_bundle_indices[rollout_offset:] rollout_pg_reordered_gpu_ids = actor_pg_reordered_gpu_ids[rollout_offset:] - if args.use_critic: - critic_pg_reordered_bundle_indices = actor_pg_reordered_bundle_indices[critic_offset:] - critic_pg_reordered_gpu_ids = actor_pg_reordered_gpu_ids[critic_offset:] - return { + result = { "actor": (pg, actor_pg_reordered_bundle_indices, actor_pg_reordered_gpu_ids), - "critic": (pg, critic_pg_reordered_bundle_indices, critic_pg_reordered_gpu_ids) if args.use_critic else None, "rollout": (pg, rollout_pg_reordered_bundle_indices, rollout_pg_reordered_gpu_ids), } + result["critic"] = result["actor"] if args.use_critic else None + + return result + def allocate_train_group(args, num_nodes, num_gpus_per_node, pg, role="actor"): return RayTrainGroup( @@ -137,19 +126,22 @@ def create_training_models(args, pgs, rollout_manager): num_gpus_per_node=args.actor_num_gpus_per_node, pg=pgs["actor"], ) + + critic_model = None if args.use_critic: + from slime.utils.arguments import parse_critic_args + + critic_args = parse_critic_args(args, args.critic_config_path) if args.critic_config_path is not None else args critic_model = allocate_train_group( - args=args, + args=critic_args, num_nodes=args.critic_num_nodes, num_gpus_per_node=args.critic_num_gpus_per_node, pg=pgs["critic"], role="critic", ) - critic_init_handle = critic_model.async_init(args, role="critic", with_ref=False) - else: - critic_model = None + critic_start_rollout_ids = ray.get(critic_model.async_init(critic_model.args, role="critic", with_ref=False)) - start_rollout_ids = ray.get( + actor_start_rollout_ids = ray.get( actor_model.async_init( args, role="actor", @@ -157,13 +149,11 @@ def create_training_models(args, pgs, rollout_manager): with_opd_teacher=args.use_opd and args.opd_type == "megatron", ) ) - + # TODO how to decide rollout start id when critic is involved? For now we just require user to specify it via args. if args.use_critic: - critic_start_rollout_ids = ray.get(critic_init_handle) - if not args.critic_train_only: - actor_model.connect(critic_model) - else: - start_rollout_ids = critic_start_rollout_ids + start_rollout_ids = critic_start_rollout_ids + else: + start_rollout_ids = actor_start_rollout_ids assert len(set(start_rollout_ids)) == 1 diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index 81e9bd9d26..83c6600a72 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -967,11 +967,7 @@ def _compute_rollout_offset(args) -> int: """Offset (in PG bundle slots) where rollout GPUs start.""" if args.debug_train_only or args.debug_rollout_only or args.colocate: return 0 - if args.critic_train_only: - return args.critic_num_nodes * args.critic_num_gpus_per_node offset = args.actor_num_nodes * args.actor_num_gpus_per_node - if args.use_critic: - offset += args.critic_num_nodes * args.critic_num_gpus_per_node return offset @@ -979,11 +975,7 @@ def _compute_megatron_num_gpus(args) -> int: """Total number of megatron (actor + critic) GPU slots in the placement group.""" if args.debug_rollout_only: return 0 - if args.critic_train_only: - return args.critic_num_nodes * args.critic_num_gpus_per_node num = args.actor_num_nodes * args.actor_num_gpus_per_node - if args.use_critic: - num += args.critic_num_nodes * args.critic_num_gpus_per_node return num diff --git a/slime/ray/train_actor.py b/slime/ray/train_actor.py index 5d4db18c6f..a8ba6ddc64 100644 --- a/slime/ray/train_actor.py +++ b/slime/ray/train_actor.py @@ -107,7 +107,7 @@ def wake_up(self, tags): raise NotImplementedError @abc.abstractmethod - def train(self, rollout_id, rollout_data_ref): + def train(self, rollout_id, rollout_data_ref, external_data=None): raise NotImplementedError @abc.abstractmethod @@ -118,10 +118,6 @@ def save_model(self, rollout_id, force_sync=False): def update_weights(self): raise NotImplementedError - @abc.abstractmethod - def connect_actor_critic(self, critic_group): - raise NotImplementedError - @abc.abstractmethod def _get_parallel_config(self): raise NotImplementedError diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index a634d1f003..787f9177d9 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1,4 +1,5 @@ import argparse +import copy import json import logging import os @@ -39,12 +40,6 @@ def add_cluster_arguments(parser): parser.add_argument( "--actor-num-gpus-per-node", type=int, default=8, help="Number of gpus per node for training actor" ) - parser.add_argument( - "--critic-num-nodes", type=int, default=None, help="Number of nodes for training actor" - ) - parser.add_argument( - "--critic-num-gpus-per-node", type=int, default=None, help="Number of gpus per node for training actor" - ) parser.add_argument( "--rollout-num-gpus", @@ -750,16 +745,21 @@ def add_algo_arguments(parser): reset_arg(parser, "--calculate-per-token-loss", action="store_true") reset_arg(parser, "--lr", type=float, default=1e-6) - parser.add_argument("--num-critic-only-steps", type=int, default=0, help="Number of critic only steps") - parser.add_argument("--critic-load", type=str, default=None, help="The checkpoint for critic model.") - parser.add_argument("--critic-save", type=str, default=None, help="The checkpoint for critic model.") - parser.add_argument("--critic-lr", type=float, default=None, help="The lr for critic model") - parser.add_argument("--critic-train-only", action="store_true", default=False, help="Only train critic") parser.add_argument( - "--critic-lr-warmup-iters", + "--num-critic-only-steps", type=int, default=0, - help="number of iterations to linearly warmup for critic model.", + help="Number of initial rollout steps that train critic only; set >= num_rollout for critic-only runs", + ) + parser.add_argument( + "--critic-config-path", + type=str, + default=None, + help=( + "Path to a structured YAML config for the critic model. The file should use " + "a top-level 'critic' key, similar to --sglang-config style, and contain " + "exactly one critic entry with Megatron/slime overrides." + ), ) parser.add_argument("--eps-clip", type=float, default=0.2, help="PPO clip range") @@ -829,6 +829,19 @@ def add_algo_arguments(parser): "This is useful for sft or custom loss function." ), ) + parser.add_argument( + "--custom-advantage-function-path", + type=str, + default=None, + help=( + "Path to a custom advantage/returns computation function. " + "When set, this function replaces the built-in compute_advantages_and_returns. " + "Signature: def custom_fn(args, rollout_data) -> None. " + "The function should set rollout_data['advantages'] and rollout_data['returns'] in-place. " + "Critic values are available in rollout_data['values']. " + "(e.g., my_module.py:my_advantage_fn)." + ), + ) parser.add_argument( "--use-kl-loss", action="store_true", default=False, help="whether to use KL loss from GRPO" ) @@ -1451,6 +1464,71 @@ def parse_args(add_custom_arguments=None): return args +def parse_critic_args(actor_args, critic_config_path): + """Parse critic-specific arguments from a YAML config, inheriting from actor_args. + + Creates an independent copy of actor_args and overrides with values from the YAML config. + This enables the critic to have a completely different parallel configuration (TP/PP/CP/EP), + micro batch size, global batch size, checkpoint path, etc. + + Args: + actor_args: The parsed actor arguments namespace to inherit from. + critic_config_path: Path to a YAML file containing critic-specific overrides. + + Returns: + A new Namespace with critic-specific configuration. + """ + critic_args = copy.deepcopy(actor_args) + + with open(critic_config_path) as f: + raw_config = yaml.safe_load(f) or {} + + if "critic" in raw_config: + critic_entries = raw_config["critic"] + assert isinstance(critic_entries, list) and len(critic_entries) == 1, ( + "critic config must contain exactly one entry under 'critic', e.g. " + "critic: [{name: default, overrides: {...}}]" + ) + critic_entry = critic_entries[0] + critic_config = critic_entry.get("overrides") or critic_entry.get("args") or {} + else: + logger.warning( + "Legacy flat critic config detected. Please wrap overrides under " + "'critic: [{name: default, overrides: {...}}]'." + ) + critic_config = raw_config + + ignored_keys = {"num_nodes", "num_gpus_per_node"} + + # Apply overrides from the YAML config. + # Unspecified keys inherit from actor_args via deepcopy. + for key, value in critic_config.items(): + if key in ignored_keys: + logger.info(f"Ignoring critic config key '{key}'; critic GPU allocation always follows actor.") + continue + if not hasattr(critic_args, key): + logger.warning(f"Critic config key '{key}' is not a known argument, setting it anyway.") + else: + # YAML safe_load doesn't parse scientific notation (e.g. 1e-5) as float. + # Coerce the value to match the type of the existing attribute. + original = getattr(critic_args, key) + if original is not None and isinstance(value, str) and isinstance(original, (int, float)): + try: + value = type(original)(value) + except (ValueError, TypeError): + pass + setattr(critic_args, key, value) + + # Critic-specific: disable features that only apply to actors + critic_args.kl_coef = 0 + critic_args.use_opd = False + critic_args.custom_advantage_function_path = None + critic_args.untie_embeddings_and_output_weights = True + logger.info(f"Parsed critic config from {critic_config_path}: overrides = {list(critic_config.keys())}") + + return critic_args + + def _resolve_eval_datasets(args) -> list[EvalDatasetConfig]: """ Build evaluation dataset configurations from either --eval-config or --eval-prompt-data. @@ -1623,22 +1701,9 @@ def slime_validate_args(args): args.debug_train_only = True args.use_critic = args.advantage_estimator == "ppo" - if args.critic_train_only: - if not args.use_critic: - raise ValueError("--critic-train-only requires --use-critic (or --advantage-estimator ppo).") - if args.actor_num_nodes != 0 or args.actor_num_gpus_per_node != 0: - raise ValueError( - "--critic-train-only requires --actor-num-nodes 0 --actor-num-gpus-per-node 0, " - f"but got actor_num_nodes={args.actor_num_nodes}, actor_num_gpus_per_node={args.actor_num_gpus_per_node}." - ) - if args.critic_num_gpus_per_node is None: - args.critic_num_gpus_per_node = args.actor_num_gpus_per_node - if args.critic_num_nodes is None: - args.critic_num_nodes = args.actor_num_nodes - if args.critic_load is None: - args.critic_load = args.load - if args.critic_lr is None: - args.critic_lr = args.lr + # Critic always uses the same GPU count as actor. + args.critic_num_gpus_per_node = args.actor_num_gpus_per_node + args.critic_num_nodes = args.actor_num_nodes if args.offload: args.offload_train = True diff --git a/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py b/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py index 533b2010bf..11c310a8ce 100644 --- a/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py +++ b/tests/test_qwen2.5_0.5B_ppo_critic_only_short.py @@ -1,4 +1,5 @@ import os +import tempfile import slime.utils.external_utils.command_utils as U @@ -16,6 +17,10 @@ def prepare(): def execute(): + critic_config = tempfile.NamedTemporaryFile("w", suffix=".yaml", delete=False) + critic_config.write("critic:\n - name: default\n overrides:\n lr: 1e-5\n") + critic_config.close() + ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " rollout_args = ( @@ -52,9 +57,8 @@ def execute(): "--kl-coef 0.00 " "--entropy-coef 0.00 " "--eps-clip 4e-4 " - "--critic-train-only " + "--num-critic-only-steps 3 " "--normalize-advantages " - "--critic-lr 1e-5 " ) optimizer_args = ( @@ -82,14 +86,13 @@ def execute(): "--accumulate-allreduce-grads-in-fp32 " "--attention-softmax-in-fp32 " "--attention-backend flash " - "--actor-num-nodes 0 " - "--actor-num-gpus-per-node 0 " - "--critic-num-nodes 1 " - "--critic-num-gpus-per-node 2 " + "--actor-num-nodes 1 " + "--actor-num-gpus-per-node 4 " "--megatron-to-hf-mode bridge " ) train_args = ( + f"--critic-config-path {critic_config.name} " f"{ckpt_args} " f"{rollout_args} " f"{optimizer_args} " diff --git a/tests/test_qwen3_4B_ppo.py b/tests/test_qwen3_4B_ppo.py index b800c19fed..0e79f5bfcb 100644 --- a/tests/test_qwen3_4B_ppo.py +++ b/tests/test_qwen3_4B_ppo.py @@ -1,4 +1,5 @@ import os +import tempfile import slime.utils.external_utils.command_utils as U @@ -21,6 +22,10 @@ def prepare(): def execute(): + critic_config = tempfile.NamedTemporaryFile("w", suffix=".yaml", delete=False) + critic_config.write("critic:\n - name: default\n overrides:\n lr: 1e-5\n") + critic_config.close() + ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/{MODEL_NAME}_torch_dist " rollout_args = ( @@ -69,7 +74,6 @@ def execute(): "--eps-clip 4e-4 " "--num-critic-only-steps 1 " "--normalize-advantages " - "--critic-lr 1e-5 " ) optimizer_args = ( @@ -102,11 +106,12 @@ def execute(): # need to comment this when using model with MLA "--attention-backend flash " "--actor-num-nodes 1 " - "--actor-num-gpus-per-node 4 " + "--actor-num-gpus-per-node 8 " "--colocate " ) train_args = ( + f"--critic-config-path {critic_config.name} " f"{ckpt_args} " f"{rollout_args} " f"{optimizer_args} " diff --git a/tests/test_qwen3_4B_ppo_train_critic_only.py b/tests/test_qwen3_4B_ppo_train_critic_only.py index 64859a5716..2d678b024b 100644 --- a/tests/test_qwen3_4B_ppo_train_critic_only.py +++ b/tests/test_qwen3_4B_ppo_train_critic_only.py @@ -1,4 +1,5 @@ import os +import tempfile import slime.utils.external_utils.command_utils as U @@ -21,6 +22,10 @@ def prepare(): def execute(): + critic_config = tempfile.NamedTemporaryFile("w", suffix=".yaml", delete=False) + critic_config.write("critic:\n - name: default\n overrides:\n lr: 1e-5\n") + critic_config.close() + ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/{MODEL_NAME}_torch_dist " rollout_args = ( @@ -67,9 +72,8 @@ def execute(): "--kl-coef 0.00 " "--entropy-coef 0.00 " "--eps-clip 4e-4 " - "--critic-train-only " + "--num-critic-only-steps 3 " "--normalize-advantages " - "--critic-lr 1e-5 " ) optimizer_args = ( @@ -83,7 +87,7 @@ def execute(): sglang_args = ( "--rollout-num-gpus-per-engine 2 " - "--rollout-num-gpus 4 " + "--rollout-num-gpus 8 " "--sglang-mem-fraction-static 0.8 " "--sglang-max-running-requests 512 " "--sglang-enable-metrics " @@ -100,13 +104,12 @@ def execute(): "--attention-softmax-in-fp32 " # need to comment this when using model with MLA "--attention-backend flash " - "--actor-num-nodes 0 " - "--actor-num-gpus-per-node 0 " - "--critic-num-nodes 1 " - "--critic-num-gpus-per-node 4 " + "--actor-num-nodes 1 " + "--actor-num-gpus-per-node 8 " ) train_args = ( + f"--critic-config-path {critic_config.name} " f"{ckpt_args} " f"{rollout_args} " f"{optimizer_args} " diff --git a/train.py b/train.py index 7bba9a3670..a470cde35f 100644 --- a/train.py +++ b/train.py @@ -26,12 +26,11 @@ def train(args): if args.offload_rollout: ray.get(rollout_manager.onload_weights.remote()) - # always update weight first so that sglang has the loaded weights from training. - if not args.critic_train_only: - actor_model.update_weights() + # Always push actor weights to rollout once weights are loaded. + actor_model.update_weights() - if args.check_weight_update_equal: - ray.get(rollout_manager.check_weights.remote(action="compare")) + if args.check_weight_update_equal: + ray.get(rollout_manager.check_weights.remote(action="compare")) if args.offload_rollout: ray.get(rollout_manager.onload_kv.remote()) @@ -40,22 +39,18 @@ def train(args): if args.num_rollout == 0 and args.eval_interval is not None: ray.get(rollout_manager.eval.remote(rollout_id=0)) - def offload_train(rollout_id): - if args.offload_train: - if args.use_critic: - critic_model.offload() - if rollout_id >= args.num_critic_only_steps and not args.critic_train_only: - actor_model.offload() + def offload_train(actor_trains_this_step): + # Each model auto-offloads after train() when offload_train is set, + # so we only need clear_memory for the non-offload case. + if not args.offload_train: + if not args.use_critic or actor_trains_this_step: + actor_model.clear_memory() else: - actor_model.offload() - else: - if args.critic_train_only: critic_model.clear_memory() - else: - actor_model.clear_memory() def save(rollout_id): - if (not args.use_critic) or (rollout_id >= args.num_critic_only_steps and not args.critic_train_only): + actor_trains_this_step = (not args.use_critic) or rollout_id >= args.num_critic_only_steps + if actor_trains_this_step: actor_model.save_model( rollout_id, force_sync=rollout_id == args.num_rollout - 1, @@ -69,7 +64,6 @@ def save(rollout_id): ray.get(rollout_manager.save.remote(rollout_id)) # train loop. - # note that for async training, one can change the position of the sync operation(ray.get). for rollout_id in range(args.start_rollout_id, args.num_rollout): if args.eval_interval is not None and rollout_id == 0 and not args.skip_eval_before_train: ray.get(rollout_manager.eval.remote(rollout_id)) @@ -79,22 +73,25 @@ def save(rollout_id): if args.offload_rollout: ray.get(rollout_manager.offload.remote()) + actor_trains_this_step = (not args.use_critic) or rollout_id >= args.num_critic_only_steps + if args.use_critic: - critic_train_handle = critic_model.async_train(rollout_id, rollout_data_ref) - if rollout_id >= args.num_critic_only_steps and not args.critic_train_only: - ray.get(actor_model.async_train(rollout_id, rollout_data_ref)) - ray.get(critic_train_handle) + value_refs = critic_model.async_train(rollout_id, rollout_data_ref) + if actor_trains_this_step: + ray.get(actor_model.async_train(rollout_id, rollout_data_ref, external_data=value_refs)) + else: + ray.get(value_refs) else: ray.get(actor_model.async_train(rollout_id, rollout_data_ref)) if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout): save(rollout_id) - offload_train(rollout_id) + offload_train(actor_trains_this_step) if args.offload_rollout: ray.get(rollout_manager.onload_weights.remote()) - if not args.critic_train_only: actor_model.update_weights() + if args.offload_rollout: ray.get(rollout_manager.onload_kv.remote()) diff --git a/train_async.py b/train_async.py index 94cc29694d..6960bd0558 100644 --- a/train_async.py +++ b/train_async.py @@ -25,12 +25,11 @@ def train(args): # create the actor and critic models actor_model, critic_model = create_training_models(args, pgs, rollout_manager) - # always update weight first so that sglang has the loaded weights from training. - if not args.critic_train_only: - actor_model.update_weights() + # Always push actor weights to rollout once weights are loaded. + actor_model.update_weights() - if args.check_weight_update_equal: - ray.get(rollout_manager.check_weights.remote(action="compare")) + if args.check_weight_update_equal: + ray.get(rollout_manager.check_weights.remote(action="compare")) # async train loop. rollout_data_next_future = rollout_manager.generate.remote(args.start_rollout_id) @@ -44,15 +43,17 @@ def train(args): rollout_data_next_future = rollout_manager.generate.remote(rollout_id + 1) if args.use_critic: - critic_train_handle = critic_model.async_train(rollout_id, rollout_data_curr_ref) - if rollout_id >= args.num_critic_only_steps and not args.critic_train_only: - ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref)) - ray.get(critic_train_handle) + actor_trains_this_step = rollout_id >= args.num_critic_only_steps + value_refs = critic_model.async_train(rollout_id, rollout_data_curr_ref) + if actor_trains_this_step: + ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref, external_data=value_refs)) + else: + ray.get(value_refs) else: ray.get(actor_model.async_train(rollout_id, rollout_data_curr_ref)) if should_run_periodic_action(rollout_id, args.save_interval, num_rollout_per_epoch, args.num_rollout): - if not args.critic_train_only: + if (not args.use_critic) or rollout_id >= args.num_critic_only_steps: actor_model.save_model( rollout_id, force_sync=rollout_id == args.num_rollout - 1, @@ -69,8 +70,7 @@ def train(args): # sync generate before update weights to prevent update weight in the middle of generation rollout_data_curr_ref = ray.get(x) if (x := rollout_data_next_future) is not None else None rollout_data_next_future = None - if not args.critic_train_only: - actor_model.update_weights() + actor_model.update_weights() if should_run_periodic_action(rollout_id, args.eval_interval, num_rollout_per_epoch): ray.get(rollout_manager.eval.remote(rollout_id))