Skip to content
Open
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
33 changes: 30 additions & 3 deletions nemo_rl/algorithms/grpo_sync.py
Original file line number Diff line number Diff line change
Expand Up @@ -262,6 +262,8 @@ def validate_sync(
return {}, {}

timer = Timer()
# >= 1 is validated in setup().
val_num_generations_per_prompt = master_config.grpo.val_num_generations_per_prompt
total_rewards: list[float] = []
total_lengths: list[float] = []
all_message_logs: list[list[dict[str, str]]] = []
Expand All @@ -276,8 +278,10 @@ def validate_sync(
for batch_idx, val_batch in enumerate(val_dataloader):
if batch_idx >= max_batches:
break
n_prompts = int(val_batch.size)
policy.prepare_val_partition(n_prompts, partition_id=partition_id)
if val_num_generations_per_prompt > 1:
val_batch = val_batch.repeat_interleave(val_num_generations_per_prompt)
n_rollouts = int(val_batch.size)
policy.prepare_val_partition(n_rollouts, partition_id=partition_id)
meta, driver_carry, rollout_metrics, _ = ray.get(
rollout_actor.rollout_to_tq.remote(
val_batch,
Expand All @@ -294,7 +298,7 @@ def validate_sync(
total_lengths.append(rollout_metrics["mean_gen_tokens_per_sample"])
all_message_logs.extend(
[{"role": r, "content": c} for r, c in zip(roles[i], contents[i])]
for i in range(n_prompts)
for i in range(n_rollouts)
)
if capture_extras:
additional_metrics = rollout_metrics
Expand All @@ -305,12 +309,35 @@ def validate_sync(
if total_rewards
else 0.0
)
# Grouped validation (val_num_generations_per_prompt > 1) additionally
# reports pass@k over each prompt's k rollouts, mirroring
# nemo_rl.algorithms.grpo.validate.
pass_k = None
if total_rewards and val_num_generations_per_prompt > 1:
assert len(total_rewards) % val_num_generations_per_prompt == 0, (
"Validation rewards must be divisible by "
"grpo.val_num_generations_per_prompt"
)
pass_k = (
(
torch.tensor(total_rewards, dtype=torch.float32).view(
-1, val_num_generations_per_prompt
)
> 0
)
.any(dim=1)
.float()
.mean()
.item()
)
avg_length = sum(total_lengths) / len(total_lengths) if total_lengths else 0.0
val_metrics = {
"accuracy": accuracy,
"avg_length": avg_length,
**additional_metrics,
}
if pass_k is not None:
val_metrics["pass_k"] = pass_k
try:
print_message_log_samples(
all_message_logs,
Expand Down
69 changes: 68 additions & 1 deletion tests/unit/algorithms/test_grpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,11 @@
setup,
validate,
)
from nemo_rl.algorithms.grpo_sync import _train_fields_for_step, grpo_train_sync
from nemo_rl.algorithms.grpo_sync import (
_train_fields_for_step,
grpo_train_sync,
validate_sync,
)
from nemo_rl.algorithms.loss import ClippedPGLossConfig, ClippedPGLossFn
from nemo_rl.algorithms.reward_functions import (
RewardShapingConfig,
Expand Down Expand Up @@ -4037,6 +4041,69 @@ def run_rollout(_policy, repeated_batch, *_args, **_kwargs):
assert val_metrics["accuracy"] == pytest.approx(0.125)
assert val_metrics["pass_k"] == pytest.approx(0.5)

def test_sync_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 rollout_to_tq_remote(repeated_batch, **_kwargs):
# Each prompt is repeated k=4 times, contiguously.
assert repeated_batch["idx"].tolist() == [0, 0, 0, 0, 1, 1, 1, 1]
driver_carry = {
# Prompt 0 passes once out of 4; prompt 1 never passes.
"total_reward": torch.tensor(
[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]
),
"turn_roles": [["user"]] * 8,
"turn_contents": [["x"]] * 8,
}
return (MagicMock(), driver_carry, {"mean_gen_tokens_per_sample": 1.0}, {})

rollout_actor = MagicMock()
rollout_actor.rollout_to_tq.remote.side_effect = rollout_to_tq_remote
policy = MagicMock()

with (
patch("nemo_rl.algorithms.grpo_sync.ray.get", side_effect=lambda x: x),
patch(
"nemo_rl.algorithms.grpo_sync._should_use_nemo_gym",
return_value=False,
),
patch("nemo_rl.algorithms.grpo_sync.print_message_log_samples"),
):
val_metrics, _ = validate_sync(
rollout_actor=rollout_actor,
policy=policy,
val_dataloader=mock_dataloader,
val_task_to_env={"math": MagicMock(spec=EnvironmentInterface)},
step=0,
master_config=mock_config,
)

# The val partition is sized for the expanded batch (2 prompts x k=4).
policy.prepare_val_partition.assert_called_once_with(8, partition_id="val")
# 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
):
Expand Down
Loading