diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index da955ac32a3..f995e326557 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -389,6 +389,16 @@ def init_train_dataloader(dataset, suffix: str = ""): os.environ["NRL_IGNORE_TP_ACCURACY_CHECK"] = "1" print(" ✓ force_on_policy_ratio enabled") + # Validate skip_reference_policy_logprobs_calculation + if grpo_config.get("skip_reference_policy_logprobs_calculation"): + assert loss_config["reference_policy_kl_penalty"] == 0, ( + "grpo.skip_reference_policy_logprobs_calculation=True requires " + "loss_fn.reference_policy_kl_penalty == 0" + ) + print( + "Reference policy logprob calculation will be skipped since `grpo.skip_reference_policy_logprobs_calculation` is set to True and `loss_fn.reference_policy_kl_penalty` is 0." + ) + # ========================== # Cluster # ========================== @@ -554,6 +564,19 @@ def init_train_dataloader(dataset, suffix: str = ""): policy_config["megatron_cfg"]["train_iters"] = total_train_iters # Define initialization functions that will be used in all paths + init_reference_model = master_config["loss_fn"]["reference_policy_kl_penalty"] > 0 + + # Auto-enable skip_reference_policy_logprobs_calculation when the reference model is not loaded. + if not init_reference_model and not master_config["grpo"].get( + "skip_reference_policy_logprobs_calculation" + ): + master_config["grpo"]["skip_reference_policy_logprobs_calculation"] = True + print( + "Auto-enabling `grpo.skip_reference_policy_logprobs_calculation=True` " + "because `loss_fn.reference_policy_kl_penalty == 0` " + "(reference model is not loaded)." + ) + def init_policy(): """Initialize policy training workers.""" t0 = time.perf_counter() @@ -565,6 +588,7 @@ def init_policy(): weights_path=weights_path, optimizer_path=optimizer_path, init_optimizer=True, + init_reference_model=init_reference_model, ) return p, time.perf_counter() - t0 @@ -1336,12 +1360,6 @@ def grpo_train( POLICY_GENERATION_STALE = True # tracks if generation needs a refit before running assert policy_generation is not None # for mypy type check - if master_config["grpo"].get("skip_reference_policy_logprobs_calculation"): - assert master_config["loss_fn"]["reference_policy_kl_penalty"] == 0 - print( - "Reference policy logprob calculation will be skipped since `grpo.skip_reference_policy_logprobs_calculation` is set to True and `loss_fn.reference_policy_kl_penalty` is 0." - ) - # Check if we need to sync KV cache scales # When fallback to policy as the policy_generation, we use getattr to check. sync_kv_scales = getattr(policy_generation, "requires_kv_scale_sync", False) @@ -1728,9 +1746,19 @@ def grpo_train( metrics_logging_data["content"] = flat_messages["content"] memory_tracker.snapshot_start_of_stage("Computing logprobs", dir()) - print("▶ Preparing for logprob inference...", flush=True) - with timer.time("logprob_inference_prep"): - policy.prepare_for_lp_inference() + # Skip prev_logprobs computation when force_on_policy_ratio=True + skip_prev_logprobs = master_config["loss_fn"].get( + "force_on_policy_ratio", False + ) + if skip_prev_logprobs: + print( + "▶ Skipping prev_logprobs (force_on_policy_ratio=True)...", + flush=True, + ) + else: + print("▶ Preparing for logprob inference...", flush=True) + with timer.time("logprob_inference_prep"): + policy.prepare_for_lp_inference() print("▶ Computing logprobs...", flush=True) with timer.time("policy_and_reference_logprobs"): @@ -1744,9 +1772,14 @@ def grpo_train( **extra_multimodal_data, } ) - train_data["prev_logprobs"] = policy.get_logprobs( - logprob_data, timer=timer - )["logprobs"] + if not skip_prev_logprobs: + train_data["prev_logprobs"] = policy.get_logprobs( + logprob_data, timer=timer + )["logprobs"] + else: + train_data["prev_logprobs"] = torch.zeros_like( + train_data["generation_logprobs"] + ) if not master_config["grpo"].get( "skip_reference_policy_logprobs_calculation" @@ -1757,10 +1790,28 @@ def grpo_train( timer=timer, )["reference_logprobs"] ) + else: + train_data["reference_policy_logprobs"] = torch.zeros_like( + train_data["prev_logprobs"] + ) del logprob_data del extra_multimodal_data + # Seq-level logprob error metrics/masking require real prev_logprobs + seq_logprob_error_threshold = master_config["grpo"].get( + "seq_logprob_error_threshold", None + ) + if skip_prev_logprobs: + # Cannot compute seq-level metrics with placeholder prev_logprobs + assert seq_logprob_error_threshold is None, ( + "seq_logprob_error_threshold requires prev_logprobs computation; " + "cannot use with force_on_policy_ratio=True" + ) + max_seq_mult_prob_error = 0.0 + num_masked_seqs = 0 + masked_correct_pct = 0.0 + else: ( max_seq_mult_prob_error, num_masked_seqs, @@ -1768,9 +1819,7 @@ def grpo_train( ) = compute_and_apply_seq_logprob_error_masking( train_data=train_data, rewards=rewards, - seq_logprob_error_threshold=master_config["grpo"][ - "seq_logprob_error_threshold" - ], + seq_logprob_error_threshold=seq_logprob_error_threshold, ) # Compute advantages with adv_estimator using correct mask and logprobs @@ -2789,23 +2838,56 @@ def async_grpo_train( train_data.to("cpu") # Training phase (same as sync version) - print("▶ Preparing for logprob inference...") - with timer.time("logprob_inference_prep"): - policy.prepare_for_lp_inference() + # Skip prev_logprobs computation when force_on_policy_ratio=True + skip_prev_logprobs = master_config["loss_fn"].get( + "force_on_policy_ratio", False + ) + if skip_prev_logprobs: + print( + "▶ Skipping prev_logprobs (force_on_policy_ratio=True)...", + flush=True, + ) + fprop_logprobs = torch.zeros_like(train_data["generation_logprobs"]) + else: + print("▶ Preparing for logprob inference...") + with timer.time("logprob_inference_prep"): + policy.prepare_for_lp_inference() - print("▶ Computing logprobs...") + print("▶ Computing logprobs...", flush=True) with timer.time("policy_and_reference_logprobs"): - fprop_logprobs = policy.get_logprobs( - train_data, - timer=timer, - )["logprobs"] - reference_logprobs = policy.get_reference_policy_logprobs( - train_data, - timer=timer, - )["reference_logprobs"] + if not skip_prev_logprobs: + fprop_logprobs = policy.get_logprobs( + train_data, + timer=timer, + )["logprobs"] train_data["prev_logprobs"] = fprop_logprobs - train_data["reference_policy_logprobs"] = reference_logprobs + if not master_config["grpo"].get( + "skip_reference_policy_logprobs_calculation" + ): + reference_logprobs = policy.get_reference_policy_logprobs( + train_data, + timer=timer, + )["reference_logprobs"] + train_data["reference_policy_logprobs"] = reference_logprobs + else: + train_data["reference_policy_logprobs"] = torch.zeros_like( + fprop_logprobs + ) + + # Seq-level logprob error metrics/masking require real prev_logprobs + seq_logprob_error_threshold = master_config["grpo"].get( + "seq_logprob_error_threshold", None + ) + if skip_prev_logprobs: + assert seq_logprob_error_threshold is None, ( + "seq_logprob_error_threshold requires prev_logprobs computation; " + "cannot use with force_on_policy_ratio=True" + ) + max_seq_mult_prob_error = 0.0 + num_masked_seqs = 0 + masked_correct_pct = 0.0 + else: ( max_seq_mult_prob_error, num_masked_seqs, @@ -2813,9 +2895,7 @@ def async_grpo_train( ) = compute_and_apply_seq_logprob_error_masking( train_data=train_data, rewards=rewards, - seq_logprob_error_threshold=master_config["grpo"][ - "seq_logprob_error_threshold" - ], + seq_logprob_error_threshold=seq_logprob_error_threshold, ) # Compute advantages with adv_estimator using correct mask and logprobs diff --git a/nemo_rl/algorithms/loss/loss_functions.py b/nemo_rl/algorithms/loss/loss_functions.py index df6ff6bc54b..9442bce999e 100755 --- a/nemo_rl/algorithms/loss/loss_functions.py +++ b/nemo_rl/algorithms/loss/loss_functions.py @@ -256,7 +256,10 @@ def __call__( token_mask = data["token_mask"][:, 1:] sample_mask = data["sample_mask"] advantages = data["advantages"][:, 1:] - prev_logprobs = data["prev_logprobs"][:, 1:] + # Skip loading prev_logprobs when force_on_policy_ratio=True (will use curr_logprobs instead) + prev_logprobs = ( + None if self.force_on_policy_ratio else data["prev_logprobs"][:, 1:] + ) generation_logprobs = data["generation_logprobs"][:, 1:] if self.reference_policy_kl_penalty != 0: reference_policy_logprobs = data["reference_policy_logprobs"][:, 1:] @@ -266,6 +269,11 @@ def __call__( mask = token_mask * sample_mask.unsqueeze(-1) + # For truly on-policy training, use curr_logprobs as prev_logprobs + # This avoids computing prev_logprobs upstream + if self.force_on_policy_ratio: + prev_logprobs = curr_logprobs.detach() + # token_mult_prob_error # See more details and other metrics in docs/guides/grpo.md#metrics lp_error = torch.abs(generation_logprobs - prev_logprobs) # noqa: F841 (precommit ignore for now) diff --git a/tests/unit/algorithms/test_grpo.py b/tests/unit/algorithms/test_grpo.py index 2ddbf001c98..363c8162867 100644 --- a/tests/unit/algorithms/test_grpo.py +++ b/tests/unit/algorithms/test_grpo.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +from contextlib import contextmanager from unittest.mock import MagicMock, patch import pytest @@ -258,6 +259,45 @@ def mock_ray_get(ref): return stack +@contextmanager +def _patched_logprob_phase(policy): + """Provide real tensors for the logprob phase of ``grpo_train``. + + Both PRs #2174 / #2178 (``skip_reference_policy_logprobs_calculation=True``) + and PR #2177 (``force_on_policy_ratio=True``) use ``torch.zeros_like(...)`` + placeholders inside ``grpo_train``. Those calls fail with + ``TypeError: zeros_like(): argument 'input' must be Tensor, not MagicMock`` + when the surrounding inputs come straight from the bare ``mock_grpo_components`` + fixture. This helper swaps in real tensors for the duration of the test and + restores the original mock return values afterwards. + """ + fake_flat = BatchedDataDict( + { + "token_ids": torch.tensor([[1, 2]]), + "advantages": torch.tensor([[0.5, 0.5]]), + "generation_logprobs": torch.tensor([[0.0, 0.0]]), + "token_loss_mask": torch.tensor([[1, 1]]), + "content": ["ok"], + } + ) + fake_lengths = torch.tensor([2]) + saved_lp = policy.get_logprobs.return_value + saved_rlp = policy.get_reference_policy_logprobs.return_value + policy.get_logprobs.return_value = {"logprobs": torch.zeros(1, 2)} + policy.get_reference_policy_logprobs.return_value = { + "reference_logprobs": torch.zeros(1, 2) + } + with patch( + "nemo_rl.algorithms.grpo.batched_message_log_to_flat_message", + return_value=(fake_flat, fake_lengths), + ): + try: + yield + finally: + policy.get_logprobs.return_value = saved_lp + policy.get_reference_policy_logprobs.return_value = saved_rlp + + @ray.remote(num_cpus=0) class MockEnvironment(EnvironmentInterface): def __init__(self, rewards: list[float]): @@ -995,6 +1035,7 @@ def init_collective(self, *_args, **_kwargs): "loss_fn": { "force_on_policy_ratio": False, "use_importance_sampling_correction": False, + "reference_policy_kl_penalty": 0.0, }, "env": {}, "grpo": { @@ -1206,6 +1247,184 @@ def fake_batched_message_log_to_flat_message(*_args, **_kwargs): ) +@pytest.mark.parametrize("train_func", [grpo_train, async_grpo_train]) +def test_grpo_train_skips_reference_policy_logprobs_when_configured( + mock_grpo_components, train_func +): + """Regression test for issue #1968 (Bug 1) and PRs #2174 / #2178. + + When ``grpo.skip_reference_policy_logprobs_calculation=True`` and + ``loss_fn.reference_policy_kl_penalty=0``, both ``grpo_train`` and + ``async_grpo_train`` MUST NOT call ``policy.get_reference_policy_logprobs``. + Without the skip guards in ``grpo.py``, training would crash inside + ``use_reference_model()`` because the reference model state was never loaded. + """ + master_config = mock_grpo_components["master_config"] + master_config["grpo"]["skip_reference_policy_logprobs_calculation"] = True + master_config["loss_fn"]["reference_policy_kl_penalty"] = 0 + master_config["grpo"]["max_num_steps"] = 1 + master_config["grpo"]["max_num_epochs"] = 1 + master_config["grpo"]["val_period"] = 0 + master_config["grpo"]["val_at_start"] = False + master_config["grpo"]["use_dynamic_sampling"] = False + + if train_func == async_grpo_train: + master_config["policy"]["generation"]["colocated"]["enabled"] = False + + grpo_save_state = _default_grpo_save_state() + mock_rollout_metrics = { + "mean_gen_tokens_per_sample": 10.0, + "max_gen_tokens": 20, + "min_gen_tokens": 5, + } + mock_batch = next(iter(mock_grpo_components["train_dataloader"])) + policy = mock_grpo_components["policy"] + + if train_func == async_grpo_train: + with ( + mock_async_grpo_infrastructure(mock_batch, mock_rollout_metrics), + _patched_logprob_phase(policy), + ): + train_func( + policy, + None, + mock_grpo_components["train_dataloader"], + mock_grpo_components["val_dataloader"], + mock_grpo_components["tokenizer"], + mock_grpo_components["loss_fn"], + mock_grpo_components["task_to_env"], + mock_grpo_components["val_task_to_env"], + mock_grpo_components["logger"], + mock_grpo_components["checkpointer"], + grpo_save_state, + master_config, + ) + else: + with ( + _patched_logprob_phase(policy), + patch( + "nemo_rl.algorithms.grpo.run_multi_turn_rollout", + return_value=(mock_batch, mock_rollout_metrics), + ), + patch( + "nemo_rl.algorithms.grpo.run_async_multi_turn_rollout", + return_value=(mock_batch, mock_rollout_metrics), + ), + patch( + "nemo_rl.algorithms.grpo.compute_and_apply_seq_logprob_error_masking", + return_value=(0.0, 0, 0.0), + ), + ): + train_func( + policy, + None, + mock_grpo_components["train_dataloader"], + mock_grpo_components["val_dataloader"], + mock_grpo_components["tokenizer"], + mock_grpo_components["loss_fn"], + mock_grpo_components["task_to_env"], + mock_grpo_components["val_task_to_env"], + mock_grpo_components["logger"], + mock_grpo_components["checkpointer"], + grpo_save_state, + master_config, + ) + + assert not policy.get_reference_policy_logprobs.called, ( + "policy.get_reference_policy_logprobs was called even though " + "skip_reference_policy_logprobs_calculation=True. " + "This indicates a regression of issue #1968 / PRs #2174, #2178." + ) + + +@pytest.mark.parametrize("train_func", [grpo_train, async_grpo_train]) +def test_grpo_train_skips_prev_logprobs_when_force_on_policy_ratio( + mock_grpo_components, train_func +): + """Regression test for PR #2177. + + When ``loss_fn.force_on_policy_ratio=True``, both ``grpo_train`` and + ``async_grpo_train`` MUST NOT call ``policy.get_logprobs`` to compute + ``prev_logprobs`` -- the importance-sampling ratio is forced to 1.0 so + the prev-policy forward pass would be wasted compute. + """ + master_config = mock_grpo_components["master_config"] + master_config["loss_fn"]["force_on_policy_ratio"] = True + master_config["grpo"]["seq_logprob_error_threshold"] = None + master_config["grpo"]["max_num_steps"] = 1 + master_config["grpo"]["max_num_epochs"] = 1 + master_config["grpo"]["val_period"] = 0 + master_config["grpo"]["val_at_start"] = False + master_config["grpo"]["use_dynamic_sampling"] = False + + if train_func == async_grpo_train: + master_config["policy"]["generation"]["colocated"]["enabled"] = False + + grpo_save_state = _default_grpo_save_state() + mock_rollout_metrics = { + "mean_gen_tokens_per_sample": 10.0, + "max_gen_tokens": 20, + "min_gen_tokens": 5, + } + mock_batch = next(iter(mock_grpo_components["train_dataloader"])) + policy = mock_grpo_components["policy"] + + if train_func == async_grpo_train: + with ( + mock_async_grpo_infrastructure(mock_batch, mock_rollout_metrics), + _patched_logprob_phase(policy), + ): + train_func( + policy, + None, + mock_grpo_components["train_dataloader"], + mock_grpo_components["val_dataloader"], + mock_grpo_components["tokenizer"], + mock_grpo_components["loss_fn"], + mock_grpo_components["task_to_env"], + mock_grpo_components["val_task_to_env"], + mock_grpo_components["logger"], + mock_grpo_components["checkpointer"], + grpo_save_state, + master_config, + ) + else: + with ( + _patched_logprob_phase(policy), + patch( + "nemo_rl.algorithms.grpo.run_multi_turn_rollout", + return_value=(mock_batch, mock_rollout_metrics), + ), + patch( + "nemo_rl.algorithms.grpo.run_async_multi_turn_rollout", + return_value=(mock_batch, mock_rollout_metrics), + ), + patch( + "nemo_rl.algorithms.grpo.compute_and_apply_seq_logprob_error_masking", + return_value=(0.0, 0, 0.0), + ), + ): + train_func( + policy, + None, + mock_grpo_components["train_dataloader"], + mock_grpo_components["val_dataloader"], + mock_grpo_components["tokenizer"], + mock_grpo_components["loss_fn"], + mock_grpo_components["task_to_env"], + mock_grpo_components["val_task_to_env"], + mock_grpo_components["logger"], + mock_grpo_components["checkpointer"], + grpo_save_state, + master_config, + ) + + assert not policy.get_logprobs.called, ( + "policy.get_logprobs was called even though force_on_policy_ratio=True. " + "This indicates a regression of PR #2177." + ) + + @pytest.fixture def mock_grpo_components(): # Create mock components diff --git a/tests/unit/algorithms/test_loss_functions.py b/tests/unit/algorithms/test_loss_functions.py index a49eff739d4..6d91f472b05 100644 --- a/tests/unit/algorithms/test_loss_functions.py +++ b/tests/unit/algorithms/test_loss_functions.py @@ -637,6 +637,60 @@ def test_clipped_pg_loss_force_on_policy_ratio(): assert metrics["probs_ratio_clamped_max"] == 1.0 +def test_clipped_pg_loss_force_on_policy_ratio_ignores_prev_logprobs(): + """Tests that force_on_policy_ratio ignores prev_logprobs from data and uses curr_logprobs instead. + + When force_on_policy_ratio=True, the loss function should use curr_logprobs.detach() + as prev_logprobs, so the actual prev_logprobs in data are irrelevant. This allows + skipping the expensive prev_logprobs computation upstream. + """ + if not torch.cuda.is_available(): + pytest.skip("No GPU available") + + device = "cuda" + data, batch_size, seq_len, vocab_size = _setup_clipped_pg_test_data(device=device) + + cfg = deepcopy(basic_pg_loss_test_config) + cfg["force_on_policy_ratio"] = True + loss_fn = ClippedPGLossFn(cfg) + + curr_lp = torch.tensor([[-1.0, -1.0, -1.0]], device=device) + input_ids = data["input_ids"] + dummy_logits = _create_exact_logits( + curr_lp, input_ids, batch_size, seq_len, vocab_size, device + ) + + # Run with correct prev_logprobs + data_1, _, _, _ = _setup_clipped_pg_test_data(device=device) + data_1["prev_logprobs"][0, 1:] = curr_lp + loss_input_1, data_1 = prepare_loss_input(dummy_logits.clone(), data_1, loss_fn) + loss_1, metrics_1 = loss_fn( + data=data_1, + global_valid_seqs=torch.sum(data_1["sample_mask"]), + global_valid_toks=torch.sum( + data_1["sample_mask"].unsqueeze(-1) * data_1["token_mask"] + ), + **loss_input_1, + ) + + # Run with wildly different prev_logprobs (should be ignored) + data_2, _, _, _ = _setup_clipped_pg_test_data(device=device) + data_2["prev_logprobs"][0, 1:] = torch.tensor([-10.0, -10.0, -10.0], device=device) + loss_input_2, data_2 = prepare_loss_input(dummy_logits.clone(), data_2, loss_fn) + loss_2, metrics_2 = loss_fn( + data=data_2, + global_valid_seqs=torch.sum(data_2["sample_mask"]), + global_valid_toks=torch.sum( + data_2["sample_mask"].unsqueeze(-1) * data_2["token_mask"] + ), + **loss_input_2, + ) + + # Both should produce identical loss and ratios since prev_logprobs is ignored + torch.testing.assert_close(loss_1, loss_2) + assert metrics_1["probs_ratio"] == metrics_2["probs_ratio"] == 1.0 + + @pytest.mark.parametrize("kl_type", ["k1", "k2", "k3"]) def test_calculate_kl(kl_type): """Tests KL calculations."""