Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions nemo_rl/algorithms/async_utils/trajectory_collector.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from nemo_rl.algorithms.grpo import (
AsyncGRPOConfig,
GRPOConfig,
_get_effort_config,
)
from nemo_rl.algorithms.grpo import (
MasterConfig as GRPOMasterConfig,
Expand Down Expand Up @@ -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),
Expand Down
21 changes: 19 additions & 2 deletions tests/unit/algorithms/test_async_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@
PENDING_PROMPTS_KEY,
RETAINED_TASK_INDICES_KEY,
)
from nemo_rl.experience.rollouts import EffortLevelsConfig


@ray.remote(num_cpus=0)
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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:
Expand Down
Loading