diff --git a/examples/configs/distillation_math.yaml b/examples/configs/distillation_math.yaml index 90e16000aa8..76ccf33e0ce 100644 --- a/examples/configs/distillation_math.yaml +++ b/examples/configs/distillation_math.yaml @@ -201,6 +201,9 @@ policy: &POLICY_BASE temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/examples/configs/evals/eval.yaml b/examples/configs/evals/eval.yaml index c492ac34cd1..820d503f27e 100644 --- a/examples/configs/evals/eval.yaml +++ b/examples/configs/evals/eval.yaml @@ -12,6 +12,9 @@ generation: temperature: 0.0 top_p: 1.0 top_k: -1 # -1 means disable + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} num_prompts_per_step: -1 # -1 means pass all prompts at once model_name: "Qwen/Qwen2.5-Math-1.5B-Instruct" stop_token_ids: null diff --git a/examples/configs/evals/mmau.yaml b/examples/configs/evals/mmau.yaml index e12c3ea0aec..7bc324d492f 100644 --- a/examples/configs/evals/mmau.yaml +++ b/examples/configs/evals/mmau.yaml @@ -11,6 +11,9 @@ generation: temperature: 0.0 top_p: 1.0 top_k: -1 + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} num_prompts_per_step: -1 model_name: "Qwen/Qwen2.5-Omni-3B" stop_token_ids: null diff --git a/examples/configs/grpo_math_1B.yaml b/examples/configs/grpo_math_1B.yaml index c4682a2c976..062abddaa4c 100644 --- a/examples/configs/grpo_math_1B.yaml +++ b/examples/configs/grpo_math_1B.yaml @@ -15,6 +15,9 @@ grpo: advantage_clip_low: null advantage_clip_high: null max_val_samples: 256 + # Validation rollouts per prompt; k > 1 also reports the pass_k metric. + # max_val_samples counts PROMPTS: total validation rollouts = max_val_samples * k. + val_num_generations_per_prompt: 1 # Early stop once this metric (e.g. accuracy or pass_k) reaches the threshold; null disables. stop_at_validation_metric: null # Required when stop_at_validation_metric is set. @@ -344,6 +347,10 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + # Validation-only sampling; defaults follow the train values above. + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # null = topology default (IPC colocated, NCCL non-colocated). diff --git a/examples/configs/ppo_math_1B.yaml b/examples/configs/ppo_math_1B.yaml index 1d79963534b..2cfb53e3263 100644 --- a/examples/configs/ppo_math_1B.yaml +++ b/examples/configs/ppo_math_1B.yaml @@ -234,6 +234,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null mcore_generation_config: diff --git a/examples/configs/recipes/llm/mopd-qwen3-1.7b-3n8g-megatron-pack.yaml b/examples/configs/recipes/llm/mopd-qwen3-1.7b-3n8g-megatron-pack.yaml index d967ea86519..e8be20b666c 100644 --- a/examples/configs/recipes/llm/mopd-qwen3-1.7b-3n8g-megatron-pack.yaml +++ b/examples/configs/recipes/llm/mopd-qwen3-1.7b-3n8g-megatron-pack.yaml @@ -2,7 +2,6 @@ defaults: ../../grpo_math_1B.yaml grpo: num_prompts_per_step: 8 num_generations_per_prompt: 4 - num_val_generations_per_prompt: 1 max_num_steps: 5 val_period: 1000 overlong_filtering: true diff --git a/examples/nemo_gym/grpo_nanov3.yaml b/examples/nemo_gym/grpo_nanov3.yaml index 8fa3175d9e4..103d470f610 100644 --- a/examples/nemo_gym/grpo_nanov3.yaml +++ b/examples/nemo_gym/grpo_nanov3.yaml @@ -2,7 +2,7 @@ grpo: num_prompts_per_step: 128 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 4 + val_num_generations_per_prompt: 4 max_rollout_turns: 1 # for multi-turn rollouts. Math Environments just have 1 turn (answering the question) max_num_epochs: 1 max_num_steps: 1000000 @@ -210,6 +210,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null mcore_generation_config: diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml index 48f284cb711..c19fb351b65 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe1.yaml @@ -13,7 +13,7 @@ checkpointing: grpo: num_prompts_per_step: 64 num_generations_per_prompt: 8 - num_val_generations_per_prompt: 1 + val_num_generations_per_prompt: 1 max_num_epochs: 100 advantage_clip_low: -100 advantage_clip_high: 100 diff --git a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml index 6764b43126d..449b873e5f2 100644 --- a/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml +++ b/examples/nemo_gym/grpo_qwen3_30ba3b_thinking_swe2.yaml @@ -13,7 +13,7 @@ checkpointing: grpo: num_prompts_per_step: 8 num_generations_per_prompt: 8 - num_val_generations_per_prompt: 1 + val_num_generations_per_prompt: 1 max_num_epochs: 100 advantage_clip_low: -100 advantage_clip_high: 100 diff --git a/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml b/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml index e7b78207ff4..b3641dd5437 100644 --- a/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml +++ b/examples/nemo_gym/grpo_workplace_assistant_nemotron_nano_v2_9b.yaml @@ -14,6 +14,7 @@ grpo: advantage_clip_low: null advantage_clip_high: null max_val_samples: null # inferred from size of val dataset. for multi evals, repeat val ds via `num_repeats` in `ng_prepare_data`. + val_num_generations_per_prompt: 1 # Early stop once this metric (e.g. accuracy or pass_k) reaches the threshold; null disables. stop_at_validation_metric: null # Required when stop_at_validation_metric is set. @@ -226,6 +227,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml b/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml index 1b2d300dc26..03317acd30e 100644 --- a/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage1_rlvr.yaml @@ -13,7 +13,7 @@ checkpointing: grpo: num_prompts_per_step: 256 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -226,6 +226,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml b/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml index 59c41c3d0b2..8a8d52ca042 100644 --- a/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage2_swe1.yaml @@ -13,7 +13,7 @@ checkpointing: grpo: num_prompts_per_step: 64 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 1 + val_num_generations_per_prompt: 1 max_rollout_turns: 1 max_num_epochs: 100 max_num_steps: 1000000 @@ -226,6 +226,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml b/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml index 72a22c7913d..4c5284f5aae 100644 --- a/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage2_swe2.yaml @@ -13,7 +13,7 @@ checkpointing: grpo: num_prompts_per_step: 16 num_generations_per_prompt: 32 - num_val_generations_per_prompt: 1 + val_num_generations_per_prompt: 1 max_rollout_turns: 1 max_num_epochs: 100 max_num_steps: 1000000 @@ -219,6 +219,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml b/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml index e56c770a2b4..f2d2b38930f 100644 --- a/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml +++ b/examples/nemo_gym/nemotron-3-super/stage3_rlhf.yaml @@ -13,7 +13,7 @@ checkpointing: grpo: num_prompts_per_step: 128 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -226,6 +226,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml index 026824cf562..18dc57404a5 100644 --- a/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/ifbench_teacher.yaml @@ -47,7 +47,7 @@ checkpointing: grpo: num_prompts_per_step: 128 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -305,6 +305,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # TP=8 EP=8: EP=TP so vllm_dp_size=1, async_engine=true works with NeMo Gym. diff --git a/examples/nemo_gym/nemotron-3-ultra/mopd.yaml b/examples/nemo_gym/nemotron-3-ultra/mopd.yaml index fe85213fdd4..203bf03e83f 100644 --- a/examples/nemo_gym/nemotron-3-ultra/mopd.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/mopd.yaml @@ -62,7 +62,7 @@ checkpointing: grpo: num_prompts_per_step: 1024 num_generations_per_prompt: 1 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -326,6 +326,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # TP=8 EP=8: EP=TP so vllm_dp_size=1, async_engine=true works with NeMo Gym. diff --git a/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml index b32979ce39a..e62f261e867 100644 --- a/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/reasoning_teacher.yaml @@ -51,7 +51,7 @@ checkpointing: grpo: num_prompts_per_step: 128 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 10 max_num_steps: 1000000 @@ -308,6 +308,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # TP=8 EP=8: EP=TP so vllm_dp_size=1, async_engine=true works with NeMo Gym. diff --git a/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml index 32fc72deb4b..d757cde7a0e 100644 --- a/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/rlhf_teacher.yaml @@ -48,7 +48,7 @@ checkpointing: grpo: num_prompts_per_step: 128 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -306,6 +306,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # TP=8 EP=8: EP=TP so vllm_dp_size=1, async_engine=true works with NeMo Gym. diff --git a/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml b/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml index 03e43b694dc..f6c3b308ad5 100644 --- a/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/student_rlvr1.yaml @@ -44,7 +44,7 @@ checkpointing: grpo: num_prompts_per_step: 512 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -302,6 +302,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # TP=8 EP=8: EP=TP so vllm_dp_size=1, async_engine=true works with NeMo Gym. diff --git a/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml b/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml index 3a9fbae3712..262d1679cd8 100644 --- a/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/student_rlvr2.yaml @@ -45,7 +45,7 @@ checkpointing: grpo: num_prompts_per_step: 512 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 1 max_num_steps: 1000000 @@ -303,6 +303,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # TP=8 EP=8: EP=TP so vllm_dp_size=1, async_engine=true works with NeMo Gym. diff --git a/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml b/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml index c8db114be15..75224b0e556 100644 --- a/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml +++ b/examples/nemo_gym/nemotron-3-ultra/swe_teacher.yaml @@ -67,7 +67,7 @@ checkpointing: grpo: num_prompts_per_step: 32 num_generations_per_prompt: 16 - num_val_generations_per_prompt: 2 + val_num_generations_per_prompt: 2 max_rollout_turns: 1 max_num_epochs: 4 max_num_steps: 1000000 @@ -325,6 +325,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 4159e5d5b60..ab4cdaea80a 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -98,6 +98,7 @@ from nemo_rl.models.generation.interfaces import ( GenerationConfig, GenerationInterface, + GenerationSamplingParams, resolve_routed_experts_dtype_name_for_model, ) from nemo_rl.models.generation.megatron import MegatronGeneration @@ -246,7 +247,13 @@ class GRPOConfig(BaseModel, extra="allow"): # 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 = False + # Counts PROMPTS, not rollouts: with val_num_generations_per_prompt = k, + # total validation rollouts = max_val_samples * k. max_val_samples: int | None = 256 # None for NeMo-Gym compatibility + # Number of independent validation rollouts generated for each prompt; + # k > 1 additionally reports pass@k over each prompt's k rollouts as the + # pass_k metric. + val_num_generations_per_prompt: int = 1 # Early stop: end training once this validation metric (e.g. accuracy, # always reported, or pass_k with grouped validation) reaches # stop_at_validation_threshold; null disables early stopping. @@ -398,6 +405,40 @@ def setup( if generation_config["backend"] == "vllm": normalize_vllm_refit_config(cast(VllmConfig, generation_config)) + # Validation-only sampling is honored only on the NeMo-Gym vLLM rollout + # path; everywhere else validation must sample exactly like training. + val_sampling_overridden = ( + generation_config["val_temperature"] != generation_config["temperature"] + or generation_config["val_top_p"] != generation_config["top_p"] + or generation_config["val_top_k"] != generation_config["top_k"] + ) + if val_sampling_overridden: + assert generation_config["backend"] == "vllm" and _should_use_nemo_gym( + master_config + ), ( + "generation.val_temperature/val_top_p/val_top_k differing from the " + "train sampling params is only supported for vLLM NeMo-Gym rollouts." + ) + # The NeMo-Gym path only stamps temperature/top_p onto requests and + # rejects any top_k at rollout time, so a val_top_k override can never + # be honored — fail here instead of at the first validation step. + assert not generation_config["val_top_k"], ( + "generation.val_top_k is not supported: the NeMo-Gym rollout path " + "only honors val_temperature/val_top_p. Leave val_top_k null." + ) + assert grpo_config.val_num_generations_per_prompt >= 1, ( + "grpo.val_num_generations_per_prompt must be >= 1" + ) + # pass_k is only reported when k > 1; catch the mismatch here instead of + # at the first validation step. + assert not ( + grpo_config.stop_at_validation_metric == "pass_k" + and grpo_config.val_num_generations_per_prompt <= 1 + ), ( + "grpo.stop_at_validation_metric='pass_k' requires " + "grpo.val_num_generations_per_prompt > 1" + ) + # Set seed for all random number generators set_seed(grpo_config.seed) @@ -3681,6 +3722,10 @@ def validate( timer = Timer(context={"worker": "validator"}) with timer.time("total_validation_time"): print(f"▶ Starting validation at step {step}...", flush=True) + # >= 1 is validated in setup(). + val_num_generations_per_prompt = ( + master_config.grpo.val_num_generations_per_prompt + ) total_rewards = [] total_lengths = [] @@ -3693,12 +3738,23 @@ def validate( if batch_idx >= max_batches: break + if val_num_generations_per_prompt > 1: + val_batch = val_batch.repeat_interleave(val_num_generations_per_prompt) + additional_metrics_to_report = dict() # 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): generation_config = master_config.policy["generation"] + # Validation-only sampling (e.g. near-greedy validation); + # defaults to the train profile via the exemplar YAML + # interpolations. Training rollouts keep policy.generation. + val_sampling_params = GenerationSamplingParams( + temperature=generation_config["val_temperature"], + top_p=generation_config["val_top_p"], + top_k=generation_config["val_top_k"], + ) nemo_gym_rollout_result = run_nemo_gym_rollout_sync( policy_generation=policy_generation, input_batch=val_batch, @@ -3706,6 +3762,7 @@ def validate( task_to_env=val_task_to_env, max_seq_len=master_config.policy["max_total_sequence_length"], generation_config=generation_config, + sampling_params=val_sampling_params, log_full_result_tables=should_log_nemo_gym_full_result_tables( wandb_enabled=master_config.logger["wandb_enabled"], wandb_config=master_config.logger["wandb"], @@ -3756,11 +3813,26 @@ def validate( all_message_logs.extend(to_env) - # Calculate validation metrics + # Calculate validation metrics. accuracy is the mean reward over all + # rollouts; grouped validation (val_num_generations_per_prompt > 1) + # additionally reports pass@k over each prompt's k rollouts as pass_k. num_samples = len(total_rewards) + pass_k = None if num_samples > 0: rewards_t = torch.tensor(total_rewards, dtype=torch.float32) accuracy = rewards_t.mean().item() + if val_num_generations_per_prompt > 1: + assert num_samples % val_num_generations_per_prompt == 0, ( + "Validation rewards must be divisible by " + "grpo.val_num_generations_per_prompt" + ) + pass_k = ( + (rewards_t.view(-1, val_num_generations_per_prompt) > 0) + .any(dim=1) + .float() + .mean() + .item() + ) else: accuracy = 0.0 @@ -3773,6 +3845,8 @@ def validate( "avg_length": avg_length, **additional_metrics_to_report, } + if pass_k is not None: + val_metrics["pass_k"] = pass_k # Print sample conversations only once at the end of validation try: diff --git a/nemo_rl/experience/rollouts.py b/nemo_rl/experience/rollouts.py index 1a6145160b8..84ff46b13d3 100644 --- a/nemo_rl/experience/rollouts.py +++ b/nemo_rl/experience/rollouts.py @@ -55,6 +55,7 @@ GenerationDatumSpec, GenerationInterface, GenerationOutputSpec, + GenerationSamplingParams, ) from nemo_rl.utils.timer import Timer @@ -2007,7 +2008,9 @@ def apply_reward_penalties( def _prepare_nemo_gym_rows( - rows: list[dict], generation_config: GenerationConfig + rows: list[dict], + generation_config: GenerationConfig, + sampling_params: GenerationSamplingParams, ) -> None: """Apply NeMo-RL sampling parameters and stable row indices in place.""" for row_index, row in enumerate(rows): @@ -2017,8 +2020,8 @@ def _prepare_nemo_gym_rows( "Each NeMo-Gym row must contain a responses_create_params dict" ) - responses_create_params["temperature"] = generation_config["temperature"] - responses_create_params["top_p"] = generation_config["top_p"] + responses_create_params["temperature"] = sampling_params.temperature + responses_create_params["top_p"] = sampling_params.top_p configured_max_tokens = generation_config["max_new_tokens"] row_max_tokens = responses_create_params.get("max_output_tokens") responses_create_params["max_output_tokens"] = ( @@ -2059,6 +2062,7 @@ async def run_async_nemo_gym_rollout( thinking_tags: list[str] | tuple[str, ...] | None = None, mask_env_flagged_samples: bool = True, returns_entire_batch: bool = False, + sampling_params: Optional[GenerationSamplingParams] = None, ) -> AsyncGenerator[NemoGymRolloutResult, None]: """Stream complete NeMo-Gym prompt groups in group-completion order. @@ -2089,6 +2093,9 @@ async def run_async_nemo_gym_rollout( returns_entire_batch: Whether to treat the input as one potentially heterogeneous group. This requires ``num_generations`` to equal the batch size and is used by synchronous callers. + sampling_params: Sampling profile stamped onto every NeMo-Gym row. + ``None`` uses the train profile from ``generation_config``; + validation passes its own profile explicitly. Yields: ``NemoGymRolloutResult`` objects in prompt-group completion order. Rows @@ -2139,9 +2146,13 @@ async def run_async_nemo_gym_rollout( assert not generation_config["stop_token_ids"], ( "Stop strings is not supported in the generation config in NeMo-Gym path!" ) + if sampling_params is None: + sampling_params = GenerationSamplingParams.from_generation_config( + generation_config + ) # Top k is not OpenAI compatible, so NeMo-Gym does not guarantee support over it. - assert not generation_config["top_k"], ( - "Top k is not supported in the generation config in NeMo-Gym path!" + assert not sampling_params.top_k, ( + "Top k is not supported in the sampling params in NeMo-Gym path!" ) if num_generations <= 0: raise ValueError("num_generations must be greater than zero") @@ -2162,7 +2173,7 @@ async def run_async_nemo_gym_rollout( run_rollouts_timer_label = f"{timer_prefix}/run_rollouts" with timer.time(total_timer_label): - _prepare_nemo_gym_rows(nemo_gym_rows, generation_config) + _prepare_nemo_gym_rows(nemo_gym_rows, generation_config, sampling_params) accumulator = _NemoGymStreamAccumulator( rows=nemo_gym_rows, num_generations=num_generations, @@ -2248,6 +2259,7 @@ def run_nemo_gym_rollout_sync( effort_config: Optional[EffortLevelsConfig] = None, reward_penalty_config: dict[str, Any] | BaseModel | None = None, thinking_tags: list[str] | tuple[str, ...] | None = None, + sampling_params: Optional[GenerationSamplingParams] = None, mask_env_flagged_samples: bool = True, ) -> NemoGymRolloutResult: """Run and return one complete NeMo-Gym batch synchronously. @@ -2271,6 +2283,9 @@ def run_nemo_gym_rollout_sync( effort_config: Optional configuration for effort-based reward shaping. reward_penalty_config: Optional reward-penalty configuration. thinking_tags: Optional opening and closing tags used by thinking penalties. + sampling_params: Sampling profile stamped onto every NeMo-Gym row. + ``None`` uses the train profile from ``generation_config``; + validation passes its own profile explicitly. mask_env_flagged_samples: Whether to carry env-driven ``mask_sample`` flags in the rollout batch for loss masking. @@ -2304,6 +2319,7 @@ async def _consume_rollout() -> NemoGymRolloutResult: thinking_tags=thinking_tags, mask_env_flagged_samples=mask_env_flagged_samples, returns_entire_batch=True, + sampling_params=sampling_params, ): pass if rollout_result is None: diff --git a/nemo_rl/models/generation/interfaces.py b/nemo_rl/models/generation/interfaces.py index a757ad3a8ff..791657c394f 100644 --- a/nemo_rl/models/generation/interfaces.py +++ b/nemo_rl/models/generation/interfaces.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. from abc import ABC, abstractmethod +from dataclasses import dataclass from typing import Any, NotRequired, Optional, TypedDict, Union import ray @@ -201,6 +202,13 @@ class GenerationConfig(TypedDict): temperature: float top_p: float top_k: int | None + # Validation-only sampling. The exemplar YAMLs default these to the train + # values above via interpolation (${.temperature}, ...), so validation + # samples exactly like training unless overridden. Only honored on the + # NeMo-Gym vLLM rollout path (guarded in grpo.setup()). + val_temperature: float + val_top_p: float + val_top_k: int | None model_name: NotRequired[str] # Not Required b/c GRPO writes this stop_token_ids: list[int] | None stop_strings: list[str] | None @@ -215,6 +223,33 @@ class GenerationConfig(TypedDict): _mtp_weights_from_refit: NotRequired[bool] +@dataclass +class GenerationSamplingParams: + """Sampling profile threaded explicitly through rollout entry points. + + Rollout callers construct one from the relevant ``GenerationConfig`` + fields (train or validation) so the sampling used for a rollout is + visible at the call site instead of flowing through config side-channels. + Named to distinguish it from ``TrainingSamplingParams`` (train-time logit + filtering) and vLLM's own ``SamplingParams``. + """ + + temperature: float + top_p: float + top_k: int | None + + @classmethod + def from_generation_config( + cls, generation_config: "GenerationConfig" + ) -> "GenerationSamplingParams": + """Build the train-time sampling profile from a generation config.""" + return cls( + temperature=generation_config["temperature"], + top_p=generation_config["top_p"], + top_k=generation_config["top_k"], + ) + + class GenerationDatumSpec(TypedDict): """Specification for input data required by generation models. diff --git a/nemo_rl/models/generation/vllm/vllm_worker_async.py b/nemo_rl/models/generation/vllm/vllm_worker_async.py index 201d0790044..d1da8a3ec1a 100644 --- a/nemo_rl/models/generation/vllm/vllm_worker_async.py +++ b/nemo_rl/models/generation/vllm/vllm_worker_async.py @@ -678,9 +678,38 @@ async def create_chat_completion( request.top_k = -1 # The request sampling params need to exactly match those as are set in NeMo RL. - # If they do not match, the inference will be off policy and destroy training stability. - assert request.temperature == generation_config["temperature"] - assert request.top_p == generation_config["top_p"] + # If they do not match, the inference will be off policy and destroy training + # stability. Validation rollouts are the one exception: they are stamped with + # the validation sampling profile (generation.val_temperature / val_top_p), + # which is metric-only and safe to serve — grpo.validate() is the only + # caller that constructs a non-train GenerationSamplingParams. Multi-turn + # agents issue their own requests, so this server-side check is the one + # chokepoint they all pass. + # vLLM resolves an unset top_p from the model's generation_config.json + # (ModelConfig.generation_config defaults to "auto"), NOT to 1.0, so a + # request omitting it would sample off-policy while passing this check. + assert request.top_p is not None, ( + "top_p must be set explicitly on NeMo-RL requests; an unset top_p is " + "resolved by vLLM from the model's generation_config.json and would " + "bypass the on-policy sampling check." + ) + request_top_p = request.top_p + is_train_sampling = ( + request.temperature == generation_config["temperature"] + and request_top_p == generation_config["top_p"] + ) + is_val_sampling = ( + request.temperature == generation_config["val_temperature"] + and request_top_p == generation_config["val_top_p"] + ) + assert is_train_sampling or is_val_sampling, ( + f"request sampling (temperature={request.temperature}, " + f"top_p={request.top_p}) matches neither the train sampling params " + f"(temperature={generation_config['temperature']}, " + f"top_p={generation_config['top_p']}) nor the validation sampling " + f"params (val_temperature={generation_config['val_temperature']}, " + f"val_top_p={generation_config['val_top_p']})" + ) try: generator = await openai_serving_chat.create_chat_completion( diff --git a/research/template_project/configs/grpo_math_1B.yaml b/research/template_project/configs/grpo_math_1B.yaml index 2862ad17e02..9d7b2b86296 100644 --- a/research/template_project/configs/grpo_math_1B.yaml +++ b/research/template_project/configs/grpo_math_1B.yaml @@ -13,6 +13,8 @@ grpo: val_at_end: false overlong_filtering: false max_val_samples: 256 + # Validation rollouts per prompt; k > 1 also reports the pass_k metric. + val_num_generations_per_prompt: 1 # Early stop once this metric (e.g. accuracy or pass_k) reaches the threshold; null disables. stop_at_validation_metric: null # Required when stop_at_validation_metric is set. @@ -286,6 +288,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null mcore_generation_config: diff --git a/tests/unit/algorithms/test_distillation.py b/tests/unit/algorithms/test_distillation.py index dfda02093b8..df9188c02b2 100644 --- a/tests/unit/algorithms/test_distillation.py +++ b/tests/unit/algorithms/test_distillation.py @@ -150,6 +150,9 @@ def val_iter(self): "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "colocated": { "enabled": False, }, @@ -851,6 +854,9 @@ def test_noncolocated_inference_requires_explicit_gpus_per_node_single_node(): "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "backend": "vllm", "colocated": { "enabled": False, # Non-colocated @@ -926,6 +932,9 @@ def test_distillation_setup_non_colocated_smoke(monkeypatch, refit_transport): "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "backend": "vllm", "refit_transport": refit_transport, "refit_cfg": None, @@ -1110,6 +1119,9 @@ def test_distillation_setup_nemo_gym_uses_deferred_vllm( "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "backend": "vllm", "vllm_kwargs": {}, "vllm_cfg": { @@ -1363,6 +1375,9 @@ def test_noncolocated_inference_requires_explicit_gpus_per_node_multi_node(): "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "backend": "vllm", "colocated": { "enabled": False, # Non-colocated diff --git a/tests/unit/algorithms/test_grpo.py b/tests/unit/algorithms/test_grpo.py index dee6cdb260b..479e5479919 100644 --- a/tests/unit/algorithms/test_grpo.py +++ b/tests/unit/algorithms/test_grpo.py @@ -50,6 +50,7 @@ dynamic_sampling, grpo_train, refit_policy_generation, + setup, validate, ) from nemo_rl.algorithms.grpo_sync import _train_fields_for_step, grpo_train_sync @@ -269,6 +270,7 @@ def val_iter(self): max_rollout_turns=1, val_period=100, val_start_at=-1, + val_num_generations_per_prompt=1, val_batch_size=1, val_at_start=False, val_at_end=False, @@ -305,6 +307,9 @@ def val_iter(self): "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "backend": "vllm", "colocated": {"enabled": True}, "vllm_cfg": {"async_engine": True}, # Support async mode @@ -3773,6 +3778,141 @@ def test_validate_works_without_logger(self, mock_grpo_components): assert "accuracy" in val_metrics assert "avg_length" in val_metrics + def test_grouped_validation_reports_pass_k(self, mock_grpo_components): + mock_batch = BatchedDataDict[DatumSpec]( + { + "message_log": [ + [{"role": "user", "content": "a", "token_ids": torch.tensor([1])}], + [{"role": "user", "content": "b", "token_ids": torch.tensor([2])}], + ], + "task_name": ["math", "math"], + "extra_env_info": [{}, {}], + "loss_multiplier": torch.tensor([1.0, 1.0]), + "idx": torch.tensor([0, 1]), + "length": torch.tensor([1, 1]), + "total_reward": torch.tensor([0.0, 0.0]), + } + ) + mock_dataloader = MagicMock(spec=StatefulDataLoader) + mock_dataloader.__iter__ = MagicMock(return_value=iter([mock_batch])) + mock_config = mock_grpo_components["master_config"] + mock_config.grpo.max_val_samples = 2 + mock_config.grpo.val_batch_size = 2 + mock_config.grpo.val_num_generations_per_prompt = 4 + + def run_rollout(_policy, repeated_batch, *_args, **_kwargs): + # Each prompt is repeated k=4 times, contiguously. + assert repeated_batch["idx"].tolist() == [0, 0, 0, 0, 1, 1, 1, 1] + # Prompt 0 passes once out of 4; prompt 1 never passes. + repeated_batch["total_reward"] = torch.tensor( + [1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0] + ) + return repeated_batch, {"mean_gen_tokens_per_sample": 1.0} + + with ( + patch( + "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_async_rollouts", + return_value=False, + ), + patch("nemo_rl.algorithms.grpo.print_message_log_samples"), + ): + val_metrics, _ = validate( + MagicMock(), + mock_dataloader, + MagicMock(), + {"math": MagicMock(spec=EnvironmentInterface)}, + step=0, + master_config=mock_config, + ) + + # accuracy stays the plain mean over all 8 rollouts; pass@4 counts + # prompts with at least one passing rollout (1 of 2). + assert val_metrics["accuracy"] == pytest.approx(0.125) + assert val_metrics["pass_k"] == pytest.approx(0.5) + + def test_validation_uses_val_sampling_params_on_gym_path( + self, mock_grpo_components + ): + mock_batch = BatchedDataDict[DatumSpec]( + { + "message_log": [ + [{"role": "user", "content": "a", "token_ids": torch.tensor([1])}], + [{"role": "user", "content": "b", "token_ids": torch.tensor([2])}], + ], + "task_name": ["math", "math"], + "extra_env_info": [{}, {}], + "loss_multiplier": torch.tensor([1.0, 1.0]), + "idx": torch.tensor([0, 1]), + "length": torch.tensor([1, 1]), + "total_reward": torch.tensor([0.0, 0.0]), + } + ) + mock_dataloader = MagicMock(spec=StatefulDataLoader) + mock_dataloader.__iter__ = MagicMock(return_value=iter([mock_batch])) + mock_config = mock_grpo_components["master_config"] + mock_config.grpo.max_val_samples = 2 + mock_config.grpo.val_batch_size = 2 + mock_config.grpo.val_num_generations_per_prompt = 2 + # Train samples at 1.0/1.0; validation runs near-greedy. + mock_config.policy["generation"].update( + {"val_temperature": 0.1, "val_top_p": 0.9, "val_top_k": None} + ) + mock_config.logger.update({"wandb_enabled": False, "wandb": {}}) + mock_config.env = {} + + def run_gym_rollout(**kwargs): + repeated_batch = kwargs["input_batch"] + # 2 prompts x k=2 validation rollouts, contiguous per prompt. + assert repeated_batch["idx"].tolist() == [0, 0, 1, 1] + repeated_batch["total_reward"] = torch.tensor([1.0, 0.0, 0.0, 0.0]) + return MagicMock( + final_batch=repeated_batch, + rollout_metrics={"mean_gen_tokens_per_sample": 1.0}, + ) + + with ( + patch( + "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.print_message_log_samples"), + ): + val_metrics, _ = validate( + MagicMock(), + mock_dataloader, + MagicMock(), + {"math": MagicMock(spec=EnvironmentInterface)}, + step=0, + master_config=mock_config, + ) + + sampling_params = mock_rollout.call_args.kwargs["sampling_params"] + assert sampling_params.temperature == pytest.approx(0.1) + assert sampling_params.top_p == pytest.approx(0.9) + assert sampling_params.top_k is None + assert val_metrics["accuracy"] == pytest.approx(0.25) + assert val_metrics["pass_k"] == pytest.approx(0.5) + + def test_setup_rejects_val_sampling_outside_gym_vllm_path( + self, mock_grpo_components + ): + master_config = mock_grpo_components["master_config"] + # Non-gym rollouts (env has no nemo_gym) with validation sampling + # different from training must be rejected at setup time. + master_config.policy["generation"].update( + {"backend": "megatron", "val_temperature": 0.1} + ) + master_config.env = {} + + with pytest.raises(AssertionError, match="only supported for vLLM NeMo-Gym"): + setup(master_config, MagicMock(), MagicMock(), None) + def test_validate_returns_empty_when_no_dataloader(self, mock_grpo_components): """Test that validate returns empty dicts when no dataloader is provided.""" mock_policy_gen = MagicMock() diff --git a/tests/unit/environments/test_code_environment.py b/tests/unit/environments/test_code_environment.py index d32550aba1e..cfdb96cc089 100644 --- a/tests/unit/environments/test_code_environment.py +++ b/tests/unit/environments/test_code_environment.py @@ -46,6 +46,9 @@ "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "vllm_cfg": { diff --git a/tests/unit/environments/test_retriever.py b/tests/unit/environments/test_retriever.py index c9413e67590..9e0c27e62f1 100644 --- a/tests/unit/environments/test_retriever.py +++ b/tests/unit/environments/test_retriever.py @@ -45,6 +45,9 @@ "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "vllm_cfg": { diff --git a/tests/unit/experience/test_rollouts.py b/tests/unit/experience/test_rollouts.py index 038384a5958..40676612602 100644 --- a/tests/unit/experience/test_rollouts.py +++ b/tests/unit/experience/test_rollouts.py @@ -357,6 +357,9 @@ def initial_multi_step_calculator_batch(rollout_tokenizer): "temperature": 0.01, # Near-greedy "top_p": 1.0, "top_k": None, + "val_temperature": 0.01, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "vllm_cfg": { @@ -1188,6 +1191,9 @@ async def _collect(): "top_k": None, "temperature": 1.0, "top_p": 1.0, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "max_new_tokens": 32, }, num_generations=2, diff --git a/tests/unit/models/generation/test_vllm_generation.py b/tests/unit/models/generation/test_vllm_generation.py index 0821f0865b6..93429ccfb5e 100644 --- a/tests/unit/models/generation/test_vllm_generation.py +++ b/tests/unit/models/generation/test_vllm_generation.py @@ -67,6 +67,9 @@ "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "vllm_cfg": { diff --git a/tests/unit/models/generation/test_vllm_large_model.py b/tests/unit/models/generation/test_vllm_large_model.py index 89eaece234c..9b18a446ebf 100644 --- a/tests/unit/models/generation/test_vllm_large_model.py +++ b/tests/unit/models/generation/test_vllm_large_model.py @@ -38,6 +38,9 @@ "temperature": 0.8, "top_p": 1.0, "top_k": None, + "val_temperature": 0.8, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "vllm_cfg": { diff --git a/tests/unit/models/generation/test_vllm_quant_backend.py b/tests/unit/models/generation/test_vllm_quant_backend.py index 6b8bd6487c4..2cb80268661 100644 --- a/tests/unit/models/generation/test_vllm_quant_backend.py +++ b/tests/unit/models/generation/test_vllm_quant_backend.py @@ -57,6 +57,9 @@ def _make_vllm_config(tokenizer, *, async_engine=False, is_eval=True): "temperature": 0.0, "top_p": 1.0, "top_k": None, + "val_temperature": 0.0, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "quant_cfg": _QUANT_CFG, diff --git a/tests/unit/models/generation/trtllm/test_trtllm_generation.py b/tests/unit/models/generation/trtllm/test_trtllm_generation.py index d9dc1eaea70..1123ebae58a 100644 --- a/tests/unit/models/generation/trtllm/test_trtllm_generation.py +++ b/tests/unit/models/generation/trtllm/test_trtllm_generation.py @@ -43,6 +43,9 @@ def _config(**trtllm_overrides): "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "stop_token_ids": None, "stop_strings": None, "_pad_token_id": 0, diff --git a/tests/unit/reference_configs/distillation_math.yaml b/tests/unit/reference_configs/distillation_math.yaml index 5b14bd9bb8b..9fe04653393 100644 --- a/tests/unit/reference_configs/distillation_math.yaml +++ b/tests/unit/reference_configs/distillation_math.yaml @@ -191,6 +191,9 @@ policy: &POLICY_BASE temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null vllm_cfg: diff --git a/tests/unit/reference_configs/eval.yaml b/tests/unit/reference_configs/eval.yaml index abe20f4d74d..95813a0a433 100644 --- a/tests/unit/reference_configs/eval.yaml +++ b/tests/unit/reference_configs/eval.yaml @@ -12,6 +12,9 @@ generation: temperature: 0.0 top_p: 1.0 top_k: -1 # -1 means disable + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} num_prompts_per_step: -1 # -1 means pass all prompts at once model_name: "Qwen/Qwen2.5-Math-1.5B-Instruct" stop_token_ids: null diff --git a/tests/unit/reference_configs/grpo_math_1B.yaml b/tests/unit/reference_configs/grpo_math_1B.yaml index 5f4ffb2270d..7f026fca478 100644 --- a/tests/unit/reference_configs/grpo_math_1B.yaml +++ b/tests/unit/reference_configs/grpo_math_1B.yaml @@ -15,6 +15,7 @@ grpo: advantage_clip_low: null advantage_clip_high: null max_val_samples: 256 + val_num_generations_per_prompt: 1 # Early stop once this metric (e.g. accuracy or pass_k) reaches the threshold; null disables. stop_at_validation_metric: null # Required when stop_at_validation_metric is set. @@ -340,6 +341,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null # null = topology default (IPC colocated, NCCL non-colocated). diff --git a/tests/unit/reference_configs/ppo_math_1B_megatron.yaml b/tests/unit/reference_configs/ppo_math_1B_megatron.yaml index 3610bfd09e8..12ebd6c7cee 100644 --- a/tests/unit/reference_configs/ppo_math_1B_megatron.yaml +++ b/tests/unit/reference_configs/ppo_math_1B_megatron.yaml @@ -222,6 +222,9 @@ policy: temperature: 1.0 top_p: 1.0 top_k: null + val_temperature: ${.temperature} + val_top_p: ${.top_p} + val_top_k: ${.top_k} stop_token_ids: null stop_strings: null mcore_generation_config: diff --git a/tests/unit/utils/test_native_checkpoint.py b/tests/unit/utils/test_native_checkpoint.py index 33240d92888..ba7b1151d64 100755 --- a/tests/unit/utils/test_native_checkpoint.py +++ b/tests/unit/utils/test_native_checkpoint.py @@ -73,6 +73,9 @@ "temperature": 1.0, "top_p": 1.0, "top_k": None, + "val_temperature": 1.0, + "val_top_p": 1.0, + "val_top_k": None, "backend": "vllm", "colocated": {"enabled": True}, },