diff --git a/pyproject.toml b/pyproject.toml index 261d2185..38d32b27 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,8 +43,6 @@ dependencies = [ # TransformAmiVersion was added to the SageMaker service model in botocore 1.37.23. "boto3>=1.37.23,<2", "packaging>=23.0,<27", - # Supports forwarding inference_ami_version from Model.deploy. - "sagemaker>=2.240.0,<3", "pyarrow>=19.0.1,<25", # lower bound to avoid https://github.com/apache/arrow/issues/45283 "PyYAML~=6.0", "Pillow>=10.2,<13", diff --git a/src/autogluon/cloud/__init__.py b/src/autogluon/cloud/__init__.py index 7ad48811..ad6c6385 100644 --- a/src/autogluon/cloud/__init__.py +++ b/src/autogluon/cloud/__init__.py @@ -1,15 +1,4 @@ import logging -import os - -os.environ.setdefault("SAGEMAKER_SUPPRESS_V2_WARNING", "1") -import sagemaker -from packaging.version import Version - -if Version(sagemaker.__version__) >= Version("3.0"): - raise ImportError( - f"SageMaker SDK >= 3.0 is currently not supported (found {sagemaker.__version__}). " - "Please downgrade: pip install -U 'sagemaker<3'" - ) from autogluon.common.utils.log_utils import _add_stream_handler diff --git a/src/autogluon/cloud/backend/backend.py b/src/autogluon/cloud/backend/backend.py index 386976a3..60d98c25 100644 --- a/src/autogluon/cloud/backend/backend.py +++ b/src/autogluon/cloud/backend/backend.py @@ -7,8 +7,6 @@ import pandas as pd -from ..endpoint.endpoint import Endpoint - def dumps_ag_args(config: Dict[str, Any]) -> str: """Serialize the remote-training config to JSON, raising a user-facing error on failure. @@ -68,12 +66,7 @@ def initialize( self.predictor_type = predictor_type self.resource_prefix = resource_prefix or f"ag-cloud-{predictor_type}" self.original_features = None - self.endpoint: Optional[Endpoint] = None - - @abstractmethod - def parse_backend_fit_kwargs(self, kwargs: Dict) -> Dict[str, Any]: - """Parse backend specific kwargs and get them ready to be sent to fit call""" - raise NotImplementedError + self.endpoint_name: Optional[str] = None @abstractmethod def attach_job(self, job_name: str) -> None: @@ -139,11 +132,6 @@ def fit(self, **kwargs) -> None: """Fit AG on the backend""" raise NotImplementedError - @abstractmethod - def parse_backend_deploy_kwargs(self, kwargs: Dict) -> Dict[str, Any]: - """Parse backend specific kwargs and get them ready to be sent to deploy call""" - raise NotImplementedError - @abstractmethod def deploy(self, **kwargs) -> None: """Deploy and endpoint""" @@ -155,13 +143,13 @@ def cleanup_deployment(self, **kwargs) -> None: raise NotImplementedError @abstractmethod - def attach_endpoint(self, endpoint: Endpoint) -> None: - """Attach the backend to an existing endpoint""" + def attach_endpoint(self, endpoint: str) -> None: + """Attach the backend to an existing endpoint by name""" raise NotImplementedError @abstractmethod - def detach_endpoint(self) -> Endpoint: - """Detach the current endpoint and return it""" + def detach_endpoint(self) -> str: + """Detach the current endpoint and return its name""" raise NotImplementedError @abstractmethod @@ -174,11 +162,6 @@ def predict_proba_real_time(self, test_data: Union[str, pd.DataFrame], **kwargs) """Realtime prediction probability with the endpoint""" raise NotImplementedError - @abstractmethod - def parse_backend_predict_kwargs(self, kwargs: Dict) -> Dict[str, Any]: - """Parse backend specific kwargs and get them ready to be sent to predict call""" - raise NotImplementedError - @abstractmethod def get_batch_inference_job_info(self, job_name: Optional[str] = None) -> Dict[str, Any]: """ diff --git a/src/autogluon/cloud/backend/multimodal_sagemaker_backend.py b/src/autogluon/cloud/backend/multimodal_sagemaker_backend.py index 48141512..f5f30993 100644 --- a/src/autogluon/cloud/backend/multimodal_sagemaker_backend.py +++ b/src/autogluon/cloud/backend/multimodal_sagemaker_backend.py @@ -1,13 +1,11 @@ -import copy import os from typing import Any, Dict, Optional, Tuple, Union import pandas as pd -from sagemaker import Predictor from autogluon.common.loaders import load_pd -from ..utils.ag_sagemaker import AutoGluonMultiModalRealtimePredictor +from ..utils.serializers import MultiModalSerializer from ..utils.utils import convert_image_path_to_encoded_bytes_in_dataframe, is_image_file, read_image_bytes_and_encode from .constant import MULTIMODL_SAGEMAKER from .sagemaker_backend import SagemakerBackend @@ -15,11 +13,12 @@ class MultiModalSagemakerBackend(SagemakerBackend): name = MULTIMODL_SAGEMAKER + # Images are sent one file per request. + _IMAGE_BATCH_ARGS = dict(content_type="application/x-image", split_type=None, batch_strategy="SingleRecord") - @property - def _realtime_predictor_cls(self) -> Predictor: - """Class used for realtime endpoint""" - return AutoGluonMultiModalRealtimePredictor + def _realtime_serializer(self): + """Serializer used for realtime endpoint requests""" + return MultiModalSerializer() def _load_predict_real_time_test_data( self, test_data: Union[str, pd.DataFrame], test_data_image_column: str @@ -87,10 +86,9 @@ def predict_real_time( test_data, content_type = self._load_predict_real_time_test_data( test_data=test_data, test_data_image_column=test_data_image_column ) - # Providing content type here because sagemaker serializer doesn't support change content type dynamically. - # Pass to `endpoint.predict()` call as `initial_args` instead + # The serializer's content type is fixed, so the per-request content type is passed explicitly. pred, _ = self._predict_real_time( - test_data=test_data, accept=accept, inference_kwargs=inference_kwargs, ContentType=content_type + test_data=test_data, accept=accept, inference_kwargs=inference_kwargs, content_type=content_type ) return pred @@ -139,10 +137,9 @@ def predict_proba_real_time( test_data, content_type = self._load_predict_real_time_test_data( test_data=test_data, test_data_image_column=test_data_image_column ) - # Providing content type here because sagemaker serializer doesn't support change content type dynamically. - # Pass to `endpoint.predict()` call as `initial_args` instead + # The serializer's content type is fixed, so the per-request content type is passed explicitly. pred, proba = self._predict_real_time( - test_data=test_data, accept=accept, inference_kwargs=inference_kwargs, ContentType=content_type + test_data=test_data, accept=accept, inference_kwargs=inference_kwargs, content_type=content_type ) if proba is None: @@ -161,8 +158,6 @@ def predict( When minimizing latency isn't a concern, then the batch transform functionality may be easier, more scalable, and more appropriate. If you want to minimize latency, use `predict_real_time()` instead. To learn more: https://docs.aws.amazon.com/sagemaker/latest/dg/batch-transform.html - This method would first create a AutoGluonSagemakerInferenceModel with the trained predictor, - then create a transformer with it, and call transform in the end. Parameters ---------- @@ -182,10 +177,10 @@ def predict( image_modality_only = self._check_image_modality_only(test_data) if image_modality_only: - processed_args = self._prepare_image_predict_args(**kwargs) - kwargs["transformer_kwargs"] = processed_args["transformer_kwargs"] - kwargs["transform_kwargs"] = processed_args["transform_kwargs"] - return super().predict(test_data, test_data_image_column=None, **kwargs) + pred, _ = self._predict( + test_data, original_features=self.original_features, **kwargs, **self._IMAGE_BATCH_ARGS + ) + return pred else: return super().predict( test_data, @@ -204,8 +199,6 @@ def predict_proba( When minimizing latency isn't a concern, then the batch transform functionality may be easier, more scalable, and more appropriate. If you want to minimize latency, use `predict_real_time()` instead. To learn more: https://docs.aws.amazon.com/sagemaker/latest/dg/batch-transform.html - This method would first create a AutoGluonSagemakerInferenceModel with the trained predictor, - then create a transformer with it, and call transform in the end. Parameters ---------- @@ -225,10 +218,11 @@ def predict_proba( image_modality_only = self._check_image_modality_only(test_data) if image_modality_only: - processed_args = self._prepare_image_predict_args(**kwargs) - kwargs["transformer_kwargs"] = processed_args["transformer_kwargs"] - kwargs["transform_kwargs"] = processed_args["transform_kwargs"] - return super().predict_proba(test_data, test_data_image_column=None, **kwargs) + include_predict = kwargs.pop("include_predict", True) + pred, pred_proba = self._predict( + test_data, original_features=self.original_features, **kwargs, **self._IMAGE_BATCH_ARGS + ) + return (pred, pred_proba) if include_predict else pred_proba else: return super().predict_proba( test_data, @@ -236,22 +230,6 @@ def predict_proba( **kwargs, ) - def _prepare_image_predict_args(self, **predict_kwargs): - split_type = None - content_type = "application/x-image" - predict_kwargs = copy.deepcopy(predict_kwargs) - transformer_kwargs = predict_kwargs.pop("transformer_kwargs", {}) - if transformer_kwargs is None: - transformer_kwargs = {} - transformer_kwargs["strategy"] = "SingleRecord" - transform_kwargs = predict_kwargs.pop("transofrm_kwargs", {}) - if transform_kwargs is None: - transform_kwargs = {} - transform_kwargs["split_type"] = split_type - transform_kwargs["content_type"] = content_type - - return {"transformer_kwargs": transformer_kwargs, "transform_kwargs": transform_kwargs} - def _check_image_modality_only(self, test_data): image_modality_only = False if isinstance(test_data, str): diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index fc11ddce..c2d1902c 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -2,35 +2,43 @@ import json import logging import os -import shutil import tarfile import tempfile from typing import Any, Dict, List, Literal, Optional, Tuple, Union import pandas as pd -import sagemaker from botocore.exceptions import ClientError -from sagemaker import Predictor -from sagemaker.serverless import ServerlessInferenceConfig from autogluon.common.loaders import load_pd from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix from ..data import FormatConverterFactory -from ..endpoint.sagemaker_endpoint import SagemakerEndpoint from ..job import SageMakerBatchTransformationJob, SageMakerFitJob from ..scripts import ScriptManager from ..utils.ag_sagemaker import ( - AutoGluonBatchPredictor, - AutoGluonNonRepackInferenceModel, - AutoGluonRealtimePredictor, - AutoGluonRepackInferenceModel, + SOURCE_DIR_TARBALL_NAME, + repack_model_with_serving_code, + script_mode_environment, + staged_serving_code, + training_script_hyperparameters, + upload_training_code, ) from ..utils.aws_utils import resolve_execution_role, setup_sagemaker_session -from ..utils.constants import VALID_ACCEPT -from ..utils.dlc_utils import infer_sagemaker_ami_version, parse_framework_version -from ..utils.misc import MostRecentInsertedOrderedDict -from ..utils.serializers import AutoGluonSerializationWrapper +from ..utils.constants import LOCAL_MODE, LOCAL_MODE_GPU, VALID_ACCEPT +from ..utils.deserializers import PandasDeserializer +from ..utils.dlc_utils import infer_sagemaker_ami_version, parse_framework_version, retrieve_image_uri +from ..utils.misc import MostRecentInsertedOrderedDict, sagemaker_timestamp, unique_name_from_base +from ..utils.sagemaker_api import ( + BATCH_PREDICT_OVERRIDE_KEYS, + DEPLOY_OVERRIDE_KEYS, + FIT_OVERRIDE_KEYS, + check_override_keys, + deep_merge, + delete_endpoint, + delete_quietly, + invoke_endpoint, +) +from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer from ..utils.tag_utils import build_tags from ..utils.utils import ( convert_image_path_to_encoded_bytes_in_dataframe, @@ -43,6 +51,42 @@ logger = logging.getLogger(__name__) +SAGEMAKER_MODEL_SERVER_WORKERS = "SAGEMAKER_MODEL_SERVER_WORKERS" +_SERVERLESS_CONFIG_FIELDS = { + "memory_size_in_mb": "MemorySizeInMB", + "max_concurrency": "MaxConcurrency", + "provisioned_concurrency": "ProvisionedConcurrency", +} + + +def _reject_local_mode(instance_type: Optional[str]) -> None: + if instance_type in (LOCAL_MODE, LOCAL_MODE_GPU): + raise ValueError( + f"instance_type={instance_type!r} (SageMaker local mode) is no longer supported. " + "Use a SageMaker instance type such as 'ml.m5.2xlarge'." + ) + + +def _s3_channel(channel_name: str, s3_uri: str) -> Dict[str, Any]: + return { + "ChannelName": channel_name, + "DataSource": { + "S3DataSource": { + "S3DataType": "S3Prefix", + "S3Uri": s3_uri, + "S3DataDistributionType": "FullyReplicated", + } + }, + } + + +def _to_request_fields(settings: Dict[str, Any], fields: Dict[str, str], arg_name: str) -> Dict[str, Any]: + """Rename the snake_case keys of a user-facing settings dict to the SageMaker API field names.""" + unknown = sorted(set(settings) - set(fields)) + if unknown: + raise ValueError(f"Unsupported `{arg_name}` key(s) {unknown}. Valid keys: {list(fields)}.") + return {fields[key]: value for key, value in settings.items()} + class SagemakerBackend(Backend): name = SAGEMAKER @@ -63,19 +107,6 @@ def __init__( **kwargs, ) - @property - def _realtime_predictor_cls(self) -> Predictor: - """Class used for realtime endpoint""" - return AutoGluonRealtimePredictor - - def _resolve_tags( - self, - kwargs: Dict[str, Any], - extra_tags: Optional[List[Dict[str, str]]] = None, - ) -> None: - """In-place: replace ``kwargs['tags']`` with the merged default + extra + user tag list.""" - kwargs["tags"] = build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=kwargs.get("tags")) - def initialize(self, role: Optional[str] = None, **kwargs) -> None: """Initialize the backend. @@ -86,27 +117,26 @@ def initialize(self, role: Optional[str] = None, **kwargs) -> None: :func:`autogluon.cloud.utils.aws_utils.resolve_execution_role` for the resolution order. """ super().initialize(**kwargs) + self.sagemaker_session = setup_sagemaker_session() try: - self.role_arn = resolve_execution_role(role, backend_name=SAGEMAKER) + self.role_arn = resolve_execution_role(role, backend_name=SAGEMAKER, session=self.sagemaker_session) except ClientError as e: logger.warning( "Failed to resolve SageMaker execution role. Pass `role=` to the predictor/model " "or run `autogluon.cloud.bootstrap()` / `register()` to persist one." ) raise e - self.sagemaker_session = setup_sagemaker_session() - self.endpoint = None self._region = self.sagemaker_session.boto_region_name self._fit_job: SageMakerFitJob = SageMakerFitJob(session=self.sagemaker_session) self._batch_transform_jobs = MostRecentInsertedOrderedDict() - def parse_backend_fit_kwargs(self, kwargs: Dict) -> Dict[str, Any]: - """Parse backend specific kwargs and get them ready to be sent to fit call""" - return dict( - autogluon_sagemaker_estimator_kwargs=kwargs.get("autogluon_sagemaker_estimator_kwargs", None), - fit_kwargs=kwargs.get("fit_kwargs", None), - extra_ag_args=kwargs.get("extra_ag_args", None), - ) + def _realtime_serializer(self): + """Serializer used for realtime endpoint requests""" + return AutoGluonSerializer() + + def _resolve_tags(self, extra_tags: Optional[List[Dict[str, str]]] = None) -> List[Dict[str, str]]: + """Tags for a created SageMaker resource: default + extra tags.""" + return build_tags(self.predictor_type, extra_tags=extra_tags) def attach_job(self, job_name: str) -> None: """ @@ -118,7 +148,7 @@ def attach_job(self, job_name: str) -> None: job_name: str The name of the job being attached """ - self._fit_job = SageMakerFitJob.attach(job_name) + self._fit_job = SageMakerFitJob.attach(job_name, session=self.sagemaker_session) @property def is_fit(self) -> bool: @@ -175,8 +205,7 @@ def fit( custom_image_uri: Optional[str] = None, timeout: int = 24 * 60 * 60, wait: bool = True, - autogluon_sagemaker_estimator_kwargs: Optional[Dict] = None, - fit_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, extra_ag_args: Optional[Dict[str, Any]] = None, extra_tags: Optional[List[Dict[str, str]]] = None, ) -> None: @@ -221,12 +250,9 @@ def fit( Whether the call should wait until the job completes To be noticed, the function won't return immediately because there are some preparations needed prior fit. Use `get_fit_job_status` to get job status. - autogluon_sagemaker_estimator_kwargs: dict, default = dict() - Any extra arguments needed to initialize AutoGluonSagemakerEstimator - Please refer to https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#sagemaker.estimator.Estimator for all options - fit_kwargs: - Any extra arguments needed to pass to fit. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#sagemaker.estimator.Estimator.fit for all options + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw ``CreateTrainingJob`` request fields (SageMaker API / boto3 PascalCase names) under the + ``"create_training_job"`` key, deep-merged over the request built by AutoGluon-Cloud. extra_ag_args: Optional[Dict[str, Any]], default = None Additional entries to merge into ``ag_args.json``. Use this to ship caller-specific metadata to the train script (e.g. ``predict_after_fit``, ``save_predictor``, or ``id_column`` / @@ -234,6 +260,8 @@ def fit( """ if data_channels.get("train_data") is None: raise ValueError("`data_channels['train_data']` is required.") + _reject_local_mode(instance_type) + overrides = check_override_keys(backend_overrides, FIT_OVERRIDE_KEYS) predictor_fit_args = copy.deepcopy(predictor_fit_args) # Resolve any path inputs (str or pathlib.Path) into DataFrames so they can be CSV-uploaded as SageMaker channels. data_channels = { @@ -251,7 +279,7 @@ def fit( logger.log(20, f"Training with framework_version=={framework_version}") if not job_name: - job_name = sagemaker.utils.unique_name_from_base(self.resource_prefix) + job_name = unique_name_from_base(self.resource_prefix) if instance_count == "auto": instance_count = 1 @@ -261,38 +289,10 @@ def fit( ) instance_count = 1 - if autogluon_sagemaker_estimator_kwargs is None: - autogluon_sagemaker_estimator_kwargs = {} - autogluon_sagemaker_estimator_kwargs = copy.deepcopy(autogluon_sagemaker_estimator_kwargs) - autogluon_sagemaker_estimator_kwargs.pop("output_path", None) - if ( - autogluon_sagemaker_estimator_kwargs.get("disable_profiler", None) is None - and autogluon_sagemaker_estimator_kwargs.get("debugger_hook_config", None) is None - ): - autogluon_sagemaker_estimator_kwargs["disable_profiler"] = True - autogluon_sagemaker_estimator_kwargs["debugger_hook_config"] = False - max_run = autogluon_sagemaker_estimator_kwargs.get("max_run", None) - if max_run is None: - autogluon_sagemaker_estimator_kwargs["max_run"] = timeout - else: - logger.warning(f"Both `max_run`: {max_run} and `timeout`: {timeout} are specified. Will ignore timeout") - - output_path = self.cloud_output_path + "/model" - code_location = self.cloud_output_path + "/code" - self._train_script_path = ScriptManager.get_train_script( backend_type=self.name, framework_version=framework_version ) entry_point = self._train_script_path - user_entry_point = autogluon_sagemaker_estimator_kwargs.pop("entry_point", None) - if user_entry_point: - logger.warning( - f"Providing a custom entry point could break the fit. Please refer to `{entry_point}` for our implementation" - ) - entry_point = user_entry_point - else: - # Avoid user passing in source_dir without specifying entry point - autogluon_sagemaker_estimator_kwargs.pop("source_dir", None) ag_args = dict( predictor_init_args=predictor_init_args, @@ -325,36 +325,95 @@ def fit( backend_type=self.name, framework_version=framework_version ), # Training and Inference should have the same framework_version ) - if fit_kwargs is None: - fit_kwargs = {} - self._resolve_tags(autogluon_sagemaker_estimator_kwargs, extra_tags) - - self._fit_job.run( - role=self.role_arn, + code_uri = upload_training_code( entry_point=entry_point, - region=self._region, - instance_type=instance_type, - instance_count=instance_count, - volume_size=volume_size, - framework_version=framework_version, - py_version=py_version, - base_job_name=f"{self.resource_prefix}-train", - output_path=output_path, - code_location=code_location, - inputs=inputs, - custom_image_uri=custom_image_uri, - wait=wait, - job_name=job_name, - autogluon_sagemaker_estimator_kwargs=autogluon_sagemaker_estimator_kwargs, - **fit_kwargs, + sagemaker_session=self.sagemaker_session, + s3_uri_prefix=f"{self.cloud_output_path}/code/{job_name}/source", ) - def parse_backend_deploy_kwargs(self, kwargs: Dict) -> Dict[str, Any]: - """Parse backend specific kwargs and get them ready to be sent to deploy call""" - model_kwargs = kwargs.get("model_kwargs", None) - deploy_kwargs = kwargs.get("deploy_kwargs", None) + request: Dict[str, Any] = { + "TrainingJobName": job_name, + "RoleArn": self.role_arn, + "AlgorithmSpecification": { + "TrainingImage": retrieve_image_uri( + framework_version, self._region, "training", instance_type, py_version, custom_image_uri + ), + "TrainingInputMode": "File", + }, + # The code is passed as an input channel (not read from S3 by the container), like SageMaker SDK v2 did, + # so training also works with `EnableNetworkIsolation`. + "HyperParameters": training_script_hyperparameters( + entry_point=entry_point, + submit_directory=f"/opt/ml/input/data/code/{SOURCE_DIR_TARBALL_NAME}", + job_name=job_name, + region=self._region, + ), + "InputDataConfig": [_s3_channel(name, uri) for name, uri in {**inputs, "code": code_uri}.items()], + "OutputDataConfig": {"S3OutputPath": self.cloud_output_path + "/model"}, + "ResourceConfig": { + "InstanceType": instance_type, + "InstanceCount": instance_count, + "VolumeSizeInGB": volume_size, + }, + "StoppingCondition": {"MaxRuntimeInSeconds": timeout}, + "ProfilerConfig": {"DisableProfiler": True}, + "Tags": self._resolve_tags(extra_tags), + } + request = deep_merge(request, overrides.get("create_training_job", {})) + + self._fit_job = SageMakerFitJob(session=self.sagemaker_session) + self._fit_job.run(training_job_request=request, framework_version=framework_version, wait=wait) - return dict(model_kwargs=model_kwargs, deploy_kwargs=deploy_kwargs) + def _create_model( + self, + model_name: str, + model_data: str, + image_uri: str, + entry_point: str, + environment: Dict[str, str], + tags: List[Dict[str, str]], + overrides: Dict[str, Dict[str, Any]], + ) -> str: + """Create a SageMaker model serving ``model_data`` (which must contain the serving code under ``code/``).""" + # PYTHONUNBUFFERED disables output buffering for endpoint logging. + container_environment = { + "PYTHONUNBUFFERED": "1", + **environment, + **script_mode_environment(entry_point, self._region), + } + request: Dict[str, Any] = { + "ModelName": model_name, + "PrimaryContainer": { + "Image": image_uri, + "ModelDataUrl": model_data, + "Environment": container_environment, + }, + "ExecutionRoleArn": self.role_arn, + "Tags": tags, + } + request = deep_merge(request, overrides.get("create_model", {})) + logger.log(20, "Creating inference model...") + self.sagemaker_session.sagemaker_client.create_model(**request) + logger.log(20, "Inference model created successfully") + return request["ModelName"] + + def _prepare_model_data( + self, + predictor_path: str, + entry_point: str, + repack: bool, + repacked_model_uri: str, + ) -> str: + """Return an S3 model tarball that contains the serving code, repacking ``predictor_path`` if needed.""" + if not repack: + return predictor_path + logger.log(20, "Repacking the serving code into the model artifact...") + return repack_model_with_serving_code( + model_data=predictor_path, + entry_point=entry_point, + repacked_model_uri=repacked_model_uri, + sagemaker_session=self.sagemaker_session, + ) def deploy( self, @@ -366,8 +425,8 @@ def deploy( custom_image_uri: Optional[str] = None, volume_size: Optional[int] = None, wait: bool = True, - model_kwargs: Optional[Dict] = None, - deploy_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + entry_point: Optional[str] = None, fm_serve_config: Optional[Dict[str, Any]] = None, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, @@ -376,8 +435,7 @@ def deploy( ) -> None: """ Deploy a predictor as a SageMaker endpoint, which can be used to do real-time inference later. - This method would first create a AutoGluonSagemakerInferenceModel with the trained predictor, - and then deploy it to the endpoint. + This method creates a SageMaker model with the trained predictor, an endpoint config, and the endpoint. Parameters ---------- @@ -407,34 +465,44 @@ def deploy( wait: Bool, default = True, Whether to wait for the endpoint to be deployed. To be noticed, the function won't return immediately because there are some preparations needed prior deployment. - model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - deploy_kwargs: - Any extra arguments needed to pass to deploy. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#sagemaker.model.Model.deploy for all options + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw request fields (SageMaker API / boto3 PascalCase names) deep-merged over the requests built by + AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"production_variant"``, + ``"create_endpoint_config"``, ``"create_endpoint"``. + entry_point: Optional[str], default = None + Serve script to use instead of the predictor type's default. fm_serve_config: Optional[Dict[str, Any]], default = None Configuration dict passed to the FM serve script via the AG_FM_SERVE_CONFIG env var. inference_mode: {"realtime", "serverless"}, default = "realtime" Endpoint type. ``"serverless"`` provisions a SageMaker Serverless Inference endpoint (no instance management, scales to zero). inference_config: Optional[Dict[str, Any]], default = None - Mode-specific overrides forwarded to `sagemaker.serverless.ServerlessInferenceConfig` - (e.g. ``memory_size_in_mb``, ``max_concurrency``). + Serverless overrides forwarded to the production variant's ``ServerlessConfig`` + (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). repack: bool, default = True - Whether the SageMaker SDK should download ``predictor_path``, inject the entry-point script, and re-upload - it. Set to False when ``predictor_path`` already contains the serve script (e.g. an artifact bundled by + Whether to download ``predictor_path``, inject the serve script, and re-upload it. Set to False when + ``predictor_path`` already contains the serve script (e.g. an artifact bundled by :meth:`FoundationModel.cache_model_artifact`) to skip the round-trip. Ignored when ``predictor_path`` is None. """ - assert self.endpoint is None, ( + assert self.endpoint_name is None, ( "There is an endpoint already attached. Either detach it with `detach` or clean it up with `cleanup_deployment`" ) + if inference_mode not in ("realtime", "serverless"): + raise ValueError(f"Unsupported inference_mode={inference_mode!r}") + overrides = check_override_keys(backend_overrides, DEPLOY_OVERRIDE_KEYS) + serverless_config = None + if inference_mode == "serverless": + preset = {"memory_size_in_mb": 4096, "max_concurrency": 5} + serverless_config = _to_request_fields( + {**preset, **(inference_config or {})}, _SERVERLESS_CONFIG_FIELDS, "inference_config" + ) if inference_mode == "serverless" and instance_type is None: # Needed to infer the container image (CPU vs GPU) downstream — serverless is CPU-only. instance_type = "ml.m5.2xlarge" + _reject_local_mode(instance_type) if not endpoint_name: - endpoint_name = sagemaker.utils.unique_name_from_base(self.resource_prefix) + endpoint_name = unique_name_from_base(self.resource_prefix) # Resolve container image if custom_image_uri: @@ -461,109 +529,99 @@ def deploy( if predictor_path: predictor_path = self._upload_predictor(predictor_path, f"endpoints/{endpoint_name}/predictor") - # Resolve entry point - if model_kwargs is None: - model_kwargs = {} - model_kwargs = copy.deepcopy(model_kwargs) - user_entry_point = model_kwargs.pop("entry_point", None) - if user_entry_point: - entry_point = user_entry_point - else: + user_entry_point = entry_point + if entry_point is None: self._serve_script_path = ScriptManager.get_serve_script( backend_type=self.name, framework_version=framework_version ) entry_point = self._serve_script_path - # Pick model class. The question is whether the tarball already contains - # the entry_point script — if yes, NonRepack uses it as-is; if no, Repack - # injects the script at deploy time. + # Decide whether the tarball already contains the entry_point script — if yes, use it as-is; + # if no, repack the script into it. if predictor_path is None: - predictor_path = self._create_serve_script_tarball(entry_point, endpoint_name) - model_cls = AutoGluonNonRepackInferenceModel - elif not repack: - model_cls = AutoGluonNonRepackInferenceModel + model_data = self._create_serve_script_tarball(entry_point, endpoint_name) else: is_default_fit_output = ( self._fit_job is not None and predictor_path == self._fit_job.get_output_path() and user_entry_point is None ) - model_cls = AutoGluonNonRepackInferenceModel if is_default_fit_output else AutoGluonRepackInferenceModel - - # Assemble env vars and deploy - predictor_cls = self._realtime_predictor_cls - user_predictor_cls = model_kwargs.pop("predictor_cls", None) - if user_predictor_cls: - logger.warning( - "Providing a custom predictor_cls could break the deployment.", - "Please refer to `AutoGluonRealtimePredictor` for how to provide a custom predictor", + model_data = self._prepare_model_data( + predictor_path, + entry_point=entry_point, + repack=repack and not is_default_fit_output, + repacked_model_uri=f"{self.cloud_output_path}/endpoints/{endpoint_name}/model/model.tar.gz", ) - predictor_cls = user_predictor_cls - - model_kwargs_env = model_kwargs.pop("env", None) - SAGEMAKER_MODEL_SERVER_WORKERS = "SAGEMAKER_MODEL_SERVER_WORKERS" - if model_kwargs_env is not None: - if ( - SAGEMAKER_MODEL_SERVER_WORKERS in model_kwargs_env - and int(model_kwargs_env[SAGEMAKER_MODEL_SERVER_WORKERS]) > 1 - ): - logger.warning( - f"Setting {SAGEMAKER_MODEL_SERVER_WORKERS} to value larger than 1 might cause running out of RAM and/or GPU RAM" - ) - else: - model_kwargs_env[SAGEMAKER_MODEL_SERVER_WORKERS] = "1" - else: - model_kwargs_env = {SAGEMAKER_MODEL_SERVER_WORKERS: "1"} + container_environment = {SAGEMAKER_MODEL_SERVER_WORKERS: "1"} if fm_serve_config is not None: - model_kwargs_env["AG_FM_SERVE_CONFIG"] = json.dumps(fm_serve_config) - + container_environment["AG_FM_SERVE_CONFIG"] = json.dumps(fm_serve_config) if inference_mode == "serverless": # Serverless containers run with `/` as cwd and a read-only root, so TorchServe's # default `logs/` path resolves to `/logs` and startup fails. Redirect to /tmp. - model_kwargs_env.setdefault("LOG_LOCATION", "/tmp") - model_kwargs_env.setdefault("METRICS_LOCATION", "/tmp") - - model = model_cls( - model_data=predictor_path, - role=self.role_arn, - region=self._region, - framework_version=framework_version, - py_version=py_version, - instance_type=instance_type, - custom_image_uri=custom_image_uri, + container_environment.setdefault("LOG_LOCATION", "/tmp") + container_environment.setdefault("METRICS_LOCATION", "/tmp") + + tags = self._resolve_tags(extra_tags) + model_name = self._create_model( + model_name=endpoint_name, + model_data=model_data, + image_uri=retrieve_image_uri( + framework_version, self._region, "inference", instance_type, py_version, custom_image_uri + ), entry_point=entry_point, - predictor_cls=predictor_cls, - env=model_kwargs_env, - **model_kwargs, - ) - deploy_kwargs = copy.deepcopy(deploy_kwargs or {}) - self._resolve_tags(deploy_kwargs, extra_tags) - inference_ami_version = infer_sagemaker_ami_version( - custom_image_uri, - instance_type, - image_scope="inference", + environment=container_environment, + tags=tags, + overrides=overrides, ) - if inference_ami_version is not None: - deploy_kwargs.setdefault("inference_ami_version", inference_ami_version) - instance_kwargs = { - "instance_type": instance_type, - "initial_instance_count": initial_instance_count, - "volume_size": volume_size, - } - user_config = inference_config or {} + variant: Dict[str, Any] = {"VariantName": "AllTraffic", "ModelName": model_name} if inference_mode == "realtime": - mode_kwargs = instance_kwargs - elif inference_mode == "serverless": - preset = {"memory_size_in_mb": 4096, "max_concurrency": 5} - mode_kwargs = {"serverless_inference_config": ServerlessInferenceConfig(**{**preset, **user_config})} + variant["InstanceType"] = instance_type + variant["InitialInstanceCount"] = initial_instance_count + if volume_size: + variant["VolumeSizeInGB"] = volume_size + inference_ami_version = infer_sagemaker_ami_version( + custom_image_uri, instance_type, image_scope="inference" + ) + if inference_ami_version is not None: + variant["InferenceAmiVersion"] = inference_ami_version else: - raise ValueError(f"Unsupported inference_mode={inference_mode!r}") + variant["ServerlessConfig"] = serverless_config + variant = deep_merge(variant, overrides.get("production_variant", {})) + + endpoint_config_request: Dict[str, Any] = { + "EndpointConfigName": endpoint_name, + "ProductionVariants": [variant], + "Tags": tags, + } + endpoint_config_request = deep_merge(endpoint_config_request, overrides.get("create_endpoint_config", {})) + endpoint_request = deep_merge( + { + "EndpointName": endpoint_name, + "EndpointConfigName": endpoint_config_request["EndpointConfigName"], + "Tags": tags, + }, + overrides.get("create_endpoint", {}), + ) logger.log(20, f"Deploying model to the endpoint (inference_mode={inference_mode})") - predictor = model.deploy(endpoint_name=endpoint_name, wait=wait, **mode_kwargs, **deploy_kwargs) - self.endpoint = SagemakerEndpoint(predictor) + client = self.sagemaker_session.sagemaker_client + try: + client.create_endpoint_config(**endpoint_config_request) + try: + client.create_endpoint(**endpoint_request) + except Exception: + delete_quietly( + client.delete_endpoint_config, EndpointConfigName=endpoint_config_request["EndpointConfigName"] + ) + raise + except Exception: + delete_quietly(client.delete_model, ModelName=model_name) + raise + self.endpoint_name = endpoint_request["EndpointName"] + if wait: + client.get_waiter("endpoint_in_service").wait(EndpointName=self.endpoint_name) def _create_serve_script_tarball(self, serve_script_path: str, endpoint_name: str) -> str: """Create a minimal model.tar.gz containing the serve script + serving_utils/ under code/.""" @@ -581,38 +639,29 @@ def cleanup_deployment(self) -> None: """ Delete endpoint, endpoint configuration and deployed model """ - assert self.endpoint is not None, "No deployed endpoint detected" - self.endpoint.delete_endpoint() - self.endpoint = None + assert self.endpoint_name is not None, "No deployed endpoint detected" + delete_endpoint(self.endpoint_name, self.sagemaker_session) + self.endpoint_name = None - def attach_endpoint(self, endpoint: Union[str, SagemakerEndpoint]) -> None: + def attach_endpoint(self, endpoint: str) -> None: """ Attach the current backend to an existing SageMaker endpoint. Parameters ---------- - endpoint: str or :class:`SagemakerEndpoint` - If str is passed, it should be the name of the endpoint being attached to. + endpoint: str + Name of the endpoint being attached to. """ - assert self.endpoint is None, ( + assert self.endpoint_name is None, ( "There is an endpoint already attached. Either detach it with `detach` or clean it up with `cleanup_deployment`" ) - if isinstance(endpoint, str): - endpoint = self._realtime_predictor_cls( - endpoint_name=endpoint, - sagemaker_session=self.sagemaker_session, - ) - self.endpoint = SagemakerEndpoint(endpoint) - elif isinstance(endpoint, SagemakerEndpoint): - self.endpoint = endpoint - else: - raise ValueError(f"Please provide either an endpoint name or an endpoint of type `{SagemakerEndpoint}`") + self.endpoint_name = endpoint - def detach_endpoint(self) -> SagemakerEndpoint: - """Detach the current endpoint and return it""" - assert self.endpoint is not None, "There is no attached endpoint" - detached_endpoint = self.endpoint - self.endpoint = None + def detach_endpoint(self) -> str: + """Detach the current endpoint and return its name""" + assert self.endpoint_name is not None, "There is no attached endpoint" + detached_endpoint = self.endpoint_name + self.endpoint_name = None return detached_endpoint def predict_real_time( @@ -693,24 +742,6 @@ def predict_proba_real_time( return proba - def parse_backend_predict_kwargs(self, kwargs: Dict) -> Dict[str, Any]: - """Parse backend specific kwargs and get them ready to be sent to predict call""" - download = kwargs.get("download", True) - persist = kwargs.get("persist", True) - save_path = kwargs.get("save_path", None) - model_kwargs = kwargs.get("model_kwargs", None) - transformer_kwargs = kwargs.get("transformer_kwargs", None) - transform_kwargs = kwargs.get("transform_kwargs", None) - - return dict( - download=download, - persist=persist, - save_path=save_path, - model_kwargs=model_kwargs, - transformer_kwargs=transformer_kwargs, - transform_kwargs=transform_kwargs, - ) - def get_batch_inference_job_info(self, job_name: Optional[str] = None) -> Optional[Dict[str, Any]]: """ Get general info of the batch inference job. @@ -767,20 +798,15 @@ def predict( instance_count: int = 1, custom_image_uri: Optional[str] = None, wait: bool = True, - download: bool = True, - persist: bool = True, - save_path: Optional[str] = None, - model_kwargs: Optional[Dict] = None, - transformer_kwargs: Optional[Dict] = None, - transform_kwargs: Optional[Dict] = None, + predictions_path: Optional[str] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Predict using SageMaker batch transform. When minimizing latency isn't a concern, then the batch transform functionality may be easier, more scalable, and more appropriate. If you want to minimize latency, use `predict_real_time()` instead. To learn more: https://docs.aws.amazon.com/sagemaker/latest/dg/batch-transform.html - This method would first create a AutoGluonSagemakerInferenceModel with the trained predictor, - then create a transformer with it, and call transform in the end. + This method creates a SageMaker model with the trained predictor and runs a transform job with it. Parameters ---------- @@ -808,33 +834,18 @@ def predict( wait: bool, default = True Whether to wait for batch transform to complete. To be noticed, the function won't return immediately because there are some preparations needed prior transform. - download: bool, default = True - Whether to download the batch transform results to the disk and load it after the batch transform finishes. - Will be ignored if `wait` is `False`. - persist: bool, default = True - Whether to persist the downloaded batch transform results on the disk. - Will be ignored if `download` is `False` - save_path: str, default = None, - Path to save the downloaded result. - Will be ignored if `download` is `False`. - If None, CloudPredictor will create one. - If `persist` is `False`, file would first be downloaded to this path and then removed. - model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - transformer_kwargs: dict - Any extra arguments needed to pass to transformer. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer for all options. - transform_kwargs: - Any extra arguments needed to pass to transform. - Please refer to - https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer.transform for all options. + predictions_path: Optional[str], default = None + S3 prefix under which the batch transform job writes its results (``/.out``). + Defaults to ``{cloud_output_path}/batch_transform//results``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw request fields (SageMaker API / boto3 PascalCase names) deep-merged over the requests built by + AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"create_transform_job"``. Returns ------- Optional Pandas.Series - Predict results in Series if `download` is True - None if `download` is False + Predict results in Series if `wait` is True + None if `wait` is False """ pred, _ = self._predict( test_data=test_data, @@ -846,12 +857,8 @@ def predict( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, - model_kwargs=model_kwargs, - transformer_kwargs=transformer_kwargs, - transform_kwargs=transform_kwargs, + predictions_path=predictions_path, + backend_overrides=backend_overrides, original_features=self.original_features, ) @@ -869,20 +876,15 @@ def predict_proba( instance_count: int = 1, custom_image_uri: Optional[str] = None, wait: bool = True, - download: bool = True, - persist: bool = True, - save_path: Optional[str] = None, - model_kwargs: Optional[Dict] = None, - transformer_kwargs: Optional[Dict] = None, - transform_kwargs: Optional[Dict] = None, + predictions_path: Optional[str] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]]: """ Predict using SageMaker batch transform. When minimizing latency isn't a concern, then the batch transform functionality may be easier, more scalable, and more appropriate. If you want to minimize latency, use `predict_real_time()` instead. To learn more: https://docs.aws.amazon.com/sagemaker/latest/dg/batch-transform.html - This method would first create a AutoGluonSagemakerInferenceModel with the trained predictor, - then create a transformer with it, and call transform in the end. + This method creates a SageMaker model with the trained predictor and runs a transform job with it. Parameters ---------- @@ -913,34 +915,19 @@ def predict_proba( wait: bool, default = True Whether to wait for batch transform to complete. To be noticed, the function won't return immediately because there are some preparations needed prior transform. - download: bool, default = True - Whether to download the batch transform results to the disk and load it after the batch transform finishes. - Will be ignored if `wait` is `False`. - persist: bool, default = True - Whether to persist the downloaded batch transform results on the disk. - Will be ignored if `download` is `False` - save_path: str, default = None, - Path to save the downloaded result. - Will be ignored if `download` is `False`. - If None, CloudPredictor will create one. - If `persist` is `False`, file would first be downloaded to this path and then removed. - model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - transformer_kwargs: dict - Any extra arguments needed to pass to transformer. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer for all options. - transform_kwargs: - Any extra arguments needed to pass to transform. - Please refer to - https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer.transform for all options. + predictions_path: Optional[str], default = None + S3 prefix under which the batch transform job writes its results (``/.out``). + Defaults to ``{cloud_output_path}/batch_transform//results``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw request fields (SageMaker API / boto3 PascalCase names) deep-merged over the requests built by + AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"create_transform_job"``. Returns ------- Optional[Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]] - If `download` is False, will return None or (None, None) if `include_predict` is True - If `download` is True and `include_predict` is True, + If `wait` is False, will return None or (None, None) if `include_predict` is True + If `wait` is True and `include_predict` is True, will return (prediction, predict_probability), where prediction is a Pandas.Series and predict_probability is a Pandas.DataFrame or a Pandas.Series that's identical to prediction when it's a regression problem. """ @@ -954,12 +941,8 @@ def predict_proba( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, - model_kwargs=model_kwargs, - transformer_kwargs=transformer_kwargs, - transform_kwargs=transform_kwargs, + predictions_path=predictions_path, + backend_overrides=backend_overrides, original_features=self.original_features, ) @@ -1018,7 +1001,7 @@ def get_fit_predict_results(self) -> pd.DataFrame: bucket, key = s3_path_to_bucket_prefix(predictions_path) with tempfile.TemporaryDirectory(prefix="ag_fit_predict_") as tmpdir: local_path = os.path.join(tmpdir, os.path.basename(key)) - self.sagemaker_session.boto_session.client("s3").download_file(bucket, key, local_path) + self.sagemaker_session.s3_client.download_file(bucket, key, local_path) return load_pd.load(local_path) def _download_ag_args_from_job(self) -> Dict[str, Any]: @@ -1029,7 +1012,7 @@ def _download_ag_args_from_job(self) -> Dict[str, Any]: """ job_name = self._fit_job.job_name assert job_name is not None, "No fit job found. Call `fit()` / `fit_predict()` first." - desc = self.sagemaker_session.describe_training_job(job_name) + desc = self.sagemaker_session.sagemaker_client.describe_training_job(TrainingJobName=job_name) channels = {c["ChannelName"]: c["DataSource"]["S3DataSource"]["S3Uri"] for c in desc["InputDataConfig"]} ag_args_uri = channels.get("ag_args") assert ag_args_uri is not None, ( @@ -1039,7 +1022,7 @@ def _download_ag_args_from_job(self) -> Dict[str, Any]: assert key.endswith(".json"), f"Expected ag_args channel to point to a .json file, got {ag_args_uri!r}" with tempfile.TemporaryDirectory(prefix="ag_args_") as tmpdir: local_path = os.path.join(tmpdir, os.path.basename(key)) - self.sagemaker_session.boto_session.client("s3").download_file(bucket, key, local_path) + self.sagemaker_session.s3_client.download_file(bucket, key, local_path) with open(local_path, "r") as f: return json.load(f) @@ -1128,15 +1111,10 @@ def _upload_fit_artifact( return inputs def _upload_serving_files(self, entry_point: str, bucket: str, key_prefix: str) -> str: - staging_dir = tempfile.mkdtemp(prefix="ag_serving_") - try: - shutil.copy(entry_point, os.path.join(staging_dir, os.path.basename(entry_point))) - shutil.copytree(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, os.path.join(staging_dir, "serving_utils")) + with staged_serving_code(entry_point) as staging_dir: return self.sagemaker_session.upload_data( path=staging_dir, bucket=bucket, key_prefix=key_prefix + "/serving" ) - finally: - shutil.rmtree(staging_dir, ignore_errors=True) def _upload_fit_image_artifact(self, image_dir_path, bucket, key_prefix): upload_image_path = None @@ -1171,7 +1149,7 @@ def _upload_predictor(self, predictor_path, key_prefix): return predictor_path def _validate_predict_real_time_args(self, accept): - assert self.endpoint is not None, "Please call `deploy()` to deploy an endpoint first." + assert self.endpoint_name is not None, "Please call `deploy()` to deploy an endpoint first." assert accept in VALID_ACCEPT, f"Invalid accept type: {accept}. Options are {VALID_ACCEPT}." def _load_predict_real_time_test_data(self, test_data, test_data_image_column): @@ -1183,11 +1161,19 @@ def _load_predict_real_time_test_data(self, test_data, test_data_image_column): return test_data - def _predict_real_time(self, test_data, accept, split_pred_proba=True, inference_kwargs=None, **initial_args): + def _predict_real_time(self, test_data, accept, split_pred_proba=True, inference_kwargs=None, content_type=None): try: if not isinstance(test_data, AutoGluonSerializationWrapper): test_data = AutoGluonSerializationWrapper(data=test_data, inference_kwargs=inference_kwargs) - prediction = self.endpoint.predict(test_data, initial_args={"Accept": accept, **initial_args}) + prediction = invoke_endpoint( + self.endpoint_name, + self.sagemaker_session, + test_data, + serializer=self._realtime_serializer(), + deserializer=PandasDeserializer(), + content_type=content_type, + accept=accept, + ) pred, pred_proba = None, None pred = prediction if split_pred_proba: @@ -1223,15 +1209,20 @@ def _predict( instance_count=1, custom_image_uri=None, wait=True, - download=True, - persist=True, - save_path=None, - model_kwargs=None, - transformer_kwargs=None, + predictions_path=None, + backend_overrides=None, split_pred_proba=True, - transform_kwargs=None, original_features=None, + content_type="text/csv", + split_type="Line", + accept="application/json", + assemble_with="Line", + batch_strategy="MultiRecord", ): + _reject_local_mode(instance_type) + overrides = check_override_keys(backend_overrides, BATCH_PREDICT_OVERRIDE_KEYS) + if predictions_path is not None and not is_s3_url(predictions_path): + raise ValueError(f"`predictions_path` must be an S3 URL, got {predictions_path!r}.") if not predictor_path: predictor_path = self._fit_job.get_output_path() assert predictor_path, "No cloud trained model found." @@ -1245,20 +1236,14 @@ def _predict( ) logger.log(20, f"Predicting with framework_version=={framework_version}") - if transform_kwargs is None: - transform_kwargs = {} - output_path = transform_kwargs.get("output_path", None) - if not output_path: - output_path = self.cloud_output_path - assert is_s3_url(output_path) - output_path = output_path + "/batch_transform" + f"/{sagemaker.utils.sagemaker_timestamp()}" + output_path = self.cloud_output_path + "/batch_transform" + f"/{sagemaker_timestamp()}" cloud_bucket, cloud_key_prefix = s3_path_to_bucket_prefix(output_path) logger.log(20, "Preparing autogluon predictor...") predictor_path = self._upload_predictor(predictor_path, cloud_key_prefix + "/predictor") if not job_name: - job_name = sagemaker.utils.unique_name_from_base(self.resource_prefix) + job_name = unique_name_from_base(self.resource_prefix) if test_data_image_column is not None: logger.warning("Batch inference with image modality could be slow because of some technical details.") @@ -1298,99 +1283,81 @@ def _predict( backend_type=self.name, framework_version=framework_version ) entry_point = self._serve_script_path - if model_kwargs is None: - model_kwargs = {} - model_kwargs = copy.deepcopy(model_kwargs) - if transformer_kwargs is None: - transformer_kwargs = {} - transformer_kwargs = copy.deepcopy(transformer_kwargs) - self._resolve_tags(transformer_kwargs) - user_entry_point = model_kwargs.pop("entry_point", None) - repack_model = False - if predictor_path != self._fit_job.get_output_path() or user_entry_point is not None: - # Not inference on cloud trained model or not using inference on cloud trained model - # Need to repack the code into model. This will slow down batch inference and deployment - repack_model = True - if user_entry_point: - entry_point = user_entry_point - - predictor_cls = AutoGluonBatchPredictor - user_predictor_cls = model_kwargs.pop("predictor_cls", None) - if user_predictor_cls: - logger.warning( - "Providing a custom predictor_cls could break the deployment. Please refer to `AutoGluonBatchPredictor` for how to provide a custom predictor" - ) - predictor_cls = user_predictor_cls + # Models not produced by this predictor's fit job don't carry our serving code yet. + repack = predictor_path != self._fit_job.get_output_path() + model_data = self._prepare_model_data( + predictor_path, + entry_point=entry_point, + repack=repack, + repacked_model_uri=f"{output_path}/model/model.tar.gz", + ) - transform_kwargs = copy.deepcopy(transform_kwargs) - content_type = transform_kwargs.pop("content_type", None) - if "split_type" not in transform_kwargs: - split_type = "Line" - else: - split_type = transform_kwargs.pop("split_type") - if not content_type: - content_type = "text/csv" + tags = self._resolve_tags() + model_name = self._create_model( + model_name=job_name, + model_data=model_data, + image_uri=retrieve_image_uri( + framework_version, self._region, "inference", instance_type, py_version, custom_image_uri + ), + entry_point=entry_point, + environment={}, + tags=tags, + overrides=overrides, + ) - if not wait: - if download: - logger.warning( - f"`download={download}` will be ignored because `wait={wait}`. Setting `download` to `False`." - ) - download = False - if not download: - if persist: - logger.warning( - f"`persist={persist}` will be ignored because `download={download}`. Setting `persist` to `False`." - ) - persist = False - if save_path: - logger.warning( - f"`save_path={save_path}` will be ignored because `download={download}`. Setting `save_path` to `None`." - ) - save_path = None + transform_input: Dict[str, Any] = { + "DataSource": {"S3DataSource": {"S3DataType": "S3Prefix", "S3Uri": test_input}}, + "ContentType": content_type, + } + if split_type is not None: + transform_input["SplitType"] = split_type + transform_output: Dict[str, Any] = { + "S3OutputPath": (predictions_path or output_path + "/results").rstrip("/"), + "Accept": accept, + } + if assemble_with is not None: + transform_output["AssembleWith"] = assemble_with + transform_resources: Dict[str, Any] = {"InstanceType": instance_type, "InstanceCount": instance_count} + transform_ami_version = infer_sagemaker_ami_version(custom_image_uri, instance_type, image_scope="transform") + if transform_ami_version is not None: + transform_resources["TransformAmiVersion"] = transform_ami_version + request = { + "TransformJobName": job_name, + "ModelName": model_name, + "TransformInput": transform_input, + "TransformOutput": transform_output, + "TransformResources": transform_resources, + "BatchStrategy": batch_strategy, + # Maximum size in MB of a single request to the container; larger inputs are split into multiple batches. + "MaxPayloadInMB": 6, + # The maximum number of HTTP requests made to each individual transform container at one time. + "MaxConcurrentTransforms": 1, + "Tags": tags, + } + request = deep_merge(request, overrides.get("create_transform_job", {})) batch_transform_job = SageMakerBatchTransformationJob(session=self.sagemaker_session) - batch_transform_job.run( - model_data=predictor_path, - role=self.role_arn, - region=self._region, - framework_version=framework_version, - py_version=py_version, - instance_count=instance_count, - instance_type=instance_type, - entry_point=entry_point, - predictor_cls=predictor_cls, - output_path=output_path + "/results", - test_input=test_input, - job_name=job_name, - split_type=split_type, - content_type=content_type, - custom_image_uri=custom_image_uri, - wait=wait, - transformer_kwargs=transformer_kwargs, - model_kwargs=model_kwargs, - repack_model=repack_model, - **transform_kwargs, - ) - self._batch_transform_jobs[job_name] = batch_transform_job + batch_transform_job.run(transform_job_request=request, model_name=model_name, wait=wait) + self._batch_transform_jobs[batch_transform_job.job_name] = batch_transform_job pred, pred_proba = None, None - if download: - results_path = self.download_predict_results(save_path=save_path) - accept = transformer_kwargs.get("accept", "application/json") - if accept == "application/x-parquet": - results = pd.read_parquet(results_path) - elif accept == "text/csv": - results = pd.read_csv(results_path) - elif accept == "application/json": - results = pd.read_json(results_path) - else: - raise ValueError(f"Unsupported accept type for batch inference results: {accept!r}") + if wait: + bucket, key = s3_path_to_bucket_prefix(batch_transform_job.get_output_path()) + with tempfile.TemporaryDirectory(prefix="ag_batch_results_") as tmpdir: + results_path = os.path.join(tmpdir, os.path.basename(key)) + self.sagemaker_session.s3_client.download_file(bucket, key, results_path) + accept = request["TransformOutput"].get("Accept") + if accept == "application/x-parquet": + results = pd.read_parquet(results_path) + elif accept == "text/csv": + results = pd.read_csv(results_path) + elif accept == "application/json": + results = pd.read_json(results_path) + else: + raise ValueError(f"Unsupported accept type for batch inference results: {accept!r}") pred = results if split_pred_proba: pred, pred_proba = split_pred_and_pred_proba(results) - if not persist: - os.remove(results_path) return pred, pred_proba @@ -1399,10 +1366,6 @@ def __getstate__(self) -> Dict[str, Any]: d = self.__dict__.copy() d["sagemaker_session"] = None d["_region"] = None - if self.endpoint is not None: - d["_endpoint_saved"] = self.endpoint.endpoint_name - d["endpoint"] = None - return d def __setstate__(self, state): @@ -1410,9 +1373,6 @@ def __setstate__(self, state): self.__dict__.update(state) self.sagemaker_session = setup_sagemaker_session() self._region = self.sagemaker_session.boto_region_name - if hasattr(self, "_endpoint_saved") and self._endpoint_saved is not None: - self.attach_endpoint(self._endpoint_saved) - self._endpoint_saved = None self._fit_job.session = self.sagemaker_session - for job in self._batch_transform_jobs: + for job in self._batch_transform_jobs.values(): job.session = self.sagemaker_session diff --git a/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py b/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py index 66f32bdf..5dbd4359 100644 --- a/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py +++ b/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py @@ -31,8 +31,7 @@ def fit( volume_size: int = 100, custom_image_uri: Optional[str] = None, wait: bool = True, - autogluon_sagemaker_estimator_kwargs: Optional[Dict] = None, - fit_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, extra_ag_args: Optional[Dict[str, Any]] = None, extra_tags: Optional[List[Dict[str, str]]] = None, ) -> None: @@ -63,8 +62,7 @@ def fit( volume_size=volume_size, custom_image_uri=custom_image_uri, wait=wait, - autogluon_sagemaker_estimator_kwargs=autogluon_sagemaker_estimator_kwargs, - fit_kwargs=fit_kwargs, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, extra_tags=extra_tags, ) @@ -138,8 +136,6 @@ def predict( When minimizing latency isn't a concern, then the batch transform functionality may be easier, more scalable, and more appropriate. If you want to minimize latency, use `predict_real_time()` instead. To learn more: https://docs.aws.amazon.com/sagemaker/latest/dg/batch-transform.html - This method would first create a AutoGluonSagemakerInferenceModel with the trained predictor, - then create a transformer with it, and call transform in the end. Parameters ---------- @@ -176,20 +172,15 @@ def predict( payload_path = os.path.join(payload_dir, "predict_payload.json") with open(payload_path, "wb") as f: f.write(AutoGluonSerializer().serialize(wrapper)) - transform_kwargs = kwargs.pop("transform_kwargs", None) or {} - transform_kwargs["content_type"] = "application/x-autogluon" - transform_kwargs["split_type"] = "None" - # Parquet output (JSON can exceed TorchServe's 6.5MB cap); assemble_with=None preserves - # parquet footer (default "Line" appends a newline that corrupts it). - transformer_kwargs = kwargs.pop("transformer_kwargs", None) or {} - transformer_kwargs.setdefault("accept", "application/x-parquet") - transformer_kwargs.setdefault("assemble_with", None) - + # Parquet output (JSON can exceed TorchServe's 6.5MB cap); assemble_with="None" preserves + # parquet footer ("Line" appends a newline that corrupts it). pred, _ = super()._predict( test_data=payload_path, split_pred_proba=False, - transform_kwargs=transform_kwargs, - transformer_kwargs=transformer_kwargs, + content_type="application/x-autogluon", + split_type="None", + accept="application/x-parquet", + assemble_with="None", **kwargs, ) return pred diff --git a/src/autogluon/cloud/endpoint/endpoint.py b/src/autogluon/cloud/endpoint/endpoint.py deleted file mode 100644 index 025b065b..00000000 --- a/src/autogluon/cloud/endpoint/endpoint.py +++ /dev/null @@ -1,26 +0,0 @@ -from abc import ABC, abstractmethod -from typing import Union - -import pandas as pd - - -class Endpoint(ABC): - @property - @abstractmethod - def endpoint_name(self) -> str: - """Name of the endpoint""" - raise NotImplementedError - - @abstractmethod - def predict(self, test_data: Union[str, pd.DataFrame], **kwargs) -> Union[pd.DataFrame, pd.Series]: - """ - Predict with the endpoint - """ - raise NotImplementedError - - @abstractmethod - def delete_endpoint(self) -> None: - """ - Delete the endpoint and cleanup artifacts - """ - raise NotImplementedError diff --git a/src/autogluon/cloud/endpoint/prediction_future.py b/src/autogluon/cloud/endpoint/prediction_future.py index 9ae85e86..2d08b68c 100644 --- a/src/autogluon/cloud/endpoint/prediction_future.py +++ b/src/autogluon/cloud/endpoint/prediction_future.py @@ -4,8 +4,6 @@ from typing import TYPE_CHECKING, Any, Callable, Literal -from ..utils.ag_sagemaker import AutoGluonSagemakerEstimator - if TYPE_CHECKING: from ..job.sagemaker_job import SageMakerFitJob @@ -42,7 +40,7 @@ def status(self) -> PredictionStatus: def result(self) -> Any: if not self._job.completed: - AutoGluonSagemakerEstimator.attach(self._job.job_name, sagemaker_session=self._job.session).logs() + self._job.wait(logs=True) if self.status() == "Failed": raise RuntimeError( f"Prediction job {self._job.job_name!r} did not complete successfully " diff --git a/src/autogluon/cloud/endpoint/sagemaker_endpoint.py b/src/autogluon/cloud/endpoint/sagemaker_endpoint.py deleted file mode 100644 index edf9c0de..00000000 --- a/src/autogluon/cloud/endpoint/sagemaker_endpoint.py +++ /dev/null @@ -1,47 +0,0 @@ -import logging -from typing import Union - -import pandas as pd -from sagemaker.predictor import Predictor - -from .endpoint import Endpoint - -logger = logging.getLogger(__name__) - - -class SagemakerEndpoint(Endpoint): - def __init__(self, endpoint: Predictor) -> None: - self._endpoint: Predictor = endpoint - - @property - def endpoint_name(self) -> str: - """Name of the endpoint""" - if self._endpoint is not None: - return self._endpoint.endpoint_name - return None - - def predict(self, test_data: Union[str, pd.DataFrame], **kwargs) -> Union[pd.DataFrame, pd.Series]: - """ - Predict with the endpoint - """ - return self._endpoint.predict(test_data, **kwargs) - - def delete_endpoint(self) -> None: - """ - Delete the endpoint and cleanup artifacts - """ - self._delete_endpoint_model() - self._delete_endpoint() - - def _delete_endpoint_model(self): - assert self._endpoint is not None, "There is no endpoint deployed yet" - logger.log(20, "Deleting endpoint model") - self._endpoint.delete_model() - logger.log(20, "Endpoint model deleted") - - def _delete_endpoint(self, delete_endpoint_config=True): - assert self._endpoint is not None, "There is no endpoint deployed yet" - logger.log(20, "Deleteing endpoint") - self._endpoint.delete_endpoint(delete_endpoint_config=delete_endpoint_config) - logger.log(20, "Endpoint deleted") - self._endpoint = None diff --git a/src/autogluon/cloud/endpoint/tabular_endpoint.py b/src/autogluon/cloud/endpoint/tabular_endpoint.py index 719fee7b..6b049503 100644 --- a/src/autogluon/cloud/endpoint/tabular_endpoint.py +++ b/src/autogluon/cloud/endpoint/tabular_endpoint.py @@ -3,12 +3,12 @@ import boto3 import pandas as pd -from sagemaker.predictor import Predictor from autogluon.common.loaders import load_pd from ..utils.aws_utils import setup_sagemaker_session from ..utils.deserializers import PandasDeserializer +from ..utils.sagemaker_api import delete_endpoint, invoke_endpoint from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer from ..utils.utils import split_pred_and_pred_proba @@ -29,16 +29,12 @@ def __init__(self, endpoint_name: str, session: Optional[boto3.Session] = None): session ``boto3.Session`` used to invoke and delete the endpoint. If ``None``, the default ambient session is used. """ - self._predictor = Predictor( - endpoint_name=endpoint_name, - sagemaker_session=setup_sagemaker_session(boto_session=session), - serializer=AutoGluonSerializer(), - deserializer=PandasDeserializer(), - ) + self._endpoint_name = endpoint_name + self._session = setup_sagemaker_session(boto_session=session) @property def endpoint_name(self) -> str: - return self._predictor.endpoint_name + return self._endpoint_name @staticmethod def _load_data(data: DataInput) -> pd.DataFrame: @@ -68,7 +64,14 @@ def _predict( train_data=train_data, inference_kwargs={"label": label, **(inference_kwargs or {})}, ) - raw = self._predictor.predict(payload, initial_args={"Accept": "application/x-parquet"}) + raw = invoke_endpoint( + self._endpoint_name, + self._session, + payload, + serializer=AutoGluonSerializer(), + deserializer=PandasDeserializer(), + accept="application/x-parquet", + ) pred, pred_proba = split_pred_and_pred_proba(raw) if pred_proba is None: pred_proba = pred @@ -124,5 +127,4 @@ def predict_proba( def delete_endpoint(self) -> None: """Delete the endpoint and its backing model + endpoint config.""" - self._predictor.delete_model() - self._predictor.delete_endpoint(delete_endpoint_config=True) + delete_endpoint(self._endpoint_name, self._session) diff --git a/src/autogluon/cloud/endpoint/timeseries_endpoint.py b/src/autogluon/cloud/endpoint/timeseries_endpoint.py index a2451510..2a83e330 100644 --- a/src/autogluon/cloud/endpoint/timeseries_endpoint.py +++ b/src/autogluon/cloud/endpoint/timeseries_endpoint.py @@ -2,12 +2,12 @@ import boto3 import pandas as pd -from sagemaker.predictor import Predictor from autogluon.common.loaders import load_pd from ..utils.aws_utils import setup_sagemaker_session from ..utils.deserializers import PandasDeserializer +from ..utils.sagemaker_api import delete_endpoint, invoke_endpoint from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer @@ -31,16 +31,12 @@ def __init__(self, endpoint_name: str, session: Optional[boto3.Session] = None): session ``boto3.Session`` used to invoke and delete the endpoint. If ``None``, the default ambient session is used. """ - self._predictor = Predictor( - endpoint_name=endpoint_name, - sagemaker_session=setup_sagemaker_session(boto_session=session), - serializer=AutoGluonSerializer(), - deserializer=PandasDeserializer(), - ) + self._endpoint_name = endpoint_name + self._session = setup_sagemaker_session(boto_session=session) @property def endpoint_name(self) -> str: - return self._predictor.endpoint_name + return self._endpoint_name def predict( self, @@ -104,9 +100,15 @@ def predict( static_features=static_features, known_covariates=known_covariates, ) - return self._predictor.predict(payload, initial_args={"Accept": "application/x-parquet"}) + return invoke_endpoint( + self._endpoint_name, + self._session, + payload, + serializer=AutoGluonSerializer(), + deserializer=PandasDeserializer(), + accept="application/x-parquet", + ) def delete_endpoint(self) -> None: """Delete the endpoint and its backing model + endpoint config.""" - self._predictor.delete_model() - self._predictor.delete_endpoint(delete_endpoint_config=True) + delete_endpoint(self._endpoint_name, self._session) diff --git a/src/autogluon/cloud/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index cafa5bbd..97deaad3 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -1,32 +1,27 @@ import logging from abc import abstractmethod -from typing import Dict, Optional, Union +from typing import Any, Dict, Optional, Union -import sagemaker - -from ..utils.ag_sagemaker import ( - AutoGluonNonRepackInferenceModel, - AutoGluonRepackInferenceModel, - AutoGluonSagemakerEstimator, -) -from ..utils.constants import LOCAL_MODE, LOCAL_MODE_GPU, MODEL_ARTIFACT_NAME -from ..utils.dlc_utils import infer_sagemaker_ami_version +from ..utils.aws_utils import setup_sagemaker_session +from ..utils.constants import MODEL_ARTIFACT_NAME +from ..utils.job_logs import TRAINING_JOB_LOG_GROUP, TRANSFORM_JOB_LOG_GROUP, wait_for_job +from ..utils.sagemaker_api import delete_quietly from .remote_job import RemoteJob logger = logging.getLogger(__name__) class SageMakerJob(RemoteJob): + _LOG_GROUP: str + def __init__(self, session=None): - self.session = session or sagemaker.session.Session() + self.session = session or setup_sagemaker_session() self._job_name = None - self._local_mode = False - self._output_path = "" # only used in local mode self._output_filename = "" @classmethod @abstractmethod - def attach(cls, job_name): + def attach(cls, job_name, session=None): """ Reattach to a job given its name. @@ -54,6 +49,11 @@ def run(self, **kwargs): """Execute the job""" raise NotImplementedError + @abstractmethod + def _describe(self) -> Dict[str, Any]: + """Return the ``Describe*Job`` response for the job.""" + raise NotImplementedError + @abstractmethod def _get_job_status(self): raise NotImplementedError @@ -72,8 +72,6 @@ def job_name(self): @property def completed(self): - if self._local_mode: - return True # We just return True here to unblock local mode. User should know if the job is done or not easily from the log. if not self.job_name: return False return self.get_job_status() == "Completed" @@ -89,10 +87,7 @@ def get_job_status(self) -> Optional[str]: """ if not self.job_name: return "NotCreated" - if not self._local_mode: - return self._get_job_status() - logger.warning("Job status not available in local mode. Please check the local log.") - return None + return self._get_job_status() def get_output_path(self) -> Optional[str]: """ @@ -119,6 +114,32 @@ def get_hyperparameters(self) -> Dict[str, Union[int, str]]: """ return self._get_hyperparameters() + def wait(self, logs: bool = True) -> str: + """Block until the job reaches a terminal state, streaming its CloudWatch logs if ``logs`` is True. + + Does not raise if the job fails. Returns the final status (Completed | Failed | Stopped). + """ + assert self.job_name, "The job has not been started" + status = wait_for_job( + self.get_job_status, + job_name=self.job_name, + log_group=self._LOG_GROUP, + logs_client=self.session.boto_session.client("logs") if logs else None, + ) + if status != "Completed": + logger.error( + f"SageMaker job {self.job_name} finished with status {status}: {self._describe().get('FailureReason')}" + ) + return status + + def _wait_until_completed(self) -> None: + """Wait for the job with logs and raise if it does not complete successfully.""" + status = self.wait(logs=True) + if status != "Completed": + raise RuntimeError( + f"SageMaker job {self.job_name} finished with status {status}: {self._describe().get('FailureReason')}" + ) + def __getstate__(self): state_dict = self.__dict__.copy() state_dict["session"] = None @@ -129,21 +150,19 @@ def __setstate__(self, state): class SageMakerFitJob(SageMakerJob): + _LOG_GROUP = TRAINING_JOB_LOG_GROUP + def __init__(self, **kwargs): super().__init__(**kwargs) self._framework_version = None self._output_filename = MODEL_ARTIFACT_NAME @classmethod - def attach(cls, job_name): + def attach(cls, job_name, session=None): # FIXME: find a way to recover framework version - logger.warning( - "Reattach to a job does not support real-time logging. Logs will be printed once the training job completes" - ) - obj = cls() + obj = cls(session=session) obj._job_name = job_name - sagemaker_estimator = AutoGluonSagemakerEstimator.attach(job_name) - sagemaker_estimator.logs() + obj._wait_until_completed() return obj @property @@ -160,81 +179,49 @@ def info(self): ) return info + def _describe(self) -> Dict[str, Any]: + return self.session.sagemaker_client.describe_training_job(TrainingJobName=self.job_name) + def _get_job_status(self): - return self.session.describe_training_job(self.job_name)["TrainingJobStatus"] + return self._describe()["TrainingJobStatus"] def _get_output_path(self): - if not self._local_mode: - return self.session.describe_training_job(self.job_name)["ModelArtifacts"]["S3ModelArtifacts"] - assert self._output_path is not None - return self._output_path + "/" + self._output_filename + return self._describe()["ModelArtifacts"]["S3ModelArtifacts"] def _get_hyperparameters(self): if self.job_name: - return self.session.describe_training_job(self.job_name)["HyperParameters"] + return self._describe().get("HyperParameters") return None def run( self, - role, - entry_point, - region, - instance_type, - instance_count, - volume_size, - framework_version, - py_version, - base_job_name, - output_path, - code_location, - inputs, - custom_image_uri, - wait, - job_name, - autogluon_sagemaker_estimator_kwargs, - **kwargs, + training_job_request: Dict[str, Any], + framework_version: Optional[str], + wait: bool, ): - self._local_mode = instance_type in (LOCAL_MODE, LOCAL_MODE_GPU) - sagemaker_estimator = AutoGluonSagemakerEstimator( - role=role, - entry_point=entry_point, - region=region, - instance_type=instance_type, - instance_count=instance_count, - volume_size=volume_size, - framework_version=framework_version, - py_version=py_version, - base_job_name=base_job_name, - output_path=output_path, - code_location=code_location, - image_uri=custom_image_uri, - **autogluon_sagemaker_estimator_kwargs, - ) + """Create the training job from a ``CreateTrainingJob`` request and optionally wait for it to finish.""" + job_name = training_job_request["TrainingJobName"] logger.log(20, f"Start sagemaker training job `{job_name}`") try: - sagemaker_estimator.fit(inputs=inputs, wait=wait, job_name=job_name, **kwargs) + self.session.sagemaker_client.create_training_job(**training_job_request) self._job_name = job_name self._framework_version = framework_version - - assert sagemaker_estimator.output_path is not None - latest_training_job = sagemaker_estimator.latest_training_job - assert latest_training_job is not None - latest_training_job_name = latest_training_job.name - assert latest_training_job_name is not None - - self._output_path = sagemaker_estimator.output_path + "/" + latest_training_job_name + if wait: + self._wait_until_completed() except Exception as e: logger.error(f"Training failed. Please check sagemaker console training jobs {job_name} for details.") raise e class SageMakerBatchTransformationJob(SageMakerJob): + _LOG_GROUP = TRANSFORM_JOB_LOG_GROUP + def __init__(self, **kwargs): super().__init__(**kwargs) self._output_filename = "" @classmethod - def attach(cls, job_name): + def attach(cls, job_name, session=None): raise NotImplementedError def info(self): @@ -246,105 +233,46 @@ def info(self): ) return info + def _describe(self) -> Dict[str, Any]: + return self.session.sagemaker_client.describe_transform_job(TransformJobName=self.job_name) + def _get_job_status(self): - return self.session.describe_transform_job(self.job_name)["TransformJobStatus"] + return self._describe()["TransformJobStatus"] def _get_output_path(self): - if not self._local_mode: - return ( - self.session.describe_transform_job(self.job_name)["TransformOutput"]["S3OutputPath"] - + "/" - + self._output_filename - ) - assert self._output_path is not None - return self._output_path + "/" + self._output_filename + return self._describe()["TransformOutput"]["S3OutputPath"] + "/" + self._output_filename + + def _delete_model(self, model_name: str) -> None: + self.session.sagemaker_client.delete_model(ModelName=model_name) def run( self, - model_data, - role, - region, - framework_version, - py_version, - instance_count, - instance_type, - entry_point, - predictor_cls, - output_path, - test_input, - job_name, - split_type, - content_type, - custom_image_uri, - wait, - model_kwargs, - transformer_kwargs, - repack_model=False, - **kwargs, + transform_job_request: Dict[str, Any], + model_name: str, + wait: bool, ): - self._local_mode = instance_type in (LOCAL_MODE, LOCAL_MODE_GPU) - if repack_model: - model_cls = AutoGluonRepackInferenceModel - else: - model_cls = AutoGluonNonRepackInferenceModel - logger.log(20, "Creating inference model...") - model = model_cls( - model_data=model_data, - role=role, - region=region, - framework_version=framework_version, - py_version=py_version, - instance_type=instance_type, - custom_image_uri=custom_image_uri, - entry_point=entry_point, - predictor_cls=predictor_cls, - **model_kwargs, - ) - logger.log(20, "Inference model created successfully") - logger.log(20, "Creating transformer...") - transform_ami_version = infer_sagemaker_ami_version( - custom_image_uri, - instance_type, - image_scope="transform", - ) - if transform_ami_version is not None: - transformer_kwargs.setdefault("transform_ami_version", transform_ami_version) - transformer = model.transformer( - instance_count=instance_count, - instance_type=instance_type, - output_path=output_path, - **transformer_kwargs, - ) - logger.log(20, "Transformer created successfully") + """Create the transform job from a ``CreateTransformJob`` request. + ``model_name`` (the model created for this job) is deleted once the job finishes (``wait=True``) or fails to + start. With ``wait=False`` the model is kept, since the job still needs it. + """ + job_name = transform_job_request["TransformJobName"] try: logger.log(20, "Transforming") - transformer.transform( - test_input, - job_name=job_name, - split_type=split_type, - content_type=content_type, - wait=wait, - **kwargs, - ) + self.session.sagemaker_client.create_transform_job(**transform_job_request) self._job_name = job_name - - assert transformer.output_path is not None - latest_transform_job = transformer.latest_transform_job - assert latest_transform_job is not None - latest_transform_job_name = latest_transform_job.name - assert latest_transform_job_name is not None - - self._output_path = transformer.output_path + "/" + latest_transform_job_name + if wait: + self._wait_until_completed() logger.log(20, "Transform done") except Exception as e: - transformer.delete_model() + delete_quietly(self.session.sagemaker_client.delete_model, ModelName=model_name) raise e - self._output_filename = test_input.split("/")[-1] + ".out" + input_uri = transform_job_request["TransformInput"]["DataSource"]["S3DataSource"]["S3Uri"] + self._output_filename = input_uri.split("/")[-1] + ".out" if wait: - transformer.delete_model() + self._delete_model(model_name) logger.log(20, f"Predict results have been saved to {self.get_output_path()}") else: logger.log( diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 89da81bc..f2ac48bb 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -23,6 +23,7 @@ from ..endpoint.timeseries_endpoint import TimeSeriesEndpoint from ..scripts.script_manager import ScriptManager from ..utils.aws_utils import resolve_cloud_output_path +from ..utils.sagemaker_api import reject_legacy_kwargs from ..utils.utils import split_pred_and_pred_proba from ..version import __version__ from .registry import get_model_config @@ -104,7 +105,7 @@ def __init__( role ARN of the SageMaker execution role used to run training and inference jobs. If ``None``, falls back to ``role_arn`` in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` / - :func:`autogluon.cloud.register`), and finally to ``sagemaker.get_execution_role()``. + :func:`autogluon.cloud.register`), and finally to the role of the current AWS identity. hyperparameters Default hyperparameters applied to inference and (when supported) training. model_artifact_uri @@ -218,10 +219,7 @@ def _deploy_backend( "problem_type": self._config.problem_type, } - model_kwargs = backend_kwargs.pop("model_kwargs", {}) - model_kwargs["entry_point"] = self._serve_script_path - - # FM deploys never want SDK repack: predictor_path is either None (script-only tarball is built locally) or a + # FM deploys never repack: predictor_path is either None (script-only tarball is built locally) or a # pre-bundled cache artifact that already contains the serve script. self._backend.deploy( predictor_path=self.model_artifact_uri, @@ -230,7 +228,7 @@ def _deploy_backend( instance_type=instance_type, custom_image_uri=custom_image_uri, wait=wait, - model_kwargs=model_kwargs, + entry_point=self._serve_script_path, fm_serve_config=fm_serve_config, inference_mode=inference_mode, inference_config=inference_config, @@ -238,7 +236,7 @@ def _deploy_backend( extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], **backend_kwargs, ) - assert self._backend.endpoint is not None + assert self._backend.endpoint_name is not None def fit( self, @@ -410,6 +408,7 @@ class TimeSeriesFoundationModel(FoundationModel): def _serve_script_path(self) -> str: return ScriptManager.SAGEMAKER_TIMESERIES_FM_SERVE_SCRIPT_PATH + @reject_legacy_kwargs def deploy( self, instance_type: Optional[str] = None, @@ -444,11 +443,10 @@ def deploy( Endpoint type. ``"serverless"`` provisions a SageMaker Serverless Inference endpoint (no instance management, scales to zero). inference_config - Mode-specific overrides forwarded to ``sagemaker.serverless.ServerlessInferenceConfig`` - (e.g. ``memory_size_in_mb``, ``max_concurrency``). + Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). **backend_kwargs - Backend-specific arguments (e.g., initial_instance_count, volume_size, - model_kwargs, deploy_kwargs). + Backend-specific arguments (e.g., ``initial_instance_count``, ``volume_size``, ``backend_overrides``; see + :meth:`autogluon.cloud.TabularCloudPredictor.deploy`). """ self._deploy_backend( instance_type=instance_type, @@ -462,7 +460,7 @@ def deploy( **backend_kwargs, ) return TimeSeriesEndpoint( - endpoint_name=self._backend.endpoint.endpoint_name, + endpoint_name=self._backend.endpoint_name, session=self._backend.sagemaker_session.boto_session, ) @@ -489,6 +487,7 @@ def _build_predictor_init_args( args["quantile_levels"] = quantile_levels return args + @reject_legacy_kwargs def predict( self, data: Union[str, Path, pd.DataFrame], @@ -552,8 +551,8 @@ def predict( :class:`JobPredictionFuture` immediately — call ``.result()`` on it later to retrieve the DataFrame, or ``.status()`` to check progress. **backend_kwargs - Additional backend-specific arguments (e.g., job_name, volume_size, - autogluon_sagemaker_estimator_kwargs). + Additional backend-specific arguments (e.g., ``job_name``, ``volume_size``, ``backend_overrides``; this + prediction runs as a SageMaker training job, see :meth:`autogluon.cloud.TabularCloudPredictor.fit`). Returns ------- @@ -623,6 +622,7 @@ class TabularFoundationModel(FoundationModel): def _serve_script_path(self) -> str: return ScriptManager.SAGEMAKER_TABULAR_FM_SERVE_SCRIPT_PATH + @reject_legacy_kwargs def deploy( self, instance_type: Optional[str] = None, @@ -664,7 +664,7 @@ def deploy( **backend_kwargs, ) return TabularEndpoint( - endpoint_name=self._backend.endpoint.endpoint_name, + endpoint_name=self._backend.endpoint_name, session=self._backend.sagemaker_session.boto_session, ) @@ -694,6 +694,7 @@ def _load_results( else: return pred_proba + @reject_legacy_kwargs def predict( self, test_data: Union[str, Path, pd.DataFrame], @@ -738,7 +739,8 @@ def predict( If True, block and return the predictions. If False, return a :class:`JobPredictionFuture` immediately — call ``.result()`` on it later to retrieve the predictions. **backend_kwargs - Additional backend-specific arguments (e.g., job_name, volume_size). + Additional backend-specific arguments (e.g., ``job_name``, ``volume_size``, ``backend_overrides``; this + prediction runs as a SageMaker training job, see :meth:`autogluon.cloud.TabularCloudPredictor.fit`). Returns ------- @@ -766,6 +768,7 @@ def predict( pred, _ = result return pred + @reject_legacy_kwargs def predict_proba( self, test_data: Union[str, Path, pd.DataFrame], @@ -812,7 +815,8 @@ def predict_proba( wait If True, block and return the result. If False, return a :class:`JobPredictionFuture` immediately. **backend_kwargs - Additional backend-specific arguments (e.g., job_name, volume_size). + Additional backend-specific arguments (e.g., ``job_name``, ``volume_size``, ``backend_overrides``; this + prediction runs as a SageMaker training job, see :meth:`autogluon.cloud.TabularCloudPredictor.fit`). Returns ------- diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index ab466771..eaaee208 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -21,8 +21,8 @@ from ..backend.backend import Backend from ..backend.backend_factory import BackendFactory from ..backend.constant import SAGEMAKER -from ..endpoint.endpoint import Endpoint from ..utils.aws_utils import resolve_cloud_output_path +from ..utils.sagemaker_api import reject_legacy_kwargs from ..utils.utils import safe_unpack_archive logger = logging.getLogger(__name__) @@ -66,7 +66,7 @@ def __init__( role: Optional[str], default = None ARN of the SageMaker execution role used to run training and inference jobs. If ``None``, falls back to ``role_arn`` in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` / - :func:`autogluon.cloud.register`), and finally to ``sagemaker.get_execution_role()``. + :func:`autogluon.cloud.register`), and finally to the role of the current AWS identity. verbosity : int, default = 2 Verbosity levels range from 0 to 4 and control how much information is printed. Higher levels correspond to more detailed print statements (you can set verbosity = 0 to suppress warnings). @@ -110,9 +110,7 @@ def endpoint_name(self) -> Optional[str]: """ Return the CloudPredictor deployed endpoint name """ - if self.backend.endpoint: - return self.backend.endpoint.endpoint_name - return None + return self.backend.endpoint_name def info(self) -> Dict[str, Any]: """ @@ -161,6 +159,7 @@ def _setup_local_output_path(self, path): ) return os.path.abspath(path) + @reject_legacy_kwargs def fit( self, train_data: Optional[Union[str, Path, pd.DataFrame]] = None, @@ -178,7 +177,7 @@ def fit( custom_image_uri: Optional[str] = None, timeout: int = 24 * 60 * 60, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> CloudPredictor: """ @@ -219,15 +218,12 @@ def fit( Whether the call should wait until the job completes To be noticed, the function won't return immediately because there are some preparations needed prior fit. Use `get_fit_job_status` to get job status. - backend_kwargs: dict, default = None - Any extra arguments needed to pass to the underneath backend. - For SageMaker backend, valid keys are: - 1. autogluon_sagemaker_estimator_kwargs - Any extra arguments needed to initialize AutoGluonSagemakerEstimator - Please refer to https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#sagemaker.estimator.Estimator for all options - 2. fit_kwargs - Any extra arguments needed to pass to fit. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#sagemaker.estimator.Estimator.fit for all options + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument. Maps ``"create_training_job"`` to raw + `CreateTrainingJob `_ + request fields in the PascalCase format of the SageMaker API and boto3, which are deep-merged over the + request built by AutoGluon-Cloud, e.g. ``{"create_training_job": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}``. + Nested dicts merge recursively; other values, including lists, replace the generated ones. Returns ------- `CloudPredictor` object. Returns self. @@ -235,11 +231,10 @@ def fit( assert not self.backend.is_fit, ( "Predictor is already fit! To fit additional models, create a new `CloudPredictor`" ) - if backend_kwargs is None: - backend_kwargs = {} - # `test_data` is an internal channel for the fit_predict path (see TabularCloudPredictor.fit_predict_proba); - # it is intentionally not part of the public `fit()` signature. + # `test_data` / `extra_ag_args` are internal channels for the fit_predict path (see + # TabularCloudPredictor.fit_predict_proba); they are intentionally not part of the public `fit()` signature. test_data = kwargs.pop("test_data", None) + extra_ag_args = kwargs.pop("extra_ag_args", None) if kwargs: raise TypeError(f"fit() got unexpected keyword arguments: {sorted(kwargs)}") predictor_fit_args = {} if predictor_fit_args is None else dict(predictor_fit_args) @@ -258,7 +253,6 @@ def fit( "AutoGluon-Tabular require autogluon.multimodal, which is being deprecated. " "Use `MultiModalCloudPredictor` for image data." ) - backend_kwargs = self.backend.parse_backend_fit_kwargs(backend_kwargs) self.backend.fit( predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, @@ -273,7 +267,8 @@ def fit( custom_image_uri=custom_image_uri, timeout=timeout, wait=wait, - **backend_kwargs, + backend_overrides=backend_overrides, + extra_ag_args=extra_ag_args, ) return self @@ -372,6 +367,7 @@ def to_local_predictor(self, predictor_path: Optional[str] = None, save_path: Op local_model_path = self.download_trained_predictor(predictor_path=predictor_path, save_path=save_path) return predictor_cls.load(local_model_path, **kwargs) + @reject_legacy_kwargs def deploy( self, predictor_path: Optional[str] = None, @@ -384,7 +380,7 @@ def deploy( wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - backend_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> None: """ Deploy a predictor to an inference endpoint. @@ -423,25 +419,20 @@ def deploy( Endpoint type. ``"serverless"`` provisions a SageMaker Serverless Inference endpoint (no instance management, scales to zero). inference_config: Optional[Dict[str, Any]], default = None - Mode-specific overrides forwarded to ``sagemaker.serverless.ServerlessInferenceConfig`` - (e.g. ``memory_size_in_mb``, ``max_concurrency``). - backend_kwargs: dict, default = None - Any extra arguments needed to pass to the underneath backend. - For SageMaker backend, valid keys are: - 1. model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - 2. deploy_kwargs - Any extra arguments needed to pass to deploy. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#sagemaker.model.Model.deploy for all options + Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument: raw request fields in the PascalCase + format of the SageMaker API and boto3, deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + ``"create_model"``, ``"production_variant"`` (the endpoint config's single production variant), + ``"create_endpoint_config"`` and ``"create_endpoint"``, e.g. + ``{"production_variant": {"ModelDataDownloadTimeoutInSeconds": 1200}}``. + Nested dicts merge recursively; other values, including lists, replace the generated ones. Only + resources created by AutoGluon-Cloud are cleaned up. """ if inference_mode == "serverless" and instance_type is not None: raise ValueError("`instance_type` must not be set when `inference_mode='serverless'`.") if instance_type is None and inference_mode == "realtime": instance_type = "ml.m5.2xlarge" - if backend_kwargs is None: - backend_kwargs = {} - backend_kwargs = self.backend.parse_backend_deploy_kwargs(backend_kwargs) self.backend.deploy( predictor_path=predictor_path, endpoint_name=endpoint_name, @@ -453,27 +444,28 @@ def deploy( wait=wait, inference_mode=inference_mode, inference_config=inference_config, - **backend_kwargs, + backend_overrides=backend_overrides, ) - def attach_endpoint(self, endpoint: Union[str, Endpoint]) -> None: + def attach_endpoint(self, endpoint: str) -> None: """ Attach the current CloudPredictor to an existing endpoint. Parameters ---------- - endpoint: str or :class:`Endpoint` - If str is passed, it should be the name of the endpoint being attached to. + endpoint: str + Name of the endpoint being attached to. """ self.backend.attach_endpoint(endpoint) - def detach_endpoint(self) -> Endpoint: + def detach_endpoint(self) -> str: """ - Detach the current endpoint and return it. + Detach the current endpoint and return its name. Returns ------- - `Endpoint` object. + str + Name of the detached endpoint. Pass it to :meth:`attach_endpoint` to attach it again. """ return self.backend.detach_endpoint() @@ -550,6 +542,7 @@ def predict_proba_real_time( test_data=test_data, test_data_image_column=test_data_image_column, accept=accept ) + @reject_legacy_kwargs def predict( self, test_data: Union[str, pd.DataFrame], @@ -561,7 +554,8 @@ def predict( instance_count: int = 1, custom_image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + predictions_path: Optional[str] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Batch inference. @@ -594,40 +588,23 @@ def predict( wait: bool, default = True Whether to wait for batch transform to complete. To be noticed, the function won't return immediately because there are some preparations needed prior transform. - backend_kwargs: dict, default = None - Any extra arguments needed to pass to the underneath backend. - For SageMaker backend, valid keys are: - 1. download: bool, default = True - Whether to download the batch transform results to the disk and load it after the batch transform finishes. - Will be ignored if `wait` is `False`. - 2. persist: bool, default = True - Whether to persist the downloaded batch transform results on the disk. - Will be ignored if `download` is `False` - 3. save_path: str, default = None, - Path to save the downloaded result. - Will be ignored if `download` is `False`. - If None, CloudPredictor will create one. - If `persist` is `False`, file would first be downloaded to this path and then removed. - 4. model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - 5. transformer_kwargs: dict - Any extra arguments needed to pass to transformer. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer for all options. - 6. transform_kwargs: - Any extra arguments needed to pass to transform. - Please refer to - https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer.transform for all options. + predictions_path: Optional[str], default = None + S3 prefix under which the batch transform job writes its results (``/.out``). + Defaults to ``{cloud_output_path}/batch_transform//results``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument: raw request fields in the PascalCase + format of the SageMaker API and boto3, deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + ``"create_model"`` and ``"create_transform_job"``, e.g. + ``{"create_transform_job": {"BatchStrategy": "SingleRecord", "MaxPayloadInMB": 20}}``. + Nested dicts merge recursively; other values, including lists, replace the generated ones. Only + resources created by AutoGluon-Cloud are cleaned up. Returns ------- Optional Pandas.Series - Predict results in Series if `download` is True - None if `download` is False + Predict results in Series if `wait` is True + None if `wait` is False """ - if backend_kwargs is None: - backend_kwargs = {} - backend_kwargs = self.backend.parse_backend_predict_kwargs(backend_kwargs) return self.backend.predict( test_data=test_data, test_data_image_column=test_data_image_column, @@ -638,9 +615,11 @@ def predict( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - **backend_kwargs, + predictions_path=predictions_path, + backend_overrides=backend_overrides, ) + @reject_legacy_kwargs def predict_proba( self, test_data: Union[str, pd.DataFrame], @@ -653,7 +632,8 @@ def predict_proba( instance_count: int = 1, custom_image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + predictions_path: Optional[str] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]]: """ Batch inference @@ -689,42 +669,25 @@ def predict_proba( wait: bool, default = True Whether to wait for batch transform to complete. To be noticed, the function won't return immediately because there are some preparations needed prior transform. - backend_kwargs: dict, default = None - Any extra arguments needed to pass to the underneath backend. - For SageMaker backend, valid keys are: - 1. download: bool, default = True - Whether to download the batch transform results to the disk and load it after the batch transform finishes. - Will be ignored if `wait` is `False`. - 2. persist: bool, default = True - Whether to persist the downloaded batch transform results on the disk. - Will be ignored if `download` is `False` - 3. save_path: str, default = None, - Path to save the downloaded result. - Will be ignored if `download` is `False`. - If None, CloudPredictor will create one. - If `persist` is `False`, file would first be downloaded to this path and then removed. - 4. model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - 5. transformer_kwargs: dict - Any extra arguments needed to pass to transformer. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer for all options. - 6. transform_kwargs: - Any extra arguments needed to pass to transform. - Please refer to - https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer.transform for all options. + predictions_path: Optional[str], default = None + S3 prefix under which the batch transform job writes its results (``/.out``). + Defaults to ``{cloud_output_path}/batch_transform//results``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument: raw request fields in the PascalCase + format of the SageMaker API and boto3, deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + ``"create_model"`` and ``"create_transform_job"``, e.g. + ``{"create_transform_job": {"BatchStrategy": "SingleRecord", "MaxPayloadInMB": 20}}``. + Nested dicts merge recursively; other values, including lists, replace the generated ones. Only + resources created by AutoGluon-Cloud are cleaned up. Returns ------- Optional[Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]] - If `download` is False, will return None or (None, None) if `include_predict` is True - If `download` is True and `include_predict` is True, + If `wait` is False, will return None or (None, None) if `include_predict` is True + If `wait` is True and `include_predict` is True, will return (prediction, predict_probability), where prediction is a Pandas.Series and predict_probability is a Pandas.DataFrame or a Pandas.Series that's identical to prediction when it's a regression problem. """ - if backend_kwargs is None: - backend_kwargs = {} - backend_kwargs = self.backend.parse_backend_predict_kwargs(backend_kwargs) return self.backend.predict_proba( test_data=test_data, test_data_image_column=test_data_image_column, @@ -736,7 +699,8 @@ def predict_proba( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - **backend_kwargs, + predictions_path=predictions_path, + backend_overrides=backend_overrides, ) def get_batch_inference_job_info(self, job_name: Optional[str] = None) -> Dict[str, Any]: diff --git a/src/autogluon/cloud/predictor/tabular_cloud_predictor.py b/src/autogluon/cloud/predictor/tabular_cloud_predictor.py index 511cd0db..c9a613ed 100644 --- a/src/autogluon/cloud/predictor/tabular_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/tabular_cloud_predictor.py @@ -7,6 +7,7 @@ import pandas as pd from ..backend.constant import SAGEMAKER, TABULAR_SAGEMAKER +from ..utils.sagemaker_api import reject_legacy_kwargs from ..utils.utils import split_pred_and_pred_proba from .cloud_predictor import CloudPredictor @@ -36,6 +37,7 @@ def _get_local_predictor_cls(self): predictor_cls = TabularPredictor return predictor_cls + @reject_legacy_kwargs def fit_predict( self, train_data: Union[str, Path, pd.DataFrame], @@ -52,7 +54,7 @@ def fit_predict( custom_image_uri: Optional[str] = None, wait: bool = True, predictions_path: Optional[str] = None, - backend_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Fit and predict in a single SageMaker training job. @@ -95,8 +97,8 @@ def fit_predict( S3 URL where predictions will be written by the training container (e.g. ``s3://my-bucket/runs/2024-05-01/predictions.csv``). Defaults to ``{cloud_output_path}/{job_name}/predictions.csv``. - backend_kwargs: Optional[dict], default = None - Backend-specific arguments. Same keys as ``fit()``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields, same as in :meth:`fit`. Returns ------- @@ -119,13 +121,14 @@ def fit_predict( custom_image_uri=custom_image_uri, wait=wait, predictions_path=predictions_path, - backend_kwargs=backend_kwargs, + backend_overrides=backend_overrides, ) if result is None: # wait=False return None pred, _ = result return pred + @reject_legacy_kwargs def fit_predict_proba( self, train_data: Union[str, Path, pd.DataFrame], @@ -143,7 +146,7 @@ def fit_predict_proba( custom_image_uri: Optional[str] = None, wait: bool = True, predictions_path: Optional[str] = None, - backend_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]]: """ Fit and predict probabilities in a single SageMaker training job. @@ -183,8 +186,8 @@ def fit_predict_proba( predictions_path: Optional[str] S3 URL where predictions will be written by the training container. Defaults to ``{cloud_output_path}/{job_name}/predictions.csv``. - backend_kwargs: Optional[dict], default = None - Backend-specific arguments. Same keys as ``fit()``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields, same as in :meth:`fit`. Returns ------- @@ -193,12 +196,9 @@ def fit_predict_proba( ``predict_probability``. Returns ``None`` when ``wait`` is False; fetch later via ``get_fit_predict_proba_results()``. """ - backend_kwargs = {} if backend_kwargs is None else dict(backend_kwargs) - extra_ag_args = dict(backend_kwargs.get("extra_ag_args") or {}) - extra_ag_args["predict_after_fit"] = True + extra_ag_args = {"predict_after_fit": True} if predictions_path is not None: extra_ag_args["predictions_path"] = predictions_path - backend_kwargs["extra_ag_args"] = extra_ag_args self.fit( train_data=train_data, @@ -213,7 +213,8 @@ def fit_predict_proba( volume_size=volume_size, custom_image_uri=custom_image_uri, wait=wait, - backend_kwargs=backend_kwargs, + backend_overrides=backend_overrides, + extra_ag_args=extra_ag_args, ) if not wait: diff --git a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py index fa470731..012f0a50 100644 --- a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py @@ -7,6 +7,7 @@ import pandas as pd from ..backend.constant import SAGEMAKER, TIMESERIES_SAGEMAKER +from ..utils.sagemaker_api import reject_legacy_kwargs from .cloud_predictor import CloudPredictor logger = logging.getLogger(__name__) @@ -34,6 +35,7 @@ def _get_local_predictor_cls(self): return TimeSeriesPredictor + @reject_legacy_kwargs def fit( self, train_data: Optional[Union[str, Path, pd.DataFrame]] = None, @@ -51,8 +53,9 @@ def fit( volume_size: int = 100, custom_image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, known_covariates: Optional[Union[str, Path, pd.DataFrame]] = None, + **kwargs, ) -> TimeSeriesCloudPredictor: """ Fit the predictor in a SageMaker training job. @@ -105,15 +108,8 @@ def fit( Whether the call should wait until the job completes To be noticed, the function won't return immediately because there are some preparations needed prior fit. Use `get_fit_job_status` to get job status. - backend_kwargs: dict, default = None - Any extra arguments needed to pass to the underneath backend. - For SageMaker backend, valid keys are: - 1. autogluon_sagemaker_estimator_kwargs - Any extra arguments needed to initialize AutoGluonSagemakerEstimator - Please refer to https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#sagemaker.estimator.Estimator for all options - 2. fit_kwargs - Any extra arguments needed to pass to fit. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#sagemaker.estimator.Estimator.fit for all options + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields under ``"create_training_job"``. See :meth:`TabularCloudPredictor.fit`. Returns ------- @@ -122,9 +118,10 @@ def fit( assert not self.backend.is_fit, ( "Predictor is already fit! To fit additional models, create a new `CloudPredictor`" ) - if backend_kwargs is None: - backend_kwargs = {} - + # `extra_ag_args` is an internal channel for `fit_predict`; it is intentionally not part of the public signature. + extra_ag_args = kwargs.pop("extra_ag_args", None) + if kwargs: + raise TypeError(f"fit() got unexpected keyword arguments: {sorted(kwargs)}") predictor_fit_args = {} if predictor_fit_args is None else dict(predictor_fit_args) data_channels = { "train_data": train_data, @@ -141,7 +138,6 @@ def fit( if data_channels["train_data"] is None: raise TypeError("fit() missing required argument: 'train_data'") - backend_kwargs = self.backend.parse_backend_fit_kwargs(backend_kwargs) self.backend.fit( predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, @@ -155,7 +151,8 @@ def fit( volume_size=volume_size, custom_image_uri=custom_image_uri, wait=wait, - **backend_kwargs, + backend_overrides=backend_overrides, + extra_ag_args=extra_ag_args, ) return self @@ -210,6 +207,7 @@ def predict_proba_real_time(self, **kwargs) -> pd.DataFrame: """ raise ValueError(f"{self.__class__.__name__} does not support predict_proba operation.") + @reject_legacy_kwargs def predict( self, data: Union[str, pd.DataFrame], @@ -222,7 +220,8 @@ def predict( instance_count: int = 1, custom_image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + predictions_path: Optional[str] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ Predict using SageMaker batch transform. @@ -261,34 +260,9 @@ def predict( wait: bool, default = True Whether to wait for batch transform to complete. To be noticed, the function won't return immediately because there are some preparations needed prior transform. - backend_kwargs: dict, default = None - Any extra arguments needed to pass to the underneath backend. - For SageMaker backend, valid keys are: - 1. download: bool, default = True - Whether to download the batch transform results to the disk and load it after the batch transform finishes. - Will be ignored if `wait` is `False`. - 2. persist: bool, default = True - Whether to persist the downloaded batch transform results on the disk. - Will be ignored if `download` is `False` - 3. save_path: str, default = None, - Path to save the downloaded result. - Will be ignored if `download` is `False`. - If None, CloudPredictor will create one. - If `persist` is `False`, file would first be downloaded to this path and then removed. - 4. model_kwargs: dict, default = dict() - Any extra arguments needed to initialize Sagemaker Model - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model for all options - 5. transformer_kwargs: dict - Any extra arguments needed to pass to transformer. - Please refer to https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer for all options. - 6. transform_kwargs: - Any extra arguments needed to pass to transform. - Please refer to - https://sagemaker.readthedocs.io/en/v2/api/inference/transformer.html#sagemaker.transformer.Transformer.transform for all options. + predictions_path, backend_overrides: + Same as in :meth:`TabularCloudPredictor.predict`. """ - if backend_kwargs is None: - backend_kwargs = {} - backend_kwargs = self.backend.parse_backend_predict_kwargs(backend_kwargs) return self.backend.predict( test_data=data, static_features=static_features, @@ -300,7 +274,8 @@ def predict( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - **backend_kwargs, + predictions_path=predictions_path, + backend_overrides=backend_overrides, ) def predict_proba( @@ -312,6 +287,7 @@ def predict_proba( """ raise ValueError(f"{self.__class__.__name__} does not support predict_proba operation.") + @reject_legacy_kwargs def fit_predict( self, train_data: Union[str, Path, pd.DataFrame], @@ -330,7 +306,7 @@ def fit_predict( volume_size: int = 100, custom_image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ Fit and predict in a single SageMaker training job. @@ -383,22 +359,17 @@ def fit_predict( Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True Whether the call should wait until the job completes. - backend_kwargs: Optional[dict], default = None - Backend-specific arguments. Same keys as ``fit()``. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields, same as in :meth:`fit`. Returns ------- Optional[pd.DataFrame] Predictions as a DataFrame. Returns ``None`` when ``wait`` is False. """ - if backend_kwargs is None: - backend_kwargs = {} - else: - backend_kwargs = dict(backend_kwargs) extra_ag_args = {"predict_after_fit": True} if predictions_path is not None: extra_ag_args["predictions_path"] = predictions_path - backend_kwargs["extra_ag_args"] = extra_ag_args self.fit( train_data=train_data, @@ -415,7 +386,8 @@ def fit_predict( volume_size=volume_size, custom_image_uri=custom_image_uri, wait=wait, - backend_kwargs=backend_kwargs, + backend_overrides=backend_overrides, + extra_ag_args=extra_ag_args, ) if not wait: diff --git a/src/autogluon/cloud/utils/ag_sagemaker.py b/src/autogluon/cloud/utils/ag_sagemaker.py index 8a247b76..619a8384 100644 --- a/src/autogluon/cloud/utils/ag_sagemaker.py +++ b/src/autogluon/cloud/utils/ag_sagemaker.py @@ -1,271 +1,105 @@ -import copy -import os - -import sagemaker -from sagemaker import fw_utils, vpc_utils -from sagemaker.estimator import Estimator -from sagemaker.model import DIR_PARAM_NAME, SCRIPT_PARAM_NAME, Model -from sagemaker.predictor import Predictor -from sagemaker.serializers import CSVSerializer - -from .deserializers import PandasDeserializer -from .dlc_utils import retrieve_image_uri, retrieve_latest_framework_version -from .serializers import AutoGluonSerializer, MultiModalSerializer - - -# SageMaker SDK v2 does not expose TransformAmiVersion through its public Transformer API. -# Remove this proxy when AG Cloud migrates Batch Transform to the SDK v3 resource API. -class _TransformAmiVersionSession: - """Delegate to a SageMaker session while adding a Batch Transform AMI.""" - - def __init__(self, session, transform_ami_version): - self._session = session - self._transform_ami_version = transform_ami_version - - def __getattr__(self, name): - return getattr(self._session, name) - - def transform(self, **kwargs): - kwargs["resource_config"] = copy.deepcopy(kwargs["resource_config"]) - kwargs["resource_config"]["TransformAmiVersion"] = self._transform_ami_version - return self._session.transform(**kwargs) - - -# Estimator documentation: https://sagemaker.readthedocs.io/en/v2/api/training/estimators.html#estimators -class AutoGluonSagemakerEstimator(Estimator): - def __init__( - self, - region, - framework_version, - py_version, - instance_type, - entry_point=None, - source_dir=None, - hyperparameters=None, - image_uri=None, - **kwargs, - ): - self.framework_version = framework_version - self.py_version = py_version - self.image_uri = image_uri - if self.image_uri is None: - self.image_uri = retrieve_image_uri( - framework_version=framework_version, - region=region, - image_scope="training", - instance_type=instance_type, - py_version=py_version, - ) - super().__init__( - entry_point=entry_point, - source_dir=source_dir, - hyperparameters=hyperparameters, - instance_type=instance_type, - image_uri=self.image_uri, - **kwargs, - ) - - def _configure_distribution(self, distributions): - return - - def create_model( - self, - region, - framework_version, - py_version, - instance_type, - source_dir=None, - entry_point=None, - role=None, - image_uri=None, - predictor_cls=None, - vpc_config_override=vpc_utils.VPC_CONFIG_DEFAULT, - repack=False, - **kwargs, - ): - image_uri = retrieve_image_uri( - framework_version=framework_version, - region=region, - image_scope="inference", - instance_type=instance_type, - py_version=py_version, - ) - if predictor_cls is None: - - def predict_wrapper(endpoint, session): - return Predictor(endpoint, session) +"""Packaging helpers for AutoGluon training and serving code on SageMaker. - predictor_cls = predict_wrapper +These replace the SageMaker SDK v2 ``Estimator`` / ``Model`` script-mode machinery: the AutoGluon DLCs still run +the SageMaker training and inference toolkits, which locate user code through the ``sagemaker_program`` / +``sagemaker_submit_directory`` hyperparameters (training) and the ``SAGEMAKER_PROGRAM`` / +``SAGEMAKER_SUBMIT_DIRECTORY`` environment variables (inference). +""" - role = role or self.role - - if "enable_network_isolation" not in kwargs: - kwargs["enable_network_isolation"] = self.enable_network_isolation() - - if repack: - model_cls = AutoGluonRepackInferenceModel - else: - model_cls = AutoGluonNonRepackInferenceModel - return model_cls( - image_uri=image_uri, - source_dir=source_dir, - entry_point=entry_point, - model_data=self.model_data, - role=role, - vpc_config=self.get_vpc_config(vpc_config_override), - sagemaker_session=self.sagemaker_session, - predictor_cls=predictor_cls, - **kwargs, - ) +import json +import os +import shutil +import tarfile +import tempfile +from contextlib import contextmanager +from typing import Dict, Iterator, Optional - @classmethod - def _prepare_init_params_from_job_description(cls, job_details, model_channel_name=None): - init_params = super()._prepare_init_params_from_job_description( - job_details, model_channel_name=model_channel_name - ) - # These parameters will not be used, but is required to reattach the job - init_params["region"] = "us-east-1" - framework_version, py_version = retrieve_latest_framework_version() - py_version = py_version[0] - init_params["framework_version"] = framework_version - init_params["py_version"] = py_version - return init_params +from autogluon.common.utils.s3_utils import s3_path_to_bucket_prefix +from .utils import safe_unpack_archive -# Documentation for Model: https://sagemaker.readthedocs.io/en/v2/api/inference/model.html#model -class AutoGluonSagemakerInferenceModel(Model): - def __init__( - self, - model_data, - role, - entry_point, - region, - framework_version, - py_version, - instance_type, - custom_image_uri=None, - env=None, - **kwargs, - ): - image_uri = custom_image_uri - if image_uri is None: - image_uri = retrieve_image_uri( - framework_version=framework_version, - region=region, - image_scope="inference", - instance_type=instance_type, - py_version=py_version, - ) - # setting PYTHONUNBUFFERED to disable output buffering for endpoints logging - if env is None: - env = {} - if "PYTHONUNBUFFERED" not in env: - env["PYTHONUNBUFFERED"] = "1" - super().__init__( - model_data=model_data, - role=role, - entry_point=entry_point, - image_uri=image_uri, - env=env, - **kwargs, - ) +SOURCE_DIR_TARBALL_NAME = "sourcedir.tar.gz" - def transformer( - self, - instance_count, - instance_type, - strategy="MultiRecord", - # Maximum size of the payload in a single HTTP request to the container in MB. Will split into multiple batches if a request is more than max_payload - max_payload=6, - max_concurrent_transforms=1, # The maximum number of HTTP requests to be made to each individual transform container at one time. - accept="application/json", - assemble_with="Line", - transform_ami_version=None, - **kwargs, - ): - transformer = super().transformer( - instance_count=instance_count, - instance_type=instance_type, - strategy=strategy, - max_payload=max_payload, - max_concurrent_transforms=max_concurrent_transforms, - accept=accept, - assemble_with=assemble_with, - **kwargs, - ) - if transform_ami_version is not None: - transformer.sagemaker_session = _TransformAmiVersionSession( - transformer.sagemaker_session, - transform_ami_version, - ) - return transformer +def upload_training_code(entry_point: str, sagemaker_session, s3_uri_prefix: str) -> str: + """Bundle the training entry point as ``sourcedir.tar.gz`` and upload it. -class AutoGluonRepackInferenceModel(AutoGluonSagemakerInferenceModel): - """ - Custom implementation to force repack of inference code into model artifacts + Returns the S3 URI of the uploaded tarball. """ + with tempfile.TemporaryDirectory(prefix="ag_train_code_") as tmpdir: + tarball_path = os.path.join(tmpdir, SOURCE_DIR_TARBALL_NAME) + with tarfile.open(tarball_path, "w:gz") as tar: + tar.add(entry_point, arcname=os.path.basename(entry_point)) + bucket, key_prefix = s3_path_to_bucket_prefix(s3_uri_prefix) + return sagemaker_session.upload_data(path=tarball_path, bucket=bucket, key_prefix=key_prefix) - def prepare_container_def( - self, - instance_type=None, - accelerator_type=None, - serverless_inference_config=None, - accept_eula=None, - model_reference_arn=None, - ): # pylint: disable=unused-argument - deploy_key_prefix = fw_utils.model_code_key_prefix(self.key_prefix, self.name, self.image_uri) - deploy_env = copy.deepcopy(self.env) - self._upload_code(deploy_key_prefix, repack=True) - deploy_env.update(self._script_mode_env_vars()) - return sagemaker.container_def( - self.image_uri, - self.repacked_model_data or self.model_data, - deploy_env, - image_config=self.image_config, - ) +def training_script_hyperparameters( + entry_point: str, submit_directory: str, job_name: str, region: str +) -> Dict[str, str]: + """Hyperparameters the SageMaker training toolkit uses to download and run the entry point. -class AutoGluonNonRepackInferenceModel(AutoGluonSagemakerInferenceModel): + Values are JSON-encoded, matching what SageMaker SDK v2 sent; the toolkit JSON-decodes them. """ - Custom implementation to force no repack of inference code into model artifacts. - This requires inference code already present in the trained artifacts, which is created during CloudPredictor training. + return { + "sagemaker_program": json.dumps(os.path.basename(entry_point)), + "sagemaker_submit_directory": json.dumps(submit_directory), + "sagemaker_container_log_level": json.dumps(20), + "sagemaker_job_name": json.dumps(job_name), + "sagemaker_region": json.dumps(region), + } + + +def script_mode_environment(entry_point: str, region: str) -> Dict[str, str]: + """Environment variables pointing the inference toolkit at the serve script under the model's ``code/`` dir.""" + return { + "SAGEMAKER_PROGRAM": os.path.basename(entry_point), + "SAGEMAKER_SUBMIT_DIRECTORY": "/opt/ml/model/code", + "SAGEMAKER_CONTAINER_LOG_LEVEL": "20", + "SAGEMAKER_REGION": region, + } + + +@contextmanager +def staged_serving_code(entry_point: str) -> Iterator[str]: + """Yield a temporary directory holding ``entry_point`` and ``serving_utils/``, i.e. the model's ``code/`` dir.""" + from ..scripts import ScriptManager # deferred: importing scripts pulls in the backend package + + staging_dir = tempfile.mkdtemp(prefix="ag_serving_") + try: + shutil.copy(entry_point, os.path.join(staging_dir, os.path.basename(entry_point))) + shutil.copytree(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, os.path.join(staging_dir, "serving_utils")) + yield staging_dir + finally: + shutil.rmtree(staging_dir, ignore_errors=True) + + +def repack_model_with_serving_code( + model_data: str, + entry_point: str, + repacked_model_uri: str, + sagemaker_session, + kms_key: Optional[str] = None, +) -> str: + """Replace ``code/`` inside the S3 ``model_data`` tarball with ``entry_point`` + ``serving_utils/`` and upload it. + + Returns ``repacked_model_uri``. """ - - def prepare_container_def( - self, - instance_type=None, - accelerator_type=None, - serverless_inference_config=None, - accept_eula=None, - model_reference_arn=None, - ): # pylint: disable=unused-argument - deploy_env = copy.deepcopy(self.env) - deploy_env.update(self._script_mode_env_vars()) - deploy_env[SCRIPT_PARAM_NAME.upper()] = os.path.basename(deploy_env[SCRIPT_PARAM_NAME.upper()]) - deploy_env[DIR_PARAM_NAME.upper()] = "/opt/ml/model/code" - - return sagemaker.container_def( - self.image_uri, - self.model_data, - deploy_env, - image_config=self.image_config, - ) - - -# Predictor documentation: https://sagemaker.readthedocs.io/en/v2/api/inference/predictors.html -class AutoGluonRealtimePredictor(Predictor): - def __init__(self, *args, **kwargs): - super().__init__(*args, serializer=AutoGluonSerializer(), deserializer=PandasDeserializer(), **kwargs) - - -class AutoGluonMultiModalRealtimePredictor(Predictor): - def __init__(self, *args, **kwargs): - super().__init__(*args, serializer=MultiModalSerializer(), deserializer=PandasDeserializer(), **kwargs) - - -# Predictor documentation: https://sagemaker.readthedocs.io/en/v2/api/inference/predictors.html -# SageMaker can only take in csv format for batch transformation because files need to be easily splitable to be batch processed. -class AutoGluonBatchPredictor(Predictor): - def __init__(self, *args, **kwargs): - super().__init__(*args, serializer=CSVSerializer(), **kwargs) + s3 = sagemaker_session.s3_client + with tempfile.TemporaryDirectory(prefix="ag_repack_") as tmpdir: + original_tarball = os.path.join(tmpdir, "original.tar.gz") + s3.download_file(*s3_path_to_bucket_prefix(model_data), original_tarball) + model_dir = os.path.join(tmpdir, "model") + safe_unpack_archive(original_tarball, model_dir) + code_dir = os.path.join(model_dir, "code") + shutil.rmtree(code_dir, ignore_errors=True) + with staged_serving_code(entry_point) as staging_dir: + shutil.copytree(staging_dir, code_dir) + + repacked_tarball = os.path.join(tmpdir, "model.tar.gz") + with tarfile.open(repacked_tarball, "w:gz") as tar: + for name in sorted(os.listdir(model_dir)): + tar.add(os.path.join(model_dir, name), arcname=name) + extra_args = {"ServerSideEncryption": "aws:kms", "SSEKMSKeyId": kms_key} if kms_key else None + s3.upload_file(repacked_tarball, *s3_path_to_bucket_prefix(repacked_model_uri), ExtraArgs=extra_args) + return repacked_model_uri diff --git a/src/autogluon/cloud/utils/aws_utils.py b/src/autogluon/cloud/utils/aws_utils.py index 8a889188..79649e52 100644 --- a/src/autogluon/cloud/utils/aws_utils.py +++ b/src/autogluon/cloud/utils/aws_utils.py @@ -1,13 +1,16 @@ import logging -from typing import Optional +import os +import re +from typing import Any, Dict, List, Optional import boto3 -import sagemaker from botocore.config import Config +from botocore.exceptions import ClientError from autogluon.common.utils.s3_utils import is_s3_url from ..config import load_config +from .misc import sagemaker_timestamp logger = logging.getLogger(__name__) @@ -28,14 +31,107 @@ def _resolve_sagemaker_region() -> Optional[str]: return entry.region -def resolve_execution_role(role: Optional[str], backend_name: str) -> str: +class AwsSession: + """A ``boto3.Session`` together with the clients AutoGluon-Cloud uses, all sharing its credentials and region.""" + + def __init__(self, boto_session: boto3.Session, sagemaker_client: Any = None) -> None: + self.boto_session = boto_session + self.sagemaker_client = sagemaker_client or boto_session.client("sagemaker") + # Realtime inference requests can take up to 60s server-side, so allow some headroom over botocore's default. + self.sagemaker_runtime_client = boto_session.client("sagemaker-runtime", config=Config(read_timeout=80)) + self.s3_client = boto_session.client("s3") + + @property + def boto_region_name(self) -> Optional[str]: + return self.boto_session.region_name + + def upload_data(self, path: str, bucket: str, key_prefix: str, extra_args: Optional[Dict[str, Any]] = None) -> str: + """Upload a local file or directory under ``s3://bucket/key_prefix``. + + A file is uploaded to ``{key_prefix}/{filename}`` and its S3 URI is returned. A directory is uploaded + recursively, preserving its structure below ``key_prefix``, and ``s3://bucket/key_prefix`` is returned. + """ + key_prefix = key_prefix.strip("/") + if os.path.isdir(path): + for dirpath, _, filenames in os.walk(path): + for name in filenames: + local_path = os.path.join(dirpath, name) + key = f"{key_prefix}/{os.path.relpath(local_path, path)}".replace(os.sep, "/") + self.s3_client.upload_file(local_path, bucket, key, ExtraArgs=extra_args) + return f"s3://{bucket}/{key_prefix}" + key = f"{key_prefix}/{os.path.basename(path)}" + self.s3_client.upload_file(path, bucket, key, ExtraArgs=extra_args) + return f"s3://{bucket}/{key}" + + def download_data(self, path: str, bucket: str, key_prefix: str) -> List[str]: + """Download every object under ``s3://bucket/key_prefix`` into the local directory ``path``. + + Objects keep their key relative to ``key_prefix``; if ``key_prefix`` is a single object, it is saved under + its file name. Returns the local paths of the downloaded files. + """ + root = os.path.realpath(path) + folder_prefix = key_prefix.rstrip("/") + "/" + downloaded = [] + for page in self.s3_client.get_paginator("list_objects_v2").paginate(Bucket=bucket, Prefix=key_prefix): + for obj in page.get("Contents", []): + key = obj["Key"] + if key == key_prefix: + relative = os.path.basename(key) + elif key.startswith(folder_prefix) and not key.endswith("/"): + relative = key[len(folder_prefix) :] + else: + continue # folder placeholder, or a sibling that merely shares the string prefix + destination = os.path.realpath(os.path.join(root, relative)) + if not destination.startswith(root + os.sep): + raise ValueError(f"S3 key {key!r} would be downloaded outside of {path!r}.") + os.makedirs(os.path.dirname(destination), exist_ok=True) + self.s3_client.download_file(bucket, key, destination) + downloaded.append(destination) + return downloaded + + +_ASSUMED_ROLE_ARN = re.compile(r"^arn:([^:]+):sts::(\d+):assumed-role/([^/]+)/.+$") + + +def get_execution_role(session: Optional[AwsSession] = None) -> str: + """Return the IAM role whose credentials ``session`` (default: a new session) uses, e.g. the execution role + inside SageMaker. + + Raises ``ValueError`` if the caller is not an assumed role (e.g. an IAM user), since only roles can be passed to + SageMaker as execution roles. + """ + session = session or setup_sagemaker_session() + caller_arn = session.boto_session.client("sts").get_caller_identity()["Arn"] + match = _ASSUMED_ROLE_ARN.match(caller_arn) + if match is None: + raise ValueError( + f"Cannot infer a SageMaker execution role from the current AWS identity {caller_arn}. Pass " + "`role=`, or run `autogluon.cloud.bootstrap()` / `register()` once to " + "persist a role." + ) + partition, account, role_name = match.groups() + try: + # The STS ARN drops the role path (e.g. `service-role/`), which only IAM knows. + return session.boto_session.client("iam").get_role(RoleName=role_name)["Role"]["Arn"] + except ClientError: + # Roles created by the SageMaker console live under `service-role/` and often lack `iam:GetRole`. + path = "service-role/" if role_name.startswith("AmazonSageMaker-ExecutionRole") else "" + role_arn = f"arn:{partition}:iam::{account}:role/{path}{role_name}" + logger.warning( + f"Could not look up role {role_name!r} in IAM, using {role_arn}. Pass `role=` if this is wrong." + ) + return role_arn + + +def resolve_execution_role(role: Optional[str], backend_name: str, *, session: Optional[AwsSession] = None) -> str: """Resolve the SageMaker execution role ARN. Resolution order: 1. ``role`` argument if provided. 2. ``role_arn`` from ``~/.autogluon/cloud.yaml`` under the matching backend slot. - 3. ``sagemaker.get_execution_role()``. + 3. The role whose credentials ``session`` uses (e.g. the execution role inside SageMaker), see + :func:`get_execution_role`. """ if role: return role @@ -45,7 +141,7 @@ def resolve_execution_role(role: Optional[str], backend_name: str) -> str: if entry is not None and entry.role_arn: logger.info(f"Using execution role from ~/.autogluon/cloud.yaml: {entry.role_arn}") return entry.role_arn - return sagemaker.get_execution_role() + return get_execution_role(session) def resolve_cloud_output_path(path: Optional[str], backend_name: str) -> Optional[str]: @@ -81,7 +177,7 @@ def resolve_cloud_output_path(path: Optional[str], backend_name: str) -> Optiona body = path[len("s3://") :] bucket, _, prefix = body.partition("/") if not prefix: - path = f"s3://{bucket}/ag-{sagemaker.utils.sagemaker_timestamp()}" + path = f"s3://{bucket}/ag-{sagemaker_timestamp()}" logger.info(f"cloud_output_path set to {path} (timestamped subfolder under bucket).") else: logger.info(f"cloud_output_path set to {path}.") @@ -130,9 +226,9 @@ def setup_sagemaker_session( read_timeout: int = 60, retries: Optional[dict] = None, **kwargs, -): +) -> AwsSession: """ - Setup a sagemaker session with a given configuration + Setup an :class:`AwsSession` with a given configuration Region resolution (only when ``boto_session`` is not provided): read from ``~/.autogluon/cloud.yaml`` if set, otherwise fall back to the boto3 default chain (env vars, @@ -180,5 +276,4 @@ def setup_sagemaker_session( "`autogluon-cloud register --region `), set the `AWS_DEFAULT_REGION` env var, " "or configure a default region in `~/.aws/config`." ) - sm_boto = boto_session.client("sagemaker", config=config) - return sagemaker.Session(boto_session=boto_session, sagemaker_client=sm_boto) + return AwsSession(boto_session=boto_session, sagemaker_client=boto_session.client("sagemaker", config=config)) diff --git a/src/autogluon/cloud/utils/deserializers.py b/src/autogluon/cloud/utils/deserializers.py index 4bdbec45..1da8abb5 100644 --- a/src/autogluon/cloud/utils/deserializers.py +++ b/src/autogluon/cloud/utils/deserializers.py @@ -2,7 +2,6 @@ from abc import ABC, abstractmethod import pandas as pd -from sagemaker.deserializers import SimpleBaseDeserializer class PandasDeserializeStrategy(ABC): @@ -68,7 +67,7 @@ def get_strategy(content_type: str) -> PandasDeserializeStrategy: return PandasDeserializeStrategyFactory.__content_type_to_strategy[content_type]() -class PandasDeserializer(SimpleBaseDeserializer): +class PandasDeserializer: """Deserialize Parquet, CSV or JSON data from an inference endpoint into a pandas dataframe.""" def __init__(self, accept=("application/x-parquet", "text/csv", "application/json")): @@ -78,7 +77,7 @@ def __init__(self, accept=("application/x-parquet", "text/csv", "application/jso accept (union[str, tuple[str]]): The MIME type (or tuple of allowable MIME types) that is expected from the inference endpoint (default: ("application/x-parquet", "text/csv","application/json")). """ - super().__init__(accept=accept) + self.accept = (accept,) if isinstance(accept, str) else tuple(accept) def deserialize(self, stream, content_type): """Deserialize CSV or JSON data from an inference endpoint into a pandas dataframe. diff --git a/src/autogluon/cloud/utils/dlc_utils.py b/src/autogluon/cloud/utils/dlc_utils.py index de5b6074..07cdefaa 100644 --- a/src/autogluon/cloud/utils/dlc_utils.py +++ b/src/autogluon/cloud/utils/dlc_utils.py @@ -96,11 +96,13 @@ def retrieve_latest_framework_version(framework_type="training"): return versions[-1] -def retrieve_image_uri(framework_version, region, image_scope, instance_type, py_version=None): +def retrieve_image_uri(framework_version, region, image_scope, instance_type, py_version=None, custom_image_uri=None): """Construct the full ECR image URI for a given AG version/region/scope. - Drop-in replacement for sagemaker.image_uris.retrieve("autogluon", ...). + Drop-in replacement for sagemaker.image_uris.retrieve("autogluon", ...). Returns ``custom_image_uri`` as-is if set. """ + if custom_image_uri: + return custom_image_uri config = _load_config() version_info = config[image_scope]["versions"][framework_version] registry = version_info["registries"][region] diff --git a/src/autogluon/cloud/utils/job_logs.py b/src/autogluon/cloud/utils/job_logs.py new file mode 100644 index 00000000..3ec13516 --- /dev/null +++ b/src/autogluon/cloud/utils/job_logs.py @@ -0,0 +1,97 @@ +"""Wait for SageMaker jobs while streaming their CloudWatch logs through the caller's own boto3 session. + +Polling uses the clients we pass in, so logs always come from the job's account and region. +""" + +from __future__ import annotations + +import logging +import time +from typing import Any, Callable, Dict, Optional + +from botocore.exceptions import BotoCoreError, ClientError + +logger = logging.getLogger(__name__) + +TRAINING_JOB_LOG_GROUP = "/aws/sagemaker/TrainingJobs" +TRANSFORM_JOB_LOG_GROUP = "/aws/sagemaker/TransformJobs" +TERMINAL_JOB_STATUSES = ("Completed", "Failed", "Stopped") + + +class LogTailer: + """Print new events from every log stream ``/...`` in ``log_group``, oldest first.""" + + def __init__(self, logs_client: Any, log_group: str, job_name: str) -> None: + self._client = logs_client + self._log_group = log_group + self._prefix = job_name + "/" + self._next_tokens: Dict[str, Optional[str]] = {} + self._enabled = True + + def poll(self) -> None: + """Print all events that arrived since the previous call. Never raises: log access problems disable tailing.""" + if not self._enabled: + return + try: + self._discover_streams() + for stream_name in sorted(self._next_tokens): + self._print_new_events(stream_name) + except ClientError as e: + if e.response.get("Error", {}).get("Code") == "ResourceNotFoundException": + return # The log group / streams appear once the container starts writing output. + self._disable(e) + except BotoCoreError as e: + self._disable(e) + + def _disable(self, error: Exception) -> None: + logger.warning(f"Unable to read job logs from CloudWatch ({error}). Waiting for the job without logs.") + self._enabled = False + + def _discover_streams(self) -> None: + kwargs = {"logGroupName": self._log_group, "logStreamNamePrefix": self._prefix} + while True: + response = self._client.describe_log_streams(**kwargs) + for stream in response.get("logStreams", []): + self._next_tokens.setdefault(stream["logStreamName"], None) + if not response.get("nextToken"): + return + kwargs["nextToken"] = response["nextToken"] + + def _print_new_events(self, stream_name: str) -> None: + # Prefix lines with the stream (e.g. instance id) only when there is more than one stream to tell apart. + label = f"[{stream_name[len(self._prefix) :]}] " if len(self._next_tokens) > 1 else "" + while True: + kwargs = {"logGroupName": self._log_group, "logStreamName": stream_name, "startFromHead": True} + if self._next_tokens[stream_name]: + kwargs["nextToken"] = self._next_tokens[stream_name] + response = self._client.get_log_events(**kwargs) + self._next_tokens[stream_name] = response["nextForwardToken"] + if not response["events"]: + return + for event in response["events"]: + print(f"{label}{event['message']}") + + +def wait_for_job( + get_status: Callable[[], str], + job_name: str, + log_group: str, + logs_client: Optional[Any] = None, + poll: float = 10, +) -> str: + """Poll ``get_status`` until the job reaches a terminal state and return that state. + + If ``logs_client`` (a boto3 ``logs`` client) is given, the job's CloudWatch logs are printed while waiting. + """ + tailer = LogTailer(logs_client, log_group, job_name) if logs_client is not None else None + while True: + status = get_status() + if tailer is not None: + tailer.poll() + if status in TERMINAL_JOB_STATUSES: + if tailer is not None: + # The last log lines can land in CloudWatch shortly after the status flips. + time.sleep(poll) + tailer.poll() + return status + time.sleep(poll) diff --git a/src/autogluon/cloud/utils/misc.py b/src/autogluon/cloud/utils/misc.py index afee43a9..26381934 100644 --- a/src/autogluon/cloud/utils/misc.py +++ b/src/autogluon/cloud/utils/misc.py @@ -1,6 +1,20 @@ +import secrets +import time from collections import OrderedDict +def sagemaker_timestamp() -> str: + """UTC timestamp with millisecond precision, e.g. ``2026-10-02-13-45-07-123``.""" + now = time.time() + return time.strftime("%Y-%m-%d-%H-%M-%S", time.gmtime(now)) + f"-{int(now * 1000) % 1000:03d}" + + +def unique_name_from_base(base: str, max_length: int = 63) -> str: + """Append a timestamp and a random suffix to ``base``, trimming it so the result fits in ``max_length``.""" + suffix = f"-{int(time.time())}-{secrets.token_hex(2)}" + return base[: max_length - len(suffix)] + suffix + + # https://stackoverflow.com/questions/9917178/last-element-in-ordereddict class MostRecentInsertedOrderedDict(OrderedDict): @property diff --git a/src/autogluon/cloud/utils/s3_utils.py b/src/autogluon/cloud/utils/s3_utils.py index ef01406a..d9fba31c 100644 --- a/src/autogluon/cloud/utils/s3_utils.py +++ b/src/autogluon/cloud/utils/s3_utils.py @@ -2,7 +2,6 @@ from typing import Optional import boto3 -import sagemaker from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix @@ -41,10 +40,10 @@ def is_s3_folder(path, session=None): This function tries to determine if a s3 path is a folder. """ assert is_s3_url(path) - if session is None: - session = sagemaker.session.Session() + s3 = boto3.client("s3") if session is None else session.s3_client bucket, prefix = s3_path_to_bucket_prefix(path) - contents = session.list_s3_files(bucket, prefix) + pages = s3.get_paginator("list_objects_v2").paginate(Bucket=bucket, Prefix=prefix) + contents = [obj["Key"] for page in pages for obj in page.get("Contents", [])] if len(contents) > 1: return False # When the folder contains only 1 object, or the prefix is a file results in a len(contents) == 1 diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py new file mode 100644 index 00000000..24f3f142 --- /dev/null +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -0,0 +1,124 @@ +"""Helpers for building SageMaker API requests and sending them through a session's boto3 clients.""" + +import copy +import functools +import logging +from typing import Any, Callable, Dict, Iterable, Mapping, Optional + +from .aws_utils import AwsSession + +logger = logging.getLogger(__name__) + +# Requests that each method sends, i.e. the valid `backend_overrides` keys, named after the boto3 client methods. +# `production_variant` is the single variant inside `create_endpoint_config`'s `ProductionVariants`. +FIT_OVERRIDE_KEYS = ("create_training_job",) +DEPLOY_OVERRIDE_KEYS = ("create_model", "production_variant", "create_endpoint_config", "create_endpoint") +BATCH_PREDICT_OVERRIDE_KEYS = ("create_model", "create_transform_job") +# Fields that link the resources AutoGluon-Cloud creates to each other. Overriding them would point a request at a +# resource we didn't create, which cleanup would then delete. +_RESERVED_OVERRIDE_FIELDS = { + "production_variant": ("ModelName",), + "create_endpoint_config": ("ProductionVariants",), + "create_endpoint": ("EndpointConfigName",), + "create_transform_job": ("ModelName",), +} + +_REMOVED_KWARGS = { + "backend_kwargs": "`backend_overrides` (and `predictions_path` to choose where `predict()` writes results)", + "autogluon_sagemaker_estimator_kwargs": "`backend_overrides={'create_training_job': ...}`", + "fit_kwargs": "`backend_overrides={'create_training_job': ...}`", + "model_kwargs": "`backend_overrides={'create_model': ...}`", + "deploy_kwargs": "`backend_overrides={'production_variant': ..., 'create_endpoint_config': ...}`", + "transformer_kwargs": "`backend_overrides={'create_transform_job': ...}`", + "transform_kwargs": "`backend_overrides={'create_transform_job': ...}`", +} + + +def reject_legacy_kwargs(func): + """Raise an actionable ``TypeError`` for kwargs removed when AutoGluon-Cloud stopped using the SageMaker Python SDK.""" + + @functools.wraps(func) + def wrapper(*args, **kwargs): + for name, value in kwargs.items(): + if name not in _REMOVED_KWARGS: + continue + if _sets_custom_entry_point(value): + raise TypeError( + f"Custom `entry_point` / `source_dir` scripts (passed via `{name}`) are no longer supported: " + "AutoGluon-Cloud always runs its own training and serving scripts. To customize the container, " + "pass `custom_image_uri`." + ) + raise TypeError(f"`{name}` was removed from {func.__qualname__}(). Use {_REMOVED_KWARGS[name]} instead.") + return func(*args, **kwargs) + + return wrapper + + +def _sets_custom_entry_point(value: Any) -> bool: + """Whether a legacy SDK kwargs dict (possibly nested, e.g. ``backend_kwargs["model_kwargs"]``) sets a script.""" + if not isinstance(value, Mapping): + return False + return any(key in ("entry_point", "source_dir") or _sets_custom_entry_point(v) for key, v in value.items()) + + +def check_override_keys(overrides: Optional[Mapping[str, Any]], allowed_keys: Iterable[str]) -> Dict[str, Any]: + """Return ``overrides`` (or ``{}``), raising if it targets a request the calling method doesn't send.""" + overrides = dict(overrides or {}) + unknown = sorted(set(overrides) - set(allowed_keys)) + if unknown: + raise ValueError(f"Unsupported `backend_overrides` key(s) {unknown}. Valid keys: {list(allowed_keys)}.") + for key, fields in _RESERVED_OVERRIDE_FIELDS.items(): + reserved = sorted(set(overrides.get(key, {})) & set(fields)) + if reserved: + raise ValueError(f"`backend_overrides[{key!r}]` cannot set {reserved}; AutoGluon-Cloud manages these.") + return overrides + + +def delete_quietly(delete: Callable[..., Any], **kwargs) -> None: + """Call a ``delete_*`` API during rollback, logging instead of raising so the original error propagates.""" + try: + delete(**kwargs) + except Exception as e: + logger.warning(f"Failed to clean up {kwargs}: {e}") + + +def deep_merge(base: Mapping[str, Any], override: Mapping[str, Any]) -> Dict[str, Any]: + """Merge ``override`` into a copy of ``base``: dicts merge recursively, every other value replaces.""" + merged = copy.deepcopy(dict(base)) + for key, value in override.items(): + if isinstance(value, Mapping) and isinstance(merged.get(key), dict): + merged[key] = deep_merge(merged[key], value) + else: + merged[key] = copy.deepcopy(value) + return merged + + +def invoke_endpoint( + endpoint_name: str, + session: AwsSession, + payload: Any, + serializer, + deserializer, + content_type: Optional[str] = None, + accept: Optional[str] = None, +) -> Any: + """Serialize ``payload``, invoke the endpoint, and deserialize the response.""" + response = session.sagemaker_runtime_client.invoke_endpoint( + EndpointName=endpoint_name, + Body=serializer.serialize(payload), + ContentType=content_type or serializer.content_type, + Accept=accept or ", ".join(deserializer.accept), + ) + return deserializer.deserialize(response["Body"], response["ContentType"]) + + +def delete_endpoint(endpoint_name: str, session: AwsSession) -> None: + """Delete an endpoint together with its endpoint config and models.""" + client = session.sagemaker_client + endpoint = client.describe_endpoint(EndpointName=endpoint_name) + endpoint_config = client.describe_endpoint_config(EndpointConfigName=endpoint["EndpointConfigName"]) + logger.info(f"Deleting endpoint {endpoint_name}") + client.delete_endpoint(EndpointName=endpoint_name) + client.delete_endpoint_config(EndpointConfigName=endpoint["EndpointConfigName"]) + for variant in endpoint_config["ProductionVariants"]: + client.delete_model(ModelName=variant["ModelName"]) diff --git a/src/autogluon/cloud/utils/serializers.py b/src/autogluon/cloud/utils/serializers.py index 1f6633a1..d265a461 100644 --- a/src/autogluon/cloud/utils/serializers.py +++ b/src/autogluon/cloud/utils/serializers.py @@ -5,7 +5,6 @@ import numpy as np import pandas as pd -from sagemaker.serializers import SimpleBaseSerializer AUTOGLUON_SERDE_VERSION = 1 @@ -34,7 +33,7 @@ class AutoGluonSerializationWrapper: known_covariates: Optional[pd.DataFrame] = field(default=None) -class AutoGluonSerializer(SimpleBaseSerializer): +class AutoGluonSerializer: """Serialize data to a buffer with data itself and optional AutoGluon inference arguments.""" def __init__(self, content_type="application/x-autogluon"): @@ -44,7 +43,7 @@ def __init__(self, content_type="application/x-autogluon"): content_type (str): The MIME type to signal to the inference endpoint when sending request data (default: "application/x-autogluon"). """ - super(AutoGluonSerializer, self).__init__(content_type=content_type) + self.content_type = content_type def serialize(self, data: AutoGluonSerializationWrapper): """Serialize data to a JSON envelope with base64-encoded parquet payloads. @@ -74,7 +73,7 @@ def serialize(self, data: AutoGluonSerializationWrapper): return json.dumps(package).encode("utf-8") -class MultiModalSerializer(SimpleBaseSerializer): +class MultiModalSerializer: """Serializer for multi-modal use case. Produces a JSON envelope containing either base64-encoded parquet (for DataFrames) or a @@ -87,11 +86,9 @@ def __init__(self, content_type="application/x-autogluon-parquet"): Args: content_type (str): The MIME type to signal to the inference endpoint when sending request data (default: "application/x-autogluon-parquet"). - To BE NOTICED, this content_type will not used by MultiModalSerializer - as it doesn't support dynamic updating. Instead, we pass expected content_type to - `initial_args` of `predict()` call to endpoints. + Requests with image data pass their own content type to the endpoint call instead. """ - super(MultiModalSerializer, self).__init__(content_type=content_type) + self.content_type = content_type def serialize(self, data): """Serialize data to a JSON envelope. diff --git a/tests/conftest.py b/tests/conftest.py index 256e322e..d79441ad 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -2,8 +2,10 @@ from datetime import datetime, timezone import boto3 +import botocore.session import pandas as pd import pytest +from botocore.validate import validate_parameters from autogluon.cloud.backend import sagemaker_backend @@ -153,3 +155,17 @@ def shared_training_job_name(): @pytest.fixture def test_helper(): return CloudTestHelper + + +@pytest.fixture(scope="session") +def assert_valid_request(): + """Check a generated SageMaker request against botocore's service model, as the boto3 client does before sending. + + Raises ``botocore.exceptions.ParamValidationError`` for unknown, missing or mistyped fields. + """ + service_model = botocore.session.get_session().get_service_model("sagemaker") + + def check(operation: str, request): + validate_parameters(request, service_model.operation_model(operation).input_shape) + + return check diff --git a/tests/unittests/general/test_aws_session.py b/tests/unittests/general/test_aws_session.py new file mode 100644 index 00000000..829c1d1f --- /dev/null +++ b/tests/unittests/general/test_aws_session.py @@ -0,0 +1,136 @@ +import os +import re +import tarfile +from unittest import mock + +import boto3 +import pytest +from botocore.exceptions import ClientError +from moto import mock_aws + +from autogluon.cloud.utils.ag_sagemaker import repack_model_with_serving_code +from autogluon.cloud.utils.aws_utils import AwsSession, get_execution_role +from autogluon.cloud.utils.misc import sagemaker_timestamp, unique_name_from_base + +BUCKET = "test-bucket" + + +@pytest.fixture +def session(monkeypatch): + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "testing") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "testing") + with mock_aws(): + boto_session = boto3.Session(region_name="us-east-1") + boto_session.client("s3").create_bucket(Bucket=BUCKET) + yield AwsSession(boto_session) + + +def _keys(session): + return sorted(obj["Key"] for obj in session.s3_client.list_objects_v2(Bucket=BUCKET).get("Contents", [])) + + +def _write(path, content="x"): + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "w") as f: + f.write(content) + + +def test_upload_file_returns_object_uri(session, tmp_path): + _write(tmp_path / "data.csv") + uri = session.upload_data(str(tmp_path / "data.csv"), BUCKET, "run/utils") + assert uri == f"s3://{BUCKET}/run/utils/data.csv" + assert _keys(session) == ["run/utils/data.csv"] + + +def test_upload_directory_keeps_structure_and_returns_prefix_uri(session, tmp_path): + _write(tmp_path / "code" / "serve.py") + _write(tmp_path / "code" / "serving_utils" / "a.py") + uri = session.upload_data(str(tmp_path / "code"), BUCKET, "run/serving") + assert uri == f"s3://{BUCKET}/run/serving" + assert _keys(session) == ["run/serving/serve.py", "run/serving/serving_utils/a.py"] + + +def test_download_prefix_skips_siblings_sharing_the_string_prefix(session, tmp_path): + for key in ["results/a.out", "results/sub/b.out", "results-other/c.out"]: + session.s3_client.put_object(Bucket=BUCKET, Key=key, Body=b"x") + downloaded = session.download_data(str(tmp_path), BUCKET, "results") + assert sorted(os.path.relpath(p, tmp_path) for p in downloaded) == ["a.out", os.path.join("sub", "b.out")] + + +def test_download_single_object_uses_its_file_name(session, tmp_path): + session.s3_client.put_object(Bucket=BUCKET, Key="bt/results/test.csv.out", Body=b"pred") + downloaded = session.download_data(str(tmp_path), BUCKET, "bt/results/test.csv.out") + assert downloaded == [os.path.join(os.path.realpath(tmp_path), "test.csv.out")] + + +def test_repack_replaces_code_dir_and_keeps_model_files(session, tmp_path): + model_dir = tmp_path / "model" + _write(model_dir / "predictor.pkl", "weights") + _write(model_dir / "code" / "stale.py") + with tarfile.open(tmp_path / "model.tar.gz", "w:gz") as tar: + for name in os.listdir(model_dir): + tar.add(model_dir / name, arcname=name) + session.s3_client.upload_file(str(tmp_path / "model.tar.gz"), BUCKET, "fit/model.tar.gz") + entry_point = tmp_path / "my_serve.py" + _write(entry_point) + + uri = repack_model_with_serving_code( + model_data=f"s3://{BUCKET}/fit/model.tar.gz", + entry_point=str(entry_point), + repacked_model_uri=f"s3://{BUCKET}/endpoints/ep/model/model.tar.gz", + sagemaker_session=session, + ) + + assert uri == f"s3://{BUCKET}/endpoints/ep/model/model.tar.gz" + session.s3_client.download_file(BUCKET, "endpoints/ep/model/model.tar.gz", str(tmp_path / "repacked.tar.gz")) + with tarfile.open(tmp_path / "repacked.tar.gz") as tar: + names = set(tar.getnames()) + assert "predictor.pkl" in names + assert "code/my_serve.py" in names + assert any(name.startswith("code/serving_utils/") for name in names) + assert "code/stale.py" not in names + + +def _session_with_caller(arn, get_role_result=None): + clients = {"sts": mock.MagicMock(), "iam": mock.MagicMock()} + clients["sts"].get_caller_identity.return_value = {"Arn": arn} + if isinstance(get_role_result, Exception): + clients["iam"].get_role.side_effect = get_role_result + else: + clients["iam"].get_role.return_value = {"Role": {"Arn": get_role_result}} + session = mock.MagicMock() + session.boto_session.client.side_effect = lambda name, **kwargs: clients[name] + return session, clients + + +def test_execution_role_resolves_path_through_iam(): + session, clients = _session_with_caller( + "arn:aws:sts::123456789012:assumed-role/MyRole/botocore-session-1", + get_role_result="arn:aws:iam::123456789012:role/team/MyRole", + ) + assert get_execution_role(session) == "arn:aws:iam::123456789012:role/team/MyRole" + clients["iam"].get_role.assert_called_once_with(RoleName="MyRole") + + +@pytest.mark.parametrize( + "role_name, expected_path", + [("MyRole", ""), ("AmazonSageMaker-ExecutionRole-20240101T000000", "service-role/")], +) +def test_execution_role_without_iam_access_falls_back_to_sts_arn(role_name, expected_path): + denied = ClientError({"Error": {"Code": "AccessDenied", "Message": "denied"}}, "GetRole") + session, _ = _session_with_caller(f"arn:aws:sts::123456789012:assumed-role/{role_name}/SageMaker", denied) + assert get_execution_role(session) == f"arn:aws:iam::123456789012:role/{expected_path}{role_name}" + + +def test_execution_role_rejects_iam_users(): + session, _ = _session_with_caller("arn:aws:iam::123456789012:user/alice") + with pytest.raises(ValueError, match="role="): + get_execution_role(session) + + +def test_resource_names(): + assert re.fullmatch(r"\d{4}-\d{2}-\d{2}-\d{2}-\d{2}-\d{2}-\d{3}", sagemaker_timestamp()) + name = unique_name_from_base("a" * 100) + assert len(name) == 63 + assert re.fullmatch(r"a+-\d+-[0-9a-f]{4}", name) + assert unique_name_from_base("ag") != unique_name_from_base("ag") diff --git a/tests/unittests/general/test_aws_utils.py b/tests/unittests/general/test_aws_utils.py index aa345e01..a73cf9aa 100644 --- a/tests/unittests/general/test_aws_utils.py +++ b/tests/unittests/general/test_aws_utils.py @@ -36,7 +36,7 @@ def _save_role_in_config(backend_name: str, role_arn: str) -> None: def test_explicit_role_wins_over_config_and_env(): _save_role_in_config("sagemaker", "arn:aws:iam::111111111111:role/from-config") explicit = "arn:aws:iam::222222222222:role/explicit" - with mock.patch("autogluon.cloud.utils.aws_utils.sagemaker.get_execution_role") as mock_env: + with mock.patch("autogluon.cloud.utils.aws_utils.get_execution_role") as mock_env: assert resolve_execution_role(explicit, backend_name="sagemaker") == explicit mock_env.assert_not_called() @@ -52,7 +52,7 @@ def test_config_role_used_when_no_explicit(): aws_utils_logger.addHandler(handler) aws_utils_logger.setLevel(logging.INFO) try: - with mock.patch("autogluon.cloud.utils.aws_utils.sagemaker.get_execution_role") as mock_env: + with mock.patch("autogluon.cloud.utils.aws_utils.get_execution_role") as mock_env: assert resolve_execution_role(None, backend_name="sagemaker") == config_role mock_env.assert_not_called() finally: @@ -65,7 +65,7 @@ def test_config_role_used_when_no_explicit(): def test_falls_back_to_env_when_no_config_or_explicit(): env_role = "arn:aws:iam::333333333333:role/from-env" with mock.patch( - "autogluon.cloud.utils.aws_utils.sagemaker.get_execution_role", + "autogluon.cloud.utils.aws_utils.get_execution_role", return_value=env_role, ) as mock_env: assert resolve_execution_role(None, backend_name="sagemaker") == env_role @@ -76,13 +76,20 @@ def test_falls_back_to_env_when_backend_missing_in_config(): _save_role_in_config("other_backend", "arn:aws:iam::111111111111:role/other") env_role = "arn:aws:iam::333333333333:role/from-env" with mock.patch( - "autogluon.cloud.utils.aws_utils.sagemaker.get_execution_role", + "autogluon.cloud.utils.aws_utils.get_execution_role", return_value=env_role, ) as mock_env: assert resolve_execution_role(None, backend_name="sagemaker") == env_role mock_env.assert_called_once() +def test_role_fallback_uses_the_backend_session(): + session = mock.sentinel.session + with mock.patch("autogluon.cloud.utils.aws_utils.get_execution_role", return_value="role") as get_role: + assert resolve_execution_role(None, backend_name="sagemaker", session=session) == "role" + get_role.assert_called_once_with(session) + + def _save_bucket_in_config(backend_name: str, bucket: str) -> None: save_config( CloudConfig( diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 7817419d..eb5142e9 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -128,7 +128,7 @@ def test_deploy_passes_artifact_uri_and_overrides_model_path_to_container_dir(): cloud_output_path="s3://b", model_artifact_uri="s3://b/cache/chronos-2/model.tar.gz", ) - fm._backend.endpoint = mock.MagicMock() # _deploy_backend asserts this is set after the call + fm._backend.endpoint_name = "ep" # _deploy_backend asserts this is set after the call fm._deploy_backend() call = fm._backend.deploy.call_args @@ -141,7 +141,7 @@ def test_deploy_passes_artifact_uri_and_overrides_model_path_to_container_dir(): def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): fm = FoundationModel("chronos-2", cloud_output_path="s3://b") - fm._backend.endpoint = mock.MagicMock() + fm._backend.endpoint_name = "ep" fm._deploy_backend() call = fm._backend.deploy.call_args @@ -154,7 +154,7 @@ def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): def test_tabular_deploy_uses_tabular_fm_handler_and_returns_tabular_endpoint(): fm = FoundationModel("mitra-classifier", cloud_output_path="s3://b") - fm._backend.endpoint = mock.MagicMock(endpoint_name="mitra-endpoint") + fm._backend.endpoint_name = "mitra-endpoint" fm._backend.sagemaker_session.boto_session = mock.sentinel.boto_session with mock.patch("autogluon.cloud.model.foundation_model.TabularEndpoint") as endpoint_cls: @@ -162,7 +162,7 @@ def test_tabular_deploy_uses_tabular_fm_handler_and_returns_tabular_endpoint(): call = fm._backend.deploy.call_args assert call.kwargs["instance_type"] == "ml.m5.4xlarge" - assert call.kwargs["model_kwargs"]["entry_point"].endswith("tabular_fm_serve.py") + assert call.kwargs["entry_point"].endswith("tabular_fm_serve.py") assert call.kwargs["fm_serve_config"] == { "ag_model_key": "MITRA", "hyperparameters": { @@ -198,7 +198,7 @@ def test_deploy_rejects_user_model_path_when_artifact_uri_set(): cloud_output_path="s3://b", model_artifact_uri="s3://b/cache/chronos-2/model.tar.gz", ) - fm._backend.endpoint = mock.MagicMock() + fm._backend.endpoint_name = "ep" with pytest.raises(ValueError, match="model_artifact_uri"): fm._deploy_backend(hyperparameters={"model_path": "my-org/something-else"}) @@ -243,16 +243,15 @@ def test_cache_model_artifact_raises_on_stale_version_without_overwrite(): fm.cache_model_artifact("s3://b/cache") -def test_sagemaker_backend_uses_nonrepack_when_repack_is_false(): - """A pre-bundled cached artifact should bypass the SDK's download/repack/re-upload path.""" +def test_sagemaker_backend_skips_repack_when_repack_is_false(): + """A pre-bundled cached artifact should bypass the download/repack/re-upload path.""" from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend sb = "autogluon.cloud.backend.sagemaker_backend" with ( mock.patch(f"{sb}.setup_sagemaker_session", return_value=mock.MagicMock(boto_region_name="us-east-1")), mock.patch(f"{sb}.resolve_execution_role", return_value="arn:aws:iam::000000000000:role/t"), - mock.patch(f"{sb}.AutoGluonNonRepackInferenceModel") as nonrepack_cls, - mock.patch(f"{sb}.AutoGluonRepackInferenceModel") as repack_cls, + mock.patch(f"{sb}.repack_model_with_serving_code") as repack, mock.patch.object(SagemakerBackend, "_upload_predictor", side_effect=lambda p, _: p), ): backend = SagemakerBackend( @@ -264,9 +263,10 @@ def test_sagemaker_backend_uses_nonrepack_when_repack_is_false(): backend.deploy( predictor_path="s3://bucket/cache/chronos-2/model.tar.gz", endpoint_name="ep", - model_kwargs={"entry_point": "stub.py"}, + entry_point="stub.py", repack=False, ) - nonrepack_cls.assert_called_once() - repack_cls.assert_not_called() + repack.assert_not_called() + container = backend.sagemaker_session.sagemaker_client.create_model.call_args.kwargs["PrimaryContainer"] + assert container["ModelDataUrl"] == "s3://bucket/cache/chronos-2/model.tar.gz" diff --git a/tests/unittests/general/test_inference_modes.py b/tests/unittests/general/test_inference_modes.py index 7ad73cf2..a6a351e4 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -1,24 +1,22 @@ -"""Verify that ``inference_mode`` translates to the right ``sagemaker.Model.deploy(...)`` kwargs.""" +"""Verify that ``inference_mode`` translates to the right SageMaker endpoint config production variant.""" from unittest import mock import pytest -from sagemaker.serverless import ServerlessInferenceConfig from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend GPU_IMAGE_URI = "123456789012.dkr.ecr.us-east-1.amazonaws.com/autogluon:1.6-cu133-amzn2023" +SB = "autogluon.cloud.backend.sagemaker_backend" @pytest.fixture -def deploy_kwargs(): - """Run ``SagemakerBackend.deploy(...)`` with AWS calls and the SDK Model class mocked, - and return the kwargs that reached ``model.deploy(...)``.""" - sb = "autogluon.cloud.backend.sagemaker_backend" +def deploy_requests(assert_valid_request): + """Run ``SagemakerBackend.deploy(...)`` with AWS calls mocked, and return the validated requests sent to + ``create_model`` / ``create_endpoint_config`` / ``create_endpoint``.""" with ( - mock.patch(f"{sb}.setup_sagemaker_session", return_value=mock.MagicMock(boto_region_name="us-east-1")), - mock.patch(f"{sb}.resolve_execution_role", return_value="arn:aws:iam::000000000000:role/test"), - mock.patch(f"{sb}.AutoGluonNonRepackInferenceModel") as model_cls, + mock.patch(f"{SB}.setup_sagemaker_session", return_value=mock.MagicMock(boto_region_name="us-east-1")), + mock.patch(f"{SB}.resolve_execution_role", return_value="arn:aws:iam::000000000000:role/test"), mock.patch.object(SagemakerBackend, "_create_serve_script_tarball", return_value="s3://stub/m.tar.gz"), ): backend = SagemakerBackend( @@ -29,50 +27,115 @@ def deploy_kwargs(): backend._fit_job = None # deploy a serve-script tarball, not a fit-job artifact def run(**kwargs): - backend.endpoint = None # allow re-deploy across cases - backend.deploy(endpoint_name="ep", model_kwargs={"entry_point": "stub.py"}, **kwargs) - return model_cls.return_value.deploy.call_args.kwargs - + backend.endpoint_name = None # allow re-deploy across cases + backend.deploy(endpoint_name="ep", entry_point="stub.py", **kwargs) + client = backend.sagemaker_session.sagemaker_client + requests = { + "model": client.create_model.call_args.kwargs, + "endpoint_config": client.create_endpoint_config.call_args.kwargs, + "endpoint": client.create_endpoint.call_args.kwargs, + } + assert_valid_request("CreateModel", requests["model"]) + assert_valid_request("CreateEndpointConfig", requests["endpoint_config"]) + assert_valid_request("CreateEndpoint", requests["endpoint"]) + return requests + + run.backend = backend yield run -def test_when_inference_mode_realtime_then_instance_kwargs_are_passed(deploy_kwargs): - captured = deploy_kwargs(instance_type="ml.m5.xlarge", initial_instance_count=2) - assert captured["instance_type"] == "ml.m5.xlarge" - assert captured["initial_instance_count"] == 2 - assert "serverless_inference_config" not in captured +def _variant(requests): + (variant,) = requests["endpoint_config"]["ProductionVariants"] + return variant + + +def test_when_inference_mode_realtime_then_instance_settings_are_in_variant(deploy_requests): + variant = _variant(deploy_requests(instance_type="ml.m5.xlarge", initial_instance_count=2)) + assert variant["InstanceType"] == "ml.m5.xlarge" + assert variant["InitialInstanceCount"] == 2 + assert "ServerlessConfig" not in variant + + +def test_when_deployed_then_model_endpoint_config_and_endpoint_are_linked(deploy_requests): + requests = deploy_requests(instance_type="ml.m5.xlarge") + assert _variant(requests)["ModelName"] == requests["model"]["ModelName"] + assert requests["endpoint"]["EndpointConfigName"] == requests["endpoint_config"]["EndpointConfigName"] + assert requests["endpoint"]["EndpointName"] == "ep" + environment = requests["model"]["PrimaryContainer"]["Environment"] + assert environment["SAGEMAKER_PROGRAM"] == "stub.py" + assert environment["SAGEMAKER_SUBMIT_DIRECTORY"] == "/opt/ml/model/code" -def test_when_cuda_13_custom_image_then_inference_ami_is_inferred(deploy_kwargs): - captured = deploy_kwargs(instance_type="ml.g4dn.xlarge", custom_image_uri=GPU_IMAGE_URI) - assert captured["inference_ami_version"] == "al2023-ami-sagemaker-inference-gpu-4-1" +def test_when_cuda_13_custom_image_then_inference_ami_is_inferred(deploy_requests): + variant = _variant(deploy_requests(instance_type="ml.g4dn.xlarge", custom_image_uri=GPU_IMAGE_URI)) + assert variant["InferenceAmiVersion"] == "al2023-ami-sagemaker-inference-gpu-4-1" -def test_when_inference_ami_is_provided_then_it_is_not_overridden(deploy_kwargs): - captured = deploy_kwargs( - instance_type="ml.g4dn.xlarge", - custom_image_uri=GPU_IMAGE_URI, - deploy_kwargs={"inference_ami_version": "custom-ami"}, +def test_when_inference_ami_is_overridden_then_override_wins(deploy_requests): + variant = _variant( + deploy_requests( + instance_type="ml.g4dn.xlarge", + custom_image_uri=GPU_IMAGE_URI, + backend_overrides={"production_variant": {"InferenceAmiVersion": "custom-ami"}}, + ) ) - assert captured["inference_ami_version"] == "custom-ami" + assert variant["InferenceAmiVersion"] == "custom-ami" + + +def test_when_inference_mode_serverless_then_preset_serverless_config_is_used(deploy_requests): + variant = _variant(deploy_requests(inference_mode="serverless")) + assert variant["ServerlessConfig"] == {"MemorySizeInMB": 4096, "MaxConcurrency": 5} + assert "InstanceType" not in variant + +def test_when_inference_config_provided_then_user_values_override_preset(deploy_requests): + variant = _variant(deploy_requests(inference_mode="serverless", inference_config={"memory_size_in_mb": 8192})) + assert variant["ServerlessConfig"]["MemorySizeInMB"] == 8192 + assert variant["ServerlessConfig"]["MaxConcurrency"] == 5 # preset wins for keys the user didn't override -def test_when_inference_mode_serverless_then_preset_serverless_config_is_used(deploy_kwargs): - captured = deploy_kwargs(inference_mode="serverless") - cfg = captured["serverless_inference_config"] - assert isinstance(cfg, ServerlessInferenceConfig) - assert cfg.memory_size_in_mb == 4096 - assert cfg.max_concurrency == 5 - assert "instance_type" not in captured +def test_when_inference_config_has_unknown_key_then_value_error_is_raised(deploy_requests): + with pytest.raises(ValueError, match="memory_size"): + deploy_requests(inference_mode="serverless", inference_config={"memory_size": 8192}) + deploy_requests.backend.sagemaker_session.sagemaker_client.create_model.assert_not_called() -def test_when_inference_config_provided_then_user_values_override_preset(deploy_kwargs): - captured = deploy_kwargs(inference_mode="serverless", inference_config={"memory_size_in_mb": 8192}) - cfg = captured["serverless_inference_config"] - assert cfg.memory_size_in_mb == 8192 - assert cfg.max_concurrency == 5 # preset wins for keys the user didn't override +@pytest.mark.parametrize("wait", [True, False]) +def test_endpoint_waiter_is_used_only_when_waiting(deploy_requests, wait): + deploy_requests(instance_type="ml.m5.xlarge", wait=wait) + client = deploy_requests.backend.sagemaker_session.sagemaker_client + if wait: + client.get_waiter.assert_called_once_with("endpoint_in_service") + client.get_waiter.return_value.wait.assert_called_once_with(EndpointName="ep") + else: + client.get_waiter.assert_not_called() -def test_when_inference_mode_is_unknown_then_value_error_is_raised(deploy_kwargs): + +def test_when_inference_mode_is_unknown_then_value_error_is_raised(deploy_requests): with pytest.raises(ValueError, match="Unsupported inference_mode"): - deploy_kwargs(inference_mode="batch") + deploy_requests(inference_mode="batch") + + +def test_when_container_environment_overridden_then_it_merges_with_defaults(deploy_requests): + requests = deploy_requests( + instance_type="ml.m5.xlarge", + backend_overrides={"create_model": {"PrimaryContainer": {"Environment": {"FOO": "bar"}}}}, + ) + environment = requests["model"]["PrimaryContainer"]["Environment"] + assert environment["FOO"] == "bar" + assert environment["SAGEMAKER_MODEL_SERVER_WORKERS"] == "1" + + +def test_when_override_targets_training_job_then_deploy_rejects_it(deploy_requests): + with pytest.raises(ValueError, match="Unsupported `backend_overrides` key"): + deploy_requests(backend_overrides={"create_training_job": {}}) + + +def test_when_endpoint_creation_fails_then_model_and_config_are_deleted(deploy_requests): + client = deploy_requests.backend.sagemaker_session.sagemaker_client + client.create_endpoint.side_effect = RuntimeError("boom") + with pytest.raises(RuntimeError, match="boom"): + deploy_requests(instance_type="ml.m5.xlarge") + client.delete_endpoint_config.assert_called_once_with(EndpointConfigName="ep") + client.delete_model.assert_called_once_with(ModelName=client.create_model.call_args.kwargs["ModelName"]) + assert deploy_requests.backend.endpoint_name is None diff --git a/tests/unittests/general/test_job_logs.py b/tests/unittests/general/test_job_logs.py new file mode 100644 index 00000000..81bdc7a8 --- /dev/null +++ b/tests/unittests/general/test_job_logs.py @@ -0,0 +1,78 @@ +from unittest import mock + +import pytest +from botocore.exceptions import ClientError + +from autogluon.cloud.utils.job_logs import LogTailer, wait_for_job + + +class FakeLogsClient: + """In-memory CloudWatch Logs client: ``streams`` maps stream name -> list of messages written so far.""" + + def __init__(self): + self.streams = {} + + def describe_log_streams(self, logGroupName, logStreamNamePrefix, nextToken=None): + if not self.streams: + raise ClientError({"Error": {"Code": "ResourceNotFoundException"}}, "DescribeLogStreams") + names = sorted(n for n in self.streams if n.startswith(logStreamNamePrefix)) + return {"logStreams": [{"logStreamName": n} for n in names]} + + def get_log_events(self, logGroupName, logStreamName, startFromHead, nextToken=None): + start = int(nextToken or 0) + messages = self.streams[logStreamName][start:] + return {"events": [{"message": m} for m in messages], "nextForwardToken": str(start + len(messages))} + + +def test_log_tailer_prints_each_event_once(capsys): + client = FakeLogsClient() + tailer = LogTailer(client, "/aws/sagemaker/TrainingJobs", "job") + + tailer.poll() # no log group yet + client.streams["job/algo-1"] = ["a", "b"] + tailer.poll() + client.streams["job/algo-1"].append("c") + client.streams["other-job/algo-1"] = ["not mine"] + tailer.poll() + + assert capsys.readouterr().out.splitlines() == ["a", "b", "c"] + + +def test_log_tailer_labels_lines_when_there_are_multiple_streams(capsys): + client = FakeLogsClient() + client.streams = {"job/i-1": ["x"], "job/i-1/data-log": ["y"]} + LogTailer(client, "/aws/sagemaker/TransformJobs", "job").poll() + assert capsys.readouterr().out.splitlines() == ["[i-1] x", "[i-1/data-log] y"] + + +def test_log_tailer_disables_itself_on_access_errors(capsys): + client = mock.MagicMock() + client.describe_log_streams.side_effect = ClientError({"Error": {"Code": "AccessDeniedException"}}, "Describe") + tailer = LogTailer(client, "/aws/sagemaker/TrainingJobs", "job") + tailer.poll() + tailer.poll() + assert client.describe_log_streams.call_count == 1 + + +@pytest.mark.parametrize("final_status", ["Completed", "Failed", "Stopped"]) +def test_wait_for_job_returns_terminal_status_and_drains_logs(final_status, capsys): + client = FakeLogsClient() + statuses = iter(["InProgress", "InProgress", final_status]) + + def get_status(): + status = next(statuses) + client.streams.setdefault("job/algo-1", []).append(status) + return status + + with mock.patch("autogluon.cloud.utils.job_logs.time.sleep") as sleep: + assert wait_for_job(get_status, job_name="job", log_group="g", logs_client=client, poll=3) == final_status + + assert capsys.readouterr().out.splitlines() == ["InProgress", "InProgress", final_status] + assert all(call.args == (3,) for call in sleep.call_args_list) + + +def test_wait_for_job_without_logs_only_polls_status(): + statuses = iter(["InProgress", "Completed"]) + with mock.patch("autogluon.cloud.utils.job_logs.time.sleep") as sleep: + assert wait_for_job(lambda: next(statuses), job_name="job", log_group="g") == "Completed" + assert sleep.call_count == 1 diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index 402ddb80..e8c8b056 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -1,9 +1,9 @@ from unittest import mock +import pandas as pd import pytest -from autogluon.cloud.job.sagemaker_job import SageMakerBatchTransformationJob -from autogluon.cloud.utils.ag_sagemaker import _TransformAmiVersionSession +from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend from autogluon.cloud.utils.dlc_utils import infer_sagemaker_ami_version GPU_IMAGE_URI = "123456789012.dkr.ecr.us-east-1.amazonaws.com/autogluon:1.6-cu133-amzn2023" @@ -45,52 +45,75 @@ def test_infer_realtime_ami_ignores_unsupported_or_already_compatible_instance_f assert infer_sagemaker_ami_version(GPU_IMAGE_URI, instance_type, "inference") is None -def test_transform_ami_session_injects_ami_without_mutating_input(): - session = mock.MagicMock() - wrapper = _TransformAmiVersionSession(session, "al2-ami-sagemaker-batch-gpu-535") - resource_config = {"InstanceCount": 1, "InstanceType": "ml.g4dn.xlarge"} - - wrapper.transform(resource_config=resource_config, job_name="job") - - assert "TransformAmiVersion" not in resource_config - assert session.transform.call_args.kwargs["resource_config"]["TransformAmiVersion"] == ( - "al2-ami-sagemaker-batch-gpu-535" - ) +SB = "autogluon.cloud.backend.sagemaker_backend" + + +@pytest.fixture +def transform_request(assert_valid_request): + """Run ``TabularSagemakerBackend._predict(...)`` with AWS calls mocked and return the ``CreateTransformJob`` request.""" + + def run(**predict_kwargs): + with ( + mock.patch(f"{SB}.setup_sagemaker_session", return_value=mock.MagicMock(boto_region_name="us-east-1")), + mock.patch(f"{SB}.resolve_execution_role", return_value="arn:aws:iam::000000000000:role/test"), + mock.patch(f"{SB}.SageMakerBatchTransformationJob") as job_cls, + mock.patch.object(TabularSagemakerBackend, "_upload_predictor", side_effect=lambda path, _: path), + mock.patch.object( + TabularSagemakerBackend, "_upload_batch_predict_data", return_value="s3://input/data.csv" + ), + mock.patch.object(TabularSagemakerBackend, "_prepare_model_data", return_value="s3://bucket/model.tar.gz"), + mock.patch.object(TabularSagemakerBackend, "_create_model", return_value="job"), + ): + backend = TabularSagemakerBackend( + local_output_path="/tmp/test", + cloud_output_path="s3://bucket/run", + predictor_type="tabular", + ) + backend._fit_job = mock.MagicMock() + backend._predict( + test_data=pd.DataFrame({"x": [1]}), + predictor_path="s3://bucket/model.tar.gz", + job_name="job", + wait=False, + **predict_kwargs, + ) + request = job_cls.return_value.run.call_args.kwargs["transform_job_request"] + assert_valid_request("CreateTransformJob", request) + return request + + return run @pytest.mark.parametrize( - ("transformer_kwargs", "expected"), + ("backend_overrides", "expected"), [ - ({}, "al2-ami-sagemaker-batch-gpu-535"), - ({"transform_ami_version": "custom-ami"}, "custom-ami"), + (None, "al2-ami-sagemaker-batch-gpu-535"), + ({"create_transform_job": {"TransformResources": {"TransformAmiVersion": "custom-ami"}}}, "custom-ami"), ], ) -def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(transformer_kwargs, expected): - sj = "autogluon.cloud.job.sagemaker_job" - transformer = mock.MagicMock(output_path="s3://output") - transformer.latest_transform_job.name = "job" - with mock.patch(f"{sj}.AutoGluonNonRepackInferenceModel") as model_cls: - model_cls.return_value.transformer.return_value = transformer - job = SageMakerBatchTransformationJob(session=mock.MagicMock()) - job.run( - model_data="s3://bucket/model.tar.gz", - role="role", - region="us-east-1", - framework_version=None, - py_version=None, - instance_count=1, - instance_type="ml.g4dn.xlarge", - entry_point="serve.py", - predictor_cls=mock.MagicMock(), - output_path="s3://output", - test_input="s3://input/data.csv", - job_name="job", - split_type="Line", - content_type="text/csv", - custom_image_uri=GPU_IMAGE_URI, - wait=False, - model_kwargs={}, - transformer_kwargs=transformer_kwargs, - ) - - assert model_cls.return_value.transformer.call_args.kwargs["transform_ami_version"] == expected +def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value( + transform_request, backend_overrides, expected +): + request = transform_request( + instance_type="ml.g4dn.xlarge", + custom_image_uri=GPU_IMAGE_URI, + backend_overrides=backend_overrides, + ) + assert request["TransformResources"]["TransformAmiVersion"] == expected + assert request["TransformResources"]["InstanceType"] == "ml.g4dn.xlarge" + + +def test_batch_transform_writes_results_to_predictions_path(transform_request): + request = transform_request(predictions_path="s3://my-bucket/preds/") + assert request["TransformOutput"]["S3OutputPath"] == "s3://my-bucket/preds" + + +def test_batch_transform_results_default_to_cloud_output_path(transform_request): + output_path = transform_request()["TransformOutput"]["S3OutputPath"] + assert output_path.startswith("s3://bucket/run/batch_transform/") + assert output_path.endswith("/results") + + +def test_batch_transform_rejects_non_s3_predictions_path(transform_request): + with pytest.raises(ValueError, match="S3 URL"): + transform_request(predictions_path="/tmp/preds") diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py new file mode 100644 index 00000000..1933a065 --- /dev/null +++ b/tests/unittests/general/test_sagemaker_api.py @@ -0,0 +1,140 @@ +from unittest import mock + +import pandas as pd +import pytest + +from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend +from autogluon.cloud.utils.sagemaker_api import ( + check_override_keys, + deep_merge, + delete_endpoint, + invoke_endpoint, + reject_legacy_kwargs, +) + +SB = "autogluon.cloud.backend.sagemaker_backend" + + +def test_deep_merge_merges_dicts_and_replaces_other_values(): + base = {"a": {"x": 1, "y": 2}, "b": [1, 2], "c": 1} + merged = deep_merge(base, {"a": {"y": 3}, "b": [3], "d": 4}) + assert merged == {"a": {"x": 1, "y": 3}, "b": [3], "c": 1, "d": 4} + assert base["a"]["y"] == 2 # input is not mutated + + +def test_check_override_keys_rejects_unknown_keys(): + with pytest.raises(ValueError, match="create_model"): + check_override_keys({"create_model": {}}, ("create_training_job",)) + assert check_override_keys(None, ("create_training_job",)) == {} + + +def test_check_override_keys_rejects_resource_references(): + with pytest.raises(ValueError, match="ModelName"): + check_override_keys({"production_variant": {"ModelName": "other"}}, ("production_variant",)) + + +def test_reject_legacy_kwargs_points_to_replacement(): + @reject_legacy_kwargs + def fit(**kwargs): + return kwargs + + assert fit(custom_image_uri="x") == {"custom_image_uri": "x"} + with pytest.raises(TypeError, match="backend_overrides"): + fit(backend_kwargs={}) + with pytest.raises(TypeError, match="no longer supported"): + fit(backend_kwargs={"model_kwargs": {"entry_point": "serve.py"}}) + + +def test_delete_endpoint_removes_endpoint_config_and_models(): + session = mock.MagicMock() + client = session.sagemaker_client + client.describe_endpoint.return_value = {"EndpointConfigName": "ep-config"} + client.describe_endpoint_config.return_value = {"ProductionVariants": [{"ModelName": "m-1"}, {"ModelName": "m-2"}]} + delete_endpoint("ep", session) + client.delete_endpoint.assert_called_once_with(EndpointName="ep") + client.delete_endpoint_config.assert_called_once_with(EndpointConfigName="ep-config") + assert [c.kwargs for c in client.delete_model.call_args_list] == [{"ModelName": "m-1"}, {"ModelName": "m-2"}] + + +def test_invoke_endpoint_uses_the_session_runtime_client(): + session = mock.MagicMock() + runtime = session.sagemaker_runtime_client + runtime.invoke_endpoint.return_value = {"Body": mock.sentinel.body, "ContentType": "text/csv"} + serializer = mock.Mock(content_type="text/csv") + deserializer = mock.Mock(accept=("application/json",)) + result = invoke_endpoint("ep", session, "payload", serializer=serializer, deserializer=deserializer) + runtime.invoke_endpoint.assert_called_once_with( + EndpointName="ep", + Body=serializer.serialize.return_value, + ContentType="text/csv", + Accept="application/json", + ) + deserializer.deserialize.assert_called_once_with(mock.sentinel.body, "text/csv") + assert result is deserializer.deserialize.return_value + + +@pytest.fixture +def fit_request(tmp_path, assert_valid_request): + """Run ``SagemakerBackend.fit(...)`` with uploads mocked and return the ``CreateTrainingJob`` request.""" + with ( + mock.patch(f"{SB}.setup_sagemaker_session", return_value=mock.MagicMock(boto_region_name="us-east-1")), + mock.patch(f"{SB}.resolve_execution_role", return_value="arn:aws:iam::000000000000:role/test"), + mock.patch(f"{SB}.upload_training_code", return_value="s3://bucket/run/code/job/source/sourcedir.tar.gz"), + mock.patch(f"{SB}.SageMakerFitJob") as fit_job_cls, + mock.patch.object( + TabularSagemakerBackend, "_upload_fit_artifact", return_value={"train_data": "s3://b/train.csv"} + ), + ): + + def run(**fit_kwargs): + backend = TabularSagemakerBackend( + local_output_path=str(tmp_path), + cloud_output_path="s3://bucket/run", + predictor_type="tabular", + ) + backend.fit( + predictor_init_args={"label": "y"}, + predictor_fit_args={}, + data_channels={"train_data": pd.DataFrame({"x": [1], "y": [0]})}, + job_name="job", + custom_image_uri="example.com/autogluon:train", + **fit_kwargs, + ) + request = fit_job_cls.return_value.run.call_args.kwargs["training_job_request"] + assert_valid_request("CreateTrainingJob", request) + return request + + yield run + + +def test_fit_builds_script_mode_training_job(fit_request): + request = fit_request(timeout=3600) + assert request["TrainingJobName"] == "job" + assert request["AlgorithmSpecification"]["TrainingImage"] == "example.com/autogluon:train" + assert request["HyperParameters"]["sagemaker_program"] == '"train.py"' + assert request["HyperParameters"]["sagemaker_submit_directory"] == '"/opt/ml/input/data/code/sourcedir.tar.gz"' + channels = {c["ChannelName"]: c["DataSource"]["S3DataSource"]["S3Uri"] for c in request["InputDataConfig"]} + assert channels["code"] == "s3://bucket/run/code/job/source/sourcedir.tar.gz" + assert "train_data" in channels + assert request["StoppingCondition"] == {"MaxRuntimeInSeconds": 3600} + assert request["OutputDataConfig"] == {"S3OutputPath": "s3://bucket/run/model"} + assert {"Key": "autogluon-cloud-module", "Value": "tabular"} in request["Tags"] + + +def test_fit_applies_overrides(fit_request): + request = fit_request( + backend_overrides={"create_training_job": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}, + ) + assert request["RetryStrategy"] == {"MaximumRetryAttempts": 2} + + +def test_fit_rejects_local_mode(fit_request): + with pytest.raises(ValueError, match="local mode"): + fit_request(instance_type="local") + + +def test_misspelled_override_field_fails_request_validation(fit_request): + from botocore.exceptions import ParamValidationError + + with pytest.raises(ParamValidationError, match="RetryStrategyy"): + fit_request(backend_overrides={"create_training_job": {"RetryStrategyy": {"MaximumRetryAttempts": 2}}}) diff --git a/tests/unittests/general/test_tabular_endpoint.py b/tests/unittests/general/test_tabular_endpoint.py index baf934ca..b4bf1525 100644 --- a/tests/unittests/general/test_tabular_endpoint.py +++ b/tests/unittests/general/test_tabular_endpoint.py @@ -6,33 +6,46 @@ from autogluon.cloud.endpoint.tabular_endpoint import TabularEndpoint from autogluon.cloud.utils.serializers import AutoGluonSerializationWrapper +TE = "autogluon.cloud.endpoint.tabular_endpoint" -def _make_endpoint(response): - endpoint = TabularEndpoint.__new__(TabularEndpoint) - endpoint._predictor = mock.MagicMock(endpoint_name="tabular-fm-endpoint") - endpoint._predictor.predict.return_value = response - return endpoint +@pytest.fixture(autouse=True) +def invoke_endpoint(): + with mock.patch(f"{TE}.invoke_endpoint") as invoke: + yield invoke -def test_predict_sends_train_data_and_returns_prediction_series(): + +@pytest.fixture +def make_endpoint(invoke_endpoint): + def make(response): + endpoint = TabularEndpoint.__new__(TabularEndpoint) + endpoint._endpoint_name = "tabular-fm-endpoint" + endpoint._session = mock.sentinel.session + invoke_endpoint.return_value = response + return endpoint + + return make + + +def test_predict_sends_train_data_and_returns_prediction_series(make_endpoint, invoke_endpoint): train_data = pd.DataFrame({"feature": [0, 1], "label": ["a", "b"]}) data = pd.DataFrame({"feature": [2, 3]}) response = pd.DataFrame({"label": ["a", "b"], "a_proba": [0.8, 0.2], "b_proba": [0.2, 0.8]}) - endpoint = _make_endpoint(response) + endpoint = make_endpoint(response) pred = endpoint.predict(data=data, train_data=train_data, label="label") assert pred.tolist() == ["a", "b"] - payload = endpoint._predictor.predict.call_args.args[0] + payload = invoke_endpoint.call_args.args[2] assert isinstance(payload, AutoGluonSerializationWrapper) pd.testing.assert_frame_equal(payload.data, data) pd.testing.assert_frame_equal(payload.train_data, train_data) assert payload.inference_kwargs == {"label": "label"} -def test_predict_proba_matches_batch_result_shape(): +def test_predict_proba_matches_batch_result_shape(make_endpoint, invoke_endpoint): response = pd.DataFrame({"label": ["a"], "a_proba": [0.7], "b_proba": [0.3]}) - endpoint = _make_endpoint(response) + endpoint = make_endpoint(response) train_data = pd.DataFrame({"feature": [0, 1], "label": ["a", "b"]}) data = pd.DataFrame({"feature": [2]}) @@ -43,8 +56,8 @@ def test_predict_proba_matches_batch_result_shape(): assert proba.iloc[0].tolist() == [0.7, 0.3] -def test_regression_predict_proba_equals_prediction(): - endpoint = _make_endpoint(pd.DataFrame({"target": [1.5, 2.5]})) +def test_regression_predict_proba_equals_prediction(make_endpoint, invoke_endpoint): + endpoint = make_endpoint(pd.DataFrame({"target": [1.5, 2.5]})) train_data = pd.DataFrame({"feature": [0, 1], "target": [0.0, 1.0]}) data = pd.DataFrame({"feature": [2, 3]}) @@ -53,8 +66,8 @@ def test_regression_predict_proba_equals_prediction(): pd.testing.assert_series_equal(pred, proba) -def test_predict_validates_label_and_feature_columns(): - endpoint = _make_endpoint(pd.DataFrame()) +def test_predict_validates_label_and_feature_columns(make_endpoint, invoke_endpoint): + endpoint = make_endpoint(pd.DataFrame()) train_data = pd.DataFrame({"feature": [0], "label": ["a"]}) with pytest.raises(ValueError, match="Label column"): @@ -64,9 +77,9 @@ def test_predict_validates_label_and_feature_columns(): endpoint.predict(data=pd.DataFrame({"other": [1]}), train_data=train_data, label="label") -def test_delete_endpoint_removes_model_endpoint_and_config(): - endpoint = _make_endpoint(pd.DataFrame()) - endpoint.delete_endpoint() +def test_delete_endpoint_removes_model_endpoint_and_config(make_endpoint, invoke_endpoint): + endpoint = make_endpoint(pd.DataFrame()) + with mock.patch(f"{TE}.delete_endpoint") as delete_endpoint: + endpoint.delete_endpoint() - endpoint._predictor.delete_model.assert_called_once_with() - endpoint._predictor.delete_endpoint.assert_called_once_with(delete_endpoint_config=True) + delete_endpoint.assert_called_once_with("tabular-fm-endpoint", mock.sentinel.session) diff --git a/tests/unittests/tabular/test_tabular_fit_predict.py b/tests/unittests/tabular/test_tabular_fit_predict.py index 827e68f1..94116a73 100644 --- a/tests/unittests/tabular/test_tabular_fit_predict.py +++ b/tests/unittests/tabular/test_tabular_fit_predict.py @@ -34,7 +34,7 @@ def cloud_predictor(): def test_when_fit_predict_then_launches_predict_job_and_returns_prediction_series(cloud_predictor): pred = cloud_predictor.fit_predict(**DEFAULT_ARGS) - extra_ag_args = cloud_predictor.fit.call_args.kwargs["backend_kwargs"]["extra_ag_args"] + extra_ag_args = cloud_predictor.fit.call_args.kwargs["extra_ag_args"] assert extra_ag_args["predict_after_fit"] is True assert "predictions_path" not in extra_ag_args # not passed -> backend fills in a default @@ -44,7 +44,7 @@ def test_when_fit_predict_then_launches_predict_job_and_returns_prediction_serie def test_when_predictions_path_given_then_forwarded_to_backend(cloud_predictor): cloud_predictor.fit_predict(**DEFAULT_ARGS, predictions_path="s3://bucket/key/predictions.csv") - extra_ag_args = cloud_predictor.fit.call_args.kwargs["backend_kwargs"]["extra_ag_args"] + extra_ag_args = cloud_predictor.fit.call_args.kwargs["extra_ag_args"] assert extra_ag_args["predictions_path"] == "s3://bucket/key/predictions.csv"