From 840e5bfd3f8766e17b00d78d59fb20593e91f6ff Mon Sep 17 00:00:00 2001 From: Ealynn Hsu Date: Tue, 21 Jul 2026 20:14:25 +0000 Subject: [PATCH 1/2] [Fix] Remove task-type from RLVR recipe, update RLVR image selection logic --- .../src/sagemaker/train/base_trainer.py | 4 ++ .../train/common_utils/finetune_utils.py | 11 ++- .../train/common_utils/test_finetune_utils.py | 44 ++++++++++++ .../unit/train/test_base_trainer_compute.py | 67 +++++++++++++++++++ 4 files changed, 125 insertions(+), 1 deletion(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 9fb911709a..fda74e8d48 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -1279,6 +1279,10 @@ def _train_hyperpod(self, training_dataset=None, validation_dataset=None, sagemaker_session=sagemaker_session, ) + # RFT/RLVR on HyperPod requires the TRAIN-specific image tag. + if training_image and "SM-HP-RFT-" in training_image and "TRAIN" not in training_image: + training_image = training_image.replace("SM-HP-RFT-", "SM-HP-RFT-TRAIN-") + if not training_image: raise ValueError( "training_image is required for HyperPod compute but could not be resolved " diff --git a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py index 9048204cc6..d8f4f2e58d 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1361,6 +1361,9 @@ def _extract_recipe_from_helm_template(template_content: str) -> str: with ``---`` separators). The HyperPod CLI expects a single-document recipe YAML. This function extracts just the ``config.yaml`` content section. + For RFT/RLVR recipes, also strips the ``task_type: storm_rbs`` field from the + Hub template. + Args: template_content: Raw Helm chart template string from S3. @@ -1388,7 +1391,13 @@ def _extract_recipe_from_helm_template(template_content: str) -> str: "The template format may have changed." ) - return textwrap.dedent(recipe_match.group(1)).strip() + result = textwrap.dedent(recipe_match.group(1)).strip() + + # Strip task_type: storm_rbs from RFT/RLVR recipes - including it causes service validation failures. + result = re.sub(r"^\s*task_type:\s*storm_rbs\s*$", "", result, flags=re.MULTILINE) + result = textwrap.dedent(result) + + return result def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str: diff --git a/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py b/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py index ead56fa65d..e137f1513b 100644 --- a/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py +++ b/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py @@ -1126,6 +1126,50 @@ def test_unparseable_template_raises(self): with pytest.raises(ValueError, match="template format may have changed"): fu._extract_recipe_from_helm_template(template) + def test_strips_task_type_storm_rbs(self): + """task_type: storm_rbs is internal RFT metadata and should be stripped.""" + template = ( + "---\n" + "# Source: grpo/templates/training-config.yaml\n" + "apiVersion: v1\n" + "data:\n" + " config.yaml: |-\n" + " run:\n" + " name: test\n" + " peft:\n" + " peft_scheme: lora\n" + " lora_tuning:\n" + " alpha: 32\n" + " task_type: storm_rbs\n" + "---\n" + ) + + extracted = fu._extract_recipe_from_helm_template(template) + + assert "task_type" not in extracted + assert "storm_rbs" not in extracted + assert "peft_scheme: lora" in extracted + + def test_preserves_non_storm_rbs_task_type(self): + """task_type: other task types (used by OSS) should NOT be stripped.""" + template = ( + "---\n" + "# Source: grpo/templates/training-config.yaml\n" + "apiVersion: v1\n" + "data:\n" + " config.yaml: |-\n" + " run:\n" + " name: test\n" + " peft:\n" + " lora_tuning:\n" + " task_type: OTHER_TASK\n" + "---\n" + ) + + extracted = fu._extract_recipe_from_helm_template(template) + + assert "task_type: OTHER_TASK" in extracted + class TestGetRecipeS3Uri: @patch(f"{_MOD}._normalize_model_name", side_effect=lambda m: m) diff --git a/sagemaker-train/tests/unit/train/test_base_trainer_compute.py b/sagemaker-train/tests/unit/train/test_base_trainer_compute.py index 4d2df0e28f..acee4e0697 100644 --- a/sagemaker-train/tests/unit/train/test_base_trainer_compute.py +++ b/sagemaker-train/tests/unit/train/test_base_trainer_compute.py @@ -190,6 +190,73 @@ def test_missing_cluster_name_raises( mock_verify.assert_not_called() + @patch("sagemaker.train.base_trainer.subprocess") + @patch("sagemaker.train.base_trainer.TrainDefaults.verify_hyperpod_caller_permissions") + @patch("sagemaker.train.base_trainer.validate_hyperpod_compute") + @patch("sagemaker.train.base_trainer.TrainDefaults.get_sagemaker_session") + @patch("sagemaker.train.base_trainer.get_hyperpod_recipe_path", return_value="recipes/test") + @patch("sagemaker.train.base_trainer.flatten_resolved_recipe", return_value={}) + def test_rft_image_tag_corrected_to_train( + self, mock_flatten, mock_get_recipe_path, mock_get_session, + mock_validate, mock_verify, mock_subprocess + ): + """SM-HP-RFT-V2-latest should be rewritten to SM-HP-RFT-TRAIN-V2-latest.""" + mock_get_session.return_value = MagicMock() + mock_subprocess.run.return_value = SimpleNamespace( + stdout="NAME: rft-job-123\n", stderr="" + ) + + trainer = _make_hyperpod_trainer(node_count=2) + # Simulate Hub resolving the wrong RFT image tag + trainer.training_image = ( + "012345678910.dkr.ecr.us-east-1.amazonaws.com/test-repo:SM-HP-RFT-TEST" + ) + + with patch( + "sagemaker.train.common_utils.finetune_utils.get_training_image", + return_value=None, + ), patch.object(trainer, "get_resolved_recipe", return_value={"training_config": {}}): + trainer._train_hyperpod(wait=False) + + start_cmd = mock_subprocess.run.call_args_list[-1].args[0] + overrides = json.loads(start_cmd[start_cmd.index("--override-parameters") + 1]) + assert overrides["container"] == ( + "012345678910.dkr.ecr.us-east-1.amazonaws.com/test-repo:SM-HP-RFT-TRAIN-TEST" + ) + + @patch("sagemaker.train.base_trainer.subprocess") + @patch("sagemaker.train.base_trainer.TrainDefaults.verify_hyperpod_caller_permissions") + @patch("sagemaker.train.base_trainer.validate_hyperpod_compute") + @patch("sagemaker.train.base_trainer.TrainDefaults.get_sagemaker_session") + @patch("sagemaker.train.base_trainer.get_hyperpod_recipe_path", return_value="recipes/test") + @patch("sagemaker.train.base_trainer.flatten_resolved_recipe", return_value={}) + def test_rft_train_image_not_double_replaced( + self, mock_flatten, mock_get_recipe_path, mock_get_session, + mock_validate, mock_verify, mock_subprocess + ): + """SM-HP-RFT-TRAIN-V2-latest should NOT be modified (already correct).""" + mock_get_session.return_value = MagicMock() + mock_subprocess.run.return_value = SimpleNamespace( + stdout="NAME: rft-job-123\n", stderr="" + ) + + trainer = _make_hyperpod_trainer(node_count=2) + trainer.training_image = ( + "012345678910.dkr.ecr.us-east-1.amazonaws.com/test-repo:SM-HP-RFT-TRAIN" + ) + + with patch( + "sagemaker.train.common_utils.finetune_utils.get_training_image", + return_value=None, + ), patch.object(trainer, "get_resolved_recipe", return_value={"training_config": {}}): + trainer._train_hyperpod(wait=False) + + start_cmd = mock_subprocess.run.call_args_list[-1].args[0] + overrides = json.loads(start_cmd[start_cmd.index("--override-parameters") + 1]) + assert overrides["container"] == ( + "012345678910.dkr.ecr.us-east-1.amazonaws.com/test-repo:SM-HP-RFT-TRAIN" + ) + @patch("sagemaker.train.base_trainer.subprocess") @patch("sagemaker.train.base_trainer.TrainDefaults.verify_hyperpod_caller_permissions") @patch("sagemaker.train.base_trainer.validate_hyperpod_compute") From d4fdf4235f0960922baaad1de85a73a389d6d197 Mon Sep 17 00:00:00 2001 From: Ealynn Hsu Date: Tue, 21 Jul 2026 21:57:03 +0000 Subject: [PATCH 2/2] Add check for RLVR/RFT before removing task_type --- .../src/sagemaker/train/base_trainer.py | 5 ++++- .../train/common_utils/finetune_utils.py | 14 +++++++++---- .../train/common_utils/test_finetune_utils.py | 21 ++++++++++++++++++- 3 files changed, 34 insertions(+), 6 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index fda74e8d48..94a3aa6e7d 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -211,7 +211,10 @@ def _fetch_full_recipe_template(self) -> Optional[Dict[str, Any]]: hp_uri = recipe_entry["HpEksPayloadTemplateS3Uri"] bucket, key = hp_uri.replace("s3://", "").split("/", 1) raw = s3_client.get_object(Bucket=bucket, Key=key)["Body"].read().decode("utf-8") - return yaml.safe_load(_extract_recipe_from_helm_template(raw)) + return yaml.safe_load(_extract_recipe_from_helm_template( + raw, + customization_technique=self._customization_technique if _is_nova_model(self._model_name) else None, + )) else: smtj_uri = resolve_s3_uri_placeholders(recipe_entry["SmtjRecipeTemplateS3Uri"], sagemaker_session) uri_path = smtj_uri.replace("s3://", "") diff --git a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py index d8f4f2e58d..3e8fe67029 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1354,7 +1354,7 @@ def _get_smhp_replicas_enum(model_name: str, customization_technique: str, train return None -def _extract_recipe_from_helm_template(template_content: str) -> str: +def _extract_recipe_from_helm_template(template_content: str, customization_technique: str = None) -> str: """Extract the training config YAML from a HyperPod Helm chart template. The HpEksPayloadTemplateS3Uri contains a full Helm chart (multi-document YAML @@ -1366,6 +1366,9 @@ def _extract_recipe_from_helm_template(template_content: str) -> str: Args: template_content: Raw Helm chart template string from S3. + customization_technique: The training technique (e.g. "RLVR", "RFT", "SFT"). + When set to "RLVR" or "RFT", strips ``task_type: storm_rbs`` from the + extracted config. Returns: str: Single-document recipe YAML content. @@ -1394,8 +1397,9 @@ def _extract_recipe_from_helm_template(template_content: str) -> str: result = textwrap.dedent(recipe_match.group(1)).strip() # Strip task_type: storm_rbs from RFT/RLVR recipes - including it causes service validation failures. - result = re.sub(r"^\s*task_type:\s*storm_rbs\s*$", "", result, flags=re.MULTILINE) - result = textwrap.dedent(result) + if customization_technique and customization_technique.upper() in ("RLVR", "RFT"): + result = re.sub(r"^\s*task_type:\s*storm_rbs\s*$", "", result, flags=re.MULTILINE) + result = textwrap.dedent(result) return result @@ -1508,7 +1512,9 @@ def get_hyperpod_recipe_path(model_name: str, customization_technique: str, trai recipe_content = response["Body"].read().decode("utf-8") # Extract the training config from the Helm chart template - recipe_content = _extract_recipe_from_helm_template(recipe_content) + # Only pass customization_technique for Nova models (task_type stripping is Nova RLVR/RFT specific) + technique_for_extraction = customization_technique if _is_nova_model(model_name) else None + recipe_content = _extract_recipe_from_helm_template(recipe_content, customization_technique=technique_for_extraction) # Inject additional overrides into spec before rendering if additional_overrides: diff --git a/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py b/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py index e137f1513b..465933d3ed 100644 --- a/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py +++ b/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py @@ -1144,7 +1144,7 @@ def test_strips_task_type_storm_rbs(self): "---\n" ) - extracted = fu._extract_recipe_from_helm_template(template) + extracted = fu._extract_recipe_from_helm_template(template, customization_technique="RLVR") assert "task_type" not in extracted assert "storm_rbs" not in extracted @@ -1170,6 +1170,25 @@ def test_preserves_non_storm_rbs_task_type(self): assert "task_type: OTHER_TASK" in extracted + def test_does_not_strip_task_type_for_non_rlvr(self): + """task_type: storm_rbs is preserved when technique is not RLVR/RFT.""" + template = ( + "---\n" + "# Source: grpo/templates/training-config.yaml\n" + "apiVersion: v1\n" + "data:\n" + " config.yaml: |-\n" + " run:\n" + " name: test\n" + " peft:\n" + " task_type: storm_rbs\n" + "---\n" + ) + + extracted = fu._extract_recipe_from_helm_template(template, customization_technique="SFT") + + assert "task_type: storm_rbs" in extracted + class TestGetRecipeS3Uri: @patch(f"{_MOD}._normalize_model_name", side_effect=lambda m: m)