From 22056fe123ebeb8c21ba0de809bed31c306734bf Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Thu, 13 Aug 2026 15:50:21 -0700 Subject: [PATCH] fix: rever preset reward function deletion from hyperparams dict --- .../src/sagemaker/train/rlvr_trainer.py | 18 +++++++++++++++--- .../train/test_rlvr_trainer_integration.py | 2 ++ 2 files changed, 17 insertions(+), 3 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/rlvr_trainer.py b/sagemaker-train/src/sagemaker/train/rlvr_trainer.py index 7435f561bb..5f03cb5b8c 100644 --- a/sagemaker-train/src/sagemaker/train/rlvr_trainer.py +++ b/sagemaker-train/src/sagemaker/train/rlvr_trainer.py @@ -240,9 +240,6 @@ def _process_hyperparameters(self): if hasattr(self.hyperparameters, 'reward_lambda_arn'): delattr(self.hyperparameters, 'reward_lambda_arn') self.hyperparameters._specs.pop('reward_lambda_arn', None) - if hasattr(self.hyperparameters, 'preset_reward_function'): - delattr(self.hyperparameters, 'preset_reward_function') - self.hyperparameters._specs.pop('preset_reward_function', None) if hasattr(self.hyperparameters, 'data_path'): delattr(self.hyperparameters, 'data_path') self.hyperparameters._specs.pop('data_path', None) @@ -410,7 +407,22 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, Returns: TrainingJob: The SageMaker training job object, or None if dry_run=True. + + Raises: + ValueError: If neither a custom reward function nor a preset reward + function hyperparameter is configured. """ + # A reward signal is required: either a custom reward function (Lambda ARN, + # evaluator ARN, or Evaluator object) or the preset_reward_function hyperparameter. + preset_reward_function = getattr(self.hyperparameters, "preset_reward_function", None) + if not self.custom_reward_function and not preset_reward_function: + raise ValueError( + "RLVR training requires a reward signal. Provide either " + "'custom_reward_function' (a Lambda ARN, evaluator ARN, or Evaluator object) " + "when initializing RLVRTrainer, or set the 'preset_reward_function' " + "hyperparameter (e.g. trainer.hyperparameters.preset_reward_function = 'prime_code')." + ) + # Dispatch based on compute type if isinstance(self.compute, HyperPodCompute): return self._train_hyperpod( diff --git a/sagemaker-train/tests/integ/train/test_rlvr_trainer_integration.py b/sagemaker-train/tests/integ/train/test_rlvr_trainer_integration.py index a022b6846f..8676259307 100644 --- a/sagemaker-train/tests/integ/train/test_rlvr_trainer_integration.py +++ b/sagemaker-train/tests/integ/train/test_rlvr_trainer_integration.py @@ -92,6 +92,8 @@ def test_rlvr_trainer_lora_complete_workflow(sagemaker_session): accept_eula=True, base_job_name=f"rlvr-lora-integ-{unique_id}", ) + + rlvr_trainer.hyperparameters.preset_reward_function = "prime_code" # Create training job training_job = rlvr_trainer.train(wait=False)