From 34414a64a3be78cdc34e71dae654d470156957f4 Mon Sep 17 00:00:00 2001 From: Mohamed Zeidan Date: Sun, 27 Sep 2026 16:07:20 -0700 Subject: [PATCH] fix: emit repack step for ModelBuilder.register/build in ModelStep (#5828, #5829) In v2, Model.register()/create() placed the Model instance (carrying sagemaker_session, role, model_data, entry_point, source_dir, ...) into the pipeline context, so ModelStep repacked it. In v3, ModelBuilder.register()/build() place the ModelBuilder into the context, but ModelStep._append_repack_model_step only accepted sagemaker.core.resources.Model, so a ModelBuilder fell through to 'No models to repack' and the model was never repacked with the user's source_code (#5828). The v3 core Model is a pydantic model with extra='forbid' and no sagemaker_session field, so it could never satisfy ModelStep's reads (#5829). Make the repack path ModelBuilder-aware: accept a ModelBuilder in the repack gate and map its attributes (model_name/role_arn/s3_model_data_url/source_code.requirements) to _RepackModelStep's parameters. Also stop passing v2-era args (dependencies/ output_path/output_kms_key) that are not part of the v3 _RepackModelStep signature and would leak into ModelTrainer; pass 'requirements' as the v3 step expects. --- .../sagemaker/mlops/workflow/model_step.py | 76 +++++++++++------ .../tests/unit/workflow/test_model_step.py | 82 ++++++++++++++++++- 2 files changed, 133 insertions(+), 25 deletions(-) diff --git a/sagemaker-mlops/src/sagemaker/mlops/workflow/model_step.py b/sagemaker-mlops/src/sagemaker/mlops/workflow/model_step.py index c493542cc1..3841fb39fb 100644 --- a/sagemaker-mlops/src/sagemaker/mlops/workflow/model_step.py +++ b/sagemaker-mlops/src/sagemaker/mlops/workflow/model_step.py @@ -211,9 +211,46 @@ def properties(self): """A Properties object representing the appropriate SageMaker response data model.""" return self._properties + @staticmethod + def _repack_inputs_for(model): + """Return the repack inputs for ``model`` mapped to _RepackModelStep's params. + + Handles both a v3 ``ModelBuilder`` (what ``ModelBuilder.register()``/``.build()`` + place into the pipeline context) and a legacy ``sagemaker.core.resources.Model``. + A ``ModelBuilder`` exposes these under different names (``role_arn``, + ``s3_model_data_url``, ``model_name``) and carries its inference requirements on + its ``source_code``; normalize them here. See GH #5828 / #5829. + """ + from sagemaker.serve.model_builder import ModelBuilder + + if isinstance(model, ModelBuilder): + source_code = getattr(model, "source_code", None) + requirements = getattr(source_code, "requirements", None) if source_code else None + return { + "name": getattr(model, "model_name", None), + "sagemaker_session": model.sagemaker_session, + "role": getattr(model, "role_arn", None), + "model_data": getattr(model, "s3_model_data_url", None), + "entry_point": getattr(model, "entry_point", None), + "source_dir": getattr(model, "source_dir", None), + "requirements": requirements, + } + # Legacy sagemaker.core.resources.Model path. + return { + "name": getattr(model, "name", None), + "sagemaker_session": getattr(model, "sagemaker_session", None), + "role": getattr(model, "role", None), + "model_data": getattr(model, "model_data", None), + "entry_point": getattr(model, "entry_point", None), + "source_dir": getattr(model, "source_dir", None), + "requirements": getattr(model, "requirements", None), + } + def _append_repack_model_step(self): """Create and append a `_RepackModelStep` for the runtime repack""" - if isinstance(self._model, Model): + from sagemaker.serve.model_builder import ModelBuilder + + if isinstance(self._model, (Model, ModelBuilder)): model_list = [self._model] else: logger.warning("No models to repack") @@ -224,27 +261,25 @@ def _append_repack_model_step(self): security_group_ids, subnets = self._resolve_repack_model_step_vpc_configs() for i, model in enumerate(model_list): + # need_runtime_repack holds the id() of the original model/builder object, + # so the membership test must run against ``model`` itself, not a wrapper. runtime_repack_flg = ( self._need_runtime_repack and id(model) in self._need_runtime_repack ) if runtime_repack_flg: - name_base = model.name or i + fields = self._repack_inputs_for(model) + name_base = fields["name"] or i repack_model_step = _RepackModelStep( name="{}-{}-{}".format(self.name, _REPACK_MODEL_NAME_BASE, name_base), sagemaker_session=( self._repack_model_step_settings.pop("sagemaker_session", None) - or self._model.sagemaker_session - or model.sagemaker_session + or fields["sagemaker_session"] ), - role=( - self._repack_model_step_settings.pop("role", None) - or self._model.role - or model.role - ), - model_data=model.model_data, - entry_point=model.entry_point, - source_dir=model.source_dir, - dependencies=model.dependencies, + role=(self._repack_model_step_settings.pop("role", None) or fields["role"]), + model_data=fields["model_data"], + entry_point=fields["entry_point"], + source_dir=fields["source_dir"], + requirements=fields["requirements"], subnets=subnets, security_group_ids=security_group_ids, description=( @@ -253,14 +288,6 @@ def _append_repack_model_step(self): ), depends_on=self.depends_on, retry_policies=self._repack_model_retry_policies, - output_path=( - self._repack_model_step_settings.pop("output_path", None) - or self._runtime_repack_output_prefix - ), - output_kms_key=( - self._repack_model_step_settings.pop("output_kms_key", None) - or model.model_kms_key - ), **self._repack_model_step_settings, ) self.steps.append(repack_model_step) @@ -296,9 +323,10 @@ def _resolve_repack_model_step_vpc_configs(self): subnets = self._repack_model_step_settings.pop("subnets", None) return security_group_ids, subnets - if self._model.vpc_config: - security_group_ids = self._model.vpc_config.get("SecurityGroupIds", None) - subnets = self._model.vpc_config.get("Subnets", None) + vpc_config = getattr(self._model, "vpc_config", None) + if vpc_config: + security_group_ids = vpc_config.get("SecurityGroupIds", None) + subnets = vpc_config.get("Subnets", None) return security_group_ids, subnets return None, None diff --git a/sagemaker-mlops/tests/unit/workflow/test_model_step.py b/sagemaker-mlops/tests/unit/workflow/test_model_step.py index 050f48cd67..3849261cfe 100644 --- a/sagemaker-mlops/tests/unit/workflow/test_model_step.py +++ b/sagemaker-mlops/tests/unit/workflow/test_model_step.py @@ -14,7 +14,7 @@ from __future__ import absolute_import -from unittest.mock import patch +from unittest.mock import Mock, patch def test_model_step_properties(): @@ -27,3 +27,83 @@ def test_model_step_properties(): step = ModelStep(name="model-step", step_args=step_args) assert step.name == "model-step" assert hasattr(step, "properties") + + +def _pipeline_session(): + from sagemaker.core.workflow.pipeline_context import PipelineSession + + ps = Mock(spec=PipelineSession) + ps.context = Mock() + ps.boto_region_name = "us-west-2" + return ps + + +class _FakeModelStepArgs: + """Mimics the _ModelStepArguments produced by ModelBuilder.register() under a + PipelineSession (a register/create_model_package request that needs a repack).""" + + def __init__(self, model, need_runtime_repack): + self.model = model + self.need_runtime_repack = need_runtime_repack + self.runtime_repack_output_prefix = "s3://bucket/prefix" + self.create_model_request = None + self.create_model_package_request = { + "InferenceSpecification": {"Containers": [{"ModelDataUrl": "s3://orig/model.tar.gz"}]} + } + + +def test_model_builder_register_appends_repack_step(): + """GH #5828/#5829: ModelBuilder.register() in a ModelStep must emit a repack step + (v2 parity) and rewire the container ModelDataUrl to the repacked artifact.""" + from sagemaker.serve.model_builder import ModelBuilder + from sagemaker.mlops.workflow import model_step as ms + + ps = _pipeline_session() + builder = Mock(spec=ModelBuilder) + builder.sagemaker_session = ps + builder.model_name = "my-model" + builder.role_arn = "arn:aws:iam::111122223333:role/R" + builder.s3_model_data_url = "s3://orig/model.tar.gz" + builder.entry_point = "inference.py" + builder.source_dir = "/code" + builder.source_code = Mock(requirements="requirements.txt") + builder.vpc_config = None + + step_args = _FakeModelStepArgs(builder, {id(builder)}) + fake_repack = Mock() + fake_repack.properties.ModelArtifacts.S3ModelArtifacts = "s3://repacked/model.tar.gz" + + with patch("sagemaker.core.workflow.utilities.validate_step_args_input"): + with patch.object(ms, "_RepackModelStep", return_value=fake_repack) as mock_repack: + step = ms.ModelStep(name="step", step_args=step_args) + + # A repack step was generated (the bug: none was on master). + assert len(step.steps) == 1 + # ModelBuilder attributes were mapped to the repack step's parameters. + _, kwargs = mock_repack.call_args + assert kwargs["role"] == "arn:aws:iam::111122223333:role/R" + assert kwargs["model_data"] == "s3://orig/model.tar.gz" + assert kwargs["entry_point"] == "inference.py" + assert kwargs["source_dir"] == "/code" + assert kwargs["requirements"] == "requirements.txt" + assert kwargs["sagemaker_session"] is ps + # The container now points at the repacked artifact. + container = step_args.create_model_package_request["InferenceSpecification"]["Containers"][0] + assert container["ModelDataUrl"] == "s3://repacked/model.tar.gz" + + +def test_no_repack_step_for_unrecognized_model_type(): + """An object that is neither a core Model nor a ModelBuilder yields no repack step.""" + from sagemaker.mlops.workflow import model_step as ms + + ps = _pipeline_session() + unknown = Mock() + unknown.sagemaker_session = ps + step_args = _FakeModelStepArgs(unknown, {id(unknown)}) + + with patch("sagemaker.core.workflow.utilities.validate_step_args_input"): + with patch.object(ms, "_RepackModelStep") as mock_repack: + step = ms.ModelStep(name="step", step_args=step_args) + + assert step.steps == [] + mock_repack.assert_not_called()