From 68b799472e76fd75bbc787eb083e51e8ed81d3b5 Mon Sep 17 00:00:00 2001 From: Zhiyu Li Date: Thu, 28 May 2026 10:23:45 -0700 Subject: [PATCH] fix(grpo): update grpo_sync seq_logprob_error_masking call site PR #2559 changed compute_and_apply_seq_logprob_error_masking to return a dict of 8 keys (previously a 3-tuple of max_seq_mult_prob_error, num_masked_seqs, masked_correct_pct). The two call sites in grpo.py were updated, but the call site in grpo_sync.py was missed, causing L1_Functional_Tests_GPU to fail on main with: File ".../nemo_rl/algorithms/grpo_sync.py", line 772, in grpo_train_sync ValueError: too many values to unpack (expected 3) Aligns grpo_sync.py with the dict-based pattern already in use at grpo.py:1935-1963 and grpo.py:3064-3091, and propagates the richer mean/min and after-mask metrics to the sync recipe. Signed-off-by: Zhiyu Li --- nemo_rl/algorithms/grpo_sync.py | 34 +++++++++++++++++++++++++-------- 1 file changed, 26 insertions(+), 8 deletions(-) diff --git a/nemo_rl/algorithms/grpo_sync.py b/nemo_rl/algorithms/grpo_sync.py index c147da2cd45..efb1f71e996 100644 --- a/nemo_rl/algorithms/grpo_sync.py +++ b/nemo_rl/algorithms/grpo_sync.py @@ -769,17 +769,37 @@ def grpo_train_sync( } ) - ( - max_seq_mult_prob_error, - num_masked_seqs, - masked_correct_pct, - ) = compute_and_apply_seq_logprob_error_masking( + seq_error_result = compute_and_apply_seq_logprob_error_masking( train_data=masking_data, rewards=rewards, seq_logprob_error_threshold=master_config.grpo[ "seq_logprob_error_threshold" ], ) + seq_logprob_error_metrics = { + "max_seq_mult_prob_error": seq_error_result[ + "max_seq_mult_prob_error" + ], + "mean_seq_mult_prob_error": seq_error_result[ + "mean_seq_mult_prob_error" + ], + "min_seq_mult_prob_error": seq_error_result[ + "min_seq_mult_prob_error" + ], + "max_seq_mult_prob_error_after_mask": seq_error_result[ + "max_seq_mult_prob_error_after_mask" + ], + "mean_seq_mult_prob_error_after_mask": seq_error_result[ + "mean_seq_mult_prob_error_after_mask" + ], + "min_seq_mult_prob_error_after_mask": seq_error_result[ + "min_seq_mult_prob_error_after_mask" + ], + "num_masked_seqs_by_logprob_error": seq_error_result[ + "num_masked_seqs" + ], + "masked_correct_pct": seq_error_result["masked_correct_pct"], + } # masking may have mutated sample_mask in place — # capture the post-masking value for delta-write. sample_mask = masking_data["sample_mask"] @@ -1011,9 +1031,7 @@ def grpo_train_sync( metrics["generation_logger_metrics"] = generation_logger_metrics total_valid_tokens += metrics["global_valid_toks"] - metrics["max_seq_mult_prob_error"] = max_seq_mult_prob_error - metrics["num_masked_seqs_by_logprob_error"] = num_masked_seqs - metrics["masked_correct_pct"] = masked_correct_pct + metrics.update(seq_logprob_error_metrics) consumed_samples += master_config.grpo["num_prompts_per_step"] timeout.mark_iteration()