Skip to content
Open
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
49 changes: 45 additions & 4 deletions sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
import logging
import json
from typing import Any, Dict, Optional, Union
import time
import boto3
from sagemaker.core.resources import ModelPackage, ModelPackageGroup
from sagemaker.core.helper.session_helper import Session
Expand Down Expand Up @@ -188,7 +187,12 @@ def _get_prod_sm_client(sagemaker_session) -> "boto3.client":
return boto3.client("sagemaker", region_name=region)


def _resolve_mlflow_resource_arn(sagemaker_session, mlflow_resource_arn: Optional[str] = None, min_mlflow_version: Optional[str] = None) -> Optional[str]:
def _resolve_mlflow_resource_arn(
sagemaker_session,
mlflow_resource_arn: Optional[str] = None,
min_mlflow_version: Optional[str] = None,
dry_run: bool = False,
) -> Optional[str]:
"""Resolve MLflow resource ARN using default experience logic.

All MLflow API calls use a raw boto3 client against prod (no custom endpoint),
Expand All @@ -199,6 +203,8 @@ def _resolve_mlflow_resource_arn(sagemaker_session, mlflow_resource_arn: Optiona
mlflow_resource_arn: Explicit ARN to use (returned as-is if provided).
min_mlflow_version: Minimum required MLflow version (e.g. "3.10").
If the resolved app's version is below this, a new app is created.
dry_run: If True, only performs read-only checks (list/describe) without
creating new apps or waiting for apps in Creating status.
"""
if mlflow_resource_arn:
return mlflow_resource_arn
Expand Down Expand Up @@ -240,10 +246,24 @@ def _resolve_mlflow_resource_arn(sagemaker_session, mlflow_resource_arn: Optiona
logger.warning("Resolved MLflow app %s is in failed state: %s. Skipping.",
resolved_arn, resolved_app.get("Status"))
resolved_app = None
elif dry_run and resolved_app.get("Status") in ["Creating", "Updating"]:
logger.warning(
"dry_run: MLflow app %s is in '%s' state. "
"Job submission would block until the app is ready.",
resolved_arn, resolved_app.get("Status"),
)
return resolved_arn

# Version check: if resolved app is below min version, create a new one as default
if resolved_app and min_mlflow_version and not _mlflow_version_meets_minimum_dict(resolved_app, min_mlflow_version):
resolved_arn = resolved_app["Arn"]
if dry_run:
logger.warning(
"dry_run: MLflow app %s has version below %s. "
"Job submission would create a new app (may take several minutes).",
resolved_arn, min_mlflow_version,
)
return resolved_arn
logger.info(
"Existing MLflow app %s has version below %s. Creating new app as default.",
resolved_arn, min_mlflow_version
Expand All @@ -258,6 +278,21 @@ def _resolve_mlflow_resource_arn(sagemaker_session, mlflow_resource_arn: Optiona
if resolved_app:
return resolved_app["Arn"]

# In dry_run mode, don't create a new app — just warn and return None
if dry_run:
if mlflow_apps_list:
# Apps exist but none are in a ready/usable state
logger.warning(
"dry_run: No MLflow app in ready state found. "
"Job submission would create a new app (may take several minutes)."
)
else:
logger.warning(
"dry_run: No MLflow app exists. "
"Job submission would create a new app (may take several minutes)."
)
return None

# Create new app
new_arn = _create_mlflow_app(sagemaker_session)
if new_arn:
Expand Down Expand Up @@ -679,6 +714,7 @@ def _get_fine_tuning_options_and_model_arn(model_name: str, customization_techni
except Exception as e:
logger.debug(f"Could not fetch subscription recipe override_params: {type(e).__name__}: {e}")

# Build union of supported instance types from both sources:
if options_dict:
return FineTuningOptions(options_dict), model_arn, is_gated_model
else:
Expand Down Expand Up @@ -969,22 +1005,27 @@ def _create_model_package_config(model_package_group_name, model, sagemaker_sess


def _create_mlflow_config(sagemaker_session, mlflow_resource_arn=None,
mlflow_experiment_name=None, mlflow_run_name=None):
mlflow_experiment_name=None, mlflow_run_name=None,
dry_run=False):
"""Create MLflow configuration with resolved resource ARN.

Args:
sagemaker_session: SageMaker session for resolving MLflow ARN
mlflow_resource_arn: MLflow resource ARN (if None, uses default experience)
mlflow_experiment_name: MLflow experiment name
mlflow_run_name: MLflow run name
dry_run: If True, only performs read-only checks without creating
new MLflow apps or waiting for apps in Creating status.

Returns:
MlflowConfig object or None if no MLflow resource ARN is resolved
"""


# Derive mlflow_resource_arn with default experience
resolved_mlflow_arn = _resolve_mlflow_resource_arn(sagemaker_session, mlflow_resource_arn)
resolved_mlflow_arn = _resolve_mlflow_resource_arn(
sagemaker_session, mlflow_resource_arn, dry_run=dry_run
)
logger.info(f"MLflow resource ARN: {resolved_mlflow_arn}")

# Create MlflowConfig using shapes
Expand Down
2 changes: 1 addition & 1 deletion sagemaker-train/src/sagemaker/train/dpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -286,7 +286,6 @@ def train(self,
)

logger.info(f"Training Job Name: {current_training_job_name}")
print(f"Training Job Name: {current_training_job_name}")

#data
input_data_config = _create_input_data_config(training_dataset or self.training_dataset,
Expand Down Expand Up @@ -316,6 +315,7 @@ def train(self,
mlflow_resource_arn=self.mlflow_resource_arn,
mlflow_experiment_name=self.mlflow_experiment_name,
mlflow_run_name=self.mlflow_run_name,
dry_run=dry_run,
)

final_hyperparameters = self.hyperparameters.to_dict()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -971,5 +971,5 @@ def _evaluate_hyperpod(self, subtask=None):

return self._submit_hyperpod_eval_job(
override_parameters=override_parameters,
base_job_name=f"eval-{self.benchmark.value}",
base_job_name=self.base_eval_name or f"eval-{self.benchmark.value}",
)
Original file line number Diff line number Diff line change
Expand Up @@ -787,5 +787,5 @@ def _evaluate_hyperpod(self):

return self._submit_hyperpod_eval_job(
override_parameters=override_parameters,
base_job_name="custom-eval",
base_job_name=self.base_eval_name or "custom-eval",
)
25 changes: 17 additions & 8 deletions sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -272,15 +272,18 @@ def train(
self,
training_dataset: Optional[Union[str, DataSet]] = None,
wait: bool = True,
dry_run: bool = False,
) -> AgentRFTJob:
"""Launch an Agentic RFT job.

Args:
training_dataset: Training dataset override.
wait: If True (default), block until job reaches terminal status.
dry_run: If True, runs validation without submitting a job.
Returns None on success.

Returns:
AgentRFTJob instance for tracking the job.
AgentRFTJob instance for tracking the job, or None if dry_run=True.
"""
sagemaker_session = TrainDefaults.get_sagemaker_session(
sagemaker_session=self.sagemaker_session
Expand All @@ -301,7 +304,11 @@ def train(

if training_dataset is not None:
self.training_dataset = training_dataset
job_config_doc = self._build_job_config_document()
job_config_doc = self._build_job_config_document(dry_run=dry_run)

if dry_run:
logger.info("Dry-run validation passed. No job submitted.")
return None

tags = _get_jumpstart_tags(self._model_name, get_sagemaker_hub_name())

Expand Down Expand Up @@ -361,14 +368,14 @@ def attach(cls, job_name: str, session=None) -> AgentRFTJob:

# ---- Private: JobConfigDocument construction ----

def _build_job_config_document(self) -> str:
def _build_job_config_document(self, dry_run: bool = False) -> str:
"""Build the JobConfigDocument JSON string conforming to v1_0_0 schema."""
config = {
"AgentConfig": self._build_agent_config(),
"InputDataConfig": self._build_input_data_config(),
"OutputDataConfig": self._build_output_data_config(),
"ModelPackageConfig": self._build_model_package_config(),
"TrainingConfig": self._build_training_config(),
"TrainingConfig": self._build_training_config(dry_run=dry_run),
}
if self.networking:
config["VpcConfig"] = {
Expand Down Expand Up @@ -444,12 +451,12 @@ def _build_model_package_config(self) -> dict:
)
return config

def _build_training_config(self) -> dict:
def _build_training_config(self, dry_run: bool = False) -> dict:
hyperparameters = getattr(self, "_final_hyperparameters", {})
config = {
"BaseModelArn": self._model_arn,
}
mlflow_config = self._build_mlflow_config()
mlflow_config = self._build_mlflow_config(dry_run=dry_run)
if mlflow_config:
config["MlflowConfig"] = mlflow_config
if self.accept_eula is not None:
Expand All @@ -462,7 +469,7 @@ def _build_training_config(self) -> dict:
config["HyperParameters"] = user_set
return config

def _build_mlflow_config(self) -> Optional[dict]:
def _build_mlflow_config(self, dry_run: bool = False) -> Optional[dict]:
arn = (
self.mlflow_app_arn.arn
if isinstance(self.mlflow_app_arn, MlflowApp)
Expand All @@ -472,7 +479,9 @@ def _build_mlflow_config(self) -> Optional[dict]:
session = self.sagemaker_session or TrainDefaults.get_sagemaker_session(
sagemaker_session=self.sagemaker_session
)
arn = _resolve_mlflow_resource_arn(session, None, min_mlflow_version=MIN_MLFLOW_VERSION)
arn = _resolve_mlflow_resource_arn(
session, None, min_mlflow_version=MIN_MLFLOW_VERSION, dry_run=dry_run
)
if not arn:
return None
logger.info("MLflow resource ARN: %s", arn)
Expand Down
1 change: 1 addition & 0 deletions sagemaker-train/src/sagemaker/train/rlaif_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -276,6 +276,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati
mlflow_resource_arn=self.mlflow_resource_arn,
mlflow_experiment_name=self.mlflow_experiment_name,
mlflow_run_name=self.mlflow_run_name,
dry_run=dry_run,
)

final_hyperparameters = self.hyperparameters.to_dict()
Expand Down
1 change: 1 addition & 0 deletions sagemaker-train/src/sagemaker/train/rlvr_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -477,6 +477,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None,
mlflow_resource_arn=self.mlflow_resource_arn,
mlflow_experiment_name=self.mlflow_experiment_name,
mlflow_run_name=self.mlflow_run_name,
dry_run=dry_run,
)

final_hyperparameters = self.hyperparameters.to_dict()
Expand Down
1 change: 1 addition & 0 deletions sagemaker-train/src/sagemaker/train/sft_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -368,6 +368,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati
mlflow_resource_arn=self.mlflow_resource_arn,
mlflow_experiment_name=self.mlflow_experiment_name,
mlflow_run_name=self.mlflow_run_name,
dry_run=dry_run,
)

final_hyperparameters = self.hyperparameters.to_dict()
Expand Down
Loading