Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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://", "")
Expand DownExpand Up@@ -1279,6 +1282,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 "
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1354,15 +1354,21 @@ 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
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.
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.
Expand All@@ -1388,7 +1394,14 @@ 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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This code is RLVR speicfic but will run for all other training techniques too

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we check training type (rlvr) before doing this check

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also this is very nova specific, might want to use is_nova flag too

@ehsu3ehsu3Jul 21, 2026

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yep, I can add that check!

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


def _render_recipe_placeholders(recipe_content: str, override_spec: dict) -> str:
Expand DownExpand Up@@ -1499,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:
Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -1126,6 +1126,69 @@ 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, customization_technique="RLVR")

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

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)
Expand Down
67 changes: 67 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_compute.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -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")
Expand Down