diff --git a/nemo_rl/algorithms/async_utils/trajectory_collector.py b/nemo_rl/algorithms/async_utils/trajectory_collector.py index dd0ff5bcd94..8d58d61d3ec 100644 --- a/nemo_rl/algorithms/async_utils/trajectory_collector.py +++ b/nemo_rl/algorithms/async_utils/trajectory_collector.py @@ -30,6 +30,7 @@ from nemo_rl.algorithms.grpo import ( AsyncGRPOConfig, GRPOConfig, + _get_effort_config, ) from nemo_rl.algorithms.grpo import ( MasterConfig as GRPOMasterConfig, @@ -1290,6 +1291,11 @@ async def _iter_rollout_groups( ), max_rollout_turns=None, greedy=False, + effort_config=( + _get_effort_config(self.master_config) + if isinstance(self.master_config, GRPOMasterConfig) + else None + ), reward_penalty_config=self.master_config.reward_penalties, length_penalty_config=self.master_config.grpo.model_dump(), thinking_tags=get_nemo_gym_thinking_tags(self.master_config.env), diff --git a/tests/unit/algorithms/test_async_utils.py b/tests/unit/algorithms/test_async_utils.py index a26387638ae..57e6f59f419 100644 --- a/tests/unit/algorithms/test_async_utils.py +++ b/tests/unit/algorithms/test_async_utils.py @@ -62,6 +62,7 @@ PENDING_PROMPTS_KEY, RETAINED_TASK_INDICES_KEY, ) +from nemo_rl.experience.rollouts import EffortLevelsConfig @ray.remote(num_cpus=0) @@ -2898,8 +2899,10 @@ async def fake_rollouts(**kwargs): assert exc.value.__cause__ is not None assert "unexpected add status" in str(exc.value.__cause__) - def test_nemo_gym_batch_retry_does_not_duplicate_buffered_groups(self, monkeypatch): - """A partial stream retry only enqueues prompt groups not already buffered.""" + def test_nemo_gym_batch_retry_forwards_effort_config_without_duplicates( + self, monkeypatch + ): + """Retries preserve effort shaping and do not re-enqueue buffered groups.""" class _ReadyResult: def __init__(self, value): @@ -2930,6 +2933,14 @@ def __init__(self): "stop_token_ids": [1], "stop_strings": ["stop"], } + collector.master_config.env["nemo_gym"] = { + "effort_levels": { + "low_weight": 0.1, + "low_penalty": 1.0, + "low_ub": 15_000, + "low_string": "{reasoning effort: efficient}", + } + } target_weight = 15 collector._generating_targets.add(target_weight) repeated_batch = BatchedDataDict( @@ -2957,6 +2968,12 @@ async def fake_rollouts(**kwargs): assert kwargs["generation_config"]["stop_token_ids"] is None assert kwargs["generation_config"]["stop_strings"] is None assert kwargs["log_full_result_tables"] is False + assert kwargs["effort_config"] == EffortLevelsConfig( + low_weight=0.1, + low_penalty=1.0, + low_ub=15_000, + low_string="{reasoning effort: efficient}", + ) rollout_calls += 1 yield _rollout_result(7) if rollout_calls == 1: