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
55 changes: 45 additions & 10 deletions sagemaker-serve/src/sagemaker/serve/bedrock_model_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,20 +107,23 @@ class BedrockModelBuilder:
model artifacts, opt into resource reuse with ``reuse_resources=True``.

Args:
model: The model to deploy. Can be a ModelTrainer, MultiTurnRLTrainer,
TrainingJob, ModelPackage instance, or an S3 URI string pointing
to model artifacts (e.g., ``"s3://bucket/checkpoint/step_4/"``).
model: The model to deploy. Can be a ModelTrainer, BaseTrainer
(SFTTrainer, DPOTrainer, RLVRTrainer, RLAIFTrainer, etc.),
MultiTurnRLTrainer, TrainingJob, ModelPackage instance, or an
S3 URI string pointing to model artifacts
(e.g., ``"s3://bucket/checkpoint/step_4/"``).
"""

def __init__(
self,
model: Optional[Union[str, ModelTrainer, MultiTurnRLTrainer, AgentRFTJob, TrainingJob, ModelPackage]] = None,
model: Optional[Union[str, ModelTrainer, BaseTrainer, MultiTurnRLTrainer, AgentRFTJob, TrainingJob, ModelPackage]] = None,
):
"""Initialize BedrockModelBuilder.

Args:
model: Model to deploy. Accepts a ModelTrainer, MultiTurnRLTrainer,
AgentRFTJob, TrainingJob, ModelPackage, or S3 URI string.
model: Model to deploy. Accepts a ModelTrainer, BaseTrainer (SFTTrainer,
DPOTrainer, RLVRTrainer, etc.), MultiTurnRLTrainer, AgentRFTJob,
TrainingJob, ModelPackage, or S3 URI string.
"""
self._bedrock_client = None
self._sagemaker_client = None
Expand Down Expand Up @@ -755,6 +758,8 @@ def _fetch_model_package(self) -> Optional[ModelPackage]:
# No valid model package ARN — _get_s3_artifacts will resolve.
return None
if isinstance(self.model, (MultiTurnRLTrainer, AgentRFTJob)):
# NOTE: Must check MultiTurnRLTrainer before BaseTrainer since it
# extends BaseTrainer but requires output_model_package_arn (raises ValueError).
arn = self.model.output_model_package_arn
if not arn:
job_name = None
Expand Down Expand Up @@ -783,6 +788,17 @@ def _fetch_model_package(self) -> Optional[ModelPackage]:
return ModelPackage.get(mp_arn)
# No model package (e.g., HyperPod) — _get_s3_artifacts will resolve.
return None
if isinstance(self.model, BaseTrainer):
training_job = getattr(self.model, '_latest_training_job', None)
if training_job:
mp_arn = getattr(training_job, 'output_model_package_arn', None)
if mp_arn and isinstance(mp_arn, str):
try:
return ModelPackage.get(mp_arn)
except Exception:
pass
# No model package — _get_s3_artifacts will resolve from training job.
return None
return None

def _get_s3_artifacts(self) -> Optional[str]:
Expand All @@ -793,9 +809,10 @@ def _get_s3_artifacts(self) -> Optional[str]:
checkpoint URI from manifest.json in training job output.
2. If model_package exists, returns the model data source S3 URI, resolving
to the hf_merged checkpoint directory if it exists (required for Bedrock import).
3. If no model_package and model is a TrainingJob, reads model_artifacts.s3_model_artifacts.
4. If no model_package and model is a ModelTrainer/BaseTrainer, reads
_latest_training_job.model_artifacts.s3_model_artifacts.
3. If no model_package and model is a TrainingJob/ModelTrainer/BaseTrainer,
reads model_artifacts.s3_model_artifacts from the training job.
4. If model_artifacts is empty (e.g., HyperPod jobs), falls back to resolving
checkpoint from the Nova manifest in the training job's s3_output_path.

Returns:
S3 URI string of the model artifacts, or None if not available.
Expand Down Expand Up @@ -825,7 +842,25 @@ def _get_s3_artifacts(self) -> Optional[str]:
s3_path = self._s3_artifacts_from_training_job(training_job)
if s3_path:
logger.info("Resolved S3 artifacts from training job: %s", s3_path)
return s3_path
return s3_path

# For HyperPod jobs, model_artifacts may not be set. Try resolving
# the checkpoint from the Nova manifest in s3_output_path.
output_data_config = getattr(training_job, "output_data_config", None)
s3_output_path = getattr(output_data_config, "s3_output_path", None)
job_name = getattr(training_job, "training_job_name", None)
if s3_output_path and job_name:
try:
checkpoint_uri = resolve_nova_checkpoint_uri(
self.boto_session.client("s3"),
s3_output_path,
job_name,
)
if checkpoint_uri:
logger.info("Resolved checkpoint from manifest: %s", checkpoint_uri)
return checkpoint_uri
except Exception as e:
logger.debug("Could not resolve checkpoint from manifest: %s", e)

return None

Expand Down
54 changes: 54 additions & 0 deletions sagemaker-serve/tests/unit/test_bedrock_model_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,60 @@ def test_unknown_type_returns_none(self):
b.model = "unknown"
assert b._fetch_model_package() is None

def test_from_base_trainer_with_model_package_arn(self):
"""BaseTrainer with _latest_training_job.output_model_package_arn resolves package."""
b = _builder()
mock_trainer = Mock()
mock_trainer._latest_training_job = Mock()
mock_trainer._latest_training_job.output_model_package_arn = "arn:pkg"
b.model = mock_trainer
expected = Mock()

class _FakeModelPackage:
@staticmethod
def get(arn):
return expected

with patch(f"{MODULE}.ModelPackage", _FakeModelPackage), \
patch(f"{MODULE}.TrainingJob", type(None)), \
patch(f"{MODULE}.ModelTrainer", type(None)), \
patch(f"{MODULE}.MultiTurnRLTrainer", type(None)), \
patch(f"{MODULE}.AgentRFTJob", type(None)), \
patch(f"{MODULE}.BaseTrainer", type(mock_trainer)):
result = b._fetch_model_package()
assert result is expected

def test_from_base_trainer_without_model_package_arn(self):
"""BaseTrainer without output_model_package_arn returns None."""
b = _builder()
mock_trainer = Mock()
mock_trainer._latest_training_job = Mock()
mock_trainer._latest_training_job.output_model_package_arn = None
b.model = mock_trainer

with patch(f"{MODULE}.TrainingJob", type(None)), \
patch(f"{MODULE}.ModelTrainer", type(None)), \
patch(f"{MODULE}.MultiTurnRLTrainer", type(None)), \
patch(f"{MODULE}.AgentRFTJob", type(None)), \
patch(f"{MODULE}.BaseTrainer", type(mock_trainer)):
result = b._fetch_model_package()
assert result is None

def test_from_base_trainer_no_training_job(self):
"""BaseTrainer with no _latest_training_job returns None."""
b = _builder()
mock_trainer = Mock()
mock_trainer._latest_training_job = None
b.model = mock_trainer

with patch(f"{MODULE}.TrainingJob", type(None)), \
patch(f"{MODULE}.ModelTrainer", type(None)), \
patch(f"{MODULE}.MultiTurnRLTrainer", type(None)), \
patch(f"{MODULE}.AgentRFTJob", type(None)), \
patch(f"{MODULE}.BaseTrainer", type(mock_trainer)):
result = b._fetch_model_package()
assert result is None


# ── _get_s3_artifacts ───────────────────────────────────────────────────────

Expand Down