diff --git a/.agents/contributor-skills/config-conventions/SKILL.md b/.agents/contributor-skills/config-conventions/SKILL.md index e54252513e2..481dee75aac 100644 --- a/.agents/contributor-skills/config-conventions/SKILL.md +++ b/.agents/contributor-skills/config-conventions/SKILL.md @@ -33,7 +33,7 @@ Use the right tool for the job. **v2 (the new convention):** **v1 (legacy, being migrated away):** -- **`typing.TypedDict` β€” v1, legacy / not-yet-migrated user-facing config.** Most nested sub-configs (e.g. `GRPOConfig`, `RewardScalingConfig`, `AsyncGRPOConfig`) are still `TypedDict`. Continue to maintain them with the same defaults rules below until they are migrated to `BaseModel`. Use `typing.NotRequired` to mark optional attributes. **Do not add new `TypedDict`-based config classes.** +- **`typing.TypedDict` β€” v1, legacy / not-yet-migrated user-facing config.** Some nested sub-configs are still `TypedDict`. Continue to maintain them with the same defaults rules below until they are migrated to `BaseModel`. Use `typing.NotRequired` to mark optional attributes. **Do not add new `TypedDict`-based config classes.** When in doubt: *is this class populated from a user-edited YAML?* If yes β†’ `BaseModel` (or legacy `TypedDict`). If no β†’ `@dataclass`. diff --git a/docs/guides/ppo.md b/docs/guides/ppo.md index d99396ec1da..8d83e01fca7 100644 --- a/docs/guides/ppo.md +++ b/docs/guides/ppo.md @@ -57,6 +57,18 @@ policy: When only one node remains for policy and generation after other resources are reserved, `gpus_per_node` reserves that many GPUs for generation and `num_nodes` must be `null` or `1`. When more than one node remains for training and generation, generation uses complete nodes: set `num_nodes` to the number of inference nodes and `gpus_per_node` equal to `cluster.gpus_per_node`. Non-colocated SGLang generation is not currently supported by PPO. +### Asynchronous PPO + +Set `ppo.async_ppo.enabled: true` to overlap rollout generation with training. A background collector fills a replay buffer on the non-colocated vLLM GPUs while the policy and value model train on their shared cluster. Values and policy/reference log probabilities are recomputed when a trajectory is sampled, then PPO runs GAE and its normal `ppo_epochs` updates before publishing one new policy version to vLLM. + +Async PPO reuses the trajectory collector, replay buffer, and weight-versioning infrastructure described in the [Async GRPO guide](async-grpo.md); this section focuses on PPO-specific behavior and constraints. + +Async PPO requires non-colocated vLLM generation with `vllm_cfg.async_engine: true`, `loss_fn.use_importance_sampling_correction: true`, and `loss_fn.force_on_policy_ratio: false`. Dynamic sampling, reward scaling, reward shaping, multiple dataloaders, NeMo Gym, colocated generation, and FP8 KV-scale synchronization are not supported yet. + +`max_trajectory_age_steps` is the normal policy-training age limit. The recommended value is `1`; larger values improve overlap but increase off-policy bias in GAE. When `policy_training_start_step > 0`, set `warmup_generation_lead_steps` to a larger value to bank additional rollout batches while the policy is frozen for critic warmup. The collector caps frozen-policy targets at `policy_training_start_step + max_trajectory_age_steps`, so their actual policy-update age remains within the normal limit. The buffer keeps these batches valid through that frontier and then restores the normal age limit. `null` uses `max_trajectory_age_steps` as the generation lead throughout. + +Async training stops at `max_num_steps`; the collector cycles the training dataloader as needed. `max_num_epochs` is not supported yet and must be set to `-1`; use `max_num_steps` to control training length. Async checkpoints save the collector dataloader and replay-buffer state together with policy and value state. By default, incomplete restored targets are retained and gap-filled. Setting `drop_incomplete_targets_on_restore: true` discards their restored rows and fills the target from subsequent dataloader prompts; it does not regenerate the original prompts. + ### Value Model Configuration ```yaml @@ -235,6 +247,14 @@ ppo: # null logs mismatch metrics without masking; set a threshold to mask sequences. seq_logprob_error_threshold: null + async_ppo: + enabled: false + max_trajectory_age_steps: 1 + warmup_generation_lead_steps: null + in_flight_weight_updates: false + recompute_kv_cache_after_weight_updates: false + drop_incomplete_targets_on_restore: false + adv_estimator: name: "gae" gae_lambda: 0.95 @@ -275,6 +295,7 @@ value_loss_fn: - **`ppo.ppo_epochs`**: Number of training updates per rollout batch - **`ppo.policy_training_start_step`**: Number of critic-only warmup steps before policy training begins - **`ppo.seq_logprob_error_threshold`**: Nullable sequence-level multiplicative probability-error threshold. PPO always logs sequence-level train/generation mismatch metrics; when this is set, sequences above the threshold are excluded from advantage and loss computation. +- **`ppo.async_ppo`**: Enables replay-buffer-based asynchronous PPO. See [Asynchronous PPO](#asynchronous-ppo) for requirements and staleness controls. - **`ppo.adv_estimator.name`**: Set to `"gae"` for GAE advantage estimation (PPO default) - **`ppo.adv_estimator.gae_lambda`**: GAE $\lambda$ parameter (bias-variance tradeoff, typically 0.95) - **`ppo.adv_estimator.gae_gamma`**: Discount factor $\gamma$ (typically 1.0 for outcome-supervised tasks) @@ -282,7 +303,7 @@ value_loss_fn: - **`value_loss_fn.cliprange`**: Clip range for value function predictions - **`loss_fn.positive_example_nll_weight`**: VAPO NLL auxiliary loss weight on correct samples (0 = disabled) -All other parameters (clipping, KL, importance sampling, dynamic sampling, reward shaping, reward scaling) work identically to GRPO. See the [GRPO Guide](grpo.md) for details. +For synchronous PPO, the remaining clipping, KL, sampling, and reward options work as documented in the [GRPO Guide](grpo.md). Async PPO has the limitations listed above. ## Metrics diff --git a/examples/configs/ppo_math_1B.yaml b/examples/configs/ppo_math_1B.yaml index efc54143d49..8ea57e8241a 100644 --- a/examples/configs/ppo_math_1B.yaml +++ b/examples/configs/ppo_math_1B.yaml @@ -22,6 +22,17 @@ ppo: batch_multiplier: 1 skip_reference_policy_logprobs_calculation: true # No KL, so skip ref logprobs + # Async PPO requires non-colocated vLLM generation and importance correction. + async_ppo: + enabled: false + max_trajectory_age_steps: 1 + # Requires policy_training_start_step > 0. null uses max_trajectory_age_steps. + warmup_generation_lead_steps: null + in_flight_weight_updates: false + recompute_kv_cache_after_weight_updates: false + # true discards partial restored rows and fills from subsequent prompts. + drop_incomplete_targets_on_restore: false + reward_shaping: enabled: true overlong_buffer_length: 2048 @@ -443,6 +454,7 @@ data: max_input_seq_length: 2048 shuffle: true num_workers: 1 + use_multiple_dataloader: false train: dataset_name: DAPOMath17K validation: diff --git a/examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.yaml b/examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.yaml new file mode 100644 index 00000000000..af3e868c992 --- /dev/null +++ b/examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.yaml @@ -0,0 +1,67 @@ +defaults: ../../ppo_math_1B.yaml +ppo: + num_prompts_per_step: 1024 + num_generations_per_prompt: 1 + # Async PPO cycles the dataloader and stops via max_num_steps. + max_num_epochs: -1 + ppo_epochs: 1 + policy_training_start_step: 5 + val_period: 1 + overlong_filtering: true + async_ppo: + enabled: true + warmup_generation_lead_steps: 2 + in_flight_weight_updates: true + reward_shaping: + enabled: false + adv_estimator: + gae_lambda_value: 1.0 + gae_lambda_policy: 1 + reward_scaling: + enabled: false +loss_fn: + ratio_clip_max: 0.2 + ratio_clip_c: 3 + use_importance_sampling_correction: true +value_loss_fn: + scale: 1.0 + cliprange: 0.5 +checkpointing: + checkpoint_dir: results/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async +policy: + model_name: Qwen/Qwen2.5-1.5B-Instruct + train_global_batch_size: 256 + max_total_sequence_length: 1024 + generation: + max_new_tokens: 512 + vllm_cfg: + async_engine: true + gpu_memory_utilization: 0.4 + max_model_len: 1024 + colocated: + enabled: false + resources: + gpus_per_node: 4 +value: + model_name: Qwen/Qwen2.5-1.5B-Instruct + train_micro_batch_size: 4 +data: + max_input_seq_length: 512 + train: + dataset_name: gsm8k + split: train + validation: + dataset_name: gsm8k + split: test + default: + system_prompt_file: examples/prompts/gsm8k.txt +env: + math: + math_verify_impl: hf_math_verify +logger: + log_dir: logs/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async + wandb: + project: nemo-rl + name: ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async +cluster: + gpus_per_node: 8 diff --git a/examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yaml b/examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yaml new file mode 100644 index 00000000000..6851d615b66 --- /dev/null +++ b/examples/configs/recipes/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.yaml @@ -0,0 +1,89 @@ +defaults: ../../ppo_math_1B_megatron.yaml +ppo: + num_prompts_per_step: 1024 + num_generations_per_prompt: 1 + # Async PPO cycles the dataloader and stops via max_num_steps. + max_num_epochs: -1 + ppo_epochs: 1 + val_period: 1 + overlong_filtering: true + async_ppo: + enabled: true + in_flight_weight_updates: true + reward_shaping: + enabled: false + adv_estimator: + gae_lambda_value: 1.0 + gae_lambda_policy: 1 + reward_scaling: + enabled: false +loss_fn: + ratio_clip_max: 0.2 + ratio_clip_c: 3 + use_importance_sampling_correction: true +value_loss_fn: + scale: 1.0 + cliprange: 0.5 +checkpointing: + checkpoint_dir: results/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async +policy: + model_name: Qwen/Qwen2.5-1.5B-Instruct + train_global_batch_size: 256 + max_total_sequence_length: 1024 + megatron_cfg: + tensor_model_parallel_size: 2 + context_parallel_size: 2 + sequence_parallel: true + optimizer: + weight_decay: 0.01 + scheduler: + start_weight_decay: 0.01 + end_weight_decay: 0.01 + lr_warmup_iters: 0 + lr_warmup_init: 0 + make_sequence_length_divisible_by: 8 + generation: + max_new_tokens: 512 + colocated: + enabled: false + resources: + gpus_per_node: 8 + num_nodes: 1 + vllm_cfg: + async_engine: true + gpu_memory_utilization: 0.4 + max_model_len: 1024 +value: + model_name: Qwen/Qwen2.5-1.5B-Instruct + train_micro_batch_size: 4 + megatron_cfg: + tensor_model_parallel_size: 2 + sequence_parallel: true + optimizer: + lr: 1.0e-05 + weight_decay: 0.01 + scheduler: + lr_warmup_iters: 0 + dynamic_batching: + enabled: true +data: + max_input_seq_length: 512 + train: + dataset_name: gsm8k + split: train + validation: + dataset_name: gsm8k + split: test + default: + system_prompt_file: examples/prompts/gsm8k.txt +env: + math: + math_verify_impl: hf_math_verify +logger: + log_dir: logs/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async + wandb: + project: nemo-rl + name: ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async +cluster: + num_nodes: 2 + gpus_per_node: 8 diff --git a/examples/nemo_gym/run_distillation_nemo_gym.py b/examples/nemo_gym/run_distillation_nemo_gym.py index 4472f715dbc..16d2603ef9f 100644 --- a/examples/nemo_gym/run_distillation_nemo_gym.py +++ b/examples/nemo_gym/run_distillation_nemo_gym.py @@ -28,11 +28,13 @@ distillation_train, setup, ) -from nemo_rl.algorithms.grpo import _should_use_nemo_gym from nemo_rl.algorithms.utils import get_tokenizer from nemo_rl.data.utils import setup_response_data from nemo_rl.distributed.virtual_cluster import init_ray -from nemo_rl.environments.nemo_gym import setup_nemo_gym_config +from nemo_rl.environments.nemo_gym import ( + setup_nemo_gym_config, + should_use_nemo_gym, +) from nemo_rl.models.generation import configure_generation_config from nemo_rl.utils.config import ( load_config, @@ -105,7 +107,7 @@ def main() -> None: setup_nemo_gym_config(config, tokenizer) # We assert here since this is right after the final config has been materialized. - assert _should_use_nemo_gym(config) + assert should_use_nemo_gym(config) # NeMo-Gym environment needs to get dp_openai_server_base_urls from # student_generation, so we don't setup env here. diff --git a/examples/nemo_gym/run_grpo_nemo_gym.py b/examples/nemo_gym/run_grpo_nemo_gym.py index a9c86a2be48..2c3a2544bed 100644 --- a/examples/nemo_gym/run_grpo_nemo_gym.py +++ b/examples/nemo_gym/run_grpo_nemo_gym.py @@ -33,7 +33,6 @@ MasterConfig, StatefulDataLoader, TokenizerType, - _should_use_nemo_gym, grpo_train, refit_policy_generation, setup, @@ -42,7 +41,10 @@ from nemo_rl.algorithms.utils import get_tokenizer from nemo_rl.data.utils import setup_response_data from nemo_rl.distributed.virtual_cluster import init_ray -from nemo_rl.environments.nemo_gym import setup_nemo_gym_config +from nemo_rl.environments.nemo_gym import ( + setup_nemo_gym_config, + should_use_nemo_gym, +) from nemo_rl.experience.rollouts import run_nemo_gym_rollout_sync from nemo_rl.models.generation import configure_generation_config from nemo_rl.utils.config import ( @@ -188,7 +190,7 @@ def main() -> None: setup_nemo_gym_config(config, tokenizer) # We assert here since this is right after the final config has been materialized. - assert _should_use_nemo_gym(config) + assert should_use_nemo_gym(config) # NeMo-Gym environment needs to get dp_openai_server_base_urls from policy_generation, so we don't setup env here. with rl_init_timer.time("data"): diff --git a/examples/run_ppo.py b/examples/run_ppo.py index 6d148ffd673..ae7e81ffa67 100644 --- a/examples/run_ppo.py +++ b/examples/run_ppo.py @@ -18,11 +18,12 @@ from omegaconf import OmegaConf -from nemo_rl.algorithms.ppo import MasterConfig, ppo_train, setup +from nemo_rl.algorithms.ppo import MasterConfig, async_ppo_train, ppo_train, setup from nemo_rl.algorithms.utils import get_tokenizer from nemo_rl.data.utils import setup_response_data from nemo_rl.distributed.virtual_cluster import init_ray from nemo_rl.models.generation import configure_generation_config +from nemo_rl.models.generation.interfaces import GenerationInterface from nemo_rl.utils.config import ( load_config, parse_hydra_overrides, @@ -44,6 +45,49 @@ def parse_args() -> tuple[argparse.Namespace, list[str]]: return args, overrides +def _validate_async_ppo_config( + config: MasterConfig, policy_generation: GenerationInterface | None +) -> None: + """Validate Async PPO requirements and unsupported features.""" + generation_config = config.policy["generation"] + backend = generation_config.get("backend") if generation_config else None + vllm_config = generation_config.get("vllm_cfg") if generation_config else None + async_engine = bool(vllm_config and vllm_config.get("async_engine")) + if backend != "vllm" or not async_engine: + raise ValueError( + "Async PPO requires policy.generation.backend=vllm and " + "policy.generation.vllm_cfg.async_engine=true" + ) + if not config.loss_fn.use_importance_sampling_correction: + raise ValueError( + "Async PPO requires loss_fn.use_importance_sampling_correction=true" + ) + if config.loss_fn.force_on_policy_ratio: + raise ValueError("Async PPO requires loss_fn.force_on_policy_ratio=false") + if generation_config["colocated"]["enabled"]: + raise ValueError("Async PPO requires non-colocated generation") + if config.ppo.max_num_epochs != -1: + raise NotImplementedError( + "Async PPO does not support an epoch limit; set " + "ppo.max_num_epochs=-1 and use ppo.max_num_steps to control " + "training length" + ) + + unsupported_features = { + "Dynamic sampling": config.ppo.use_dynamic_sampling, + "Reward scaling": config.ppo.reward_scaling.enabled, + "Reward shaping": config.ppo.reward_shaping.enabled, + "Multiple dataloaders": config.data["use_multiple_dataloader"], + "NeMo Gym rollout": bool(config.env.get("should_use_nemo_gym")), + "FP8 KV-scale synchronization": bool( + getattr(policy_generation, "requires_kv_scale_sync", False) + ), + } + for feature, enabled in unsupported_features.items(): + if enabled: + raise NotImplementedError(f"{feature} is not supported with async PPO") + + def main() -> None: """Main entry point.""" # Parse arguments @@ -112,28 +156,48 @@ def main() -> None: master_config, ) = setup(config, tokenizer, dataset, val_dataset) - print("πŸš€ Running synchronous PPO training") + async_config = config.ppo.async_ppo + async_ppo_enabled = async_config.enabled + if async_ppo_enabled: + _validate_async_ppo_config(config, policy_generation) - # Run standard PPO training. The checkpointer owns background - # async-checkpoint finalization threads; the context manager guarantees they - # are flushed (rename + delete) on exit. with checkpointer: - ppo_train( - policy, - policy_generation, - value_model, - dataloader, - val_dataloader, - tokenizer, - loss_fn, - value_loss_fn, - task_to_env, - val_task_to_env, - logger, - checkpointer, - ppo_state, - master_config, - ) + if async_ppo_enabled: + print("πŸš€ Running asynchronous PPO training") + async_ppo_train( + policy, + policy_generation, + value_model, + dataloader, + val_dataloader, + tokenizer, + loss_fn, + value_loss_fn, + task_to_env, + val_task_to_env, + logger, + checkpointer, + ppo_state, + master_config, + ) + else: + print("πŸš€ Running synchronous PPO training") + ppo_train( + policy, + policy_generation, + value_model, + dataloader, + val_dataloader, + tokenizer, + loss_fn, + value_loss_fn, + task_to_env, + val_task_to_env, + logger, + checkpointer, + ppo_state, + master_config, + ) if __name__ == "__main__": diff --git a/nemo_rl/algorithms/async_utils/interfaces.py b/nemo_rl/algorithms/async_utils/interfaces.py index 824f718b6e7..c7e275bbe49 100644 --- a/nemo_rl/algorithms/async_utils/interfaces.py +++ b/nemo_rl/algorithms/async_utils/interfaces.py @@ -47,7 +47,8 @@ def sample( loses its last chance to be used for its intended training step. Returns: - Dictionary with 'trajectories' and 'avg_trajectory_age' keys, or None if insufficient data + Dictionary with ``trajectories`` and ``avg_trajectory_age``, or + None if insufficient data. """ ... diff --git a/nemo_rl/algorithms/async_utils/replay_buffer.py b/nemo_rl/algorithms/async_utils/replay_buffer.py index f85d170bbc9..2b67a4ce93d 100644 --- a/nemo_rl/algorithms/async_utils/replay_buffer.py +++ b/nemo_rl/algorithms/async_utils/replay_buffer.py @@ -41,13 +41,20 @@ class ReplayBufferImpl(ReplayBufferProtocol): """Replay buffer storing per-prompt groups. A single entry corresponds to 1 prompt repeated by - grpo.num_generations_per_prompt (required to compute per-prompt advantages). + the algorithm's ``num_generations_per_prompt`` setting. """ - def __init__(self, max_size: int): + def __init__( + self, + max_size: int, + drop_incomplete_targets_on_restore: bool, + ) -> None: if max_size <= 0: raise ValueError(f"max_size must be positive, got {max_size}") self.max_size = max_size + # True discards partial restored rows. The dataloader is not rewound, + # so replacement rollouts come from subsequent prompts. + self._drop_incomplete_targets_on_restore = drop_incomplete_targets_on_restore self.trajectories = [] # List[dict[str, Any]] # If trajectory_version is 1 and target_weight_version is 4 it means that weight version 1 was used for generating a trajectory and this trajectory will be used for training when weight version is 4. self.trajectory_versions = [] # it is the weight-version used for generation of a trajectory @@ -448,6 +455,12 @@ def load_state_dict( "last_target_weight_already_generated" ] + # Filter stale rows before checking target completeness. Otherwise a + # target can look complete, lose stale rows, and remain partially + # restored even when incomplete targets should be dropped. + if max_age_steps is not None and self.trajectories: + self._remove_stale_trajectories(max_age_steps) + if current_training_step is not None and num_prompts_per_step is not None: self._prepare_for_training_step( current_step=current_training_step, @@ -456,11 +469,6 @@ def load_state_dict( elif num_prompts_per_step is not None and self.trajectories: self._remove_incomplete_target_steps(num_prompts_per_step) - if max_age_steps is not None and self.trajectories: - self._remove_stale_trajectories(max_age_steps) - if current_training_step is None and num_prompts_per_step is not None: - self._remove_incomplete_target_steps(num_prompts_per_step) - self._truncate_to_max_size(current_training_step) print( @@ -520,11 +528,33 @@ def _prepare_for_training_step( " Complete targets: " f"{sorted(complete_targets) if complete_targets else 'none'}" ) - for target in sorted(incomplete_targets): + if incomplete_targets and self._drop_incomplete_targets_on_restore: print( - f" Incomplete target {target}: " - f"{target_counts[target]}/{num_prompts_per_step}" + " Dropping incomplete restored targets; replacements will use " + "subsequent prompts: " + + ", ".join( + f"{target}={target_counts[target]}/{num_prompts_per_step}" + for target in sorted(incomplete_targets) + ) ) + indices_to_keep = [ + i + for i, target in enumerate(self.target_weight_versions) + if target not in incomplete_targets + ] + self.trajectories = [self.trajectories[i] for i in indices_to_keep] + self.trajectory_versions = [ + self.trajectory_versions[i] for i in indices_to_keep + ] + self.target_weight_versions = [ + self.target_weight_versions[i] for i in indices_to_keep + ] + else: + for target in sorted(incomplete_targets): + print( + f" Incomplete target {target}: " + f"{target_counts[target]}/{num_prompts_per_step}" + ) # Let the collector ask each target from current_step onward how many # trajectories are still needed, so incomplete restored batches can be diff --git a/nemo_rl/algorithms/async_utils/trajectory_collector.py b/nemo_rl/algorithms/async_utils/trajectory_collector.py index 652a86fe62c..b15f3828f8a 100644 --- a/nemo_rl/algorithms/async_utils/trajectory_collector.py +++ b/nemo_rl/algorithms/async_utils/trajectory_collector.py @@ -27,12 +27,27 @@ from torchdata.stateful_dataloader import StatefulDataLoader from transformers import PreTrainedTokenizerBase -from nemo_rl.algorithms.grpo import MasterConfig +from nemo_rl.algorithms.grpo import ( + AsyncGRPOConfig, + GRPOConfig, +) +from nemo_rl.algorithms.grpo import ( + MasterConfig as GRPOMasterConfig, +) from nemo_rl.algorithms.opd import resolve_reference_aliases, teacher_seq_pad_multiple +from nemo_rl.algorithms.ppo import ( + AsyncPPOConfig, + PPOConfig, +) +from nemo_rl.algorithms.ppo import ( + MasterConfig as PPOMasterConfig, +) +from nemo_rl.data.dataloader import CyclingDataLoader from nemo_rl.data.interfaces import DatumSpec from nemo_rl.data.multimodal_utils import PackedTensor from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.environments.interfaces import EnvironmentInterface +from nemo_rl.environments.nemo_gym import should_use_nemo_gym from nemo_rl.experience.interfaces import ( NEMO_GYM_TASK_INDEX_KEY, NEXT_NEMO_GYM_TASK_INDEX_KEY, @@ -66,7 +81,7 @@ def __init__( policy_generation: GenerationInterface, tokenizer: TokenizerType, task_to_env: dict[str, EnvironmentInterface], - master_config: MasterConfig, + master_config: GRPOMasterConfig | PPOMasterConfig, replay_buffer: Any, start_step: int = 0, teacher_worker_groups: Optional[dict[str, Any]] = None, @@ -74,11 +89,39 @@ def __init__( on_policy_distillation_cfg: Optional[dict[str, Any]] = None, next_nemo_gym_task_index: int = 0, processor: Any = None, - ): + ) -> None: self.policy_generation = policy_generation self.tokenizer = tokenizer self.task_to_env = task_to_env self.master_config = master_config + algorithm_config: GRPOConfig | PPOConfig + async_config: AsyncGRPOConfig | AsyncPPOConfig + if isinstance(master_config, GRPOMasterConfig): + algorithm_config = master_config.grpo + grpo_async_config = algorithm_config.async_grpo + assert grpo_async_config is not None + async_config = grpo_async_config + self._deduplicate_multimodal_data = ( + algorithm_config.deduplicate_multimodal_data + ) + self._debug_payload_metrics = algorithm_config.debug_payload_metrics + self._max_generation_failures = async_config.max_generation_failures + elif isinstance(master_config, PPOMasterConfig): + algorithm_config = master_config.ppo + async_config = algorithm_config.async_ppo + self._deduplicate_multimodal_data = False + self._debug_payload_metrics = False + self._max_generation_failures = 0 + else: + raise TypeError( + "master_config must be a GRPO or PPO MasterConfig, got " + f"{type(master_config).__name__}" + ) + self.algorithm_config = algorithm_config + self.async_config = async_config + self._num_prompts_per_step = int(algorithm_config.num_prompts_per_step) + self._num_generations_per_prompt = algorithm_config.num_generations_per_prompt + self._max_rollout_turns = algorithm_config.max_rollout_turns self.replay_buffer = replay_buffer self.teacher_worker_groups = teacher_worker_groups or {} self.alias_to_group_alias = alias_to_group_alias or {} @@ -100,6 +143,8 @@ def __init__( self.running = False self.data_exhausted = False self.collection_failed = False + self.collection_error: Optional[str] = None + self._failure_lock: _threading.Lock = _threading.Lock() self._pg_lock: _threading.Lock = _threading.Lock() @@ -112,8 +157,10 @@ def __init__( self.current_weight_version: int = start_step self.initial_weight_version: int = start_step - self.dataloader: StatefulDataLoader | None = None + self.dataloader: StatefulDataLoader | CyclingDataLoader | None = None self.collection_thread: _threading.Thread | None = None + self._generation_lead_steps = self.async_config.max_trajectory_age_steps + self._max_trajectory_age_steps = self.async_config.max_trajectory_age_steps # Track when generation limits cause collection to pause self._last_limit_warning_version: int | None = None @@ -135,14 +182,9 @@ def __init__( # Timer for efficiency metrics self._efficiency_timer = ThreadSafeTimer(context={"worker": "collector"}) - # Failure tracking for rollout batch workers. _failure_lock guards both - # _failure_count and _fatal_error_message. - self._failure_lock: _threading.Lock = _threading.Lock() + # Failure tracking for rollout batch workers. self._failure_count: int = 0 self._fatal_error_message: str | None = None - self._max_generation_failures = ( - self.master_config.grpo.async_grpo.max_generation_failures - ) def _calculate_target_weights(self, generation_weight_version: int) -> list[int]: """Calculate target weight versions for given generation weight version. @@ -154,31 +196,33 @@ def _calculate_target_weights(self, generation_weight_version: int) -> list[int] Example: generation_weight_version = 10 - max_trajectory_age_steps = 4 + generation_lead_steps = 4 + + Generation lead usually equals maximum trajectory age, but PPO critic + warmup can temporarily configure them independently. Returns: [11, 12, 13, 14] # Meaning this generation server can create trajectories for training step 11, 12, 13, 14 """ - # Read async config strictly from grpo.async_grpo - max_trajectory_age = self.master_config.grpo.async_grpo.max_trajectory_age_steps + generation_lead = self._generation_lead_steps if generation_weight_version == self.initial_weight_version: return [ i for i in range( self.initial_weight_version, - self.initial_weight_version + max_trajectory_age + 1, + self.initial_weight_version + generation_lead + 1, ) ] - return [generation_weight_version + i for i in range(1, max_trajectory_age + 1)] + return [generation_weight_version + i for i in range(1, generation_lead + 1)] def _get_next_target_for_generation( self, generation_weight_version: int ) -> Optional[int]: """Get the next target weight that needs generation (if any).""" target_weights = self._calculate_target_weights(generation_weight_version) - num_prompts = int(self.master_config.grpo.num_prompts_per_step) - max_age_steps = int(self.master_config.grpo.async_grpo.max_trajectory_age_steps) + num_prompts = self._num_prompts_per_step + max_age_steps = self._max_trajectory_age_steps last_consumed_target = ray.get( self.replay_buffer.get_last_target_weight_already_generated.remote() ) @@ -221,20 +265,45 @@ def set_weight_version(self, version: int) -> None: else: print(f"πŸ”„ Updated weight version to {version}") + def set_generation_window( + self, + *, + weight_version: int, + generation_lead_steps: int, + max_trajectory_age_steps: int, + ) -> None: + """Update the PPO generation version, lead, and buffer-validity age.""" + if generation_lead_steps < 1: + raise ValueError("generation_lead_steps must be at least 1") + if max_trajectory_age_steps < generation_lead_steps: + raise ValueError( + "max_trajectory_age_steps must be greater than or equal to " + "generation_lead_steps" + ) + + with self._generation_check_lock: + self.current_weight_version = weight_version + self._generation_lead_steps = generation_lead_steps + self._max_trajectory_age_steps = max_trajectory_age_steps + + self._generation_limit_cleared.set() + print( + f"πŸ”„ Updated generation window: version={weight_version}, " + f"lead={generation_lead_steps}, max_age={max_trajectory_age_steps}" + ) + def _should_pause_for_generation_limits(self) -> bool: """Check if collection should be paused due to generation limits.""" try: target_weights = self._calculate_target_weights(self.current_weight_version) - num_prompts = int(self.master_config.grpo.num_prompts_per_step) - max_age_steps = int( - self.master_config.grpo.async_grpo.max_trajectory_age_steps - ) + num_prompts = self._num_prompts_per_step + max_age_steps = self._max_trajectory_age_steps last_consumed_target = ray.get( self.replay_buffer.get_last_target_weight_already_generated.remote() ) - # Check if any target weight in our range needs generation with self._generation_check_lock: + # Check if any target weight in our range needs generation for target_weight in target_weights: if target_weight <= last_consumed_target: continue @@ -255,7 +324,9 @@ def _should_pause_for_generation_limits(self) -> bool: except Exception: return False - def start_collection(self, dataloader: StatefulDataLoader) -> None: + def start_collection( + self, dataloader: StatefulDataLoader | CyclingDataLoader + ) -> None: """Start collecting trajectories from dataloader.""" self.running = True self.dataloader = dataloader @@ -276,13 +347,24 @@ def get_status(self) -> dict: """Return a snapshot of the collector's internal state for driver-side diagnostics.""" with self._threads_lock: inflight_workers = len(self._inflight_threads) + with self._failure_lock: + collection_failed = self.collection_failed + collection_error = self.collection_error return { "running": self.running, "data_exhausted": self.data_exhausted, - "errored": self.collection_failed, + "errored": collection_failed, + "error": collection_error, "inflight_workers": inflight_workers, } + def _mark_collection_failed(self, error: Exception) -> None: + """Record the first collection-loop failure.""" + with self._failure_lock: + if not self.collection_failed: + self.collection_failed = True + self.collection_error = f"{type(error).__name__}: {error}" + def _collection_loop(self): """Run the collection loop in background thread.""" dataloader_exhausted = False @@ -313,14 +395,9 @@ def _collection_loop(self): # Only log warning once per weight version if self._last_limit_warning_version != self.current_weight_version: - max_trajectory_age = ( - self.master_config.grpo.async_grpo.max_trajectory_age_steps + target_weights = self._calculate_target_weights( + self.current_weight_version ) - target_weights = [ - self.current_weight_version + i - for i in range(max_trajectory_age) - ] - print( f"⏸️ Pausing collection: all target weights {target_weights} for weight version {self.current_weight_version} " f"already exist in buffer. Waiting for weight update..." @@ -342,13 +419,12 @@ def _collection_loop(self): else: # for-loop completed without break β†’ dataloader iterator exhausted dataloader_exhausted = True - except Exception as e: print(f"❌ Error in trajectory collection: {e}") import traceback traceback.print_exc() - self.collection_failed = True + self._mark_collection_failed(e) finally: self.running = False if dataloader_exhausted: @@ -384,12 +460,10 @@ def _process_batch(self, batch: BatchedDataDict[DatumSpec]) -> None: worker_started = False try: generation_weight_version = self.current_weight_version - num_generations = self.master_config.grpo.num_generations_per_prompt + num_generations = self._num_generations_per_prompt num_prompts_in_batch = batch.size - num_prompts_per_step = int(self.master_config.grpo.num_prompts_per_step) - max_age_steps = int( - self.master_config.grpo.async_grpo.max_trajectory_age_steps - ) + num_prompts_per_step = self._num_prompts_per_step + max_age_steps = self._max_trajectory_age_steps # Get the next target weight that needs generation target_weight = self._get_next_target_for_generation( @@ -428,9 +502,7 @@ def _process_batch(self, batch: BatchedDataDict[DatumSpec]) -> None: ) # Generate all prompt groups needed for this target in one batched worker. - from nemo_rl.algorithms.grpo import _should_use_nemo_gym - - use_nemo_gym = _should_use_nemo_gym(self.master_config) + use_nemo_gym = should_use_nemo_gym(self.master_config) if not self._refit_pause_cleared.is_set() and self.running: with self._threads_lock: @@ -446,21 +518,19 @@ def _process_batch(self, batch: BatchedDataDict[DatumSpec]) -> None: rollout_batch = batch.slice(0, num_prompts_to_generate) if use_nemo_gym: self._stamp_nemo_gym_task_indices(rollout_batch) - if self.master_config.grpo.deduplicate_multimodal_data: + if self._deduplicate_multimodal_data: attach_initial_nemo_gym_image_payloads( rollout_batch, self.processor ) repeated_batch = rollout_batch.repeat_interleave( num_generations, - share_immutable_media=( - self.master_config.grpo.deduplicate_multimodal_data - ), + share_immutable_media=self._deduplicate_multimodal_data, ) print_multimodal_payload_metrics( collect_multimodal_payload_metrics( repeated_batch, "prompt_repeat_async", - enabled=self.master_config.grpo.debug_payload_metrics, + enabled=self._debug_payload_metrics, ) ) @@ -572,8 +642,7 @@ def prepare_for_refit(self) -> None: is_async_engine = False else: is_async_engine = False - async_grpo_config = self.master_config.grpo.async_grpo - in_flight_weight_updates = async_grpo_config.in_flight_weight_updates + in_flight_weight_updates = self.async_config.in_flight_weight_updates if is_async_engine and in_flight_weight_updates: # async engines support in-flight weight updates @@ -602,8 +671,7 @@ def resume_after_refit(self) -> None: # Invalidate&recompute vLLM caches after the weight updates (in-flight or not) if # recompute_kv_cache_after_weight_updates is True (AREAL-style implementation). # Otherwise, keep using the stale KV caches (Magistral-style implementation). - async_cfg = self.master_config.grpo.async_grpo - if async_cfg.recompute_kv_cache_after_weight_updates: + if self.async_config.recompute_kv_cache_after_weight_updates: try: print( "πŸ”„ Invalidating generation backend KV caches after weight update" @@ -870,10 +938,8 @@ async def _iter_rollout_groups( mask_env_flagged_samples=should_mask_flagged_samples( self.master_config.env ), - deduplicate_multimodal_data=( - self.master_config.grpo.deduplicate_multimodal_data - ), - debug_payload_metrics=self.master_config.grpo.debug_payload_metrics, + deduplicate_multimodal_data=self._deduplicate_multimodal_data, + debug_payload_metrics=self._debug_payload_metrics, ): task_index = rollout_result.task_index if task_index is None: @@ -896,11 +962,9 @@ async def _iter_rollout_groups( task_to_env=self.task_to_env, max_seq_len=self.master_config.policy["max_total_sequence_length"], num_generations=num_generations, - max_rollout_turns=self.master_config.grpo.max_rollout_turns, + max_rollout_turns=self._max_rollout_turns, greedy=False, - deduplicate_multimodal_data=( - self.master_config.grpo.deduplicate_multimodal_data - ), + deduplicate_multimodal_data=self._deduplicate_multimodal_data, ): yield rollout_result @@ -1073,7 +1137,7 @@ async def _enqueue_rollout_group( target_weight_version, ), "replay_push", - enabled=self.master_config.grpo.debug_payload_metrics, + enabled=self._debug_payload_metrics, ) ) status = await self.replay_buffer.add.remote( diff --git a/nemo_rl/algorithms/distillation.py b/nemo_rl/algorithms/distillation.py index c0ce48dce57..91e00346576 100644 --- a/nemo_rl/algorithms/distillation.py +++ b/nemo_rl/algorithms/distillation.py @@ -26,12 +26,7 @@ from transformers import AutoConfig, AutoTokenizer from transformers.tokenization_utils_base import PreTrainedTokenizerBase -from nemo_rl.algorithms.grpo import ( - _should_use_async_rollouts, - _should_use_nemo_gym, - aggregate_rollout_metrics, - refit_policy_generation, -) +from nemo_rl.algorithms.grpo import aggregate_rollout_metrics, refit_policy_generation from nemo_rl.algorithms.loss import ( DistillationLossConfig, DistillationLossDataDict, @@ -59,6 +54,7 @@ NemoGymConfig, get_nemo_gym_uv_cache_dir, get_nemo_gym_venv_dir, + should_use_nemo_gym, ) from nemo_rl.experience.rollouts import ( run_async_multi_turn_rollout, @@ -67,6 +63,7 @@ ) from nemo_rl.models.generation.interfaces import ( GenerationInterface, + should_use_async_rollouts, ) from nemo_rl.models.generation.vllm import VllmConfig, VllmGeneration from nemo_rl.models.generation.vllm.config import ( @@ -337,7 +334,7 @@ def setup( # ========================== print("\nβ–Ά Setting up compute cluster...", flush=True) colocated_inference = generation_config["colocated"]["enabled"] - enable_nemo_gym = bool(env_configs) and _should_use_nemo_gym(master_config) + enable_nemo_gym = bool(env_configs) and should_use_nemo_gym(master_config) nemo_gym_actor: Optional[EnvironmentInterface] = None if enable_nemo_gym: nemo_gym_num_nodes = env_configs.get("nemo_gym", {}).get("num_gpu_nodes", 0) @@ -707,7 +704,7 @@ def distillation_train( NEED_REFIT = False POLICY_GENERATION_STALE = True # tracks if generation needs a refit before running assert student_generation is not None # for mypy type check - use_nemo_gym = _should_use_nemo_gym(master_config) + use_nemo_gym = should_use_nemo_gym(master_config) if use_nemo_gym: print("β–Ά Using NeMo-Gym rollouts for distillation", flush=True) @@ -827,7 +824,7 @@ def distillation_train( del nemo_gym_rollout_result # Use async rollouts if vLLM async engine is enabled - elif _should_use_async_rollouts(master_config): + elif should_use_async_rollouts(master_config.policy["generation"]): ( repeated_batch, rollout_metrics, @@ -1196,7 +1193,7 @@ def validate( ) return {}, {} - use_nemo_gym = _should_use_nemo_gym(master_config) + use_nemo_gym = should_use_nemo_gym(master_config) timer = Timer() with timer.time("total_validation_time"): @@ -1239,7 +1236,7 @@ def validate( for key, value in gen_metrics.items(): validation_rollout_metrics.setdefault(key, []).append(value) # Use async rollouts if vLLM async engine is enabled - elif _should_use_async_rollouts(master_config): + elif should_use_async_rollouts(master_config.policy["generation"]): val_batch, gen_metrics = run_async_multi_turn_rollout( policy_generation, val_batch, diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index d8d5f26763d..b92ffdc82ca 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -85,7 +85,7 @@ prepare_segment_topology, ) from nemo_rl.environments.interfaces import EnvironmentInterface -from nemo_rl.environments.nemo_gym import spinup_nemo_gym_actor +from nemo_rl.environments.nemo_gym import should_use_nemo_gym, spinup_nemo_gym_actor from nemo_rl.experience.interfaces import ( NEXT_NEMO_GYM_TASK_INDEX_KEY, ) @@ -105,6 +105,7 @@ GenerationInterface, GenerationSamplingParams, resolve_routed_experts_dtype_name_for_model, + should_use_async_rollouts, ) from nemo_rl.models.generation.megatron import MegatronGeneration from nemo_rl.models.generation.sglang.config import SGLangConfig @@ -508,7 +509,7 @@ def setup( or generation_config["val_top_k"] != generation_config["top_k"] ) if val_sampling_overridden: - assert generation_config["backend"] == "vllm" and _should_use_nemo_gym( + assert generation_config["backend"] == "vllm" and should_use_nemo_gym( master_config ), ( "generation.val_temperature/val_top_p/val_top_k differing from the " @@ -729,7 +730,7 @@ def init_train_dataloader(dataset, suffix: str = ""): # NeMo Gym is initialized inside setup() (rather than by the caller) so its # spinup can overlap with vLLM model loading via deferred model load. - enable_nemo_gym = _should_use_nemo_gym(master_config) + enable_nemo_gym = should_use_nemo_gym(master_config) _raise_if_reward_penalties_enabled_without_nemo_gym( master_config, enable_nemo_gym=enable_nemo_gym ) @@ -781,7 +782,7 @@ def _spinup_nemo_gym(base_urls, model_name): opd_teacher_nodes = 0 enable_opd_teachers = opd_module.is_non_colocated_teachers_enabled(master_config) if enable_opd_teachers: - assert _should_use_async_rollouts(master_config), ( + assert should_use_async_rollouts(generation_config), ( "Non-colocated OPD teachers require async GRPO (vLLM backend with async_engine enabled)." ) from nemo_rl.models.policy.teacher_worker_group import ( @@ -1331,7 +1332,7 @@ def initialize_generation_with_policy( assert policy_config["dtensor_cfg"]["enabled"] == False, ( "DTensor backend is not supported with kv cache fp8 enabled." ) - assert not _should_use_async_rollouts(master_config), ( + assert not should_use_async_rollouts(generation_config), ( "Async rollouts is not supported with kv cache fp8 enabled." ) assert policy_config["megatron_cfg"]["pipeline_model_parallel_size"] == 1, ( @@ -1953,7 +1954,7 @@ def _resolve_message_level_advantage_penalties( # The is_invalid_tool_call / has_malformed_thinking flags these penalties rely on # are only populated by the NeMo-Gym environment. Without that path the penalties # would silently no-op, so fail loudly instead. - if not _should_use_nemo_gym(master_config): + if not should_use_nemo_gym(master_config): raise ValueError( "grpo.invalid_tool_call_advantage / grpo.malformed_thinking_advantage require " "the NeMo-Gym path (env.should_use_nemo_gym=true); they are not supported with " @@ -2111,47 +2112,6 @@ def _apply_configured_message_level_advantage_penalties( ) -def _should_use_async_rollouts(master_config: MasterConfig) -> bool: - """Determine if async rollouts should be used based on the configuration. - - Dynamo is intrinsically async because all rollouts use its HTTP frontend. - SGLang only uses async rollouts when configured with ``policy.generation.use_async_rollouts``. - vLLM uses async rollouts when ``vllm_cfg.async_engine`` is enabled. - TRT-LLM always requires ``trtllm_cfg.async_engine=true``. - Megatron Inference always use async rollouts and does not need a parameter. - """ - generation_config = master_config.policy["generation"] - if generation_config is None: - return False - backend = generation_config.get("backend", "") - - if backend == "dynamo": - return True - - if backend == "sglang": - return bool(generation_config.get("use_async_rollouts", False)) - - if backend == "vllm": - return bool(generation_config.get("vllm_cfg", {}).get("async_engine", False)) - - if backend == "trtllm": - assert generation_config.get("trtllm_cfg", {}).get("async_engine", False), ( - "TRT-LLM backend requires trtllm_cfg.async_engine=true; the " - "synchronous engine path (async_engine=false) is no longer supported." - ) - return True - - if backend == "megatron": - mcore_cfg = generation_config.get("mcore_generation_config", {}) - assert mcore_cfg.get("async_engine") is None, ( - "Megatron Inference always uses the async engine. The parameter " - "policy.generation.mcore_generation_config.async_engine was removed." - ) - return True - - return False - - def _preserve_router_replay_routed_experts( target: BatchedDataDict, flat_messages: BatchedDataDict, @@ -2210,43 +2170,6 @@ def _apply_mask_sample_filter(repeated_batch: BatchedDataDict[DatumSpec]) -> int return num_masked -def _should_use_nemo_gym(master_config: MasterConfig) -> bool: - """Determine if NeMo-Gym should be used for rollouts and validation based on the configuration.""" - env_config = master_config.env - should_use_nemo_gym = bool(env_config.get("should_use_nemo_gym")) - if not should_use_nemo_gym: - return should_use_nemo_gym - - # Validate the setup for training with NeMo-Gym. - generation_config = master_config.policy["generation"] - assert _should_use_async_rollouts(master_config), ( - "❌ Error: In order to use NeMo-Gym, you must use a generation backend with `async_engine: true`!" - ) - - # We piggyback off of `_should_use_async_rollouts` to guarantee the existence of these configs. - if generation_config["backend"] == "vllm": - should_expose_http_server = generation_config["vllm_cfg"].get( - "expose_http_server" - ) - elif generation_config["backend"] == "megatron": - should_expose_http_server = generation_config["mcore_generation_config"].get( - "expose_http_server" - ) - elif generation_config["backend"] == "trtllm": - should_expose_http_server = generation_config["trtllm_cfg"].get( - "expose_http_server" - ) - elif generation_config["backend"] == "dynamo": - should_expose_http_server = generation_config["vllm_cfg"]["expose_http_server"] - else: - should_expose_http_server = False - assert should_expose_http_server, ( - "In order to use NeMo-Gym, you must expose the generation server via `expose_http_server: true`!" - ) - - return should_use_nemo_gym - - def _should_log_nemo_gym_responses(master_config: MasterConfig) -> bool: """Whether NeMo Gym is responsible for full response logging. @@ -2916,7 +2839,7 @@ def grpo_train( with timer.time("data_processing"): if ( master_config.grpo.deduplicate_multimodal_data - and _should_use_nemo_gym(master_config) + and should_use_nemo_gym(master_config) ): attach_initial_nemo_gym_image_payloads(batch, processor) # Repeat batch items @@ -3008,7 +2931,7 @@ def grpo_train( if policy_generation is not None: policy_generation.clear_logger_metrics() # Use NeMo-Gym rollouts if enabled. We cascade NeMo-Gym first since NeMo-Gym requires async rollouts. - if _should_use_nemo_gym(master_config): + if should_use_nemo_gym(master_config): # configure_generation_config auto-fills stop_token_ids from the EOS # token, but run_async_nemo_gym_rollout asserts these are unset because # NeMo-Gym manages its own stop criteria. Clear them here so the @@ -3052,7 +2975,7 @@ def grpo_train( del nemo_gym_rollout_result # Use async rollouts when enabled by config/backend defaults. - elif _should_use_async_rollouts(master_config): + elif should_use_async_rollouts(master_config.policy["generation"]): ( repeated_batch, rollout_metrics, @@ -3822,7 +3745,13 @@ def grpo_train( metrics["global_valid_toks"] / total_time / total_num_gpus ) performance_metrics = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=master_config.grpo.num_prompts_per_step, + num_generations_per_prompt=master_config.grpo.num_generations_per_prompt, + is_async_rl=master_config.grpo.async_grpo.enabled, ) if payload_metrics: @@ -3934,7 +3863,7 @@ def validate( # Generate responses (updates the LLMMessageLogType in batch_with_msg_logs) # Use async rollouts when enabled by config/backend defaults. # We cascade NeMo-Gym first since NeMo-Gym also uses async rollouts. - if _should_use_nemo_gym(master_config): + if should_use_nemo_gym(master_config): if master_config.grpo.deduplicate_multimodal_data: attach_initial_nemo_gym_image_payloads(val_batch, processor) generation_config = master_config.policy["generation"] @@ -3974,7 +3903,7 @@ def validate( val_batch = nemo_gym_rollout_result.final_batch gen_metrics = nemo_gym_rollout_result.rollout_metrics additional_metrics_to_report = gen_metrics - elif _should_use_async_rollouts(master_config): + elif should_use_async_rollouts(master_config.policy["generation"]): val_batch, gen_metrics = run_async_multi_turn_rollout( policy_generation, val_batch, @@ -4184,7 +4113,7 @@ def async_grpo_train( "Async GRPO supports the vLLM, Megatron, TRT-LLM, and Dynamo generation backends; " f"got policy.generation.backend={backend!r}." ) - assert _should_use_async_rollouts(master_config), ( + assert should_use_async_rollouts(generation_config), ( "Async GRPO requires Dynamo, Megatron, or an async vLLM or TRT-LLM " "generation engine. Set policy.generation.backend=dynamo, " "policy.generation.vllm_cfg.async_engine=true (vLLM), or " @@ -4300,7 +4229,8 @@ def async_grpo_train( ) replay_buffer = ReplayBuffer.options(runtime_env=_replay_runtime_env).remote( - max_size=optimal_buffer_size + max_size=optimal_buffer_size, + drop_incomplete_targets_on_restore=False, ) last_checkpoint_path = checkpointer.get_latest_checkpoint_path() @@ -5361,7 +5291,13 @@ def async_grpo_train( metrics["global_valid_toks"] / total_time / total_num_gpus ) performance_metrics = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=master_config.grpo.num_prompts_per_step, + num_generations_per_prompt=master_config.grpo.num_generations_per_prompt, + is_async_rl=master_config.grpo.async_grpo.enabled, ) collector_efficiency = ray.get( diff --git a/nemo_rl/algorithms/grpo_sync.py b/nemo_rl/algorithms/grpo_sync.py index 3d843e36b05..17b65751075 100644 --- a/nemo_rl/algorithms/grpo_sync.py +++ b/nemo_rl/algorithms/grpo_sync.py @@ -55,7 +55,6 @@ _policy_dtype, _resolve_logprob_skip_flags, _should_log_nemo_gym_responses, - _should_use_nemo_gym, _validation_early_stop_message, compute_and_apply_seq_logprob_error_masking, refit_policy_generation, @@ -78,6 +77,7 @@ from nemo_rl.data_plane.schema import DP_CALIB_INPUT_FIELDS, DP_TRAIN_FIELDS from nemo_rl.distributed.batched_data_dict import BatchedDataDict from nemo_rl.environments.interfaces import EnvironmentInterface +from nemo_rl.environments.nemo_gym import should_use_nemo_gym from nemo_rl.experience.sync_rollout_actor import SyncRolloutActor from nemo_rl.models.generation.interfaces import GenerationInterface from nemo_rl.models.generation.megatron import MegatronGeneration @@ -266,7 +266,7 @@ def validate_sync( total_lengths: list[float] = [] all_message_logs: list[list[dict[str, str]]] = [] additional_metrics: dict[str, Any] = {} - capture_extras = _should_use_nemo_gym(master_config) + capture_extras = should_use_nemo_gym(master_config) with timer.time("total_validation_time"): print(f"β–Ά Starting validation at step {step}...", flush=True) @@ -1348,7 +1348,13 @@ def grpo_train_sync( metrics["global_valid_toks"] / total_time / total_num_gpus ) performance_metrics = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=master_config.grpo.num_prompts_per_step, + num_generations_per_prompt=master_config.grpo.num_generations_per_prompt, + is_async_rl=False, ) logger.log_metrics(metrics, total_steps + 1, prefix="train") diff --git a/nemo_rl/algorithms/ppo.py b/nemo_rl/algorithms/ppo.py index 30374bbb313..c743b3c0e81 100644 --- a/nemo_rl/algorithms/ppo.py +++ b/nemo_rl/algorithms/ppo.py @@ -14,13 +14,14 @@ import gc import os import time +import traceback import warnings from typing import Any, NotRequired, Optional, TypedDict, TypeVar, cast import numpy as np import ray import torch -from pydantic import BaseModel +from pydantic import BaseModel, Field, model_validator from torchdata.stateful_dataloader import StatefulDataLoader from transformers import AutoProcessor from transformers.tokenization_utils_base import PreTrainedTokenizerBase @@ -31,8 +32,7 @@ ) from nemo_rl.algorithms.grpo import ( RewardScalingConfig, - _should_use_async_rollouts, - _should_use_nemo_gym, + aggregate_rollout_metrics, compute_and_apply_seq_logprob_error_masking, extract_initial_prompt_messages, refit_policy_generation, @@ -49,9 +49,14 @@ RewardShapingConfig, apply_reward_shaping, ) -from nemo_rl.algorithms.utils import print_performance_metrics, set_seed +from nemo_rl.algorithms.utils import ( + print_efficiency_summary, + print_performance_metrics, + set_seed, +) from nemo_rl.data import DataConfig from nemo_rl.data.collate_fn import rl_collate_fn +from nemo_rl.data.dataloader import CyclingDataLoader from nemo_rl.data.datasets import AllTaskProcessedDataset from nemo_rl.data.interfaces import DatumSpec from nemo_rl.data.llm_message_utils import ( @@ -68,12 +73,16 @@ prepare_segment_topology, ) from nemo_rl.environments.interfaces import EnvironmentInterface +from nemo_rl.environments.nemo_gym import should_use_nemo_gym from nemo_rl.experience.rollouts import ( run_async_multi_turn_rollout, run_multi_turn_rollout, run_nemo_gym_rollout_sync, ) -from nemo_rl.models.generation.interfaces import GenerationInterface +from nemo_rl.models.generation.interfaces import ( + GenerationInterface, + should_use_async_rollouts, +) from nemo_rl.models.generation.sglang.config import SGLangConfig from nemo_rl.models.generation.sglang.sglang_generation import SGLangGeneration from nemo_rl.models.generation.vllm import VllmConfig, VllmGeneration @@ -96,6 +105,7 @@ from nemo_rl.utils.memory_tracker import MemoryTracker from nemo_rl.utils.nsys import maybe_gpu_profile_step from nemo_rl.utils.timer import TimeoutChecker, Timer +from nemo_rl.utils.venvs import make_actor_runtime_env # =============================================================================== # Configuration @@ -103,6 +113,43 @@ TokenizerType = TypeVar("TokenizerType", bound=PreTrainedTokenizerBase) +class AsyncPPOConfig(BaseModel, extra="allow"): + """Configuration for asynchronous PPO training.""" + + # Enables the replay-buffer training loop. + enabled: bool = False + # Maximum generation-version age accepted for training. + max_trajectory_age_steps: int = Field(default=1, ge=1) + # Number of future target steps generation may fill during critic warmup. + # None uses max_trajectory_age_steps as the generation lead. + warmup_generation_lead_steps: int | None = Field(default=None, ge=1) + # Allows weight updates while rollout requests are still in flight. + in_flight_weight_updates: bool = False + # Recomputes the KV cache after weight updates. + recompute_kv_cache_after_weight_updates: bool = False + # Drops partial restored targets; replacement rollouts use subsequent prompts. + drop_incomplete_targets_on_restore: bool = False + + @model_validator(mode="after") + def validate_settings(self) -> "AsyncPPOConfig": + if ( + self.warmup_generation_lead_steps is not None + and self.warmup_generation_lead_steps < self.max_trajectory_age_steps + ): + raise ValueError( + "warmup_generation_lead_steps must be greater than or equal " + "to max_trajectory_age_steps" + ) + return self + + @property + def resolved_warmup_generation_lead_steps(self) -> int: + """Resolve the optional warmup generation lead.""" + if self.warmup_generation_lead_steps is None: + return self.max_trajectory_age_steps + return self.warmup_generation_lead_steps + + class AdvEstimatorConfig(TypedDict): """Configuration for PPO advantage estimator (GAE or raw_reward).""" @@ -118,45 +165,67 @@ class AdvEstimatorConfig(TypedDict): length_adaptive_alpha: NotRequired[float] -class PPOConfig(TypedDict): - num_prompts_per_step: int - num_generations_per_prompt: int - max_num_epochs: int - max_num_steps: int - max_rollout_turns: int - val_period: int - val_batch_size: int - val_at_start: bool +class PPOConfig(BaseModel, extra="allow"): + num_prompts_per_step: int = 32 + num_generations_per_prompt: int = 16 + max_num_epochs: int = 100000 + max_num_steps: int = 100000 + max_rollout_turns: int = 1 + val_period: int = 20 + val_batch_size: int = 256 + val_at_start: bool = True # Whether to run validation on the last training step. Setting this to True ensures the # final checkpoint has validation metrics, which is required for get_best_checkpoint_path(). - val_at_end: bool - max_val_samples: int - skip_reference_policy_logprobs_calculation: NotRequired[bool] - seed: int - overlong_filtering: bool + val_at_end: bool = False + max_val_samples: int = 256 + skip_reference_policy_logprobs_calculation: bool = True + seed: int = 42 + overlong_filtering: bool = False # whether to enable dynamic sampling, i.e. # whether to discard prompts whose rewards have zero standard deviation - use_dynamic_sampling: bool + use_dynamic_sampling: bool = False # When using dynamic sampling, the maximum number of batches to generate # before throwing an error - dynamic_sampling_max_gen_batches: NotRequired[int] + dynamic_sampling_max_gen_batches: int = 10 # When using dynamic sampling, generation prompt batch size will equal # num_prompts_per_step * batch_multiplier - batch_multiplier: NotRequired[float] - ppo_epochs: int - reward_shaping: RewardShapingConfig - reward_scaling: RewardScalingConfig - # By default advantages are calculated on CPU. Setting this flag to true leverages GPU for their computation. - calculate_advantages_on_gpu: NotRequired[bool] + batch_multiplier: float = 1.0 + ppo_epochs: int = 4 + reward_shaping: RewardShapingConfig = Field(default_factory=RewardShapingConfig) + reward_scaling: RewardScalingConfig = Field(default_factory=RewardScalingConfig) # Advantage estimator configuration (gae or raw_reward) - adv_estimator: AdvEstimatorConfig + adv_estimator: AdvEstimatorConfig = Field( + default_factory=lambda: AdvEstimatorConfig( + name="gae", + gae_lambda=0.95, + gae_gamma=1.0, + normalize_advantages=True, + gae_lambda_value=None, + gae_lambda_policy=None, + length_adaptive_alpha=0.0, + ) + ) # Number of PPO steps of critic-only warmup before policy training begins. # Value model trains from step 0; policy training is skipped for # total_steps < this value. Default 0 (train from start). - policy_training_start_step: NotRequired[int] + policy_training_start_step: int = 0 # Nullable sequence-level multiplicative probability-error threshold. # None logs metrics without masking; values above the threshold are excluded. - seq_logprob_error_threshold: float | None + seq_logprob_error_threshold: float | None = None + # Asynchronous PPO uses a replay buffer with non-colocated generation. + async_ppo: AsyncPPOConfig = Field(default_factory=AsyncPPOConfig) + + @model_validator(mode="after") + def validate_async_warmup_settings(self) -> "PPOConfig": + if ( + self.async_ppo.enabled + and self.policy_training_start_step == 0 + and self.async_ppo.warmup_generation_lead_steps is not None + ): + raise ValueError( + "warmup_generation_lead_steps requires policy_training_start_step > 0" + ) + return self class PPOSaveState(TypedDict): @@ -295,7 +364,7 @@ def setup( # Policy optimizer state first appears after critic warmup, so a cached # checkpoint layout cannot represent both the warmup and training states. assert not ( - ppo_config["policy_training_start_step"] > 0 + ppo_config.policy_training_start_step > 0 and master_config.checkpointing["enabled"] and master_config.checkpointing["save_optimizer"] and "checkpoint" in policy_megatron_config @@ -337,7 +406,7 @@ def setup( ) # Set seed for all random number generators - set_seed(ppo_config["seed"]) + set_seed(ppo_config.seed) # ========================== # Logger @@ -360,9 +429,9 @@ def setup( # Data # ========================== # Validate batch_multiplier - batch_multiplier = ppo_config["batch_multiplier"] - dataloader_batch_size = ppo_config["num_prompts_per_step"] - if not ppo_config["use_dynamic_sampling"]: + batch_multiplier = ppo_config.batch_multiplier + dataloader_batch_size = ppo_config.num_prompts_per_step + if not ppo_config.use_dynamic_sampling: assert batch_multiplier == 1, ( "batch_multiplier>1 can only be used if use_dynamic_sampling=True" ) @@ -385,17 +454,13 @@ def setup( # Load validation dataset if provided val_dataloader: Optional[StatefulDataLoader] = None # If validation is enabled, load the validation dataloader - if ( - ppo_config["val_period"] > 0 - or ppo_config["val_at_start"] - or ppo_config["val_at_end"] - ): + if ppo_config.val_period > 0 or ppo_config.val_at_start or ppo_config.val_at_end: assert val_dataset is not None, ( "Validation dataset is required if validation is enabled" ) val_dataloader = StatefulDataLoader( val_dataset, - batch_size=ppo_config["val_batch_size"], + batch_size=ppo_config.val_batch_size, shuffle=False, collate_fn=rl_collate_fn, num_workers=data_config["num_workers"], @@ -414,8 +479,7 @@ def setup( # Validate force_on_policy_ratio if loss_config.force_on_policy_ratio: assert ( - ppo_config["num_prompts_per_step"] - * ppo_config["num_generations_per_prompt"] + ppo_config.num_prompts_per_step * ppo_config.num_generations_per_prompt == policy_config["train_global_batch_size"] ), ( "force_on_policy_ratio requires train_global_batch_size == num_prompts_per_step * num_generations_per_prompt" @@ -674,25 +738,21 @@ def setup( # per outer step. So total ticks = (outer steps) * ppo_epochs. # Scale train_iters accordingly so the configured warmup/decay horizon # matches the actual scheduler-step count. - ppo_epochs = ppo_config["ppo_epochs"] - if policy_config.get("megatron_cfg", {}).get("enabled", False): - total_train_iters = ( - min( - ppo_config["max_num_steps"], - ppo_config["max_num_epochs"] * len(dataloader), - ) - * ppo_epochs + ppo_epochs = ppo_config.ppo_epochs + async_config = ppo_config.async_ppo + if async_config.enabled: + outer_training_steps = ppo_config.max_num_steps + else: + outer_training_steps = min( + ppo_config.max_num_steps, + ppo_config.max_num_epochs * len(dataloader), ) + total_train_iters = outer_training_steps * ppo_epochs + + if policy_config.get("megatron_cfg", {}).get("enabled", False): policy_config["megatron_cfg"]["train_iters"] = total_train_iters if value_config.get("megatron_cfg", {}).get("enabled", False): - total_train_iters = ( - min( - ppo_config["max_num_steps"], - ppo_config["max_num_epochs"] * len(dataloader), - ) - * ppo_epochs - ) value_config["megatron_cfg"]["train_iters"] = total_train_iters # Define initialization functions that will be used in all paths @@ -806,7 +866,7 @@ def initialize_generation_with_policy( assert policy_config["dtensor_cfg"]["enabled"] == False, ( "DTensor backend is not supported with kv cache fp8 enabled." ) - assert not _should_use_async_rollouts(master_config), ( + assert not should_use_async_rollouts(generation_config), ( "Async rollouts is not supported with kv cache fp8 enabled." ) assert policy_config["megatron_cfg"]["pipeline_model_parallel_size"] == 1, ( @@ -973,8 +1033,8 @@ def dynamic_sampling( # Required batch size for training train_prompts_size = ( - master_config.ppo["num_prompts_per_step"] - * master_config.ppo["num_generations_per_prompt"] + master_config.ppo.num_prompts_per_step + * master_config.ppo.num_generations_per_prompt ) # Store the baseline, std and total_reward for the current unfiltered batch. repeated_batch["baseline"] = baseline @@ -985,7 +1045,7 @@ def dynamic_sampling( # Dynamic sampling algorithm (used in DAPO algorithm) # This block implements dynamic sampling by selecting prompt groups with non-zero std. # If sampled prompts (with non-zero std) are fewer than num_prompts_per_step * num_generations_per_prompt, continue sampling until dynamic_sampling_max_gen_batches is reached. - if master_config.ppo["use_dynamic_sampling"]: + if master_config.ppo.use_dynamic_sampling: with timer.time("dynamic_sampling"): # Get the prompt indices with non-zero std non_zero_std_mask = std != 0.0 @@ -1028,9 +1088,9 @@ def dynamic_sampling( # If the generation samples size is smaller than a fixed threshold (train_prompts_size), keep generating by processing the next batch if filtered_prompts_size < train_prompts_size: - dynamic_sampling_max_gen_batches = master_config.ppo[ - "dynamic_sampling_max_gen_batches" - ] + dynamic_sampling_max_gen_batches = ( + master_config.ppo.dynamic_sampling_max_gen_batches + ) assert dynamic_sampling_max_gen_batches > 0, ( "When using ppo.use_dynamic_sampling, ppo.dynamic_sampling_max_gen_batches must be > 0" ) @@ -1056,7 +1116,7 @@ def dynamic_sampling( batch_to_return = ( filtered_repeated_batch - if master_config.ppo["use_dynamic_sampling"] + if master_config.ppo.use_dynamic_sampling else repeated_batch ) return batch_to_return, is_batch_complete, batch_cache, dynamic_sampling_metrics @@ -1082,7 +1142,7 @@ def _create_advantage_estimator(master_config: MasterConfig): ppo_config = master_config.ppo loss_config = master_config.loss_fn - adv_estimator_config = ppo_config["adv_estimator"] + adv_estimator_config = ppo_config.adv_estimator adv_estimator_name = adv_estimator_config["name"] if adv_estimator_name == "gae": @@ -1102,6 +1162,35 @@ def _create_advantage_estimator(master_config: MasterConfig): return adv_estimator +def _compute_critic_metrics(value_results: dict[str, Any]) -> dict[str, Any]: + """Aggregate value-model metrics under the ``critic/`` namespace.""" + value_mb_metrics = value_results.get("all_mb_metrics", {}) + critic_metrics: dict[str, Any] = { + "critic/grad_norm": value_results["grad_norm"].numpy(), + "critic/loss": value_results["loss"].numpy(), + } + for key, value in value_mb_metrics.items(): + metric_name = f"critic/{key}" + if key in {"lr", "wd", "global_valid_seqs", "global_valid_toks", "grad_norm"}: + critic_metrics[metric_name] = np.mean(value).item() + elif key == "values_min": + critic_metrics[metric_name] = np.min(value).item() + elif key == "values_max": + critic_metrics[metric_name] = np.max(value).item() + elif isinstance(value, (np.ndarray, list)): + critic_metrics[metric_name] = np.sum(value).item() + else: + raise ValueError(f"Unsupported value-model metric: {key}") + returns_mean = critic_metrics.get("critic/returns_mean", 0) + values_mean = critic_metrics.get("critic/values_mean", 0) + returns_sq_mean = critic_metrics.get("critic/returns_sq_mean", 0) + residual_sq_mean = critic_metrics.get("critic/residual_sq_mean", 0) + returns_var = returns_sq_mean - returns_mean**2 + residual_var = residual_sq_mean - (returns_mean - values_mean) ** 2 + critic_metrics["critic/explained_var"] = 1.0 - residual_var / max(returns_var, 1e-8) + return critic_metrics + + # =============================================================================== # Training & Validation # =============================================================================== @@ -1149,7 +1238,7 @@ def ppo_train( POLICY_GENERATION_STALE = True # tracks if generation needs a refit before running assert policy_generation is not None # for mypy type check - if master_config.ppo.get("skip_reference_policy_logprobs_calculation"): + if master_config.ppo.skip_reference_policy_logprobs_calculation: assert master_config.loss_fn.reference_policy_kl_penalty == 0 print( "Reference policy logprob calculation will be skipped since `ppo.skip_reference_policy_logprobs_calculation` is set to True and `loss_fn.reference_policy_kl_penalty` is 0." @@ -1161,19 +1250,19 @@ def ppo_train( # common config/state current_step = ppo_save_state["current_step"] total_steps = ppo_save_state["total_steps"] - max_num_steps = master_config.ppo["max_num_steps"] + max_num_steps = master_config.ppo.max_num_steps current_epoch = ppo_save_state["current_epoch"] - max_num_epochs = master_config.ppo["max_num_epochs"] - ppo_epochs = master_config.ppo["ppo_epochs"] + max_num_epochs = master_config.ppo.max_num_epochs + ppo_epochs = master_config.ppo.ppo_epochs # Number of PPO steps to train only the critic before starting policy # training. Despite the legacy name, this is compared against total_steps # (not current_epoch) to match veRL's critic_warmup semantics. - policy_training_start_step = master_config.ppo["policy_training_start_step"] + policy_training_start_step = master_config.ppo.policy_training_start_step consumed_samples = ppo_save_state["consumed_samples"] total_valid_tokens = ppo_save_state.get("total_valid_tokens", 0) - val_at_start = master_config.ppo["val_at_start"] - val_at_end = master_config.ppo["val_at_end"] - val_period = master_config.ppo["val_period"] + val_at_start = master_config.ppo.val_at_start + val_at_end = master_config.ppo.val_at_end + val_period = master_config.ppo.val_period colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] # Initialize advantage estimator @@ -1232,7 +1321,7 @@ def ppo_train( with timer.time("data_processing"): repeated_batch: BatchedDataDict[DatumSpec] = ( batch.repeat_interleave( - master_config.ppo["num_generations_per_prompt"] + master_config.ppo.num_generations_per_prompt ) ) batched_flat, input_lengths = batched_message_log_to_flat_message( @@ -1300,7 +1389,7 @@ def ppo_train( if policy_generation is not None: policy_generation.clear_logger_metrics() - if _should_use_nemo_gym(master_config): + if should_use_nemo_gym(master_config): generation_config = master_config.policy["generation"] nemo_gym_rollout_result = run_nemo_gym_rollout_sync( policy_generation=policy_generation, @@ -1321,7 +1410,7 @@ def ppo_train( rollout_metrics = nemo_gym_rollout_result.rollout_metrics del nemo_gym_rollout_result - elif _should_use_async_rollouts(master_config): + elif should_use_async_rollouts(master_config.policy["generation"]): ( repeated_batch, rollout_metrics, @@ -1333,7 +1422,7 @@ def ppo_train( max_seq_len=master_config.policy[ "max_total_sequence_length" ], - max_rollout_turns=master_config.ppo["max_rollout_turns"], + max_rollout_turns=master_config.ppo.max_rollout_turns, greedy=False, ) else: @@ -1345,7 +1434,7 @@ def ppo_train( max_seq_len=master_config.policy[ "max_total_sequence_length" ], - max_rollout_turns=master_config.ppo["max_rollout_turns"], + max_rollout_turns=master_config.ppo.max_rollout_turns, greedy=False, ) policy_generation.finish_generation() @@ -1357,9 +1446,9 @@ def ppo_train( logger.log_metrics(rollout_metrics, total_steps + 1, prefix="train") repeated_batch = scale_rewards( - repeated_batch, master_config.ppo["reward_scaling"] + repeated_batch, master_config.ppo.reward_scaling ) - reward_shaping_config = master_config.ppo["reward_shaping"] + reward_shaping_config = master_config.ppo.reward_shaping if reward_shaping_config.enabled: repeated_batch = apply_reward_shaping( repeated_batch, reward_shaping_config @@ -1372,7 +1461,7 @@ def ppo_train( rewards = repeated_batch["total_reward"] with timer.time("data_processing"): - use_overlong_filtering = master_config.ppo["overlong_filtering"] + use_overlong_filtering = master_config.ppo.overlong_filtering if use_overlong_filtering: loss_multiplier = repeated_batch["loss_multiplier"].clone() truncated = repeated_batch["truncated"] @@ -1454,9 +1543,7 @@ def ppo_train( logprob_data, timer=timer )["logprobs"] - if not master_config.ppo.get( - "skip_reference_policy_logprobs_calculation" - ): + if not master_config.ppo.skip_reference_policy_logprobs_calculation: train_data["reference_policy_logprobs"] = ( policy.get_reference_policy_logprobs( logprob_data, @@ -1475,9 +1562,9 @@ def ppo_train( ) = _apply_ppo_seq_logprob_error_masking( train_data=train_data, rewards=rewards, - seq_logprob_error_threshold=master_config.ppo[ - "seq_logprob_error_threshold" - ], + seq_logprob_error_threshold=( + master_config.ppo.seq_logprob_error_threshold + ), ) # Build prompt IDs for advantage estimation (groups responses from same prompt). @@ -1657,45 +1744,7 @@ def ppo_train( # Extract critic metrics from value training results if value_results is not None: - value_mb_metrics = value_results.get("all_mb_metrics", {}) - critic_metrics = { - "critic/grad_norm": value_results["grad_norm"].numpy(), - "critic/loss": value_results["loss"].numpy(), - } - - for k, v in value_mb_metrics.items(): - if k in { - "lr", - "wd", - "global_valid_seqs", - "global_valid_toks", - "grad_norm", - }: - critic_metrics["critic/" + k] = np.mean(v).item() - elif k in {"values_min"}: - critic_metrics["critic/" + k] = np.min(v).item() - elif k in {"values_max"}: - critic_metrics["critic/" + k] = np.max(v).item() - elif isinstance(v, (np.ndarray, list)): - critic_metrics["critic/" + k] = np.sum(v).item() - else: - raise ValueError( - f"Unknown metric for value don't know how to handle: {k}" - ) - - # Compute explained variance from sufficient statistics: - # EV = 1 - Var(returns - values) / Var(returns) - r_mean = critic_metrics.get("critic/returns_mean", 0) - v_mean = critic_metrics.get("critic/values_mean", 0) - r_sq = critic_metrics.get("critic/returns_sq_mean", 0) - res_sq = critic_metrics.get("critic/residual_sq_mean", 0) - var_returns = r_sq - r_mean**2 - var_residual = res_sq - (r_mean - v_mean) ** 2 - critic_metrics["critic/explained_var"] = 1.0 - var_residual / max( - var_returns, 1e-8 - ) - - metrics.update(critic_metrics) + metrics.update(_compute_critic_metrics(value_results)) metrics.update( { "reward": rewards.numpy(), @@ -1750,7 +1799,7 @@ def ppo_train( total_valid_tokens += metrics["global_valid_toks"] ## Checkpointing - consumed_samples += master_config.ppo["num_prompts_per_step"] + consumed_samples += master_config.ppo.num_prompts_per_step timeout.mark_iteration() should_save_by_step = ( @@ -1897,8 +1946,8 @@ def ppo_train( total_time = timing_metrics.get("total_step_time", 0) number_of_samples_per_step = ( - master_config.ppo["num_prompts_per_step"] - * master_config.ppo["num_generations_per_prompt"] + master_config.ppo.num_prompts_per_step + * master_config.ppo.num_generations_per_prompt ) total_num_gpus = ( master_config.cluster["num_nodes"] @@ -1924,6 +1973,11 @@ def ppo_train( metrics, timing_metrics, master_config, + num_prompts_per_step=master_config.ppo.num_prompts_per_step, + num_generations_per_prompt=( + master_config.ppo.num_generations_per_prompt + ), + is_async_rl=master_config.ppo.async_ppo.enabled, ) logger.log_metrics(metrics, total_steps + 1, prefix="train") @@ -1972,47 +2026,1055 @@ def ppo_train( checkpointer.shutdown() -def validate( - policy_generation: GenerationInterface, +def _async_ppo_generation_lead_steps( + *, + step: int, + policy_training_start_step: int, + max_trajectory_age_steps: int, + warmup_generation_lead_steps: int, +) -> int: + """Return the collector lead without crossing the safe warmup frontier.""" + if step >= policy_training_start_step: + return max_trajectory_age_steps + + max_warmup_target = policy_training_start_step + max_trajectory_age_steps + remaining_to_frontier = max_warmup_target - step + return max( + max_trajectory_age_steps, + min(warmup_generation_lead_steps, remaining_to_frontier), + ) + + +def _async_ppo_buffer_max_age( + *, + step: int, + policy_training_start_step: int, + max_trajectory_age_steps: int, + warmup_generation_lead_steps: int, +) -> int: + """Keep frozen-policy rollouts valid through their safe training frontier.""" + warmup_rollout_frontier = policy_training_start_step + max_trajectory_age_steps + if policy_training_start_step > 0 and step <= warmup_rollout_frontier: + return warmup_generation_lead_steps + return max_trajectory_age_steps + + +def async_ppo_train( + policy: ColocatablePolicyInterface, + policy_generation: Optional[GenerationInterface], + value_model: ValueInterface, + dataloader: StatefulDataLoader, val_dataloader: Optional[StatefulDataLoader], - tokenizer, + tokenizer: TokenizerType, + loss_fn: LossFunction, + value_loss_fn: LossFunction, + task_to_env: dict[str, EnvironmentInterface], val_task_to_env: Optional[dict[str, EnvironmentInterface]], - step: int, + logger: Logger, + checkpointer: CheckpointManager, + ppo_save_state: PPOSaveState, master_config: MasterConfig, - logger: Optional[Logger] = None, -) -> tuple[dict[str, Any], dict[str, Any]]: - """Run validation on the validation dataset.""" - if val_dataloader is None: - assert val_dataloader is not None or master_config.ppo["val_period"] == 0, ( - "val_dataloader is None, so ppo.val_period must be 0" +) -> None: + """Run PPO while a background collector fills a replay buffer.""" + colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] + async_config = master_config.ppo.async_ppo + max_trajectory_age_steps = async_config.max_trajectory_age_steps + warmup_generation_lead_steps = async_config.resolved_warmup_generation_lead_steps + policy_training_start_step = master_config.ppo.policy_training_start_step + if master_config.ppo.ppo_epochs < 1: + raise ValueError("ppo.ppo_epochs must be at least 1") + if max_trajectory_age_steps > 1: + print( + "⚠️ WARNING: max_trajectory_age_steps > 1 increases off-policy " + "bias in GAE. The validated/recommended value is 1." ) - print(" ⚠️ No validation dataloader provided, skipping validation", flush=True) - return {}, {} + if not async_config.in_flight_weight_updates: + print( + "⚠️ WARNING: In-flight weight updates must be enabled for async " + "PPO with max_trajectory_age_steps > 1. Without in-flight weight " + "updates, a larger trajectory age provides no performance benefit." + ) - timer = Timer() - with timer.time("total_validation_time"): - print(f"β–Ά Starting validation at step {step}...", flush=True) + # Import async utilities only when needed (heavy Ray actors). + from nemo_rl.algorithms.async_utils import AsyncTrajectoryCollector, ReplayBuffer - total_rewards = [] - total_lengths = [] - all_message_logs = [] # Collect all message logs + timer = Timer(context={"worker": "driver"}) + training_wall_start = time.perf_counter() + timeout = TimeoutChecker( + timeout=master_config.checkpointing["checkpoint_must_save_by"], + fit_last_save_time=True, + ) + timeout.start_iterations() - max_batches = ( - master_config.ppo["max_val_samples"] // master_config.ppo["val_batch_size"] + # PPO async always uses non-colocated vLLM generation, so a refit is always + # required and the generation engine is a real (non-None) actor. + assert policy_generation is not None + + if master_config.ppo.skip_reference_policy_logprobs_calculation: + if master_config.loss_fn.reference_policy_kl_penalty != 0: + raise ValueError( + "Skipping reference logprobs requires " + "loss_fn.reference_policy_kl_penalty=0" + ) + + # ------------------------------------------------------------------ + # Training state. `step` is the global monotonic training step; it is what + # max_num_steps bounds and what the replay-buffer weight versioning tracks. + # ------------------------------------------------------------------ + step = ppo_save_state["total_steps"] + weight_version = step + consumed_samples = ppo_save_state["consumed_samples"] + total_valid_tokens = ppo_save_state.get("total_valid_tokens", 0) + max_num_steps = master_config.ppo.max_num_steps + ppo_epochs = master_config.ppo.ppo_epochs + val_period = master_config.ppo.val_period + val_at_start = master_config.ppo.val_at_start + val_at_end = master_config.ppo.val_at_end + num_prompts_per_step = master_config.ppo.num_prompts_per_step + ft_save_period = master_config.checkpointing.get("ft_save_period") + max_training_steps = max_num_steps + + replay_buffer: Any = None + trajectory_collector: Any = None + + def _shutdown_workers(*, propagate_checkpoint_error: bool) -> None: + """Finalize pending saves and stop async PPO workers.""" + checkpoint_error = None + try: + checkpointer.shutdown() + except Exception as error: + checkpoint_error = error + print(f"Error finalizing pending checkpoint: {error}") + + print("πŸ›‘ Stopping trajectory collection...") + for actor, actor_name in ( + (trajectory_collector, "trajectory collector"), + (replay_buffer, "replay buffer"), + ): + if actor is None: + continue + try: + ray.kill(actor) + except Exception as error: + print(f"Error stopping {actor_name}: {error}") + + for env_dict in (task_to_env, val_task_to_env): + if env_dict is None: + continue + for task_name, env in env_dict.items(): + print(f"πŸ›‘ Shutting down environment {task_name}...") + try: + ray.get(env.shutdown.remote(), timeout=10) + except Exception: + try: + ray.kill(env) + except Exception as error: + print(f"Error shutting down environment {task_name}: {error}") + + print("πŸ›‘ Shutting down generation workers...") + try: + policy_generation.shutdown() + except Exception as error: + print(f"Error shutting down generation workers: {error}") + if policy is not policy_generation: + print("πŸ›‘ Shutting down policy workers...") + try: + policy.shutdown() + except Exception as error: + print(f"Error shutting down policy workers: {error}") + print("πŸ›‘ Shutting down value workers...") + try: + value_model.shutdown() + except Exception as error: + print(f"Error shutting down value workers: {error}") + + if checkpoint_error is not None and propagate_checkpoint_error: + raise checkpoint_error + + if step >= max_training_steps: + print( + f"Training is already complete at step {step} " + f"(configured limit: {max_training_steps})" ) - for batch_idx, val_batch in enumerate(val_dataloader): - if batch_idx >= max_batches: - break + _shutdown_workers(propagate_checkpoint_error=True) + return - additional_metrics_to_report = dict() + adv_estimator = _create_advantage_estimator(master_config) + + # ------------------------------------------------------------------ + # Spin up the replay buffer + trajectory collector Ray actors. + # ------------------------------------------------------------------ + late_arrival_slack = 2 + buffer_age = max( + max_trajectory_age_steps, + warmup_generation_lead_steps, + ) + optimal_buffer_size = num_prompts_per_step * buffer_age * late_arrival_slack + print("πŸ“Š Async PPO buffer requirements:") + print(f" - num_prompts_per_step: {num_prompts_per_step}") + print(f" - max_trajectory_age_steps: {max_trajectory_age_steps}") + print(f" - warmup_generation_lead_steps: {warmup_generation_lead_steps}") + print(f" - optimal_buffer_size: {optimal_buffer_size}") + + replay_buffer = ReplayBuffer.options( + runtime_env=make_actor_runtime_env( + "nemo_rl.algorithms.async_utils.ReplayBuffer" + ) + ).remote( + max_size=optimal_buffer_size, + drop_incomplete_targets_on_restore=( + async_config.drop_incomplete_targets_on_restore + ), + ) + + last_checkpoint_path = checkpointer.get_latest_checkpoint_path() + if last_checkpoint_path is not None: + replay_buffer_path = os.path.join(last_checkpoint_path, "replay_buffer.pt") + if os.path.exists(replay_buffer_path): + print(f"πŸ“¦ Restoring replay buffer from checkpoint: {replay_buffer_path}") + restore_max_age = _async_ppo_buffer_max_age( + step=step, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=warmup_generation_lead_steps, + ) + ray.get( + replay_buffer.load_from_path.remote( + replay_buffer_path, + num_prompts_per_step=num_prompts_per_step, + current_training_step=step, + max_age_steps=restore_max_age, + ) + ) + print("βœ… Replay buffer restored from checkpoint") + else: + print( + f"⚠️ No replay buffer checkpoint found at {replay_buffer_path}. " + "Starting with an empty replay buffer." + ) + + trajectory_collector = AsyncTrajectoryCollector.options( + runtime_env=make_actor_runtime_env( + "nemo_rl.algorithms.async_utils.AsyncTrajectoryCollector" + ) + ).remote( + policy_generation=policy_generation, + tokenizer=tokenizer, + task_to_env=task_to_env, + master_config=master_config, + replay_buffer=replay_buffer, + start_step=step, + ) + + def _raise_if_collector_stopped(waiting_for: str) -> None: + ray.get(trajectory_collector.check_health.remote()) + status = ray.get(trajectory_collector.get_status.remote()) + if status["errored"]: + raise RuntimeError( + f"Trajectory collector failed while {waiting_for}: " + f"{status.get('error') or status}" + ) + if ( + not status["running"] + and status["inflight_workers"] == 0 + and status["data_exhausted"] + ): + raise RuntimeError( + "Trajectory collector exhausted data before the configured " + f"training limit while {waiting_for}: {status}" + ) + + try: + # Refit first so resumed runs cannot generate with stale base weights. + print("⏳ Preparing policy generation for training (initial refit)...") + refit_policy_generation(policy, policy_generation, colocated_inference) + policy.offload_to_cpu() - val_batch, gen_metrics = run_multi_turn_rollout( + if val_at_start and step == 0: + print("\nπŸ” Running initial validation...") + val_metrics, validation_timings = validate( policy_generation, - val_batch, + val_dataloader, tokenizer, val_task_to_env, + step=0, + master_config=master_config, + logger=logger, + ) + policy_generation.finish_generation() + logger.log_metrics(val_metrics, step, prefix="validation") + logger.log_metrics(validation_timings, step, prefix="timing/validation") + + policy_generation.clear_logger_metrics() + + initial_generation_lead = _async_ppo_generation_lead_steps( + step=step, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=warmup_generation_lead_steps, + ) + initial_buffer_max_age = _async_ppo_buffer_max_age( + step=step, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=warmup_generation_lead_steps, + ) + ray.get( + trajectory_collector.set_generation_window.remote( + weight_version=weight_version, + generation_lead_steps=initial_generation_lead, + max_trajectory_age_steps=initial_buffer_max_age, + ) + ) + ray.get( + trajectory_collector.start_collection.remote(CyclingDataLoader(dataloader)) + ) + print("πŸ“¦ Started continuous background trajectory collection") + + print(f"⏳ Waiting for replay buffer to be ready for step {step}...") + timer.start("init/total") + wait_iterations = 0 + while True: + current_step_ready = ray.get( + replay_buffer.has_complete_batch.remote( + step, num_prompts_per_step, initial_buffer_max_age + ) + ) + if current_step_ready: + # The initial collector is the only window that can generate + # both `step` and `step + 1`. Fill both before the first refit. + need_lookahead = step + 1 < max_training_steps + if need_lookahead: + lookahead_step_ready = ray.get( + replay_buffer.has_complete_batch.remote( + step + 1, + num_prompts_per_step, + initial_buffer_max_age, + ) + ) + if not lookahead_step_ready: + if wait_iterations % 10 == 0: + print( + f" Pipeline barrier: step {step} ready but " + f"step {step + 1} not yet β€” waiting for lookahead fill" + ) + _raise_if_collector_stopped( + "waiting for the initial replay lookahead batch" + ) + wait_iterations += 1 + time.sleep(1.0) + continue + break + if wait_iterations % 10 == 0: + buffer_size_current = ray.get(replay_buffer.size.remote()) + print( + f" Wait iteration {wait_iterations}: " + f"buffer_size={buffer_size_current}, " + f"step {step} ready={current_step_ready}" + ) + _raise_if_collector_stopped("waiting for the initial replay batch") + wait_iterations += 1 + time.sleep(1.0) + timer.stop("init/total") + print(f"βœ… Buffer ready for step {step}! Starting async PPO training loop...") + except Exception: + _shutdown_workers(propagate_checkpoint_error=False) + raise + + # ------------------------------------------------------------------ + # Main loop + # ------------------------------------------------------------------ + loop_failed = False + try: + while step < max_training_steps: + ray.get(trajectory_collector.check_health.remote()) + print(f"\n{'=' * 25} Step {step + 1}/{max_training_steps} {'=' * 25}") + maybe_gpu_profile_step(policy, step + 1) + if policy != policy_generation: + maybe_gpu_profile_step(policy_generation, step + 1) + + metrics: dict[str, Any] = {} + val_metrics, validation_timings = None, None + + with timer.time("total_step_time"): + # ---- 1. Sample a fixed batch of trajectories from the buffer ---- + print("πŸ“¦ Sampling from replay buffer...") + with timer.time("exposed_generation"): + current_buffer_max_age = _async_ppo_buffer_max_age( + step=step, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=(warmup_generation_lead_steps), + ) + sample_result = ray.get( + replay_buffer.sample.remote( + num_prompt_groups=num_prompts_per_step, + current_weight_version=weight_version, + max_age_steps=current_buffer_max_age, + ) + ) + if ( + sample_result is None + or len(sample_result["trajectories"]) != num_prompts_per_step + ): + print( + "⏳ Buffer empty or not enough groups for a full step, " + "waiting..." + ) + _raise_if_collector_stopped( + f"waiting for the replay batch for step {step}" + ) + with timer.time("idle/buffer_starvation"): + time.sleep(0.5) + continue + + trajectories = sample_result["trajectories"] + avg_trajectory_age = sample_result["avg_trajectory_age"] + print( + f"βœ… Sampled {len(trajectories)} trajectory groups " + f"(average age: {avg_trajectory_age:.2f} steps)" + ) + + per_prompt_batches = [t["batch"] for t in trajectories] + repeated_batch = BatchedDataDict.from_batches(per_prompt_batches) + + per_group_metrics: dict[str, list] = {} + for t in trajectories: + for k, v in t["rollout_metrics"].items(): + per_group_metrics.setdefault(k, []).append(v) + rollout_metrics = aggregate_rollout_metrics(per_group_metrics) + + expected_batch_size = ( + master_config.ppo.num_prompts_per_step + * master_config.ppo.num_generations_per_prompt + ) + if repeated_batch.size != expected_batch_size: + raise RuntimeError( + f"Unexpected training batch size: got {repeated_batch.size}, " + f"expected {expected_batch_size}" + ) + + # ---- 2. Build PPO training data (rewards + inline loss mask) ---- + print("β–Ά Processing rewards...") + with timer.time("data_processing"): + rewards = repeated_batch["total_reward"] + + use_overlong_filtering = master_config.ppo.overlong_filtering + if use_overlong_filtering: + loss_multiplier = repeated_batch["loss_multiplier"].clone() + truncated = repeated_batch["truncated"] + if isinstance(truncated, list): + truncated = torch.tensor(truncated, dtype=torch.bool) + loss_multiplier[truncated] = 0 + repeated_batch["loss_multiplier"] = loss_multiplier + + # PPO's inline loss-mask setup (unmask all assistant messages), + # matching sync ppo_train β€” deliberately NOT GRPO's helper, + # which only unmasks generated assistant messages. + for message_log in repeated_batch["message_log"]: + for message in message_log: + if message["role"] == "assistant": + message["token_loss_mask"] = torch.ones_like( + message["token_ids"] + ) + else: + message["token_loss_mask"] = torch.zeros_like( + message["token_ids"] + ) + if "generation_logprobs" not in message: + message["generation_logprobs"] = torch.zeros_like( + message["token_ids"], dtype=torch.float32 + ) + + flat_messages, input_lengths = batched_message_log_to_flat_message( + repeated_batch["message_log"], + pad_value_dict={"token_ids": tokenizer.pad_token_id}, + make_sequence_length_divisible_by=master_config.policy[ + "make_sequence_length_divisible_by" + ], + ) + + train_data = BatchedDataDict[ClippedPGLossDataDict]( + { + "input_ids": flat_messages["token_ids"], + "input_lengths": input_lengths, + "generation_logprobs": flat_messages["generation_logprobs"], + "rewards": repeated_batch["total_reward"], + "token_mask": flat_messages["token_loss_mask"], + "sample_mask": repeated_batch["loss_multiplier"], + } + ) + extra_multimodal_data = flat_messages.get_multimodal_dict( + as_tensors=False + ) + train_data.update(extra_multimodal_data) + train_data.to("cpu") + + # ---- 3. Value forward (critic on GPU, then offloaded) ---- + # GPU state entering here: policy OFF, value OFF (see refit/step + # end below). Load value only. + print("β–Ά Computing values...") + with timer.time("value_inference"): + value_model.prepare_for_inference() + train_data["values"] = value_model.get_values(train_data)[ + "values" + ].squeeze(-1) + value_model.finish_inference() + + # ---- 4. Policy / reference logprobs (policy on GPU, then off) ---- + print("β–Ά Computing logprobs...") + with timer.time("logprob_inference_prep"): + policy.prepare_for_lp_inference() + with timer.time("policy_and_reference_logprobs"): + logprob_data = BatchedDataDict[ClippedPGLossDataDict]( + { + "input_ids": train_data["input_ids"], + "input_lengths": train_data["input_lengths"], + **extra_multimodal_data, + } + ) + train_data["prev_logprobs"] = policy.get_logprobs( + logprob_data, timer=timer + )["logprobs"] + if not master_config.ppo.skip_reference_policy_logprobs_calculation: + train_data["reference_policy_logprobs"] = ( + policy.get_reference_policy_logprobs( + logprob_data, + timer=timer, + )["reference_logprobs"] + ) + del logprob_data + del extra_multimodal_data + policy.finish_inference() + + # ---- 5. Sequence-level train/inference mismatch diagnostics ---- + ( + advantage_mask, + seq_logprob_error_metrics, + ) = _apply_ppo_seq_logprob_error_masking( + train_data=train_data, + rewards=rewards, + seq_logprob_error_threshold=( + master_config.ppo.seq_logprob_error_threshold + ), + ) + + # ---- 6. GAE advantages/returns (uses fresh values) ---- + with timer.time("advantage_calculation"): + print("β–Ά Computing advantages...") + initial_prompt_message_logs = extract_initial_prompt_messages( + repeated_batch["message_log"], + repeated_batch["length"], + ) + prompt_batched_flat, _ = batched_message_log_to_flat_message( + initial_prompt_message_logs, + pad_value_dict={"token_ids": tokenizer.pad_token_id}, + ) + prompt_ids_for_adv = prompt_batched_flat["token_ids"] + del initial_prompt_message_logs + del prompt_batched_flat + + adv_kwargs = dict( + prompt_ids=prompt_ids_for_adv, + rewards=train_data["rewards"], + mask=advantage_mask, + reference_logprobs=train_data.get("reference_policy_logprobs"), + logprobs=train_data["prev_logprobs"], + ) + if "values" in train_data: + adv_kwargs["values"] = train_data["values"] + result = adv_estimator.compute_advantage(**adv_kwargs) + if isinstance(result, tuple): + advantages, returns = result + else: + advantages, returns = result, None + del prompt_ids_for_adv + train_data["advantages"] = advantages + if returns is not None: + train_data["returns"] = returns + + # ---- 7. ppo_epochs inner loop (critic, then actor) ---- + # Each epoch: value on GPU -> train -> off. Then, once past critic + # warmup, policy on GPU -> train -> off (except the last epoch, + # which leaves the policy on GPU for the refit broadcast below). + # During warmup (step < policy_training_start_step) the policy is + # frozen: it is never loaded/trained here, exactly as in sync + # ppo_train, so train_results stays None for the step. + is_policy_training_step = step >= policy_training_start_step + train_results = None + value_results = None + for epoch in range(ppo_epochs): + print(f"β–Ά PPO epoch {epoch + 1}/{ppo_epochs}...") + with timer.time("value_training_prep"): + value_model.prepare_for_training() + with timer.time("value_training"): + value_results = value_model.train( + train_data, + value_loss_fn, + timer=timer, + ) + value_model.finish_training() + + if is_policy_training_step: + if ( + step == policy_training_start_step + and policy_training_start_step > 0 + and epoch == 0 + ): + print( + f" βœ“ Critic warmup complete ({policy_training_start_step} " + "steps). Starting policy training.", + flush=True, + ) + with timer.time("training_prep"): + policy.prepare_for_training() + with timer.time("policy_training"): + train_results = policy.train( + train_data, loss_fn, timer=timer + ) + if epoch < ppo_epochs - 1: + policy.offload_to_cpu() + + # ---- 8. Refit once after all PPO epochs ---- + # Warmup still advances the replay-buffer version, but skips the + # transfer because the policy has not changed. + generation_logger_metrics = None + print("πŸ”„ Coordinating with trajectory collector before refit...") + next_weight_version = weight_version + 1 + with timer.time("idle/refit_bubble"): + with timer.time("exposed_generation"): + ray.get(trajectory_collector.prepare_for_refit.remote()) + generation_logger_metrics = policy_generation.get_logger_metrics() + with timer.time("weight_sync"): + if is_policy_training_step: + refit_policy_generation( + policy, policy_generation, colocated_inference + ) + else: + print( + "β–Ά Critic warmup: skipping policy weight transfer " + "(policy frozen; generation already up to date)" + ) + weight_version = next_weight_version + next_generation_lead = _async_ppo_generation_lead_steps( + step=weight_version, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=(warmup_generation_lead_steps), + ) + next_buffer_max_age = _async_ppo_buffer_max_age( + step=weight_version, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=(warmup_generation_lead_steps), + ) + ray.get( + trajectory_collector.set_generation_window.remote( + weight_version=weight_version, + generation_lead_steps=next_generation_lead, + max_trajectory_age_steps=next_buffer_max_age, + ) + ) + ray.get(trajectory_collector.resume_after_refit.remote()) + # Only the policy-training path leaves the policy resident on GPU; + # during warmup it is already offloaded, so skip the redundant call. + if is_policy_training_step: + policy.offload_to_cpu() + + policy_generation.clear_logger_metrics() + + # ---- Validation ---- + is_last_step = step + 1 == max_training_steps + if (val_period > 0 and (step + 1) % val_period == 0) or ( + val_at_end and is_last_step + ): + with timer.time("idle/validation"): + ray.get(trajectory_collector.pause.remote()) + # Policy weights were synced by the refit above. + policy_generation.prepare_for_generation() + val_metrics, validation_timings = validate( + policy_generation, + val_dataloader, + tokenizer, + val_task_to_env, + step=step + 1, + master_config=master_config, + logger=logger, + ) + policy_generation.finish_generation() + logger.log_metrics( + validation_timings, step + 1, prefix="timing/validation" + ) + logger.log_metrics(val_metrics, step + 1, prefix="validation") + gc.collect() + torch.cuda.empty_cache() + ray.get(trajectory_collector.resume.remote()) + + # ---- Metrics ---- + flat_advantages = train_data["advantages"] + flat_messages_content = flat_messages.get("content", []) + del flat_messages + response_advantages = torch.masked_select( + flat_advantages, advantage_mask.bool() + ) + + metrics.update( + { + "reward": rewards.numpy(), + "mean_prompt_length": repeated_batch["length"].numpy(), + "total_num_tokens": input_lengths.numpy(), + "advantages/mean": torch.mean(response_advantages) + .detach() + .item() + if response_advantages.numel() > 0 + else 0.0, + "advantages/max": torch.max(response_advantages).detach().item() + if response_advantages.numel() > 0 + else 0.0, + "advantages/min": torch.min(response_advantages).detach().item() + if response_advantages.numel() > 0 + else 0.0, + } + ) + # Policy metrics are absent during critic warmup (train_results is + # None because the policy was not trained this step). + if train_results is not None: + metrics["loss"] = train_results["loss"].numpy() + metrics["grad_norm"] = train_results["grad_norm"].numpy() + if "moe_metrics" in train_results: + metrics.update( + { + f"moe/{k}": v + for k, v in train_results["moe_metrics"].items() + } + ) + metrics.update(train_results["all_mb_metrics"]) + if value_results is not None: + metrics.update(_compute_critic_metrics(value_results)) + + for k, v in metrics.items(): + if k in {"probs_ratio_min", "probs_ratio_clamped_min"}: + valid_values = [x for x in v if not np.isinf(x)] + metrics[k] = ( + np.min(valid_values).item() if valid_values else -1.0 + ) + elif k in {"probs_ratio_max", "probs_ratio_clamped_max"}: + valid_values = [x for x in v if not np.isinf(x)] + metrics[k] = ( + np.max(valid_values).item() if valid_values else -1.0 + ) + elif k in { + "lr", + "wd", + "reward", + "global_valid_seqs", + "global_valid_toks", + "mean_prompt_length", + }: + metrics[k] = np.mean(v).item() + elif isinstance(v, (np.ndarray, list)): + metrics[k] = np.sum(v).item() + + metrics.update(rollout_metrics) + if generation_logger_metrics is not None: + metrics["generation_logger_metrics"] = generation_logger_metrics + if "global_valid_toks" in metrics: + total_valid_tokens += metrics["global_valid_toks"] + # Always log seq-level error metrics (useful for tuning threshold). + metrics.update(seq_logprob_error_metrics) + + # ---- Checkpointing ---- + consumed_samples += master_config.ppo.num_prompts_per_step + timeout.mark_iteration() + should_save_by_step = ( + is_last_step + or (step + 1) % master_config.checkpointing["save_period"] == 0 + or (ft_save_period is not None and (step + 1) % ft_save_period == 0) + ) + should_save_by_timeout = timeout.check_save() + if master_config.checkpointing["enabled"] and ( + should_save_by_step or should_save_by_timeout + ): + ppo_save_state["current_step"] = step + 1 + ppo_save_state["total_steps"] = step + 1 + ppo_save_state["total_valid_tokens"] = total_valid_tokens + if val_metrics is not None: + ppo_save_state["val_reward"] = val_metrics["accuracy"] + elif "val_reward" in ppo_save_state: + del ppo_save_state["val_reward"] + ppo_save_state["consumed_samples"] = consumed_samples + + # Record the top-k ranking metric into the save state so + # get_best_checkpoint_path / top-k pruning work (parity with + # sync ppo_train and async_grpo_train). + full_metric_name = master_config.checkpointing["metric_name"] + if full_metric_name is not None: + assert full_metric_name.startswith( + "train:" + ) or full_metric_name.startswith("val:"), ( + f"metric_name={full_metric_name} must start with 'val:' or 'train:',\n" + f'followed by the corresponding name in the "val" or "train" metrics dictionary.' + ) + prefix, metric_name = full_metric_name.split(":", 1) + metrics_source = metrics if prefix == "train" else val_metrics + if not metrics_source: + warnings.warn( + f"You asked to save checkpoints based on {metric_name} but no {prefix} metrics were collected. " + "This checkpoint will not be saved as top-k.", + stacklevel=2, + ) + if full_metric_name in ppo_save_state: + del ppo_save_state[full_metric_name] + elif metric_name not in metrics_source: + raise ValueError( + f"Metric {metric_name} not found in {prefix} metrics" + ) + else: + ppo_save_state[full_metric_name] = metrics_source[ + metric_name + ] + + with timer.time("checkpointing"): + print(f"Saving checkpoint for step {step + 1}...") + checkpoint_path = checkpointer.init_tmp_checkpoint( + step + 1, ppo_save_state, master_config + ) + policy.prepare_for_training() + policy.save_checkpoint( + weights_path=os.path.join( + checkpoint_path, "policy", "weights" + ), + optimizer_path=( + os.path.join(checkpoint_path, "policy", "optimizer") + if ( + checkpointer.save_optimizer + and step >= policy_training_start_step + ) + else None + ), + tokenizer_path=os.path.join( + checkpoint_path, "policy", "tokenizer" + ), + checkpointing_cfg=master_config.checkpointing, + ) + policy.offload_to_cpu() + + value_model.prepare_for_training() + value_model.save_checkpoint( + weights_path=os.path.join( + checkpoint_path, "value", "weights" + ), + optimizer_path=( + os.path.join(checkpoint_path, "value", "optimizer") + if checkpointer.save_optimizer + else None + ), + tokenizer_path=os.path.join( + checkpoint_path, "value", "tokenizer" + ), + checkpointing_cfg=master_config.checkpointing, + ) + value_model.finish_training() + + dataloader_state = ray.get( + trajectory_collector.get_dataloader_state.remote() + ) + torch.save( + dataloader_state, + os.path.join(checkpoint_path, "train_dataloader.pt"), + ) + print("πŸ“¦ Saving replay buffer state...") + num_buffered_trajectories = ray.get( + replay_buffer.save_to_path.remote( + os.path.join(checkpoint_path, "replay_buffer.pt") + ) + ) + print( + "βœ… Saved replay buffer with " + f"{num_buffered_trajectories} trajectories" + ) + checkpointer.begin_finalization( + checkpoint_path, + wait_fn=policy.finalize_async_save, + ) + + # ---- Logging ---- + log_data = { + "content": flat_messages_content, + "rewards": rewards.tolist(), + "input_lengths": input_lengths.tolist(), + "token_ids": train_data["input_ids"].tolist(), + "token_loss_mask": train_data["token_mask"].tolist(), + "sample_loss_mask": train_data["sample_mask"].tolist(), + "advantages": train_data["advantages"].tolist(), + "generation_logprobs": train_data["generation_logprobs"].tolist(), + "prev_logprobs": train_data["prev_logprobs"].tolist(), + } + logger.log_batched_dict_as_jsonl( + log_data, f"train_data_step{step + 1}.jsonl" + ) + del log_data + del flat_messages_content + + timing_metrics: dict[str, float] = timer.get_timing_metrics( + reduction_op="sum" + ) # type: ignore + + buffer_size_current = ray.get(replay_buffer.size.remote()) + metrics["buffer_size"] = buffer_size_current + metrics["avg_trajectory_age"] = avg_trajectory_age + + # Track the worst-mismatch example plot (parity with sync PPO). + if metrics.get("token_mult_prob_error", 0) > 1.05: + logger.log_plot_token_mult_prob_error( + { + "prompt_lengths": repeated_batch["length"], + "full_lengths": input_lengths, + "generation_logprobs": train_data["generation_logprobs"], + "prev_logprobs": train_data["prev_logprobs"], + "token_mask": train_data["token_mask"], + "sample_mask": train_data["sample_mask"], + }, + step + 1, + name="train/token_mult_prob_error_plot_sample", + ) + del train_data + + print("\nπŸ“Š Training Results:") + if "loss" in metrics: + print(f" β€’ Loss: {metrics['loss']:.4f}") + print(f" β€’ Generation KL Error: {metrics.get('gen_kl_error', 'N/A')}") + else: + print(" β€’ (critic warmup: policy not trained this step)") + if "critic/loss" in metrics: + print(f" β€’ Critic Loss: {metrics['critic/loss']:.4f}") + print(f" β€’ Avg Reward: {np.mean(rewards.numpy()):.4f}") + print(f" β€’ Buffer Size: {buffer_size_current}") + print( + f" β€’ Avg Trajectory Age (gen-version): {avg_trajectory_age:.2f} steps" + ) + + total_time = timing_metrics.get("total_step_time", 0) + total_num_gpus = ( + master_config.cluster["num_nodes"] + * master_config.cluster["gpus_per_node"] + ) + if total_time > 0 and "global_valid_toks" in metrics: + timing_metrics["valid_tokens_per_sec_per_gpu"] = ( + metrics["global_valid_toks"] / total_time / total_num_gpus + ) + performance_metrics = print_performance_metrics( + train_results if train_results is not None else (value_results or {}), + metrics, + timing_metrics, + master_config, + num_prompts_per_step=master_config.ppo.num_prompts_per_step, + num_generations_per_prompt=( + master_config.ppo.num_generations_per_prompt + ), + is_async_rl=master_config.ppo.async_ppo.enabled, + ) + + collector_efficiency = ray.get( + trajectory_collector.get_efficiency_metrics.remote() + ) + driver_efficiency = { + cat: timer.reduce(cat, "sum") + for cat in [ + "init/total", + "idle/buffer_starvation", + "idle/refit_bubble", + "idle/validation", + ] + if cat in timer._timers + } + merged_efficiency = {**driver_efficiency} + for cat, dur in collector_efficiency.items(): + merged_efficiency[cat] = merged_efficiency.get(cat, 0.0) + dur + total_wall_time = time.perf_counter() - training_wall_start + efficiency_loggable = print_efficiency_summary( + merged_efficiency, total_wall_time, step + 1 + ) + + logger.log_metrics(performance_metrics, step + 1, prefix="performance") + logger.log_metrics(metrics, step + 1, prefix="train") + logger.log_metrics(efficiency_loggable, step + 1, prefix="") + logger.log_metrics( + timing_metrics, + step + 1, + prefix="timing/train", + step_finished=True, + ) + + timer.reset() + step += 1 + if should_save_by_timeout: + print("Timeout has been reached, stopping training early", flush=True) + return + if step >= max_training_steps: + print( + "Configured step/epoch limit has been reached, stopping training", + flush=True, + ) + return + + except Exception as e: + loop_failed = True + print(f"❌ Error in async PPO loop: {e}") + traceback.print_exc() + raise + + finally: + _shutdown_workers(propagate_checkpoint_error=not loop_failed) + print("Async PPO training complete!") + + +def validate( + policy_generation: GenerationInterface, + val_dataloader: Optional[StatefulDataLoader], + tokenizer, + val_task_to_env: Optional[dict[str, EnvironmentInterface]], + step: int, + master_config: MasterConfig, + logger: Optional[Logger] = None, +) -> tuple[dict[str, Any], dict[str, Any]]: + """Run validation on the validation dataset.""" + if val_dataloader is None: + assert val_dataloader is not None or master_config.ppo.val_period == 0, ( + "val_dataloader is None, so ppo.val_period must be 0" + ) + print(" ⚠️ No validation dataloader provided, skipping validation", flush=True) + return {}, {} + + timer = Timer() + with timer.time("total_validation_time"): + print(f"β–Ά Starting validation at step {step}...", flush=True) + + total_rewards = [] + total_lengths = [] + all_message_logs = [] # Collect all message logs + + max_batches = ( + master_config.ppo.max_val_samples // master_config.ppo.val_batch_size + ) + for batch_idx, val_batch in enumerate(val_dataloader): + if batch_idx >= max_batches: + break + + additional_metrics_to_report = dict() + + rollout_fn = ( + run_async_multi_turn_rollout + if should_use_async_rollouts(master_config.policy["generation"]) + else run_multi_turn_rollout + ) + val_batch, gen_metrics = rollout_fn( + policy_generation=policy_generation, + input_batch=val_batch, + tokenizer=tokenizer, + task_to_env=val_task_to_env, max_seq_len=master_config.policy["max_total_sequence_length"], - max_rollout_turns=master_config.ppo["max_rollout_turns"], + max_rollout_turns=master_config.ppo.max_rollout_turns, greedy=False, ) diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index d7b0a24add6..f9fa674366f 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -41,7 +41,7 @@ import time from collections import deque from functools import partial -from typing import Any, Awaitable, Callable, Optional, Union, cast +from typing import Any, Awaitable, Callable, Optional, Union import ray import torch @@ -67,6 +67,7 @@ from nemo_rl.data_plane import KVBatchMeta from nemo_rl.data_plane.schema import DP_CALIB_INPUT_FIELDS from nemo_rl.distributed.batched_data_dict import BatchedDataDict +from nemo_rl.environments.nemo_gym import should_use_nemo_gym from nemo_rl.experience.failures import RolloutStall from nemo_rl.experience.rollout_manager import RolloutOutcome from nemo_rl.models.generation.sglang.sglang_generation import SGLangGeneration @@ -1611,15 +1612,11 @@ async def _sync_weights( """ self._rollout_permitted.clear() - # TODO(#2625): abort unconditionally once gym-path abort is validated; - # for now only the native path aborts. Local import dodges the grpo.py - # circular dep (as in async_utils/trajectory_collector.py). - from nemo_rl.algorithms.grpo import MasterConfig as GrpoMasterConfig - from nemo_rl.algorithms.grpo import _should_use_nemo_gym - + # TODO(#2625): Abort unconditionally once Gym-path abort is validated; + # for now only the native path aborts stale in-flight requests. aborted_stale_inflight_groups = ( 0 - if _should_use_nemo_gym(cast(GrpoMasterConfig, self._master_config)) + if should_use_nemo_gym(self._master_config) else await self._abort_stale_inflight() ) diff --git a/nemo_rl/algorithms/single_controller_utils/setup.py b/nemo_rl/algorithms/single_controller_utils/setup.py index 4123022f8e1..0346d441c75 100644 --- a/nemo_rl/algorithms/single_controller_utils/setup.py +++ b/nemo_rl/algorithms/single_controller_utils/setup.py @@ -39,7 +39,6 @@ GRPOSaveState, _create_advantage_estimator, _get_grpo_save_state, - _should_use_nemo_gym, ) from nemo_rl.algorithms.grpo import MasterConfig as GrpoMasterConfig from nemo_rl.algorithms.loss import ClippedPGLossFn @@ -62,7 +61,7 @@ _get_node_ip_local, ) from nemo_rl.environments.interfaces import EnvironmentInterface -from nemo_rl.environments.nemo_gym import spinup_nemo_gym_actor +from nemo_rl.environments.nemo_gym import should_use_nemo_gym, spinup_nemo_gym_actor from nemo_rl.experience.rollout_manager import ( RolloutManager, RolloutRetryPolicy, @@ -574,7 +573,7 @@ def setup_single_controller( # Setup Dataset & Environments # ========================== # TODO: add validate dataset wiring. - use_nemo_gym = _should_use_nemo_gym(cast(GrpoMasterConfig, master_config)) + use_nemo_gym = should_use_nemo_gym(master_config) if use_nemo_gym and generation_config["backend"] != "vllm": raise NotImplementedError( "SC NeMo-Gym integration currently supports the vllm backend " diff --git a/nemo_rl/algorithms/utils.py b/nemo_rl/algorithms/utils.py index 28179b8a126..bca9766c404 100644 --- a/nemo_rl/algorithms/utils.py +++ b/nemo_rl/algorithms/utils.py @@ -521,8 +521,12 @@ def print_performance_metrics( metrics: dict[str, Any], timing_metrics: dict[str, float], master_config: dict, + *, + num_prompts_per_step: int, + num_generations_per_prompt: int, + is_async_rl: bool, ) -> dict[str, float]: - """Print performance metrics for GRPO.""" + """Print performance metrics for an RL training step.""" # ===================================================== # Generate Token Imbalance Visualization @@ -739,10 +743,10 @@ def visualize_per_worker_timeline( total_num_gpus = num_nodes * gpus_per_node colocated_inference = master_config.policy["generation"]["colocated"]["enabled"] - # Idle Time from Training Worker (Async GRPO only) + # Idle time from the training worker in async RL. if ( "exposed_generation" in timing_metrics - and master_config.grpo.async_grpo.enabled + and is_async_rl and not colocated_inference ): exposed_generation_time = timing_metrics["exposed_generation"] @@ -764,15 +768,6 @@ def visualize_per_worker_timeline( training_worker_idle_time_ratio ) - # Detect which algorithm config key is being used - grpo_config = getattr(master_config, "grpo", None) - algo_config = grpo_config or getattr(master_config, "ppo", None) or {} - if isinstance(algo_config, dict): - num_prompts_per_step = algo_config.get("num_prompts_per_step", 1) - num_generations_per_prompt = algo_config.get("num_generations_per_prompt", 1) - else: - num_prompts_per_step = algo_config.num_prompts_per_step - num_generations_per_prompt = algo_config.num_generations_per_prompt number_of_samples_per_step = num_prompts_per_step * num_generations_per_prompt if colocated_inference: diff --git a/nemo_rl/data/dataloader.py b/nemo_rl/data/dataloader.py index b1da3406ce0..ddef3fdf118 100644 --- a/nemo_rl/data/dataloader.py +++ b/nemo_rl/data/dataloader.py @@ -12,8 +12,40 @@ # See the License for the specific language governing permissions and # limitations under the License. +from collections.abc import Iterator +from typing import Any + from torchdata.stateful_dataloader import StatefulDataLoader +from nemo_rl.distributed.batched_data_dict import BatchedDataDict + + +class CyclingDataLoader: + """Repeat a stateful dataloader until its consumer stops.""" + + def __init__(self, dataloader: StatefulDataLoader) -> None: + self.dataloader = dataloader + + def __iter__(self) -> Iterator[BatchedDataDict]: + consecutive_empty_epochs = 0 + while True: + produced_this_epoch = False + for batch in self.dataloader: + produced_this_epoch = True + yield batch + + if produced_this_epoch: + consecutive_empty_epochs = 0 + else: + consecutive_empty_epochs += 1 + if consecutive_empty_epochs >= 2: + raise RuntimeError( + "Dataloader yielded no batches for two consecutive epochs" + ) + + def state_dict(self) -> dict[str, Any]: + return self.dataloader.state_dict() + class MultipleDataloaderWrapper: """Wrapper for multiple dataloaders. diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index 773db4626d6..f17275fdd28 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -19,7 +19,7 @@ from collections.abc import AsyncGenerator from copy import deepcopy from pathlib import Path -from typing import Any, Dict, List, NotRequired, Optional, TypedDict +from typing import Any, Dict, List, NotRequired, Optional, Protocol, TypedDict import ray import torch @@ -47,14 +47,14 @@ RolloutDataFailure, http_status_is_infra, ) -from nemo_rl.models.policy import TokenizerConfig +from nemo_rl.models.generation.interfaces import should_use_async_rollouts +from nemo_rl.models.policy import PolicyConfig, TokenizerConfig from nemo_rl.utils.routed_experts_codec import decode_routed_experts from nemo_rl.utils.timer import Timer from nemo_rl.utils.venvs import create_local_venv_on_each_node -# Kept local (not imported from models.generation) so the gym actor stays free of -# generation-module imports. Must cover every name resolve_routed_experts_dtype -# can produce. +# Kept local so the Gym actor does not depend on model-config dtype resolution. +# Must cover every name resolve_routed_experts_dtype can produce. _ROUTED_EXPERTS_DTYPES = { "int8": torch.int8, "int16": torch.int16, @@ -70,6 +70,54 @@ DEFAULT_THINKING_TAGS = ["", ""] +class NemoGymCompatibleConfig(Protocol): + """Configuration fields required to select the NeMo Gym rollout path.""" + + @property + def env(self) -> dict[str, Any]: ... + + @property + def policy(self) -> PolicyConfig: ... + + +def should_use_nemo_gym(master_config: NemoGymCompatibleConfig) -> bool: + """Determine whether NeMo Gym should handle rollouts and validation.""" + should_use_gym = bool(master_config.env.get("should_use_nemo_gym")) + if not should_use_gym: + return False + + generation_config = master_config.policy["generation"] + assert should_use_async_rollouts(generation_config), ( + "❌ Error: In order to use NeMo-Gym, you must use a generation " + "backend with `async_engine: true`!" + ) + + if generation_config["backend"] == "vllm": + should_expose_http_server = generation_config.get("vllm_cfg", {}).get( + "expose_http_server" + ) + elif generation_config["backend"] == "megatron": + should_expose_http_server = generation_config.get( + "mcore_generation_config", {} + ).get("expose_http_server") + elif generation_config["backend"] == "trtllm": + should_expose_http_server = generation_config.get("trtllm_cfg", {}).get( + "expose_http_server" + ) + elif generation_config["backend"] == "dynamo": + should_expose_http_server = generation_config.get("vllm_cfg", {}).get( + "expose_http_server" + ) + else: + should_expose_http_server = False + assert should_expose_http_server, ( + "In order to use NeMo-Gym, you must expose the generation server via " + "`expose_http_server: true`!" + ) + + return True + + def _has_nan_generation_logprobs(result: dict) -> bool: """Return whether a postprocessed rollout contains NaN policy logprobs.""" return any( diff --git a/nemo_rl/experience/sync_rollout_actor.py b/nemo_rl/experience/sync_rollout_actor.py index 1bc5c6c391e..9f7f7e9ea03 100644 --- a/nemo_rl/experience/sync_rollout_actor.py +++ b/nemo_rl/experience/sync_rollout_actor.py @@ -210,16 +210,14 @@ def rollout_to_tq( uses for compute (rewards, masks, lengths, prompt_ids_for_adv, …) β€” stays on the driver, never crosses an actor boundary. """ - # Lazy imports β€” avoid pulling grpo into this module at load. - from nemo_rl.algorithms.grpo import ( - _should_use_async_rollouts, - _should_use_nemo_gym, - ) + # Lazy imports keep rollout-specific dependencies off the actor startup path. from nemo_rl.algorithms.utils import get_gdpo_reward_component_keys from nemo_rl.data.llm_message_utils import ( MESSAGE_LOG_BULK_FIELDS, decompose_message_log, ) + from nemo_rl.environments.nemo_gym import should_use_nemo_gym + from nemo_rl.models.generation.interfaces import should_use_async_rollouts # Per-step generation-side metric hooks: snapshot once on the # first DS iter so backends with per-step deltas have a stable @@ -245,7 +243,7 @@ def rollout_to_tq( ) # Rollout dispatch (mirrors grpo_sync.py:294-349). - if _should_use_nemo_gym(cfg): + if should_use_nemo_gym(cfg): r = run_nemo_gym_rollout_sync( **common, max_seq_len=None, @@ -270,7 +268,7 @@ def rollout_to_tq( else: runner = ( run_async_multi_turn_rollout - if _should_use_async_rollouts(cfg) + if should_use_async_rollouts(cfg.policy["generation"]) else run_multi_turn_rollout ) final_batch, rollout_metrics = runner( diff --git a/nemo_rl/models/generation/interfaces.py b/nemo_rl/models/generation/interfaces.py index f8c7d25ca02..fa4f490bf68 100644 --- a/nemo_rl/models/generation/interfaces.py +++ b/nemo_rl/models/generation/interfaces.py @@ -226,6 +226,41 @@ class GenerationConfig(TypedDict): _debug_payload_metrics: NotRequired[bool] +def should_use_async_rollouts( + generation_config: GenerationConfig | None, +) -> bool: + """Determine whether a generation backend uses asynchronous rollouts.""" + if generation_config is None: + return False + backend = generation_config.get("backend", "") + + if backend == "dynamo": + return True + + if backend == "sglang": + return bool(generation_config.get("use_async_rollouts", False)) + + if backend == "vllm": + return bool(generation_config.get("vllm_cfg", {}).get("async_engine", False)) + + if backend == "trtllm": + assert generation_config.get("trtllm_cfg", {}).get("async_engine", False), ( + "TRT-LLM backend requires trtllm_cfg.async_engine=true; the " + "synchronous engine path (async_engine=false) is no longer supported." + ) + return True + + if backend == "megatron": + mcore_cfg = generation_config.get("mcore_generation_config", {}) + assert mcore_cfg.get("async_engine") is None, ( + "Megatron Inference always uses the async engine. The parameter " + "policy.generation.mcore_generation_config.async_engine was removed." + ) + return True + + return False + + @dataclass class GenerationSamplingParams: """Sampling profile threaded explicitly through rollout entry points. diff --git a/tests/functional/L1_Functional_Tests_PPO.sh b/tests/functional/L1_Functional_Tests_PPO.sh index c330c6b07ec..45893e9cd86 100755 --- a/tests/functional/L1_Functional_Tests_PPO.sh +++ b/tests/functional/L1_Functional_Tests_PPO.sh @@ -38,6 +38,7 @@ run_test fast uv run --no-sync bash ./tests/functional/ppo_automodel.sh run_test fast uv run --no-sync bash ./tests/functional/ppo_megatron.sh run_test fast uv run --no-sync bash ./tests/functional/ppo_non_colocated.sh run_test fast uv run --no-sync bash ./tests/functional/ppo_megatron_non_colocated.sh +run_test fast uv run --no-sync bash ./tests/functional/ppo_async_megatron.sh cd ${PROJECT_ROOT}/tests if compgen -G ".coverage*" > /dev/null; then diff --git a/tests/functional/ppo_async_megatron.sh b/tests/functional/ppo_async_megatron.sh new file mode 100644 index 00000000000..6d53c48e7bc --- /dev/null +++ b/tests/functional/ppo_async_megatron.sh @@ -0,0 +1,108 @@ +#!/bin/bash + +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +PROJECT_ROOT=$(realpath "${SCRIPT_DIR}/../..") +EXP_NAME=$(basename "$0" .sh) +EXP_DIR="${SCRIPT_DIR}/${EXP_NAME}" +CKPT_DIR="${EXP_DIR}/checkpoints" +export PYTHONPATH="${PROJECT_ROOT}:${PYTHONPATH:-}" + +rm -rf "${EXP_DIR}" +mkdir -p "${EXP_DIR}" + +TRAIN_CMD=( + uv run coverage run -a + --data-file="${PROJECT_ROOT}/tests/.coverage" + --source="${PROJECT_ROOT}/nemo_rl" + "${PROJECT_ROOT}/examples/run_ppo.py" + --config "${PROJECT_ROOT}/examples/configs/ppo_math_1B_megatron.yaml" + policy.model_name=Qwen/Qwen2.5-0.5B + value.model_name=Qwen/Qwen2.5-0.5B + ppo.num_prompts_per_step=2 + ppo.num_generations_per_prompt=4 + ppo.ppo_epochs=2 + ppo.max_num_epochs=-1 + ppo.policy_training_start_step=1 + ppo.val_at_start=false + ppo.val_period=0 + ppo.val_at_end=true + ppo.max_val_samples=8 + ppo.val_batch_size=8 + ppo.reward_scaling.enabled=false + ppo.reward_shaping.enabled=false + ppo.seq_logprob_error_threshold=1000 + ppo.async_ppo.enabled=true + ppo.async_ppo.max_trajectory_age_steps=1 + ppo.async_ppo.warmup_generation_lead_steps=2 + policy.train_global_batch_size=4 + policy.logprob_batch_size=4 + policy.train_micro_batch_size=1 + +policy.megatron_cfg.scheduler.override_opt_param_scheduler=true + policy.generation.colocated.enabled=false + policy.generation.colocated.resources.gpus_per_node=1 + policy.generation.colocated.resources.num_nodes=1 + policy.generation.vllm_cfg.async_engine=true + loss_fn.use_importance_sampling_correction=true + value.train_global_batch_size=4 + value.train_micro_batch_size=1 + data.use_multiple_dataloader=false + +value.megatron_cfg.scheduler.override_opt_param_scheduler=true + cluster.gpus_per_node=2 + logger.tensorboard_enabled=true + logger.wandb_enabled=false + logger.monitor_gpus=true + checkpointing.enabled=true + checkpointing.checkpoint_dir="${CKPT_DIR}" + checkpointing.metric_name=null + checkpointing.save_period=1 +) + +cd "${PROJECT_ROOT}" + +"${TRAIN_CMD[@]}" \ + ppo.max_num_steps=2 \ + logger.log_dir="${EXP_DIR}/logs_run1" \ + "$@" \ + 2>&1 | tee "${EXP_DIR}/run1.log" + +grep -q "Separate PPO clusters initialized" "${EXP_DIR}/run1.log" +grep -q "Updated generation window: version=0, lead=2, max_age=2" "${EXP_DIR}/run1.log" +grep -q "Updated generation window: version=1, lead=1, max_age=2" "${EXP_DIR}/run1.log" +test "$(grep -c "PPO epoch 2/2" "${EXP_DIR}/run1.log")" -eq 2 +test -f "${CKPT_DIR}/step_1/replay_buffer.pt" +test -f "${CKPT_DIR}/step_2/replay_buffer.pt" + +"${TRAIN_CMD[@]}" \ + ppo.max_num_steps=4 \ + logger.log_dir="${EXP_DIR}/logs_run2" \ + "$@" \ + 2>&1 | tee "${EXP_DIR}/run2.log" + +grep -q "Restoring replay buffer from checkpoint" "${EXP_DIR}/run2.log" +grep -q "ReplayBuffer restored:" "${EXP_DIR}/run2.log" +grep -q "Updated generation window: version=3, lead=1, max_age=1" "${EXP_DIR}/run2.log" +test "$(grep -c "PPO epoch 2/2" "${EXP_DIR}/run2.log")" -eq 2 +test -d "${CKPT_DIR}/step_4/policy/weights" +test -d "${CKPT_DIR}/step_4/value/weights" + +for run_spec in "run1 1 1" "run2 2 2"; do + read -r run expected_policy_steps expected_max_age <<< "${run_spec}" + metrics="${EXP_DIR}/metrics_${run}.json" + uv run tests/json_dump_tb_logs.py "${EXP_DIR}/logs_${run}" \ + --output_path "${metrics}" + uv run tests/check_metrics.py "${metrics}" \ + 'len(data["train/reward"]) == 2' \ + "len(data[\"train/loss\"]) == ${expected_policy_steps}" \ + 'len(data["train/critic/loss"]) == 2' \ + 'min(data["train/probs_ratio_clamped_min"]) > 0.79' \ + 'max(data["train/probs_ratio_clamped_min"]) < 1.21' \ + 'min(data["train/probs_ratio_clamped_max"]) > 0.79' \ + 'max(data["train/probs_ratio_clamped_max"]) < 1.29' \ + 'max(data["train/token_mult_prob_error"]) < 1.05' \ + 'max(data["train/critic/loss"]) < 6.0' \ + 'min(data["train/critic/loss"]) >= 0' \ + "max(data[\"train/avg_trajectory_age\"]) <= ${expected_max_age}" \ + 'len(data["validation/accuracy"]) == 1' +done diff --git a/tests/functional/ppo_non_colocated.sh b/tests/functional/ppo_non_colocated.sh index 82fd897c052..9c7692728ba 100755 --- a/tests/functional/ppo_non_colocated.sh +++ b/tests/functional/ppo_non_colocated.sh @@ -57,7 +57,7 @@ uv run tests/check_metrics.py $JSON_METRICS \ 'max(data["train/probs_ratio_clamped_min"]) < 1.21' \ 'min(data["train/probs_ratio_clamped_max"]) > 0.79' \ 'max(data["train/probs_ratio_clamped_max"]) < 1.29' \ - 'max(data["train/critic/loss"]) < 6.0' \ + 'max(data["train/critic/loss"]) < 8.0' \ 'min(data["train/critic/loss"]) >= 0' \ 'max(data["train/critic/explained_var"]) <= 1.0001' \ 'max(data["train/critic/grad_norm"]) < 350' diff --git a/tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.sh b/tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.sh new file mode 100755 index 00000000000..673ede47b23 --- /dev/null +++ b/tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.sh @@ -0,0 +1,51 @@ +#!/bin/bash +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd) +source $SCRIPT_DIR/common.env + +# ===== BEGIN CONFIG ===== +NUM_NODES=1 +STEPS_PER_RUN=40 +MAX_STEPS=40 +NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) # Round up +NUM_MINUTES=60 +# ===== END CONFIG ===== + +exit_if_max_steps_reached + +# Run the experiment +cd $PROJECT_ROOT +uv run examples/run_ppo.py \ + --config $CONFIG_PATH \ + ppo.max_num_steps=$MAX_STEPS \ + logger.log_dir=$LOG_DIR \ + logger.wandb_enabled=True \ + logger.wandb.project=nemo-rl \ + logger.wandb.name=$EXP_NAME \ + logger.monitor_gpus=True \ + logger.tensorboard_enabled=True \ + checkpointing.enabled=True \ + checkpointing.checkpoint_dir=$CKPT_DIR \ + $@ \ + 2>&1 | tee $RUN_LOG + +# Verify that policy training and generation use separate clusters. +grep -q "Separate PPO clusters initialized" $RUN_LOG +grep -q "collective communication" $RUN_LOG + +# Convert tensorboard logs to json +uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS + +# Only run metrics if the target step is reached +if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | map(tonumber) | max' $JSON_METRICS) -ge $MAX_STEPS ]]; then + uv run tests/check_metrics.py $JSON_METRICS \ + 'median(data["train/token_mult_prob_error"]) < 1.1' \ + 'data["train/token_mult_prob_error"]["40"] < 1.1' \ + 'median(data["train/max_seq_mult_prob_error"]) < 1.2' \ + 'max(data["train/avg_trajectory_age"]) <= 2.0' \ + 'data["train/avg_trajectory_age"]["40"] <= 1.0' \ + 'data["train/reward"]["40"] > 0.75' \ + 'data["validation/accuracy"]["40"] > 0.65' + + # Clean up checkpoint directory after successful run to save space. + rm -rf "$CKPT_DIR" +fi diff --git a/tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.sh b/tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.sh new file mode 100755 index 00000000000..4babe4e1f06 --- /dev/null +++ b/tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.sh @@ -0,0 +1,50 @@ +#!/bin/bash +SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd) +source $SCRIPT_DIR/common.env + +# ===== BEGIN CONFIG ===== +NUM_NODES=2 +STEPS_PER_RUN=40 +MAX_STEPS=40 +NUM_RUNS=$(( (MAX_STEPS + STEPS_PER_RUN - 1) / STEPS_PER_RUN )) # Round up +NUM_MINUTES=60 +# ===== END CONFIG ===== + +exit_if_max_steps_reached + +# Run the experiment +cd $PROJECT_ROOT +uv run examples/run_ppo.py \ + --config $CONFIG_PATH \ + ppo.max_num_steps=$MAX_STEPS \ + logger.log_dir=$LOG_DIR \ + logger.wandb_enabled=True \ + logger.wandb.project=nemo-rl \ + logger.wandb.name=$EXP_NAME \ + logger.monitor_gpus=True \ + logger.tensorboard_enabled=True \ + checkpointing.enabled=True \ + checkpointing.checkpoint_dir=$CKPT_DIR \ + $@ \ + 2>&1 | tee $RUN_LOG + +# Verify that policy training and generation use separate clusters. +grep -q "Separate PPO clusters initialized" $RUN_LOG +grep -q "collective communication" $RUN_LOG + +# Convert tensorboard logs to json +uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS + +# Only run metrics if the target step is reached +if [[ $(jq 'to_entries | .[] | select(.key == "train/loss") | .value | keys | map(tonumber) | max' $JSON_METRICS) -ge $MAX_STEPS ]]; then + uv run tests/check_metrics.py $JSON_METRICS \ + 'median(data["train/token_mult_prob_error"]) < 1.1' \ + 'data["train/token_mult_prob_error"]["40"] < 1.1' \ + 'median(data["train/max_seq_mult_prob_error"]) < 1.2' \ + 'max(data["train/avg_trajectory_age"]) <= 1.0' \ + 'data["train/reward"]["40"] > 0.75' \ + 'data["validation/accuracy"]["40"] > 0.65' + + # Clean up checkpoint directory after successful run to save space. + rm -rf "$CKPT_DIR" +fi diff --git a/tests/test_suites/nightly.txt b/tests/test_suites/nightly.txt index cfde74908d1..001c95235cd 100644 --- a/tests/test_suites/nightly.txt +++ b/tests/test_suites/nightly.txt @@ -283,12 +283,18 @@ tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-valuetp2sp.sh # Non-colocated AutoModel PPO with vLLM on a dedicated GPU split (Qwen2.5-1.5B, GSM8K) tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated.sh +# Asynchronous non-colocated AutoModel PPO with one-step-old rollouts (Qwen2.5-1.5B, GSM8K) +tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-automodel-noncolocated-async.sh + # Megatron PPO with value sequence parallelism (TP2+SP) and dynamic batching (Qwen2.5-1.5B, GSM8K) tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-megatron-valuetp2sp-dynbatch.sh # Non-colocated Megatron PPO with dedicated training and vLLM nodes (Qwen2.5-1.5B, GSM8K) tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated.sh +# Asynchronous non-colocated Megatron PPO with cross-node refit (Qwen2.5-1.5B, GSM8K) +tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-2n8g-megatron-valuetp2sp-dynbatch-noncolocated-async.sh + # Megatron PPO with value SP + pipeline + context parallelism + sequence packing (TP2+SP+PP2+CP2) (Qwen2.5-1.5B, GSM8K) tests/test_suites/llm/ppo-qwen2.5-1.5b-gsm8k-1n8g-megatron-valuetp2sp-pp2cp2-pack.sh diff --git a/tests/unit/algorithms/test_async_utils.py b/tests/unit/algorithms/test_async_utils.py index bf1f320f084..33cd5274430 100644 --- a/tests/unit/algorithms/test_async_utils.py +++ b/tests/unit/algorithms/test_async_utils.py @@ -32,7 +32,6 @@ os.environ["TMPDIR"] = _temp_dir # System temp dir import nemo_rl.algorithms.async_utils.trajectory_collector as trajectory_collector_mod -import nemo_rl.algorithms.grpo as grpo_mod from nemo_rl.algorithms.async_utils import ( AsyncTrajectoryCollector, ReplayBuffer, @@ -130,7 +129,7 @@ def _state( } def test_local_restore_prepares_current_step_for_gap_fill(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) state = self._state( trajectory_versions=[0, 1, 1, 2], target_weight_versions=[1, 2, 2, 3], @@ -151,8 +150,40 @@ def test_local_restore_prepares_current_step_for_gap_fill(self): assert buffer.get_trajectories_needed(2, 2) == 0 assert buffer.get_trajectories_needed(3, 2) == 1 + @pytest.mark.parametrize( + ("drop_incomplete_targets_on_restore", "expected_targets", "expected_needed"), + [ + (True, [2, 2], 2), + (False, [2, 2, 3], 1), + ], + ) + def test_local_restore_can_drop_incomplete_frontier( + self, + drop_incomplete_targets_on_restore, + expected_targets, + expected_needed, + ): + buffer = ReplayBufferImpl( + max_size=10, + drop_incomplete_targets_on_restore=drop_incomplete_targets_on_restore, + ) + state = self._state( + trajectory_versions=[1, 1, 2], + target_weight_versions=[2, 2, 3], + last_target_weight_already_generated=3, + ) + + buffer.load_state_dict( + state, + num_prompts_per_step=2, + current_training_step=2, + ) + + assert buffer.get_debug_info()["target_weight_versions"] == expected_targets + assert buffer.get_trajectories_needed(3, 2) == expected_needed + def test_local_restore_empty_state_resets_generation_watermark(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) state = self._state( trajectory_versions=[], target_weight_versions=[], @@ -170,7 +201,7 @@ def test_local_restore_empty_state_resets_generation_watermark(self): assert buffer.get_trajectories_needed(5, 2) == 2 def test_local_restore_removes_stale_trajectories(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) state = self._state( trajectory_versions=[0, 1, 4], target_weight_versions=[5, 5, 5], @@ -189,8 +220,42 @@ def test_local_restore_removes_stale_trajectories(self): assert not buffer.has_complete_batch(5, 2) assert buffer.get_trajectories_needed(5, 2) == 1 + @pytest.mark.parametrize( + ("drop_incomplete_targets_on_restore", "expected_targets", "expected_needed"), + [ + (True, [], 2), + (False, [5], 1), + ], + ) + def test_local_restore_drops_target_made_incomplete_by_stale_filter( + self, + drop_incomplete_targets_on_restore, + expected_targets, + expected_needed, + ): + buffer = ReplayBufferImpl( + max_size=10, + drop_incomplete_targets_on_restore=drop_incomplete_targets_on_restore, + ) + state = self._state( + trajectory_versions=[4, 1], + target_weight_versions=[5, 5], + last_target_weight_already_generated=5, + ) + + buffer.load_state_dict( + state, + num_prompts_per_step=2, + current_training_step=5, + max_age_steps=1, + ) + + assert buffer.get_debug_info()["target_weight_versions"] == expected_targets + assert buffer.size() == len(expected_targets) + assert buffer.get_trajectories_needed(5, 2) == expected_needed + def test_local_restore_truncates_after_resume_cleanup(self): - buffer = ReplayBufferImpl(max_size=2) + buffer = ReplayBufferImpl(max_size=2, drop_incomplete_targets_on_restore=False) state = self._state( trajectory_versions=[0, 1, 2, 3], target_weight_versions=[1, 2, 3, 4], @@ -208,7 +273,7 @@ def test_local_restore_truncates_after_resume_cleanup(self): assert buffer.get_debug_info()["target_weight_versions"] == [2, 3] def test_local_restore_without_current_step_rechecks_after_stale_removal(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) state = self._state( trajectory_versions=[0, 4, 4], target_weight_versions=[5, 5, 6], @@ -225,7 +290,7 @@ def test_local_restore_without_current_step_rechecks_after_stale_removal(self): assert buffer.get_last_target_weight_already_generated() == -1 def test_local_sample_evicts_stale_restored_trajectories(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) assert ( buffer.add( {"batch": {"data": "stale"}, "rollout_metrics": {}}, @@ -254,7 +319,7 @@ def test_local_sample_evicts_stale_restored_trajectories(self): assert buffer.size() == 0 def test_local_debug_info_reports_starvation_diagnostics(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) assert ( buffer.add( { @@ -308,7 +373,7 @@ def test_local_debug_info_reports_starvation_diagnostics(self): assert diagnostics["num_trajectories_sampled"] == 2 def test_local_load_state_dict_validates_checkpoint_shape(self): - buffer = ReplayBufferImpl(max_size=10) + buffer = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) with pytest.raises(ValueError, match="Checkpoint missing required keys"): buffer.load_state_dict( @@ -336,7 +401,7 @@ def test_local_actor_side_checkpoint_preserves_compact_media_and_resume_metadata torch.tensor([[1.0, 2.0]]), dim_to_pack=0 ).enable_deduplication() compact_media = PackedTensor.concat([compact_media_row] * 2) - source = ReplayBufferImpl(max_size=10) + source = ReplayBufferImpl(max_size=10, drop_incomplete_targets_on_restore=False) assert ( source.add( { @@ -364,7 +429,9 @@ def test_local_actor_side_checkpoint_preserves_compact_media_and_resume_metadata assert source.save_to_path(str(checkpoint_path)) == 2 - restored = ReplayBufferImpl(max_size=10) + restored = ReplayBufferImpl( + max_size=10, drop_incomplete_targets_on_restore=False + ) metadata = restored.load_from_path( str(checkpoint_path), num_prompts_per_step=1, @@ -394,7 +461,9 @@ class TestReplayBuffer: def test_replay_buffer_initialization(self): """Test ReplayBuffer initialization.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) size = ray.get(buffer.size.remote()) assert size == 0 @@ -408,7 +477,9 @@ def test_replay_buffer_initialization(self): def test_replay_buffer_push_and_size(self): """Test pushing trajectories to buffer.""" - buffer = ReplayBuffer.remote(max_size=3) + buffer = ReplayBuffer.remote( + max_size=3, drop_incomplete_targets_on_restore=False + ) # Create mock trajectories trajectory1 = {"batch": {"data": "test1"}, "rollout_metrics": {"reward": 1.0}} @@ -439,7 +510,9 @@ def test_replay_buffer_push_and_size(self): def test_replay_buffer_max_size_limit(self): """Test that buffer respects max size limit.""" - buffer = ReplayBuffer.remote(max_size=2) + buffer = ReplayBuffer.remote( + max_size=2, drop_incomplete_targets_on_restore=False + ) # Fill buffer to capacity trajectory1 = {"batch": {"data": "test1"}, "rollout_metrics": {"reward": 1.0}} @@ -470,7 +543,9 @@ def test_replay_buffer_max_size_limit(self): def test_replay_buffer_sampling_basic(self): """Test basic trajectory sampling.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) # Push trajectories with different weight versions trajectories = [] @@ -509,7 +584,9 @@ def test_replay_buffer_sampling_basic(self): def test_replay_buffer_sampling_insufficient_trajectories(self): """Test sampling when insufficient trajectories are available.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) # Push only one trajectory trajectory = {"batch": {"data": "test"}, "rollout_metrics": {"reward": 1.0}} @@ -533,7 +610,9 @@ def test_replay_buffer_sampling_insufficient_trajectories(self): def test_replay_buffer_watermark_advances_only_after_consumption(self): """Test buffering alone does not mark a target as consumed.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) trajectory1 = {"batch": {"data": "test1"}, "rollout_metrics": {}} trajectory2 = {"batch": {"data": "test2"}, "rollout_metrics": {}} @@ -571,7 +650,9 @@ def test_replay_buffer_watermark_advances_only_after_consumption(self): def test_replay_buffer_starvation_diagnostics_nemo_gym_turn_keys(self): """NeMo Gym uses turns_per_sample/* in rollout_metrics; diagnostics must read them.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) t1 = { "batch": {"data": "a"}, "rollout_metrics": { @@ -607,7 +688,9 @@ def test_replay_buffer_starvation_diagnostics_nemo_gym_turn_keys(self): def test_replay_buffer_age_filtering(self): """Test that old trajectories are evicted.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) # Push trajectories with different ages old_trajectory = {"batch": {"data": "old"}, "rollout_metrics": {"reward": 1.0}} @@ -644,7 +727,9 @@ def test_replay_buffer_age_filtering(self): def test_replay_buffer_target_weight_matching(self): """Test that sampling only returns trajectories intended for current step.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) # Push trajectories intended for different target steps trajectory1 = { @@ -680,7 +765,9 @@ def test_replay_buffer_target_weight_matching(self): def test_replay_buffer_get_existing_target_weights(self): """Test getting existing target weight versions.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) # Initially empty existing_weights = ray.get(buffer.get_existing_target_weights.remote()) @@ -704,7 +791,9 @@ def test_replay_buffer_get_existing_target_weights(self): def test_replay_buffer_clear(self): """Test clearing the buffer.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) # Push some trajectories trajectory = {"batch": {"data": "test"}, "rollout_metrics": {"reward": 1.0}} @@ -732,7 +821,9 @@ def test_replay_buffer_clear(self): def test_replay_buffer_state_dict(self): """Test state_dict serialization for checkpointing.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) trajectory1 = {"batch": {"data": "test1"}, "rollout_metrics": {"reward": 1.0}} trajectory2 = {"batch": {"data": "test2"}, "rollout_metrics": {"reward": 2.0}} @@ -757,7 +848,9 @@ def test_replay_buffer_state_dict(self): def test_replay_buffer_load_state_dict(self): """Test load_state_dict restoration from checkpoint.""" - buffer1 = ReplayBuffer.remote(max_size=10) + buffer1 = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) trajectory1 = {"batch": {"data": "test1"}, "rollout_metrics": {"reward": 1.0}} trajectory2 = {"batch": {"data": "test2"}, "rollout_metrics": {"reward": 2.0}} @@ -772,7 +865,9 @@ def test_replay_buffer_load_state_dict(self): state = ray.get(buffer1.state_dict.remote()) ray.kill(buffer1) - buffer2 = ReplayBuffer.remote(max_size=10) + buffer2 = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) assert ray.get(buffer2.size.remote()) == 0 ray.get(buffer2.load_state_dict.remote(state)) @@ -787,7 +882,9 @@ def test_replay_buffer_load_state_dict(self): def test_replay_buffer_state_dict_round_trip_sampling(self): """Test save/restore preserves sampling behavior.""" - buffer1 = ReplayBuffer.remote(max_size=10) + buffer1 = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) for i in range(3): trajectory = { @@ -803,7 +900,9 @@ def test_replay_buffer_state_dict_round_trip_sampling(self): state = ray.get(buffer1.state_dict.remote()) ray.kill(buffer1) - buffer2 = ReplayBuffer.remote(max_size=10) + buffer2 = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) ray.get(buffer2.load_state_dict.remote(state)) sample_result = ray.get( @@ -821,7 +920,9 @@ def test_replay_buffer_state_dict_round_trip_sampling(self): def test_replay_buffer_load_state_dict_max_size_change(self): """Test load_state_dict truncates after resume cleanup.""" - buffer1 = ReplayBuffer.remote(max_size=5) + buffer1 = ReplayBuffer.remote( + max_size=5, drop_incomplete_targets_on_restore=False + ) for i in range(4): trajectory = { @@ -837,7 +938,9 @@ def test_replay_buffer_load_state_dict_max_size_change(self): state = ray.get(buffer1.state_dict.remote()) ray.kill(buffer1) - buffer2 = ReplayBuffer.remote(max_size=2) + buffer2 = ReplayBuffer.remote( + max_size=2, drop_incomplete_targets_on_restore=False + ) ray.get( buffer2.load_state_dict.remote( state, @@ -855,7 +958,9 @@ def test_replay_buffer_load_state_dict_max_size_change(self): def test_replay_buffer_load_empty_state_resets_generation_watermark(self): """Test empty restore can generate from the current step.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) state = { "trajectories": [], @@ -881,7 +986,9 @@ def test_replay_buffer_load_empty_state_resets_generation_watermark(self): def test_replay_buffer_restore_removes_stale_trajectories(self): """Test stale restored trajectories do not make a step look complete.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) state = { "trajectories": [ @@ -914,7 +1021,9 @@ def test_replay_buffer_restore_removes_stale_trajectories(self): def test_replay_buffer_readiness_ignores_stale_trajectories(self): """Test readiness helpers match sample's age-window filtering.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) for version in [0, 1, 4]: ray.get( @@ -946,7 +1055,9 @@ def test_replay_buffer_readiness_ignores_stale_trajectories(self): def test_replay_buffer_load_state_dict_missing_keys(self): """Test load_state_dict raises for missing required keys.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) incomplete_state = { "trajectories": [], @@ -960,7 +1071,9 @@ def test_replay_buffer_load_state_dict_missing_keys(self): def test_replay_buffer_load_state_dict_inconsistent_lengths(self): """Test load_state_dict raises for inconsistent parallel lists.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) bad_state = { "trajectories": [{"batch": {"data": "test"}}], @@ -976,7 +1089,9 @@ def test_replay_buffer_load_state_dict_inconsistent_lengths(self): def test_replay_buffer_restore_for_training_step_gap_fill_accounting(self): """Test resume cleanup keeps incomplete future targets for gap filling.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) state = { "trajectories": [ @@ -1014,7 +1129,9 @@ def test_replay_buffer_remove_incomplete_resets_watermark_before_first_remaining self, ): """Test fallback cleanup does not skip gaps after removing partial targets.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) state = { "trajectories": [ @@ -1040,7 +1157,9 @@ def test_replay_buffer_remove_incomplete_resets_watermark_before_first_remaining def test_replay_buffer_checkpoint_with_torch_save(self, tmp_path): """Actor-side compact replay checkpoint survives a config flag flip.""" - buffer1 = ReplayBuffer.remote(max_size=10) + buffer1 = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) pixel_row = PackedTensor( torch.tensor([[1.0, 2.0]]), dim_to_pack=0 ).enable_deduplication() @@ -1064,7 +1183,9 @@ def test_replay_buffer_checkpoint_with_torch_save(self, tmp_path): ray.kill(buffer1) - buffer2 = ReplayBuffer.remote(max_size=10) + buffer2 = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) restore_metadata = ray.get(buffer2.load_from_path.remote(str(checkpoint_path))) assert restore_metadata == { @@ -1088,11 +1209,11 @@ def test_replay_buffer_checkpoint_with_torch_save(self, tmp_path): ray.kill(buffer2) def test_resume_deadlock_precondition_detectable(self): - """Regression: restored buffer can expose the async-GRPO resume deadlock. + """Regression: restored buffer can expose an async resume deadlock. After PR #2651 introduced replay-buffer checkpointing, resuming from a - checkpoint where target N is complete but target N+1 is absent caused an - async-GRPO deadlock: + checkpoint where target N is complete but target N+1 is absent can + deadlock Async GRPO or Async PPO: 1. Startup wait sees has_complete_batch(N) == True and breaks immediately. 2. Training consumes all target-N trajectories and triggers a refit. @@ -1110,7 +1231,9 @@ def test_resume_deadlock_precondition_detectable(self): max_age = 1 # Build a pre-checkpoint buffer: 8 trajectories for target 30, none for 31. - buffer1 = ReplayBuffer.remote(max_size=20) + buffer1 = ReplayBuffer.remote( + max_size=20, drop_incomplete_targets_on_restore=False + ) for _ in range(num_prompts): ray.get( buffer1.add.remote( @@ -1124,7 +1247,9 @@ def test_resume_deadlock_precondition_detectable(self): ray.kill(buffer1) # Restore at step 30, simulating a checkpoint resume. - buffer2 = ReplayBuffer.remote(max_size=20) + buffer2 = ReplayBuffer.remote( + max_size=20, drop_incomplete_targets_on_restore=False + ) ray.get( buffer2.load_state_dict.remote( state, @@ -1192,7 +1317,6 @@ def _prime_collection_loop(self, collector): collector.running = True def test_collection_loop_marks_data_exhausted_on_natural_completion(self): - """for...else path: iterator drains cleanly -> data_exhausted, not errored.""" collector = self.create_local_collector() self._prime_collection_loop(collector) processed = [] @@ -1251,6 +1375,7 @@ def _boom(batch): assert collector.data_exhausted is False status = collector.get_status() assert status["errored"] is True + assert status["error"] == "RuntimeError: collection blew up" assert status["data_exhausted"] is False assert status["running"] is False @@ -1296,6 +1421,77 @@ def create_mock_config(self) -> MasterConfig: }, ) + def test_collector_selects_ppo_config(self): + """The shared collector derives PPO settings from its master config.""" + from nemo_rl.algorithms.ppo import ( + AsyncPPOConfig, + PPOConfig, + ) + from nemo_rl.algorithms.ppo import ( + MasterConfig as PPOMasterConfig, + ) + + async_config = AsyncPPOConfig( + max_trajectory_age_steps=3, + warmup_generation_lead_steps=5, + ) + master_config = PPOMasterConfig.model_construct( + policy={"make_sequence_length_divisible_by": 1}, + ppo=PPOConfig.model_construct( + num_prompts_per_step=2, + num_generations_per_prompt=4, + max_rollout_turns=1, + async_ppo=async_config, + ), + ) + collector_cls = AsyncTrajectoryCollector.__ray_metadata__.modified_class + collector = collector_cls( + policy_generation=MockGenerationInterface(), + tokenizer=mock.MagicMock(), + task_to_env={}, + master_config=master_config, + replay_buffer=mock.MagicMock(), + ) + + assert collector.algorithm_config is master_config.ppo + assert collector.async_config is async_config + assert collector.async_config.max_trajectory_age_steps == 3 + + collector.set_generation_window( + weight_version=2, + generation_lead_steps=3, + max_trajectory_age_steps=5, + ) + assert collector.current_weight_version == 2 + assert collector._generation_lead_steps == 3 + assert collector._max_trajectory_age_steps == 5 + assert collector._calculate_target_weights(2) == [3, 4, 5] + + def test_collector_grpo_window_remains_fixed(self): + collector = self.create_local_collector() + + assert collector.current_weight_version == 0 + assert collector._generation_lead_steps == 2 + assert collector._max_trajectory_age_steps == 2 + assert collector._calculate_target_weights(0) == [0, 1, 2] + + collector.set_weight_version(5) + + assert collector.current_weight_version == 5 + assert collector._generation_lead_steps == 2 + assert collector._max_trajectory_age_steps == 2 + assert collector._calculate_target_weights(5) == [6, 7] + + def test_collector_rejects_generation_lead_above_validity_age(self): + collector = self.create_local_collector() + + with pytest.raises(ValueError, match="max_trajectory_age_steps"): + collector.set_generation_window( + weight_version=1, + generation_lead_steps=3, + max_trajectory_age_steps=2, + ) + def create_mock_batch(self, size: int = 2) -> BatchedDataDict[DatumSpec]: """Create a mock batch for testing.""" message_logs = [] @@ -1317,7 +1513,9 @@ def create_mock_batch(self, size: int = 2) -> BatchedDataDict[DatumSpec]: def test_async_trajectory_collector_initialization(self): """Test AsyncTrajectoryCollector initialization.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -1343,7 +1541,9 @@ def test_async_trajectory_collector_initialization(self): def test_async_trajectory_collector_weight_version_updates(self): """Test weight version updates in trajectory collector.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -1370,7 +1570,9 @@ def test_async_trajectory_collector_weight_version_updates(self): def test_async_trajectory_collector_pause_resume(self): """Test pause and resume functionality.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -1396,7 +1598,9 @@ def test_async_trajectory_collector_pause_resume(self): def test_async_trajectory_collector_prepare_for_refit(self): """Test prepare for refit functionality.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -1474,7 +1678,9 @@ def test_dynamo_prepare_for_refit_drains_pending_generations(self): def test_calculate_target_weights(self): """Test target weight calculation logic.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -1664,7 +1870,9 @@ async def capture_batch(**kwargs): collector._get_next_target_for_generation = reserve_target collector._run_rollout_batch_worker = capture_batch monkeypatch.setattr(trajectory_collector_mod.ray, "get", lambda value: value) - monkeypatch.setattr(grpo_mod, "_should_use_nemo_gym", lambda config: True) + monkeypatch.setattr( + trajectory_collector_mod, "should_use_nemo_gym", lambda config: True + ) monkeypatch.setattr( trajectory_collector_mod._threading, "Thread", RecordingThread ) @@ -1723,7 +1931,9 @@ async def capture_batch(**kwargs): collector._get_next_target_for_generation = reserve_target collector._run_rollout_batch_worker = capture_batch monkeypatch.setattr(trajectory_collector_mod.ray, "get", lambda value: value) - monkeypatch.setattr(grpo_mod, "_should_use_nemo_gym", lambda config: False) + monkeypatch.setattr( + trajectory_collector_mod, "should_use_nemo_gym", lambda config: False + ) monkeypatch.setattr( trajectory_collector_mod._threading, "Thread", RecordingThread ) @@ -1982,7 +2192,9 @@ def test_rollouts_state_retrieval(self): def test_dataloader_state_retrieval(self): """Test getting dataloader state for checkpointing.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -2229,7 +2441,9 @@ def create_mock_batch(self, size: int = 2) -> BatchedDataDict[DatumSpec]: def test_buffer_and_collector_integration(self): """Test that buffer and collector work together correctly.""" - buffer = ReplayBuffer.remote(max_size=10) + buffer = ReplayBuffer.remote( + max_size=10, drop_incomplete_targets_on_restore=False + ) mock_generation = MockGenerationInterface() mock_tokenizer = mock.MagicMock() mock_env = MockEnvironment.remote(rewards=[1.0, 2.0]) @@ -2263,7 +2477,9 @@ def test_buffer_and_collector_integration(self): def test_concurrent_operations(self): """Test that concurrent operations don't cause race conditions.""" - buffer = ReplayBuffer.remote(max_size=5) + buffer = ReplayBuffer.remote( + max_size=5, drop_incomplete_targets_on_restore=False + ) # Push trajectories concurrently from multiple threads def push_trajectory(buffer, trajectory_id): @@ -2308,11 +2524,15 @@ def test_error_handling(self): """Test error handling in async utilities.""" # Test with invalid buffer size with pytest.raises(Exception): - buffer = ReplayBuffer.remote(max_size=-1) + buffer = ReplayBuffer.remote( + max_size=-1, drop_incomplete_targets_on_restore=False + ) ray.get(buffer.size.remote()) # Test buffer operations - buffer = ReplayBuffer.remote(max_size=1) + buffer = ReplayBuffer.remote( + max_size=1, drop_incomplete_targets_on_restore=False + ) # Test sampling from empty buffer sample_result = ray.get( diff --git a/tests/unit/algorithms/test_distillation.py b/tests/unit/algorithms/test_distillation.py index 236acabe097..8fb307fd0fc 100644 --- a/tests/unit/algorithms/test_distillation.py +++ b/tests/unit/algorithms/test_distillation.py @@ -643,10 +643,10 @@ def capture_log(data, filename): return_value=(mock_batch, mock_rollout_metrics), ), patch( - "nemo_rl.algorithms.distillation._should_use_nemo_gym", return_value=False + "nemo_rl.algorithms.distillation.should_use_nemo_gym", return_value=False ), patch( - "nemo_rl.algorithms.distillation._should_use_async_rollouts", + "nemo_rl.algorithms.distillation.should_use_async_rollouts", return_value=False, ), patch("nemo_rl.algorithms.distillation.print_message_log_samples"), @@ -699,10 +699,10 @@ def test_validate_works_without_logger(mock_components): return_value=(mock_batch, mock_rollout_metrics), ), patch( - "nemo_rl.algorithms.distillation._should_use_nemo_gym", return_value=False + "nemo_rl.algorithms.distillation.should_use_nemo_gym", return_value=False ), patch( - "nemo_rl.algorithms.distillation._should_use_async_rollouts", + "nemo_rl.algorithms.distillation.should_use_async_rollouts", return_value=False, ), patch("nemo_rl.algorithms.distillation.print_message_log_samples"), @@ -1338,7 +1338,7 @@ def test_nemo_gym_distillation_runner_uses_setup_actor(): side_effect=lambda cfg, _: cfg, ), patch.object(runner, "setup_nemo_gym_config"), - patch.object(runner, "_should_use_nemo_gym", return_value=True), + patch.object(runner, "should_use_nemo_gym", return_value=True), patch.object(runner, "setup_response_data", return_value=(MagicMock(), None)), patch.object(runner, "init_ray"), patch.object(runner, "setup") as mock_setup, diff --git a/tests/unit/algorithms/test_grpo.py b/tests/unit/algorithms/test_grpo.py index ffc0aebd5b6..47f4ce294fc 100644 --- a/tests/unit/algorithms/test_grpo.py +++ b/tests/unit/algorithms/test_grpo.py @@ -46,8 +46,6 @@ _resolve_logprob_skip_flags, _resolve_message_level_advantage_penalties, _save_async_replay_buffer_checkpoint, - _should_use_async_rollouts, - _should_use_nemo_gym, _validate_multimodal_dedup_capability, _validate_use_kl_in_reward_compat, aggregate_rollout_metrics, @@ -74,10 +72,12 @@ EnvironmentInterface, EnvironmentReturn, ) +from nemo_rl.environments.nemo_gym import should_use_nemo_gym from nemo_rl.experience.interfaces import NEXT_NEMO_GYM_TASK_INDEX_KEY from nemo_rl.experience.rollouts import calculate_rewards from nemo_rl.models.generation import configure_generation_config from nemo_rl.models.generation.dynamo import DynamoConfig +from nemo_rl.models.generation.interfaces import should_use_async_rollouts from nemo_rl.models.generation.megatron import MegatronGeneration from nemo_rl.utils.config import load_config, register_omegaconf_resolvers from nemo_rl.utils.timer import Timer @@ -612,7 +612,7 @@ def test_apply_configured_message_level_advantage_penalties_noops_when_disabled( ] master_config = mock_grpo_components["master_config"] - with patch("nemo_rl.algorithms.grpo._should_use_nemo_gym") as should_use_nemo_gym: + with patch("nemo_rl.algorithms.grpo.should_use_nemo_gym") as should_use_nemo_gym: _apply_configured_message_level_advantage_penalties( train_data, message_logs, master_config ) @@ -648,7 +648,7 @@ def test_apply_configured_message_level_advantage_penalties_uses_config( master_config.grpo.malformed_thinking_advantage = -6.0 with patch( - "nemo_rl.algorithms.grpo._should_use_nemo_gym", return_value=True + "nemo_rl.algorithms.grpo.should_use_nemo_gym", return_value=True ) as should_use_nemo_gym: _apply_configured_message_level_advantage_penalties( train_data, message_logs, master_config, log_config=True @@ -669,7 +669,7 @@ def test_resolve_message_level_advantage_penalties_requires_nemo_gym( master_config = mock_grpo_components["master_config"] master_config.grpo.invalid_tool_call_advantage = -5.0 - with patch("nemo_rl.algorithms.grpo._should_use_nemo_gym", return_value=False): + with patch("nemo_rl.algorithms.grpo.should_use_nemo_gym", return_value=False): with pytest.raises(ValueError, match="NeMo-Gym path"): _resolve_message_level_advantage_penalties(master_config) @@ -1332,15 +1332,36 @@ def test_async_grpo_propagates_main_loop_collector_failure(mock_grpo_components) }, False, ), + ({"backend": "megatron", "mcore_generation_config": {}}, True), ], ) def test_should_use_async_rollouts_selects_backend_specific_config( generation_config, expected ): + assert should_use_async_rollouts(generation_config) is expected + + +def test_should_use_async_rollouts_rejects_removed_megatron_async_engine(): + with pytest.raises(AssertionError, match="async_engine was removed"): + should_use_async_rollouts( + { + "backend": "megatron", + "mcore_generation_config": {"async_engine": False}, + } + ) + + +def test_should_use_nemo_gym_accepts_megatron_always_async(): master_config = MagicMock() - master_config.policy = {"generation": generation_config} + master_config.env = {"should_use_nemo_gym": True} + master_config.policy = { + "generation": { + "backend": "megatron", + "mcore_generation_config": {"expose_http_server": True}, + } + } - assert _should_use_async_rollouts(master_config) is expected + assert should_use_nemo_gym(master_config) @pytest.mark.parametrize("backend", ["dynamo", "vllm"]) @@ -1460,10 +1481,10 @@ def test_should_use_nemo_gym_requires_dynamo_token_wrapper() -> None: } with pytest.raises(AssertionError, match="expose_http_server: true"): - _should_use_nemo_gym(master_config) + should_use_nemo_gym(master_config) master_config.policy["generation"]["vllm_cfg"]["expose_http_server"] = True - assert _should_use_nemo_gym(master_config) is True + assert should_use_nemo_gym(master_config) is True @contextmanager @@ -2834,7 +2855,7 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): fake_batched_message_log_to_flat_message, ) monkeypatch.setattr( - grpo_mod, "_should_use_async_rollouts", lambda *_args, **_kwargs: True + grpo_mod, "should_use_async_rollouts", lambda *_args, **_kwargs: True ) monkeypatch.setattr( grpo_mod, @@ -4396,10 +4417,10 @@ def capture_log(data, filename): with patch("nemo_rl.algorithms.grpo.run_multi_turn_rollout") as mock_rollout: mock_rollout.return_value = (mock_batch, mock_rollout_metrics) with patch( - "nemo_rl.algorithms.grpo._should_use_nemo_gym", return_value=False + "nemo_rl.algorithms.grpo.should_use_nemo_gym", return_value=False ): with patch( - "nemo_rl.algorithms.grpo._should_use_async_rollouts", + "nemo_rl.algorithms.grpo.should_use_async_rollouts", return_value=False, ): with patch("nemo_rl.algorithms.grpo.print_message_log_samples"): @@ -4473,10 +4494,10 @@ def test_validate_works_without_logger(self, mock_grpo_components): with patch("nemo_rl.algorithms.grpo.run_multi_turn_rollout") as mock_rollout: mock_rollout.return_value = (mock_batch, mock_rollout_metrics) with patch( - "nemo_rl.algorithms.grpo._should_use_nemo_gym", return_value=False + "nemo_rl.algorithms.grpo.should_use_nemo_gym", return_value=False ): with patch( - "nemo_rl.algorithms.grpo._should_use_async_rollouts", + "nemo_rl.algorithms.grpo.should_use_async_rollouts", return_value=False, ): with patch("nemo_rl.algorithms.grpo.print_message_log_samples"): @@ -4531,9 +4552,9 @@ def run_rollout(_policy, repeated_batch, *_args, **_kwargs): "nemo_rl.algorithms.grpo.run_multi_turn_rollout", side_effect=run_rollout, ), - patch("nemo_rl.algorithms.grpo._should_use_nemo_gym", return_value=False), + patch("nemo_rl.algorithms.grpo.should_use_nemo_gym", return_value=False), patch( - "nemo_rl.algorithms.grpo._should_use_async_rollouts", + "nemo_rl.algorithms.grpo.should_use_async_rollouts", return_value=False, ), patch("nemo_rl.algorithms.grpo.print_message_log_samples"), @@ -4597,7 +4618,7 @@ def run_gym_rollout(**kwargs): "nemo_rl.algorithms.grpo.run_nemo_gym_rollout_sync", side_effect=run_gym_rollout, ) as mock_rollout, - patch("nemo_rl.algorithms.grpo._should_use_nemo_gym", return_value=True), + patch("nemo_rl.algorithms.grpo.should_use_nemo_gym", return_value=True), patch("nemo_rl.algorithms.grpo.print_message_log_samples"), ): val_metrics, _ = validate( diff --git a/tests/unit/algorithms/test_ppo.py b/tests/unit/algorithms/test_ppo.py index 0d6a9c9388b..46bca0ec802 100644 --- a/tests/unit/algorithms/test_ppo.py +++ b/tests/unit/algorithms/test_ppo.py @@ -28,6 +28,7 @@ MseValueLossConfig, MseValueLossFn, ) +from nemo_rl.algorithms.ppo import PPOConfig from nemo_rl.algorithms.reward_functions import RewardShapingConfig from nemo_rl.data import DataConfig from nemo_rl.distributed.batched_data_dict import BatchedDataDict @@ -638,14 +639,14 @@ def test_create_advantage_estimator_gae(): # because the estimator accesses .use_kl_in_reward / .reference_policy_kl_* # as attributes, not dict keys. master_config = SimpleNamespace( - ppo={ - "adv_estimator": { + ppo=PPOConfig( + adv_estimator={ "name": "gae", **_make_gae_config( gae_lambda=0.95, gae_gamma=1.0, normalize_advantages=True ), }, - }, + ), loss_fn=_make_loss_config(kl_penalty=0.0), ) @@ -661,9 +662,9 @@ def test_create_advantage_estimator_raw_reward(): from nemo_rl.algorithms.ppo import _create_advantage_estimator master_config = SimpleNamespace( - ppo={ - "adv_estimator": {"name": "raw_reward", "normalize_advantages": True}, - }, + ppo=PPOConfig( + adv_estimator={"name": "raw_reward", "normalize_advantages": True}, + ), loss_fn={"reference_policy_kl_penalty": 0.0}, ) @@ -678,7 +679,7 @@ def test_create_advantage_estimator_rejects_unsupported_name(): from nemo_rl.algorithms.ppo import _create_advantage_estimator master_config = SimpleNamespace( - ppo={"adv_estimator": {"name": "grpo"}}, + ppo=PPOConfig(adv_estimator={"name": "grpo"}), loss_fn={"reference_policy_kl_penalty": 0.0}, ) @@ -686,19 +687,17 @@ def test_create_advantage_estimator_rejects_unsupported_name(): _create_advantage_estimator(master_config) -def test_create_advantage_estimator_requires_adv_estimator_key(): - """No more silent default β€” missing `adv_estimator` should KeyError.""" +def test_create_advantage_estimator_uses_ppo_config_default(): + """The PPO schema provides the centralized default estimator config.""" from types import SimpleNamespace from nemo_rl.algorithms.ppo import _create_advantage_estimator - master_config = SimpleNamespace( - ppo={}, - loss_fn={}, - ) + master_config = SimpleNamespace(ppo=PPOConfig(), loss_fn=_make_loss_config()) - with pytest.raises(KeyError): - _create_advantage_estimator(master_config) + assert isinstance( + _create_advantage_estimator(master_config), GeneralizedAdvantageEstimator + ) def _make_ppo_loop_batch( @@ -743,10 +742,13 @@ def _make_ppo_loop_batch( def _run_mock_ppo_train( monkeypatch, *, + async_mode: bool = False, + checkpoint_path: str | None = None, max_num_steps: int, ppo_epochs: int, seq_logprob_error_threshold: float | None, policy_training_start_step: int = 0, + warmup_generation_lead_steps: int | None = None, overlong_filtering: bool = False, truncated_samples: tuple[bool, bool] = (False, False), ): @@ -787,12 +789,24 @@ def compute_advantage(self, **kwargs): return mask.clone(), mask.clone() class DummyTimer: + def __init__(self, *_args, **_kwargs): + self._timers = {} + def time(self, *_args, **_kwargs): return nullcontext() + def start(self, *_args, **_kwargs): + pass + + def stop(self, *_args, **_kwargs): + pass + def get_timing_metrics(self, **_kwargs): return {"total_step_time": 1.0} + def reduce(self, *_args, **_kwargs): + return 0.0 + def reset(self): pass @@ -847,6 +861,12 @@ def __len__(self): } value_model = MagicMock() + value_model.prepare_for_inference.side_effect = lambda: events.append( + "value_inference_prep" + ) + value_model.finish_inference.side_effect = lambda: events.append( + "value_inference_finish" + ) value_model.finish_training.side_effect = lambda: events.append("value_finish") value_model.get_values.return_value = {"values": torch.zeros(2, 3, 1)} value_model.train.side_effect = lambda *_args, **_kwargs: ( @@ -872,10 +892,15 @@ def fake_rollout(*_args, input_batch, **_kwargs): monkeypatch.setattr(ppo_mod, "TimeoutChecker", DummyTimeoutChecker) monkeypatch.setattr(ppo_mod, "MemoryTracker", DummyMemoryTracker) monkeypatch.setattr(ppo_mod, "maybe_gpu_profile_step", lambda *_args: None) - monkeypatch.setattr(ppo_mod, "print_performance_metrics", lambda *_args: {}) + monkeypatch.setattr( + ppo_mod, "print_performance_metrics", lambda *_args, **_kwargs: {} + ) + monkeypatch.setattr( + ppo_mod, "print_efficiency_summary", lambda *_args, **_kwargs: {} + ) monkeypatch.setattr(ppo_mod, "scale_rewards", lambda batch, _config: batch) - monkeypatch.setattr(ppo_mod, "_should_use_nemo_gym", lambda _config: False) - monkeypatch.setattr(ppo_mod, "_should_use_async_rollouts", lambda _config: False) + monkeypatch.setattr(ppo_mod, "should_use_nemo_gym", lambda _config: False) + monkeypatch.setattr(ppo_mod, "should_use_async_rollouts", lambda _config: False) monkeypatch.setattr(ppo_mod, "run_multi_turn_rollout", fake_rollout) monkeypatch.setattr(ppo_mod, "refit_policy_generation", refit) monkeypatch.setattr(ppo_mod, "batched_message_log_to_flat_message", fake_flatten) @@ -891,22 +916,26 @@ def fake_rollout(*_args, input_batch, **_kwargs): ) master_config = SimpleNamespace( - ppo={ - "max_num_steps": max_num_steps, - "max_num_epochs": 1, - "max_rollout_turns": 1, - "num_prompts_per_step": 2, - "num_generations_per_prompt": 1, - "overlong_filtering": overlong_filtering, - "policy_training_start_step": policy_training_start_step, - "ppo_epochs": ppo_epochs, - "reward_scaling": {"enabled": False}, - "reward_shaping": RewardShapingConfig(enabled=False), - "seq_logprob_error_threshold": seq_logprob_error_threshold, - "val_at_start": False, - "val_at_end": False, - "val_period": 0, - }, + ppo=PPOConfig( + async_ppo={ + "enabled": async_mode, + "warmup_generation_lead_steps": warmup_generation_lead_steps, + }, + max_num_steps=max_num_steps, + max_num_epochs=-1 if async_mode else 1, + max_rollout_turns=1, + num_prompts_per_step=2, + num_generations_per_prompt=1, + overlong_filtering=overlong_filtering, + policy_training_start_step=policy_training_start_step, + ppo_epochs=ppo_epochs, + reward_scaling={"enabled": False}, + reward_shaping=RewardShapingConfig(enabled=False), + seq_logprob_error_threshold=seq_logprob_error_threshold, + val_at_start=False, + val_at_end=False, + val_period=0, + ), policy={ "generation": { "backend": "vllm", @@ -918,7 +947,7 @@ def fake_rollout(*_args, input_batch, **_kwargs): }, loss_fn=_make_loss_config(), checkpointing={ - "enabled": False, + "enabled": checkpoint_path is not None, "checkpoint_must_save_by": None, "save_period": 100, "metric_name": None, @@ -929,27 +958,89 @@ def fake_rollout(*_args, input_batch, **_kwargs): logger = MagicMock() checkpointer = MagicMock() checkpointer.save_optimizer = False + checkpointer.get_latest_checkpoint_path.return_value = checkpoint_path + checkpointer.init_tmp_checkpoint.return_value = checkpoint_path dataloader = DummyLoader( [_make_ppo_loop_batch(truncated_samples) for _ in range(max_num_steps)] ) tokenizer = SimpleNamespace(pad_token_id=0) + replay_actor = None - ppo_mod.ppo_train( - policy, - policy_generation, - value_model, - dataloader, - None, - tokenizer, - MagicMock(), - MagicMock(), - {}, - None, - logger, - checkpointer, - ppo_mod._default_ppo_save_state(), - master_config, - ) + if async_mode: + from nemo_rl.algorithms import async_utils + + master_config.policy["generation"]["vllm_cfg"]["async_engine"] = True + master_config.loss_fn.use_importance_sampling_correction = True + + replay_actor = MagicMock() + replay_actor.has_complete_batch.remote.return_value = True + replay_actor.sample.remote.return_value = { + "trajectories": [ + { + "batch": _make_ppo_loop_batch(truncated_samples).slice(i, i + 1), + "rollout_metrics": {}, + } + for i in range(2) + ], + "avg_trajectory_age": 0.0, + } + replay_actor.size.remote.return_value = 2 + replay_actor.save_to_path.remote.return_value = 2 + + collector_actor = MagicMock() + collector_actor.check_health.remote.return_value = None + collector_actor.get_status.remote.return_value = { + "errored": False, + "running": True, + "inflight_workers": 0, + "data_exhausted": False, + } + collector_actor.get_efficiency_metrics.remote.return_value = {} + collector_actor.get_dataloader_state.remote.return_value = {} + + replay_type = MagicMock() + replay_type.options.return_value.remote.return_value = replay_actor + collector_type = MagicMock() + collector_type.options.return_value.remote.return_value = collector_actor + monkeypatch.setattr(async_utils, "ReplayBuffer", replay_type) + monkeypatch.setattr(async_utils, "AsyncTrajectoryCollector", collector_type) + monkeypatch.setattr(ppo_mod, "make_actor_runtime_env", lambda _actor: {}) + monkeypatch.setattr(ppo_mod.ray, "get", lambda value, **_kwargs: value) + monkeypatch.setattr(ppo_mod.ray, "kill", lambda _actor: None) + + ppo_mod.async_ppo_train( + policy, + policy_generation, + value_model, + dataloader, + None, + tokenizer, + MagicMock(), + MagicMock(), + {}, + None, + logger, + checkpointer, + ppo_mod._default_ppo_save_state(), + master_config, + ) + else: + ppo_mod.ppo_train( + policy, + policy_generation, + value_model, + dataloader, + None, + tokenizer, + MagicMock(), + MagicMock(), + {}, + None, + logger, + checkpointer, + ppo_mod._default_ppo_save_state(), + master_config, + ) return SimpleNamespace( policy=policy, @@ -959,6 +1050,7 @@ def fake_rollout(*_args, input_batch, **_kwargs): checkpointer=checkpointer, advantage_estimator=advantage_estimator, refit=refit, + replay_actor=replay_actor, events=events, ) @@ -996,33 +1088,91 @@ def test_ppo_train_noncolocated_refit_offload_lifecycle(monkeypatch): ] -def test_ppo_train_critic_warmup_reuses_generation_until_policy_update(monkeypatch): +@pytest.mark.parametrize("async_mode", [False, True]) +def test_ppo_train_critic_warmup_reuses_generation_until_policy_update( + monkeypatch, async_mode +): harness = _run_mock_ppo_train( monkeypatch, + async_mode=async_mode, max_num_steps=2, ppo_epochs=1, seq_logprob_error_threshold=None, policy_training_start_step=1, ) - assert harness.refit.call_count == 1 - assert harness.policy_generation.prepare_for_generation.call_count == 1 + assert harness.refit.call_count == (2 if async_mode else 1) + assert harness.policy_generation.prepare_for_generation.call_count == ( + 0 if async_mode else 1 + ) assert harness.policy.train.call_count == 1 assert harness.value_model.train.call_count == 2 - assert harness.policy_generation.finish_generation.call_count == 2 + assert harness.policy_generation.finish_generation.call_count == ( + 0 if async_mode else 2 + ) - rollout_indices = [ - index for index, event in enumerate(harness.events) if event == "rollout" - ] - assert harness.events[rollout_indices[1] - 1 : rollout_indices[1] + 1] == [ - "generation_prepare", - "rollout", - ] + if not async_mode: + rollout_indices = [ + index for index, event in enumerate(harness.events) if event == "rollout" + ] + assert harness.events[rollout_indices[1] - 1 : rollout_indices[1] + 1] == [ + "generation_prepare", + "rollout", + ] + + +@pytest.mark.parametrize( + ( + "policy_training_start_step", + "warmup_generation_lead_steps", + "expected_restore_max_age", + ), + [ + (0, None, 1), + (2, 3, 3), + ], +) +def test_async_ppo_checkpoints_replay_buffer_inside_actor( + monkeypatch, + tmp_path, + policy_training_start_step, + warmup_generation_lead_steps, + expected_restore_max_age, +): + checkpoint_path = tmp_path / "step_0" + checkpoint_path.mkdir() + replay_buffer_path = checkpoint_path / "replay_buffer.pt" + replay_buffer_path.touch() + harness = _run_mock_ppo_train( + monkeypatch, + async_mode=True, + checkpoint_path=str(checkpoint_path), + max_num_steps=1, + ppo_epochs=1, + seq_logprob_error_threshold=None, + policy_training_start_step=policy_training_start_step, + warmup_generation_lead_steps=warmup_generation_lead_steps, + ) -def test_ppo_train_excludes_overlong_samples_from_advantage(monkeypatch): + harness.replay_actor.load_from_path.remote.assert_called_once_with( + str(replay_buffer_path), + num_prompts_per_step=2, + current_training_step=0, + max_age_steps=expected_restore_max_age, + ) + harness.replay_actor.load_state_dict.remote.assert_not_called() + harness.replay_actor.save_to_path.remote.assert_called_once_with( + str(replay_buffer_path) + ) + harness.replay_actor.state_dict.remote.assert_not_called() + + +@pytest.mark.parametrize("async_mode", [False, True]) +def test_ppo_train_excludes_overlong_samples_from_advantage(monkeypatch, async_mode): harness = _run_mock_ppo_train( monkeypatch, + async_mode=async_mode, max_num_steps=1, ppo_epochs=1, seq_logprob_error_threshold=None, @@ -1036,13 +1186,15 @@ def test_ppo_train_excludes_overlong_samples_from_advantage(monkeypatch): ) -def test_ppo_train_rejects_all_masked_batch(monkeypatch): +@pytest.mark.parametrize("async_mode", [False, True]) +def test_ppo_train_rejects_all_masked_batch(monkeypatch, async_mode): with pytest.raises( RuntimeError, match="no valid response tokens after filtering", ): _run_mock_ppo_train( monkeypatch, + async_mode=async_mode, max_num_steps=1, ppo_epochs=1, seq_logprob_error_threshold=None, @@ -1051,9 +1203,13 @@ def test_ppo_train_rejects_all_masked_batch(monkeypatch): ) -def test_ppo_train_wires_logprob_mask_to_advantage_training_and_metrics(monkeypatch): +@pytest.mark.parametrize("async_mode", [False, True]) +def test_ppo_train_wires_logprob_mask_to_advantage_training_and_metrics( + monkeypatch, async_mode +): harness = _run_mock_ppo_train( monkeypatch, + async_mode=async_mode, max_num_steps=1, ppo_epochs=1, seq_logprob_error_threshold=1.5, @@ -1084,6 +1240,67 @@ def test_ppo_train_wires_logprob_mask_to_advantage_training_and_metrics(monkeypa assert final_train_metrics[0]["advantages/max"] == pytest.approx(1.0) +def test_compute_critic_metrics_aggregates_and_namespaces_results(): + from nemo_rl.algorithms.ppo import _compute_critic_metrics + + metrics = _compute_critic_metrics( + { + "grad_norm": torch.tensor(1.5), + "loss": torch.tensor(0.25), + "all_mb_metrics": { + "lr": [0.1, 0.3], + "values_min": [-2.0, -1.0], + "values_max": [1.0, 3.0], + "num_valid_tokens": [2, 3], + "returns_mean": [1.0], + "values_mean": [0.5], + "returns_sq_mean": [5.0], + "residual_sq_mean": [1.25], + }, + } + ) + + assert metrics["critic/grad_norm"] == pytest.approx(1.5) + assert metrics["critic/loss"] == pytest.approx(0.25) + assert metrics["critic/lr"] == pytest.approx(0.2) + assert metrics["critic/values_min"] == pytest.approx(-2.0) + assert metrics["critic/values_max"] == pytest.approx(3.0) + assert metrics["critic/num_valid_tokens"] == 5 + assert metrics["critic/explained_var"] == pytest.approx(0.75) + + +def test_compute_critic_metrics_handles_zero_return_variance(): + from nemo_rl.algorithms.ppo import _compute_critic_metrics + + metrics = _compute_critic_metrics( + { + "grad_norm": torch.tensor(0.0), + "loss": torch.tensor(0.0), + "all_mb_metrics": { + "returns_mean": [2.0], + "values_mean": [2.0], + "returns_sq_mean": [4.0], + "residual_sq_mean": [0.0], + }, + } + ) + + assert metrics["critic/explained_var"] == pytest.approx(1.0) + + +def test_compute_critic_metrics_rejects_unsupported_metric_type(): + from nemo_rl.algorithms.ppo import _compute_critic_metrics + + with pytest.raises(ValueError, match="Unsupported value-model metric"): + _compute_critic_metrics( + { + "grad_norm": torch.tensor(0.0), + "loss": torch.tensor(0.0), + "all_mb_metrics": {"unexpected": 1.0}, + } + ) + + # ============================================================================ # Tests for non-colocated setup # ============================================================================ @@ -1158,27 +1375,27 @@ def _make_noncolocated_setup_config( value_loss_fn=MseValueLossConfig(), env=env_config, data=data_config, - ppo={ - "max_num_steps": 1, - "max_num_epochs": 1, - "num_prompts_per_step": 1, - "num_generations_per_prompt": 1, - "max_rollout_turns": 1, - "val_period": 0, - "val_batch_size": 1, - "val_at_start": False, - "val_at_end": False, - "max_val_samples": 1, - "seed": 42, - "overlong_filtering": False, - "use_dynamic_sampling": False, - "batch_multiplier": 1, - "ppo_epochs": 1, - "policy_training_start_step": 0, - "reward_shaping": {"enabled": False}, - "reward_scaling": {"enabled": False}, - "adv_estimator": {"name": "raw_reward"}, - }, + ppo=PPOConfig( + max_num_steps=1, + max_num_epochs=1, + num_prompts_per_step=1, + num_generations_per_prompt=1, + max_rollout_turns=1, + val_period=0, + val_batch_size=1, + val_at_start=False, + val_at_end=False, + max_val_samples=1, + seed=42, + overlong_filtering=False, + use_dynamic_sampling=False, + batch_multiplier=1, + ppo_epochs=1, + policy_training_start_step=0, + reward_shaping={"enabled": False}, + reward_scaling={"enabled": False}, + adv_estimator={"name": "raw_reward"}, + ), logger={"num_val_samples_to_print": 0}, cluster={ "num_nodes": total_nodes, @@ -1740,6 +1957,30 @@ def test_noncolocated_vllm_builds_separate_clusters_and_collective(monkeypatch): generation.prepare_refit_info.assert_called_once_with({"state": "dict"}) +@pytest.mark.parametrize( + ("async_enabled", "expected_train_iters"), + [(False, 3), (True, 30)], +) +def test_megatron_train_iters_matches_ppo_training_limit( + monkeypatch, async_enabled, expected_train_iters +): + """Async PPO cycles data until max_num_steps; sync PPO also honors epochs.""" + from nemo_rl.algorithms.ppo import AsyncPPOConfig + + config = _make_noncolocated_setup_config() + config.policy["dtensor_cfg"]["enabled"] = False + config.policy["megatron_cfg"]["enabled"] = True + config.ppo.max_num_steps = 10 + config.ppo.max_num_epochs = -1 if async_enabled else 1 + config.ppo.ppo_epochs = 3 + config.ppo.async_ppo = AsyncPPOConfig(enabled=async_enabled) + + _run_noncolocated_setup(monkeypatch, config) + + assert config.policy["megatron_cfg"]["train_iters"] == expected_train_iters + assert config.value["megatron_cfg"]["train_iters"] == expected_train_iters + + def test_colocated_setup_keeps_single_cluster_and_skips_collective(monkeypatch): """The default colocated setup remains unchanged by the cluster split.""" config = _make_noncolocated_setup_config() @@ -1823,3 +2064,532 @@ def test_noncolocated_reward_model_node_leaves_shared_train_inference_node( generation.init_collective.assert_called_once_with( "127.0.0.1", 1234, 8, train_world_size=6 ) + + +def _make_async_ppo_config() -> SimpleNamespace: + from nemo_rl.algorithms.ppo import AsyncPPOConfig + + return SimpleNamespace( + policy={ + "generation": { + "backend": "vllm", + "colocated": {"enabled": False}, + "vllm_cfg": {"async_engine": True}, + } + }, + loss_fn=ClippedPGLossConfig( + use_importance_sampling_correction=True, + reference_policy_kl_penalty=0, + ), + ppo=PPOConfig( + async_ppo=AsyncPPOConfig(enabled=True), + max_num_epochs=-1, + policy_training_start_step=0, + ppo_epochs=1, + use_dynamic_sampling=False, + reward_scaling={"enabled": False}, + reward_shaping=RewardShapingConfig(enabled=False), + ), + data={"use_multiple_dataloader": False}, + env={}, + checkpointing={"checkpoint_must_save_by": None}, + ) + + +def _call_async_ppo_until_guard( + master_config: SimpleNamespace, *, requires_kv_scale_sync: bool = False +) -> None: + from nemo_rl.algorithms.ppo import async_ppo_train + + generation = MagicMock() + generation.requires_kv_scale_sync = requires_kv_scale_sync + async_ppo_train( + policy=MagicMock(), + policy_generation=generation, + value_model=MagicMock(), + dataloader=MagicMock(), + val_dataloader=None, + tokenizer=MagicMock(), + loss_fn=MagicMock(), + value_loss_fn=MagicMock(), + task_to_env={}, + val_task_to_env=None, + logger=MagicMock(), + checkpointer=MagicMock(), + ppo_save_state=MagicMock(), + master_config=master_config, + ) + + +def _validate_async_ppo_entry_config( + master_config: SimpleNamespace, *, requires_kv_scale_sync: bool = False +) -> None: + from examples.run_ppo import _validate_async_ppo_config + + generation = MagicMock() + generation.requires_kv_scale_sync = requires_kv_scale_sync + _validate_async_ppo_config(master_config, generation) + + +@pytest.mark.parametrize( + ("mutate", "message"), + [ + ( + lambda cfg: cfg.policy["generation"].update(backend="sglang"), + "backend=vllm.*async_engine=true", + ), + ( + lambda cfg: cfg.policy["generation"]["vllm_cfg"].update(async_engine=False), + "backend=vllm.*async_engine=true", + ), + ( + lambda cfg: setattr( + cfg.loss_fn, "use_importance_sampling_correction", False + ), + "importance_sampling_correction", + ), + ( + lambda cfg: setattr(cfg.loss_fn, "force_on_policy_ratio", True), + "force_on_policy_ratio", + ), + ( + lambda cfg: cfg.policy["generation"]["colocated"].update(enabled=True), + "non-colocated", + ), + ], +) +def test_async_ppo_launcher_entry_guards(mutate, message): + config = _make_async_ppo_config() + mutate(config) + with pytest.raises(ValueError, match=message): + _validate_async_ppo_entry_config(config) + + +@pytest.mark.parametrize( + ("mutate", "message"), + [ + (lambda cfg: setattr(cfg.ppo, "ppo_epochs", 0), "ppo_epochs"), + ( + lambda cfg: ( + setattr(cfg.ppo, "skip_reference_policy_logprobs_calculation", True), + setattr(cfg.loss_fn, "reference_policy_kl_penalty", 0.1), + ), + "Skipping reference logprobs", + ), + ], +) +def test_async_ppo_training_loop_guards(mutate, message): + config = _make_async_ppo_config() + mutate(config) + with pytest.raises(ValueError, match=message): + _call_async_ppo_until_guard(config) + + +@pytest.mark.parametrize( + ("mutate", "message"), + [ + ( + lambda cfg: setattr(cfg.ppo, "use_dynamic_sampling", True), + "Dynamic sampling", + ), + ( + lambda cfg: setattr(cfg.ppo.reward_scaling, "enabled", True), + "Reward scaling", + ), + ( + lambda cfg: setattr(cfg.ppo.reward_shaping, "enabled", True), + "Reward shaping", + ), + ( + lambda cfg: cfg.data.update(use_multiple_dataloader=True), + "Multiple dataloaders", + ), + ( + lambda cfg: cfg.env.update(should_use_nemo_gym=True), + "NeMo Gym", + ), + ( + lambda cfg: setattr(cfg.ppo, "max_num_epochs", 2), + "max_num_epochs=-1", + ), + ], +) +def test_async_ppo_rejects_unsupported_features(mutate, message): + config = _make_async_ppo_config() + mutate(config) + with pytest.raises(NotImplementedError, match=message): + _validate_async_ppo_entry_config(config) + + +def test_async_ppo_rejects_fp8_kv_scale_sync(): + config = _make_async_ppo_config() + with pytest.raises(NotImplementedError, match="FP8 KV-scale"): + _validate_async_ppo_entry_config(config, requires_kv_scale_sync=True) + + +def test_async_ppo_config_allows_kv_cache_recompute_without_inflight_updates(): + from nemo_rl.algorithms.ppo import AsyncPPOConfig + + config = AsyncPPOConfig( + in_flight_weight_updates=False, + recompute_kv_cache_after_weight_updates=True, + ) + + assert not config.in_flight_weight_updates + assert config.recompute_kv_cache_after_weight_updates + + +def test_async_ppo_config_defaults(): + from nemo_rl.algorithms.ppo import AsyncPPOConfig + + config = AsyncPPOConfig() + + assert not config.in_flight_weight_updates + assert not config.drop_incomplete_targets_on_restore + + +def test_async_ppo_config_warmup_lead_defaults_to_training_age(): + from nemo_rl.algorithms.ppo import AsyncPPOConfig + + config = AsyncPPOConfig(max_trajectory_age_steps=3) + + assert config.warmup_generation_lead_steps is None + assert config.resolved_warmup_generation_lead_steps == 3 + + +def test_async_ppo_config_rejects_warmup_lead_below_training_age(): + from pydantic import ValidationError + + from nemo_rl.algorithms.ppo import AsyncPPOConfig + + with pytest.raises(ValidationError, match="warmup_generation_lead_steps"): + AsyncPPOConfig( + max_trajectory_age_steps=2, + warmup_generation_lead_steps=1, + ) + + +def test_ppo_config_rejects_warmup_lead_without_critic_warmup(): + from pydantic import ValidationError + + with pytest.raises(ValidationError, match="policy_training_start_step > 0"): + PPOConfig( + policy_training_start_step=0, + async_ppo={ + "enabled": True, + "warmup_generation_lead_steps": 2, + }, + ) + + +def test_ppo_config_allows_warmup_lead_with_critic_warmup(): + config = PPOConfig( + policy_training_start_step=1, + async_ppo={ + "enabled": True, + "warmup_generation_lead_steps": 2, + }, + ) + + assert config.async_ppo.warmup_generation_lead_steps == 2 + + +@pytest.mark.parametrize( + ("step", "expected_lead", "expected_buffer_age"), + [ + (0, 4, 4), + (1, 4, 4), + (2, 3, 4), + (3, 2, 4), + (4, 1, 4), + (5, 1, 4), + (6, 1, 1), + ], +) +def test_async_ppo_warmup_window_has_fixed_safe_frontier( + step, expected_lead, expected_buffer_age +): + from nemo_rl.algorithms.ppo import ( + _async_ppo_buffer_max_age, + _async_ppo_generation_lead_steps, + ) + + generation_lead = _async_ppo_generation_lead_steps( + step=step, + policy_training_start_step=4, + max_trajectory_age_steps=1, + warmup_generation_lead_steps=4, + ) + buffer_age = _async_ppo_buffer_max_age( + step=step, + policy_training_start_step=4, + max_trajectory_age_steps=1, + warmup_generation_lead_steps=4, + ) + + assert generation_lead == expected_lead + if step <= 4: + assert step + generation_lead <= 5 + assert buffer_age == expected_buffer_age + + +def test_async_ppo_warmup_window_is_disabled_without_critic_warmup(): + from nemo_rl.algorithms.ppo import ( + _async_ppo_buffer_max_age, + _async_ppo_generation_lead_steps, + ) + + assert ( + _async_ppo_generation_lead_steps( + step=0, + policy_training_start_step=0, + max_trajectory_age_steps=1, + warmup_generation_lead_steps=4, + ) + == 1 + ) + assert ( + _async_ppo_buffer_max_age( + step=0, + policy_training_start_step=0, + max_trajectory_age_steps=1, + warmup_generation_lead_steps=4, + ) + == 1 + ) + + +def test_async_ppo_consumes_frozen_policy_rollout_at_safe_warmup_frontier(): + """A rollout banked at version 0 remains usable through the W+A frontier.""" + from nemo_rl.algorithms.async_utils.replay_buffer import ReplayBufferImpl + from nemo_rl.algorithms.ppo import ( + _async_ppo_buffer_max_age, + _async_ppo_generation_lead_steps, + ) + + policy_training_start_step = 2 + max_trajectory_age_steps = 1 + warmup_generation_lead_steps = 3 + frontier = policy_training_start_step + max_trajectory_age_steps + + assert ( + _async_ppo_generation_lead_steps( + step=0, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=warmup_generation_lead_steps, + ) + == frontier + ) + + buffer = ReplayBufferImpl( + max_size=4, + drop_incomplete_targets_on_restore=False, + ) + frozen_rollout = { + "batch": {"data": "frozen-policy"}, + "rollout_metrics": {}, + } + assert ( + buffer.add( + frozen_rollout, + weight_version=0, + target_weight_version=frontier, + ) + == "success" + ) + + frontier_max_age = _async_ppo_buffer_max_age( + step=frontier, + policy_training_start_step=policy_training_start_step, + max_trajectory_age_steps=max_trajectory_age_steps, + warmup_generation_lead_steps=warmup_generation_lead_steps, + ) + sample = buffer.sample( + num_prompt_groups=1, + current_weight_version=frontier, + max_age_steps=frontier_max_age, + ) + + assert sample is not None + assert sample["trajectories"] == [frozen_rollout] + assert sample["avg_trajectory_age"] == frontier + assert buffer.size() == 0 + + +def test_async_ppo_completed_resume_exits_before_actor_start(monkeypatch): + from nemo_rl.algorithms import ppo + + config = _make_async_ppo_config() + config.ppo.max_num_steps = 10 + config.ppo.max_num_epochs = -1 + config.ppo.val_period = 0 + config.ppo.val_at_start = False + config.ppo.val_at_end = False + config.ppo.num_prompts_per_step = 1 + config.ppo.skip_reference_policy_logprobs_calculation = False + config.checkpointing = { + "checkpoint_must_save_by": None, + "ft_save_period": None, + } + policy = MagicMock() + generation = MagicMock() + generation.requires_kv_scale_sync = False + value_model = MagicMock() + checkpointer = MagicMock() + refit = MagicMock() + monkeypatch.setattr(ppo, "refit_policy_generation", refit) + + ppo.async_ppo_train( + policy=policy, + policy_generation=generation, + value_model=value_model, + dataloader=[MagicMock()], + val_dataloader=None, + tokenizer=MagicMock(), + loss_fn=MagicMock(), + value_loss_fn=MagicMock(), + task_to_env={}, + val_task_to_env=None, + logger=MagicMock(), + checkpointer=checkpointer, + ppo_save_state={ + "total_steps": 10, + "consumed_samples": 10, + "current_epoch": 1, + "current_step": 0, + "total_valid_tokens": 0, + }, + master_config=config, + ) + + refit.assert_not_called() + checkpointer.shutdown.assert_called_once() + generation.shutdown.assert_called_once() + policy.shutdown.assert_called_once() + value_model.shutdown.assert_called_once() + + +def test_async_ppo_initial_refit_failure_cleans_up_actors(monkeypatch): + from unittest.mock import call + + from nemo_rl.algorithms import async_utils, ppo + + config = _make_async_ppo_config() + config.policy["make_sequence_length_divisible_by"] = 1 + config.ppo.max_num_steps = 2 + config.ppo.max_num_epochs = -1 + config.ppo.val_period = 0 + config.ppo.val_at_start = False + config.ppo.val_at_end = False + config.ppo.num_prompts_per_step = 1 + config.ppo.max_rollout_turns = 1 + config.ppo.skip_reference_policy_logprobs_calculation = False + config.ppo.adv_estimator = { + "name": "raw_reward", + "normalize_advantages": False, + } + config.checkpointing = { + "checkpoint_must_save_by": None, + "ft_save_period": None, + } + + replay_actor = MagicMock() + collector_actor = MagicMock() + replay_type = MagicMock() + replay_type.options.return_value.remote.return_value = replay_actor + collector_type = MagicMock() + collector_type.options.return_value.remote.return_value = collector_actor + monkeypatch.setattr(async_utils, "ReplayBuffer", replay_type) + monkeypatch.setattr(async_utils, "AsyncTrajectoryCollector", collector_type) + monkeypatch.setattr( + ppo, + "make_actor_runtime_env", + lambda _actor: {"py_executable": "/tmp/fake-venv/bin/python"}, + ) + monkeypatch.setattr( + ppo, + "refit_policy_generation", + MagicMock(side_effect=RuntimeError("initial refit failed")), + ) + ray_kill = MagicMock() + monkeypatch.setattr(ppo.ray, "kill", ray_kill) + + policy = MagicMock() + generation = MagicMock() + generation.requires_kv_scale_sync = False + value_model = MagicMock() + checkpointer = MagicMock() + checkpointer.get_latest_checkpoint_path.return_value = None + + with pytest.raises(RuntimeError, match="initial refit failed"): + ppo.async_ppo_train( + policy=policy, + policy_generation=generation, + value_model=value_model, + dataloader=[MagicMock()], + val_dataloader=None, + tokenizer=MagicMock(), + loss_fn=MagicMock(), + value_loss_fn=MagicMock(), + task_to_env={}, + val_task_to_env=None, + logger=MagicMock(), + checkpointer=checkpointer, + ppo_save_state={ + "total_steps": 0, + "consumed_samples": 0, + "current_epoch": 0, + "current_step": 0, + "total_valid_tokens": 0, + }, + master_config=config, + ) + + ray_kill.assert_has_calls( + [call(collector_actor), call(replay_actor)], any_order=True + ) + checkpointer.shutdown.assert_called_once() + generation.shutdown.assert_called_once() + policy.shutdown.assert_called_once() + value_model.shutdown.assert_called_once() + + +@pytest.mark.parametrize("async_engine", [False, True]) +def test_validate_dispatches_rollout_by_engine_mode(monkeypatch, async_engine): + from nemo_rl.algorithms import ppo + + rollout_result = ( + { + "total_reward": torch.tensor([1.0]), + "message_log": [[{"role": "assistant", "content": "ok"}]], + }, + {"mean_gen_tokens_per_sample": 1.0}, + ) + async_rollout = MagicMock(return_value=rollout_result) + sync_rollout = MagicMock(return_value=rollout_result) + monkeypatch.setattr(ppo, "run_async_multi_turn_rollout", async_rollout) + monkeypatch.setattr(ppo, "run_multi_turn_rollout", sync_rollout) + + config = _make_async_ppo_config() + config.policy["generation"]["vllm_cfg"]["async_engine"] = async_engine + config.policy["max_total_sequence_length"] = 16 + config.ppo.max_val_samples = 1 + config.ppo.val_batch_size = 1 + config.ppo.max_rollout_turns = 1 + config.logger = {"num_val_samples_to_print": 0} + + ppo.validate( + policy_generation=MagicMock(), + val_dataloader=[MagicMock()], + tokenizer=MagicMock(), + val_task_to_env={}, + step=1, + master_config=config, + logger=None, + ) + + selected_rollout = async_rollout if async_engine else sync_rollout + unselected_rollout = sync_rollout if async_engine else async_rollout + selected_rollout.assert_called_once() + unselected_rollout.assert_not_called() diff --git a/tests/unit/algorithms/test_utils.py b/tests/unit/algorithms/test_utils.py index a44244a6de3..74b18615f2d 100755 --- a/tests/unit/algorithms/test_utils.py +++ b/tests/unit/algorithms/test_utils.py @@ -19,6 +19,8 @@ import torch from nemo_rl.algorithms.grpo import AsyncGRPOConfig, GRPOConfig, MasterConfig +from nemo_rl.algorithms.ppo import MasterConfig as PPOMasterConfig +from nemo_rl.algorithms.ppo import PPOConfig from nemo_rl.algorithms.utils import ( EFFICIENCY_CATEGORIES, WALL_CLOCK_EFFICIENCY_CATEGORIES, @@ -247,6 +249,26 @@ def _base_master_config(colocated: bool): ) +def _base_ppo_master_config(colocated: bool): + return PPOMasterConfig.model_construct( + cluster={"num_nodes": 2, "gpus_per_node": 8}, + policy={ + "generation": { + "temperature": 1.0, + "top_p": 1.0, + "top_k": None, + "colocated": { + "enabled": colocated, + "resources": {"num_nodes": 1, "gpus_per_node": 8}, + }, + } + }, + ppo=PPOConfig.model_construct( + num_prompts_per_step=8, num_generations_per_prompt=10 + ), + ) + + def test_sync_colocated_throughput_flops_and_imbalance(capsys): master_config = _base_master_config(colocated=True) @@ -274,7 +296,13 @@ def test_sync_colocated_throughput_flops_and_imbalance(capsys): } perf = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=8, + num_generations_per_prompt=10, + is_async_rl=False, ) # Validate key throughput metrics @@ -343,7 +371,13 @@ def test_train_elapsed_seconds_used_for_flops_calculation(capsys): } perf = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=8, + num_generations_per_prompt=10, + is_async_rl=False, ) assert math.isclose(perf["train_flops_per_gpu"], 500.0 / 8, rel_tol=1e-6) @@ -375,9 +409,17 @@ def test_async_non_colocated_idle_ratio_and_generation_time(capsys): train_results = {} perf = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=8, + num_generations_per_prompt=10, + is_async_rl=True, ) + assert "training_worker_idle_time_ratio" in perf + # Throughput checks assert math.isclose(perf["samples_per_sec_per_gpu"], 0.5, rel_tol=1e-6) assert math.isclose( @@ -427,7 +469,13 @@ def test_minimal_inputs_no_counts_no_flops(capsys): train_results = {} perf = print_performance_metrics( - train_results, metrics, timing_metrics, master_config + train_results, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=8, + num_generations_per_prompt=10, + is_async_rl=False, ) # Core metrics exist @@ -459,7 +507,15 @@ def test_empty_per_worker_token_counts_skips_imbalance(capsys): "per_worker_token_counts": {}, } - perf = print_performance_metrics({}, metrics, timing_metrics, master_config) + perf = print_performance_metrics( + {}, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=8, + num_generations_per_prompt=10, + is_async_rl=False, + ) assert "average_token_imbalance" not in perf @@ -468,6 +524,30 @@ def test_empty_per_worker_token_counts_skips_imbalance(capsys): assert "Throughputs (per GPU)" in out +def test_async_ppo_metrics_use_async_flag_without_grpo_config(): + master_config = _base_ppo_master_config(colocated=False) + timing_metrics = { + "policy_and_reference_logprobs": 2.0, + "policy_training": 4.0, + "total_step_time": 10.0, + "exposed_generation": 2.0, + "prepare_for_generation/total": 1.0, + } + metrics = {"total_num_tokens": 6050.0} + + perf = print_performance_metrics( + {}, + metrics, + timing_metrics, + master_config, + num_prompts_per_step=8, + num_generations_per_prompt=10, + is_async_rl=True, + ) + + assert "training_worker_idle_time_ratio" in perf + + # ============================================================================ # Tests for calculate_baseline_and_std_per_prompt function # ============================================================================ diff --git a/tests/unit/data/test_dataloader.py b/tests/unit/data/test_dataloader.py new file mode 100644 index 00000000000..da6a7a574e3 --- /dev/null +++ b/tests/unit/data/test_dataloader.py @@ -0,0 +1,63 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from collections.abc import Iterator + +import pytest + +from nemo_rl.data.dataloader import CyclingDataLoader + + +class _EpochListLoader: + def __init__(self, sizes: list[int]) -> None: + self.sizes = list(sizes) + self.iter_count = 0 + + def __iter__(self) -> Iterator[int]: + size = self.sizes[min(self.iter_count, len(self.sizes) - 1)] + self.iter_count += 1 + return iter(range(size)) + + def state_dict(self) -> dict[str, int]: + return {"iter_count": self.iter_count} + + +def test_cycling_dataloader_cycles_after_resume_boundary(): + dataloader = _EpochListLoader([0, 3]) + iterator = iter(CyclingDataLoader(dataloader)) + + assert [next(iterator) for _ in range(3)] == [0, 1, 2] + assert dataloader.iter_count == 2 + + +def test_cycling_dataloader_cycles_multiple_epochs(): + dataloader = _EpochListLoader([2]) + iterator = iter(CyclingDataLoader(dataloader)) + + assert [next(iterator) for _ in range(5)] == [0, 1, 0, 1, 0] + assert dataloader.iter_count == 3 + + +def test_cycling_dataloader_rejects_empty_dataset(): + dataloader = _EpochListLoader([0]) + + with pytest.raises(RuntimeError, match="two consecutive epochs"): + next(iter(CyclingDataLoader(dataloader))) + assert dataloader.iter_count == 2 + + +def test_cycling_dataloader_delegates_checkpoint_state(): + dataloader = _EpochListLoader([1]) + + assert CyclingDataLoader(dataloader).state_dict() == {"iter_count": 0} diff --git a/tests/unit/reference_configs/ppo_math_1B_megatron.yaml b/tests/unit/reference_configs/ppo_math_1B_megatron.yaml index 5bbad873a7f..8ffb0081ff9 100644 --- a/tests/unit/reference_configs/ppo_math_1B_megatron.yaml +++ b/tests/unit/reference_configs/ppo_math_1B_megatron.yaml @@ -22,6 +22,14 @@ ppo: batch_multiplier: 1 skip_reference_policy_logprobs_calculation: true # No KL, so skip ref logprobs + async_ppo: + enabled: false + max_trajectory_age_steps: 1 + warmup_generation_lead_steps: null + in_flight_weight_updates: false + recompute_kv_cache_after_weight_updates: false + drop_incomplete_targets_on_restore: false + reward_shaping: enabled: true overlong_buffer_length: 2048 @@ -399,6 +407,7 @@ data: max_input_seq_length: 2048 shuffle: true num_workers: 1 + use_multiple_dataloader: false train: dataset_name: DAPOMath17K validation: diff --git a/tests/unit/single_controller/test_single_controller.py b/tests/unit/single_controller/test_single_controller.py index d03aa6c9821..fc663aaf35c 100644 --- a/tests/unit/single_controller/test_single_controller.py +++ b/tests/unit/single_controller/test_single_controller.py @@ -235,7 +235,7 @@ def test_sync_weights_honors_recompute_kv_cache_config( ctrl._rollout_manager = SimpleNamespace(set_weight_version=MagicMock()) ctrl._trainer_version = 3 ctrl._inflight_by_group_id = {} - # env={} -> _should_use_nemo_gym is False, so _sync_weights takes the native + # env={} -> should_use_nemo_gym is False, so _sync_weights takes the native # abort path (empty registry -> no-op) instead of the gym gate. ctrl._master_config = SimpleNamespace(env={}) @@ -264,7 +264,7 @@ def test_sync_weights_calibrates_and_forwards_fp8_kv_scales() -> None: ctrl._rollout_manager = SimpleNamespace(set_weight_version=MagicMock()) ctrl._trainer_version = 3 ctrl._inflight_by_group_id = {} - # env={} -> _should_use_nemo_gym is False, so _sync_weights takes the native + # env={} -> should_use_nemo_gym is False, so _sync_weights takes the native # abort path (empty registry -> no-op) instead of the gym gate. ctrl._master_config = SimpleNamespace(env={}) calibration_data = BatchedDataDict( diff --git a/tests/unit/single_controller/test_single_controller_setup.py b/tests/unit/single_controller/test_single_controller_setup.py index e720b28ac69..74f32e8b5c5 100644 --- a/tests/unit/single_controller/test_single_controller_setup.py +++ b/tests/unit/single_controller/test_single_controller_setup.py @@ -454,7 +454,7 @@ def test_megatron_train_iters_not_set_when_disabled(self, patched_factories): assert "train_iters" not in mc.policy.get("megatron_cfg", {}) def test_nemo_gym_wires_env_handle(self, patched_factories): - """When _should_use_nemo_gym is True the nemo-gym actor is spun up and stored.""" + """When should_use_nemo_gym is True the nemo-gym actor is spun up and stored.""" mc = _make_master_config(colocated=True, backend="vllm") mc.policy["generation"]["model_name"] = "test-model" mc.policy["generation"]["stop_strings"] = None @@ -467,7 +467,7 @@ def test_nemo_gym_wires_env_handle(self, patched_factories): fake_gym_actor = MagicMock(name="nemo_gym_actor") with ( - patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True), + patch.object(sc_setup_mod, "should_use_nemo_gym", return_value=True), patch.object( sc_setup_mod, "spinup_nemo_gym_actor", return_value=fake_gym_actor ) as mock_spinup, @@ -551,7 +551,7 @@ def test_nemo_gym_uses_deferred_vllm_load(self, patched_factories): patched_factories["setup_response_data"].return_value = (list(range(8)), None) with ( - patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True), + patch.object(sc_setup_mod, "should_use_nemo_gym", return_value=True), patch.object( sc_setup_mod, "spinup_nemo_gym_actor", return_value=MagicMock() ), @@ -577,7 +577,7 @@ def test_nemo_gym_records_timing_metrics(self, patched_factories): patched_factories["setup_response_data"].return_value = (list(range(8)), None) with ( - patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True), + patch.object(sc_setup_mod, "should_use_nemo_gym", return_value=True), patch.object( sc_setup_mod, "spinup_nemo_gym_actor", return_value=MagicMock() ), @@ -603,7 +603,7 @@ def test_nemo_gym_noncolocated_finishes_deferred_load(self, patched_factories): patched_factories["setup_response_data"].return_value = (list(range(8)), None) with ( - patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True), + patch.object(sc_setup_mod, "should_use_nemo_gym", return_value=True), patch.object( sc_setup_mod, "spinup_nemo_gym_actor", return_value=MagicMock() ), @@ -648,7 +648,7 @@ def test_nemo_gym_generation_init_time_includes_reserve_time( ) with ( - patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True), + patch.object(sc_setup_mod, "should_use_nemo_gym", return_value=True), patch.object( sc_setup_mod, "spinup_nemo_gym_actor", return_value=MagicMock() ), @@ -672,7 +672,7 @@ def test_nemo_gym_rejects_non_vllm_backend(self, patched_factories, backend): ) with ( - patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True), + patch.object(sc_setup_mod, "should_use_nemo_gym", return_value=True), patch.object(sc_setup_mod, "spinup_nemo_gym_actor") as mock_spinup, pytest.raises(NotImplementedError, match="vllm"), ): diff --git a/tests/unit/single_controller/test_train_pump.py b/tests/unit/single_controller/test_train_pump.py index 9a0b37dd8ea..1df434d9860 100644 --- a/tests/unit/single_controller/test_train_pump.py +++ b/tests/unit/single_controller/test_train_pump.py @@ -302,7 +302,7 @@ def test_train_pump_drives_mcore_training_step( master_config = MasterConfig.model_construct( policy={"train_global_batch_size": train_gbs}, - # _sync_weights gates stale-abort on _should_use_nemo_gym(env); empty + # _sync_weights gates stale-abort on should_use_nemo_gym(env); empty # env -> native path (nemo_gym disabled). env={}, grpo=GRPOConfig.model_construct( diff --git a/tests/unit/test_config_v2.py b/tests/unit/test_config_v2.py index a2c8aa17213..335f33db32c 100644 --- a/tests/unit/test_config_v2.py +++ b/tests/unit/test_config_v2.py @@ -27,7 +27,13 @@ from nemo_rl.algorithms.distillation import MasterConfig as DistillationMasterConfig from nemo_rl.algorithms.dpo import MasterConfig as DPOMasterConfig -from nemo_rl.algorithms.grpo import MasterConfig as GRPOMasterConfig +from nemo_rl.algorithms.grpo import ( + AsyncGRPOConfig, +) +from nemo_rl.algorithms.grpo import ( + MasterConfig as GRPOMasterConfig, +) +from nemo_rl.algorithms.ppo import AsyncPPOConfig, PPOConfig from nemo_rl.algorithms.ppo import MasterConfig as PPOMasterConfig from nemo_rl.algorithms.rm import MasterConfig as RMMasterConfig from nemo_rl.algorithms.sft import MasterConfig as SFTMasterConfig @@ -132,6 +138,11 @@ def test_config_v2_same_as_v1(config_file): master_config_class = RMMasterConfig config_v2 = master_config_class(**config_v1) + if "grpo" in config_v1 and "async_grpo" in config_v1["grpo"]: + assert isinstance(config_v2.grpo.async_grpo, AsyncGRPOConfig) + if "ppo" in config_v1 and "async_ppo" in config_v1["ppo"]: + assert isinstance(config_v2.ppo, PPOConfig) + assert isinstance(config_v2.ppo.async_ppo, AsyncPPOConfig) config_v2 = config_v2.model_dump() # Check v1 keys missing from v2, and differing values diff --git a/tests/unit/test_recipes_and_test_suites.py b/tests/unit/test_recipes_and_test_suites.py index 26eb39eb917..dd80daf6e30 100644 --- a/tests/unit/test_recipes_and_test_suites.py +++ b/tests/unit/test_recipes_and_test_suites.py @@ -256,7 +256,7 @@ def test_all_recipe_yamls_accounted_for_in_test_suites( ) -def test_nightly_compute_stays_below_4024_hours(nightly_test_suite, tracker): +def test_nightly_compute_stays_below_4048_hours(nightly_test_suite, tracker): command = f"DRYRUN=1 HF_HOME=... HF_DATASETS_CACHE=... CONTAINER= ACCOUNT= PARTITION= ./tools/launch {' '.join(nightly_test_suite)}" print(f"Running command: {command}") @@ -288,10 +288,10 @@ def test_nightly_compute_stays_below_4024_hours(nightly_test_suite, tracker): f"Last line of output was not as expected: '{last_line}'" ) total_gpu_hours = float(last_line.split(":")[-1].strip()) - # The managed Dynamo 3x8 H100 SWE1 test adds 96 GPU-hours to the former - # 3928-hour limit. - assert total_gpu_hours <= 4024, ( - f"Total GPU hours exceeded 4024: {last_line}. We should revisit the test suites to reduce the total GPU hours." + # Dynamo adds 96 GPU-hours and the two Async PPO nightlies add 24 to the + # former 3928-hour limit. + assert total_gpu_hours <= 4048, ( + f"Total GPU hours exceeded 4048: {last_line}. We should revisit the test suites to reduce the total GPU hours." ) tracker.track("total_nightly_gpu_hours", total_gpu_hours)