From 9c0c3add5b71f8e6f9b2c451c8d50c6ec64938e3 Mon Sep 17 00:00:00 2001 From: Nemo Assist Date: Tue, 10 Mar 2026 18:28:09 +0000 Subject: [PATCH] fix: skip_reference_policy_logprobs_calculation=true crashes training Fixes #1968: Setting skip_reference_policy_logprobs_calculation=true with reference_policy_kl_penalty=0 crashed training in three ways: Bug 1: use_reference_model() context manager crash when reference model was never initialized (AttributeError on reference_state_dict). Fix: Added early-return guard in use_reference_model() for all three worker types (megatron, dtensor v1, dtensor v2) - yields without swapping when reference model is None/missing. Bug 2: Async GRPO path unconditionally called get_reference_policy_logprobs() without checking the skip flag. Fix: Added the same skip guard as the sync path, setting zeros_like for reference_policy_logprobs when skipping. Bug 3: Missing reference_policy_logprobs key in train_data causing shape mismatches downstream in loss computation. Fix: Both sync and async paths now explicitly set train_data['reference_policy_logprobs'] = zeros_like(prev_logprobs) when skipping. Also added a _has_reference_model() helper and zeros fallback in base_policy_worker.get_reference_policy_logprobs() as defense-in-depth. --- nemo_rl/algorithms/grpo.py | 20 ++++++++++++++----- .../policy/workers/base_policy_worker.py | 18 +++++++++++++++++ .../policy/workers/dtensor_policy_worker.py | 5 +++++ .../workers/dtensor_policy_worker_v2.py | 5 +++++ .../policy/workers/megatron_policy_worker.py | 5 +++++ 5 files changed, 48 insertions(+), 5 deletions(-) diff --git a/nemo_rl/algorithms/grpo.py b/nemo_rl/algorithms/grpo.py index 996579eaad7..8d9ae9ccd48 100644 --- a/nemo_rl/algorithms/grpo.py +++ b/nemo_rl/algorithms/grpo.py @@ -1769,6 +1769,10 @@ 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 @@ -2799,12 +2803,18 @@ def async_grpo_train( train_data, timer=timer, )["logprobs"] - reference_logprobs = policy.get_reference_policy_logprobs( - train_data, - timer=timer, - )["reference_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) ( max_seq_mult_prob_error, diff --git a/nemo_rl/models/policy/workers/base_policy_worker.py b/nemo_rl/models/policy/workers/base_policy_worker.py index 34f772a175e..3c0d3a08ee2 100644 --- a/nemo_rl/models/policy/workers/base_policy_worker.py +++ b/nemo_rl/models/policy/workers/base_policy_worker.py @@ -139,6 +139,15 @@ def get_reference_policy_logprobs( We use the convention that the logprob of the first token is 0 so that the sequence length is maintained. The logprob of input token i is specified at position i in the output logprobs tensor. """ + # When reference model was never initialized (e.g., + # skip_reference_policy_logprobs_calculation=true), return zeros + # matching the expected shape to avoid crashes downstream. + if not self._has_reference_model(): + logprobs = self.get_logprobs(data=data, micro_batch_size=micro_batch_size) + return_data = BatchedDataDict[ReferenceLogprobOutputSpec]() + return_data["reference_logprobs"] = torch.zeros_like(logprobs["logprobs"]).cpu() + return return_data + with self.use_reference_model(): reference_logprobs = self.get_logprobs( data=data, micro_batch_size=micro_batch_size @@ -148,6 +157,15 @@ def get_reference_policy_logprobs( return_data["reference_logprobs"] = reference_logprobs["logprobs"].cpu() return return_data + def _has_reference_model(self) -> bool: + """Check if a reference model has been initialized.""" + # DTensor v2 uses reference_model_state_dict, others use reference_state_dict + for attr in ("reference_model_state_dict", "reference_state_dict"): + val = getattr(self, attr, None) + if val is not None: + return True + return False + def finish_training(self, *args: Any, **kwargs: Any) -> None: # Placeholder implementation pass diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker.py b/nemo_rl/models/policy/workers/dtensor_policy_worker.py index 5dd9f3a1bfa..dcad0ff4fbf 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker.py @@ -1667,6 +1667,11 @@ def use_reference_model(self) -> Generator[None, None, None]: is different from the current policy, making filtered logprobs incompatible. On exit: Restores original references and re-flips cuda/cpu, restores sampling_params. """ + # If reference model was never initialized, yield without swapping + if not hasattr(self, "reference_model_state_dict") or self.reference_model_state_dict is None: + yield + return + with torch.no_grad(): # Save train model state_dict curr_state_dict = get_cpu_state_dict( diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py index 71d5552c66d..0a248d2b7fe 100644 --- a/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py +++ b/nemo_rl/models/policy/workers/dtensor_policy_worker_v2.py @@ -785,6 +785,11 @@ def use_reference_model(self) -> Generator[None, None, None]: is different from the current policy, making filtered logprobs incompatible. On exit: Restores original references and re-flips cuda/cpu, restores sampling_params. """ + # If reference model was never initialized, yield without swapping + if not hasattr(self, "reference_model_state_dict") or self.reference_model_state_dict is None: + yield + return + with torch.no_grad(): # Save train model state_dict curr_state_dict = get_cpu_state_dict( diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py index 3a77d406e49..97568a39a4d 100644 --- a/nemo_rl/models/policy/workers/megatron_policy_worker.py +++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py @@ -534,6 +534,11 @@ def use_reference_model(self): is different from the current policy, making filtered logprobs incompatible. On exit: Restores original references and re-flips cuda/cpu, restores sampling_params. """ + # If reference model was never initialized, yield without swapping + if not hasattr(self, "reference_state_dict") or self.reference_state_dict is None: + yield + return + ## disable overlap param gather when swapping weights if self.should_disable_forward_pre_hook: self.disable_forward_pre_hook()