diff --git a/examples/arctic_rl/run_gsm8k_grpo_arl_zorro_yes.sh b/examples/arctic_rl/run_gsm8k_grpo_arl_zorro_yes.sh new file mode 100755 index 00000000000..65b2aee2fce --- /dev/null +++ b/examples/arctic_rl/run_gsm8k_grpo_arl_zorro_yes.sh @@ -0,0 +1,87 @@ +#!/bin/bash + +set -x + +export PYTHONUNBUFFERED=1 +export HYDRA_FULL_ERROR=1 +export RAY_DEDUP_LOGS=0 +export HF_HUB_OFFLINE=1 +export HF_HOME=/checkpoint/huggingface +# we want to make sure this runs on non-gpu client +export CUDA_VISIBLE_DEVICES= + +BSZ=1024 +UBS=32 +ROLL_N=5 +PROMPT_LENGTH=512 +RESPONSE_LENGTH=1024 +MAX_STEPS=100 + + +BSZ=8 +UBS=2 +ROLL_N=2 +PROMPT_LENGTH=512 +RESPONSE_LENGTH=1024 +MAX_STEPS=4 + + +experiment_name="qwen3-0.6B_arctic_gsm8k_grpo" + +python3 -m verl.trainer.main_ppo \ + algorithm.adv_estimator=grpo \ + data.train_files=/code/shared/gsm8k/train.parquet \ + data.val_files=/code/shared/gsm8k/test.parquet \ + data.train_batch_size=$BSZ \ + data.max_prompt_length=512 \ + data.max_response_length=1024 \ + data.filter_overlong_prompts=True \ + data.truncation='error' \ + data.shuffle=False \ + +data.seed=42 \ + actor_rollout_ref.actor.data_loader_seed=42 \ + reward.num_workers=1 \ + actor_rollout_ref.rollout.agent.num_workers=4 \ + actor_rollout_ref.model.path=Qwen/Qwen3-0.6B \ + actor_rollout_ref.actor.optim.lr=1e-6 \ + actor_rollout_ref.model.use_remove_padding=False \ + actor_rollout_ref.actor.ppo_mini_batch_size=$BSZ \ + actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=$UBS \ + actor_rollout_ref.actor.use_kl_loss=False \ + actor_rollout_ref.actor.kl_loss_coef=0.001 \ + actor_rollout_ref.actor.kl_loss_type=low_var_kl \ + actor_rollout_ref.actor.entropy_coeff=0 \ + actor_rollout_ref.model.enable_gradient_checkpointing=True \ + +actor_rollout_ref.model.override_config.attn_implementation=flash_attention_3 \ + actor_rollout_ref.actor.strategy=fsdp2 \ + actor_rollout_ref.actor.fsdp_config.param_offload=False \ + actor_rollout_ref.actor.fsdp_config.optimizer_offload=False \ + actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=$UBS \ + actor_rollout_ref.rollout.tensor_model_parallel_size=1 \ + actor_rollout_ref.rollout.name=arctic \ + actor_rollout_ref.rollout.gpu_memory_utilization=0.6 \ + actor_rollout_ref.rollout.enforce_eager=True \ + actor_rollout_ref.rollout.n=$ROLL_N \ + actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=$UBS \ + actor_rollout_ref.ref.fsdp_config.param_offload=False \ + actor_rollout_ref.ref.strategy=fsdp2 \ + algorithm.use_kl_in_reward=False \ + trainer.use_legacy_worker_impl=disable \ + trainer.use_arctic_rl=True \ + trainer.critic_warmup=0 \ + trainer.logger=console \ + trainer.experiment_name=$experiment_name \ + trainer.project_name=arctic_gsm8k_grpo \ + trainer.val_before_train=False \ + trainer.n_gpus_per_node=1 \ + trainer.nnodes=1 \ + trainer.save_freq=-1 \ + trainer.test_freq=-1 \ + trainer.total_epochs=15 \ + trainer.total_training_steps=${MAX_STEPS} \ + arctic_rl.colocate=False \ + arctic_rl.training_gpus=2\ + arctic_rl.sampling_gpus=2\ + arctic_rl.log_prob_gpus=0\ + arctic_rl.use_zorro=True \ + "$@" 2>&1 | tee $experiment_name.log diff --git a/verl/experimental/agent_loop/agent_loop.py b/verl/experimental/agent_loop/agent_loop.py index 8879960f128..8e7118db65e 100644 --- a/verl/experimental/agent_loop/agent_loop.py +++ b/verl/experimental/agent_loop/agent_loop.py @@ -918,13 +918,14 @@ def __init__( worker_group: RayWorkerGroup = None, rollout_resource_pool: RayResourcePool = None, reward_loop_worker_handles: list[ray.actor.ActorHandle] = None, - ): + **kwargs, + ): self.config = config self.rollout_config, self.model_config = _get_rollout_and_model_config(config) self.worker_group = worker_group self.rollout_resource_pool = rollout_resource_pool self.reward_loop_worker_handles = reward_loop_worker_handles - + self.kwargs = kwargs assert worker_group is not None or self.rollout_config.nnodes > 0, "nnodes must be > 0 in standalone mode" # for recipe to change @@ -941,9 +942,10 @@ async def create( worker_group: RayWorkerGroup = None, rollout_resource_pool: RayResourcePool = None, reward_loop_worker_handles: list[ray.actor.ActorHandle] = None, + **kwargs, ): """Create agent loop manager.""" - instance = cls(config, worker_group, rollout_resource_pool, reward_loop_worker_handles) + instance = cls(config, worker_group, rollout_resource_pool, reward_loop_worker_handles, **kwargs) await instance._initialize_llm_servers() await instance._init_global_load_balancer() await instance._init_agent_loop_workers() @@ -968,6 +970,7 @@ async def _initialize_llm_servers(self): config=self.rollout_config, model_config=self.model_config, gpus_per_node=self.rollout_config.n_gpus_per_node, + **self.kwargs, ) for replica_rank in range(num_replicas) ] diff --git a/verl/single_controller/ray/base.py b/verl/single_controller/ray/base.py index 2f6ee47064f..a7872b189a1 100644 --- a/verl/single_controller/ray/base.py +++ b/verl/single_controller/ray/base.py @@ -187,8 +187,9 @@ class ResourcePoolManager: resource_pool_spec: dict[str, list[int]] mapping: dict[int, str] resource_pool_dict: dict[str, RayResourcePool] = field(default_factory=dict) + gpu_resource_pool_dict: dict[str, RayResourcePool] = field(default_factory=dict) - def create_resource_pool(self): + def create_resource_pool(self, use_gpu: bool = True): """Create Ray resource pools for distributed training. Initializes resource pools based on the resource pool specification, @@ -202,10 +203,11 @@ def create_resource_pool(self): # For Megatron backend, we recommend using max_colocate_count>1 # that can utilize different WorkerGroup for differnt models resource_pool = RayResourcePool( - process_on_nodes=process_on_nodes, use_gpu=True, max_colocate_count=3, name_prefix=resource_pool_name + process_on_nodes=process_on_nodes, use_gpu=use_gpu, max_colocate_count=3, name_prefix=resource_pool_name ) self.resource_pool_dict[resource_pool_name] = resource_pool - + if use_gpu: + self.gpu_resource_pool_dict[resource_pool_name] = resource_pool self._check_resource_available() def get_resource_pool(self, role) -> RayResourcePool: @@ -214,7 +216,8 @@ def get_resource_pool(self, role) -> RayResourcePool: def get_n_gpus(self) -> int: """Get the number of gpus in this cluster.""" - return sum([n_gpus for process_on_nodes in self.resource_pool_spec.values() for n_gpus in process_on_nodes]) + process_on_gpu_nodes = [process_on_nodes for pool_name, process_on_nodes in self.resource_pool_spec.items() if pool_name in self.gpu_resource_pool_dict] + return sum([n_gpus for process_on_nodes in process_on_gpu_nodes for n_gpus in process_on_nodes]) def _check_resource_available(self): """Check if the resource pool can be satisfied in this ray cluster.""" @@ -226,9 +229,7 @@ def _check_resource_available(self): # check total required gpus can be satisfied total_available_gpus = sum(node_available_gpus.values()) - total_required_gpus = sum( - [n_gpus for process_on_nodes in self.resource_pool_spec.values() for n_gpus in process_on_nodes] - ) + total_required_gpus = self.get_n_gpus() if total_available_gpus < total_required_gpus: raise ValueError( f"Total available GPUs {total_available_gpus} is less than total desired GPUs {total_required_gpus}" diff --git a/verl/trainer/config/ppo_trainer.yaml b/verl/trainer/config/ppo_trainer.yaml index fd9b59862ae..e2908368851 100644 --- a/verl/trainer/config/ppo_trainer.yaml +++ b/verl/trainer/config/ppo_trainer.yaml @@ -203,6 +203,9 @@ trainer: # mode: "auto", "enable", or "disable" use_legacy_worker_impl: auto + # whether to use arctic rl + use_arctic_rl: False + # profiler configs global_profiler: @@ -310,3 +313,16 @@ ray_kwargs: # Path to save Ray timeline JSON for performance profiling timeline_json_file: null + + +# config for arctic rl +arctic_rl: + + # whether to use colocate mode + colocate: False + + training_gpus: 1 + sampling_gpus: 1 + log_prob_gpus: 1 + + use_zorro: False diff --git a/verl/trainer/main_ppo.py b/verl/trainer/main_ppo.py index 2c84374d245..262a318be8b 100644 --- a/verl/trainer/main_ppo.py +++ b/verl/trainer/main_ppo.py @@ -134,6 +134,10 @@ def add_actor_rollout_worker(self, config): actor_rollout_cls = ActorRolloutRefWorker ray_worker_group_cls = RayWorkerGroup + if config.trainer.get("use_arctic_rl", False): + from verl.workers.arctic_workers import ActorRolloutRefWorker + actor_rollout_cls = ActorRolloutRefWorker + lora_rank = config.actor_rollout_ref.model.get("lora", {}).get("rank", 0) if lora_rank <= 0: lora_rank = config.actor_rollout_ref.model.get("lora_rank", 0) @@ -340,7 +344,9 @@ def run(self, config): train_sampler = create_rl_sampler(config.data, train_dataset) # Initialize the PPO trainer. - trainer = RayPPOTrainer( + from verl.trainer.ppo.arctic_trainer import ArcticPPOTrainer + ppo_trainer_cls = RayPPOTrainer if not config.trainer.use_arctic_rl else ArcticPPOTrainer + trainer = ppo_trainer_cls( config=config, tokenizer=tokenizer, processor=processor, @@ -356,7 +362,12 @@ def run(self, config): trainer.init_workers() # Start the training process. - trainer.fit() + try: + trainer.fit() + finally: + # Ensure remote services shutdown gracefully + if hasattr(trainer, "destroy"): + trainer.destroy() def create_rl_dataset(data_paths, data_config, tokenizer, processor, is_train=True, max_samples: int = -1): diff --git a/verl/trainer/ppo/arctic_rl_client.py b/verl/trainer/ppo/arctic_rl_client.py new file mode 100644 index 00000000000..1ab23798905 --- /dev/null +++ b/verl/trainer/ppo/arctic_rl_client.py @@ -0,0 +1,209 @@ +import os +from typing import Any +import torch +from transformers import AutoModelForCausalLM, AutoConfig, AutoTokenizer +from deepspeed.utils import OnDevice +import ray +from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy +from ray.util.placement_group import placement_group +from verl.workers.rollout.replica import TokenOutput + +def create_arctic_rl_client(config): + sched_pg = placement_group([{"GPU": 0, "CPU": 1}]) + return ray.remote( + num_cpus=0, + num_gpus=0, + scheduling_strategy=PlacementGroupSchedulingStrategy( + placement_group=sched_pg, + placement_group_capture_child_tasks=True, + ), + )(ArcticRLClientWrapper).remote(config) + + +class ArcticRLClientWrapper: + """Thin wrapper around ArcticTraining's ArcticRLClient + """ + + def __init__(self, config): + self.config = config + self._client = None + self.tokenizer = None + self.use_zorro = self.config.arctic_rl.use_zorro + + def is_zorro_enabled(self): + return self.use_zorro + + def _create_ds_config(self, n_gpus: int) -> dict[str, Any]: + actor_cfg = self.config.actor_rollout_ref.actor + data_cfg = self.config.data + + micro_batch_size = actor_cfg.ppo_micro_batch_size_per_gpu or 1 + train_batch_size = data_cfg.train_batch_size * self.config.actor_rollout_ref.rollout.n + grad_accum_steps = max(1, train_batch_size // (micro_batch_size * n_gpus)) + train_seq_parallel_size = actor_cfg.fsdp_config.get("ulysses_sequence_parallel_size", 1) + return { + "train_micro_batch_size_per_gpu": micro_batch_size, + "train_batch_size": train_batch_size, + "gradient_accumulation_steps": grad_accum_steps, + "sequence_parallel_size": train_seq_parallel_size, + "zero_optimization": {"stage": 1}, + } + + def _create_ds_worker_config(self): + + if self.is_zorro_enabled(): + # XXX: can't find where it's configured + use_unpad = True + + return dict( + use_zorro=True, + response_len=self.config.data.max_response_length, + max_token_len=self.config.actor_rollout_ref.rollout.max_num_batched_tokens, + rollout_n=self.config.actor_rollout_ref.rollout.n, + temperature=self.config.actor_rollout_ref.rollout.temperature, + use_unpad=use_unpad, + ) + else: + return {} + + + def initialize(self, model_name: str): + from arctic_training.arctic_rl import ArcticRLClient, ArcticRLClientConfig + + n_training_gpus = self.config.arctic_rl.get("training_gpus", self.config.trainer.n_gpus_per_node) + n_sampling_gpus = self.config.arctic_rl.get("sampling_gpus", self.config.trainer.n_gpus_per_node) + n_log_prob_gpus = self.config.arctic_rl.get("log_prob_gpus", self.config.trainer.n_gpus_per_node) + colocate = self.config.arctic_rl.get("colocate", False) + attn_implementation = self.config.actor_rollout_ref.model.override_config.get( + 'attn_implementation', 'eager' + ) + + actor_cfg = self.config.actor_rollout_ref.actor + optim_cfg = actor_cfg.optim + data_cfg = self.config.data + + max_length = data_cfg.max_prompt_length + data_cfg.max_response_length + + rollout_cfg = self.config.actor_rollout_ref.rollout + vllm_config = { + "tensor_parallel_size": rollout_cfg.tensor_model_parallel_size, + "gpu_memory_utilization": rollout_cfg.gpu_memory_utilization, + "max_model_len": rollout_cfg.get("max_model_len") or max_length, + "max_num_seqs": rollout_cfg.max_num_seqs, + "enforce_eager": rollout_cfg.enforce_eager, + "enable_chunked_prefill": rollout_cfg.enable_chunked_prefill, + } + if rollout_cfg.get("quantization"): + vllm_config["quantization"] = rollout_cfg.quantization + + rl_config = ArcticRLClientConfig( + host="localhost", + port=7000, + backend="local", + training_gpus=n_training_gpus, + sampling_gpus=n_sampling_gpus, + log_prob_gpus=n_log_prob_gpus, + colocate=colocate, + log_prob_engine="deepspeed", + model_name=model_name, + ds_config=self._create_ds_config(n_training_gpus), + log_prob_ds_config=self._create_ds_config(n_log_prob_gpus), + training_config={ + "optimizer": { + "lr": optim_cfg.lr, + "weight_decay": optim_cfg.weight_decay, + "betas": list(optim_cfg.betas), + }, + "lr_scheduler": {"warmup_ratio": optim_cfg.lr_warmup_steps_ratio}, + "training_horizon": self.config.trainer.total_epochs, + "max_length": max_length, + "model_config": None, + "attn_implementation": attn_implementation, + }, + ds_worker_config=self._create_ds_worker_config(), + vllm_config=vllm_config, + ) + + # ArcticRLClient is constructed as a ray remote actor with num_gpus=0, + # which causes CUDA_VISIBLE_DEVICES to be empty. + if colocate: + num_visible = n_training_gpus + n_sampling_gpus + n_log_prob_gpus + else: + num_visible = rl_config.training_gpus + rl_config.sampling_gpus + rl_config.log_prob_gpus + os.environ["CUDA_VISIBLE_DEVICES"] = ",".join(str(i) for i in range(num_visible)) + + self._client = ArcticRLClient(rl_config) + self.tokenizer = AutoTokenizer.from_pretrained(model_name) + + # TODO: Just for debugging - remove later + _default_sampling_params = { + "temperature": 0.0, + "top_p": 1.0, + "top_k": 0, + "max_tokens": 1024, + } + + async def generate(self, prompt_ids, sampling_params) -> list: + prompts = [self.tokenizer.decode(prompt_ids)] # TODO: pass prompt_ids directly + merged_params = {**self._default_sampling_params, **sampling_params} + return await self._client.async_generate(prompts=prompts, sampling_params=merged_params) + + + def compute_ref_log_prob(self, payload: dict): + payload["processing"] = {"post": ["compute_logprobs", "compute_entropy"], "loss_fn": None} + response = self._client.fwd_no_grad(payload, reference_model=True) + response["batch"]["log_probs"] = response["batch"].pop("logprobs") + print(f"[ArcticRLWrapper] compute_ref_log_prob OUTPUT: {response.keys()=}") + return response + + + def compute_log_prob(self, payload: dict): + payload["processing"] = {"post": ["compute_logprobs", "compute_entropy"], "loss_fn": None} + response = self._client.fwd_no_grad(payload, reference_model=False) + response["batch"]["log_probs"] = response["batch"].pop("logprobs") + print(f"[ArcticRLWrapper] compute_log_prob OUTPUT: {response.keys()=}") + return response + + def update_actor(self, payload: dict): + payload["processing"] = { + "post": ["apply_temperature", "compute_logprobs", "compute_entropy"], + # "loss_fn": "verl_grpo" + "loss_fn": "grpo" + } + def _left_pad(t: torch.Tensor, seq_len: int) -> torch.Tensor: + """Left-pad a response-only tensor to full sequence length with zeros.""" + pad_len = seq_len - t.shape[-1] + if pad_len <= 0: + return t + pad = torch.zeros(*t.shape[:-1], pad_len, dtype=t.dtype, device=t.device) + return torch.cat([pad, t], dim=-1) + + seq_len = payload["batch"]["input_ids"].shape[-1] + for name in ["old_log_probs", "advantages", "response_mask", "ref_log_prob"]: + if name in payload["batch"]: + payload["batch"][name] = _left_pad(payload["batch"][name], seq_len) + + payload["batch"]["loss_mask"] = payload["batch"]["response_mask"] + + fwd_bwd_response = self._client.fwd_bwd(payload) + print(f"[ArcticRLWrapper] update_actor OUTPUT: {fwd_bwd_response.keys()=}") + step_response = self._client.step() + print(f"[ArcticRLWrapper] update_actor STEP OUTPUT: {step_response.keys()=}") + step_response["metrics"].update(**fwd_bwd_response["metrics"]) + return step_response + + + def save_checkpoint(self): + response = self._client.save_checkpoint() + print(f"[ArcticRLClientWrapper] save_checkpoint OUTPUT: {response.keys()=}") + return response + + def update_weights(self): + # return None # TODO: Implement this + response = self._client.sync_weights() + print(f"[ArcticRLClientWrapper] update_weights OUTPUT: {response.keys()=}") + return response + + def destroy(self): + if self._client is not None: + self._client.shutdown() diff --git a/verl/trainer/ppo/arctic_trainer.py b/verl/trainer/ppo/arctic_trainer.py new file mode 100644 index 00000000000..970f30ec1bd --- /dev/null +++ b/verl/trainer/ppo/arctic_trainer.py @@ -0,0 +1,44 @@ +from typing import Optional +from torch.utils.data import Dataset, Sampler +from verl.trainer.ppo.ray_trainer import RayPPOTrainer +from verl.single_controller.ray import RayWorkerGroup, ResourcePoolManager +from verl.trainer.ppo.utils import Role, WorkerType +from verl.trainer.ppo.arctic_rl_client import create_arctic_rl_client + + +class ArcticPPOTrainer(RayPPOTrainer): + def __init__( + self, + config, + tokenizer, + role_worker_mapping: dict[Role, WorkerType], + resource_pool_manager: ResourcePoolManager, + ray_worker_group_cls: type[RayWorkerGroup] = RayWorkerGroup, + processor=None, + train_dataset: Optional[Dataset] = None, + val_dataset: Optional[Dataset] = None, + collate_fn=None, + train_sampler: Optional[Sampler] = None, + device_name=None, + ): + super().__init__(config=config, + tokenizer=tokenizer, + processor=processor, + role_worker_mapping=role_worker_mapping, + resource_pool_manager=resource_pool_manager, + ray_worker_group_cls=ray_worker_group_cls, + train_dataset=train_dataset, + val_dataset=val_dataset, + collate_fn=collate_fn, + train_sampler=train_sampler, + device_name=device_name) + + self.use_gpu = False + self.rl_client = create_arctic_rl_client(config=config) + self.rl_client.initialize.remote(model_name="Qwen/Qwen3-0.6B") + self.wg_kwargs["arctic_rl_client"] = self.rl_client + + + def destroy(self): + # self.actor_rollout_wg.destroy() + self.rl_client.destroy.remote() \ No newline at end of file diff --git a/verl/trainer/ppo/ray_trainer.py b/verl/trainer/ppo/ray_trainer.py index e178ffc143d..478bfd07908 100644 --- a/verl/trainer/ppo/ray_trainer.py +++ b/verl/trainer/ppo/ray_trainer.py @@ -309,6 +309,9 @@ def __init__( self.checkpoint_manager = None + self.wg_kwargs = {} + self.use_gpu = True + def _create_dataloader(self, train_dataset, val_dataset, collate_fn, train_sampler: Optional[Sampler]): """ Creates the train and validation dataloaders. @@ -682,7 +685,7 @@ def init_workers(self): 1. Ray resource pools from configuration 2. Worker groups for each role (actor, critic, etc.) """ - self.resource_pool_manager.create_resource_pool() + self.resource_pool_manager.create_resource_pool(use_gpu=self.use_gpu) self.resource_pool_to_cls = {pool: {} for pool in self.resource_pool_manager.resource_pool_dict.values()} @@ -694,6 +697,7 @@ def init_workers(self): cls=self.role_worker_mapping[actor_role], config=self.config.actor_rollout_ref, role=str(actor_role), + **self.wg_kwargs, ) self.resource_pool_to_cls[actor_rollout_resource_pool][str(actor_role)] = actor_rollout_cls else: @@ -840,6 +844,7 @@ def init_workers(self): worker_group=self.actor_rollout_wg, rollout_resource_pool=actor_rollout_resource_pool, reward_loop_worker_handles=reward_loop_worker_handles, + **self.wg_kwargs, ) checkpoint_engine_config = omega_conf_to_dataclass(self.config.actor_rollout_ref.rollout.checkpoint_engine) self.checkpoint_manager = CheckpointEngineManager( @@ -1589,7 +1594,8 @@ def fit(self): metrics.update(compute_timing_metrics(batch=batch, timing_raw=timing_raw)) # TODO: implement actual tflpo and theoretical tflpo n_gpus = self.resource_pool_manager.get_n_gpus() - metrics.update(compute_throughout_metrics(batch=batch, timing_raw=timing_raw, n_gpus=n_gpus)) + # To support serverless/tinker-like training, we need to support 0 GPUs training + metrics.update(compute_throughout_metrics(batch=batch, timing_raw=timing_raw, n_gpus=max(n_gpus, 1))) # compute variance proxy metrics gradient_norm = metrics.get("actor/grad_norm", None) metrics.update(compute_variance_proxy_metrics(batch=batch, gradient_norm=gradient_norm)) diff --git a/verl/workers/arctic_workers.py b/verl/workers/arctic_workers.py new file mode 100644 index 00000000000..8100518ef3d --- /dev/null +++ b/verl/workers/arctic_workers.py @@ -0,0 +1,715 @@ +from pathlib import Path +import torch +from verl.single_controller.base.decorator import Dispatch, make_nd_compute_dataproto_dispatch_fn, register +from verl.single_controller.base import Worker +from verl.utils.profiler import DistProfiler, DistProfilerExtension +from omegaconf import DictConfig +from tensordict import TensorDict +from transformers import AutoModelForCausalLM, AutoConfig, AutoTokenizer +from deepspeed.utils import OnDevice +from verl.utils import tensordict_utils as tu +from verl.utils import hf_tokenizer +import os +import ray +from verl.utils.config import omega_conf_to_dataclass +from verl.utils.device import ( + get_device_id, + get_device_name, + get_nccl_backend, + get_torch_device, + set_expandable_segments, +) +from codetiming import Timer +import os +from contextlib import nullcontext +from functools import partial +from itertools import chain + +import torch +from codetiming import Timer +from omegaconf import DictConfig, open_dict +from tensordict import NonTensorData, TensorDict +from torch.distributed.device_mesh import init_device_mesh +import torch.nn.functional as F + +try: + from verl.workers.engine.mindspeed.transformer_impl import repatch +except ImportError: + repatch = None +from verl.checkpoint_engine import CheckpointEngineRegistry +from verl.single_controller.base import Worker +from verl.single_controller.base.decorator import Dispatch, make_nd_compute_dataproto_dispatch_fn, register +from verl.utils import tensordict_utils as tu +from verl.utils.config import omega_conf_to_dataclass +from verl.utils.device import get_device_name, set_expandable_segments +from verl.utils.distributed import initialize_global_process_group_ray +from verl.utils.flops_counter import FlopsCounter +from verl.utils.memory_utils import aggressive_empty_cache +from verl.utils.metric.utils import Metric +from verl.utils.profiler import DistProfiler, DistProfilerExtension, ProfilerConfig, log_gpu_memory_usage +from verl.utils.py_functional import append_to_dict +from verl.utils.tensordict_utils import maybe_fix_3d_position_ids +from verl.utils.torch_functional import allgather_dict_into_dict +from verl.workers.config import ActorConfig, HFModelConfig, RolloutConfig, TrainingWorkerConfig +from verl.workers.rollout.base import BaseRollout, get_rollout_class +from verl.workers.utils.losses import ppo_loss +from torch import Tensor +from verl.workers.engine.utils import postprocess_batch_func + + + +def create_meta_model(name_or_path: str): + model_config = AutoConfig.from_pretrained(name_or_path) + with OnDevice(dtype=torch.float16, device='meta'): + meta_model = AutoModelForCausalLM.from_config(model_config) + return meta_model + + +def no_padding_2_padding_prompt_response(tensor: torch.Tensor, data: TensorDict, pad_token_id) -> torch.Tensor: + """Convert jagged tensor into a left padded prompt and right padded prompt of [bsz, max_response_len], which looks like + tensor([ + [pad...prompt | response...pad], + [pad...prompt | response...pad], + [pad...prompt | response...pad] + ]) + + Args: + tensor: a nested tensor or a 1D tensor in shape (total_nnz,), + total_nnz is the total number of tokens across all sequences in the batch + data: TensorDict with "prompts", "responses", "attention_mask" + pad_token_id: token to pad with + + Returns: + tensor: sliced prompt+response tensor of shape [bsz, max_response_len] w/ left and right padding + + """ + # print(f"{tensor.is_nested=}") + values = tensor.values() if tensor.is_nested else tensor + prompt_ids = data["prompts"] + response_ids = data["responses"] + attention_mask = data["attention_mask"] + # print(f"{prompt_ids.shape=}") + # print(f"{response_ids.shape=}") + # print(f"{attention_mask.shape=}") + # print(f"{attention_mask=}") + + max_prompt_len = tu.get_non_tensor_data(data=data, key="max_prompt_len", default=-1) + max_response_len = tu.get_non_tensor_data(data=data, key="max_response_len", default=-1) + # print(f"data {max_prompt_len=}") + # print(f"data {max_response_len=}") + + # print(f"{prompt_ids.is_nested=}") + if prompt_ids.is_nested: + prompt_lens = prompt_ids.offsets().diff() + response_lens = response_ids.offsets().diff() + if max_prompt_len < 0: + max_prompt_len = prompt_lens.max().item() + if max_response_len < 0: + max_response_len = response_lens.max().item() + else: + assert not attention_mask.is_nested + prompt_lens = attention_mask[:, : prompt_ids.shape[1]].sum(dim=1) + response_lens = attention_mask[:, prompt_ids.shape[1] :].sum(dim=1) + max_prompt_len = prompt_ids.shape[1] + max_response_len = response_ids.shape[1] + + sequence_lens = prompt_lens + response_lens + sequence_offsets = sequence_lens.cumsum(dim=0) + # print(f"{data=}") + # print(f"{prompt_lens=}") + # print(f"{response_lens=}") + # print(f"{max_prompt_len=}") + # print(f"{max_response_len=}") + # print(f"{sequence_offsets=}") + # print(f"{values=}") + # print(f"{values.shape=}") + assert sequence_offsets[-1].item() == values.shape[0], f"{sequence_offsets[-1].item()} != {values.shape[0]}" + + input_ids_list = [] + for prompt_len, resp_len, seq_offset in zip(prompt_lens, response_lens, sequence_offsets, strict=True): + prompt_pad_size = max_prompt_len - prompt_len + response_pad_size = max_response_len - resp_len + prompt = values[seq_offset - prompt_len - resp_len: seq_offset - resp_len] + response = values[seq_offset - resp_len: seq_offset] + prompt_padded_left = F.pad(prompt, (prompt_pad_size, 0), value=pad_token_id) + response_padded_right = F.pad(response, (0, response_pad_size), value=pad_token_id) + input_ids_list.append(torch.cat((prompt_padded_left, response_padded_right))) + + output = torch.stack(input_ids_list, dim=0) + #print(f"{output=}") + return output, max_prompt_len, max_response_len + + +def prepare_model_inputs_remove_padding(micro_batch: TensorDict): + from verl.utils import tensordict_utils as tu + from verl.utils.dataset.dataset_utils import DatasetPadMode + from verl.utils.debug import log_gpu_memory_usage + from verl.utils.device import get_device_id, get_device_name + from verl.utils.model import extract_multi_modal_inputs + from verl.utils.torch_functional import logprobs_from_logits + import verl.utils.torch_functional as verl_F + + use_remove_padding = tu.get_non_tensor_data(data=micro_batch, key="use_remove_padding", default=True) + pad_mode = tu.get_non_tensor_data(data=micro_batch, key="pad_mode", default=DatasetPadMode.NO_PADDING) + use_fused_kernels = tu.get_non_tensor_data(data=micro_batch, key="use_fused_kernels", default=False) + temperature = micro_batch["temperature"] + temperature_item = temperature + if use_fused_kernels: + assert not isinstance(temperature, torch.Tensor), ( + "use_fused_kernels does not support per sample temperature yet" + ) + assert pad_mode == DatasetPadMode.NO_PADDING, f"pad_mode {pad_mode} not supported" + + multi_modal_inputs = extract_multi_modal_inputs(micro_batch.get("multi_modal_inputs", [])) + input_ids = micro_batch["input_ids"] + position_ids = micro_batch["position_ids"] + + if not isinstance(temperature, torch.Tensor): + temperature = torch.tensor([temperature] * input_ids.shape[0], device=input_ids.device) + + temperature = temperature.to(torch.float32) + assert temperature.shape[0] == input_ids.shape[0] + + # args used to get outputs + output_args = {} + + # support per sample temperature + # temperature (bsz,) + # input_ids (bsz, j1) + temperature_rmpad = verl_F.expand_as_nested(temperature, input_ids).values() # (total_nnz,) + temperature_rmpad = temperature_rmpad.unsqueeze(0) # (1, total_nnz) + + if pad_mode == DatasetPadMode.NO_PADDING: + input_ids_rmpad = input_ids.values().unsqueeze(0) # (1, total_nnz) + if position_ids.dim() == 3: + position_ids_rmpad = position_ids.values().unsqueeze(1) # (4, 1, total_nnz) + else: + position_ids_rmpad = position_ids.values().unsqueeze(0) # (1, total_nnz) + else: + raise NotImplementedError(f"pad_mode {pad_mode} not implemented") + + # for compute the log_prob + input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=1) # (1, total_nnz) + + # pad and slice the inputs if sp > 1 + + input_ids_rmpad_rolled = input_ids_rmpad_rolled.squeeze(0) # ((total_nnz / sp) + pad) + temperature_rmpad = temperature_rmpad.squeeze(0) + output_args["input_ids_rmpad_rolled"] = input_ids_rmpad_rolled + output_args["temperature_rmpad"] = temperature_rmpad + + # only pass input_ids and position_ids to enable flash_attn_varlen + + model_inputs = { + "input_ids": input_ids_rmpad, + "attention_mask": None, + "position_ids": position_ids_rmpad, + "labels": input_ids_rmpad, + } + + extra_args = {} + if use_fused_kernels: + extra_args["temperature"] = temperature_item + extra_args["return_dict"] = True + + model_inputs.update(multi_modal_inputs) + model_inputs.update(extra_args) + + return model_inputs, output_args + + +def prepand_max_prompt_len_zeros(tensor: Tensor, max_prompt_len): + prepand = torch.zeros([tensor.shape[0], max_prompt_len], dtype=torch.int64, device=tensor.device) + return torch.cat([prepand, tensor], dim=1) + + +def make_njt(data: TensorDict, tensor: Tensor) -> Tensor: + cu_seqlens = data["input_ids"].offsets() + seq_lengths = cu_seqlens.diff() # (bsz,) + starts = torch.zeros_like(seq_lengths, dtype=torch.int64) # (bsz,) + tensor = torch.nested.narrow(tensor, 1, starts, seq_lengths, layout=torch.jagged) + tensor = torch.cat([t for t in tensor.unbind()]) + tensor = torch.nested.nested_tensor_from_jagged(tensor, cu_seqlens) + return tensor + + +def prepare_padded_dss_batch_dict(data: TensorDict, pad_token_id) -> dict: + input_ids = data['input_ids'] + position_ids = data['position_ids'] + + orig_iput_ids_shape = input_ids.shape + orig_position_ids_shape = position_ids.shape + input_ids, max_prompt_len, max_response_len = no_padding_2_padding_prompt_response(tensor=input_ids, data=data, pad_token_id=pad_token_id) + # XXX: 0 pad on pos ids is odd, check the original - perhaps need to re-build pos ids? + position_ids, _, _= no_padding_2_padding_prompt_response(tensor=position_ids, data=data, pad_token_id=0) + attention_mask = data['attention_mask'] + + # print(f"{input_ids.shape=} {position_ids.shape=} {attention_mask.shape=} {orig_iput_ids_shape=} {orig_position_ids_shape=}") + + dss_batch_dict = dict( + input_ids=input_ids, + position_ids=position_ids, + attention_mask=attention_mask, + labels=input_ids, + ) + + return dss_batch_dict, max_prompt_len, max_response_len + + +class ActorRolloutRefWorker(Worker, DistProfilerExtension): + def __init__(self, config: DictConfig, role: str, **kwargs): + Worker.__init__(self) + self.config = config + self.role = role + self._is_actor = self.role in ["actor", "actor_rollout", "actor_rollout_ref"] + self._is_rollout = self.role in ["rollout", "actor_rollout", "actor_rollout_ref"] + self._is_ref = self.role in ["ref", "actor_rollout_ref"] + + self.arctic_rl_client = kwargs.get("arctic_rl_client", None) + assert self.arctic_rl_client is not None, "arctic_rl_client is required" + + self.use_zorro = ray.get(self.arctic_rl_client.is_zorro_enabled.remote()) + + DistProfilerExtension.__init__(self, DistProfiler(rank=self.rank, config=None, tool_config=None)) + + if self._is_actor: + model_config: HFModelConfig = omega_conf_to_dataclass(self.config.model) + actor_config: ActorConfig = omega_conf_to_dataclass(self.config.actor) + actor_config.model_config = model_config + actor_training_config = TrainingWorkerConfig( + model_type="language_model", + model_config=actor_config.model_config, + engine_config=actor_config.engine, + optimizer_config=actor_config.optim, + checkpoint_config=actor_config.checkpoint, + ) + self.actor_config = actor_config + + assert self.config.actor.use_dynamic_bsz == self.config.rollout.log_prob_use_dynamic_bsz + + # assign engine configs + actor_training_config.engine_config.use_dynamic_bsz = self.config.actor.use_dynamic_bsz + actor_training_config.engine_config.infer_max_token_len_per_gpu = ( + self.config.rollout.log_prob_max_token_len_per_gpu + ) + actor_training_config.engine_config.infer_micro_batch_size_per_gpu = ( + self.config.rollout.log_prob_micro_batch_size_per_gpu + ) + actor_training_config.engine_config.max_token_len_per_gpu = self.config.actor.ppo_max_token_len_per_gpu + actor_training_config.engine_config.micro_batch_size_per_gpu = ( + self.config.actor.ppo_micro_batch_size_per_gpu + ) + actor_training_config.engine_config.use_remove_padding = model_config.use_remove_padding + + if self.config.actor.use_dynamic_bsz: + assert self.config.rollout.log_prob_max_token_len_per_gpu is not None + assert self.config.actor.ppo_max_token_len_per_gpu is not None + else: + assert self.config.rollout.log_prob_micro_batch_size_per_gpu is not None + assert self.config.actor.ppo_micro_batch_size_per_gpu is not None + + trust_remote_code=self.config.model.get("trust_remote_code", False) + self.tokenizer = hf_tokenizer(self.config.model.path, trust_remote_code=trust_remote_code) + if self.tokenizer.pad_token_id is None: + self.tokenizer.pad_token_id = self.tokenizer.eos_token_id + self.pad_token_id = self.tokenizer.pad_token_id + + self.device_name = get_device_name() + self.flops_counter = FlopsCounter(model_config.hf_config) + + + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def init_model(self): + self._register_dispatch_collect_info("actor", dp_rank=self.rank, is_collect=True) + self._register_dispatch_collect_info("ref", dp_rank=self.rank, is_collect=True) + self._register_dispatch_collect_info("rollout", dp_rank=self.rank, is_collect=True) + + return + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def destroy(self): + self.dss_training_engine.destroy() + self.arctic_inference_engine.destroy() + return + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def set_loss_fn(self, loss_fn): + return + + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def to(self, device, model=True, optimizer=True, grad=True): + """Manual control of load/offload""" + return + + + def _update_config_params(self, data: TensorDict): + default_keys = dict( + use_remove_padding=self.config.model.use_remove_padding, + use_dynamic_bsz=self.config.actor.use_dynamic_bsz, + max_token_len_per_gpu=self.config.actor.ppo_max_token_len_per_gpu, + micro_batch_size_per_gpu=self.config.actor.ppo_micro_batch_size_per_gpu, + use_fused_kernels=self.config.actor.use_fused_kernels, + ) + + for key, val in default_keys.items(): + if key not in data.keys(): + tu.assign_non_tensor(data, **{key: val}) + + + def compute_any_log_prob(self, data: TensorDict, compute_log_prob_fn) -> TensorDict: + # print(f"compute_ref_log_prob data: {data}") + batch, max_prompt_len, max_response_len = prepare_padded_dss_batch_dict(data, self.pad_token_id) + + self._update_config_params(data) + + #max_token_len_per_gpu = self.actor_config.ppo_max_token_len_per_gpu + + meta = dict( + rollout_n=self.config.rollout.n, + max_prompt_len=max_prompt_len, + max_response_len=max_response_len, + max_token_len_per_gpu=data["max_token_len_per_gpu"], + temperature=data["temperature"], + ) + + payload = dict(batch=batch, meta=meta) + + response = ray.get(compute_log_prob_fn.remote(payload)) + + # print(f"compute_any_log_prob: {response['batch']['entropy'].shape=} {response['batch']['log_probs'].shape=}") + + #batch_output = postprocess_log_prob_output(data=data, entropy=entropy, log_probs=log_probs) + #model_output = batch_output.pop("model_output", {}) + + # verl wants a full [bs, max_prompt_len+max_response_len] tensors and jagged + entropy = prepand_max_prompt_len_zeros(response['batch']['entropy'], max_prompt_len) + log_probs = prepand_max_prompt_len_zeros(response['batch']['log_probs'], max_prompt_len) + # print(f"compute_any_log_prob: {entropy.shape=} {log_probs.shape=}") + entropy = make_njt(data, entropy) + log_probs = make_njt(data, log_probs) + # print(f"compute_any_log_prob: {entropy.shape=} {log_probs.shape=}") + + model_output = dict(entropy=entropy, log_probs=log_probs) + metrics = response['metrics'] + # TODO: fix me - mfu is not computed here + metrics["mfu"] = 0.0 + + final_output = tu.get_tensordict(tensor_dict=model_output, non_tensor_dict={"metrics": metrics}) + + return final_output + + + @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="ref")) + @DistProfiler.annotate(color="olive", role="ref_compute_log_prob") + def compute_ref_log_prob(self, data: TensorDict) -> TensorDict: + return self.compute_any_log_prob(data, self.arctic_rl_client.compute_ref_log_prob) + + + @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="actor")) + @DistProfiler.annotate(color="blue", role="actor_compute_log_prob") + def compute_log_prob(self, data: TensorDict) -> TensorDict: + return self.compute_any_log_prob(data, self.arctic_rl_client.compute_log_prob) + + + def _postprocess_output(self, output, *, global_token_num, delta_time, forward_only, images_seqlens): + """ + + Args: + output: a dictionary containing loss, model_outputs and metrics + + Returns: + + """ + # TODO: whether to log memory + # metrics["perf/max_memory_allocated_gb"] = get_torch_device().max_memory_allocated() / (1024 ** 3) + # metrics["perf/max_memory_reserved_gb"] = get_torch_device().max_memory_reserved() / (1024 ** 3) + # metrics["perf/cpu_memory_used_gb"] = psutil.virtual_memory().used / (1024 ** 3) + + metrics: dict = output.pop("metrics") + # perform all gather in dp group to ensure that it's correct. + # Here each metric in metrics can be a list (micro-batch metrics) or a singleton + # we should always sum the loss of each micro-batch as we scale by global_bsz/global_token + loss = torch.sum(torch.tensor(output.pop("loss"), device=self.device_name)) + + # For grad_norm, we do not perform all reduce because it is already been done when clipping grad + grad_norm = metrics.pop("grad_norm", None) + lr = metrics.pop("lr", None) + + final_metrics = metrics + + final_metrics["loss"] = loss + if grad_norm is not None: + final_metrics["grad_norm"] = grad_norm + if lr is not None: + final_metrics["lr"] = lr + + # TODO: confirm the mtp loss IS same across dp + for k, v in final_metrics.items(): + if k.startswith("mtp_losses"): + flatten_v = [sublist[0] for sublist in v] # sublist should be single element + final_metrics[k] = sum(flatten_v) / len(flatten_v) + # compute mfu + if global_token_num is not None: + estimated_flops, promised_flops = self.flops_counter.estimate_flops( + global_token_num, delta_time, images_seqlens=images_seqlens + ) + final_metrics["mfu"] = estimated_flops / promised_flops + if forward_only: + final_metrics["mfu"] /= 3.0 + # model outputs + model_output = output.pop("model_output", {}) + # We only return final_metrics + final_output = tu.get_tensordict(tensor_dict=model_output, non_tensor_dict={"metrics": final_metrics}) + return final_output + + + @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="train"), blocking=False) + def train_actor_global_batch(self, data: TensorDict) -> TensorDict: + """Train a global batch + + Args: + data: + + Returns: + + """ + disable_auto_offload = tu.pop(data, key="disable_auto_offload", default=False) + + # update + global_token_num = data["input_ids"].offsets().diff().tolist() # (total_nnz,) + tu.assign_non_tensor( + data, + global_token_num=NonTensorData(global_token_num), + update_lr_scheduler=True, + disable_auto_offload=disable_auto_offload, + ) + + # global_token_num should be a list of number of tokens of each seq in this batch + global_token_num = tu.get(data, key="global_token_num") + disable_auto_offload = tu.get(data, key="disable_auto_offload", default=False) + images_seqlens = tu.get(data, key="images_seqlens", default=None) + + # inject engineering parameters if not specified + default_keys = dict( + use_remove_padding=self.config.model.use_remove_padding, + use_dynamic_bsz=self.config.actor.use_dynamic_bsz, + max_token_len_per_gpu=self.config.actor.ppo_max_token_len_per_gpu, + micro_batch_size_per_gpu=self.config.actor.ppo_micro_batch_size_per_gpu, + use_fused_kernels=self.config.actor.use_fused_kernels, + ) + + for key, val in default_keys.items(): + if key not in data.keys(): + tu.assign_non_tensor(data, **{key: val}) + + with ( + Timer(name="train_batch", logger=None) as timer, + ): + # XXX: what's missing is the loss function to be run on the dss side + # arctic-verl/verl/workers/engine/fsdp/transformer_impl.py:1098 forward_step + # the loss function is arctic-verl/verl/workers/utils/losses.py:97 ppo_loss + # from verl.workers.utils.losses import ppo_loss <- need to adapt to pass a gazillion of config variables + + # from verl.utils.tensordict_utils import chunk_tensordict + # batch = chunk_tensordict(data, 1) + # print(f"update_actor data: {data}") + + # XXX: fix me + input_ids = data['input_ids'] + position_ids = data['position_ids'] + #input_ids = input_ids.unbind() + + # XXX: move to init + + input_ids, max_prompt_len, max_response_len = no_padding_2_padding_prompt_response(tensor=input_ids, data=data, pad_token_id=self.pad_token_id) + # XXX: 0 pad on pos ids is odd, check the original - perhaps need to re-build pos ids? + position_ids, _, _= no_padding_2_padding_prompt_response(tensor=position_ids, data=data, pad_token_id=0) + # print(f"{input_ids.shape=}") + # print(f"{input_ids=}") + + #input_ids = torch.nested.to_padded_tensor(input_ids, padding=4.2) + #position_ids = torch.nested.to_padded_tensor(position_ids, padding=4.2) + + # print(f"{data['attention_mask'].shape=}") + # print(f"{data['attention_mask']=}") + # print(f"{input_ids.shape=}") + # print(f"{input_ids=}") + # print(f"{position_ids.shape=}") + # print(f"{position_ids=}") + # XXX: fixme + # batch = batch[0] + + #dss_batch_dict, output_args = prepare_model_inputs_remove_padding(data) + # print(f"{dss_batch_dict=}") + #print(f"{output_args=}") + #import pdb; pdb.set_trace() + + batch = dict( + input_ids=input_ids, + position_ids=position_ids, + attention_mask=data['attention_mask'], + labels=input_ids, + prompts=data["prompts"], + responses=data["responses"], + response_mask=data["response_mask"], + old_log_probs=data["old_log_probs"], + advantages=data["advantages"], + ) + if self.config.actor.use_kl_loss: + batch["ref_log_prob"] = data["ref_log_prob"] + + # print(f"{batch=}") + + # TODO: move to init since globally constant + meta = dict( + rollout_n=self.config.rollout.n, + max_prompt_len=max_prompt_len, + max_response_len=max_response_len, + max_token_len_per_gpu=data["max_token_len_per_gpu"], + temperature=data["temperature"], + use_zorro=self.use_zorro, + global_batch_size=data["global_batch_size"], + rollout_is_weights=data.get("rollout_is_weights", None), + batch_num_tokens=data["loss_mask"].sum(), + ) + + # we need to serialize the config object to dict + # dataclasses.asdict only returns keys that are defined at init (vars will do more) - but perhaps we want `asdict`? + actor_config_as_dict = vars(self.config.actor) + # print(f"update_actor: {self.actor_config=}") + # print(f"update_actor: {actor_config_as_dict=}") + import json + def safe_serialize(obj): + return json.loads(json.dumps(obj, default=lambda o: None)) + actor_config_as_dict = safe_serialize(actor_config_as_dict) + + policy_loss_config = safe_serialize(vars(self.config.actor.policy_loss)) + + meta.update(dict(actor_config=actor_config_as_dict, policy_loss_config=policy_loss_config)) + # print(f"update_actor: {post_process_inputs=}") + + + payload = dict(batch=batch, meta=meta) + response = ray.get(self.arctic_rl_client.update_actor.remote(payload)) + # output = ray.get(self.arctic_rl_client.update_actor.remote(dss_batch_dict, post_process_inputs)) + # print(f"update_actor: {loss=}") + metrics = response['metrics'] + loss = metrics.pop("loss") + # print(f"update_actor: {metrics=}") + + + from verl.utils.metric import AggregationType, Metric + # XXX: fix me - we need to aggregate the metrics + metrics = {k:Metric(value=v[0] if isinstance(v, list) else v, aggregation=AggregationType.MEAN) for k,v in metrics.items()} + metrics["lr"] = metrics.pop("last_lr") + delta_time = timer.last + + # XXX: fix me + # metrics = { + # 'actor/pg_clipfrac': None, + # 'actor/ppo_kl': None, + # 'actor/pg_clipfrac_lower': None, + # 'actor/pg_loss': None, + # 'kl_loss': None, + # 'kl_coef': None, + # 'grad_norm': None, + # } + + # print(f"{data=}") + # print(f"{data["input_ids"].shape=}") + + # expected output so far + # + # output={ + # 'model_output': { + # 'log_probs': NestedTensor(size=(1,j18), offsets=tensor([ 0,401], device='cuda:0'), grad_fn=, contiguous=True) + # }, + # 'loss': [-0.9999991059303284], + # 'metrics': { + # 'actor/pg_clipfrac': , + # 'actor/ppo_kl': , + # 'actor/pg_clipfrac_lower': , + # 'actor/pg_loss': , + # 'kl_loss': , + # 'kl_coef': [0.001], + # 'grad_norm': 16.321151733398438, + # } + # } + + model_output = {} + output = dict( + model_output=model_output, + metrics=metrics, + loss=loss, + ) + + actor_output = self._postprocess_output( + output, + global_token_num=global_token_num, + delta_time=delta_time, + forward_only=False, + images_seqlens=images_seqlens, + ).cpu() + + output_metrics = tu.get(actor_output, "metrics") + + metrics = {} + for key, val in output_metrics.items(): + # print(f"metrics {key=} {val=}") + + # flattn dp and micro batch + if isinstance(val, list): + output_metrics[key] = ( + Metric.aggregate_dp(val) + if isinstance(val[0], Metric) + else list(chain.from_iterable(val)) + ) + + append_to_dict(metrics, output_metrics) + + output = tu.get_tensordict(tensor_dict={}, non_tensor_dict={"metrics": metrics}).cpu() + + return output + + + @register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="actor")) + @DistProfiler.annotate(color="red", role="actor_update") + def update_actor(self, data: TensorDict) -> TensorDict: + # output = self.actor.train_global_batch(data=data) + output = self.train_actor_global_batch(data=data) + return output.cpu() if output is not None else None + + # TODO: Load Checkpoint API + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def load_checkpoint(self, local_path, hdfs_path=None, del_local_after_load=False): + assert "actor" in self.role, "load_checkpoint only support actor role" + return + + + # TODO: Save Checkpoint API + @register(dispatch_mode=Dispatch.ONE_TO_ALL) + def save_checkpoint(self, local_path, hdfs_path=None, global_step=0, max_ckpt_to_keep=None): + assert "actor" in self.role, "save_checkpoint only support actor role" + ray.get(self.arctic_rl_client.save_checkpoint.remote()) + return + + # TODO: Update Weights API + @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False) + async def update_weights(self, global_steps: int = None): + """Update weights from trainer to rollout. + + 1. For sync training with colocated trainer and rollout, update rollout directly from model engine. + - before update_weights: rollout should be in sleep mode. + - after update_weights: rollout should be in wake_up mode. + 2. For async training with disaggregated trainer and rollout, send_weights only by checkpoint engine. + """ + ray.get(self.arctic_rl_client.update_weights.remote()) + return + + # TODO: CheckpointManager API Begin + @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False) + async def sleep_replicas(self): + """Sleep all rollout replicas: free weight and kv_cache device memory.""" + return + # TODO: CheckpointManager API \ No newline at end of file diff --git a/verl/workers/engine/utils.py b/verl/workers/engine/utils.py index ebc5d430d81..2547c747e23 100644 --- a/verl/workers/engine/utils.py +++ b/verl/workers/engine/utils.py @@ -89,9 +89,12 @@ def prepare_micro_batches( else: total_data_size = len(data) micro_batch_size_per_gpu = data["micro_batch_size_per_gpu"] + # assert total_data_size % (force_group_size * micro_batch_size_per_gpu) == 0, ( + # "data size must be divisible by force_group_size * micro_batch_size_per_gpu" + # ) assert total_data_size % (force_group_size * micro_batch_size_per_gpu) == 0, ( - "data size must be divisible by force_group_size * micro_batch_size_per_gpu" - ) + f"data size {total_data_size} must be divisible by force_group_size {force_group_size} * micro_batch_size_per_gpu {micro_batch_size_per_gpu}" + ) micro_batches = tu.chunk_tensordict(data, total_data_size // (micro_batch_size_per_gpu * force_group_size)) batch_idx_list = None return micro_batches, batch_idx_list diff --git a/verl/workers/rollout/arctic_rollout/__init__.py b/verl/workers/rollout/arctic_rollout/__init__.py new file mode 100644 index 00000000000..cf453f8b2f6 --- /dev/null +++ b/verl/workers/rollout/arctic_rollout/__init__.py @@ -0,0 +1,3 @@ +from .arctic_rollout import ArcticReplica + +__all__ = ["ArcticReplica"] diff --git a/verl/workers/rollout/arctic_rollout/arctic_rollout.py b/verl/workers/rollout/arctic_rollout/arctic_rollout.py new file mode 100644 index 00000000000..2f187e2e95f --- /dev/null +++ b/verl/workers/rollout/arctic_rollout/arctic_rollout.py @@ -0,0 +1,321 @@ +import ray +from typing import Any, Optional +from verl.workers.rollout.vllm_rollout.vllm_async_server import vLLMHttpServer + +import argparse +from typing import Any, Optional +from verl.trainer.ppo.arctic_rl_client import ArcticRLClientWrapper +from collections.abc import AsyncGenerator + +import ray +from ray.actor import ActorHandle +from vllm import SamplingParams +from vllm.inputs import TokensPrompt +from vllm.lora.request import LoRARequest +from vllm.outputs import RequestOutput, CompletionOutput + +from verl.utils.tokenizer import normalize_token_ids +from verl.workers.config import HFModelConfig, RolloutConfig +from verl.workers.rollout.replica import RolloutMode, RolloutReplica, TokenOutput +from verl.workers.rollout.vllm_rollout.utils import ( + VLLM_LORA_INT_ID, + VLLM_LORA_NAME, + VLLM_LORA_PATH, +) +from transformers import AutoTokenizer +from verl.utils.config import omega_conf_to_dataclass +from verl.workers.rollout.utils import get_max_position_embeddings + + +class ArcticLLMEngine: + def __init__( + self, + replica_rank: int, + arctic_rl_client: ArcticRLClientWrapper, + ): + self.replica_rank = replica_rank + self.arctic_rl_client = arctic_rl_client + self.tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen3-0.6B") + + async def generate( + self, + prompt: TokensPrompt, + sampling_params: dict[str, Any], + request_id: str, + lora_request: Optional[LoRARequest] = None, + priority: int = 0, + ) -> AsyncGenerator[RequestOutput, None]: + gen_batch_output = await self.arctic_rl_client.generate.remote( + prompt_ids=prompt['prompt_token_ids'], + sampling_params=sampling_params, + ) + # print(f"arctic_async_server: {gen_batch_output=}, {type(gen_batch_output)=}") + # gen_batch_output = await ray.get(gen_batch_output) + + raw_prompt = self.tokenizer.decode(prompt['prompt_token_ids']) + completed_outputs = [] + for i, output in enumerate(gen_batch_output): + completed_outputs.append(CompletionOutput( + index=i, + text=output['text'], + token_ids=output['token_ids'], + finish_reason=output['finish_reason'], + cumulative_logprob=None, + logprobs=None + ) + ) + + yield RequestOutput( + request_id=request_id, + outputs=completed_outputs, + prompt=raw_prompt, + prompt_logprobs=None, + prompt_token_ids=prompt['prompt_token_ids'], + # finished=completed_output.finish_reason == "stop", + finished=True, + ) + + +class ArcticLLMServer(vLLMHttpServer): + """vLLM http server in single node, this is equivalent to launch server with command line: + ``` + vllm serve --tensor-parallel-size=8 ... + ``` + """ + + def __init__( + self, + config: RolloutConfig, + model_config: HFModelConfig, + rollout_mode: RolloutMode, + arctic_rl_client: ArcticRLClientWrapper, + workers: list[ActorHandle] = [], + replica_rank: int = 0, + node_rank: int = 0, + gpus_per_node: int = 1, + nnodes: int = 1, + cuda_visible_devices: str = "0", + ): + """ + Args: + config (RolloutConfig): full config. + model_config (HFModelConfig): model config. + rollout_mode (RolloutMode): rollout mode. + replica_rank (int): replica rank, a replica may contain multiple nodes. + node_rank (int): node rank. + gpus_per_node (int): number of gpus per node. + nnodes (int): number of nodes. + cuda_visible_devices (str): cuda visible devices. + """ + self.config: RolloutConfig = omega_conf_to_dataclass(config) + self.model_config: HFModelConfig = omega_conf_to_dataclass(model_config, dataclass_type=HFModelConfig) + max_position_embeddings = get_max_position_embeddings(self.model_config.hf_config) + if self.config.max_model_len is None: + self.config.max_model_len = max_position_embeddings + else: + if self.config.max_model_len > max_position_embeddings: + raise ValueError( + f"max_model_len ({self.config.max_model_len}) should be less than or equal to " + f"max_position_embeddings ({max_position_embeddings})" + ) + + self.rollout_mode = rollout_mode + self.workers = workers + + self.replica_rank = replica_rank + self.node_rank = node_rank + self.gpus_per_node = gpus_per_node + self.nnodes = nnodes + # model weights version, set by ServerAdapter when update weights. + self.global_steps = None + + if self.rollout_mode != RolloutMode.HYBRID and self.config.load_format == "dummy": + # logger.warning(f"rollout mode is {self.rollout_mode}, load_format is dummy, set to auto") + self.config.load_format = "auto" + + + self._master_address = None + self._master_port = None + self._dp_rpc_port = None + self._dp_master_port = None + + self.engine = ArcticLLMEngine(replica_rank, arctic_rl_client) + + # logger.info( + # f"vLLMHttpServer, replica_rank: {self.replica_rank}, node_rank: {self.node_rank}, " + # f"{get_visible_devices_keyword()}: {cuda_visible_devices}, " + # f"master_address: {self._master_address}, master_port: {self._master_port}, " + # f"data_parallel_rpc_port: {self._dp_rpc_port}, data_parallel_master_port: {self._dp_master_port}" + # ) + + def get_master_address(self): pass + + def get_server_address(self): pass + + @property + def lora_as_adapter(self) -> bool: pass + + async def collective_rpc( + self, + **kwargs, + ): + pass + + async def launch_server(self, master_address: str = None, master_port: int = None, dp_rpc_port: int = None): + pass + + async def run_server(self, args: argparse.Namespace): + pass + + + async def generate( + self, + prompt_ids: list[int], + sampling_params: dict[str, Any], + request_id: str, + image_data: Optional[list[Any]] = None, + video_data: Optional[list[Any]] = None, + priority: int = 0, + ) -> TokenOutput: + """Generate sequence with token-in-token-out.""" + prompt_ids = normalize_token_ids(prompt_ids) + + # Calculate the maximum possible new tokens based on available context space + # This serves as a safety upper bound + max_possible_tokens = self.config.max_model_len - len(prompt_ids) + if max_possible_tokens < 0: + raise ValueError( + f"Prompt length ({len(prompt_ids)}) exceeds the model's maximum context length " + f"({self.config.max_model_len})." + ) + + # Determine max_tokens from sampling_params or use configured response_length as default + if "max_tokens" in sampling_params: + max_tokens = sampling_params.pop("max_tokens") + elif "max_new_tokens" in sampling_params: + # support sglang-style 'max_new_tokens' param + max_tokens = sampling_params.pop("max_new_tokens") + else: + # Default to a calculation that considers configured lengths + max_tokens = self.config.response_length + self.config.prompt_length - len(prompt_ids) + + # Clamp max_tokens to the valid range [0, max_possible_tokens] + max_tokens = max(0, min(max_tokens, max_possible_tokens)) + + assert max_tokens <= max_possible_tokens, ( + f"max_tokens {max_tokens} exceeds available context space {max_possible_tokens}" + ) + sampling_params["logprobs"] = 0 if sampling_params.pop("logprobs", False) else None + sampling_params.setdefault("repetition_penalty", self.config.get("repetition_penalty", 1.0)) + # sampling_params = SamplingParams(max_tokens=max_tokens, **sampling_params) + sampling_params["max_tokens"] = max_tokens + multi_modal_data = {} + if image_data is not None: + multi_modal_data["image"] = image_data + if video_data is not None: + multi_modal_data["video"] = video_data + # import pdb; pdb.set_trace() + prompt = TokensPrompt(prompt_token_ids=prompt_ids, multi_modal_data=multi_modal_data) + + # Add lora request + lora_request = None + if self.lora_as_adapter: + # Make sure we also check that the lora is already loaded in the engine + lora_loaded = VLLM_LORA_INT_ID in await self.engine.list_loras() + if lora_loaded: + lora_request = LoRARequest( + lora_name=VLLM_LORA_NAME, lora_int_id=VLLM_LORA_INT_ID, lora_path=VLLM_LORA_PATH + ) + # import pdb; pdb.set_trace() + generator = self.engine.generate( + prompt=prompt, + sampling_params=sampling_params, + request_id=request_id, + lora_request=lora_request, + priority=priority, + ) + + # print(f"arctic_async_server: {generator=}, {type(generator)=}") + + # Get final response + final_res: Optional[RequestOutput] = None + async for output in generator: + final_res = output + assert final_res is not None + + token_ids = final_res.outputs[0].token_ids + log_probs = None + if sampling_params["logprobs"] is not None: + log_probs = [logprobs[token_ids[i]].logprob for i, logprobs in enumerate(final_res.outputs[0].logprobs)] + + routed_experts = None + if self.config.enable_rollout_routing_replay: + routed_experts = final_res.outputs[0].routed_experts + + # Determine stop reason from finish_reason + finish_reason = final_res.outputs[0].finish_reason + if finish_reason == "abort": + stop_reason = "aborted" + elif finish_reason in ("stop", "length"): + stop_reason = "completed" + else: + stop_reason = finish_reason # for more stop reason in the future + + num_preempted = None + + if hasattr(final_res.outputs[0], "num_preempted"): + num_preempted = final_res.outputs[0].num_preempted + + return TokenOutput( + token_ids=token_ids, + log_probs=log_probs, + routed_experts=routed_experts, + stop_reason=stop_reason, + num_preempted=num_preempted, + extra_info={"global_steps": self.global_steps}, + ) + + + + +class ArcticReplica(RolloutReplica): + def __init__( + self, + replica_rank: int, + config: RolloutConfig, + model_config: HFModelConfig, + gpus_per_node: int = 1, + is_reward_model: bool = False, + **kwargs, + ): + super().__init__(replica_rank, config, model_config, gpus_per_node, is_reward_model) + self.server_class = ray.remote(ArcticLLMServer) + self.arctic_rl_client = kwargs.get("arctic_rl_client", None) + assert self.arctic_rl_client is not None, "arctic_rl_client is required" + + + def rollout_worker_use_gpu(self) -> bool: + return False + + + async def launch_servers(self): + server = self.server_class.options( + ).remote( + replica_rank=self.replica_rank, + config=self.config, + model_config=self.model_config, + rollout_mode=self.rollout_mode, + arctic_rl_client=self.arctic_rl_client, + ) + self.servers.append(server) + self._server_handle = server + + + async def wake_up(self): + pass + + async def sleep(self): + pass + + async def abort_request(self, request_id: str) -> dict[str, Any]: + return {"aborted": True, "request_id": 0} diff --git a/verl/workers/rollout/replica.py b/verl/workers/rollout/replica.py index 969c6208083..2557eb74d7a 100644 --- a/verl/workers/rollout/replica.py +++ b/verl/workers/rollout/replica.py @@ -348,11 +348,16 @@ def _load_trtllm(): return TRTLLMReplica +def _load_arctic(): + from verl.workers.rollout.arctic_rollout.arctic_rollout import ArcticReplica + + return ArcticReplica # Register built-in types RolloutReplicaRegistry.register("vllm", _load_vllm) RolloutReplicaRegistry.register("sglang", _load_sglang) RolloutReplicaRegistry.register("trtllm", _load_trtllm) +RolloutReplicaRegistry.register("arctic", _load_arctic) # Original function for backward compatibility