From b120e673a713f6a43b9a68cb49656fc9f3722a7b Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Thu, 1 Oct 2026 16:15:45 +0000 Subject: [PATCH 01/16] Migrate to SageMaker Python SDK v3 (sagemaker-core) --- pyproject.toml | 5 +- src/autogluon/cloud/__init__.py | 11 - src/autogluon/cloud/backend/backend.py | 27 +- .../backend/multimodal_sagemaker_backend.py | 60 +- .../cloud/backend/sagemaker_backend.py | 843 +++++++++--------- .../backend/timeseries_sagemaker_backend.py | 43 +- src/autogluon/cloud/endpoint/endpoint.py | 26 - .../cloud/endpoint/prediction_future.py | 4 +- .../cloud/endpoint/sagemaker_endpoint.py | 47 - .../cloud/endpoint/tabular_endpoint.py | 24 +- .../cloud/endpoint/timeseries_endpoint.py | 24 +- src/autogluon/cloud/job/sagemaker_job.py | 248 ++---- src/autogluon/cloud/model/foundation_model.py | 129 +-- .../cloud/predictor/cloud_predictor.py | 256 +++--- .../predictor/tabular_cloud_predictor.py | 53 +- .../predictor/timeseries_cloud_predictor.py | 126 +-- .../cloud/scripts/sagemaker_scripts/train.py | 2 +- src/autogluon/cloud/utils/ag_sagemaker.py | 378 +++----- src/autogluon/cloud/utils/aws_utils.py | 11 +- src/autogluon/cloud/utils/deserializers.py | 2 +- src/autogluon/cloud/utils/s3_utils.py | 4 +- src/autogluon/cloud/utils/sagemaker_api.py | 211 +++++ src/autogluon/cloud/utils/serializers.py | 2 +- tests/unittests/general/test_aws_utils.py | 8 +- .../general/test_foundation_model.py | 28 +- .../unittests/general/test_inference_modes.py | 111 ++- tests/unittests/general/test_sagemaker_ami.py | 73 +- tests/unittests/general/test_sagemaker_api.py | 155 ++++ .../general/test_tabular_endpoint.py | 51 +- tests/unittests/tabular/test_tabular.py | 14 +- .../tabular/test_tabular_fit_predict.py | 4 +- tests/unittests/timeseries/test_timeseries.py | 16 +- 32 files changed, 1551 insertions(+), 1445 deletions(-) delete mode 100644 src/autogluon/cloud/endpoint/endpoint.py delete mode 100644 src/autogluon/cloud/endpoint/sagemaker_endpoint.py create mode 100644 src/autogluon/cloud/utils/sagemaker_api.py create mode 100644 tests/unittests/general/test_sagemaker_api.py diff --git a/pyproject.toml b/pyproject.toml index 261d2185..c0aedbd4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,8 +43,9 @@ 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", + # SageMaker Python SDK v3 core layer (resources, shapes, session helpers). 2.1.0 adds TransformAmiVersion. + # We deliberately don't depend on the `sagemaker` meta-package, which also pulls in torch/mlflow via sagemaker-serve. + "sagemaker-core>=2.1.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..0153917d 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -2,35 +2,49 @@ 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 sagemaker.core.common_utils import sagemaker_timestamp, unique_name_from_base +from sagemaker.core.resources import Endpoint, EndpointConfig, Model 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, + create_serve_script_tarball, + repack_model_with_serving_code, + resolve_image_uri, + 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.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 from ..utils.misc import MostRecentInsertedOrderedDict -from ..utils.serializers import AutoGluonSerializationWrapper +from ..utils.sagemaker_api import ( + BATCH_PREDICT_OVERRIDE_KEYS, + DEPLOY_OVERRIDE_KEYS, + FIT_OVERRIDE_KEYS, + apply_overrides, + bind_core_session, + delete_endpoint, + invoke_endpoint, + normalize_tags, + normalize_vpc_config, + script_mode_environment, + to_request_tags, + validate_sagemaker_overrides, +) +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 +57,29 @@ logger = logging.getLogger(__name__) +SAGEMAKER_MODEL_SERVER_WORKERS = "SAGEMAKER_MODEL_SERVER_WORKERS" + + +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 { + "channel_name": channel_name, + "data_source": { + "s3_data_source": { + "s3_data_type": "S3Prefix", + "s3_uri": s3_uri, + "s3_data_distribution_type": "FullyReplicated", + } + }, + } + class SagemakerBackend(Backend): name = SAGEMAKER @@ -63,20 +100,22 @@ def __init__( **kwargs, ) - @property - def _realtime_predictor_cls(self) -> Predictor: - """Class used for realtime endpoint""" - return AutoGluonRealtimePredictor + def _realtime_serializer(self): + """Serializer used for realtime endpoint requests""" + return AutoGluonSerializer() - def _resolve_tags( + def _resolve_tags(self, extra_tags: Optional[List[Dict[str, str]]] = None) -> List[Dict[str, str]]: + """Tags for a created SageMaker resource, in sagemaker-core request format: default + extra + user tags.""" + return to_request_tags(build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=self.tags)) + + def initialize( self, - kwargs: Dict[str, Any], - extra_tags: Optional[List[Dict[str, str]]] = None, + role: Optional[str] = None, + vpc_config: Optional[Dict[str, List[str]]] = None, + kms_key: Optional[str] = None, + tags: Optional[Dict[str, str]] = None, + **kwargs, ) -> 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. Parameters @@ -84,6 +123,13 @@ def initialize(self, role: Optional[str] = None, **kwargs) -> None: role SageMaker execution role ARN. See :func:`autogluon.cloud.utils.aws_utils.resolve_execution_role` for the resolution order. + vpc_config + ``{"subnets": [...], "security_group_ids": [...]}`` applied to every training job, model and + transform job created by this backend. + kms_key + KMS key used to encrypt the S3 outputs and ML storage volumes of the created SageMaker resources. + tags + ``{"key": "value"}`` tags added to every created SageMaker resource. """ super().initialize(**kwargs) try: @@ -94,19 +140,20 @@ def initialize(self, role: Optional[str] = None, **kwargs) -> None: "or run `autogluon.cloud.bootstrap()` / `register()` to persist one." ) raise e + self.vpc_config = normalize_vpc_config(vpc_config) + self.kms_key = kms_key + self.tags = normalize_tags(tags) self.sagemaker_session = setup_sagemaker_session() - self.endpoint = None + self.endpoint_name: Optional[str] = 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), - ) + @property + def _boto_session(self): + boto_session = self.sagemaker_session.boto_session + bind_core_session(boto_session) + return boto_session def attach_job(self, job_name: str) -> None: """ @@ -118,7 +165,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: @@ -172,11 +219,13 @@ def fit( instance_type: str = "ml.m5.2xlarge", instance_count: Union[int, str] = 1, volume_size: int = 256, - custom_image_uri: Optional[str] = None, + 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, + environment: Optional[Dict[str, str]] = None, + use_spot_instances: bool = False, + max_wait: Optional[int] = None, + sagemaker_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: @@ -204,7 +253,7 @@ def fit( Training container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, AutoGluon Cloud creates one with a predictor- or model-specific prefix. @@ -215,18 +264,24 @@ def fit( volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training (default: 256). Must be large enough to store training data if File Mode is used (which is the default). + image_uri: Optional[str], default = None + Custom training container image. If None, the official AutoGluon DLC for ``framework_version`` is used. timeout: int, default = 24*60*60 Timeout in seconds for training. This timeout doesn't include time for pre-processing or launching up the training job. wait: bool, default = True 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 + environment: Optional[Dict[str, str]], default = None + Environment variables set in the training container. + use_spot_instances: bool, default = False + Whether to use managed spot training. + max_wait: Optional[int], default = None + Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires + ``use_spot_instances=True``. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw ``CreateTrainingJob`` request fields (sagemaker-core snake_case 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 +289,10 @@ def fit( """ if data_channels.get("train_data") is None: raise ValueError("`data_channels['train_data']` is required.") + _reject_local_mode(instance_type) + overrides = validate_sagemaker_overrides(sagemaker_overrides, FIT_OVERRIDE_KEYS) + if max_wait is not None and not use_spot_instances: + raise ValueError("`max_wait` requires `use_spot_instances=True`.") 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 = { @@ -241,9 +300,9 @@ def fit( for k, v in data_channels.items() if v is not None } - if custom_image_uri: + if image_uri: framework_version, py_version = None, None - logger.log(20, f"Training with custom_image_uri=={custom_image_uri}") + logger.log(20, f"Training with image_uri=={image_uri}") else: framework_version, py_version = parse_framework_version( framework_version, "training", minimum_version="0.6.0" @@ -251,7 +310,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 +320,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 +356,116 @@ 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, + source_dir=None, + 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) + stopping_condition: Dict[str, Any] = {"max_runtime_in_seconds": timeout} + if use_spot_instances: + stopping_condition["max_wait_time_in_seconds"] = max_wait or timeout + request: Dict[str, Any] = { + "training_job_name": job_name, + "role_arn": self.role_arn, + "algorithm_specification": { + "training_image": resolve_image_uri( + image_uri, framework_version, py_version, self._region, "training", instance_type + ), + "training_input_mode": "File", + }, + "hyper_parameters": training_script_hyperparameters( + entry_point=entry_point, submit_directory=code_uri, job_name=job_name, region=self._region + ), + "input_data_config": [_s3_channel(name, uri) for name, uri in inputs.items()], + "output_data_config": {"s3_output_path": self.cloud_output_path + "/model"}, + "resource_config": { + "instance_type": instance_type, + "instance_count": instance_count, + "volume_size_in_gb": volume_size, + }, + "stopping_condition": stopping_condition, + "profiler_config": {"disable_profiler": True}, + "tags": self._resolve_tags(extra_tags), + } + if environment: + request["environment"] = dict(environment) + if use_spot_instances: + request["enable_managed_spot_training"] = True + if self.vpc_config is not None: + request["vpc_config"] = self.vpc_config + if self.kms_key is not None: + request["output_data_config"]["kms_key_id"] = self.kms_key + request["resource_config"]["volume_kms_key_id"] = self.kms_key + request = apply_overrides(request, overrides, "create_training_job") + + self._fit_job.run(training_job_request=request, framework_version=framework_version, wait=wait) + + 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] = { + "model_name": model_name, + "primary_container": { + "image": image_uri, + "model_data_url": model_data, + "environment": container_environment, + }, + "execution_role_arn": self.role_arn, + "tags": tags, + } + if self.vpc_config is not None: + request["vpc_config"] = self.vpc_config + request = apply_overrides(request, overrides, "create_model") + logger.log(20, "Creating inference model...") + Model.create(**request, session=self._boto_session, region=self._region) + logger.log(20, "Inference model created successfully") + return request["model_name"] + + 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, + kms_key=self.kms_key, + ) - return dict(model_kwargs=model_kwargs, deploy_kwargs=deploy_kwargs) + @staticmethod + def _model_server_environment(environment: Optional[Dict[str, str]]) -> Dict[str, str]: + environment = dict(environment or {}) + if SAGEMAKER_MODEL_SERVER_WORKERS in environment and int(environment[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: + environment[SAGEMAKER_MODEL_SERVER_WORKERS] = "1" + return environment def deploy( self, @@ -363,11 +474,12 @@ def deploy( framework_version: str = "latest", instance_type: Optional[str] = "ml.m5.2xlarge", initial_instance_count: int = 1, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, volume_size: Optional[int] = None, wait: bool = True, - model_kwargs: Optional[Dict] = None, - deploy_kwargs: Optional[Dict] = None, + environment: Optional[Dict[str, str]] = None, + sagemaker_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 +488,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 ---------- @@ -392,12 +503,12 @@ def deploy( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. instance_type: str, default = 'ml.m5.2xlarge' Instance to be deployed for the endpoint initial_instance_count: int, default = 1, Initial number of instances to be deployed for the endpoint - custom_image_uri: Optional[str], default = None, + image_uri: Optional[str], default = None, Custom image to use to deploy endpoint with. If not specified, with use official DLC image: https://aws.github.io/deep-learning-containers/reference/available_images/#autogluon @@ -407,39 +518,43 @@ 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 + environment: Optional[Dict[str, str]], default = None + Environment variables set in the inference container. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw request fields (sagemaker-core snake_case 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 ``serverless_config`` + (``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`" ) + overrides = validate_sagemaker_overrides(sagemaker_overrides, DEPLOY_OVERRIDE_KEYS) 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: + if image_uri: framework_version, py_version = None, None - logger.log(20, f"Deploying with custom_image_uri=={custom_image_uri}") + logger.log(20, f"Deploying with image_uri=={image_uri}") else: framework_version, py_version = parse_framework_version( framework_version, "inference", minimum_version="0.6.0" @@ -461,158 +576,130 @@ 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 = self._model_server_environment(environment) 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=unique_name_from_base(endpoint_name), + model_data=model_data, + image_uri=resolve_image_uri( + image_uri, framework_version, py_version, self._region, "inference", instance_type + ), 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] = {"variant_name": "AllTraffic", "model_name": model_name} if inference_mode == "realtime": - mode_kwargs = instance_kwargs + variant["instance_type"] = instance_type + variant["initial_instance_count"] = initial_instance_count + if volume_size: + variant["volume_size_in_gb"] = volume_size + inference_ami_version = infer_sagemaker_ami_version(image_uri, instance_type, image_scope="inference") + if inference_ami_version is not None: + variant["inference_ami_version"] = inference_ami_version elif inference_mode == "serverless": preset = {"memory_size_in_mb": 4096, "max_concurrency": 5} - mode_kwargs = {"serverless_inference_config": ServerlessInferenceConfig(**{**preset, **user_config})} + variant["serverless_config"] = {**preset, **(inference_config or {})} else: raise ValueError(f"Unsupported inference_mode={inference_mode!r}") + variant = apply_overrides(variant, overrides, "production_variant") + + endpoint_config_request: Dict[str, Any] = { + "endpoint_config_name": endpoint_name, + "production_variants": [variant], + "tags": tags, + } + if self.kms_key is not None and inference_mode == "realtime": + endpoint_config_request["kms_key_id"] = self.kms_key + endpoint_config_request = apply_overrides(endpoint_config_request, overrides, "create_endpoint_config") + endpoint_request = apply_overrides( + { + "endpoint_name": endpoint_name, + "endpoint_config_name": endpoint_config_request["endpoint_config_name"], + "tags": tags, + }, + overrides, + "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) + boto_session = self._boto_session + EndpointConfig.create(**endpoint_config_request, session=boto_session, region=self._region) + endpoint = Endpoint.create(**endpoint_request, session=boto_session, region=self._region) + self.endpoint_name = endpoint_request["endpoint_name"] + if wait: + endpoint.wait_for_status("InService") 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/.""" - + """Create a minimal model.tar.gz containing the serve script + serving_utils/ under code/ and upload it.""" tarball_dir = tempfile.mkdtemp(prefix="ag_serve_") - tarball_path = os.path.join(tarball_dir, "model.tar.gz") - with tarfile.open(tarball_path, "w:gz") as tar: - tar.add(serve_script_path, arcname=f"code/{os.path.basename(serve_script_path)}") - tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils") + tarball_path = create_serve_script_tarball(serve_script_path, tarball_dir) s3_key = f"endpoints/{endpoint_name}/model/model.tar.gz" - s3_path = self._upload_predictor(tarball_path, s3_key) - return s3_path + return self._upload_predictor(tarball_path, s3_key) 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._boto_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}`") - - 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 + if not isinstance(endpoint, str): + raise ValueError(f"Please provide the endpoint name as a string, got {type(endpoint).__name__}.") + self.endpoint_name = endpoint + + 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 +780,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. @@ -765,22 +834,20 @@ def predict( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - custom_image_uri: Optional[str] = None, + 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, + environment: Optional[Dict[str, str]] = None, + sagemaker_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 ---------- @@ -797,14 +864,16 @@ def predict( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None - Name of the launched training job. + Name of the launched transform job. If None, AutoGluon Cloud creates one with a predictor- or model-specific prefix. instance_count: int, default = 1, Number of instances used to do batch transform. instance_type: str, default = 'ml.m5.2xlarge' Instance to be used for batch transform. + image_uri: Optional[str], default = None + Custom inference container image. If None, the official AutoGluon DLC is used. 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. @@ -819,16 +888,11 @@ def predict( 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. + environment: Optional[Dict[str, str]], default = None + Environment variables set in the inference container. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw request fields (sagemaker-core snake_case names) deep-merged over the requests built by + AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"create_transform_job"``. Returns ------- @@ -844,14 +908,13 @@ def predict( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, download=download, persist=persist, save_path=save_path, - model_kwargs=model_kwargs, - transformer_kwargs=transformer_kwargs, - transform_kwargs=transform_kwargs, + environment=environment, + sagemaker_overrides=sagemaker_overrides, original_features=self.original_features, ) @@ -862,79 +925,17 @@ def predict_proba( test_data: Union[str, pd.DataFrame], test_data_image_column: Optional[str] = None, include_predict: bool = True, - predictor_path: Optional[str] = None, - framework_version: str = "latest", - job_name: Optional[str] = None, - instance_type: str = "ml.m5.2xlarge", - 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, + **kwargs, ) -> 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. + Predict probabilities using SageMaker batch transform. + Accepts the same arguments as :meth:`predict`. Parameters ---------- - test_data: Union(str, pandas.DataFrame) - The test data to be inferenced. Can be a pandas.DataFrame, or a local path to a csv. - test_data_image_column: str, default = None - If test_data involves image modality, you must specify the column name corresponding to image paths. - The path MUST be an abspath include_predict: bool, default = True Whether to include predict result along with predict_proba results. This flag can save you time from making two calls to get both the prediction and the probability as batch inference involves noticeable overhead. - predictor_path: str - Path to the predictor tarball you want to use to predict. - Path can be both a local path or a S3 location. - If None, will use the most recent trained predictor trained with `fit()`. - framework_version: str, default = `latest` - Inference container version of autogluon. - If `latest`, will use the latest available container version. - If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. - job_name: str, default = None - Name of the launched training job. - If None, AutoGluon Cloud creates one with a predictor- or model-specific prefix. - instance_count: int, default = 1, - Number of instances used to do batch transform. - instance_type: str, default = 'ml.m5.2xlarge' - Instance to be used for batch transform. - 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. - Returns ------- @@ -947,20 +948,8 @@ def predict_proba( pred, pred_proba = self._predict( test_data=test_data, test_data_image_column=test_data_image_column, - predictor_path=predictor_path, - framework_version=framework_version, - job_name=job_name, - instance_type=instance_type, - 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, original_features=self.original_features, + **kwargs, ) if include_predict: @@ -1029,8 +1018,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) - channels = {c["ChannelName"]: c["DataSource"]["S3DataSource"]["S3Uri"] for c in desc["InputDataConfig"]} + channels = self._fit_job.get_input_channels() ag_args_uri = channels.get("ag_args") assert ag_args_uri is not None, ( f"Training job {job_name!r} has no `ag_args` input channel — cannot recover predictions_path." @@ -1128,15 +1116,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 +1154,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 +1166,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._boto_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: @@ -1221,44 +1212,44 @@ def _predict( job_name=None, instance_type="ml.m5.2xlarge", instance_count=1, - custom_image_uri=None, + image_uri=None, wait=True, download=True, persist=True, save_path=None, - model_kwargs=None, - transformer_kwargs=None, + environment=None, + sagemaker_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 = validate_sagemaker_overrides(sagemaker_overrides, BATCH_PREDICT_OVERRIDE_KEYS) if not predictor_path: predictor_path = self._fit_job.get_output_path() assert predictor_path, "No cloud trained model found." - if custom_image_uri: + if image_uri: framework_version, py_version = None, None - logger.log(20, f"Predicting with custom_image_uri=={custom_image_uri}") + logger.log(20, f"Predicting with image_uri=={image_uri}") else: framework_version, py_version = parse_framework_version( framework_version, "inference", minimum_version="0.6.0" ) 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,38 +1289,14 @@ 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 - - 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" + # 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", + ) if not wait: if download: @@ -1349,35 +1316,58 @@ def _predict( ) save_path = None - 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, + tags = self._resolve_tags() + model_name = self._create_model( + model_name=job_name, + model_data=model_data, + image_uri=resolve_image_uri( + image_uri, framework_version, py_version, self._region, "inference", 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, + environment=dict(environment or {}), + tags=tags, + overrides=overrides, ) + + transform_input: Dict[str, Any] = { + "data_source": {"s3_data_source": {"s3_data_type": "S3Prefix", "s3_uri": test_input}}, + "content_type": content_type, + } + if split_type is not None: + transform_input["split_type"] = split_type + transform_output: Dict[str, Any] = {"s3_output_path": output_path + "/results", "accept": accept} + if assemble_with is not None: + transform_output["assemble_with"] = assemble_with + transform_resources: Dict[str, Any] = {"instance_type": instance_type, "instance_count": instance_count} + transform_ami_version = infer_sagemaker_ami_version(image_uri, instance_type, image_scope="transform") + if transform_ami_version is not None: + transform_resources["transform_ami_version"] = transform_ami_version + if self.kms_key is not None: + transform_output["kms_key_id"] = self.kms_key + transform_resources["volume_kms_key_id"] = self.kms_key + request = { + "transform_job_name": job_name, + "model_name": model_name, + "transform_input": transform_input, + "transform_output": transform_output, + "transform_resources": transform_resources, + "batch_strategy": batch_strategy, + # Maximum size in MB of a single request to the container; larger inputs are split into multiple batches. + "max_payload_in_mb": 6, + # The maximum number of HTTP requests made to each individual transform container at one time. + "max_concurrent_transforms": 1, + "tags": tags, + } + request = apply_overrides(request, overrides, "create_transform_job") + + batch_transform_job = SageMakerBatchTransformationJob(session=self.sagemaker_session) + batch_transform_job.run(transform_job_request=request, wait=wait) self._batch_transform_jobs[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") + accept = request["transform_output"].get("accept") if accept == "application/x-parquet": results = pd.read_parquet(results_path) elif accept == "text/csv": @@ -1399,10 +1389,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 +1396,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..4b7b94dd 100644 --- a/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py +++ b/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py @@ -1,6 +1,6 @@ import logging import os -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, Optional, Union import pandas as pd @@ -24,23 +24,15 @@ def fit( data_channels: Dict[str, Optional[Union[str, pd.DataFrame]]], id_column: str, timestamp_column: str, - framework_version: str = "latest", - job_name: Optional[str] = None, - instance_type: str = "ml.m5.2xlarge", - instance_count: int = 1, 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, extra_ag_args: Optional[Dict[str, Any]] = None, - extra_tags: Optional[List[Dict[str, str]]] = None, + **kwargs, ) -> None: """Fit a TimeSeriesPredictor in SageMaker. ``id_column`` / ``timestamp_column`` are forwarded to the training script via ``ag_args.json``. ``known_covariates`` (if present in ``data_channels``) is only honored when - ``extra_ag_args["predict_after_fit"]`` is True. + ``extra_ag_args["predict_after_fit"]`` is True. Other arguments are forwarded to ``SagemakerBackend.fit()``. """ extra_ag_args = {**(extra_ag_args or {}), "id_column": id_column, "timestamp_column": timestamp_column} if data_channels.get("known_covariates") is not None and not extra_ag_args.get("predict_after_fit", False): @@ -56,17 +48,9 @@ def fit( predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, data_channels=data_channels, - framework_version=framework_version, - job_name=job_name, - instance_type=instance_type, - instance_count=instance_count, volume_size=volume_size, - custom_image_uri=custom_image_uri, - wait=wait, - autogluon_sagemaker_estimator_kwargs=autogluon_sagemaker_estimator_kwargs, - fit_kwargs=fit_kwargs, extra_ag_args=extra_ag_args, - extra_tags=extra_tags, + **kwargs, ) def predict_real_time( @@ -138,8 +122,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 +158,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..47c6c4bb 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).boto_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..fb72e0e0 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).boto_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..4a76ad60 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -1,16 +1,13 @@ import logging from abc import abstractmethod -from typing import Dict, Optional, Union +from typing import Any, Dict, Optional, Union -import sagemaker +from sagemaker.core.resources import Model, TrainingJob, TransformJob +from sagemaker.core.utils.exceptions import FailedStatusError -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.sagemaker_api import bind_core_session from .remote_job import RemoteJob logger = logging.getLogger(__name__) @@ -18,15 +15,13 @@ class SageMakerJob(RemoteJob): 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): + """Return the sagemaker-core resource describing the job.""" + raise NotImplementedError + @abstractmethod def _get_job_status(self): raise NotImplementedError @@ -66,14 +66,18 @@ def _get_output_path(self): def _get_hyperparameters(self): raise NotImplementedError + @property + def _boto_session(self): + boto_session = self.session.boto_session + bind_core_session(boto_session) + return boto_session + @property def job_name(self): return self._job_name @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 +93,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 +120,17 @@ def get_hyperparameters(self) -> Dict[str, Union[int, str]]: """ return self._get_hyperparameters() + def wait(self, logs: bool = True) -> None: + """Block until the job reaches a terminal state, streaming its CloudWatch logs if ``logs`` is True. + + Does not raise if the job fails; check :meth:`get_job_status` afterwards. + """ + assert self.job_name, "The job has not been started" + try: + self._describe().wait(logs=logs) + except FailedStatusError as e: + logger.error(f"SageMaker job {self.job_name} did not complete successfully: {e}") + def __getstate__(self): state_dict = self.__dict__.copy() state_dict["session"] = None @@ -135,15 +147,14 @@ def __init__(self, **kwargs): 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(logs=True) return obj @property @@ -160,69 +171,50 @@ def info(self): ) return info + def _describe(self) -> TrainingJob: + return TrainingJob.get( + training_job_name=self.job_name, + session=self._boto_session, + region=self.session.boto_region_name, + ) + def _get_job_status(self): - return self.session.describe_training_job(self.job_name)["TrainingJobStatus"] + return self._describe().training_job_status 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().model_artifacts.s3_model_artifacts def _get_hyperparameters(self): if self.job_name: - return self.session.describe_training_job(self.job_name)["HyperParameters"] + return self._describe().hyper_parameters return None + def get_input_channels(self) -> Dict[str, str]: + """Map each input channel name of the training job to its S3 URI.""" + return { + channel.channel_name: channel.data_source.s3_data_source.s3_uri + for channel in self._describe().input_data_config + } + 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 ``TrainingJob.create`` request and optionally wait for it to finish.""" + job_name = training_job_request["training_job_name"] logger.log(20, f"Start sagemaker training job `{job_name}`") try: - sagemaker_estimator.fit(inputs=inputs, wait=wait, job_name=job_name, **kwargs) + training_job = TrainingJob.create( + **training_job_request, + session=self._boto_session, + region=self.session.boto_region_name, + ) 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: + training_job.wait(logs=True) except Exception as e: logger.error(f"Training failed. Please check sagemaker console training jobs {job_name} for details.") raise e @@ -234,7 +226,7 @@ def __init__(self, **kwargs): self._output_filename = "" @classmethod - def attach(cls, job_name): + def attach(cls, job_name, session=None): raise NotImplementedError def info(self): @@ -246,105 +238,55 @@ def info(self): ) return info + def _describe(self) -> TransformJob: + return TransformJob.get( + transform_job_name=self.job_name, + session=self._boto_session, + region=self.session.boto_region_name, + ) + def _get_job_status(self): - return self.session.describe_transform_job(self.job_name)["TransformJobStatus"] + return self._describe().transform_job_status 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().transform_output.s3_output_path + "/" + self._output_filename + + def _delete_model(self, model_name: str) -> None: + bind_core_session(self.session.boto_session) + Model(model_name=model_name).delete() 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], + 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 ``TransformJob.create`` request. + The SageMaker model referenced by the request 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["transform_job_name"] + model_name = transform_job_request["model_name"] try: logger.log(20, "Transforming") - transformer.transform( - test_input, - job_name=job_name, - split_type=split_type, - content_type=content_type, - wait=wait, - **kwargs, + transform_job = TransformJob.create( + **transform_job_request, + session=self._boto_session, + region=self.session.boto_region_name, ) 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: + transform_job.wait(logs=True) logger.log(20, "Transform done") except Exception as e: - transformer.delete_model() + self._delete_model(model_name) raise e - self._output_filename = test_input.split("/")[-1] + ".out" + input_uri = transform_job_request["transform_input"]["data_source"]["s3_data_source"]["s3_uri"] + 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..3a0865d5 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 @@ -83,6 +84,9 @@ def __init__( role: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, model_artifact_uri: Optional[str] = None, + vpc_config: Optional[Dict[str, List[str]]] = None, + kms_key: Optional[str] = None, + tags: Optional[Dict[str, str]] = None, backend: Literal["sagemaker"] = "sagemaker", ): """ @@ -104,12 +108,18 @@ 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 S3 URI of a pre-bundled ``model.tar.gz`` produced by :meth:`cache_model_artifact`. When set, deploys skip the runtime HuggingFace download and load weights from the bundled artifact. + vpc_config + VPC for the created SageMaker jobs and models, as ``{"subnets": [...], "security_group_ids": [...]}``. + kms_key + KMS key ID/ARN used to encrypt S3 outputs and ML storage volumes of the created SageMaker resources. + tags + Tags added to every SageMaker resource created for this model, e.g. ``{"team": "forecasting"}``. backend Cloud backend to use. """ @@ -118,6 +128,7 @@ def __init__( self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend) self._config = get_model_config(model_id) self._hyperparameter_overrides = hyperparameters or {} + self._infra_settings = {"vpc_config": vpc_config, "kms_key": kms_key, "tags": tags} self._tmpdir = tempfile.TemporaryDirectory(prefix="ag_fm_") backend_name = self._backend_map.get(backend) @@ -133,6 +144,7 @@ def __init__( predictor_type=self._predictor_type, resource_prefix=f"ag-cloud-{self.model_id}", role=role, + **self._infra_settings, ) def _get_hyperparameters( @@ -189,11 +201,11 @@ def _deploy_backend( endpoint_name: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, framework_version: str = "latest", - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - **backend_kwargs, + **kwargs, ) -> None: """Shared deploy logic. Subclasses call this then wrap the endpoint.""" if inference_mode == "serverless" and instance_type is not None: @@ -218,27 +230,24 @@ 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, endpoint_name=endpoint_name, framework_version=framework_version, instance_type=instance_type, - custom_image_uri=custom_image_uri, + image_uri=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, repack=False, extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], - **backend_kwargs, + **kwargs, ) - assert self._backend.endpoint is not None + assert self._backend.endpoint_name is not None def fit( self, @@ -356,6 +365,7 @@ def cache_model_artifact(self, cache_path: str, *, overwrite: bool = False) -> S model_artifact_uri=cache_key, cloud_output_path=self.cloud_output_path, role=self._backend.role_arn, + **self._infra_settings, ) def to_dict(self) -> Dict[str, Any]: @@ -410,17 +420,20 @@ 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, endpoint_name: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, framework_version: str = "latest", - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - **backend_kwargs, + environment: Optional[Dict[str, str]] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + **kwargs, ) -> TimeSeriesEndpoint: """ Deploy model to an inference endpoint. @@ -436,7 +449,7 @@ def deploy( Model hyperparameters for inference. Overrides values passed to the constructor. framework_version Container framework version. If 'latest', uses the most recent available. - custom_image_uri + image_uri Custom Docker image URI for the inference container. wait Whether to block until the endpoint is ready. @@ -444,25 +457,31 @@ 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``). - **backend_kwargs - Backend-specific arguments (e.g., initial_instance_count, volume_size, - model_kwargs, deploy_kwargs). + Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). + environment + Environment variables set in the inference container. + sagemaker_overrides + Raw SageMaker request fields deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + ``"create_model"``, ``"production_variant"``, ``"create_endpoint_config"``, ``"create_endpoint"``. See + :meth:`autogluon.cloud.TabularCloudPredictor.deploy`. + **kwargs + Additional deployment arguments (``initial_instance_count``, ``volume_size``). """ self._deploy_backend( instance_type=instance_type, endpoint_name=endpoint_name, hyperparameters=hyperparameters, framework_version=framework_version, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, inference_mode=inference_mode, inference_config=inference_config, - **backend_kwargs, + environment=environment, + sagemaker_overrides=sagemaker_overrides, + **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 +508,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], @@ -503,9 +523,9 @@ def predict( hyperparameters: Optional[Dict[str, Any]] = None, instance_type: Optional[str] = None, framework_version: str = "latest", - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - **backend_kwargs, + **kwargs, ) -> Union[pd.DataFrame, JobPredictionFuture]: """ Run batch prediction for time series. @@ -545,15 +565,15 @@ def predict( Instance type for the prediction job. If None, uses registry default. framework_version Container framework version. - custom_image_uri + image_uri Custom Docker image URI for the container. wait If True, block and return a DataFrame. If False, return a :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). + **kwargs + Additional job arguments accepted by :meth:`autogluon.cloud.TimeSeriesCloudPredictor.fit` (e.g. + ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``, ``sagemaker_overrides``). Returns ------- @@ -588,11 +608,11 @@ def predict( timestamp_column=timestamp_column, framework_version=framework_version, instance_type=instance_type, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, extra_ag_args=extra_ag_args, extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], - **backend_kwargs, + **kwargs, ) if not wait: @@ -623,17 +643,20 @@ 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, endpoint_name: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, framework_version: str = "latest", - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, inference_mode: Literal["realtime"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - **backend_kwargs, + environment: Optional[Dict[str, str]] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + **kwargs, ) -> TabularEndpoint: """Deploy the tabular foundation model to an inference endpoint. @@ -658,13 +681,15 @@ def deploy( endpoint_name=endpoint_name, hyperparameters=hyperparameters, framework_version=framework_version, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, inference_mode="realtime", - **backend_kwargs, + environment=environment, + sagemaker_overrides=sagemaker_overrides, + **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 +719,7 @@ def _load_results( else: return pred_proba + @reject_legacy_kwargs def predict( self, test_data: Union[str, Path, pd.DataFrame], @@ -704,9 +730,9 @@ def predict( hyperparameters: Optional[Dict[str, Any]] = None, instance_type: Optional[str] = None, framework_version: str = "latest", - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - **backend_kwargs, + **kwargs, ) -> Union[pd.Series, JobPredictionFuture]: """ Run batch prediction for tabular tasks. @@ -732,13 +758,14 @@ def predict( Instance type for the prediction job. If None, uses registry default. framework_version Container framework version. - custom_image_uri + image_uri Custom Docker image URI for the container. wait 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). + **kwargs + Additional job arguments accepted by :meth:`autogluon.cloud.TabularCloudPredictor.fit` (e.g. + ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``, ``sagemaker_overrides``). Returns ------- @@ -754,9 +781,9 @@ def predict( hyperparameters=hyperparameters, instance_type=instance_type, framework_version=framework_version, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - **backend_kwargs, + **kwargs, ) if not wait: return JobPredictionFuture( @@ -766,6 +793,7 @@ def predict( pred, _ = result return pred + @reject_legacy_kwargs def predict_proba( self, test_data: Union[str, Path, pd.DataFrame], @@ -777,9 +805,9 @@ def predict_proba( hyperparameters: Optional[Dict[str, Any]] = None, instance_type: Optional[str] = None, framework_version: str = "latest", - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - **backend_kwargs, + **kwargs, ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series], JobPredictionFuture]: """ Run batch prediction returning class probabilities. @@ -807,12 +835,13 @@ def predict_proba( Instance type for the prediction job. If None, uses registry default. framework_version Container framework version. - custom_image_uri + image_uri Custom Docker image URI for the container. 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). + **kwargs + Additional job arguments accepted by :meth:`autogluon.cloud.TabularCloudPredictor.fit` (e.g. + ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``, ``sagemaker_overrides``). Returns ------- @@ -831,7 +860,7 @@ def predict_proba( extra_ag_args: Dict[str, Any] = {"predict_after_fit": True, "save_predictor": False} if predictions_path is not None: extra_ag_args["predictions_path"] = predictions_path - backend_kwargs["leaderboard"] = False + kwargs["leaderboard"] = False self._backend.fit( predictor_init_args=self._build_predictor_init_args(label=label), @@ -839,11 +868,11 @@ def predict_proba( data_channels={"train_data": train_data, "tuning_data": tuning_data, "test_data": test_data}, framework_version=framework_version, instance_type=instance_type, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, extra_ag_args=extra_ag_args, extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], - **backend_kwargs, + **kwargs, ) if not wait: diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index ab466771..a9980b87 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -7,7 +7,7 @@ from abc import ABC, abstractmethod from datetime import datetime from pathlib import Path -from typing import Any, Dict, Literal, Optional, Tuple, Union +from typing import Any, Dict, List, Literal, Optional, Tuple, Union import boto3 import pandas as pd @@ -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__) @@ -38,6 +38,9 @@ def __init__( cloud_output_path: Optional[str] = None, backend: str = SAGEMAKER, role: Optional[str] = None, + vpc_config: Optional[Dict[str, List[str]]] = None, + kms_key: Optional[str] = None, + tags: Optional[Dict[str, str]] = None, verbosity: int = 2, ) -> None: """ @@ -66,7 +69,15 @@ 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. + vpc_config: Optional[Dict[str, List[str]]], default = None + VPC to run training jobs, models and batch transform jobs in, as + ``{"subnets": ["subnet-..."], "security_group_ids": ["sg-..."]}``. + kms_key: Optional[str], default = None + KMS key ID/ARN used to encrypt S3 outputs and ML storage volumes of all created SageMaker resources. + Note that SageMaker rejects volume KMS keys for instance types with local NVMe storage (e.g. ``ml.g5``). + tags: Optional[Dict[str, str]], default = None + Tags added to every SageMaker resource created by this predictor, e.g. ``{"team": "forecasting"}``. 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). @@ -88,6 +99,9 @@ def __init__( cloud_output_path=self.cloud_output_path, predictor_type=self.predictor_type, role=role, + vpc_config=vpc_config, + kms_key=kms_key, + tags=tags, ) @property @@ -110,9 +124,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 +173,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, @@ -175,10 +188,13 @@ def fit( instance_type: str = "ml.m5.2xlarge", instance_count: Union[int, str] = "auto", volume_size: int = 256, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, timeout: int = 24 * 60 * 60, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + environment: Optional[Dict[str, str]] = None, + use_spot_instances: bool = False, + max_wait: Optional[int] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> CloudPredictor: """ @@ -201,7 +217,7 @@ def fit( Training container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with a predictor-specific prefix. @@ -213,21 +229,26 @@ def fit( volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training. Must be large enough to store training data if File Mode is used (which is the default). + image_uri: Optional[str], default = None + Custom training container image. If None, the official AutoGluon DLC for ``framework_version`` is used. timeout: int, default = 24*60*60 Timeout in seconds for training. This timeout doesn't include time for pre-processing or launching up the training job. wait: bool, default = True 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 + environment: Optional[Dict[str, str]], default = None + Environment variables set in the training container. + use_spot_instances: bool, default = False + Whether to train on managed spot instances. + max_wait: Optional[int], default = None + Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires + ``use_spot_instances=True``. + sagemaker_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 snake_case (as in ``sagemaker.core.shapes``), which are deep-merged over the request + built by AutoGluon-Cloud, e.g. ``{"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}``. Returns ------- `CloudPredictor` object. Returns self. @@ -235,11 +256,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 +278,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, @@ -270,10 +289,14 @@ def fit( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - custom_image_uri=custom_image_uri, + image_uri=image_uri, timeout=timeout, wait=wait, - **backend_kwargs, + environment=environment, + use_spot_instances=use_spot_instances, + max_wait=max_wait, + sagemaker_overrides=sagemaker_overrides, + extra_ag_args=extra_ag_args, ) return self @@ -372,6 +395,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, @@ -379,12 +403,13 @@ def deploy( framework_version: str = "latest", instance_type: Optional[str] = None, initial_instance_count: int = 1, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, volume_size: Optional[int] = None, wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - backend_kwargs: Optional[Dict] = None, + environment: Optional[Dict[str, str]] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> None: """ Deploy a predictor to an inference endpoint. @@ -402,14 +427,14 @@ def deploy( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. instance_type: Optional[str], default = None Instance to be deployed for the endpoint. Defaults to ``ml.m5.2xlarge``. Must be ``None`` when ``inference_mode="serverless"``. initial_instance_count: int, default = 1, Initial number of instances to be deployed for the endpoint. Ignored when ``inference_mode="serverless"``. - custom_image_uri: Optional[str], default = None, + image_uri: Optional[str], default = None, Custom image to use to deploy endpoint with. If not specified, with use official DLC image: https://github.com/aws/deep-learning-containers/blob/master/available_images.md#autogluon-inference-containers @@ -423,57 +448,54 @@ 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``). + environment: Optional[Dict[str, str]], default = None + Environment variables set in the inference container. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case + (as in ``sagemaker.core.shapes``), 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": {"model_data_download_timeout_in_seconds": 1200}}``. """ 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, framework_version=framework_version, instance_type=instance_type, initial_instance_count=initial_instance_count, - custom_image_uri=custom_image_uri, + image_uri=image_uri, volume_size=volume_size, wait=wait, inference_mode=inference_mode, inference_config=inference_config, - **backend_kwargs, + environment=environment, + sagemaker_overrides=sagemaker_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 +572,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], @@ -559,9 +582,13 @@ def predict( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + download: bool = True, + persist: bool = True, + save_path: Optional[str] = None, + environment: Optional[Dict[str, str]] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Batch inference. @@ -583,9 +610,9 @@ def predict( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None - Name of the launched training job. + Name of the launched batch transform job. If None, CloudPredictor creates one with a predictor-specific prefix. instance_count: int, default = 1, Number of instances used to do batch transform. @@ -594,30 +621,24 @@ 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. + 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. + environment: Optional[Dict[str, str]], default = None + Environment variables set in the inference container. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case + (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + ``"create_model"`` and ``"create_transform_job"``, e.g. + ``{"create_transform_job": {"batch_strategy": "SingleRecord", "max_payload_in_mb": 20}}``. Returns ------- @@ -625,9 +646,6 @@ def predict( Predict results in Series if `download` is True None if `download` 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, @@ -636,11 +654,16 @@ def predict( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - **backend_kwargs, + download=download, + persist=persist, + save_path=save_path, + environment=environment, + sagemaker_overrides=sagemaker_overrides, ) + @reject_legacy_kwargs def predict_proba( self, test_data: Union[str, pd.DataFrame], @@ -651,9 +674,13 @@ def predict_proba( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + download: bool = True, + persist: bool = True, + save_path: Optional[str] = None, + environment: Optional[Dict[str, str]] = None, + sagemaker_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 @@ -678,9 +705,9 @@ def predict_proba( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None - Name of the launched training job. + Name of the launched batch transform job. If None, CloudPredictor creates one with a predictor-specific prefix. instance_count: int, default = 1, Number of instances used to do batch transform. @@ -689,30 +716,24 @@ 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. + 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. + environment: Optional[Dict[str, str]], default = None + Environment variables set in the inference container. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case + (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + ``"create_model"`` and ``"create_transform_job"``, e.g. + ``{"create_transform_job": {"batch_strategy": "SingleRecord", "max_payload_in_mb": 20}}``. Returns ------- @@ -722,9 +743,6 @@ def predict_proba( 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, @@ -734,9 +752,13 @@ def predict_proba( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - **backend_kwargs, + download=download, + persist=persist, + save_path=save_path, + environment=environment, + sagemaker_overrides=sagemaker_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..694cd1bf 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], @@ -49,10 +51,13 @@ def fit_predict( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 256, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, predictions_path: Optional[str] = None, - backend_kwargs: Optional[Dict] = None, + environment: Optional[Dict[str, str]] = None, + use_spot_instances: bool = False, + max_wait: Optional[int] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Fit and predict in a single SageMaker training job. @@ -78,7 +83,7 @@ def fit_predict( Whether to include the leaderboard in the output artifact. framework_version: str, default = `latest` Training container version of autogluon. If `latest`, will use the latest available container version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-tabular``. instance_type: str, default = 'ml.m5.2xlarge' @@ -87,7 +92,7 @@ def fit_predict( Number of instances used to fit the predictor. volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training. - custom_image_uri: Optional[str], default = None + image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True Whether the call should wait until the job completes. @@ -95,8 +100,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()``. + environment, use_spot_instances, max_wait, sagemaker_overrides: + Same as in :meth:`fit`. Returns ------- @@ -116,16 +121,20 @@ def fit_predict( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, predictions_path=predictions_path, - backend_kwargs=backend_kwargs, + environment=environment, + use_spot_instances=use_spot_instances, + max_wait=max_wait, + sagemaker_overrides=sagemaker_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], @@ -140,10 +149,13 @@ def fit_predict_proba( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 256, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, predictions_path: Optional[str] = None, - backend_kwargs: Optional[Dict] = None, + environment: Optional[Dict[str, str]] = None, + use_spot_instances: bool = False, + max_wait: Optional[int] = None, + sagemaker_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. @@ -167,7 +179,7 @@ def fit_predict_proba( leaderboard: bool, default = True Whether to include the leaderboard in the output artifact. framework_version: str, default = `latest` - Training container version of autogluon. If `custom_image_uri` is set, this argument is ignored. + Training container version of autogluon. If `image_uri` is set, this argument is ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-tabular``. instance_type: str, default = 'ml.m5.2xlarge' @@ -176,15 +188,15 @@ def fit_predict_proba( Number of instances used to fit the predictor. volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training. - custom_image_uri: Optional[str], default = None + image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True Whether the call should wait until the job completes. 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()``. + environment, use_spot_instances, max_wait, sagemaker_overrides: + Same as in :meth:`fit`. Returns ------- @@ -193,12 +205,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, @@ -211,9 +220,13 @@ def fit_predict_proba( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - backend_kwargs=backend_kwargs, + environment=environment, + use_spot_instances=use_spot_instances, + max_wait=max_wait, + sagemaker_overrides=sagemaker_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..ada55bb8 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, @@ -49,10 +51,14 @@ def fit( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 100, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, known_covariates: Optional[Union[str, Path, pd.DataFrame]] = None, + environment: Optional[Dict[str, str]] = None, + use_spot_instances: bool = False, + max_wait: Optional[int] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + **kwargs, ) -> TimeSeriesCloudPredictor: """ Fit the predictor in a SageMaker training job. @@ -88,7 +94,7 @@ def fit( Training container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-timeseries``. @@ -99,21 +105,21 @@ def fit( volume_size: int, default = 100 Size in GB of the EBS volume to use for storing input data during training. Must be large enough to store training data if File Mode is used (which is the default). - custom_image_uri: Optional[str], default = None + image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True 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 + environment: Optional[Dict[str, str]], default = None + Environment variables set in the training container. + use_spot_instances: bool, default = False + Whether to train on managed spot instances. + max_wait: Optional[int], default = None + Maximum seconds to wait for spot capacity plus training time. Requires ``use_spot_instances=True``. + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw ``CreateTrainingJob`` request fields under the ``"create_training_job"`` key. See + :meth:`TabularCloudPredictor.fit` for details. Returns ------- @@ -122,8 +128,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`, 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 = { @@ -141,7 +149,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, @@ -153,9 +160,13 @@ def fit( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - **backend_kwargs, + environment=environment, + use_spot_instances=use_spot_instances, + max_wait=max_wait, + sagemaker_overrides=sagemaker_overrides, + extra_ag_args=extra_ag_args, ) return self @@ -210,6 +221,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], @@ -220,9 +232,13 @@ def predict( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + download: bool = True, + persist: bool = True, + save_path: Optional[str] = None, + environment: Optional[Dict[str, str]] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ Predict using SageMaker batch transform. @@ -250,7 +266,7 @@ def predict( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-timeseries``. @@ -261,34 +277,11 @@ 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. + image_uri: Optional[str], default = None + Custom inference container image. If set, ``framework_version`` is ignored. + download, persist, save_path, environment, sagemaker_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, @@ -298,9 +291,13 @@ def predict( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - **backend_kwargs, + download=download, + persist=persist, + save_path=save_path, + environment=environment, + sagemaker_overrides=sagemaker_overrides, ) def predict_proba( @@ -312,6 +309,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], @@ -328,9 +326,12 @@ def fit_predict( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 100, - custom_image_uri: Optional[str] = None, + image_uri: Optional[str] = None, wait: bool = True, - backend_kwargs: Optional[Dict] = None, + environment: Optional[Dict[str, str]] = None, + use_spot_instances: bool = False, + max_wait: Optional[int] = None, + sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ Fit and predict in a single SageMaker training job. @@ -370,7 +371,7 @@ def fit_predict( names ``item_id`` and ``timestamp``, regardless of the ``id_column`` / ``timestamp_column`` passed in. framework_version: str, default = `latest` Training container version of autogluon. If `latest`, will use the latest available container version. - If `custom_image_uri` is set, this argument will be ignored. + If `image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-timeseries``. instance_type: str, default = 'ml.m5.2xlarge' @@ -379,26 +380,21 @@ def fit_predict( Number of instances used to fit the predictor. volume_size: int, default = 100 Size in GB of the EBS volume to use for storing input data during training. - custom_image_uri: Optional[str], default = None + image_uri: Optional[str], default = None 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()``. + environment, use_spot_instances, max_wait, sagemaker_overrides: + 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, @@ -413,9 +409,13 @@ def fit_predict( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - custom_image_uri=custom_image_uri, + image_uri=image_uri, wait=wait, - backend_kwargs=backend_kwargs, + environment=environment, + use_spot_instances=use_spot_instances, + max_wait=max_wait, + sagemaker_overrides=sagemaker_overrides, + extra_ag_args=extra_ag_args, ) if not wait: diff --git a/src/autogluon/cloud/scripts/sagemaker_scripts/train.py b/src/autogluon/cloud/scripts/sagemaker_scripts/train.py index f4efaac3..bdba3ea5 100644 --- a/src/autogluon/cloud/scripts/sagemaker_scripts/train.py +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/train.py @@ -81,7 +81,7 @@ def prepare_data(data_file, predictor_type, ag_args, static_features_df=None): print(f"Args: {args}") - # See SageMaker-specific environment variables: https://sagemaker.readthedocs.io/en/v2/overview.html#prepare-a-training-script + # See SageMaker-specific environment variables: https://github.com/aws/sagemaker-training-toolkit/blob/master/ENVIRONMENT_VARIABLES.md os.makedirs(args.output_data_dir, mode=0o777, exist_ok=True) ag_args_file = get_input_path(args.ag_args) diff --git a/src/autogluon/cloud/utils/ag_sagemaker.py b/src/autogluon/cloud/utils/ag_sagemaker.py index 8a247b76..0be0d2b5 100644 --- a/src/autogluon/cloud/utils/ag_sagemaker.py +++ b/src/autogluon/cloud/utils/ag_sagemaker.py @@ -1,271 +1,125 @@ -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) - - predictor_cls = predict_wrapper - - 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, - ) - - @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 - - -# 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, - ) - - 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 +"""Packaging helpers for AutoGluon training and serving code on SageMaker. +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). +""" -class AutoGluonRepackInferenceModel(AutoGluonSagemakerInferenceModel): +import json +import os +import shutil +import tarfile +import tempfile +from contextlib import contextmanager +from typing import Dict, Iterator, Optional + +from sagemaker.core.common_utils import repack_model + +from .dlc_utils import retrieve_image_uri + +SOURCE_DIR_TARBALL_NAME = "sourcedir.tar.gz" + + +def resolve_image_uri( + image_uri: Optional[str], + framework_version: Optional[str], + py_version: Optional[str], + region: str, + image_scope: str, + instance_type: str, +) -> str: + """Return ``image_uri`` if set, otherwise the official AutoGluon DLC for the given version and instance.""" + if image_uri: + return image_uri + return retrieve_image_uri( + framework_version=framework_version, + region=region, + image_scope=image_scope, + instance_type=instance_type, + py_version=py_version, + ) + + +def upload_training_code(entry_point: str, source_dir: Optional[str], sagemaker_session, s3_uri_prefix: str) -> str: + """Bundle the training entry point (or ``source_dir`` containing it) as ``sourcedir.tar.gz`` and upload it. + + Returns the S3 URI of the uploaded tarball. """ - Custom implementation to force repack of inference code into model artifacts + from autogluon.common.utils.s3_utils import s3_path_to_bucket_prefix + + 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: + if source_dir: + for name in os.listdir(source_dir): + tar.add(os.path.join(source_dir, name), arcname=name) + else: + 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 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. + + Values are JSON-encoded, matching what SageMaker SDK v2 sent; the toolkit JSON-decodes them. """ - - 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, - ) - - -class AutoGluonNonRepackInferenceModel(AutoGluonSagemakerInferenceModel): + 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), + } + + +@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 ``model_data`` tarball with ``entry_point`` + ``serving_utils/`` and upload it. + + Returns ``repacked_model_uri``. """ - 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. - """ - - 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, + with staged_serving_code(entry_point) as code_dir: + repack_model( + inference_script=entry_point, + source_directory=code_dir, + dependencies=[], + model_uri=model_data, + repacked_model_uri=repacked_model_uri, + sagemaker_session=sagemaker_session, + kms_key=kms_key, ) + return repacked_model_uri -# 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) - +def create_serve_script_tarball(entry_point: str, output_dir: str) -> str: + """Create a minimal ``model.tar.gz`` that only contains the serving code under ``code/``.""" + from ..scripts import ScriptManager # deferred: importing scripts pulls in the backend package -# 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) + tarball_path = os.path.join(output_dir, "model.tar.gz") + with tarfile.open(tarball_path, "w:gz") as tar: + tar.add(entry_point, arcname=f"code/{os.path.basename(entry_point)}") + tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils") + return tarball_path diff --git a/src/autogluon/cloud/utils/aws_utils.py b/src/autogluon/cloud/utils/aws_utils.py index 8a889188..b2ba81b5 100644 --- a/src/autogluon/cloud/utils/aws_utils.py +++ b/src/autogluon/cloud/utils/aws_utils.py @@ -2,8 +2,9 @@ from typing import Optional import boto3 -import sagemaker from botocore.config import Config +from sagemaker.core.common_utils import sagemaker_timestamp +from sagemaker.core.helper.session_helper import Session, get_execution_role from autogluon.common.utils.s3_utils import is_s3_url @@ -35,7 +36,7 @@ def resolve_execution_role(role: Optional[str], backend_name: str) -> str: 1. ``role`` argument if provided. 2. ``role_arn`` from ``~/.autogluon/cloud.yaml`` under the matching backend slot. - 3. ``sagemaker.get_execution_role()``. + 3. ``sagemaker.core.helper.session_helper.get_execution_role()`` (the caller's own role, e.g. on SageMaker). """ if role: return role @@ -45,7 +46,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() def resolve_cloud_output_path(path: Optional[str], backend_name: str) -> Optional[str]: @@ -81,7 +82,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}.") @@ -181,4 +182,4 @@ def setup_sagemaker_session( "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 Session(boto_session=boto_session, sagemaker_client=sm_boto) diff --git a/src/autogluon/cloud/utils/deserializers.py b/src/autogluon/cloud/utils/deserializers.py index 4bdbec45..87f6fa8b 100644 --- a/src/autogluon/cloud/utils/deserializers.py +++ b/src/autogluon/cloud/utils/deserializers.py @@ -2,7 +2,7 @@ from abc import ABC, abstractmethod import pandas as pd -from sagemaker.deserializers import SimpleBaseDeserializer +from sagemaker.core.deserializers import SimpleBaseDeserializer class PandasDeserializeStrategy(ABC): diff --git a/src/autogluon/cloud/utils/s3_utils.py b/src/autogluon/cloud/utils/s3_utils.py index ef01406a..28d90d1b 100644 --- a/src/autogluon/cloud/utils/s3_utils.py +++ b/src/autogluon/cloud/utils/s3_utils.py @@ -2,7 +2,7 @@ from typing import Optional import boto3 -import sagemaker +from sagemaker.core.helper.session_helper import Session from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix @@ -42,7 +42,7 @@ def is_s3_folder(path, session=None): """ assert is_s3_url(path) if session is None: - session = sagemaker.session.Session() + session = Session() bucket, prefix = s3_path_to_bucket_prefix(path) contents = session.list_s3_files(bucket, prefix) if 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..67b6744a --- /dev/null +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -0,0 +1,211 @@ +"""Helpers for calling SageMaker through sagemaker-core (SageMaker Python SDK v3) resource classes. + +AutoGluon-Cloud builds each SageMaker API request as a plain snake_case dict (the sagemaker-core field +names, i.e. the AWS API fields in snake_case), merges the user's ``sagemaker_overrides`` on top, and +passes the result to the corresponding ``sagemaker.core.resources..create()`` call, which +validates it against the typed shapes before sending. +""" + +from __future__ import annotations + +import copy +import functools +import logging +import os +from typing import Any, Dict, Iterable, List, Mapping, Optional, Union + +import boto3 +from pydantic import BaseModel + +logger = logging.getLogger(__name__) + +# Keys accepted in ``sagemaker_overrides`` for each public method. Every key names the request it is merged into. +# ``production_variant`` is the single ``ProductionVariant`` inside ``create_endpoint_config.production_variants``; +# it gets its own key because deep-merging into a list is ambiguous. +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") + +# Removed SageMaker SDK v2 passthrough kwargs and where their contents moved. +_LEGACY_KWARG_HINTS = { + "backend_kwargs": ( + "`backend_kwargs` was removed when AutoGluon-Cloud migrated to SageMaker SDK v3. Use the named arguments " + "(e.g. `environment`, `use_spot_instances`, `download`, `save_path`), the predictor constructor " + "(`vpc_config`, `kms_key`, `tags`), or `sagemaker_overrides` for raw SageMaker API request fields." + ), + "custom_image_uri": "`custom_image_uri` was renamed to `image_uri`.", + "autogluon_sagemaker_estimator_kwargs": ( + "SageMaker SDK v3 has no Estimator. Use `sagemaker_overrides={'create_training_job': {...}}`." + ), + "fit_kwargs": "SageMaker SDK v3 has no Estimator. Use `sagemaker_overrides={'create_training_job': {...}}`.", + "model_kwargs": "Use `sagemaker_overrides={'create_model': {...}}` and `environment=`.", + "deploy_kwargs": ( + "Use `sagemaker_overrides={'production_variant': {...}, 'create_endpoint_config': {...}, " + "'create_endpoint': {...}}`." + ), + "transformer_kwargs": "Use `sagemaker_overrides={'create_transform_job': {...}}`.", + "transform_kwargs": "Use `sagemaker_overrides={'create_transform_job': {...}}`.", +} + + +def reject_legacy_kwargs(func): + """Turn removed v2-era kwargs into an actionable ``TypeError`` instead of Python's generic one.""" + + @functools.wraps(func) + def wrapper(*args, **kwargs): + legacy = [name for name in kwargs if name in _LEGACY_KWARG_HINTS] + if legacy: + hints = " ".join(_LEGACY_KWARG_HINTS[name] for name in legacy) + raise TypeError(f"{func.__qualname__}() got removed keyword argument(s) {legacy}. {hints}") + return func(*args, **kwargs) + + return wrapper + + +def _to_plain(value: Any) -> Any: + """Recursively convert sagemaker-core shape objects to plain snake_case dicts (only explicitly set fields).""" + if isinstance(value, BaseModel): + return {k: _to_plain(v) for k, v in value.model_dump(exclude_unset=True).items()} + if isinstance(value, Mapping): + return {k: _to_plain(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [_to_plain(v) for v in value] + return value + + +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 _to_plain(override).items(): + if isinstance(value, dict) and isinstance(merged.get(key), dict): + merged[key] = deep_merge(merged[key], value) + else: + merged[key] = value + return merged + + +def validate_sagemaker_overrides( + overrides: Optional[Mapping[str, Any]], allowed_keys: Iterable[str] +) -> Dict[str, Dict[str, Any]]: + """Check that ``overrides`` only targets requests the calling method actually sends.""" + if not overrides: + return {} + allowed_keys = tuple(allowed_keys) + unknown = sorted(set(overrides) - set(allowed_keys)) + if unknown: + raise ValueError(f"Unsupported `sagemaker_overrides` key(s) {unknown}. Valid keys here: {list(allowed_keys)}.") + for key, value in overrides.items(): + if not isinstance(value, (Mapping, BaseModel)): + raise TypeError(f"`sagemaker_overrides[{key!r}]` must be a dict, got {type(value).__name__}.") + return {key: _to_plain(value) for key, value in overrides.items()} + + +def apply_overrides(request: Dict[str, Any], overrides: Mapping[str, Any], key: str) -> Dict[str, Any]: + """Return ``request`` with ``overrides[key]`` (if any) deep-merged on top.""" + if key not in overrides: + return request + logger.debug(f"Applying sagemaker_overrides[{key!r}]: {overrides[key]}") + return deep_merge(request, overrides[key]) + + +def normalize_tags(tags: Optional[Union[Mapping[str, str], List[Dict[str, str]]]]) -> List[Dict[str, str]]: + """Accept user tags as ``{"key": "value"}`` (or a ``[{"Key", "Value"}]`` list) and return the list form.""" + if not tags: + return [] + if isinstance(tags, Mapping): + return [{"Key": str(k), "Value": str(v)} for k, v in tags.items()] + return [{"Key": t["Key"], "Value": t["Value"]} for t in tags] + + +def to_request_tags(tags: List[Dict[str, str]]) -> List[Dict[str, str]]: + """Convert ``[{"Key", "Value"}]`` tags to the sagemaker-core ``Tag`` field names.""" + return [{"key": t["Key"], "value": t["Value"]} for t in tags] + + +def normalize_vpc_config(vpc_config: Optional[Mapping[str, Any]]) -> Optional[Dict[str, List[str]]]: + """Validate a user ``vpc_config`` of the form ``{"subnets": [...], "security_group_ids": [...]}``.""" + if vpc_config is None: + return None + expected = {"subnets", "security_group_ids"} + if set(vpc_config) != expected: + raise ValueError( + f"`vpc_config` must have exactly the keys {sorted(expected)}, got {sorted(vpc_config)}. " + "Example: {'subnets': ['subnet-123'], 'security_group_ids': ['sg-123']}." + ) + return {key: list(vpc_config[key]) for key in sorted(expected)} + + +def bind_core_session(boto_session: boto3.Session) -> None: + """Point sagemaker-core at ``boto_session``. + + sagemaker-core caches one set of boto clients per process (``SageMakerClient`` is a singleton) and ignores the + ``session`` argument of resource methods once that cache exists. Several AutoGluon-Cloud objects can use + different sessions (e.g. an endpoint handle created with an explicit ``session``), so we rebuild the cached + clients whenever the requested session differs from the cached one. Not thread-safe across sessions. + """ + from sagemaker.core.utils.utils import SageMakerClient, SingletonMeta + + current = SingletonMeta._instances.get(SageMakerClient) + if current is not None and current.session is boto_session and current.region_name == boto_session.region_name: + return + SingletonMeta._instances.pop(SageMakerClient, None) + SageMakerClient(session=boto_session, region_name=boto_session.region_name) + + +def script_mode_environment(entry_point: str, region: str) -> Dict[str, str]: + """Environment variables telling the AutoGluon DLC's inference toolkit which bundled script to load. + + The model tarball always carries the serving code under ``code/`` (see ``SagemakerBackend``), which SageMaker + extracts to ``/opt/ml/model/code``. + """ + return { + "SAGEMAKER_PROGRAM": os.path.basename(entry_point), + "SAGEMAKER_SUBMIT_DIRECTORY": "/opt/ml/model/code", + "SAGEMAKER_CONTAINER_LOG_LEVEL": "20", + "SAGEMAKER_REGION": region, + } + + +def invoke_endpoint( + endpoint_name: str, + boto_session: boto3.Session, + payload: Any, + serializer, + deserializer, + content_type: Optional[str] = None, + accept: Optional[str] = None, +) -> Any: + """Serialize ``payload``, invoke a SageMaker endpoint, and deserialize the response. + + ``content_type`` / ``accept`` default to the serializer's ``CONTENT_TYPE`` and the deserializer's ``ACCEPT``. + """ + from sagemaker.core.resources import Endpoint + + bind_core_session(boto_session) + response = Endpoint(endpoint_name=endpoint_name).invoke( + body=serializer.serialize(payload), + content_type=content_type or serializer.CONTENT_TYPE, + accept=accept or ", ".join(deserializer.ACCEPT), + session=boto_session, + region=boto_session.region_name, + ) + return deserializer.deserialize(response.body, response.content_type) + + +def delete_endpoint(endpoint_name: str, boto_session: boto3.Session) -> None: + """Delete a SageMaker endpoint together with its endpoint config and the models it serves.""" + from sagemaker.core.resources import Endpoint, EndpointConfig, Model + + bind_core_session(boto_session) + endpoint = Endpoint.get(endpoint_name=endpoint_name, session=boto_session, region=boto_session.region_name) + endpoint_config = EndpointConfig.get( + endpoint_config_name=endpoint.endpoint_config_name, session=boto_session, region=boto_session.region_name + ) + model_names = [variant.model_name for variant in endpoint_config.production_variants] + + logger.info(f"Deleting endpoint {endpoint_name}") + endpoint.delete() + endpoint_config.delete() + for model_name in model_names: + logger.info(f"Deleting endpoint model {model_name}") + Model(model_name=model_name).delete() diff --git a/src/autogluon/cloud/utils/serializers.py b/src/autogluon/cloud/utils/serializers.py index 1f6633a1..e595fbc0 100644 --- a/src/autogluon/cloud/utils/serializers.py +++ b/src/autogluon/cloud/utils/serializers.py @@ -5,7 +5,7 @@ import numpy as np import pandas as pd -from sagemaker.serializers import SimpleBaseSerializer +from sagemaker.core.serializers import SimpleBaseSerializer AUTOGLUON_SERDE_VERSION = 1 diff --git a/tests/unittests/general/test_aws_utils.py b/tests/unittests/general/test_aws_utils.py index aa345e01..28faba2b 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,7 +76,7 @@ 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 diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 7817419d..4c7a45f3 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,19 @@ 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}.bind_core_session"), + mock.patch(f"{sb}.repack_model_with_serving_code") as repack, + mock.patch(f"{sb}.Model") as model_cls, + mock.patch(f"{sb}.EndpointConfig"), + mock.patch(f"{sb}.Endpoint"), mock.patch.object(SagemakerBackend, "_upload_predictor", side_effect=lambda p, _: p), ): backend = SagemakerBackend( @@ -264,9 +267,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 = model_cls.create.call_args.kwargs["primary_container"] + assert container["model_data_url"] == "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..cfdee967 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -1,24 +1,26 @@ -"""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(): + """Run ``SagemakerBackend.deploy(...)`` with AWS calls and sagemaker-core resources mocked, + and return the requests that reached ``Model.create`` / ``EndpointConfig.create`` / ``Endpoint.create``.""" 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(f"{SB}.bind_core_session"), + mock.patch(f"{SB}.Model") as model_cls, + mock.patch(f"{SB}.EndpointConfig") as endpoint_config_cls, + mock.patch(f"{SB}.Endpoint") as endpoint_cls, mock.patch.object(SagemakerBackend, "_create_serve_script_tarball", return_value="s3://stub/m.tar.gz"), ): backend = SagemakerBackend( @@ -29,50 +31,79 @@ 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) + return { + "model": model_cls.create.call_args.kwargs, + "endpoint_config": endpoint_config_cls.create.call_args.kwargs, + "endpoint": endpoint_cls.create.call_args.kwargs, + } 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"]["production_variants"] + return variant -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_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["instance_type"] == "ml.m5.xlarge" + assert variant["initial_instance_count"] == 2 + assert "serverless_config" not in variant -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_deployed_then_model_endpoint_config_and_endpoint_are_linked(deploy_requests): + requests = deploy_requests(instance_type="ml.m5.xlarge") + assert _variant(requests)["model_name"] == requests["model"]["model_name"] + assert requests["endpoint"]["endpoint_config_name"] == requests["endpoint_config"]["endpoint_config_name"] + assert requests["endpoint"]["endpoint_name"] == "ep" + environment = requests["model"]["primary_container"]["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_requests): + variant = _variant(deploy_requests(instance_type="ml.g4dn.xlarge", image_uri=GPU_IMAGE_URI)) + assert variant["inference_ami_version"] == "al2023-ami-sagemaker-inference-gpu-4-1" + + +def test_when_inference_ami_is_overridden_then_override_wins(deploy_requests): + variant = _variant( + deploy_requests( + instance_type="ml.g4dn.xlarge", + image_uri=GPU_IMAGE_URI, + sagemaker_overrides={"production_variant": {"inference_ami_version": "custom-ami"}}, + ) ) - assert captured["inference_ami_version"] == "custom-ami" + assert variant["inference_ami_version"] == "custom-ami" -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_mode_serverless_then_preset_serverless_config_is_used(deploy_requests): + variant = _variant(deploy_requests(inference_mode="serverless")) + assert variant["serverless_config"] == {"memory_size_in_mb": 4096, "max_concurrency": 5} + assert "instance_type" not in variant -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 +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["serverless_config"]["memory_size_in_mb"] == 8192 + assert variant["serverless_config"]["max_concurrency"] == 5 # preset wins for keys the user didn't override -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_environment_given_then_it_reaches_the_container(deploy_requests): + requests = deploy_requests(instance_type="ml.m5.xlarge", environment={"FOO": "bar"}) + environment = requests["model"]["primary_container"]["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 `sagemaker_overrides` key"): + deploy_requests(sagemaker_overrides={"create_training_job": {}}) diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index 402ddb80..2970d2e8 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -2,8 +2,7 @@ import pytest -from autogluon.cloud.job.sagemaker_job import SageMakerBatchTransformationJob -from autogluon.cloud.utils.ag_sagemaker import _TransformAmiVersionSession +from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend 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 +44,44 @@ 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.mark.parametrize( - ("transformer_kwargs", "expected"), + ("sagemaker_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": {"transform_resources": {"transform_ami_version": "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", +def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(sagemaker_overrides, expected): + 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(SagemakerBackend, "_upload_predictor", side_effect=lambda path, _: path), + mock.patch.object(SagemakerBackend, "_upload_batch_predict_data", return_value="s3://input/data.csv"), + mock.patch.object(SagemakerBackend, "_prepare_model_data", return_value="s3://bucket/model.tar.gz"), + mock.patch.object(SagemakerBackend, "_create_model", return_value="job"), + ): + backend = SagemakerBackend( + local_output_path="/tmp/test", + cloud_output_path="s3://bucket/run", + predictor_type="tabular", + ) + backend._fit_job = mock.MagicMock() + backend._predict( + test_data="s3://input/data.csv", + predictor_path="s3://bucket/model.tar.gz", job_name="job", - split_type="Line", - content_type="text/csv", - custom_image_uri=GPU_IMAGE_URI, + instance_type="ml.g4dn.xlarge", + image_uri=GPU_IMAGE_URI, wait=False, - model_kwargs={}, - transformer_kwargs=transformer_kwargs, + download=False, + persist=False, + sagemaker_overrides=sagemaker_overrides, ) - assert model_cls.return_value.transformer.call_args.kwargs["transform_ami_version"] == expected + request = job_cls.return_value.run.call_args.kwargs["transform_job_request"] + assert request["transform_resources"]["transform_ami_version"] == expected + assert request["transform_resources"]["instance_type"] == "ml.g4dn.xlarge" diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py new file mode 100644 index 00000000..a5ce47b4 --- /dev/null +++ b/tests/unittests/general/test_sagemaker_api.py @@ -0,0 +1,155 @@ +from unittest import mock + +import boto3 +import pandas as pd +import pytest +from sagemaker.core.shapes import StoppingCondition + +from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend +from autogluon.cloud.utils.sagemaker_api import ( + bind_core_session, + deep_merge, + normalize_tags, + normalize_vpc_config, + reject_legacy_kwargs, + validate_sagemaker_overrides, +) + +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_deep_merge_accepts_sagemaker_core_shapes(): + merged = deep_merge( + {"stopping_condition": {"max_runtime_in_seconds": 10}}, + {"stopping_condition": StoppingCondition(max_wait_time_in_seconds=20)}, + ) + assert merged == {"stopping_condition": {"max_runtime_in_seconds": 10, "max_wait_time_in_seconds": 20}} + + +def test_validate_sagemaker_overrides_rejects_unknown_keys(): + with pytest.raises(ValueError, match="create_model"): + validate_sagemaker_overrides({"create_model": {}}, ("create_training_job",)) + with pytest.raises(TypeError, match="must be a dict"): + validate_sagemaker_overrides({"create_training_job": 1}, ("create_training_job",)) + assert validate_sagemaker_overrides(None, ("create_training_job",)) == {} + + +def test_reject_legacy_kwargs_points_to_replacement(): + @reject_legacy_kwargs + def fit(**kwargs): + return kwargs + + assert fit(image_uri="x") == {"image_uri": "x"} + with pytest.raises(TypeError, match="renamed to `image_uri`"): + fit(custom_image_uri="x") + with pytest.raises(TypeError, match="sagemaker_overrides"): + fit(backend_kwargs={}) + + +def test_normalize_tags_and_vpc_config(): + assert normalize_tags({"team": "ts"}) == [{"Key": "team", "Value": "ts"}] + assert normalize_tags(None) == [] + assert normalize_vpc_config({"subnets": ("s-1",), "security_group_ids": ["sg-1"]}) == { + "security_group_ids": ["sg-1"], + "subnets": ["s-1"], + } + with pytest.raises(ValueError, match="vpc_config"): + normalize_vpc_config({"Subnets": ["s-1"]}) + + +def test_bind_core_session_rebinds_when_session_changes(): + from sagemaker.core.utils.utils import SageMakerClient, SingletonMeta + + first = boto3.Session(region_name="us-east-1") + second = boto3.Session(region_name="eu-west-1") + bind_core_session(first) + assert SageMakerClient().session is first + bind_core_session(first) + cached = SingletonMeta._instances[SageMakerClient] + bind_core_session(first) + assert SingletonMeta._instances[SageMakerClient] is cached # no rebuild for the same session + bind_core_session(second) + assert SageMakerClient().session is second + assert SageMakerClient().region_name == "eu-west-1" + + +@pytest.fixture +def fit_request(tmp_path): + """Run ``SagemakerBackend.fit(...)`` with uploads mocked and return the ``TrainingJob.create`` 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.object(SagemakerBackend, "_upload_fit_artifact", return_value={"train_data": "s3://b/train.csv"}), + ): + + def run(backend_kwargs=None, **fit_kwargs): + backend = SagemakerBackend( + local_output_path=str(tmp_path), + cloud_output_path="s3://bucket/run", + predictor_type="tabular", + **(backend_kwargs or {}), + ) + backend._fit_job = mock.MagicMock() + backend.fit( + predictor_init_args={"label": "y"}, + predictor_fit_args={}, + data_channels={"train_data": pd.DataFrame({"x": [1], "y": [0]})}, + job_name="job", + image_uri="example.com/autogluon:train", + **fit_kwargs, + ) + return backend._fit_job.run.call_args.kwargs["training_job_request"] + + yield run + + +def test_fit_builds_script_mode_training_job(fit_request): + request = fit_request(timeout=3600) + assert request["training_job_name"] == "job" + assert request["algorithm_specification"]["training_image"] == "example.com/autogluon:train" + assert request["hyper_parameters"]["sagemaker_program"] == '"train.py"' + assert request["hyper_parameters"]["sagemaker_submit_directory"] == ( + '"s3://bucket/run/code/job/source/sourcedir.tar.gz"' + ) + assert request["input_data_config"][0]["channel_name"] == "train_data" + assert request["stopping_condition"] == {"max_runtime_in_seconds": 3600} + assert request["output_data_config"] == {"s3_output_path": "s3://bucket/run/model"} + assert {"key": "autogluon-cloud-module", "value": "tabular"} in request["tags"] + assert "vpc_config" not in request + + +def test_fit_applies_infra_settings_spot_and_overrides(fit_request): + request = fit_request( + backend_kwargs={ + "vpc_config": {"subnets": ["s-1"], "security_group_ids": ["sg-1"]}, + "kms_key": "kms-1", + "tags": {"team": "ts"}, + }, + timeout=3600, + environment={"FOO": "bar"}, + use_spot_instances=True, + sagemaker_overrides={"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}, + ) + assert request["vpc_config"] == {"security_group_ids": ["sg-1"], "subnets": ["s-1"]} + assert request["output_data_config"]["kms_key_id"] == "kms-1" + assert request["resource_config"]["volume_kms_key_id"] == "kms-1" + assert {"key": "team", "value": "ts"} in request["tags"] + assert request["environment"] == {"FOO": "bar"} + assert request["enable_managed_spot_training"] is True + assert request["stopping_condition"]["max_wait_time_in_seconds"] == 3600 + assert request["retry_strategy"] == {"maximum_retry_attempts": 2} + + +def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): + with pytest.raises(ValueError, match="local mode"): + fit_request(instance_type="local") + with pytest.raises(ValueError, match="use_spot_instances"): + fit_request(max_wait=100) 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.py b/tests/unittests/tabular/test_tabular.py index be021305..5d54f800 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -47,7 +47,7 @@ def test_tabular_train(test_helper, framework_version, shared_training_job_name) predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), job_name=shared_training_job_name, ) info = predictor.info() @@ -74,7 +74,7 @@ def test_tabular_endpoint_lifecycle(test_helper, framework_version, shared_train predictor.deploy( framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=predictor.endpoint_name)["EndpointArn"] test_helper.assert_ag_cloud_tags(endpoint_arn, module="tabular") @@ -106,7 +106,7 @@ def test_tabular_batch_predict(test_helper, framework_version, shared_training_j pred, pred_proba = predictor.predict_proba( _TEST_DATA, framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) assert isinstance(pred, pd.Series) assert isinstance(pred_proba, pd.DataFrame) @@ -129,7 +129,7 @@ def test_tabular_deploy_trained_artifact(test_helper, framework_version, shared_ predictor.deploy( predictor_path=artifact_path, framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) test_helper.test_endpoint(predictor, _TEST_DATA) predictor.cleanup_deployment() @@ -152,7 +152,7 @@ def test_tabular_predict_trained_artifact(test_helper, framework_version, shared _TEST_DATA, predictor_path=artifact_path, framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) assert isinstance(pred, pd.Series) assert isinstance(pred_proba, pd.DataFrame) @@ -180,7 +180,7 @@ def test_tabular_foundation_model_predict(test_helper, framework_version): label="class", include_predict=True, framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), predictions_path=predictions_path, ) @@ -227,7 +227,7 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): "mitra-classifier", cloud_output_path=(f"s3://autogluon-cloud-ci/test-tabular-fm-deploy/{framework_version}/{timestamp}"), ) - endpoint = model.deploy(custom_image_uri=inference_custom_image_uri) + endpoint = model.deploy(image_uri=inference_custom_image_uri) try: endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=endpoint.endpoint_name)[ "EndpointArn" 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" diff --git a/tests/unittests/timeseries/test_timeseries.py b/tests/unittests/timeseries/test_timeseries.py index 8e42598f..48ea755d 100644 --- a/tests/unittests/timeseries/test_timeseries.py +++ b/tests/unittests/timeseries/test_timeseries.py @@ -59,7 +59,7 @@ def retail_sales_dataset(): def _deploy_kwargs(test_helper, framework_version: str) -> dict: return { "framework_version": framework_version, - "custom_image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + "image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), } @@ -68,7 +68,7 @@ def _predict_kwargs(test_helper, framework_version: str, ds: dict) -> dict: "static_features": ds["static_features"], "known_covariates": ds["known_covariates"], "framework_version": framework_version, - "custom_image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + "image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), } @@ -102,7 +102,7 @@ def test_timeseries_train(test_helper, framework_version, shared_training_job_na timestamp_column=ds["timestamp_column"], static_features=ds["static_features"], framework_version=framework_version, - custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), + image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), job_name=shared_training_job_name, ) info = predictor.info() @@ -255,7 +255,7 @@ def test_timeseries_fit_predict_chronos( id_column=ds["id_column"], timestamp_column=ds["timestamp_column"], framework_version=framework_version, - custom_image_uri=training_custom_image_uri, + image_uri=training_custom_image_uri, predictions_path=predictions_path, ) @@ -340,7 +340,7 @@ def test_foundation_model_cache_artifact_then_deploy_serverless(test_helper, fra assert cached_model.model_artifact_uri.startswith("s3://") endpoint = cached_model.deploy( - custom_image_uri=inference_custom_image_uri, + image_uri=inference_custom_image_uri, inference_mode="serverless", inference_config={"memory_size_in_mb": 6144}, ) @@ -466,9 +466,9 @@ def test_timeseries_endpoint_payload_formats(test_helper, framework_version, pla predictor_init_args=dict(target="target", prediction_length=_PLAIN_PREDICTION_LENGTH), predictor_fit_args=dict(presets="medium_quality", time_limit=60), framework_version=framework_version, - custom_image_uri=training_custom_image_uri, + image_uri=training_custom_image_uri, ) - cloud_predictor.deploy(framework_version=framework_version, custom_image_uri=inference_custom_image_uri) + cloud_predictor.deploy(framework_version=framework_version, image_uri=inference_custom_image_uri) try: format_pairs = list( itertools.product( @@ -512,7 +512,7 @@ def test_foundation_model_deploy(test_helper, framework_version, retail_sales_da "chronos-bolt-tiny", cloud_output_path=f"s3://autogluon-cloud-ci/test-fm-deploy-{device}/{framework_version}/{timestamp}", ) - endpoint = model.deploy(custom_image_uri=inference_custom_image_uri, **deploy_kwargs) + endpoint = model.deploy(image_uri=inference_custom_image_uri, **deploy_kwargs) try: endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=endpoint.endpoint_name)[ "EndpointArn" From ebd318c2962380f3c754bddac5304bc98f1f0a50 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 08:33:26 +0000 Subject: [PATCH 02/16] Pin sagemaker-core>=2.18 and stream job logs through our own session --- pyproject.toml | 6 +- src/autogluon/cloud/job/sagemaker_job.py | 46 +++++++---- src/autogluon/cloud/utils/job_logs.py | 99 ++++++++++++++++++++++++ tests/unittests/general/test_job_logs.py | 78 +++++++++++++++++++ 4 files changed, 213 insertions(+), 16 deletions(-) create mode 100644 src/autogluon/cloud/utils/job_logs.py create mode 100644 tests/unittests/general/test_job_logs.py diff --git a/pyproject.toml b/pyproject.toml index c0aedbd4..ad0073cf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,9 +43,11 @@ dependencies = [ # TransformAmiVersion was added to the SageMaker service model in botocore 1.37.23. "boto3>=1.37.23,<2", "packaging>=23.0,<27", - # SageMaker Python SDK v3 core layer (resources, shapes, session helpers). 2.1.0 adds TransformAmiVersion. + # SageMaker Python SDK v3 core layer (resources, shapes, session helpers). + # >=2.18: 2.13.0-2.17.0 sign all `sagemaker` API calls with default-chain credentials, ignoring the passed + # session (https://github.com/aws/sagemaker-python-sdk/issues/5986). # We deliberately don't depend on the `sagemaker` meta-package, which also pulls in torch/mlflow via sagemaker-serve. - "sagemaker-core>=2.1.0,<3", + "sagemaker-core>=2.18.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/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index 4a76ad60..1c9dd319 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -3,10 +3,10 @@ from typing import Any, Dict, Optional, Union from sagemaker.core.resources import Model, TrainingJob, TransformJob -from sagemaker.core.utils.exceptions import FailedStatusError 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 bind_core_session from .remote_job import RemoteJob @@ -14,6 +14,8 @@ class SageMakerJob(RemoteJob): + _LOG_GROUP: str + def __init__(self, session=None): self.session = session or setup_sagemaker_session() self._job_name = None @@ -120,16 +122,31 @@ def get_hyperparameters(self) -> Dict[str, Union[int, str]]: """ return self._get_hyperparameters() - def wait(self, logs: bool = True) -> None: + 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; check :meth:`get_job_status` afterwards. + Does not raise if the job fails. Returns the final status (Completed | Failed | Stopped). """ assert self.job_name, "The job has not been started" - try: - self._describe().wait(logs=logs) - except FailedStatusError as e: - logger.error(f"SageMaker job {self.job_name} did not complete successfully: {e}") + 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().failure_reason}" + ) + 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().failure_reason}" + ) def __getstate__(self): state_dict = self.__dict__.copy() @@ -141,6 +158,8 @@ def __setstate__(self, state): class SageMakerFitJob(SageMakerJob): + _LOG_GROUP = TRAINING_JOB_LOG_GROUP + def __init__(self, **kwargs): super().__init__(**kwargs) self._framework_version = None @@ -149,9 +168,6 @@ def __init__(self, **kwargs): @classmethod 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(session=session) obj._job_name = job_name obj.wait(logs=True) @@ -206,7 +222,7 @@ def run( job_name = training_job_request["training_job_name"] logger.log(20, f"Start sagemaker training job `{job_name}`") try: - training_job = TrainingJob.create( + TrainingJob.create( **training_job_request, session=self._boto_session, region=self.session.boto_region_name, @@ -214,13 +230,15 @@ def run( self._job_name = job_name self._framework_version = framework_version if wait: - training_job.wait(logs=True) + 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 = "" @@ -269,14 +287,14 @@ def run( model_name = transform_job_request["model_name"] try: logger.log(20, "Transforming") - transform_job = TransformJob.create( + TransformJob.create( **transform_job_request, session=self._boto_session, region=self.session.boto_region_name, ) self._job_name = job_name if wait: - transform_job.wait(logs=True) + self._wait_until_completed() logger.log(20, "Transform done") except Exception as e: self._delete_model(model_name) diff --git a/src/autogluon/cloud/utils/job_logs.py b/src/autogluon/cloud/utils/job_logs.py new file mode 100644 index 00000000..b0594bd2 --- /dev/null +++ b/src/autogluon/cloud/utils/job_logs.py @@ -0,0 +1,99 @@ +"""Wait for SageMaker jobs while streaming their CloudWatch logs through the caller's own boto3 session. + +sagemaker-core's ``TrainingJob.wait(logs=True)`` reads logs through a process-wide CloudWatch client built from the +default credential chain and region, ignoring the session the job was created with. Polling here 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/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 From 5883ddaf47da88d746cb2a86fe853c3d649af8f3 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 08:43:26 +0000 Subject: [PATCH 03/16] Fix sagemaker-core acronym field serialization and unit test setup --- src/autogluon/cloud/utils/sagemaker_api.py | 21 ++++++++++++++++ tests/unittests/general/test_sagemaker_ami.py | 15 ++++++------ tests/unittests/general/test_sagemaker_api.py | 24 ++++++++++++++++--- 3 files changed, 50 insertions(+), 10 deletions(-) diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index 67b6744a..26da0f87 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -48,6 +48,27 @@ } +def _fix_core_field_name_serialization() -> None: + """Teach sagemaker-core the API names of fields whose acronyms its snake_case -> PascalCase conversion mangles. + + sagemaker-core serializes nested request shapes with a naive conversion, e.g. ``memory_size_in_mb`` becomes + ``MemorySizeInMb`` instead of ``MemorySizeInMB``, which botocore rejects. The real member names are in its own + shape metadata, so register every name the naive conversion gets wrong. + """ + from sagemaker.core.utils import utils as core_utils + from sagemaker.core.utils.code_injection.shape_dag import SHAPE_DAG + + for shape in SHAPE_DAG.values(): + for member in shape.get("members") or []: + name = member["name"] + snake = core_utils.pascal_to_snake(name) + if core_utils.snake_to_pascal(snake) != name: + core_utils.SPECIAL_SNAKE_TO_PASCAL_MAPPINGS.setdefault(snake, name) + + +_fix_core_field_name_serialization() + + def reject_legacy_kwargs(func): """Turn removed v2-era kwargs into an actionable ``TypeError`` instead of Python's generic one.""" diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index 2970d2e8..5c626f39 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -1,8 +1,9 @@ from unittest import mock +import pandas as pd import pytest -from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend +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" @@ -59,19 +60,19 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(sag 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(SagemakerBackend, "_upload_predictor", side_effect=lambda path, _: path), - mock.patch.object(SagemakerBackend, "_upload_batch_predict_data", return_value="s3://input/data.csv"), - mock.patch.object(SagemakerBackend, "_prepare_model_data", return_value="s3://bucket/model.tar.gz"), - mock.patch.object(SagemakerBackend, "_create_model", return_value="job"), + 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 = SagemakerBackend( + 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="s3://input/data.csv", + test_data=pd.DataFrame({"x": [1]}), predictor_path="s3://bucket/model.tar.gz", job_name="job", instance_type="ml.g4dn.xlarge", diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index a5ce47b4..1a00548c 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -5,7 +5,7 @@ import pytest from sagemaker.core.shapes import StoppingCondition -from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend +from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend from autogluon.cloud.utils.sagemaker_api import ( bind_core_session, deep_merge, @@ -87,11 +87,13 @@ def fit_request(tmp_path): 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.object(SagemakerBackend, "_upload_fit_artifact", return_value={"train_data": "s3://b/train.csv"}), + mock.patch.object( + TabularSagemakerBackend, "_upload_fit_artifact", return_value={"train_data": "s3://b/train.csv"} + ), ): def run(backend_kwargs=None, **fit_kwargs): - backend = SagemakerBackend( + backend = TabularSagemakerBackend( local_output_path=str(tmp_path), cloud_output_path="s3://bucket/run", predictor_type="tabular", @@ -153,3 +155,19 @@ def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): fit_request(instance_type="local") with pytest.raises(ValueError, match="use_spot_instances"): fit_request(max_wait=100) + + +def test_core_serializes_acronym_field_names_with_api_casing(): + from sagemaker.core.shapes import ProductionVariant + from sagemaker.core.utils.utils import serialize + + variant = ProductionVariant( + variant_name="AllTraffic", + serverless_config={"memory_size_in_mb": 4096, "max_concurrency": 5}, + enable_ssm_access=True, + ) + assert serialize(variant) == { + "VariantName": "AllTraffic", + "ServerlessConfig": {"MemorySizeInMB": 4096, "MaxConcurrency": 5}, + "EnableSSMAccess": True, + } From c4683e84c91290f051fc16d55a6fd18e7625def6 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 08:55:37 +0000 Subject: [PATCH 04/16] Simplify SageMaker request helpers and isolate sagemaker-core workarounds --- .../cloud/backend/sagemaker_backend.py | 47 ++-- src/autogluon/cloud/job/sagemaker_job.py | 2 +- src/autogluon/cloud/model/foundation_model.py | 6 +- src/autogluon/cloud/utils/ag_sagemaker.py | 10 + src/autogluon/cloud/utils/sagemaker_api.py | 209 ++++-------------- .../cloud/utils/sagemaker_core_workarounds.py | 31 +++ src/autogluon/cloud/utils/tag_utils.py | 21 +- tests/conftest.py | 2 +- .../general/test_foundation_model.py | 4 +- tests/unittests/general/test_sagemaker_api.py | 50 ++--- tests/unittests/general/test_tags.py | 38 ++-- 11 files changed, 151 insertions(+), 269 deletions(-) create mode 100644 src/autogluon/cloud/utils/sagemaker_core_workarounds.py diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 0153917d..aaecd78a 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -10,6 +10,7 @@ from botocore.exceptions import ClientError from sagemaker.core.common_utils import sagemaker_timestamp, unique_name_from_base from sagemaker.core.resources import Endpoint, EndpointConfig, Model +from sagemaker.core.shapes import VpcConfig from autogluon.common.loaders import load_pd from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix @@ -21,6 +22,7 @@ create_serve_script_tarball, repack_model_with_serving_code, resolve_image_uri, + script_mode_environment, staged_serving_code, training_script_hyperparameters, upload_training_code, @@ -34,18 +36,14 @@ BATCH_PREDICT_OVERRIDE_KEYS, DEPLOY_OVERRIDE_KEYS, FIT_OVERRIDE_KEYS, - apply_overrides, - bind_core_session, + check_override_keys, + deep_merge, delete_endpoint, invoke_endpoint, - normalize_tags, - normalize_vpc_config, - script_mode_environment, - to_request_tags, - validate_sagemaker_overrides, ) +from ..utils.sagemaker_core_workarounds import bind_core_session from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer -from ..utils.tag_utils import build_tags +from ..utils.tag_utils import build_tags, to_request_tags from ..utils.utils import ( convert_image_path_to_encoded_bytes_in_dataframe, is_image_file, @@ -104,7 +102,7 @@ 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]]: + def _resolve_tags(self, extra_tags: Optional[Dict[str, str]] = None) -> List[Dict[str, str]]: """Tags for a created SageMaker resource, in sagemaker-core request format: default + extra + user tags.""" return to_request_tags(build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=self.tags)) @@ -140,9 +138,11 @@ def initialize( "or run `autogluon.cloud.bootstrap()` / `register()` to persist one." ) raise e - self.vpc_config = normalize_vpc_config(vpc_config) + if vpc_config is not None: + VpcConfig(**vpc_config) # fail early on malformed configs + self.vpc_config = vpc_config self.kms_key = kms_key - self.tags = normalize_tags(tags) + self.tags = dict(tags or {}) self.sagemaker_session = setup_sagemaker_session() self.endpoint_name: Optional[str] = None self._region = self.sagemaker_session.boto_region_name @@ -227,7 +227,7 @@ def fit( max_wait: Optional[int] = None, sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, extra_ag_args: Optional[Dict[str, Any]] = None, - extra_tags: Optional[List[Dict[str, str]]] = None, + extra_tags: Optional[Dict[str, str]] = None, ) -> None: """ Fit the predictor with SageMaker. @@ -290,7 +290,7 @@ def fit( if data_channels.get("train_data") is None: raise ValueError("`data_channels['train_data']` is required.") _reject_local_mode(instance_type) - overrides = validate_sagemaker_overrides(sagemaker_overrides, FIT_OVERRIDE_KEYS) + overrides = check_override_keys(sagemaker_overrides, FIT_OVERRIDE_KEYS) if max_wait is not None and not use_spot_instances: raise ValueError("`max_wait` requires `use_spot_instances=True`.") predictor_fit_args = copy.deepcopy(predictor_fit_args) @@ -398,7 +398,7 @@ def fit( if self.kms_key is not None: request["output_data_config"]["kms_key_id"] = self.kms_key request["resource_config"]["volume_kms_key_id"] = self.kms_key - request = apply_overrides(request, overrides, "create_training_job") + request = deep_merge(request, overrides.get("create_training_job", {})) self._fit_job.run(training_job_request=request, framework_version=framework_version, wait=wait) @@ -431,7 +431,7 @@ def _create_model( } if self.vpc_config is not None: request["vpc_config"] = self.vpc_config - request = apply_overrides(request, overrides, "create_model") + request = deep_merge(request, overrides.get("create_model", {})) logger.log(20, "Creating inference model...") Model.create(**request, session=self._boto_session, region=self._region) logger.log(20, "Inference model created successfully") @@ -484,7 +484,7 @@ def deploy( inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, repack: bool = True, - extra_tags: Optional[List[Dict[str, str]]] = None, + extra_tags: Optional[Dict[str, str]] = None, ) -> None: """ Deploy a predictor as a SageMaker endpoint, which can be used to do real-time inference later. @@ -543,7 +543,7 @@ def deploy( assert self.endpoint_name is None, ( "There is an endpoint already attached. Either detach it with `detach` or clean it up with `cleanup_deployment`" ) - overrides = validate_sagemaker_overrides(sagemaker_overrides, DEPLOY_OVERRIDE_KEYS) + overrides = check_override_keys(sagemaker_overrides, DEPLOY_OVERRIDE_KEYS) 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" @@ -636,7 +636,7 @@ def deploy( variant["serverless_config"] = {**preset, **(inference_config or {})} else: raise ValueError(f"Unsupported inference_mode={inference_mode!r}") - variant = apply_overrides(variant, overrides, "production_variant") + variant = deep_merge(variant, overrides.get("production_variant", {})) endpoint_config_request: Dict[str, Any] = { "endpoint_config_name": endpoint_name, @@ -645,15 +645,14 @@ def deploy( } if self.kms_key is not None and inference_mode == "realtime": endpoint_config_request["kms_key_id"] = self.kms_key - endpoint_config_request = apply_overrides(endpoint_config_request, overrides, "create_endpoint_config") - endpoint_request = apply_overrides( + endpoint_config_request = deep_merge(endpoint_config_request, overrides.get("create_endpoint_config", {})) + endpoint_request = deep_merge( { "endpoint_name": endpoint_name, "endpoint_config_name": endpoint_config_request["endpoint_config_name"], "tags": tags, }, - overrides, - "create_endpoint", + overrides.get("create_endpoint", {}), ) logger.log(20, f"Deploying model to the endpoint (inference_mode={inference_mode})") @@ -1228,7 +1227,7 @@ def _predict( batch_strategy="MultiRecord", ): _reject_local_mode(instance_type) - overrides = validate_sagemaker_overrides(sagemaker_overrides, BATCH_PREDICT_OVERRIDE_KEYS) + overrides = check_override_keys(sagemaker_overrides, BATCH_PREDICT_OVERRIDE_KEYS) if not predictor_path: predictor_path = self._fit_job.get_output_path() assert predictor_path, "No cloud trained model found." @@ -1358,7 +1357,7 @@ def _predict( "max_concurrent_transforms": 1, "tags": tags, } - request = apply_overrides(request, overrides, "create_transform_job") + request = deep_merge(request, overrides.get("create_transform_job", {})) batch_transform_job = SageMakerBatchTransformationJob(session=self.sagemaker_session) batch_transform_job.run(transform_job_request=request, wait=wait) diff --git a/src/autogluon/cloud/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index 1c9dd319..edbea55e 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -7,7 +7,7 @@ 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 bind_core_session +from ..utils.sagemaker_core_workarounds import bind_core_session from .remote_job import RemoteJob logger = logging.getLogger(__name__) diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 3a0865d5..f4358080 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -244,7 +244,7 @@ def _deploy_backend( inference_mode=inference_mode, inference_config=inference_config, repack=False, - extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], + extra_tags={"autogluon-cloud-model-id": self.model_id}, **kwargs, ) assert self._backend.endpoint_name is not None @@ -611,7 +611,7 @@ def predict( image_uri=image_uri, wait=wait, extra_ag_args=extra_ag_args, - extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], + extra_tags={"autogluon-cloud-model-id": self.model_id}, **kwargs, ) @@ -871,7 +871,7 @@ def predict_proba( image_uri=image_uri, wait=wait, extra_ag_args=extra_ag_args, - extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], + extra_tags={"autogluon-cloud-model-id": self.model_id}, **kwargs, ) diff --git a/src/autogluon/cloud/utils/ag_sagemaker.py b/src/autogluon/cloud/utils/ag_sagemaker.py index 0be0d2b5..e2f55c4e 100644 --- a/src/autogluon/cloud/utils/ag_sagemaker.py +++ b/src/autogluon/cloud/utils/ag_sagemaker.py @@ -76,6 +76,16 @@ def training_script_hyperparameters( } +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.""" diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index 26da0f87..193f7137 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -1,192 +1,70 @@ -"""Helpers for calling SageMaker through sagemaker-core (SageMaker Python SDK v3) resource classes. - -AutoGluon-Cloud builds each SageMaker API request as a plain snake_case dict (the sagemaker-core field -names, i.e. the AWS API fields in snake_case), merges the user's ``sagemaker_overrides`` on top, and -passes the result to the corresponding ``sagemaker.core.resources..create()`` call, which -validates it against the typed shapes before sending. -""" - -from __future__ import annotations +"""Helpers for building SageMaker API requests and calling them through sagemaker-core.""" import copy import functools import logging -import os -from typing import Any, Dict, Iterable, List, Mapping, Optional, Union +from typing import Any, Dict, Iterable, Mapping, Optional import boto3 -from pydantic import BaseModel +from sagemaker.core.resources import Endpoint, EndpointConfig, Model + +from .sagemaker_core_workarounds import bind_core_session logger = logging.getLogger(__name__) -# Keys accepted in ``sagemaker_overrides`` for each public method. Every key names the request it is merged into. -# ``production_variant`` is the single ``ProductionVariant`` inside ``create_endpoint_config.production_variants``; -# it gets its own key because deep-merging into a list is ambiguous. +# Requests that each method sends, i.e. the valid `sagemaker_overrides` keys. `production_variant` is the single +# variant inside `create_endpoint_config.production_variants`. 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") -# Removed SageMaker SDK v2 passthrough kwargs and where their contents moved. -_LEGACY_KWARG_HINTS = { - "backend_kwargs": ( - "`backend_kwargs` was removed when AutoGluon-Cloud migrated to SageMaker SDK v3. Use the named arguments " - "(e.g. `environment`, `use_spot_instances`, `download`, `save_path`), the predictor constructor " - "(`vpc_config`, `kms_key`, `tags`), or `sagemaker_overrides` for raw SageMaker API request fields." - ), - "custom_image_uri": "`custom_image_uri` was renamed to `image_uri`.", - "autogluon_sagemaker_estimator_kwargs": ( - "SageMaker SDK v3 has no Estimator. Use `sagemaker_overrides={'create_training_job': {...}}`." - ), - "fit_kwargs": "SageMaker SDK v3 has no Estimator. Use `sagemaker_overrides={'create_training_job': {...}}`.", - "model_kwargs": "Use `sagemaker_overrides={'create_model': {...}}` and `environment=`.", - "deploy_kwargs": ( - "Use `sagemaker_overrides={'production_variant': {...}, 'create_endpoint_config': {...}, " - "'create_endpoint': {...}}`." - ), - "transformer_kwargs": "Use `sagemaker_overrides={'create_transform_job': {...}}`.", - "transform_kwargs": "Use `sagemaker_overrides={'create_transform_job': {...}}`.", +_REMOVED_KWARGS = { + "backend_kwargs": "named arguments, the constructor's `vpc_config` / `kms_key` / `tags`, or `sagemaker_overrides`", + "custom_image_uri": "`image_uri`", + "autogluon_sagemaker_estimator_kwargs": "`sagemaker_overrides={'create_training_job': ...}`", + "fit_kwargs": "`sagemaker_overrides={'create_training_job': ...}`", + "model_kwargs": "`environment` or `sagemaker_overrides={'create_model': ...}`", + "deploy_kwargs": "`sagemaker_overrides={'production_variant': ..., 'create_endpoint_config': ...}`", + "transformer_kwargs": "`sagemaker_overrides={'create_transform_job': ...}`", + "transform_kwargs": "`sagemaker_overrides={'create_transform_job': ...}`", } -def _fix_core_field_name_serialization() -> None: - """Teach sagemaker-core the API names of fields whose acronyms its snake_case -> PascalCase conversion mangles. - - sagemaker-core serializes nested request shapes with a naive conversion, e.g. ``memory_size_in_mb`` becomes - ``MemorySizeInMb`` instead of ``MemorySizeInMB``, which botocore rejects. The real member names are in its own - shape metadata, so register every name the naive conversion gets wrong. - """ - from sagemaker.core.utils import utils as core_utils - from sagemaker.core.utils.code_injection.shape_dag import SHAPE_DAG - - for shape in SHAPE_DAG.values(): - for member in shape.get("members") or []: - name = member["name"] - snake = core_utils.pascal_to_snake(name) - if core_utils.snake_to_pascal(snake) != name: - core_utils.SPECIAL_SNAKE_TO_PASCAL_MAPPINGS.setdefault(snake, name) - - -_fix_core_field_name_serialization() - - def reject_legacy_kwargs(func): - """Turn removed v2-era kwargs into an actionable ``TypeError`` instead of Python's generic one.""" + """Raise an actionable ``TypeError`` for kwargs removed in the SageMaker SDK v3 migration.""" @functools.wraps(func) def wrapper(*args, **kwargs): - legacy = [name for name in kwargs if name in _LEGACY_KWARG_HINTS] - if legacy: - hints = " ".join(_LEGACY_KWARG_HINTS[name] for name in legacy) - raise TypeError(f"{func.__qualname__}() got removed keyword argument(s) {legacy}. {hints}") + for name in kwargs: + if name in _REMOVED_KWARGS: + raise TypeError( + f"`{name}` was removed from {func.__qualname__}(). Use {_REMOVED_KWARGS[name]} instead." + ) return func(*args, **kwargs) return wrapper -def _to_plain(value: Any) -> Any: - """Recursively convert sagemaker-core shape objects to plain snake_case dicts (only explicitly set fields).""" - if isinstance(value, BaseModel): - return {k: _to_plain(v) for k, v in value.model_dump(exclude_unset=True).items()} - if isinstance(value, Mapping): - return {k: _to_plain(v) for k, v in value.items()} - if isinstance(value, (list, tuple)): - return [_to_plain(v) for v in value] - return value +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 `sagemaker_overrides` key(s) {unknown}. Valid keys: {list(allowed_keys)}.") + return overrides 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 _to_plain(override).items(): - if isinstance(value, dict) and isinstance(merged.get(key), dict): + 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] = value + merged[key] = copy.deepcopy(value) return merged -def validate_sagemaker_overrides( - overrides: Optional[Mapping[str, Any]], allowed_keys: Iterable[str] -) -> Dict[str, Dict[str, Any]]: - """Check that ``overrides`` only targets requests the calling method actually sends.""" - if not overrides: - return {} - allowed_keys = tuple(allowed_keys) - unknown = sorted(set(overrides) - set(allowed_keys)) - if unknown: - raise ValueError(f"Unsupported `sagemaker_overrides` key(s) {unknown}. Valid keys here: {list(allowed_keys)}.") - for key, value in overrides.items(): - if not isinstance(value, (Mapping, BaseModel)): - raise TypeError(f"`sagemaker_overrides[{key!r}]` must be a dict, got {type(value).__name__}.") - return {key: _to_plain(value) for key, value in overrides.items()} - - -def apply_overrides(request: Dict[str, Any], overrides: Mapping[str, Any], key: str) -> Dict[str, Any]: - """Return ``request`` with ``overrides[key]`` (if any) deep-merged on top.""" - if key not in overrides: - return request - logger.debug(f"Applying sagemaker_overrides[{key!r}]: {overrides[key]}") - return deep_merge(request, overrides[key]) - - -def normalize_tags(tags: Optional[Union[Mapping[str, str], List[Dict[str, str]]]]) -> List[Dict[str, str]]: - """Accept user tags as ``{"key": "value"}`` (or a ``[{"Key", "Value"}]`` list) and return the list form.""" - if not tags: - return [] - if isinstance(tags, Mapping): - return [{"Key": str(k), "Value": str(v)} for k, v in tags.items()] - return [{"Key": t["Key"], "Value": t["Value"]} for t in tags] - - -def to_request_tags(tags: List[Dict[str, str]]) -> List[Dict[str, str]]: - """Convert ``[{"Key", "Value"}]`` tags to the sagemaker-core ``Tag`` field names.""" - return [{"key": t["Key"], "value": t["Value"]} for t in tags] - - -def normalize_vpc_config(vpc_config: Optional[Mapping[str, Any]]) -> Optional[Dict[str, List[str]]]: - """Validate a user ``vpc_config`` of the form ``{"subnets": [...], "security_group_ids": [...]}``.""" - if vpc_config is None: - return None - expected = {"subnets", "security_group_ids"} - if set(vpc_config) != expected: - raise ValueError( - f"`vpc_config` must have exactly the keys {sorted(expected)}, got {sorted(vpc_config)}. " - "Example: {'subnets': ['subnet-123'], 'security_group_ids': ['sg-123']}." - ) - return {key: list(vpc_config[key]) for key in sorted(expected)} - - -def bind_core_session(boto_session: boto3.Session) -> None: - """Point sagemaker-core at ``boto_session``. - - sagemaker-core caches one set of boto clients per process (``SageMakerClient`` is a singleton) and ignores the - ``session`` argument of resource methods once that cache exists. Several AutoGluon-Cloud objects can use - different sessions (e.g. an endpoint handle created with an explicit ``session``), so we rebuild the cached - clients whenever the requested session differs from the cached one. Not thread-safe across sessions. - """ - from sagemaker.core.utils.utils import SageMakerClient, SingletonMeta - - current = SingletonMeta._instances.get(SageMakerClient) - if current is not None and current.session is boto_session and current.region_name == boto_session.region_name: - return - SingletonMeta._instances.pop(SageMakerClient, None) - SageMakerClient(session=boto_session, region_name=boto_session.region_name) - - -def script_mode_environment(entry_point: str, region: str) -> Dict[str, str]: - """Environment variables telling the AutoGluon DLC's inference toolkit which bundled script to load. - - The model tarball always carries the serving code under ``code/`` (see ``SagemakerBackend``), which SageMaker - extracts to ``/opt/ml/model/code``. - """ - return { - "SAGEMAKER_PROGRAM": os.path.basename(entry_point), - "SAGEMAKER_SUBMIT_DIRECTORY": "/opt/ml/model/code", - "SAGEMAKER_CONTAINER_LOG_LEVEL": "20", - "SAGEMAKER_REGION": region, - } - - def invoke_endpoint( endpoint_name: str, boto_session: boto3.Session, @@ -196,12 +74,7 @@ def invoke_endpoint( content_type: Optional[str] = None, accept: Optional[str] = None, ) -> Any: - """Serialize ``payload``, invoke a SageMaker endpoint, and deserialize the response. - - ``content_type`` / ``accept`` default to the serializer's ``CONTENT_TYPE`` and the deserializer's ``ACCEPT``. - """ - from sagemaker.core.resources import Endpoint - + """Serialize ``payload``, invoke the endpoint, and deserialize the response.""" bind_core_session(boto_session) response = Endpoint(endpoint_name=endpoint_name).invoke( body=serializer.serialize(payload), @@ -214,19 +87,15 @@ def invoke_endpoint( def delete_endpoint(endpoint_name: str, boto_session: boto3.Session) -> None: - """Delete a SageMaker endpoint together with its endpoint config and the models it serves.""" - from sagemaker.core.resources import Endpoint, EndpointConfig, Model - + """Delete an endpoint together with its endpoint config and models.""" bind_core_session(boto_session) - endpoint = Endpoint.get(endpoint_name=endpoint_name, session=boto_session, region=boto_session.region_name) + region = boto_session.region_name + endpoint = Endpoint.get(endpoint_name=endpoint_name, session=boto_session, region=region) endpoint_config = EndpointConfig.get( - endpoint_config_name=endpoint.endpoint_config_name, session=boto_session, region=boto_session.region_name + endpoint_config_name=endpoint.endpoint_config_name, session=boto_session, region=region ) - model_names = [variant.model_name for variant in endpoint_config.production_variants] - logger.info(f"Deleting endpoint {endpoint_name}") endpoint.delete() endpoint_config.delete() - for model_name in model_names: - logger.info(f"Deleting endpoint model {model_name}") - Model(model_name=model_name).delete() + for variant in endpoint_config.production_variants: + Model(model_name=variant.model_name).delete() diff --git a/src/autogluon/cloud/utils/sagemaker_core_workarounds.py b/src/autogluon/cloud/utils/sagemaker_core_workarounds.py new file mode 100644 index 00000000..38b82e73 --- /dev/null +++ b/src/autogluon/cloud/utils/sagemaker_core_workarounds.py @@ -0,0 +1,31 @@ +"""Temporary workarounds for sagemaker-core bugs. Delete each one once fixed upstream.""" + +import boto3 +from sagemaker.core.utils import utils as core_utils +from sagemaker.core.utils.code_injection.shape_dag import SHAPE_DAG + + +def _register_acronym_field_names() -> None: + # sagemaker-core serializes nested shapes with a naive snake_case -> PascalCase conversion, so e.g. + # `memory_size_in_mb` is sent as `MemorySizeInMb` instead of `MemorySizeInMB`. Register the real API names. + for shape in SHAPE_DAG.values(): + for member in shape.get("members") or []: + snake = core_utils.pascal_to_snake(member["name"]) + if core_utils.snake_to_pascal(snake) != member["name"]: + core_utils.SPECIAL_SNAKE_TO_PASCAL_MAPPINGS.setdefault(snake, member["name"]) + + +_register_acronym_field_names() + + +def bind_core_session(boto_session: boto3.Session) -> None: + """Make sagemaker-core's process-wide client cache use ``boto_session``. + + The cache ignores the ``session`` argument of resource methods once it exists, so rebuild it whenever a different + session is requested. Not thread-safe across sessions. + """ + current = core_utils.SingletonMeta._instances.get(core_utils.SageMakerClient) + if current is not None and current.session is boto_session: + return + core_utils.SingletonMeta._instances.pop(core_utils.SageMakerClient, None) + core_utils.SageMakerClient(session=boto_session, region_name=boto_session.region_name) diff --git a/src/autogluon/cloud/utils/tag_utils.py b/src/autogluon/cloud/utils/tag_utils.py index 3fe697b7..6a5d859c 100644 --- a/src/autogluon/cloud/utils/tag_utils.py +++ b/src/autogluon/cloud/utils/tag_utils.py @@ -10,18 +10,19 @@ def build_tags( module: str, - extra_tags: Optional[List[Dict[str, str]]] = None, - user_tags: Optional[List[Dict[str, str]]] = None, -) -> List[Dict[str, str]]: - """Final tag list for a SageMaker resource: defaults + extras + user, with user winning on key collision. + extra_tags: Optional[Dict[str, str]] = None, + user_tags: Optional[Dict[str, str]] = None, +) -> Dict[str, str]: + """Final tags for a SageMaker resource: defaults + extras + user, with user winning on key collision. Defaults are skipped entirely when ``AG_CLOUD_DISABLE_DEFAULT_TAGS`` is truthy, so customers in tag-restricted AWS orgs can opt out without losing other functionality. """ if os.environ.get(DISABLE_DEFAULT_TAGS_ENV, "").lower() in ("1", "true", "yes"): - return list(user_tags or []) - base = [{"Key": "autogluon-cloud-module", "Value": module}] + list(extra_tags or []) - if not user_tags: - return base - user_keys = {t["Key"] for t in user_tags} - return [t for t in base if t["Key"] not in user_keys] + list(user_tags) + return dict(user_tags or {}) + return {"autogluon-cloud-module": module, **(extra_tags or {}), **(user_tags or {})} + + +def to_request_tags(tags: Dict[str, str]) -> List[Dict[str, str]]: + """Convert ``{key: value}`` tags to the list format of SageMaker API requests.""" + return [{"key": key, "value": value} for key, value in tags.items()] diff --git a/tests/conftest.py b/tests/conftest.py index 256e322e..ef760562 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -132,7 +132,7 @@ def tag_resources_with_ci_run(): build_tags = sagemaker_backend.build_tags def build_tags_with_ci_run(*args, **kwargs): - return build_tags(*args, **kwargs) + [{"Key": CI_RUN_TAG, "Value": run_id}] + return {**build_tags(*args, **kwargs), CI_RUN_TAG: run_id} with pytest.MonkeyPatch.context() as mp: mp.setattr(sagemaker_backend, "build_tags", build_tags_with_ci_run) diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 4c7a45f3..19b189e7 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -136,7 +136,7 @@ def test_deploy_passes_artifact_uri_and_overrides_model_path_to_container_dir(): assert call.kwargs["repack"] is False serve_cfg = call.kwargs["fm_serve_config"] assert serve_cfg["hyperparameters"]["model_path"] == "/opt/ml/model/weights" - assert {"Key": "autogluon-cloud-model-id", "Value": "chronos-2"} in call.kwargs["extra_tags"] + assert call.kwargs["extra_tags"] == {"autogluon-cloud-model-id": "chronos-2"} def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): @@ -149,7 +149,7 @@ def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): assert call.kwargs["repack"] is False serve_cfg = call.kwargs["fm_serve_config"] assert serve_cfg["hyperparameters"]["model_path"] == "autogluon/chronos-2" - assert {"Key": "autogluon-cloud-model-id", "Value": "chronos-2"} in call.kwargs["extra_tags"] + assert call.kwargs["extra_tags"] == {"autogluon-cloud-model-id": "chronos-2"} def test_tabular_deploy_uses_tabular_fm_handler_and_returns_tabular_endpoint(): diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index 1a00548c..3a1c9f49 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -3,17 +3,10 @@ import boto3 import pandas as pd import pytest -from sagemaker.core.shapes import StoppingCondition from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend -from autogluon.cloud.utils.sagemaker_api import ( - bind_core_session, - deep_merge, - normalize_tags, - normalize_vpc_config, - reject_legacy_kwargs, - validate_sagemaker_overrides, -) +from autogluon.cloud.utils.sagemaker_api import check_override_keys, deep_merge, reject_legacy_kwargs +from autogluon.cloud.utils.sagemaker_core_workarounds import bind_core_session SB = "autogluon.cloud.backend.sagemaker_backend" @@ -25,20 +18,10 @@ def test_deep_merge_merges_dicts_and_replaces_other_values(): assert base["a"]["y"] == 2 # input is not mutated -def test_deep_merge_accepts_sagemaker_core_shapes(): - merged = deep_merge( - {"stopping_condition": {"max_runtime_in_seconds": 10}}, - {"stopping_condition": StoppingCondition(max_wait_time_in_seconds=20)}, - ) - assert merged == {"stopping_condition": {"max_runtime_in_seconds": 10, "max_wait_time_in_seconds": 20}} - - -def test_validate_sagemaker_overrides_rejects_unknown_keys(): +def test_check_override_keys_rejects_unknown_keys(): with pytest.raises(ValueError, match="create_model"): - validate_sagemaker_overrides({"create_model": {}}, ("create_training_job",)) - with pytest.raises(TypeError, match="must be a dict"): - validate_sagemaker_overrides({"create_training_job": 1}, ("create_training_job",)) - assert validate_sagemaker_overrides(None, ("create_training_job",)) == {} + check_override_keys({"create_model": {}}, ("create_training_job",)) + assert check_override_keys(None, ("create_training_job",)) == {} def test_reject_legacy_kwargs_points_to_replacement(): @@ -47,32 +30,20 @@ def fit(**kwargs): return kwargs assert fit(image_uri="x") == {"image_uri": "x"} - with pytest.raises(TypeError, match="renamed to `image_uri`"): + with pytest.raises(TypeError, match="Use `image_uri` instead"): fit(custom_image_uri="x") with pytest.raises(TypeError, match="sagemaker_overrides"): fit(backend_kwargs={}) -def test_normalize_tags_and_vpc_config(): - assert normalize_tags({"team": "ts"}) == [{"Key": "team", "Value": "ts"}] - assert normalize_tags(None) == [] - assert normalize_vpc_config({"subnets": ("s-1",), "security_group_ids": ["sg-1"]}) == { - "security_group_ids": ["sg-1"], - "subnets": ["s-1"], - } - with pytest.raises(ValueError, match="vpc_config"): - normalize_vpc_config({"Subnets": ["s-1"]}) - - def test_bind_core_session_rebinds_when_session_changes(): from sagemaker.core.utils.utils import SageMakerClient, SingletonMeta first = boto3.Session(region_name="us-east-1") second = boto3.Session(region_name="eu-west-1") bind_core_session(first) - assert SageMakerClient().session is first - bind_core_session(first) cached = SingletonMeta._instances[SageMakerClient] + assert cached.session is first bind_core_session(first) assert SingletonMeta._instances[SageMakerClient] is cached # no rebuild for the same session bind_core_session(second) @@ -140,7 +111,7 @@ def test_fit_applies_infra_settings_spot_and_overrides(fit_request): use_spot_instances=True, sagemaker_overrides={"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}, ) - assert request["vpc_config"] == {"security_group_ids": ["sg-1"], "subnets": ["s-1"]} + assert request["vpc_config"] == {"subnets": ["s-1"], "security_group_ids": ["sg-1"]} assert request["output_data_config"]["kms_key_id"] == "kms-1" assert request["resource_config"]["volume_kms_key_id"] == "kms-1" assert {"key": "team", "value": "ts"} in request["tags"] @@ -150,6 +121,11 @@ def test_fit_applies_infra_settings_spot_and_overrides(fit_request): assert request["retry_strategy"] == {"maximum_retry_attempts": 2} +def test_fit_rejects_malformed_vpc_config(fit_request): + with pytest.raises(ValueError, match="security_group_ids"): + fit_request(backend_kwargs={"vpc_config": {"subnets": ["s-1"]}}) + + def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): with pytest.raises(ValueError, match="local mode"): fit_request(instance_type="local") diff --git a/tests/unittests/general/test_tags.py b/tests/unittests/general/test_tags.py index f52777a7..45199775 100644 --- a/tests/unittests/general/test_tags.py +++ b/tests/unittests/general/test_tags.py @@ -2,40 +2,36 @@ import pytest -from autogluon.cloud.utils.tag_utils import DISABLE_DEFAULT_TAGS_ENV, build_tags +from autogluon.cloud.utils.tag_utils import DISABLE_DEFAULT_TAGS_ENV, build_tags, to_request_tags def test_when_no_extras_or_user_then_only_module_tag_is_returned(): - assert build_tags("timeseries") == [{"Key": "autogluon-cloud-module", "Value": "timeseries"}] + assert build_tags("timeseries") == {"autogluon-cloud-module": "timeseries"} -def test_when_extra_tags_provided_then_appended_after_module(): - tags = build_tags("timeseries", extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": "chronos-2"}]) - assert tags == [ - {"Key": "autogluon-cloud-module", "Value": "timeseries"}, - {"Key": "autogluon-cloud-model-id", "Value": "chronos-2"}, - ] +def test_when_extra_tags_provided_then_added_to_module_tag(): + tags = build_tags("timeseries", extra_tags={"autogluon-cloud-model-id": "chronos-2"}) + assert tags == {"autogluon-cloud-module": "timeseries", "autogluon-cloud-model-id": "chronos-2"} def test_when_user_tag_collides_with_default_then_user_wins(): - tags = build_tags("timeseries", user_tags=[{"Key": "autogluon-cloud-module", "Value": "override"}]) - assert tags == [{"Key": "autogluon-cloud-module", "Value": "override"}] + tags = build_tags("timeseries", user_tags={"autogluon-cloud-module": "override"}) + assert tags == {"autogluon-cloud-module": "override"} -def test_when_user_tags_unique_then_appended_after_defaults(): - tags = build_tags("tabular", user_tags=[{"Key": "Owner", "Value": "team"}]) - assert tags == [ - {"Key": "autogluon-cloud-module", "Value": "tabular"}, - {"Key": "Owner", "Value": "team"}, - ] +def test_when_user_tags_unique_then_added_to_defaults(): + tags = build_tags("tabular", user_tags={"Owner": "team"}) + assert tags == {"autogluon-cloud-module": "tabular", "Owner": "team"} @pytest.mark.parametrize("value", ["1", "true", "True", "yes"]) def test_when_disable_env_var_set_then_defaults_and_extras_are_skipped(monkeypatch, value): """Extras are AG-cloud defaults too — opt-out drops them along with module.""" monkeypatch.setenv(DISABLE_DEFAULT_TAGS_ENV, value) - assert build_tags("timeseries") == [] - assert build_tags("timeseries", extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": "chronos-2"}]) == [] - assert build_tags("timeseries", user_tags=[{"Key": "Owner", "Value": "team"}]) == [ - {"Key": "Owner", "Value": "team"} - ] + assert build_tags("timeseries") == {} + assert build_tags("timeseries", extra_tags={"autogluon-cloud-model-id": "chronos-2"}) == {} + assert build_tags("timeseries", user_tags={"Owner": "team"}) == {"Owner": "team"} + + +def test_to_request_tags_uses_api_field_names(): + assert to_request_tags({"Owner": "team"}) == [{"key": "Owner", "value": "team"}] From 783a9c2f30e50d7764037f540256158e6f5620ef Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 10:05:41 +0000 Subject: [PATCH 05/16] Add reusable SageMakerConfig backend settings and bind prediction futures to their jobs - Replace role/vpc_config/kms_key/tags constructor args with backend=SageMakerConfig(...), shared by cloud predictors and foundation models; add explicit region and split output_kms_key/volume_kms_key. - Rename sagemaker_overrides to backend_overrides and accept it on foundation-model predict. - Bind JobPredictionFuture to the submitted job and upload fit inputs under a per-job prefix so later submissions don't overwrite earlier ones. --- README.md | 5 + docs/api.rst | 21 +- docs/api/setup.rst | 8 + docs/tutorials/setup.md | 76 ++++++- src/autogluon/cloud/__init__.py | 5 +- src/autogluon/cloud/backend/backend.py | 59 +++-- .../cloud/backend/backend_factory.py | 39 +++- .../cloud/backend/sagemaker_backend.py | 159 ++++++------- src/autogluon/cloud/config.py | 43 +++- src/autogluon/cloud/model/foundation_model.py | 106 ++++----- .../cloud/predictor/cloud_predictor.py | 68 +++--- .../predictor/tabular_cloud_predictor.py | 12 +- .../predictor/timeseries_cloud_predictor.py | 18 +- src/autogluon/cloud/utils/aws_utils.py | 12 +- src/autogluon/cloud/utils/sagemaker_api.py | 18 +- tests/unittests/general/test_aws_utils.py | 21 +- .../unittests/general/test_backend_config.py | 209 ++++++++++++++++++ .../general/test_foundation_model.py | 11 +- .../unittests/general/test_inference_modes.py | 27 ++- tests/unittests/general/test_sagemaker_ami.py | 10 +- tests/unittests/general/test_sagemaker_api.py | 39 ++-- .../test_tabular_foundation_model_predict.py | 19 +- 22 files changed, 716 insertions(+), 269 deletions(-) create mode 100644 tests/unittests/general/test_backend_config.py diff --git a/README.md b/README.md index 41e8ae87..448ddaae 100644 --- a/README.md +++ b/README.md @@ -35,6 +35,11 @@ bootstrap() See the [Setup tutorial](https://auto.gluon.ai/cloud/stable/tutorials/setup.html) for the full walkthrough, including how to register an existing role and bucket instead. +Pass `backend=SageMakerConfig(...)` to a cloud predictor or foundation model to set the region, +execution role, VPC, output and volume encryption keys, and resource tags. The same config can be +reused across workflows; each object creates its own backend state. Instance sizes, container +environment variables, and inference modes remain named arguments to individual operations. + ## ⚙️ Train your own model Train an AutoGluon predictor on your data and serve it from a SageMaker endpoint — same API as local AutoGluon, all heavy lifting on AWS. Full walkthrough: [tabular](https://auto.gluon.ai/cloud/stable/tutorials/predictor-tabular.html), [time series](https://auto.gluon.ai/cloud/stable/tutorials/predictor-timeseries.html). diff --git a/docs/api.rst b/docs/api.rst index f1d4483c..042bc97d 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -1,3 +1,5 @@ +:orphan: + API === @@ -7,48 +9,53 @@ API .. autosummary:: :toctree: api :template: custom_class.rst - :methods: + + SageMakerConfig + +.. autosummary:: + :toctree: api + :template: custom_class.rst + + FoundationModel + +.. autosummary:: + :toctree: api + :template: custom_class.rst TabularCloudPredictor .. autosummary:: :toctree: api :template: custom_class.rst - :methods: TabularFoundationModel .. autosummary:: :toctree: api :template: custom_class.rst - :methods: TabularEndpoint .. autosummary:: :toctree: api :template: custom_class.rst - :methods: TimeSeriesCloudPredictor .. autosummary:: :toctree: api :template: custom_class.rst - :methods: TimeSeriesFoundationModel .. autosummary:: :toctree: api :template: custom_class.rst - :methods: TimeSeriesEndpoint .. autosummary:: :toctree: api :template: custom_class.rst - :methods: MultiModalCloudPredictor diff --git a/docs/api/setup.rst b/docs/api/setup.rst index d515ade9..1af4cfe2 100644 --- a/docs/api/setup.rst +++ b/docs/api/setup.rst @@ -5,6 +5,14 @@ Functions for managing AutoGluon-Cloud's AWS configuration. See the :doc:`Setup .. currentmodule:: autogluon.cloud +Use :class:`SageMakerConfig` for reusable settings passed to predictors and foundation models. + +.. autosummary:: + :toctree: . + :template: custom_class.rst + + SageMakerConfig + .. autosummary:: :toctree: . :nosignatures: diff --git a/docs/tutorials/setup.md b/docs/tutorials/setup.md index 2594235e..2add942b 100644 --- a/docs/tutorials/setup.md +++ b/docs/tutorials/setup.md @@ -17,7 +17,7 @@ SageMaker compute and S3 storage are billed to your AWS account. AutoGluon-Cloud There are three ways to supply these resources — if you're unsure, start with option 1. -### 1. Create new resources with {func}`~autogluon.cloud.bootstrap` +## 1. Create new resources with {func}`~autogluon.cloud.bootstrap` Run this if you don't yet have an IAM role and S3 bucket set up for SageMaker. The role and bucket are provisioned on your account from a {repo-file}`CloudFormation template ` and saved under `~/.autogluon/cloud.yaml` for future calls. @@ -38,7 +38,7 @@ autogluon-cloud bootstrap ::: :::: -### 2. Use existing resources with {func}`~autogluon.cloud.register` +## 2. Use existing resources with {func}`~autogluon.cloud.register` Run this if you already have an IAM role and S3 bucket that you want to use with AutoGluon-Cloud. The values are saved under `~/.autogluon/cloud.yaml` for future calls. @@ -68,21 +68,87 @@ autogluon-cloud register \ The role must trust the `sagemaker.amazonaws.com` principal and grant the permissions AutoGluon-Cloud needs to run SageMaker jobs plus read/write access to your bucket — for example, a [SageMaker execution role](https://docs.aws.amazon.com/sagemaker/latest/dg/sagemaker-roles.html). For the exact set of permissions, see the {repo-file}`CloudFormation template ` that {func}`~autogluon.cloud.bootstrap` uses. The `region` where the jobs are executed must match the bucket's region. -### 3. Pass resources on each call +## 3. Pass resources on each call Skip the saved config entirely and provide the role and bucket every time you create a `CloudPredictor` or `FoundationModel`. ```python -from autogluon.cloud import TabularCloudPredictor +from autogluon.cloud import SageMakerConfig, TabularCloudPredictor predictor = TabularCloudPredictor( cloud_output_path="s3://my-autogluon-bucket/output", - role="arn:aws:iam::222222222222:role/MyAutoGluonRole", + backend=SageMakerConfig( + role_arn="arn:aws:iam::222222222222:role/MyAutoGluonRole", + region="us-east-1", + ), ) ``` Useful for one-off scripts or when you need different roles and buckets per call. The same role and bucket requirements as option 2 apply. +## Share backend settings across workflows + +{class}`~autogluon.cloud.SageMakerConfig` works with both cloud predictors and foundation models. It holds +the region, execution role, VPC, encryption keys, and resource tags. You can reuse it across objects; +each object gets its own backend, jobs, and endpoint state. + +```python +from autogluon.cloud import SageMakerConfig, TabularCloudPredictor, TimeSeriesFoundationModel + +backend = SageMakerConfig( + region="us-east-1", + role_arn="arn:aws:iam::222222222222:role/MyAutoGluonRole", + vpc_config={"subnets": ["subnet-..."], "security_group_ids": ["sg-..."]}, + output_kms_key="arn:aws:kms:us-east-1:222222222222:key/...", + tags={"team": "forecasting"}, +) + +predictor = TabularCloudPredictor( + backend=backend, + cloud_output_path="s3://my-autogluon-bucket/training", +) +model = TimeSeriesFoundationModel( + "chronos-2", + backend=backend, + cloud_output_path="s3://my-autogluon-bucket/inference", +) +``` + +The role and region you set explicitly take precedence over the saved configuration. Leaving them +unset uses the existing saved-config and AWS-identity fallbacks. `backend="sagemaker"` is shorthand +for `backend=SageMakerConfig()`. + +`output_kms_key` encrypts training artifacts, batch transform outputs, and repacked or cached model +artifacts in S3. `volume_kms_key` separately controls training, batch transform, and realtime endpoint +storage encryption; leave it unset for instances with local NVMe storage. Resource sizes, +container environment variables, spot training, and serverless settings remain arguments to the +individual `fit()`, `predict()`, and `deploy()` calls. + +## Advanced provider settings + +Use `backend_overrides` for SageMaker request fields without a named argument. It maps request names +to fields in the snake_case format used by `sagemaker-core`: + +```python +predictions = model.predict( + data, + prediction_length=24, + backend_overrides={ + "create_training_job": { + "retry_strategy": {"maximum_retry_attempts": 2}, + }, + }, +) +``` + +Foundation-model predictions and predictor training use `create_training_job`. Predictor batch +transform uses `create_model` and `create_transform_job`. Deployment accepts `create_model`, +`production_variant`, `create_endpoint_config`, and `create_endpoint`. + +Only requests used by the operation are accepted. Nested dictionaries merge recursively over the +generated request; other values, including lists, replace the generated value. Overrides take +precedence over backend settings and named arguments. + ## Managing the saved config Once {func}`~autogluon.cloud.bootstrap` or {func}`~autogluon.cloud.register` has written to `~/.autogluon/cloud.yaml`, you may want to check that the role and bucket are still healthy before a long training run, or clean everything up when you're done with AutoGluon-Cloud. Two helper commands cover both: diff --git a/src/autogluon/cloud/__init__.py b/src/autogluon/cloud/__init__.py index ad6c6385..bf61c90e 100644 --- a/src/autogluon/cloud/__init__.py +++ b/src/autogluon/cloud/__init__.py @@ -3,19 +3,22 @@ from autogluon.common.utils.log_utils import _add_stream_handler from .cloud_setup import bootstrap, register, status, teardown +from .config import SageMakerConfig from .endpoint.tabular_endpoint import TabularEndpoint from .endpoint.timeseries_endpoint import TimeSeriesEndpoint -from .model.foundation_model import TabularFoundationModel, TimeSeriesFoundationModel +from .model.foundation_model import FoundationModel, TabularFoundationModel, TimeSeriesFoundationModel from .predictor import MultiModalCloudPredictor, TabularCloudPredictor, TimeSeriesCloudPredictor _add_stream_handler() logging.getLogger(__name__).setLevel(logging.INFO) __all__ = [ + "FoundationModel", "MultiModalCloudPredictor", "TabularCloudPredictor", "TabularEndpoint", "TabularFoundationModel", + "SageMakerConfig", "TimeSeriesCloudPredictor", "TimeSeriesEndpoint", "TimeSeriesFoundationModel", diff --git a/src/autogluon/cloud/backend/backend.py b/src/autogluon/cloud/backend/backend.py index 60d98c25..63270723 100644 --- a/src/autogluon/cloud/backend/backend.py +++ b/src/autogluon/cloud/backend/backend.py @@ -3,10 +3,13 @@ import json import os from abc import ABC, abstractmethod -from typing import Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Union import pandas as pd +if TYPE_CHECKING: + from ..endpoint.prediction_future import JobPredictionFuture + def dumps_ag_args(config: Dict[str, Any]) -> str: """Serialize the remote-training config to JSON, raising a user-facing error on failure. @@ -39,28 +42,14 @@ def dumps_ag_args(config: Dict[str, Any]) -> str: class Backend(ABC): name = "backend" - def __init__(self, **kwargs) -> None: - self.initialize(**kwargs) - - @property - def cloud_output_path(self) -> str: - if not self._cloud_output_path: - raise ValueError( - "No `cloud_output_path` was provided and no bucket is configured in " - "~/.autogluon/cloud.yaml. Either pass `cloud_output_path=` explicitly, or run " - "`autogluon.cloud.bootstrap()` / `register(bucket=...)` once to persist a bucket." - ) - return self._cloud_output_path - - def initialize( + def __init__( self, + *, local_output_path: str, predictor_type: str, cloud_output_path: Optional[str] = None, resource_prefix: Optional[str] = None, - **kwargs, ) -> None: - """Initialize the backend.""" self.local_output_path = local_output_path self._cloud_output_path = cloud_output_path self.predictor_type = predictor_type @@ -68,6 +57,16 @@ def initialize( self.original_features = None self.endpoint_name: Optional[str] = None + @property + def cloud_output_path(self) -> str: + if not self._cloud_output_path: + raise ValueError( + "No `cloud_output_path` was provided and no bucket is configured in " + "~/.autogluon/cloud.yaml. Either pass `cloud_output_path=` explicitly, or run " + "`autogluon.cloud.bootstrap()` / `register(bucket=...)` once to persist a bucket." + ) + return self._cloud_output_path + @abstractmethod def attach_job(self, job_name: str) -> None: """ @@ -124,21 +123,33 @@ def prepare_args(self, path: str, **kwargs): with open(path, "w") as f: f.write(payload) - def _construct_ag_args(**kwargs): + def _construct_ag_args(self, **kwargs): raise NotImplementedError @abstractmethod - def fit(self, **kwargs) -> None: + def fit( + self, + *, + predictor_init_args: Dict[str, Any], + predictor_fit_args: Dict[str, Any], + data_channels: Dict[str, Optional[Union[str, pd.DataFrame]]], + **kwargs, + ) -> None: """Fit AG on the backend""" raise NotImplementedError @abstractmethod - def deploy(self, **kwargs) -> None: + def deploy( + self, + predictor_path: Optional[str] = None, + endpoint_name: Optional[str] = None, + **kwargs, + ) -> None: """Deploy and endpoint""" raise NotImplementedError @abstractmethod - def cleanup_deployment(self, **kwargs) -> None: + def cleanup_deployment(self) -> None: """Delete endpoint, and cleanup other artifacts""" raise NotImplementedError @@ -199,3 +210,9 @@ def get_fit_predict_results(self) -> pd.DataFrame: """ raise NotImplementedError(f"{self.__class__.__name__} does not support `fit_predict`.") + + def get_prediction_future( + self, *, result_transform: Optional[Callable[[pd.DataFrame], Any]] = None + ) -> JobPredictionFuture: + """Return a pending result bound to the most recently submitted prediction job.""" + raise NotImplementedError(f"{self.__class__.__name__} does not support prediction futures.") diff --git a/src/autogluon/cloud/backend/backend_factory.py b/src/autogluon/cloud/backend/backend_factory.py index a4204f0d..0b864a54 100644 --- a/src/autogluon/cloud/backend/backend_factory.py +++ b/src/autogluon/cloud/backend/backend_factory.py @@ -1,3 +1,6 @@ +from typing import Optional, Union + +from ..config import SageMakerConfig from .backend import Backend from .multimodal_sagemaker_backend import MultiModalSagemakerBackend from .sagemaker_backend import SagemakerBackend @@ -6,6 +9,7 @@ class BackendFactory: + _CONFIGS = {SageMakerConfig.name: SageMakerConfig} _BACKENDS = { SagemakerBackend.name: SagemakerBackend, TabularSagemakerBackend.name: TabularSagemakerBackend, @@ -13,6 +17,21 @@ class BackendFactory: TimeSeriesSagemakerBackend.name: TimeSeriesSagemakerBackend, } + @staticmethod + def resolve_config(backend: Union[str, SageMakerConfig]) -> SageMakerConfig: + """Normalize a backend name or reusable configuration without creating resources.""" + if isinstance(backend, SageMakerConfig): + return backend + if not isinstance(backend, str): + raise TypeError("`backend` must be a backend name or SageMakerConfig.") + if backend in ("ray", "ray_aws"): + raise ValueError("The Ray backend was removed in AutoGluon-Cloud v0.7.0. Use backend='sagemaker' instead.") + if backend not in BackendFactory._CONFIGS: + raise ValueError( + f"Unsupported backend {backend!r}. Supported backends: {sorted(BackendFactory._CONFIGS)}." + ) + return BackendFactory._CONFIGS[backend]() + @staticmethod def get_backend_cls(backend: str) -> type[Backend]: if backend in BackendFactory._BACKENDS: @@ -20,6 +39,20 @@ def get_backend_cls(backend: str) -> type[Backend]: raise ValueError(f"{backend} not supported. Supported backends: {sorted(BackendFactory._BACKENDS)}") @staticmethod - def get_backend(backend: str, **init_args) -> Backend: - """Return the corresponding backend""" - return BackendFactory.get_backend_cls(backend)(**init_args) + def get_backend( + backend: str, + *, + local_output_path: str, + cloud_output_path: Optional[str], + predictor_type: str, + config: Optional[SageMakerConfig] = None, + resource_prefix: Optional[str] = None, + ) -> Backend: + """Create a backend with its own execution state from reusable settings.""" + return BackendFactory.get_backend_cls(backend)( + local_output_path=local_output_path, + cloud_output_path=cloud_output_path, + predictor_type=predictor_type, + config=config, + resource_prefix=resource_prefix, + ) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index aaecd78a..429cab33 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -4,7 +4,9 @@ import os import tarfile import tempfile -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from dataclasses import replace +from functools import partial +from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union import pandas as pd from botocore.exceptions import ClientError @@ -15,7 +17,9 @@ from autogluon.common.loaders import load_pd from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix +from ..config import SageMakerConfig from ..data import FormatConverterFactory +from ..endpoint.prediction_future import JobPredictionFuture from ..job import SageMakerBatchTransformationJob, SageMakerFitJob from ..scripts import ScriptManager from ..utils.ag_sagemaker import ( @@ -85,70 +89,47 @@ class SagemakerBackend(Backend): def __init__( self, local_output_path: str, - cloud_output_path: str, + cloud_output_path: Optional[str], predictor_type: str, - role: Optional[str] = None, - **kwargs, + *, + config: Optional[SageMakerConfig] = None, + resource_prefix: Optional[str] = None, ) -> None: - self.initialize( + super().__init__( local_output_path=local_output_path, cloud_output_path=cloud_output_path, predictor_type=predictor_type, - role=role, - **kwargs, + resource_prefix=resource_prefix, ) - - def _realtime_serializer(self): - """Serializer used for realtime endpoint requests""" - return AutoGluonSerializer() - - def _resolve_tags(self, extra_tags: Optional[Dict[str, str]] = None) -> List[Dict[str, str]]: - """Tags for a created SageMaker resource, in sagemaker-core request format: default + extra + user tags.""" - return to_request_tags(build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=self.tags)) - - def initialize( - self, - role: Optional[str] = None, - vpc_config: Optional[Dict[str, List[str]]] = None, - kms_key: Optional[str] = None, - tags: Optional[Dict[str, str]] = None, - **kwargs, - ) -> None: - """Initialize the backend. - - Parameters - ---------- - role - SageMaker execution role ARN. See - :func:`autogluon.cloud.utils.aws_utils.resolve_execution_role` for the resolution order. - vpc_config - ``{"subnets": [...], "security_group_ids": [...]}`` applied to every training job, model and - transform job created by this backend. - kms_key - KMS key used to encrypt the S3 outputs and ML storage volumes of the created SageMaker resources. - tags - ``{"key": "value"}`` tags added to every created SageMaker resource. - """ - super().initialize(**kwargs) + if config is not None and not isinstance(config, SageMakerConfig): + raise TypeError("`config` must be a SageMakerConfig.") + config = copy.deepcopy(config or SageMakerConfig()) + if config.vpc_config is not None: + VpcConfig(**config.vpc_config) # fail before creating a session or resolving the role + self.sagemaker_session = setup_sagemaker_session(region=config.region) try: - self.role_arn = resolve_execution_role(role, backend_name=SAGEMAKER) + self.role_arn = resolve_execution_role( + config.role_arn, 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 " + "Failed to resolve SageMaker execution role. Pass `backend=SageMakerConfig(role_arn=)` " "or run `autogluon.cloud.bootstrap()` / `register()` to persist one." ) raise e - if vpc_config is not None: - VpcConfig(**vpc_config) # fail early on malformed configs - self.vpc_config = vpc_config - self.kms_key = kms_key - self.tags = dict(tags or {}) - self.sagemaker_session = setup_sagemaker_session() - self.endpoint_name: Optional[str] = None self._region = self.sagemaker_session.boto_region_name + self.config = replace(config, region=self._region, role_arn=self.role_arn) self._fit_job: SageMakerFitJob = SageMakerFitJob(session=self.sagemaker_session) self._batch_transform_jobs = MostRecentInsertedOrderedDict() + def _realtime_serializer(self): + """Serializer used for realtime endpoint requests""" + return AutoGluonSerializer() + + def _resolve_tags(self, extra_tags: Optional[Dict[str, str]] = None) -> List[Dict[str, str]]: + """Tags for a created SageMaker resource, in sagemaker-core request format: default + extra + user tags.""" + return to_request_tags(build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=self.config.tags)) + @property def _boto_session(self): boto_session = self.sagemaker_session.boto_session @@ -225,7 +206,7 @@ def fit( environment: Optional[Dict[str, str]] = None, use_spot_instances: bool = False, max_wait: Optional[int] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, extra_ag_args: Optional[Dict[str, Any]] = None, extra_tags: Optional[Dict[str, str]] = None, ) -> None: @@ -279,7 +260,7 @@ def fit( max_wait: Optional[int], default = None Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires ``use_spot_instances=True``. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Raw ``CreateTrainingJob`` request fields (sagemaker-core snake_case 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 @@ -290,7 +271,7 @@ 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(sagemaker_overrides, FIT_OVERRIDE_KEYS) + overrides = check_override_keys(backend_overrides, FIT_OVERRIDE_KEYS) if max_wait is not None and not use_spot_instances: raise ValueError("`max_wait` requires `use_spot_instances=True`.") predictor_fit_args = copy.deepcopy(predictor_fit_args) @@ -348,6 +329,7 @@ def fit( ag_args_path = os.path.join(self.local_output_path, "utils", "ag_args.json") self.prepare_args(path=ag_args_path, **ag_args) inputs = self._upload_fit_artifact( + job_name=job_name, data_channels=data_channels, label=label, ag_args=ag_args_path, @@ -393,13 +375,15 @@ def fit( request["environment"] = dict(environment) if use_spot_instances: request["enable_managed_spot_training"] = True - if self.vpc_config is not None: - request["vpc_config"] = self.vpc_config - if self.kms_key is not None: - request["output_data_config"]["kms_key_id"] = self.kms_key - request["resource_config"]["volume_kms_key_id"] = self.kms_key + if self.config.vpc_config is not None: + request["vpc_config"] = self.config.vpc_config + if self.config.output_kms_key is not None: + request["output_data_config"]["kms_key_id"] = self.config.output_kms_key + if self.config.volume_kms_key is not None: + request["resource_config"]["volume_kms_key_id"] = self.config.volume_kms_key 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) def _create_model( @@ -429,8 +413,8 @@ def _create_model( "execution_role_arn": self.role_arn, "tags": tags, } - if self.vpc_config is not None: - request["vpc_config"] = self.vpc_config + if self.config.vpc_config is not None: + request["vpc_config"] = self.config.vpc_config request = deep_merge(request, overrides.get("create_model", {})) logger.log(20, "Creating inference model...") Model.create(**request, session=self._boto_session, region=self._region) @@ -453,7 +437,7 @@ def _prepare_model_data( entry_point=entry_point, repacked_model_uri=repacked_model_uri, sagemaker_session=self.sagemaker_session, - kms_key=self.kms_key, + kms_key=self.config.output_kms_key, ) @staticmethod @@ -478,7 +462,7 @@ def deploy( volume_size: Optional[int] = None, wait: bool = True, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = 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", @@ -520,7 +504,7 @@ def deploy( To be noticed, the function won't return immediately because there are some preparations needed prior deployment. environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Raw request fields (sagemaker-core snake_case names) deep-merged over the requests built by AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"production_variant"``, ``"create_endpoint_config"``, ``"create_endpoint"``. @@ -543,7 +527,7 @@ def deploy( assert self.endpoint_name is None, ( "There is an endpoint already attached. Either detach it with `detach` or clean it up with `cleanup_deployment`" ) - overrides = check_override_keys(sagemaker_overrides, DEPLOY_OVERRIDE_KEYS) + overrides = check_override_keys(backend_overrides, DEPLOY_OVERRIDE_KEYS) 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" @@ -643,8 +627,8 @@ def deploy( "production_variants": [variant], "tags": tags, } - if self.kms_key is not None and inference_mode == "realtime": - endpoint_config_request["kms_key_id"] = self.kms_key + if self.config.volume_kms_key is not None and inference_mode == "realtime": + endpoint_config_request["kms_key_id"] = self.config.volume_kms_key endpoint_config_request = deep_merge(endpoint_config_request, overrides.get("create_endpoint_config", {})) endpoint_request = deep_merge( { @@ -839,7 +823,7 @@ def predict( persist: bool = True, save_path: Optional[str] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Predict using SageMaker batch transform. @@ -889,7 +873,7 @@ def predict( If `persist` is `False`, file would first be downloaded to this path and then removed. environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Raw request fields (sagemaker-core snake_case names) deep-merged over the requests built by AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"create_transform_job"``. @@ -913,7 +897,7 @@ def predict( persist=persist, save_path=save_path, environment=environment, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, original_features=self.original_features, ) @@ -1000,7 +984,21 @@ def download_predict_results(self, job_name: Optional[str] = None, save_path: Op def get_fit_predict_results(self) -> pd.DataFrame: """Read predictions produced by a completed ``fit_predict`` job from S3.""" - ag_args = self._download_ag_args_from_job() + return self._load_fit_predict_results(self._fit_job) + + def get_prediction_future( + self, *, result_transform: Optional[Callable[[pd.DataFrame], Any]] = None + ) -> JobPredictionFuture: + """Bind waiting and result loading to the current job, independent of later submissions.""" + job = self._fit_job + if job.job_name is None: + raise ValueError("No prediction job found. Submit a prediction first.") + load_results = partial(self._load_fit_predict_results, job) + result_loader = load_results if result_transform is None else lambda: result_transform(load_results()) + return JobPredictionFuture(job=job, result_loader=result_loader) + + def _load_fit_predict_results(self, job: SageMakerFitJob) -> pd.DataFrame: + ag_args = self._download_ag_args_from_job(job) predictions_path = ag_args.get("predictions_path") assert predictions_path is not None, "No fit_predict job found. Call `fit_predict()` first." bucket, key = s3_path_to_bucket_prefix(predictions_path) @@ -1009,15 +1007,16 @@ def get_fit_predict_results(self) -> pd.DataFrame: self.sagemaker_session.boto_session.client("s3").download_file(bucket, key, local_path) return load_pd.load(local_path) - def _download_ag_args_from_job(self) -> Dict[str, Any]: + def _download_ag_args_from_job(self, job: Optional[SageMakerFitJob] = None) -> Dict[str, Any]: """Fetch and parse the ``ag_args.json`` that was uploaded as the ``ag_args`` channel. Each training job carries the exact config it was launched with as an input channel, making this the authoritative source — independent of local-disk lifetime. """ - job_name = self._fit_job.job_name + job = self._fit_job if job is None else job + job_name = job.job_name assert job_name is not None, "No fit job found. Call `fit()` / `fit_predict()` first." - channels = self._fit_job.get_input_channels() + channels = job.get_input_channels() ag_args_uri = channels.get("ag_args") assert ag_args_uri is not None, ( f"Training job {job_name!r} has no `ag_args` input channel — cannot recover predictions_path." @@ -1064,6 +1063,7 @@ def _find_common_path_and_replace_image_column(self, data, image_column): def _upload_fit_artifact( self, + job_name: str, data_channels, label, ag_args, @@ -1071,7 +1071,7 @@ def _upload_fit_artifact( image_column=None, ): cloud_bucket, cloud_key_prefix = s3_path_to_bucket_prefix(self.cloud_output_path) - util_key_prefix = cloud_key_prefix + "/utils" + util_key_prefix = f"{cloud_key_prefix}/{job_name}/utils" # Image-column mode: rewrite image paths to be container-relative; common image directories # are zipped and uploaded as separate train_images / tune_images channels below. @@ -1217,7 +1217,7 @@ def _predict( persist=True, save_path=None, environment=None, - sagemaker_overrides=None, + backend_overrides=None, split_pred_proba=True, original_features=None, content_type="text/csv", @@ -1227,7 +1227,7 @@ def _predict( batch_strategy="MultiRecord", ): _reject_local_mode(instance_type) - overrides = check_override_keys(sagemaker_overrides, BATCH_PREDICT_OVERRIDE_KEYS) + overrides = check_override_keys(backend_overrides, BATCH_PREDICT_OVERRIDE_KEYS) if not predictor_path: predictor_path = self._fit_job.get_output_path() assert predictor_path, "No cloud trained model found." @@ -1341,9 +1341,10 @@ def _predict( transform_ami_version = infer_sagemaker_ami_version(image_uri, instance_type, image_scope="transform") if transform_ami_version is not None: transform_resources["transform_ami_version"] = transform_ami_version - if self.kms_key is not None: - transform_output["kms_key_id"] = self.kms_key - transform_resources["volume_kms_key_id"] = self.kms_key + if self.config.output_kms_key is not None: + transform_output["kms_key_id"] = self.config.output_kms_key + if self.config.volume_kms_key is not None: + transform_resources["volume_kms_key_id"] = self.config.volume_kms_key request = { "transform_job_name": job_name, "model_name": model_name, @@ -1393,7 +1394,7 @@ def __getstate__(self) -> Dict[str, Any]: def __setstate__(self, state): """Custom implementation of the unpickle process""" self.__dict__.update(state) - self.sagemaker_session = setup_sagemaker_session() + self.sagemaker_session = setup_sagemaker_session(region=self.config.region) self._region = self.sagemaker_session.boto_region_name self._fit_job.session = self.sagemaker_session for job in self._batch_transform_jobs.values(): diff --git a/src/autogluon/cloud/config.py b/src/autogluon/cloud/config.py index 31ca9cbf..d8fb48f9 100644 --- a/src/autogluon/cloud/config.py +++ b/src/autogluon/cloud/config.py @@ -1,4 +1,4 @@ -"""Persistent config for AutoGluon-Cloud. +"""Backend settings and persistent resource identifiers for AutoGluon-Cloud. Stores resource identifiers (region, stack name, bucket, IAM role ARN) at ``~/.autogluon/cloud.yaml`` so users don't need to re-specify them every @@ -20,13 +20,52 @@ import stat from dataclasses import asdict, dataclass, field from pathlib import Path -from typing import Dict, Optional +from typing import ClassVar, Dict, List, Optional import yaml CONFIG_DIR_ENV = "AG_CONFIG_DIR" +@dataclass(kw_only=True) +class SageMakerConfig: + """Reusable SageMaker settings for predictors and foundation models. + + Pass this as ``backend=`` to a cloud predictor or foundation model. Each + object creates its own backend and jobs; sharing this config does not share + execution state. Resource sizes and other operation settings remain named + arguments to ``fit()``, ``predict()`` and ``deploy()``. + + Parameters + ---------- + region + AWS region. If omitted, use the region in ``~/.autogluon/cloud.yaml``, + then the boto3 default region. + role_arn + SageMaker execution role ARN. If omitted, use the saved role, then the + role of the current AWS identity. + vpc_config + Networking for training jobs and models, as + ``{"subnets": [...], "security_group_ids": [...]}``. + output_kms_key + KMS key for training artifacts, batch transform outputs, and repacked + or cached model artifacts in S3. + volume_kms_key + KMS key for training, batch transform and realtime endpoint storage + volumes. Leave unset for instance types with local NVMe storage. + tags + Tags added to every SageMaker resource created by this backend. + """ + + name: ClassVar[str] = "sagemaker" + region: Optional[str] = None + role_arn: Optional[str] = None + vpc_config: Optional[Dict[str, List[str]]] = None + output_kms_key: Optional[str] = None + volume_kms_key: Optional[str] = None + tags: Dict[str, str] = field(default_factory=dict) + + def get_config_dir() -> Path: override = os.environ.get(CONFIG_DIR_ENV) if override: diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index f4358080..ae925ede 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -7,6 +7,7 @@ import tarfile import tempfile from abc import abstractmethod +from functools import partial from pathlib import Path from typing import Any, Dict, List, Literal, Optional, Tuple, Union @@ -18,6 +19,7 @@ from ..backend.backend_factory import BackendFactory from ..backend.constant import SAGEMAKER, TABULAR_SAGEMAKER, TIMESERIES_SAGEMAKER +from ..config import SageMakerConfig from ..endpoint.prediction_future import JobPredictionFuture from ..endpoint.tabular_endpoint import TabularEndpoint from ..endpoint.timeseries_endpoint import TimeSeriesEndpoint @@ -81,13 +83,9 @@ def __init__( model_id: str, *, cloud_output_path: Optional[str] = None, - role: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, model_artifact_uri: Optional[str] = None, - vpc_config: Optional[Dict[str, List[str]]] = None, - kms_key: Optional[str] = None, - tags: Optional[Dict[str, str]] = None, - backend: Literal["sagemaker"] = "sagemaker", + backend: Union[str, SageMakerConfig] = SAGEMAKER, ): """ Parameters @@ -105,46 +103,35 @@ def __init__( * ``None`` (default) — use the bucket saved in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` / :func:`autogluon.cloud.register`) and append a timestamped subfolder. Raises if no bucket is configured. - 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 the role of the current AWS identity. hyperparameters Default hyperparameters applied to inference and (when supported) training. model_artifact_uri S3 URI of a pre-bundled ``model.tar.gz`` produced by :meth:`cache_model_artifact`. When set, deploys skip the runtime HuggingFace download and load weights from the bundled artifact. - vpc_config - VPC for the created SageMaker jobs and models, as ``{"subnets": [...], "security_group_ids": [...]}``. - kms_key - KMS key ID/ARN used to encrypt S3 outputs and ML storage volumes of the created SageMaker resources. - tags - Tags added to every SageMaker resource created for this model, e.g. ``{"team": "forecasting"}``. backend - Cloud backend to use. + Backend name or reusable :class:`~autogluon.cloud.SageMakerConfig` with region, execution role, + networking, encryption and tags. ``"sagemaker"`` uses default settings. """ + backend_config = BackendFactory.resolve_config(backend) + backend_name = self._backend_map.get(backend_config.name) + if backend_name is None: + raise ValueError( + f"Backend {backend_config.name!r} is not supported for {self.__class__.__name__}. " + f"Available: {list(self._backend_map.keys())}" + ) self.model_id = model_id self.model_artifact_uri = model_artifact_uri - self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend) + self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend_config.name) self._config = get_model_config(model_id) self._hyperparameter_overrides = hyperparameters or {} - self._infra_settings = {"vpc_config": vpc_config, "kms_key": kms_key, "tags": tags} self._tmpdir = tempfile.TemporaryDirectory(prefix="ag_fm_") - - backend_name = self._backend_map.get(backend) - if backend_name is None: - raise ValueError( - f"Backend '{backend}' is not supported for {self.__class__.__name__}. " - f"Available: {list(self._backend_map.keys())}" - ) self._backend = BackendFactory.get_backend( backend=backend_name, local_output_path=self._tmpdir.name, cloud_output_path=self.cloud_output_path, predictor_type=self._predictor_type, resource_prefix=f"ag-cloud-{self.model_id}", - role=role, - **self._infra_settings, + config=backend_config, ) def _get_hyperparameters( @@ -352,11 +339,14 @@ def cache_model_artifact(self, cache_path: str, *, overwrite: bool = False) -> S tar.add(serve_script, arcname=f"code/{serve_script.name}") tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils") logger.info(f"Uploading to {cache_key}") + extra_args = {"Metadata": {_AG_CLOUD_VERSION_METADATA_KEY: __version__}} + if self._backend.config.output_kms_key is not None: + extra_args.update(ServerSideEncryption="aws:kms", SSEKMSKeyId=self._backend.config.output_kms_key) s3.upload_file( str(tarball), bucket, key, - ExtraArgs={"Metadata": {_AG_CLOUD_VERSION_METADATA_KEY: __version__}}, + ExtraArgs=extra_args, ) return self.__class__( @@ -364,12 +354,11 @@ def cache_model_artifact(self, cache_path: str, *, overwrite: bool = False) -> S hyperparameters=self._hyperparameter_overrides or None, model_artifact_uri=cache_key, cloud_output_path=self.cloud_output_path, - role=self._backend.role_arn, - **self._infra_settings, + backend=self._backend.config, ) def to_dict(self) -> Dict[str, Any]: - """Serialize the model identity. Runtime context (``role``, ``cloud_output_path``) is excluded so configs can + """Serialize the model identity. Runtime context (``backend``, ``cloud_output_path``) is excluded so configs can be shared across users.""" out: Dict[str, Any] = {"model_id": self.model_id} if self._hyperparameter_overrides: @@ -384,7 +373,7 @@ def to_json(self) -> str: @classmethod def from_dict(cls, config: Dict[str, Any], **runtime_context: Any) -> Self: - """Restore from :meth:`to_dict` output. Pass ``role`` / ``cloud_output_path`` as ``runtime_context``.""" + """Restore from :meth:`to_dict` output. Pass ``backend`` / ``cloud_output_path`` as ``runtime_context``.""" return cls(**config, **runtime_context) @classmethod @@ -432,7 +421,7 @@ def deploy( inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> TimeSeriesEndpoint: """ @@ -460,7 +449,7 @@ def deploy( Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). environment Environment variables set in the inference container. - sagemaker_overrides + backend_overrides Raw SageMaker request fields deep-merged over the requests built by AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"production_variant"``, ``"create_endpoint_config"``, ``"create_endpoint"``. See :meth:`autogluon.cloud.TabularCloudPredictor.deploy`. @@ -477,7 +466,7 @@ def deploy( inference_mode=inference_mode, inference_config=inference_config, environment=environment, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, **kwargs, ) return TimeSeriesEndpoint( @@ -525,6 +514,7 @@ def predict( framework_version: str = "latest", image_uri: Optional[str] = None, wait: bool = True, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> Union[pd.DataFrame, JobPredictionFuture]: """ @@ -571,9 +561,12 @@ def predict( If True, block and return a DataFrame. If False, return a :class:`JobPredictionFuture` immediately — call ``.result()`` on it later to retrieve the DataFrame, or ``.status()`` to check progress. + backend_overrides + Raw provider request fields. This prediction path uses a SageMaker training job; + the valid key is ``"create_training_job"``. See :meth:`autogluon.cloud.TabularCloudPredictor.fit`. **kwargs Additional job arguments accepted by :meth:`autogluon.cloud.TimeSeriesCloudPredictor.fit` (e.g. - ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``, ``sagemaker_overrides``). + ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``). Returns ------- @@ -610,16 +603,14 @@ def predict( instance_type=instance_type, image_uri=image_uri, wait=wait, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, extra_tags={"autogluon-cloud-model-id": self.model_id}, **kwargs, ) if not wait: - return JobPredictionFuture( - job=self._backend._fit_job, - result_loader=self._backend.get_fit_predict_results, - ) + return self._backend.get_prediction_future() return self._backend.get_fit_predict_results() @@ -655,7 +646,7 @@ def deploy( inference_mode: Literal["realtime"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> TabularEndpoint: """Deploy the tabular foundation model to an inference endpoint. @@ -685,7 +676,7 @@ def deploy( wait=wait, inference_mode="realtime", environment=environment, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, **kwargs, ) return TabularEndpoint( @@ -707,8 +698,14 @@ def _build_predictor_fit_args(self, hyperparameters: Optional[Dict[str, Any]] = def _load_results( self, *, include_predict: bool, predict_only: bool = False ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]: - # The training container writes [pred, _proba...]; regression has only the pred column. raw = self._backend.get_fit_predict_results() + return self._format_results(raw, include_predict=include_predict, predict_only=predict_only) + + @staticmethod + def _format_results( + raw: pd.DataFrame, *, include_predict: bool, predict_only: bool = False + ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]: + # The training container writes [pred, _proba...]; regression has only the pred column. pred, pred_proba = split_pred_and_pred_proba(raw) if pred_proba is None: # regression: proba mirrors pred, matching TabularPredictor.predict_proba pred_proba = pred @@ -732,6 +729,7 @@ def predict( framework_version: str = "latest", image_uri: Optional[str] = None, wait: bool = True, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> Union[pd.Series, JobPredictionFuture]: """ @@ -763,9 +761,12 @@ def predict( wait If True, block and return the predictions. If False, return a :class:`JobPredictionFuture` immediately — call ``.result()`` on it later to retrieve the predictions. + backend_overrides + Raw provider request fields. This prediction path uses a SageMaker training job; + the valid key is ``"create_training_job"``. See :meth:`autogluon.cloud.TabularCloudPredictor.fit`. **kwargs Additional job arguments accepted by :meth:`autogluon.cloud.TabularCloudPredictor.fit` (e.g. - ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``, ``sagemaker_overrides``). + ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``). Returns ------- @@ -783,12 +784,12 @@ def predict( framework_version=framework_version, image_uri=image_uri, wait=wait, + backend_overrides=backend_overrides, **kwargs, ) if not wait: - return JobPredictionFuture( - job=self._backend._fit_job, - result_loader=lambda: self._load_results(include_predict=True, predict_only=True), + return self._backend.get_prediction_future( + result_transform=partial(self._format_results, include_predict=True, predict_only=True), ) pred, _ = result return pred @@ -807,6 +808,7 @@ def predict_proba( framework_version: str = "latest", image_uri: Optional[str] = None, wait: bool = True, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series], JobPredictionFuture]: """ @@ -839,9 +841,11 @@ def predict_proba( Custom Docker image URI for the container. wait If True, block and return the result. If False, return a :class:`JobPredictionFuture` immediately. + backend_overrides + Raw provider request fields under ``"create_training_job"``. See :meth:`predict`. **kwargs Additional job arguments accepted by :meth:`autogluon.cloud.TabularCloudPredictor.fit` (e.g. - ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``, ``sagemaker_overrides``). + ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``). Returns ------- @@ -870,14 +874,14 @@ def predict_proba( instance_type=instance_type, image_uri=image_uri, wait=wait, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, extra_tags={"autogluon-cloud-model-id": self.model_id}, **kwargs, ) if not wait: - return JobPredictionFuture( - job=self._backend._fit_job, - result_loader=lambda: self._load_results(include_predict=include_predict), + return self._backend.get_prediction_future( + result_transform=partial(self._format_results, include_predict=include_predict), ) return self._load_results(include_predict=include_predict) diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index a9980b87..74aee3d7 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -7,7 +7,7 @@ from abc import ABC, abstractmethod from datetime import datetime from pathlib import Path -from typing import Any, Dict, List, Literal, Optional, Tuple, Union +from typing import Any, Dict, Literal, Optional, Tuple, Union import boto3 import pandas as pd @@ -21,6 +21,7 @@ from ..backend.backend import Backend from ..backend.backend_factory import BackendFactory from ..backend.constant import SAGEMAKER +from ..config import SageMakerConfig from ..utils.aws_utils import resolve_cloud_output_path from ..utils.sagemaker_api import reject_legacy_kwargs from ..utils.utils import safe_unpack_archive @@ -36,11 +37,7 @@ def __init__( self, local_output_path: Optional[str] = None, cloud_output_path: Optional[str] = None, - backend: str = SAGEMAKER, - role: Optional[str] = None, - vpc_config: Optional[Dict[str, List[str]]] = None, - kms_key: Optional[str] = None, - tags: Optional[Dict[str, str]] = None, + backend: Union[str, SageMakerConfig] = SAGEMAKER, verbosity: int = 2, ) -> None: """ @@ -63,21 +60,10 @@ def __init__( * ``None`` (default) — use the bucket saved in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` / :func:`autogluon.cloud.register`) and append a timestamped subfolder. Raises if no bucket is configured. - backend: str, default = "sagemaker" - The backend to use. Currently only "sagemaker" is supported. - SageMaker backend supports training, deploying and batch inference on Amazon SageMaker. Only single instance training is supported. - 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 the role of the current AWS identity. - vpc_config: Optional[Dict[str, List[str]]], default = None - VPC to run training jobs, models and batch transform jobs in, as - ``{"subnets": ["subnet-..."], "security_group_ids": ["sg-..."]}``. - kms_key: Optional[str], default = None - KMS key ID/ARN used to encrypt S3 outputs and ML storage volumes of all created SageMaker resources. - Note that SageMaker rejects volume KMS keys for instance types with local NVMe storage (e.g. ``ml.g5``). - tags: Optional[Dict[str, str]], default = None - Tags added to every SageMaker resource created by this predictor, e.g. ``{"team": "forecasting"}``. + backend: Union[str, SageMakerConfig], default = "sagemaker" + Backend name or reusable :class:`~autogluon.cloud.SageMakerConfig` with region, execution role, + networking, encryption and tags. ``"sagemaker"`` uses default settings. + Only single instance training is supported. 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). @@ -87,21 +73,17 @@ def __init__( self.verbosity = verbosity cloud_logger = logging.getLogger("autogluon.cloud") set_logger_verbosity(self.verbosity, logger=cloud_logger) + config = BackendFactory.resolve_config(backend) + if config.name not in self.backend_map: + raise ValueError(f"Unsupported backend {config.name!r}. Supported backends: {sorted(self.backend_map)}.") self.local_output_path = self._setup_local_output_path(local_output_path) - if backend in ("ray", "ray_aws"): - raise ValueError("The Ray backend was removed in AutoGluon-Cloud v0.7.0. Use backend='sagemaker' instead.") - if backend not in self.backend_map: - raise ValueError(f"Unsupported backend {backend!r}. Supported backends: {sorted(self.backend_map)}.") - self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend) + self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=config.name) self.backend: Backend = BackendFactory.get_backend( - backend=self.backend_map[backend], + backend=self.backend_map[config.name], local_output_path=self.local_output_path, cloud_output_path=self.cloud_output_path, predictor_type=self.predictor_type, - role=role, - vpc_config=vpc_config, - kms_key=kms_key, - tags=tags, + config=config, ) @property @@ -194,7 +176,7 @@ def fit( environment: Optional[Dict[str, str]] = None, use_spot_instances: bool = False, max_wait: Optional[int] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> CloudPredictor: """ @@ -244,7 +226,7 @@ def fit( max_wait: Optional[int], default = None Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires ``use_spot_instances=True``. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + 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 snake_case (as in ``sagemaker.core.shapes``), which are deep-merged over the request @@ -295,7 +277,7 @@ def fit( environment=environment, use_spot_instances=use_spot_instances, max_wait=max_wait, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) @@ -409,7 +391,7 @@ def deploy( inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> None: """ Deploy a predictor to an inference endpoint. @@ -451,7 +433,7 @@ def deploy( Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: ``"create_model"``, ``"production_variant"`` (the endpoint config's single production variant), @@ -474,7 +456,7 @@ def deploy( inference_mode=inference_mode, inference_config=inference_config, environment=environment, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, ) def attach_endpoint(self, endpoint: str) -> None: @@ -588,7 +570,7 @@ def predict( persist: bool = True, save_path: Optional[str] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Batch inference. @@ -634,7 +616,7 @@ def predict( If `persist` is `False`, file would first be downloaded to this path and then removed. environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: ``"create_model"`` and ``"create_transform_job"``, e.g. @@ -660,7 +642,7 @@ def predict( persist=persist, save_path=save_path, environment=environment, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, ) @reject_legacy_kwargs @@ -680,7 +662,7 @@ def predict_proba( persist: bool = True, save_path: Optional[str] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = 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 @@ -729,7 +711,7 @@ def predict_proba( If `persist` is `False`, file would first be downloaded to this path and then removed. environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: ``"create_model"`` and ``"create_transform_job"``, e.g. @@ -758,7 +740,7 @@ def predict_proba( persist=persist, save_path=save_path, environment=environment, - sagemaker_overrides=sagemaker_overrides, + 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 694cd1bf..15cdf736 100644 --- a/src/autogluon/cloud/predictor/tabular_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/tabular_cloud_predictor.py @@ -57,7 +57,7 @@ def fit_predict( environment: Optional[Dict[str, str]] = None, use_spot_instances: bool = False, max_wait: Optional[int] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ Fit and predict in a single SageMaker training job. @@ -100,7 +100,7 @@ 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``. - environment, use_spot_instances, max_wait, sagemaker_overrides: + environment, use_spot_instances, max_wait, backend_overrides: Same as in :meth:`fit`. Returns @@ -127,7 +127,7 @@ def fit_predict( environment=environment, use_spot_instances=use_spot_instances, max_wait=max_wait, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, ) if result is None: # wait=False return None @@ -155,7 +155,7 @@ def fit_predict_proba( environment: Optional[Dict[str, str]] = None, use_spot_instances: bool = False, max_wait: Optional[int] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = 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. @@ -195,7 +195,7 @@ 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``. - environment, use_spot_instances, max_wait, sagemaker_overrides: + environment, use_spot_instances, max_wait, backend_overrides: Same as in :meth:`fit`. Returns @@ -225,7 +225,7 @@ def fit_predict_proba( environment=environment, use_spot_instances=use_spot_instances, max_wait=max_wait, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) diff --git a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py index ada55bb8..15c5c6bb 100644 --- a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py @@ -57,7 +57,7 @@ def fit( environment: Optional[Dict[str, str]] = None, use_spot_instances: bool = False, max_wait: Optional[int] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> TimeSeriesCloudPredictor: """ @@ -117,7 +117,7 @@ def fit( Whether to train on managed spot instances. max_wait: Optional[int], default = None Maximum seconds to wait for spot capacity plus training time. Requires ``use_spot_instances=True``. - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None Raw ``CreateTrainingJob`` request fields under the ``"create_training_job"`` key. See :meth:`TabularCloudPredictor.fit` for details. @@ -165,7 +165,7 @@ def fit( environment=environment, use_spot_instances=use_spot_instances, max_wait=max_wait, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) @@ -238,7 +238,7 @@ def predict( persist: bool = True, save_path: Optional[str] = None, environment: Optional[Dict[str, str]] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ Predict using SageMaker batch transform. @@ -279,7 +279,7 @@ def predict( To be noticed, the function won't return immediately because there are some preparations needed prior transform. image_uri: Optional[str], default = None Custom inference container image. If set, ``framework_version`` is ignored. - download, persist, save_path, environment, sagemaker_overrides: + download, persist, save_path, environment, backend_overrides: Same as in :meth:`TabularCloudPredictor.predict`. """ return self.backend.predict( @@ -297,7 +297,7 @@ def predict( persist=persist, save_path=save_path, environment=environment, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, ) def predict_proba( @@ -331,7 +331,7 @@ def fit_predict( environment: Optional[Dict[str, str]] = None, use_spot_instances: bool = False, max_wait: Optional[int] = None, - sagemaker_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ Fit and predict in a single SageMaker training job. @@ -384,7 +384,7 @@ 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. - environment, use_spot_instances, max_wait, sagemaker_overrides: + environment, use_spot_instances, max_wait, backend_overrides: Same as in :meth:`fit`. Returns @@ -414,7 +414,7 @@ def fit_predict( environment=environment, use_spot_instances=use_spot_instances, max_wait=max_wait, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) diff --git a/src/autogluon/cloud/utils/aws_utils.py b/src/autogluon/cloud/utils/aws_utils.py index b2ba81b5..0f24f77a 100644 --- a/src/autogluon/cloud/utils/aws_utils.py +++ b/src/autogluon/cloud/utils/aws_utils.py @@ -29,7 +29,7 @@ def _resolve_sagemaker_region() -> Optional[str]: return entry.region -def resolve_execution_role(role: Optional[str], backend_name: str) -> str: +def resolve_execution_role(role: Optional[str], backend_name: str, *, session: Optional[Session] = None) -> str: """Resolve the SageMaker execution role ARN. Resolution order: @@ -46,7 +46,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 get_execution_role() + return get_execution_role(sagemaker_session=session) def resolve_cloud_output_path(path: Optional[str], backend_name: str) -> Optional[str]: @@ -130,12 +130,13 @@ def setup_sagemaker_session( connect_timeout: int = 60, read_timeout: int = 60, retries: Optional[dict] = None, + region: Optional[str] = None, **kwargs, ): """ Setup a sagemaker session with a given configuration - Region resolution (only when ``boto_session`` is not provided): read from + Region resolution (only when ``boto_session`` is not provided): use ``region``, then read from ``~/.autogluon/cloud.yaml`` if set, otherwise fall back to the boto3 default chain (env vars, shared config). Raises if no region can be resolved at all. @@ -144,6 +145,9 @@ def setup_sagemaker_session( boto_session Pre-built ``boto3.Session`` to wrap. If provided, region resolution is skipped and the session is used as-is. + region + Explicit AWS region. Takes precedence over saved configuration and the boto3 default region. + Ignored when ``boto_session`` is provided. config A botocore.Config object providing the intended configuration https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html @@ -174,7 +178,7 @@ def setup_sagemaker_session( retries = {"max_attempts": 20} config = Config(connect_timeout=connect_timeout, read_timeout=read_timeout, retries=retries, **kwargs) if boto_session is None: - boto_session = boto3.Session(region_name=_resolve_sagemaker_region()) + boto_session = boto3.Session(region_name=region or _resolve_sagemaker_region()) if boto_session.region_name is None: raise ValueError( "AWS region could not be resolved. Set it in `~/.autogluon/cloud.yaml` (e.g. via " diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index 193f7137..64f62c28 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -12,21 +12,21 @@ logger = logging.getLogger(__name__) -# Requests that each method sends, i.e. the valid `sagemaker_overrides` keys. `production_variant` is the single +# Requests that each method sends, i.e. the valid `backend_overrides` keys. `production_variant` is the single # variant inside `create_endpoint_config.production_variants`. 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") _REMOVED_KWARGS = { - "backend_kwargs": "named arguments, the constructor's `vpc_config` / `kms_key` / `tags`, or `sagemaker_overrides`", + "backend_kwargs": "named arguments, `backend=SageMakerConfig(...)`, or `backend_overrides`", "custom_image_uri": "`image_uri`", - "autogluon_sagemaker_estimator_kwargs": "`sagemaker_overrides={'create_training_job': ...}`", - "fit_kwargs": "`sagemaker_overrides={'create_training_job': ...}`", - "model_kwargs": "`environment` or `sagemaker_overrides={'create_model': ...}`", - "deploy_kwargs": "`sagemaker_overrides={'production_variant': ..., 'create_endpoint_config': ...}`", - "transformer_kwargs": "`sagemaker_overrides={'create_transform_job': ...}`", - "transform_kwargs": "`sagemaker_overrides={'create_transform_job': ...}`", + "autogluon_sagemaker_estimator_kwargs": "`backend_overrides={'create_training_job': ...}`", + "fit_kwargs": "`backend_overrides={'create_training_job': ...}`", + "model_kwargs": "`environment` or `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': ...}`", } @@ -50,7 +50,7 @@ def check_override_keys(overrides: Optional[Mapping[str, Any]], allowed_keys: It overrides = dict(overrides or {}) unknown = sorted(set(overrides) - set(allowed_keys)) if unknown: - raise ValueError(f"Unsupported `sagemaker_overrides` key(s) {unknown}. Valid keys: {list(allowed_keys)}.") + raise ValueError(f"Unsupported `backend_overrides` key(s) {unknown}. Valid keys: {list(allowed_keys)}.") return overrides diff --git a/tests/unittests/general/test_aws_utils.py b/tests/unittests/general/test_aws_utils.py index 28faba2b..5abd8a65 100644 --- a/tests/unittests/general/test_aws_utils.py +++ b/tests/unittests/general/test_aws_utils.py @@ -9,7 +9,7 @@ CloudConfig, save_config, ) -from autogluon.cloud.utils.aws_utils import resolve_cloud_output_path, resolve_execution_role +from autogluon.cloud.utils.aws_utils import resolve_cloud_output_path, resolve_execution_role, setup_sagemaker_session @pytest.fixture(autouse=True) @@ -83,6 +83,25 @@ def test_falls_back_to_env_when_backend_missing_in_config(): 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(sagemaker_session=session) + + +@pytest.mark.parametrize("explicit_region", [None, "eu-west-1"]) +def test_session_region_uses_explicit_setting_before_saved_config(explicit_region): + import boto3 + + _save_role_in_config("sagemaker", "role") + with mock.patch("autogluon.cloud.utils.aws_utils.boto3.Session", wraps=boto3.Session) as session_cls: + session = setup_sagemaker_session(region=explicit_region) + expected_region = explicit_region or "us-east-1" + session_cls.assert_called_once_with(region_name=expected_region) + assert session.boto_region_name == expected_region + + def _save_bucket_in_config(backend_name: str, bucket: str) -> None: save_config( CloudConfig( diff --git a/tests/unittests/general/test_backend_config.py b/tests/unittests/general/test_backend_config.py new file mode 100644 index 00000000..cf97bf3d --- /dev/null +++ b/tests/unittests/general/test_backend_config.py @@ -0,0 +1,209 @@ +"""Shared backend configuration and job-bound foundation-model results, without AWS calls.""" + +import pickle +from pathlib import Path +from unittest import mock + +import pandas as pd +import pytest + +from autogluon.cloud import FoundationModel, SageMakerConfig, TabularCloudPredictor, TimeSeriesCloudPredictor +from autogluon.cloud.backend.backend_factory import BackendFactory +from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend +from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend +from autogluon.cloud.backend.timeseries_sagemaker_backend import TimeSeriesSagemakerBackend +from autogluon.cloud.job.sagemaker_job import SageMakerFitJob + +SB = "autogluon.cloud.backend.sagemaker_backend" + + +@pytest.fixture(autouse=True) +def stub_aws(monkeypatch): + session_factory = mock.Mock( + side_effect=lambda *, region=None: mock.MagicMock(boto_region_name=region or "us-east-1") + ) + monkeypatch.setattr(f"{SB}.setup_sagemaker_session", session_factory) + monkeypatch.setattr( + f"{SB}.resolve_execution_role", + lambda role, backend_name, session=None: role or "arn:aws:iam::0:role/default", + ) + monkeypatch.setattr("autogluon.cloud.utils.aws_utils._s3_prefix_has_objects", lambda *_: False) + return session_factory + + +@pytest.mark.parametrize( + "predictor_cls,model_id", + [(TabularCloudPredictor, "mitra-classifier"), (TimeSeriesCloudPredictor, "chronos-2")], +) +def test_shared_config_creates_independent_predictor_and_model_backends(tmp_path, predictor_cls, model_id): + config = SageMakerConfig( + region="eu-west-1", + role_arn="arn:aws:iam::0:role/custom", + vpc_config={"subnets": ["subnet-1"], "security_group_ids": ["sg-1"]}, + output_kms_key="output-key", + tags={"team": "forecasting"}, + ) + predictor = predictor_cls(backend=config, cloud_output_path="s3://b/predictor", local_output_path=str(tmp_path)) + model = FoundationModel(model_id, backend=config, cloud_output_path="s3://b/model") + + assert predictor.backend.config == model._backend.config == config + assert predictor.backend is not model._backend + assert predictor.backend._fit_job is not model._backend._fit_job + assert predictor.backend.sagemaker_session is not model._backend.sagemaker_session + + predictor.backend.config.tags["team"] = "other" + predictor.backend.config.vpc_config["subnets"].append("subnet-2") + predictor.backend.attach_endpoint("endpoint") + assert config.tags == model._backend.config.tags == {"team": "forecasting"} + assert ( + config.vpc_config + == model._backend.config.vpc_config + == { + "subnets": ["subnet-1"], + "security_group_ids": ["sg-1"], + } + ) + assert model._backend.endpoint_name is None + + +def test_default_backend_name_uses_the_same_config_resolver(tmp_path): + predictor = TabularCloudPredictor(local_output_path=str(tmp_path), cloud_output_path="s3://b/run") + model = FoundationModel("chronos-2", cloud_output_path="s3://b/model") + assert predictor.backend.config == model._backend.config + assert predictor.backend.config.role_arn == "arn:aws:iam::0:role/default" + assert predictor.backend.config.region == "us-east-1" + + +@pytest.mark.parametrize("backend", ["unknown", "ray", "ray_aws"]) +def test_invalid_backend_rejected_by_both_entry_points(tmp_path, backend): + with pytest.raises(ValueError): + TabularCloudPredictor(backend=backend, local_output_path=str(tmp_path)) + with pytest.raises(ValueError): + FoundationModel("chronos-2", backend=backend) + + +def test_backend_resolver_rejects_untyped_dict(): + with pytest.raises(TypeError, match="SageMakerConfig"): + BackendFactory.resolve_config({"region": "us-east-1"}) + + +def test_predictor_reload_keeps_the_original_region(tmp_path, stub_aws): + predictor = TabularCloudPredictor( + backend=SageMakerConfig(region="eu-west-1"), + cloud_output_path="s3://b/run", + local_output_path=str(tmp_path), + ) + restored = pickle.loads(pickle.dumps(predictor)) + assert restored.backend.config == predictor.backend.config + stub_aws.assert_called_with(region="eu-west-1") + + +def test_real_training_submissions_keep_job_objects_and_inputs_separate(tmp_path, monkeypatch): + """Later predictions must not overwrite the first job's handle, config or data channels.""" + backend = TabularSagemakerBackend( + local_output_path=str(tmp_path), cloud_output_path="s3://b/run", predictor_type="tabular" + ) + uploads = {} + + def upload(path, bucket, key_prefix): + uri = f"s3://{bucket}/{key_prefix}/{Path(path).name}" + if Path(path).is_file(): + uploads[uri] = Path(path).read_bytes() + return uri + + def run(job, training_job_request, framework_version, wait): + job._job_name = training_job_request["training_job_name"] + job.request = training_job_request + + backend.sagemaker_session.upload_data.side_effect = upload + monkeypatch.setattr(SageMakerFitJob, "run", run) + monkeypatch.setattr(SageMakerFitJob, "_get_job_status", lambda self: "Completed") + monkeypatch.setattr(SB + ".upload_training_code", lambda **kwargs: "s3://b/code") + monkeypatch.setattr(backend, "_load_fit_predict_results", lambda job: pd.DataFrame({"job": [job.job_name]})) + futures, requests = [], [] + for name, value in [("first", 1), ("second", 2)]: + backend.fit( + predictor_init_args={"label": "y"}, + predictor_fit_args={}, + data_channels={"train_data": pd.DataFrame({"x": [value], "y": [0]})}, + job_name=name, + image_uri="example.com/autogluon:train", + wait=False, + extra_ag_args={"predict_after_fit": True, "save_predictor": False}, + ) + futures.append(backend.get_prediction_future()) + requests.append(backend._fit_job.request) + assert [future.job_name for future in futures] == ["first", "second"] + assert [future.result()["job"].iloc[0] for future in futures] == ["first", "second"] + channels = [ + { + channel["channel_name"]: channel["data_source"]["s3_data_source"]["s3_uri"] + for channel in request["input_data_config"] + } + for request in requests + ] + assert channels[0]["ag_args"] != channels[1]["ag_args"] + assert channels[0]["train_data"] != channels[1]["train_data"] + assert b"/first/predictions.csv" in uploads[channels[0]["ag_args"]] + assert b"/second/predictions.csv" in uploads[channels[1]["ag_args"]] + assert uploads[channels[0]["train_data"]] != uploads[channels[1]["train_data"]] + + +@pytest.mark.parametrize( + "model_id,operation,include_predict", + [ + ("chronos-2", "predict", None), + ("mitra-classifier", "predict", None), + ("mitra-classifier", "predict_proba", True), + ("mitra-classifier", "predict_proba", False), + ], +) +def test_async_model_results_stay_bound_to_the_submitted_job(monkeypatch, model_id, operation, include_predict): + model = FoundationModel(model_id, cloud_output_path="s3://b/model") + frames = [ + pd.DataFrame({"target": ["a"], "a_proba": [0.8], "b_proba": [0.2]}), + pd.DataFrame({"target": ["b"], "a_proba": [0.1], "b_proba": [0.9]}), + ] + jobs = [] + + def submit(self, **kwargs): + assert kwargs["backend_overrides"] == {"create_training_job": {"retry_strategy": {}}} + job = mock.Mock(job_name=f"job-{len(jobs)}", completed=True) + job.frame = frames[len(jobs)] + jobs.append(job) + self._fit_job = job + + monkeypatch.setattr(TabularSagemakerBackend, "fit", submit) + monkeypatch.setattr(TimeSeriesSagemakerBackend, "fit", submit) + monkeypatch.setattr(SagemakerBackend, "_load_fit_predict_results", lambda self, job: job.frame) + kwargs = {"wait": False, "backend_overrides": {"create_training_job": {"retry_strategy": {}}}} + if model_id == "chronos-2": + kwargs["data"] = pd.DataFrame({"target": [1.0]}) + else: + kwargs.update( + train_data=pd.DataFrame({"feature": [1], "target": ["a"]}), + test_data=pd.DataFrame({"feature": [2]}), + label="target", + ) + if include_predict is not None: + kwargs["include_predict"] = include_predict + + first = getattr(model, operation)(**kwargs) + second = getattr(model, operation)(**kwargs) + assert first.job_name == "job-0" + assert second.job_name == "job-1" + first_result, second_result = first.result(), second.result() + if model_id == "chronos-2": + pd.testing.assert_frame_equal(first_result, frames[0]) + pd.testing.assert_frame_equal(second_result, frames[1]) + elif operation == "predict": + assert first_result.tolist() == ["a"] + assert second_result.tolist() == ["b"] + else: + if include_predict: + first_pred, first_result = first_result + second_pred, second_result = second_result + assert first_pred.tolist() == ["a"] + assert second_pred.tolist() == ["b"] + assert first_result["a"].tolist() == [0.8] + assert second_result["a"].tolist() == [0.1] diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 19b189e7..941342f4 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -7,6 +7,7 @@ import pandas as pd import pytest +from autogluon.cloud import SageMakerConfig from autogluon.cloud.model import FoundationModel @@ -19,7 +20,7 @@ def _stub_aws(monkeypatch): ) monkeypatch.setattr( "autogluon.cloud.backend.backend_factory.BackendFactory.get_backend", - lambda **kwargs: mock.MagicMock(role_arn="arn:aws:iam::0:role/stub"), + lambda **kwargs: mock.MagicMock(role_arn="arn:aws:iam::0:role/stub", config=kwargs["config"]), ) @@ -46,7 +47,7 @@ def test_to_dict_excludes_runtime_context(): fm = FoundationModel( "chronos-2", cloud_output_path="s3://my-bucket/runs/", - role="arn:aws:iam::0:role/runtime", + backend=SageMakerConfig(role_arn="arn:aws:iam::0:role/runtime"), ) d = fm.to_dict() assert "role" not in d @@ -213,7 +214,8 @@ def test_cache_model_artifact_uploads_with_version_metadata(monkeypatch): """On cache miss, upload_file runs with the version metadata key — that's the cache-invalidation contract.""" from autogluon.cloud.version import __version__ - fm = FoundationModel("chronos-2", cloud_output_path="s3://b") + backend_config = SageMakerConfig(region="eu-west-1", output_kms_key="output-key", tags={"team": "ts"}) + fm = FoundationModel("chronos-2", cloud_output_path="s3://b", backend=backend_config) s3 = mock.MagicMock() fm._backend.sagemaker_session.boto_session.client.return_value = s3 monkeypatch.setattr("autogluon.cloud.model.foundation_model._s3_head_or_none", lambda *_: None) @@ -226,9 +228,12 @@ def test_cache_model_artifact_uploads_with_version_metadata(monkeypatch): new_fm = fm.cache_model_artifact("s3://b/cache") assert new_fm.model_artifact_uri == "s3://b/cache/chronos-2/model.tar.gz" + assert new_fm._backend.config == fm._backend.config == backend_config s3.upload_file.assert_called_once() metadata = s3.upload_file.call_args.kwargs["ExtraArgs"]["Metadata"] assert metadata == {"autogluon-cloud-version": __version__} + assert s3.upload_file.call_args.kwargs["ExtraArgs"]["SSEKMSKeyId"] == "output-key" + assert s3.upload_file.call_args.kwargs["ExtraArgs"]["ServerSideEncryption"] == "aws:kms" def test_cache_model_artifact_raises_on_stale_version_without_overwrite(): diff --git a/tests/unittests/general/test_inference_modes.py b/tests/unittests/general/test_inference_modes.py index cfdee967..9ea80108 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -4,6 +4,7 @@ import pytest +from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend GPU_IMAGE_URI = "123456789012.dkr.ecr.us-east-1.amazonaws.com/autogluon:1.6-cu133-amzn2023" @@ -30,8 +31,10 @@ def deploy_requests(): ) backend._fit_job = None # deploy a serve-script tarball, not a fit-job artifact - def run(**kwargs): + def run(backend_config=None, **kwargs): backend.endpoint_name = None # allow re-deploy across cases + if backend_config is not None: + backend.config = backend_config backend.deploy(endpoint_name="ep", entry_point="stub.py", **kwargs) return { "model": model_cls.create.call_args.kwargs, @@ -74,7 +77,7 @@ def test_when_inference_ami_is_overridden_then_override_wins(deploy_requests): deploy_requests( instance_type="ml.g4dn.xlarge", image_uri=GPU_IMAGE_URI, - sagemaker_overrides={"production_variant": {"inference_ami_version": "custom-ami"}}, + backend_overrides={"production_variant": {"inference_ami_version": "custom-ami"}}, ) ) assert variant["inference_ami_version"] == "custom-ami" @@ -105,5 +108,21 @@ def test_when_environment_given_then_it_reaches_the_container(deploy_requests): def test_when_override_targets_training_job_then_deploy_rejects_it(deploy_requests): - with pytest.raises(ValueError, match="Unsupported `sagemaker_overrides` key"): - deploy_requests(sagemaker_overrides={"create_training_job": {}}) + with pytest.raises(ValueError, match="Unsupported `backend_overrides` key"): + deploy_requests(backend_overrides={"create_training_job": {}}) + + +@pytest.mark.parametrize( + ("inference_mode", "volume_key"), + [("realtime", None), ("realtime", "volume-key"), ("serverless", "volume-key")], +) +def test_endpoint_volume_encryption_is_independent_of_output_encryption(deploy_requests, inference_mode, volume_key): + requests = deploy_requests( + inference_mode=inference_mode, + backend_config=SageMakerConfig(output_kms_key="output-key", volume_kms_key=volume_key), + ) + config = requests["endpoint_config"] + if inference_mode == "realtime" and volume_key is not None: + assert config["kms_key_id"] == volume_key + else: + assert "kms_key_id" not in config diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index 5c626f39..8d0def76 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -3,6 +3,7 @@ import pandas as pd import pytest +from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend from autogluon.cloud.utils.dlc_utils import infer_sagemaker_ami_version @@ -49,13 +50,13 @@ def test_infer_realtime_ami_ignores_unsupported_or_already_compatible_instance_f @pytest.mark.parametrize( - ("sagemaker_overrides", "expected"), + ("backend_overrides", "expected"), [ (None, "al2-ami-sagemaker-batch-gpu-535"), ({"create_transform_job": {"transform_resources": {"transform_ami_version": "custom-ami"}}}, "custom-ami"), ], ) -def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(sagemaker_overrides, expected): +def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(backend_overrides, expected): 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"), @@ -69,6 +70,7 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(sag local_output_path="/tmp/test", cloud_output_path="s3://bucket/run", predictor_type="tabular", + config=SageMakerConfig(output_kms_key="output-key"), ) backend._fit_job = mock.MagicMock() backend._predict( @@ -80,9 +82,11 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(sag wait=False, download=False, persist=False, - sagemaker_overrides=sagemaker_overrides, + backend_overrides=backend_overrides, ) request = job_cls.return_value.run.call_args.kwargs["transform_job_request"] assert request["transform_resources"]["transform_ami_version"] == expected assert request["transform_resources"]["instance_type"] == "ml.g4dn.xlarge" + assert request["transform_output"]["kms_key_id"] == "output-key" + assert "volume_kms_key_id" not in request["transform_resources"] diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index 3a1c9f49..8fdc5d3e 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -4,6 +4,7 @@ import pandas as pd import pytest +from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend from autogluon.cloud.utils.sagemaker_api import check_override_keys, deep_merge, reject_legacy_kwargs from autogluon.cloud.utils.sagemaker_core_workarounds import bind_core_session @@ -32,7 +33,7 @@ def fit(**kwargs): assert fit(image_uri="x") == {"image_uri": "x"} with pytest.raises(TypeError, match="Use `image_uri` instead"): fit(custom_image_uri="x") - with pytest.raises(TypeError, match="sagemaker_overrides"): + with pytest.raises(TypeError, match="backend_overrides"): fit(backend_kwargs={}) @@ -58,19 +59,19 @@ def fit_request(tmp_path): 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(backend_kwargs=None, **fit_kwargs): + def run(backend_config=None, **fit_kwargs): backend = TabularSagemakerBackend( local_output_path=str(tmp_path), cloud_output_path="s3://bucket/run", predictor_type="tabular", - **(backend_kwargs or {}), + config=backend_config, ) - backend._fit_job = mock.MagicMock() backend.fit( predictor_init_args={"label": "y"}, predictor_fit_args={}, @@ -79,7 +80,7 @@ def run(backend_kwargs=None, **fit_kwargs): image_uri="example.com/autogluon:train", **fit_kwargs, ) - return backend._fit_job.run.call_args.kwargs["training_job_request"] + return fit_job_cls.return_value.run.call_args.kwargs["training_job_request"] yield run @@ -101,19 +102,20 @@ def test_fit_builds_script_mode_training_job(fit_request): def test_fit_applies_infra_settings_spot_and_overrides(fit_request): request = fit_request( - backend_kwargs={ - "vpc_config": {"subnets": ["s-1"], "security_group_ids": ["sg-1"]}, - "kms_key": "kms-1", - "tags": {"team": "ts"}, - }, + backend_config=SageMakerConfig( + vpc_config={"subnets": ["s-1"], "security_group_ids": ["sg-1"]}, + output_kms_key="output-key", + volume_kms_key="volume-key", + tags={"team": "ts"}, + ), timeout=3600, environment={"FOO": "bar"}, use_spot_instances=True, - sagemaker_overrides={"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}, + backend_overrides={"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}, ) assert request["vpc_config"] == {"subnets": ["s-1"], "security_group_ids": ["sg-1"]} - assert request["output_data_config"]["kms_key_id"] == "kms-1" - assert request["resource_config"]["volume_kms_key_id"] == "kms-1" + assert request["output_data_config"]["kms_key_id"] == "output-key" + assert request["resource_config"]["volume_kms_key_id"] == "volume-key" assert {"key": "team", "value": "ts"} in request["tags"] assert request["environment"] == {"FOO": "bar"} assert request["enable_managed_spot_training"] is True @@ -123,7 +125,16 @@ def test_fit_applies_infra_settings_spot_and_overrides(fit_request): def test_fit_rejects_malformed_vpc_config(fit_request): with pytest.raises(ValueError, match="security_group_ids"): - fit_request(backend_kwargs={"vpc_config": {"subnets": ["s-1"]}}) + fit_request(backend_config=SageMakerConfig(vpc_config={"subnets": ["s-1"]})) + + +def test_output_encryption_does_not_set_volume_key_on_nvme_instance(fit_request): + request = fit_request( + backend_config=SageMakerConfig(output_kms_key="output-key"), + instance_type="ml.g5.xlarge", + ) + assert request["output_data_config"]["kms_key_id"] == "output-key" + assert "volume_kms_key_id" not in request["resource_config"] def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): diff --git a/tests/unittests/general/test_tabular_foundation_model_predict.py b/tests/unittests/general/test_tabular_foundation_model_predict.py index e6ee928b..74966fd3 100644 --- a/tests/unittests/general/test_tabular_foundation_model_predict.py +++ b/tests/unittests/general/test_tabular_foundation_model_predict.py @@ -33,10 +33,21 @@ def _stub_aws(monkeypatch): "autogluon.cloud.model.foundation_model.resolve_cloud_output_path", lambda path, backend_name: path or "s3://stub/output", ) - monkeypatch.setattr( - "autogluon.cloud.backend.backend_factory.BackendFactory.get_backend", - lambda **kwargs: mock.MagicMock(role_arn="arn:aws:iam::0:role/stub"), - ) + + def make_backend(**kwargs): + backend = mock.MagicMock(role_arn="arn:aws:iam::0:role/stub", config=kwargs["config"]) + + def get_prediction_future(*, result_transform=None): + def load(): + raw = backend.get_fit_predict_results() + return raw if result_transform is None else result_transform(raw) + + return JobPredictionFuture(job=backend._fit_job, result_loader=load) + + backend.get_prediction_future.side_effect = get_prediction_future + return backend + + monkeypatch.setattr("autogluon.cloud.backend.backend_factory.BackendFactory.get_backend", make_backend) def _make_fm(model_id="mitra-classifier", result=CLASSIFICATION_FRAME): From 7e43a3df1da93d7f82fb2abaf5cdea930b97406a Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 10:12:25 +0000 Subject: [PATCH 06/16] Call SageMaker through each backend's boto3 clients instead of sagemaker-core resources sagemaker-core resource classes route every call through a process-wide client that ignores the session argument, which breaks per-object region/credentials. Send requests through the session's own sagemaker / sagemaker-runtime clients instead. - Build requests in SageMaker API / boto3 PascalCase; backend_overrides fields use the same format as the AWS API reference. Override keys stay the boto3 method names. - SageMakerConfig.vpc_config and inference_config keep their snake_case keys and are mapped explicitly; unknown keys raise before any resource is created. - Wait for endpoints with botocore's endpoint_in_service waiter. - Remove the sagemaker-core session-binding and acronym-serialization workarounds. - Unit tests validate generated requests against botocore's service model. --- docs/tutorials/setup.md | 5 +- .../cloud/backend/sagemaker_backend.py | 204 ++++++++++-------- .../cloud/endpoint/tabular_endpoint.py | 2 +- .../cloud/endpoint/timeseries_endpoint.py | 2 +- src/autogluon/cloud/job/sagemaker_job.py | 74 +++---- .../cloud/predictor/cloud_predictor.py | 22 +- src/autogluon/cloud/utils/job_logs.py | 4 +- src/autogluon/cloud/utils/sagemaker_api.py | 46 ++-- .../cloud/utils/sagemaker_core_workarounds.py | 31 --- src/autogluon/cloud/utils/tag_utils.py | 2 +- tests/conftest.py | 16 ++ .../unittests/general/test_backend_config.py | 10 +- .../general/test_foundation_model.py | 8 +- .../unittests/general/test_inference_modes.py | 77 ++++--- tests/unittests/general/test_sagemaker_ami.py | 15 +- tests/unittests/general/test_sagemaker_api.py | 124 ++++++----- tests/unittests/general/test_tags.py | 2 +- 17 files changed, 322 insertions(+), 322 deletions(-) delete mode 100644 src/autogluon/cloud/utils/sagemaker_core_workarounds.py diff --git a/docs/tutorials/setup.md b/docs/tutorials/setup.md index 2add942b..6f13dced 100644 --- a/docs/tutorials/setup.md +++ b/docs/tutorials/setup.md @@ -127,7 +127,8 @@ individual `fit()`, `predict()`, and `deploy()` calls. ## Advanced provider settings Use `backend_overrides` for SageMaker request fields without a named argument. It maps request names -to fields in the snake_case format used by `sagemaker-core`: +(the boto3 SageMaker client methods) to request fields in the PascalCase format of the +[SageMaker API](https://docs.aws.amazon.com/sagemaker/latest/APIReference/Welcome.html) and boto3: ```python predictions = model.predict( @@ -135,7 +136,7 @@ predictions = model.predict( prediction_length=24, backend_overrides={ "create_training_job": { - "retry_strategy": {"maximum_retry_attempts": 2}, + "RetryStrategy": {"MaximumRetryAttempts": 2}, }, }, ) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 429cab33..9d36754e 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -11,8 +11,6 @@ import pandas as pd from botocore.exceptions import ClientError from sagemaker.core.common_utils import sagemaker_timestamp, unique_name_from_base -from sagemaker.core.resources import Endpoint, EndpointConfig, Model -from sagemaker.core.shapes import VpcConfig from autogluon.common.loaders import load_pd from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix @@ -45,7 +43,6 @@ delete_endpoint, invoke_endpoint, ) -from ..utils.sagemaker_core_workarounds import bind_core_session from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer from ..utils.tag_utils import build_tags, to_request_tags from ..utils.utils import ( @@ -60,6 +57,12 @@ logger = logging.getLogger(__name__) SAGEMAKER_MODEL_SERVER_WORKERS = "SAGEMAKER_MODEL_SERVER_WORKERS" +_VPC_CONFIG_FIELDS = {"subnets": "Subnets", "security_group_ids": "SecurityGroupIds"} +_SERVERLESS_CONFIG_FIELDS = { + "memory_size_in_mb": "MemorySizeInMB", + "max_concurrency": "MaxConcurrency", + "provisioned_concurrency": "ProvisionedConcurrency", +} def _reject_local_mode(instance_type: Optional[str]) -> None: @@ -72,17 +75,33 @@ def _reject_local_mode(instance_type: Optional[str]) -> None: def _s3_channel(channel_name: str, s3_uri: str) -> Dict[str, Any]: return { - "channel_name": channel_name, - "data_source": { - "s3_data_source": { - "s3_data_type": "S3Prefix", - "s3_uri": s3_uri, - "s3_data_distribution_type": "FullyReplicated", + "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()} + + +def _vpc_request(vpc_config: Dict[str, List[str]]) -> Dict[str, List[str]]: + """``SageMakerConfig.vpc_config`` in the ``VpcConfig`` request format.""" + missing = sorted(set(_VPC_CONFIG_FIELDS) - set(vpc_config)) + if missing: + raise ValueError(f"`vpc_config` is missing required key(s) {missing}.") + return _to_request_fields(vpc_config, _VPC_CONFIG_FIELDS, "vpc_config") + + class SagemakerBackend(Backend): name = SAGEMAKER @@ -105,7 +124,7 @@ def __init__( raise TypeError("`config` must be a SageMakerConfig.") config = copy.deepcopy(config or SageMakerConfig()) if config.vpc_config is not None: - VpcConfig(**config.vpc_config) # fail before creating a session or resolving the role + _vpc_request(config.vpc_config) # fail before creating a session or resolving the role self.sagemaker_session = setup_sagemaker_session(region=config.region) try: self.role_arn = resolve_execution_role( @@ -127,15 +146,9 @@ def _realtime_serializer(self): return AutoGluonSerializer() def _resolve_tags(self, extra_tags: Optional[Dict[str, str]] = None) -> List[Dict[str, str]]: - """Tags for a created SageMaker resource, in sagemaker-core request format: default + extra + user tags.""" + """Tags for a created SageMaker resource, in SageMaker API request format: default + extra + user tags.""" return to_request_tags(build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=self.config.tags)) - @property - def _boto_session(self): - boto_session = self.sagemaker_session.boto_session - bind_core_session(boto_session) - return boto_session - def attach_job(self, job_name: str) -> None: """ Attach to a existing training job. @@ -261,7 +274,7 @@ def fit( Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires ``use_spot_instances=True``. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Raw ``CreateTrainingJob`` request fields (sagemaker-core snake_case names) under the + 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 @@ -345,42 +358,42 @@ def fit( s3_uri_prefix=f"{self.cloud_output_path}/code/{job_name}/source", ) - stopping_condition: Dict[str, Any] = {"max_runtime_in_seconds": timeout} + stopping_condition: Dict[str, Any] = {"MaxRuntimeInSeconds": timeout} if use_spot_instances: - stopping_condition["max_wait_time_in_seconds"] = max_wait or timeout + stopping_condition["MaxWaitTimeInSeconds"] = max_wait or timeout request: Dict[str, Any] = { - "training_job_name": job_name, - "role_arn": self.role_arn, - "algorithm_specification": { - "training_image": resolve_image_uri( + "TrainingJobName": job_name, + "RoleArn": self.role_arn, + "AlgorithmSpecification": { + "TrainingImage": resolve_image_uri( image_uri, framework_version, py_version, self._region, "training", instance_type ), - "training_input_mode": "File", + "TrainingInputMode": "File", }, - "hyper_parameters": training_script_hyperparameters( + "HyperParameters": training_script_hyperparameters( entry_point=entry_point, submit_directory=code_uri, job_name=job_name, region=self._region ), - "input_data_config": [_s3_channel(name, uri) for name, uri in inputs.items()], - "output_data_config": {"s3_output_path": self.cloud_output_path + "/model"}, - "resource_config": { - "instance_type": instance_type, - "instance_count": instance_count, - "volume_size_in_gb": volume_size, + "InputDataConfig": [_s3_channel(name, uri) for name, uri in inputs.items()], + "OutputDataConfig": {"S3OutputPath": self.cloud_output_path + "/model"}, + "ResourceConfig": { + "InstanceType": instance_type, + "InstanceCount": instance_count, + "VolumeSizeInGB": volume_size, }, - "stopping_condition": stopping_condition, - "profiler_config": {"disable_profiler": True}, - "tags": self._resolve_tags(extra_tags), + "StoppingCondition": stopping_condition, + "ProfilerConfig": {"DisableProfiler": True}, + "Tags": self._resolve_tags(extra_tags), } if environment: - request["environment"] = dict(environment) + request["Environment"] = dict(environment) if use_spot_instances: - request["enable_managed_spot_training"] = True + request["EnableManagedSpotTraining"] = True if self.config.vpc_config is not None: - request["vpc_config"] = self.config.vpc_config + request["VpcConfig"] = _vpc_request(self.config.vpc_config) if self.config.output_kms_key is not None: - request["output_data_config"]["kms_key_id"] = self.config.output_kms_key + request["OutputDataConfig"]["KmsKeyId"] = self.config.output_kms_key if self.config.volume_kms_key is not None: - request["resource_config"]["volume_kms_key_id"] = self.config.volume_kms_key + request["ResourceConfig"]["VolumeKmsKeyId"] = self.config.volume_kms_key request = deep_merge(request, overrides.get("create_training_job", {})) self._fit_job = SageMakerFitJob(session=self.sagemaker_session) @@ -404,22 +417,22 @@ def _create_model( **script_mode_environment(entry_point, self._region), } request: Dict[str, Any] = { - "model_name": model_name, - "primary_container": { - "image": image_uri, - "model_data_url": model_data, - "environment": container_environment, + "ModelName": model_name, + "PrimaryContainer": { + "Image": image_uri, + "ModelDataUrl": model_data, + "Environment": container_environment, }, - "execution_role_arn": self.role_arn, - "tags": tags, + "ExecutionRoleArn": self.role_arn, + "Tags": tags, } if self.config.vpc_config is not None: - request["vpc_config"] = self.config.vpc_config + request["VpcConfig"] = _vpc_request(self.config.vpc_config) request = deep_merge(request, overrides.get("create_model", {})) logger.log(20, "Creating inference model...") - Model.create(**request, session=self._boto_session, region=self._region) + self.sagemaker_session.sagemaker_client.create_model(**request) logger.log(20, "Inference model created successfully") - return request["model_name"] + return request["ModelName"] def _prepare_model_data( self, @@ -505,7 +518,7 @@ def deploy( environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Raw request fields (sagemaker-core snake_case names) deep-merged over the requests built by + 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 @@ -516,7 +529,7 @@ 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 - Serverless overrides forwarded to the production variant's ``serverless_config`` + Serverless overrides forwarded to the production variant's ``ServerlessConfig`` (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). repack: bool, default = True Whether to download ``predictor_path``, inject the serve script, and re-upload it. Set to False when @@ -528,6 +541,12 @@ def deploy( "There is an endpoint already attached. Either detach it with `detach` or clean it up with `cleanup_deployment`" ) 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" @@ -606,46 +625,45 @@ def deploy( overrides=overrides, ) - variant: Dict[str, Any] = {"variant_name": "AllTraffic", "model_name": model_name} + variant: Dict[str, Any] = {"VariantName": "AllTraffic", "ModelName": model_name} if inference_mode == "realtime": - variant["instance_type"] = instance_type - variant["initial_instance_count"] = initial_instance_count + variant["InstanceType"] = instance_type + variant["InitialInstanceCount"] = initial_instance_count if volume_size: - variant["volume_size_in_gb"] = volume_size + variant["VolumeSizeInGB"] = volume_size inference_ami_version = infer_sagemaker_ami_version(image_uri, instance_type, image_scope="inference") if inference_ami_version is not None: - variant["inference_ami_version"] = inference_ami_version + variant["InferenceAmiVersion"] = inference_ami_version elif inference_mode == "serverless": - preset = {"memory_size_in_mb": 4096, "max_concurrency": 5} - variant["serverless_config"] = {**preset, **(inference_config or {})} + variant["ServerlessConfig"] = serverless_config else: raise ValueError(f"Unsupported inference_mode={inference_mode!r}") variant = deep_merge(variant, overrides.get("production_variant", {})) endpoint_config_request: Dict[str, Any] = { - "endpoint_config_name": endpoint_name, - "production_variants": [variant], - "tags": tags, + "EndpointConfigName": endpoint_name, + "ProductionVariants": [variant], + "Tags": tags, } if self.config.volume_kms_key is not None and inference_mode == "realtime": - endpoint_config_request["kms_key_id"] = self.config.volume_kms_key + endpoint_config_request["KmsKeyId"] = self.config.volume_kms_key endpoint_config_request = deep_merge(endpoint_config_request, overrides.get("create_endpoint_config", {})) endpoint_request = deep_merge( { - "endpoint_name": endpoint_name, - "endpoint_config_name": endpoint_config_request["endpoint_config_name"], - "tags": tags, + "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})") - boto_session = self._boto_session - EndpointConfig.create(**endpoint_config_request, session=boto_session, region=self._region) - endpoint = Endpoint.create(**endpoint_request, session=boto_session, region=self._region) - self.endpoint_name = endpoint_request["endpoint_name"] + client = self.sagemaker_session.sagemaker_client + client.create_endpoint_config(**endpoint_config_request) + client.create_endpoint(**endpoint_request) + self.endpoint_name = endpoint_request["EndpointName"] if wait: - endpoint.wait_for_status("InService") + 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/ and upload it.""" @@ -659,7 +677,7 @@ def cleanup_deployment(self) -> None: Delete endpoint, endpoint configuration and deployed model """ assert self.endpoint_name is not None, "No deployed endpoint detected" - delete_endpoint(self.endpoint_name, self._boto_session) + delete_endpoint(self.endpoint_name, self.sagemaker_session) self.endpoint_name = None def attach_endpoint(self, endpoint: str) -> None: @@ -874,7 +892,7 @@ def predict( environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Raw request fields (sagemaker-core snake_case names) deep-merged over the requests built by + 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 @@ -1171,7 +1189,7 @@ def _predict_real_time(self, test_data, accept, split_pred_proba=True, inference test_data = AutoGluonSerializationWrapper(data=test_data, inference_kwargs=inference_kwargs) prediction = invoke_endpoint( self.endpoint_name, - self._boto_session, + self.sagemaker_session, test_data, serializer=self._realtime_serializer(), deserializer=PandasDeserializer(), @@ -1329,34 +1347,34 @@ def _predict( ) transform_input: Dict[str, Any] = { - "data_source": {"s3_data_source": {"s3_data_type": "S3Prefix", "s3_uri": test_input}}, - "content_type": content_type, + "DataSource": {"S3DataSource": {"S3DataType": "S3Prefix", "S3Uri": test_input}}, + "ContentType": content_type, } if split_type is not None: - transform_input["split_type"] = split_type - transform_output: Dict[str, Any] = {"s3_output_path": output_path + "/results", "accept": accept} + transform_input["SplitType"] = split_type + transform_output: Dict[str, Any] = {"S3OutputPath": output_path + "/results", "Accept": accept} if assemble_with is not None: - transform_output["assemble_with"] = assemble_with - transform_resources: Dict[str, Any] = {"instance_type": instance_type, "instance_count": instance_count} + transform_output["AssembleWith"] = assemble_with + transform_resources: Dict[str, Any] = {"InstanceType": instance_type, "InstanceCount": instance_count} transform_ami_version = infer_sagemaker_ami_version(image_uri, instance_type, image_scope="transform") if transform_ami_version is not None: - transform_resources["transform_ami_version"] = transform_ami_version + transform_resources["TransformAmiVersion"] = transform_ami_version if self.config.output_kms_key is not None: - transform_output["kms_key_id"] = self.config.output_kms_key + transform_output["KmsKeyId"] = self.config.output_kms_key if self.config.volume_kms_key is not None: - transform_resources["volume_kms_key_id"] = self.config.volume_kms_key + transform_resources["VolumeKmsKeyId"] = self.config.volume_kms_key request = { - "transform_job_name": job_name, - "model_name": model_name, - "transform_input": transform_input, - "transform_output": transform_output, - "transform_resources": transform_resources, - "batch_strategy": batch_strategy, + "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. - "max_payload_in_mb": 6, + "MaxPayloadInMB": 6, # The maximum number of HTTP requests made to each individual transform container at one time. - "max_concurrent_transforms": 1, - "tags": tags, + "MaxConcurrentTransforms": 1, + "Tags": tags, } request = deep_merge(request, overrides.get("create_transform_job", {})) @@ -1367,7 +1385,7 @@ def _predict( pred, pred_proba = None, None if download: results_path = self.download_predict_results(save_path=save_path) - accept = request["transform_output"].get("accept") + accept = request["TransformOutput"].get("Accept") if accept == "application/x-parquet": results = pd.read_parquet(results_path) elif accept == "text/csv": diff --git a/src/autogluon/cloud/endpoint/tabular_endpoint.py b/src/autogluon/cloud/endpoint/tabular_endpoint.py index 47c6c4bb..6b049503 100644 --- a/src/autogluon/cloud/endpoint/tabular_endpoint.py +++ b/src/autogluon/cloud/endpoint/tabular_endpoint.py @@ -30,7 +30,7 @@ def __init__(self, endpoint_name: str, session: Optional[boto3.Session] = None): ``boto3.Session`` used to invoke and delete the endpoint. If ``None``, the default ambient session is used. """ self._endpoint_name = endpoint_name - self._session = setup_sagemaker_session(boto_session=session).boto_session + self._session = setup_sagemaker_session(boto_session=session) @property def endpoint_name(self) -> str: diff --git a/src/autogluon/cloud/endpoint/timeseries_endpoint.py b/src/autogluon/cloud/endpoint/timeseries_endpoint.py index fb72e0e0..2a83e330 100644 --- a/src/autogluon/cloud/endpoint/timeseries_endpoint.py +++ b/src/autogluon/cloud/endpoint/timeseries_endpoint.py @@ -32,7 +32,7 @@ def __init__(self, endpoint_name: str, session: Optional[boto3.Session] = None): ``boto3.Session`` used to invoke and delete the endpoint. If ``None``, the default ambient session is used. """ self._endpoint_name = endpoint_name - self._session = setup_sagemaker_session(boto_session=session).boto_session + self._session = setup_sagemaker_session(boto_session=session) @property def endpoint_name(self) -> str: diff --git a/src/autogluon/cloud/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index edbea55e..ad4a88fb 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -2,12 +2,9 @@ from abc import abstractmethod from typing import Any, Dict, Optional, Union -from sagemaker.core.resources import Model, TrainingJob, TransformJob - 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_core_workarounds import bind_core_session from .remote_job import RemoteJob logger = logging.getLogger(__name__) @@ -52,8 +49,8 @@ def run(self, **kwargs): raise NotImplementedError @abstractmethod - def _describe(self): - """Return the sagemaker-core resource describing the job.""" + def _describe(self) -> Dict[str, Any]: + """Return the ``Describe*Job`` response for the job.""" raise NotImplementedError @abstractmethod @@ -68,12 +65,6 @@ def _get_output_path(self): def _get_hyperparameters(self): raise NotImplementedError - @property - def _boto_session(self): - boto_session = self.session.boto_session - bind_core_session(boto_session) - return boto_session - @property def job_name(self): return self._job_name @@ -136,7 +127,7 @@ def wait(self, logs: bool = True) -> str: ) if status != "Completed": logger.error( - f"SageMaker job {self.job_name} finished with status {status}: {self._describe().failure_reason}" + f"SageMaker job {self.job_name} finished with status {status}: {self._describe().get('FailureReason')}" ) return status @@ -145,7 +136,7 @@ def _wait_until_completed(self) -> None: status = self.wait(logs=True) if status != "Completed": raise RuntimeError( - f"SageMaker job {self.job_name} finished with status {status}: {self._describe().failure_reason}" + f"SageMaker job {self.job_name} finished with status {status}: {self._describe().get('FailureReason')}" ) def __getstate__(self): @@ -187,29 +178,25 @@ def info(self): ) return info - def _describe(self) -> TrainingJob: - return TrainingJob.get( - training_job_name=self.job_name, - session=self._boto_session, - region=self.session.boto_region_name, - ) + 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._describe().training_job_status + return self._describe()["TrainingJobStatus"] def _get_output_path(self): - return self._describe().model_artifacts.s3_model_artifacts + return self._describe()["ModelArtifacts"]["S3ModelArtifacts"] def _get_hyperparameters(self): if self.job_name: - return self._describe().hyper_parameters + return self._describe().get("HyperParameters") return None def get_input_channels(self) -> Dict[str, str]: """Map each input channel name of the training job to its S3 URI.""" return { - channel.channel_name: channel.data_source.s3_data_source.s3_uri - for channel in self._describe().input_data_config + channel["ChannelName"]: channel["DataSource"]["S3DataSource"]["S3Uri"] + for channel in self._describe()["InputDataConfig"] } def run( @@ -218,15 +205,11 @@ def run( framework_version: Optional[str], wait: bool, ): - """Create the training job from a ``TrainingJob.create`` request and optionally wait for it to finish.""" - job_name = training_job_request["training_job_name"] + """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: - TrainingJob.create( - **training_job_request, - session=self._boto_session, - region=self.session.boto_region_name, - ) + self.session.sagemaker_client.create_training_job(**training_job_request) self._job_name = job_name self._framework_version = framework_version if wait: @@ -256,42 +239,33 @@ def info(self): ) return info - def _describe(self) -> TransformJob: - return TransformJob.get( - transform_job_name=self.job_name, - session=self._boto_session, - region=self.session.boto_region_name, - ) + 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._describe().transform_job_status + return self._describe()["TransformJobStatus"] def _get_output_path(self): - return self._describe().transform_output.s3_output_path + "/" + self._output_filename + return self._describe()["TransformOutput"]["S3OutputPath"] + "/" + self._output_filename def _delete_model(self, model_name: str) -> None: - bind_core_session(self.session.boto_session) - Model(model_name=model_name).delete() + self.session.sagemaker_client.delete_model(ModelName=model_name) def run( self, transform_job_request: Dict[str, Any], wait: bool, ): - """Create the transform job from a ``TransformJob.create`` request. + """Create the transform job from a ``CreateTransformJob`` request. The SageMaker model referenced by the request 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["transform_job_name"] - model_name = transform_job_request["model_name"] + job_name = transform_job_request["TransformJobName"] + model_name = transform_job_request["ModelName"] try: logger.log(20, "Transforming") - TransformJob.create( - **transform_job_request, - session=self._boto_session, - region=self.session.boto_region_name, - ) + self.session.sagemaker_client.create_transform_job(**transform_job_request) self._job_name = job_name if wait: self._wait_until_completed() @@ -300,7 +274,7 @@ def run( self._delete_model(model_name) raise e - input_uri = transform_job_request["transform_input"]["data_source"]["s3_data_source"]["s3_uri"] + input_uri = transform_job_request["TransformInput"]["DataSource"]["S3DataSource"]["S3Uri"] self._output_filename = input_uri.split("/")[-1] + ".out" if wait: diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index 74aee3d7..06827d17 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -229,8 +229,8 @@ def fit( 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 snake_case (as in ``sagemaker.core.shapes``), which are deep-merged over the request - built by AutoGluon-Cloud, e.g. ``{"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}``. + 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}}}``. Returns ------- `CloudPredictor` object. Returns self. @@ -434,11 +434,11 @@ def deploy( environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case - (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + 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": {"model_data_download_timeout_in_seconds": 1200}}``. + ``{"production_variant": {"ModelDataDownloadTimeoutInSeconds": 1200}}``. """ if inference_mode == "serverless" and instance_type is not None: raise ValueError("`instance_type` must not be set when `inference_mode='serverless'`.") @@ -617,10 +617,10 @@ def predict( environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case - (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + 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": {"batch_strategy": "SingleRecord", "max_payload_in_mb": 20}}``. + ``{"create_transform_job": {"BatchStrategy": "SingleRecord", "MaxPayloadInMB": 20}}``. Returns ------- @@ -712,10 +712,10 @@ def predict_proba( environment: Optional[Dict[str, str]], default = None Environment variables set in the inference container. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Escape hatch for SageMaker settings without a dedicated argument: raw request fields in snake_case - (as in ``sagemaker.core.shapes``), deep-merged over the requests built by AutoGluon-Cloud. Valid keys: + 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": {"batch_strategy": "SingleRecord", "max_payload_in_mb": 20}}``. + ``{"create_transform_job": {"BatchStrategy": "SingleRecord", "MaxPayloadInMB": 20}}``. Returns ------- diff --git a/src/autogluon/cloud/utils/job_logs.py b/src/autogluon/cloud/utils/job_logs.py index b0594bd2..3ec13516 100644 --- a/src/autogluon/cloud/utils/job_logs.py +++ b/src/autogluon/cloud/utils/job_logs.py @@ -1,8 +1,6 @@ """Wait for SageMaker jobs while streaming their CloudWatch logs through the caller's own boto3 session. -sagemaker-core's ``TrainingJob.wait(logs=True)`` reads logs through a process-wide CloudWatch client built from the -default credential chain and region, ignoring the session the job was created with. Polling here uses the clients -we pass in, so logs always come from the job's account and region. +Polling uses the clients we pass in, so logs always come from the job's account and region. """ from __future__ import annotations diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index 64f62c28..b9c7fad3 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -1,19 +1,16 @@ -"""Helpers for building SageMaker API requests and calling them through sagemaker-core.""" +"""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, Dict, Iterable, Mapping, Optional -import boto3 -from sagemaker.core.resources import Endpoint, EndpointConfig, Model - -from .sagemaker_core_workarounds import bind_core_session +from sagemaker.core.helper.session_helper import Session logger = logging.getLogger(__name__) -# Requests that each method sends, i.e. the valid `backend_overrides` keys. `production_variant` is the single -# variant inside `create_endpoint_config.production_variants`. +# 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") @@ -67,7 +64,7 @@ def deep_merge(base: Mapping[str, Any], override: Mapping[str, Any]) -> Dict[str def invoke_endpoint( endpoint_name: str, - boto_session: boto3.Session, + session: Session, payload: Any, serializer, deserializer, @@ -75,27 +72,22 @@ def invoke_endpoint( accept: Optional[str] = None, ) -> Any: """Serialize ``payload``, invoke the endpoint, and deserialize the response.""" - bind_core_session(boto_session) - response = Endpoint(endpoint_name=endpoint_name).invoke( - body=serializer.serialize(payload), - content_type=content_type or serializer.CONTENT_TYPE, - accept=accept or ", ".join(deserializer.ACCEPT), - session=boto_session, - region=boto_session.region_name, + 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.content_type) + return deserializer.deserialize(response["Body"], response["ContentType"]) -def delete_endpoint(endpoint_name: str, boto_session: boto3.Session) -> None: +def delete_endpoint(endpoint_name: str, session: Session) -> None: """Delete an endpoint together with its endpoint config and models.""" - bind_core_session(boto_session) - region = boto_session.region_name - endpoint = Endpoint.get(endpoint_name=endpoint_name, session=boto_session, region=region) - endpoint_config = EndpointConfig.get( - endpoint_config_name=endpoint.endpoint_config_name, session=boto_session, region=region - ) + 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}") - endpoint.delete() - endpoint_config.delete() - for variant in endpoint_config.production_variants: - Model(model_name=variant.model_name).delete() + 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/sagemaker_core_workarounds.py b/src/autogluon/cloud/utils/sagemaker_core_workarounds.py deleted file mode 100644 index 38b82e73..00000000 --- a/src/autogluon/cloud/utils/sagemaker_core_workarounds.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Temporary workarounds for sagemaker-core bugs. Delete each one once fixed upstream.""" - -import boto3 -from sagemaker.core.utils import utils as core_utils -from sagemaker.core.utils.code_injection.shape_dag import SHAPE_DAG - - -def _register_acronym_field_names() -> None: - # sagemaker-core serializes nested shapes with a naive snake_case -> PascalCase conversion, so e.g. - # `memory_size_in_mb` is sent as `MemorySizeInMb` instead of `MemorySizeInMB`. Register the real API names. - for shape in SHAPE_DAG.values(): - for member in shape.get("members") or []: - snake = core_utils.pascal_to_snake(member["name"]) - if core_utils.snake_to_pascal(snake) != member["name"]: - core_utils.SPECIAL_SNAKE_TO_PASCAL_MAPPINGS.setdefault(snake, member["name"]) - - -_register_acronym_field_names() - - -def bind_core_session(boto_session: boto3.Session) -> None: - """Make sagemaker-core's process-wide client cache use ``boto_session``. - - The cache ignores the ``session`` argument of resource methods once it exists, so rebuild it whenever a different - session is requested. Not thread-safe across sessions. - """ - current = core_utils.SingletonMeta._instances.get(core_utils.SageMakerClient) - if current is not None and current.session is boto_session: - return - core_utils.SingletonMeta._instances.pop(core_utils.SageMakerClient, None) - core_utils.SageMakerClient(session=boto_session, region_name=boto_session.region_name) diff --git a/src/autogluon/cloud/utils/tag_utils.py b/src/autogluon/cloud/utils/tag_utils.py index 6a5d859c..abe5a313 100644 --- a/src/autogluon/cloud/utils/tag_utils.py +++ b/src/autogluon/cloud/utils/tag_utils.py @@ -25,4 +25,4 @@ def build_tags( def to_request_tags(tags: Dict[str, str]) -> List[Dict[str, str]]: """Convert ``{key: value}`` tags to the list format of SageMaker API requests.""" - return [{"key": key, "value": value} for key, value in tags.items()] + return [{"Key": key, "Value": value} for key, value in tags.items()] diff --git a/tests/conftest.py b/tests/conftest.py index ef760562..9d22c5f0 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_backend_config.py b/tests/unittests/general/test_backend_config.py index cf97bf3d..e7d42ac3 100644 --- a/tests/unittests/general/test_backend_config.py +++ b/tests/unittests/general/test_backend_config.py @@ -112,7 +112,7 @@ def upload(path, bucket, key_prefix): return uri def run(job, training_job_request, framework_version, wait): - job._job_name = training_job_request["training_job_name"] + job._job_name = training_job_request["TrainingJobName"] job.request = training_job_request backend.sagemaker_session.upload_data.side_effect = upload @@ -137,8 +137,8 @@ def run(job, training_job_request, framework_version, wait): assert [future.result()["job"].iloc[0] for future in futures] == ["first", "second"] channels = [ { - channel["channel_name"]: channel["data_source"]["s3_data_source"]["s3_uri"] - for channel in request["input_data_config"] + channel["ChannelName"]: channel["DataSource"]["S3DataSource"]["S3Uri"] + for channel in request["InputDataConfig"] } for request in requests ] @@ -167,7 +167,7 @@ def test_async_model_results_stay_bound_to_the_submitted_job(monkeypatch, model_ jobs = [] def submit(self, **kwargs): - assert kwargs["backend_overrides"] == {"create_training_job": {"retry_strategy": {}}} + assert kwargs["backend_overrides"] == {"create_training_job": {"RetryStrategy": {}}} job = mock.Mock(job_name=f"job-{len(jobs)}", completed=True) job.frame = frames[len(jobs)] jobs.append(job) @@ -176,7 +176,7 @@ def submit(self, **kwargs): monkeypatch.setattr(TabularSagemakerBackend, "fit", submit) monkeypatch.setattr(TimeSeriesSagemakerBackend, "fit", submit) monkeypatch.setattr(SagemakerBackend, "_load_fit_predict_results", lambda self, job: job.frame) - kwargs = {"wait": False, "backend_overrides": {"create_training_job": {"retry_strategy": {}}}} + kwargs = {"wait": False, "backend_overrides": {"create_training_job": {"RetryStrategy": {}}}} if model_id == "chronos-2": kwargs["data"] = pd.DataFrame({"target": [1.0]}) else: diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 941342f4..8409ecf5 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -256,11 +256,7 @@ def test_sagemaker_backend_skips_repack_when_repack_is_false(): 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}.bind_core_session"), mock.patch(f"{sb}.repack_model_with_serving_code") as repack, - mock.patch(f"{sb}.Model") as model_cls, - mock.patch(f"{sb}.EndpointConfig"), - mock.patch(f"{sb}.Endpoint"), mock.patch.object(SagemakerBackend, "_upload_predictor", side_effect=lambda p, _: p), ): backend = SagemakerBackend( @@ -277,5 +273,5 @@ def test_sagemaker_backend_skips_repack_when_repack_is_false(): ) repack.assert_not_called() - container = model_cls.create.call_args.kwargs["primary_container"] - assert container["model_data_url"] == "s3://bucket/cache/chronos-2/model.tar.gz" + 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 9ea80108..1d51f700 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -12,16 +12,12 @@ @pytest.fixture -def deploy_requests(): - """Run ``SagemakerBackend.deploy(...)`` with AWS calls and sagemaker-core resources mocked, - and return the requests that reached ``Model.create`` / ``EndpointConfig.create`` / ``Endpoint.create``.""" +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}.bind_core_session"), - mock.patch(f"{SB}.Model") as model_cls, - mock.patch(f"{SB}.EndpointConfig") as endpoint_config_cls, - mock.patch(f"{SB}.Endpoint") as endpoint_cls, mock.patch.object(SagemakerBackend, "_create_serve_script_tarball", return_value="s3://stub/m.tar.gz"), ): backend = SagemakerBackend( @@ -36,40 +32,46 @@ def run(backend_config=None, **kwargs): if backend_config is not None: backend.config = backend_config backend.deploy(endpoint_name="ep", entry_point="stub.py", **kwargs) - return { - "model": model_cls.create.call_args.kwargs, - "endpoint_config": endpoint_config_cls.create.call_args.kwargs, - "endpoint": endpoint_cls.create.call_args.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 _variant(requests): - (variant,) = requests["endpoint_config"]["production_variants"] + (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["instance_type"] == "ml.m5.xlarge" - assert variant["initial_instance_count"] == 2 - assert "serverless_config" not in variant + 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)["model_name"] == requests["model"]["model_name"] - assert requests["endpoint"]["endpoint_config_name"] == requests["endpoint_config"]["endpoint_config_name"] - assert requests["endpoint"]["endpoint_name"] == "ep" - environment = requests["model"]["primary_container"]["environment"] + 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_requests): variant = _variant(deploy_requests(instance_type="ml.g4dn.xlarge", image_uri=GPU_IMAGE_URI)) - assert variant["inference_ami_version"] == "al2023-ami-sagemaker-inference-gpu-4-1" + assert variant["InferenceAmiVersion"] == "al2023-ami-sagemaker-inference-gpu-4-1" def test_when_inference_ami_is_overridden_then_override_wins(deploy_requests): @@ -77,22 +79,39 @@ def test_when_inference_ami_is_overridden_then_override_wins(deploy_requests): deploy_requests( instance_type="ml.g4dn.xlarge", image_uri=GPU_IMAGE_URI, - backend_overrides={"production_variant": {"inference_ami_version": "custom-ami"}}, + backend_overrides={"production_variant": {"InferenceAmiVersion": "custom-ami"}}, ) ) - assert variant["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["serverless_config"] == {"memory_size_in_mb": 4096, "max_concurrency": 5} - assert "instance_type" not in variant + 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["serverless_config"]["memory_size_in_mb"] == 8192 - assert variant["serverless_config"]["max_concurrency"] == 5 # preset wins for keys the user didn't override + assert variant["ServerlessConfig"]["MemorySizeInMB"] == 8192 + assert variant["ServerlessConfig"]["MaxConcurrency"] == 5 # preset wins for keys the user didn't override + + +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() + + +@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_requests): @@ -102,7 +121,7 @@ def test_when_inference_mode_is_unknown_then_value_error_is_raised(deploy_reques def test_when_environment_given_then_it_reaches_the_container(deploy_requests): requests = deploy_requests(instance_type="ml.m5.xlarge", environment={"FOO": "bar"}) - environment = requests["model"]["primary_container"]["environment"] + environment = requests["model"]["PrimaryContainer"]["Environment"] assert environment["FOO"] == "bar" assert environment["SAGEMAKER_MODEL_SERVER_WORKERS"] == "1" @@ -123,6 +142,6 @@ def test_endpoint_volume_encryption_is_independent_of_output_encryption(deploy_r ) config = requests["endpoint_config"] if inference_mode == "realtime" and volume_key is not None: - assert config["kms_key_id"] == volume_key + assert config["KmsKeyId"] == volume_key else: - assert "kms_key_id" not in config + assert "KmsKeyId" not in config diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index 8d0def76..59970477 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -53,10 +53,12 @@ def test_infer_realtime_ami_ignores_unsupported_or_already_compatible_instance_f ("backend_overrides", "expected"), [ (None, "al2-ami-sagemaker-batch-gpu-535"), - ({"create_transform_job": {"transform_resources": {"transform_ami_version": "custom-ami"}}}, "custom-ami"), + ({"create_transform_job": {"TransformResources": {"TransformAmiVersion": "custom-ami"}}}, "custom-ami"), ], ) -def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(backend_overrides, expected): +def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value( + backend_overrides, expected, assert_valid_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"), @@ -86,7 +88,8 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value(bac ) request = job_cls.return_value.run.call_args.kwargs["transform_job_request"] - assert request["transform_resources"]["transform_ami_version"] == expected - assert request["transform_resources"]["instance_type"] == "ml.g4dn.xlarge" - assert request["transform_output"]["kms_key_id"] == "output-key" - assert "volume_kms_key_id" not in request["transform_resources"] + assert_valid_request("CreateTransformJob", request) + assert request["TransformResources"]["TransformAmiVersion"] == expected + assert request["TransformResources"]["InstanceType"] == "ml.g4dn.xlarge" + assert request["TransformOutput"]["KmsKeyId"] == "output-key" + assert "VolumeKmsKeyId" not in request["TransformResources"] diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index 8fdc5d3e..8fc97d79 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -1,13 +1,17 @@ from unittest import mock -import boto3 import pandas as pd import pytest from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend -from autogluon.cloud.utils.sagemaker_api import check_override_keys, deep_merge, reject_legacy_kwargs -from autogluon.cloud.utils.sagemaker_core_workarounds import bind_core_session +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" @@ -37,24 +41,37 @@ def fit(**kwargs): fit(backend_kwargs={}) -def test_bind_core_session_rebinds_when_session_changes(): - from sagemaker.core.utils.utils import SageMakerClient, SingletonMeta - - first = boto3.Session(region_name="us-east-1") - second = boto3.Session(region_name="eu-west-1") - bind_core_session(first) - cached = SingletonMeta._instances[SageMakerClient] - assert cached.session is first - bind_core_session(first) - assert SingletonMeta._instances[SageMakerClient] is cached # no rebuild for the same session - bind_core_session(second) - assert SageMakerClient().session is second - assert SageMakerClient().region_name == "eu-west-1" +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): - """Run ``SagemakerBackend.fit(...)`` with uploads mocked and return the ``TrainingJob.create`` request.""" +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"), @@ -80,24 +97,26 @@ def run(backend_config=None, **fit_kwargs): image_uri="example.com/autogluon:train", **fit_kwargs, ) - return fit_job_cls.return_value.run.call_args.kwargs["training_job_request"] + 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["training_job_name"] == "job" - assert request["algorithm_specification"]["training_image"] == "example.com/autogluon:train" - assert request["hyper_parameters"]["sagemaker_program"] == '"train.py"' - assert request["hyper_parameters"]["sagemaker_submit_directory"] == ( + 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"] == ( '"s3://bucket/run/code/job/source/sourcedir.tar.gz"' ) - assert request["input_data_config"][0]["channel_name"] == "train_data" - assert request["stopping_condition"] == {"max_runtime_in_seconds": 3600} - assert request["output_data_config"] == {"s3_output_path": "s3://bucket/run/model"} - assert {"key": "autogluon-cloud-module", "value": "tabular"} in request["tags"] - assert "vpc_config" not in request + assert request["InputDataConfig"][0]["ChannelName"] == "train_data" + assert request["StoppingCondition"] == {"MaxRuntimeInSeconds": 3600} + assert request["OutputDataConfig"] == {"S3OutputPath": "s3://bucket/run/model"} + assert {"Key": "autogluon-cloud-module", "Value": "tabular"} in request["Tags"] + assert "VpcConfig" not in request def test_fit_applies_infra_settings_spot_and_overrides(fit_request): @@ -111,21 +130,25 @@ def test_fit_applies_infra_settings_spot_and_overrides(fit_request): timeout=3600, environment={"FOO": "bar"}, use_spot_instances=True, - backend_overrides={"create_training_job": {"retry_strategy": {"maximum_retry_attempts": 2}}}, + backend_overrides={"create_training_job": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}, ) - assert request["vpc_config"] == {"subnets": ["s-1"], "security_group_ids": ["sg-1"]} - assert request["output_data_config"]["kms_key_id"] == "output-key" - assert request["resource_config"]["volume_kms_key_id"] == "volume-key" - assert {"key": "team", "value": "ts"} in request["tags"] - assert request["environment"] == {"FOO": "bar"} - assert request["enable_managed_spot_training"] is True - assert request["stopping_condition"]["max_wait_time_in_seconds"] == 3600 - assert request["retry_strategy"] == {"maximum_retry_attempts": 2} + assert request["VpcConfig"] == {"Subnets": ["s-1"], "SecurityGroupIds": ["sg-1"]} + assert request["OutputDataConfig"]["KmsKeyId"] == "output-key" + assert request["ResourceConfig"]["VolumeKmsKeyId"] == "volume-key" + assert {"Key": "team", "Value": "ts"} in request["Tags"] + assert request["Environment"] == {"FOO": "bar"} + assert request["EnableManagedSpotTraining"] is True + assert request["StoppingCondition"]["MaxWaitTimeInSeconds"] == 3600 + assert request["RetryStrategy"] == {"MaximumRetryAttempts": 2} -def test_fit_rejects_malformed_vpc_config(fit_request): - with pytest.raises(ValueError, match="security_group_ids"): - fit_request(backend_config=SageMakerConfig(vpc_config={"subnets": ["s-1"]})) +@pytest.mark.parametrize( + "vpc_config", + [{"subnets": ["s-1"]}, {"subnets": ["s-1"], "security_group_ids": ["sg-1"], "security_groups": ["sg-2"]}], +) +def test_fit_rejects_malformed_vpc_config(fit_request, vpc_config): + with pytest.raises(ValueError, match="security_group"): + fit_request(backend_config=SageMakerConfig(vpc_config=vpc_config)) def test_output_encryption_does_not_set_volume_key_on_nvme_instance(fit_request): @@ -133,8 +156,8 @@ def test_output_encryption_does_not_set_volume_key_on_nvme_instance(fit_request) backend_config=SageMakerConfig(output_kms_key="output-key"), instance_type="ml.g5.xlarge", ) - assert request["output_data_config"]["kms_key_id"] == "output-key" - assert "volume_kms_key_id" not in request["resource_config"] + assert request["OutputDataConfig"]["KmsKeyId"] == "output-key" + assert "VolumeKmsKeyId" not in request["ResourceConfig"] def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): @@ -144,17 +167,8 @@ def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): fit_request(max_wait=100) -def test_core_serializes_acronym_field_names_with_api_casing(): - from sagemaker.core.shapes import ProductionVariant - from sagemaker.core.utils.utils import serialize +def test_misspelled_override_field_fails_request_validation(fit_request): + from botocore.exceptions import ParamValidationError - variant = ProductionVariant( - variant_name="AllTraffic", - serverless_config={"memory_size_in_mb": 4096, "max_concurrency": 5}, - enable_ssm_access=True, - ) - assert serialize(variant) == { - "VariantName": "AllTraffic", - "ServerlessConfig": {"MemorySizeInMB": 4096, "MaxConcurrency": 5}, - "EnableSSMAccess": True, - } + with pytest.raises(ParamValidationError, match="RetryStrategyy"): + fit_request(backend_overrides={"create_training_job": {"RetryStrategyy": {"MaximumRetryAttempts": 2}}}) diff --git a/tests/unittests/general/test_tags.py b/tests/unittests/general/test_tags.py index 45199775..85192e9f 100644 --- a/tests/unittests/general/test_tags.py +++ b/tests/unittests/general/test_tags.py @@ -34,4 +34,4 @@ def test_when_disable_env_var_set_then_defaults_and_extras_are_skipped(monkeypat def test_to_request_tags_uses_api_field_names(): - assert to_request_tags({"Owner": "team"}) == [{"key": "Owner", "value": "team"}] + assert to_request_tags({"Owner": "team"}) == [{"Key": "Owner", "Value": "team"}] From ca921ebf4f1219f4824d84bccca50954c049e750 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 11:23:39 +0000 Subject: [PATCH 07/16] Drop the sagemaker-core dependency The remaining uses were thin helpers around boto3. Replace them with small local implementations so all AWS calls go through the backend's own boto3 session. - AwsSession wraps a boto3.Session with cached sagemaker / sagemaker-runtime / s3 clients plus upload_data / download_data. - get_execution_role derives the role from the caller's STS identity and resolves its path via IAM (falling back to service-role/ for SageMaker console roles); IAM users get an actionable error instead of a guess. - repack_model_with_serving_code repacks the model tarball directly. - Move sagemaker_timestamp / unique_name_from_base to utils.misc and drop the serializer/deserializer base classes. --- pyproject.toml | 5 - .../cloud/backend/sagemaker_backend.py | 7 +- src/autogluon/cloud/utils/ag_sagemaker.py | 34 +++-- src/autogluon/cloud/utils/aws_utils.py | 114 +++++++++++++-- src/autogluon/cloud/utils/deserializers.py | 5 +- src/autogluon/cloud/utils/misc.py | 14 ++ src/autogluon/cloud/utils/s3_utils.py | 7 +- src/autogluon/cloud/utils/sagemaker_api.py | 12 +- src/autogluon/cloud/utils/serializers.py | 13 +- tests/unittests/general/test_aws_session.py | 136 ++++++++++++++++++ tests/unittests/general/test_aws_utils.py | 2 +- tests/unittests/general/test_sagemaker_api.py | 4 +- 12 files changed, 296 insertions(+), 57 deletions(-) create mode 100644 tests/unittests/general/test_aws_session.py diff --git a/pyproject.toml b/pyproject.toml index ad0073cf..38d32b27 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,11 +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", - # SageMaker Python SDK v3 core layer (resources, shapes, session helpers). - # >=2.18: 2.13.0-2.17.0 sign all `sagemaker` API calls with default-chain credentials, ignoring the passed - # session (https://github.com/aws/sagemaker-python-sdk/issues/5986). - # We deliberately don't depend on the `sagemaker` meta-package, which also pulls in torch/mlflow via sagemaker-serve. - "sagemaker-core>=2.18.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/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 9d36754e..d4f7c44b 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -10,7 +10,6 @@ import pandas as pd from botocore.exceptions import ClientError -from sagemaker.core.common_utils import sagemaker_timestamp, unique_name_from_base from autogluon.common.loaders import load_pd from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix @@ -33,7 +32,7 @@ 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 -from ..utils.misc import MostRecentInsertedOrderedDict +from ..utils.misc import MostRecentInsertedOrderedDict, sagemaker_timestamp, unique_name_from_base from ..utils.sagemaker_api import ( BATCH_PREDICT_OVERRIDE_KEYS, DEPLOY_OVERRIDE_KEYS, @@ -1022,7 +1021,7 @@ def _load_fit_predict_results(self, job: SageMakerFitJob) -> 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, job: Optional[SageMakerFitJob] = None) -> Dict[str, Any]: @@ -1043,7 +1042,7 @@ def _download_ag_args_from_job(self, job: Optional[SageMakerFitJob] = None) -> D 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) diff --git a/src/autogluon/cloud/utils/ag_sagemaker.py b/src/autogluon/cloud/utils/ag_sagemaker.py index e2f55c4e..3c6e672e 100644 --- a/src/autogluon/cloud/utils/ag_sagemaker.py +++ b/src/autogluon/cloud/utils/ag_sagemaker.py @@ -14,9 +14,10 @@ from contextlib import contextmanager from typing import Dict, Iterator, Optional -from sagemaker.core.common_utils import repack_model +from autogluon.common.utils.s3_utils import s3_path_to_bucket_prefix from .dlc_utils import retrieve_image_uri +from .utils import safe_unpack_archive SOURCE_DIR_TARBALL_NAME = "sourcedir.tar.gz" @@ -46,8 +47,6 @@ def upload_training_code(entry_point: str, source_dir: Optional[str], sagemaker_ Returns the S3 URI of the uploaded tarball. """ - from autogluon.common.utils.s3_utils import s3_path_to_bucket_prefix - 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: @@ -107,20 +106,27 @@ def repack_model_with_serving_code( sagemaker_session, kms_key: Optional[str] = None, ) -> str: - """Replace ``code/`` inside the ``model_data`` tarball with ``entry_point`` + ``serving_utils/`` and upload it. + """Replace ``code/`` inside the S3 ``model_data`` tarball with ``entry_point`` + ``serving_utils/`` and upload it. Returns ``repacked_model_uri``. """ - with staged_serving_code(entry_point) as code_dir: - repack_model( - inference_script=entry_point, - source_directory=code_dir, - dependencies=[], - model_uri=model_data, - repacked_model_uri=repacked_model_uri, - sagemaker_session=sagemaker_session, - kms_key=kms_key, - ) + 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 0f24f77a..3cdf42a4 100644 --- a/src/autogluon/cloud/utils/aws_utils.py +++ b/src/autogluon/cloud/utils/aws_utils.py @@ -1,14 +1,16 @@ import logging -from typing import Optional +import os +import re +from typing import Any, Dict, List, Optional import boto3 from botocore.config import Config -from sagemaker.core.common_utils import sagemaker_timestamp -from sagemaker.core.helper.session_helper import Session, get_execution_role +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__) @@ -29,14 +31,107 @@ def _resolve_sagemaker_region() -> Optional[str]: return entry.region -def resolve_execution_role(role: Optional[str], backend_name: str, *, session: Optional[Session] = None) -> 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 " + "`backend=SageMakerConfig(role_arn=)`, 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_arn` 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.core.helper.session_helper.get_execution_role()`` (the caller's own role, e.g. on SageMaker). + 3. The role whose credentials ``session`` uses (e.g. the execution role inside SageMaker), see + :func:`get_execution_role`. """ if role: return role @@ -46,7 +141,7 @@ def resolve_execution_role(role: Optional[str], backend_name: str, *, session: O 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 get_execution_role(sagemaker_session=session) + return get_execution_role(session) def resolve_cloud_output_path(path: Optional[str], backend_name: str) -> Optional[str]: @@ -132,9 +227,9 @@ def setup_sagemaker_session( retries: Optional[dict] = None, region: Optional[str] = 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): use ``region``, then read from ``~/.autogluon/cloud.yaml`` if set, otherwise fall back to the boto3 default chain (env vars, @@ -185,5 +280,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 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 87f6fa8b..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.core.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/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 28d90d1b..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 -from sagemaker.core.helper.session_helper import Session 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 = 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 index b9c7fad3..6be4e3a1 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -5,7 +5,7 @@ import logging from typing import Any, Dict, Iterable, Mapping, Optional -from sagemaker.core.helper.session_helper import Session +from .aws_utils import AwsSession logger = logging.getLogger(__name__) @@ -28,7 +28,7 @@ def reject_legacy_kwargs(func): - """Raise an actionable ``TypeError`` for kwargs removed in the SageMaker SDK v3 migration.""" + """Raise an actionable ``TypeError`` for kwargs removed when AutoGluon-Cloud stopped using the SageMaker Python SDK.""" @functools.wraps(func) def wrapper(*args, **kwargs): @@ -64,7 +64,7 @@ def deep_merge(base: Mapping[str, Any], override: Mapping[str, Any]) -> Dict[str def invoke_endpoint( endpoint_name: str, - session: Session, + session: AwsSession, payload: Any, serializer, deserializer, @@ -75,13 +75,13 @@ def invoke_endpoint( 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), + 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: Session) -> None: +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) diff --git a/src/autogluon/cloud/utils/serializers.py b/src/autogluon/cloud/utils/serializers.py index e595fbc0..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.core.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/unittests/general/test_aws_session.py b/tests/unittests/general/test_aws_session.py new file mode 100644 index 00000000..224f1df3 --- /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="SageMakerConfig\\(role_arn="): + 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 5abd8a65..1c3e94f1 100644 --- a/tests/unittests/general/test_aws_utils.py +++ b/tests/unittests/general/test_aws_utils.py @@ -87,7 +87,7 @@ 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(sagemaker_session=session) + get_role.assert_called_once_with(session) @pytest.mark.parametrize("explicit_region", [None, "eu-west-1"]) diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index 8fc97d79..490f245b 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -56,8 +56,8 @@ 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",)) + 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", From 7b0778f00d8f2ea923a3384d37daefbc71943d31 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 11:51:34 +0000 Subject: [PATCH 08/16] Trim SageMaker SDK migration to the minimal API change Keep the public API as close to master as possible so the PR only swaps the SageMaker SDK for direct boto3 calls. - Restore the `role=` constructor arg and `custom_image_uri`; remove SageMakerConfig, vpc/kms/user tags and the environment / use_spot_instances / max_wait args (all reachable via backend_overrides). - Revert unrelated changes: prediction-future binding, per-job upload prefix, docs/API surface tweaks, internal backend refactors. --- README.md | 5 - docs/api.rst | 21 +- docs/api/setup.rst | 8 - docs/tutorials/setup.md | 77 +---- src/autogluon/cloud/__init__.py | 5 +- src/autogluon/cloud/backend/backend.py | 59 ++-- .../cloud/backend/backend_factory.py | 39 +-- .../cloud/backend/sagemaker_backend.py | 284 +++++++++--------- .../backend/timeseries_sagemaker_backend.py | 22 +- src/autogluon/cloud/config.py | 43 +-- src/autogluon/cloud/job/sagemaker_job.py | 7 - src/autogluon/cloud/model/foundation_model.py | 179 +++++------ .../cloud/predictor/cloud_predictor.py | 89 +++--- .../predictor/tabular_cloud_predictor.py | 36 +-- .../predictor/timeseries_cloud_predictor.py | 58 ++-- .../cloud/scripts/sagemaker_scripts/train.py | 2 +- src/autogluon/cloud/utils/ag_sagemaker.py | 21 +- src/autogluon/cloud/utils/aws_utils.py | 12 +- src/autogluon/cloud/utils/sagemaker_api.py | 5 +- src/autogluon/cloud/utils/tag_utils.py | 21 +- tests/conftest.py | 2 +- tests/unittests/general/test_aws_session.py | 2 +- tests/unittests/general/test_aws_utils.py | 14 +- .../unittests/general/test_backend_config.py | 209 ------------- .../general/test_foundation_model.py | 15 +- .../unittests/general/test_inference_modes.py | 32 +- tests/unittests/general/test_sagemaker_ami.py | 6 +- tests/unittests/general/test_sagemaker_api.py | 51 +--- .../test_tabular_foundation_model_predict.py | 19 +- tests/unittests/general/test_tags.py | 38 +-- tests/unittests/tabular/test_tabular.py | 14 +- tests/unittests/timeseries/test_timeseries.py | 16 +- 32 files changed, 412 insertions(+), 999 deletions(-) delete mode 100644 tests/unittests/general/test_backend_config.py diff --git a/README.md b/README.md index 448ddaae..41e8ae87 100644 --- a/README.md +++ b/README.md @@ -35,11 +35,6 @@ bootstrap() See the [Setup tutorial](https://auto.gluon.ai/cloud/stable/tutorials/setup.html) for the full walkthrough, including how to register an existing role and bucket instead. -Pass `backend=SageMakerConfig(...)` to a cloud predictor or foundation model to set the region, -execution role, VPC, output and volume encryption keys, and resource tags. The same config can be -reused across workflows; each object creates its own backend state. Instance sizes, container -environment variables, and inference modes remain named arguments to individual operations. - ## ⚙️ Train your own model Train an AutoGluon predictor on your data and serve it from a SageMaker endpoint — same API as local AutoGluon, all heavy lifting on AWS. Full walkthrough: [tabular](https://auto.gluon.ai/cloud/stable/tutorials/predictor-tabular.html), [time series](https://auto.gluon.ai/cloud/stable/tutorials/predictor-timeseries.html). diff --git a/docs/api.rst b/docs/api.rst index 042bc97d..f1d4483c 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -1,5 +1,3 @@ -:orphan: - API === @@ -9,53 +7,48 @@ API .. autosummary:: :toctree: api :template: custom_class.rst - - SageMakerConfig - -.. autosummary:: - :toctree: api - :template: custom_class.rst - - FoundationModel - -.. autosummary:: - :toctree: api - :template: custom_class.rst + :methods: TabularCloudPredictor .. autosummary:: :toctree: api :template: custom_class.rst + :methods: TabularFoundationModel .. autosummary:: :toctree: api :template: custom_class.rst + :methods: TabularEndpoint .. autosummary:: :toctree: api :template: custom_class.rst + :methods: TimeSeriesCloudPredictor .. autosummary:: :toctree: api :template: custom_class.rst + :methods: TimeSeriesFoundationModel .. autosummary:: :toctree: api :template: custom_class.rst + :methods: TimeSeriesEndpoint .. autosummary:: :toctree: api :template: custom_class.rst + :methods: MultiModalCloudPredictor diff --git a/docs/api/setup.rst b/docs/api/setup.rst index 1af4cfe2..d515ade9 100644 --- a/docs/api/setup.rst +++ b/docs/api/setup.rst @@ -5,14 +5,6 @@ Functions for managing AutoGluon-Cloud's AWS configuration. See the :doc:`Setup .. currentmodule:: autogluon.cloud -Use :class:`SageMakerConfig` for reusable settings passed to predictors and foundation models. - -.. autosummary:: - :toctree: . - :template: custom_class.rst - - SageMakerConfig - .. autosummary:: :toctree: . :nosignatures: diff --git a/docs/tutorials/setup.md b/docs/tutorials/setup.md index 6f13dced..2594235e 100644 --- a/docs/tutorials/setup.md +++ b/docs/tutorials/setup.md @@ -17,7 +17,7 @@ SageMaker compute and S3 storage are billed to your AWS account. AutoGluon-Cloud There are three ways to supply these resources — if you're unsure, start with option 1. -## 1. Create new resources with {func}`~autogluon.cloud.bootstrap` +### 1. Create new resources with {func}`~autogluon.cloud.bootstrap` Run this if you don't yet have an IAM role and S3 bucket set up for SageMaker. The role and bucket are provisioned on your account from a {repo-file}`CloudFormation template ` and saved under `~/.autogluon/cloud.yaml` for future calls. @@ -38,7 +38,7 @@ autogluon-cloud bootstrap ::: :::: -## 2. Use existing resources with {func}`~autogluon.cloud.register` +### 2. Use existing resources with {func}`~autogluon.cloud.register` Run this if you already have an IAM role and S3 bucket that you want to use with AutoGluon-Cloud. The values are saved under `~/.autogluon/cloud.yaml` for future calls. @@ -68,88 +68,21 @@ autogluon-cloud register \ The role must trust the `sagemaker.amazonaws.com` principal and grant the permissions AutoGluon-Cloud needs to run SageMaker jobs plus read/write access to your bucket — for example, a [SageMaker execution role](https://docs.aws.amazon.com/sagemaker/latest/dg/sagemaker-roles.html). For the exact set of permissions, see the {repo-file}`CloudFormation template ` that {func}`~autogluon.cloud.bootstrap` uses. The `region` where the jobs are executed must match the bucket's region. -## 3. Pass resources on each call +### 3. Pass resources on each call Skip the saved config entirely and provide the role and bucket every time you create a `CloudPredictor` or `FoundationModel`. ```python -from autogluon.cloud import SageMakerConfig, TabularCloudPredictor +from autogluon.cloud import TabularCloudPredictor predictor = TabularCloudPredictor( cloud_output_path="s3://my-autogluon-bucket/output", - backend=SageMakerConfig( - role_arn="arn:aws:iam::222222222222:role/MyAutoGluonRole", - region="us-east-1", - ), + role="arn:aws:iam::222222222222:role/MyAutoGluonRole", ) ``` Useful for one-off scripts or when you need different roles and buckets per call. The same role and bucket requirements as option 2 apply. -## Share backend settings across workflows - -{class}`~autogluon.cloud.SageMakerConfig` works with both cloud predictors and foundation models. It holds -the region, execution role, VPC, encryption keys, and resource tags. You can reuse it across objects; -each object gets its own backend, jobs, and endpoint state. - -```python -from autogluon.cloud import SageMakerConfig, TabularCloudPredictor, TimeSeriesFoundationModel - -backend = SageMakerConfig( - region="us-east-1", - role_arn="arn:aws:iam::222222222222:role/MyAutoGluonRole", - vpc_config={"subnets": ["subnet-..."], "security_group_ids": ["sg-..."]}, - output_kms_key="arn:aws:kms:us-east-1:222222222222:key/...", - tags={"team": "forecasting"}, -) - -predictor = TabularCloudPredictor( - backend=backend, - cloud_output_path="s3://my-autogluon-bucket/training", -) -model = TimeSeriesFoundationModel( - "chronos-2", - backend=backend, - cloud_output_path="s3://my-autogluon-bucket/inference", -) -``` - -The role and region you set explicitly take precedence over the saved configuration. Leaving them -unset uses the existing saved-config and AWS-identity fallbacks. `backend="sagemaker"` is shorthand -for `backend=SageMakerConfig()`. - -`output_kms_key` encrypts training artifacts, batch transform outputs, and repacked or cached model -artifacts in S3. `volume_kms_key` separately controls training, batch transform, and realtime endpoint -storage encryption; leave it unset for instances with local NVMe storage. Resource sizes, -container environment variables, spot training, and serverless settings remain arguments to the -individual `fit()`, `predict()`, and `deploy()` calls. - -## Advanced provider settings - -Use `backend_overrides` for SageMaker request fields without a named argument. It maps request names -(the boto3 SageMaker client methods) to request fields in the PascalCase format of the -[SageMaker API](https://docs.aws.amazon.com/sagemaker/latest/APIReference/Welcome.html) and boto3: - -```python -predictions = model.predict( - data, - prediction_length=24, - backend_overrides={ - "create_training_job": { - "RetryStrategy": {"MaximumRetryAttempts": 2}, - }, - }, -) -``` - -Foundation-model predictions and predictor training use `create_training_job`. Predictor batch -transform uses `create_model` and `create_transform_job`. Deployment accepts `create_model`, -`production_variant`, `create_endpoint_config`, and `create_endpoint`. - -Only requests used by the operation are accepted. Nested dictionaries merge recursively over the -generated request; other values, including lists, replace the generated value. Overrides take -precedence over backend settings and named arguments. - ## Managing the saved config Once {func}`~autogluon.cloud.bootstrap` or {func}`~autogluon.cloud.register` has written to `~/.autogluon/cloud.yaml`, you may want to check that the role and bucket are still healthy before a long training run, or clean everything up when you're done with AutoGluon-Cloud. Two helper commands cover both: diff --git a/src/autogluon/cloud/__init__.py b/src/autogluon/cloud/__init__.py index bf61c90e..ad6c6385 100644 --- a/src/autogluon/cloud/__init__.py +++ b/src/autogluon/cloud/__init__.py @@ -3,22 +3,19 @@ from autogluon.common.utils.log_utils import _add_stream_handler from .cloud_setup import bootstrap, register, status, teardown -from .config import SageMakerConfig from .endpoint.tabular_endpoint import TabularEndpoint from .endpoint.timeseries_endpoint import TimeSeriesEndpoint -from .model.foundation_model import FoundationModel, TabularFoundationModel, TimeSeriesFoundationModel +from .model.foundation_model import TabularFoundationModel, TimeSeriesFoundationModel from .predictor import MultiModalCloudPredictor, TabularCloudPredictor, TimeSeriesCloudPredictor _add_stream_handler() logging.getLogger(__name__).setLevel(logging.INFO) __all__ = [ - "FoundationModel", "MultiModalCloudPredictor", "TabularCloudPredictor", "TabularEndpoint", "TabularFoundationModel", - "SageMakerConfig", "TimeSeriesCloudPredictor", "TimeSeriesEndpoint", "TimeSeriesFoundationModel", diff --git a/src/autogluon/cloud/backend/backend.py b/src/autogluon/cloud/backend/backend.py index 63270723..60d98c25 100644 --- a/src/autogluon/cloud/backend/backend.py +++ b/src/autogluon/cloud/backend/backend.py @@ -3,13 +3,10 @@ import json import os from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union import pandas as pd -if TYPE_CHECKING: - from ..endpoint.prediction_future import JobPredictionFuture - def dumps_ag_args(config: Dict[str, Any]) -> str: """Serialize the remote-training config to JSON, raising a user-facing error on failure. @@ -42,14 +39,28 @@ def dumps_ag_args(config: Dict[str, Any]) -> str: class Backend(ABC): name = "backend" - def __init__( + def __init__(self, **kwargs) -> None: + self.initialize(**kwargs) + + @property + def cloud_output_path(self) -> str: + if not self._cloud_output_path: + raise ValueError( + "No `cloud_output_path` was provided and no bucket is configured in " + "~/.autogluon/cloud.yaml. Either pass `cloud_output_path=` explicitly, or run " + "`autogluon.cloud.bootstrap()` / `register(bucket=...)` once to persist a bucket." + ) + return self._cloud_output_path + + def initialize( self, - *, local_output_path: str, predictor_type: str, cloud_output_path: Optional[str] = None, resource_prefix: Optional[str] = None, + **kwargs, ) -> None: + """Initialize the backend.""" self.local_output_path = local_output_path self._cloud_output_path = cloud_output_path self.predictor_type = predictor_type @@ -57,16 +68,6 @@ def __init__( self.original_features = None self.endpoint_name: Optional[str] = None - @property - def cloud_output_path(self) -> str: - if not self._cloud_output_path: - raise ValueError( - "No `cloud_output_path` was provided and no bucket is configured in " - "~/.autogluon/cloud.yaml. Either pass `cloud_output_path=` explicitly, or run " - "`autogluon.cloud.bootstrap()` / `register(bucket=...)` once to persist a bucket." - ) - return self._cloud_output_path - @abstractmethod def attach_job(self, job_name: str) -> None: """ @@ -123,33 +124,21 @@ def prepare_args(self, path: str, **kwargs): with open(path, "w") as f: f.write(payload) - def _construct_ag_args(self, **kwargs): + def _construct_ag_args(**kwargs): raise NotImplementedError @abstractmethod - def fit( - self, - *, - predictor_init_args: Dict[str, Any], - predictor_fit_args: Dict[str, Any], - data_channels: Dict[str, Optional[Union[str, pd.DataFrame]]], - **kwargs, - ) -> None: + def fit(self, **kwargs) -> None: """Fit AG on the backend""" raise NotImplementedError @abstractmethod - def deploy( - self, - predictor_path: Optional[str] = None, - endpoint_name: Optional[str] = None, - **kwargs, - ) -> None: + def deploy(self, **kwargs) -> None: """Deploy and endpoint""" raise NotImplementedError @abstractmethod - def cleanup_deployment(self) -> None: + def cleanup_deployment(self, **kwargs) -> None: """Delete endpoint, and cleanup other artifacts""" raise NotImplementedError @@ -210,9 +199,3 @@ def get_fit_predict_results(self) -> pd.DataFrame: """ raise NotImplementedError(f"{self.__class__.__name__} does not support `fit_predict`.") - - def get_prediction_future( - self, *, result_transform: Optional[Callable[[pd.DataFrame], Any]] = None - ) -> JobPredictionFuture: - """Return a pending result bound to the most recently submitted prediction job.""" - raise NotImplementedError(f"{self.__class__.__name__} does not support prediction futures.") diff --git a/src/autogluon/cloud/backend/backend_factory.py b/src/autogluon/cloud/backend/backend_factory.py index 0b864a54..a4204f0d 100644 --- a/src/autogluon/cloud/backend/backend_factory.py +++ b/src/autogluon/cloud/backend/backend_factory.py @@ -1,6 +1,3 @@ -from typing import Optional, Union - -from ..config import SageMakerConfig from .backend import Backend from .multimodal_sagemaker_backend import MultiModalSagemakerBackend from .sagemaker_backend import SagemakerBackend @@ -9,7 +6,6 @@ class BackendFactory: - _CONFIGS = {SageMakerConfig.name: SageMakerConfig} _BACKENDS = { SagemakerBackend.name: SagemakerBackend, TabularSagemakerBackend.name: TabularSagemakerBackend, @@ -17,21 +13,6 @@ class BackendFactory: TimeSeriesSagemakerBackend.name: TimeSeriesSagemakerBackend, } - @staticmethod - def resolve_config(backend: Union[str, SageMakerConfig]) -> SageMakerConfig: - """Normalize a backend name or reusable configuration without creating resources.""" - if isinstance(backend, SageMakerConfig): - return backend - if not isinstance(backend, str): - raise TypeError("`backend` must be a backend name or SageMakerConfig.") - if backend in ("ray", "ray_aws"): - raise ValueError("The Ray backend was removed in AutoGluon-Cloud v0.7.0. Use backend='sagemaker' instead.") - if backend not in BackendFactory._CONFIGS: - raise ValueError( - f"Unsupported backend {backend!r}. Supported backends: {sorted(BackendFactory._CONFIGS)}." - ) - return BackendFactory._CONFIGS[backend]() - @staticmethod def get_backend_cls(backend: str) -> type[Backend]: if backend in BackendFactory._BACKENDS: @@ -39,20 +20,6 @@ def get_backend_cls(backend: str) -> type[Backend]: raise ValueError(f"{backend} not supported. Supported backends: {sorted(BackendFactory._BACKENDS)}") @staticmethod - def get_backend( - backend: str, - *, - local_output_path: str, - cloud_output_path: Optional[str], - predictor_type: str, - config: Optional[SageMakerConfig] = None, - resource_prefix: Optional[str] = None, - ) -> Backend: - """Create a backend with its own execution state from reusable settings.""" - return BackendFactory.get_backend_cls(backend)( - local_output_path=local_output_path, - cloud_output_path=cloud_output_path, - predictor_type=predictor_type, - config=config, - resource_prefix=resource_prefix, - ) + def get_backend(backend: str, **init_args) -> Backend: + """Return the corresponding backend""" + return BackendFactory.get_backend_cls(backend)(**init_args) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index d4f7c44b..37a73682 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -4,9 +4,7 @@ import os import tarfile import tempfile -from dataclasses import replace -from functools import partial -from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union +from typing import Any, Dict, List, Literal, Optional, Tuple, Union import pandas as pd from botocore.exceptions import ClientError @@ -14,13 +12,10 @@ from autogluon.common.loaders import load_pd from autogluon.common.utils.s3_utils import is_s3_url, s3_path_to_bucket_prefix -from ..config import SageMakerConfig from ..data import FormatConverterFactory -from ..endpoint.prediction_future import JobPredictionFuture from ..job import SageMakerBatchTransformationJob, SageMakerFitJob from ..scripts import ScriptManager from ..utils.ag_sagemaker import ( - create_serve_script_tarball, repack_model_with_serving_code, resolve_image_uri, script_mode_environment, @@ -43,7 +38,7 @@ invoke_endpoint, ) from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer -from ..utils.tag_utils import build_tags, to_request_tags +from ..utils.tag_utils import build_tags from ..utils.utils import ( convert_image_path_to_encoded_bytes_in_dataframe, is_image_file, @@ -56,7 +51,6 @@ logger = logging.getLogger(__name__) SAGEMAKER_MODEL_SERVER_WORKERS = "SAGEMAKER_MODEL_SERVER_WORKERS" -_VPC_CONFIG_FIELDS = {"subnets": "Subnets", "security_group_ids": "SecurityGroupIds"} _SERVERLESS_CONFIG_FIELDS = { "memory_size_in_mb": "MemorySizeInMB", "max_concurrency": "MaxConcurrency", @@ -93,50 +87,45 @@ def _to_request_fields(settings: Dict[str, Any], fields: Dict[str, str], arg_nam return {fields[key]: value for key, value in settings.items()} -def _vpc_request(vpc_config: Dict[str, List[str]]) -> Dict[str, List[str]]: - """``SageMakerConfig.vpc_config`` in the ``VpcConfig`` request format.""" - missing = sorted(set(_VPC_CONFIG_FIELDS) - set(vpc_config)) - if missing: - raise ValueError(f"`vpc_config` is missing required key(s) {missing}.") - return _to_request_fields(vpc_config, _VPC_CONFIG_FIELDS, "vpc_config") - - class SagemakerBackend(Backend): name = SAGEMAKER def __init__( self, local_output_path: str, - cloud_output_path: Optional[str], + cloud_output_path: str, predictor_type: str, - *, - config: Optional[SageMakerConfig] = None, - resource_prefix: Optional[str] = None, + role: Optional[str] = None, + **kwargs, ) -> None: - super().__init__( + self.initialize( local_output_path=local_output_path, cloud_output_path=cloud_output_path, predictor_type=predictor_type, - resource_prefix=resource_prefix, + role=role, + **kwargs, ) - if config is not None and not isinstance(config, SageMakerConfig): - raise TypeError("`config` must be a SageMakerConfig.") - config = copy.deepcopy(config or SageMakerConfig()) - if config.vpc_config is not None: - _vpc_request(config.vpc_config) # fail before creating a session or resolving the role - self.sagemaker_session = setup_sagemaker_session(region=config.region) + + def initialize(self, role: Optional[str] = None, **kwargs) -> None: + """Initialize the backend. + + Parameters + ---------- + role + SageMaker execution role ARN. See + :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( - config.role_arn, backend_name=SAGEMAKER, session=self.sagemaker_session - ) + 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 `backend=SageMakerConfig(role_arn=)` " + "Failed to resolve SageMaker execution role. Pass `role=` to the predictor/model " "or run `autogluon.cloud.bootstrap()` / `register()` to persist one." ) raise e self._region = self.sagemaker_session.boto_region_name - self.config = replace(config, region=self._region, role_arn=self.role_arn) self._fit_job: SageMakerFitJob = SageMakerFitJob(session=self.sagemaker_session) self._batch_transform_jobs = MostRecentInsertedOrderedDict() @@ -144,9 +133,9 @@ def _realtime_serializer(self): """Serializer used for realtime endpoint requests""" return AutoGluonSerializer() - def _resolve_tags(self, extra_tags: Optional[Dict[str, str]] = None) -> List[Dict[str, str]]: - """Tags for a created SageMaker resource, in SageMaker API request format: default + extra + user tags.""" - return to_request_tags(build_tags(self.predictor_type, extra_tags=extra_tags, user_tags=self.config.tags)) + 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: """ @@ -212,15 +201,12 @@ def fit( instance_type: str = "ml.m5.2xlarge", instance_count: Union[int, str] = 1, volume_size: int = 256, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, timeout: int = 24 * 60 * 60, wait: bool = True, - environment: Optional[Dict[str, str]] = None, - use_spot_instances: bool = False, - max_wait: Optional[int] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, extra_ag_args: Optional[Dict[str, Any]] = None, - extra_tags: Optional[Dict[str, str]] = None, + extra_tags: Optional[List[Dict[str, str]]] = None, ) -> None: """ Fit the predictor with SageMaker. @@ -246,7 +232,7 @@ def fit( Training container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, AutoGluon Cloud creates one with a predictor- or model-specific prefix. @@ -257,21 +243,12 @@ def fit( volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training (default: 256). Must be large enough to store training data if File Mode is used (which is the default). - image_uri: Optional[str], default = None - Custom training container image. If None, the official AutoGluon DLC for ``framework_version`` is used. timeout: int, default = 24*60*60 Timeout in seconds for training. This timeout doesn't include time for pre-processing or launching up the training job. wait: bool, default = True 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the training container. - use_spot_instances: bool, default = False - Whether to use managed spot training. - max_wait: Optional[int], default = None - Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires - ``use_spot_instances=True``. 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. @@ -284,8 +261,6 @@ def fit( raise ValueError("`data_channels['train_data']` is required.") _reject_local_mode(instance_type) overrides = check_override_keys(backend_overrides, FIT_OVERRIDE_KEYS) - if max_wait is not None and not use_spot_instances: - raise ValueError("`max_wait` requires `use_spot_instances=True`.") 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 = { @@ -293,9 +268,9 @@ def fit( for k, v in data_channels.items() if v is not None } - if image_uri: + if custom_image_uri: framework_version, py_version = None, None - logger.log(20, f"Training with image_uri=={image_uri}") + logger.log(20, f"Training with custom_image_uri=={custom_image_uri}") else: framework_version, py_version = parse_framework_version( framework_version, "training", minimum_version="0.6.0" @@ -341,7 +316,6 @@ def fit( ag_args_path = os.path.join(self.local_output_path, "utils", "ag_args.json") self.prepare_args(path=ag_args_path, **ag_args) inputs = self._upload_fit_artifact( - job_name=job_name, data_channels=data_channels, label=label, ag_args=ag_args_path, @@ -352,20 +326,16 @@ def fit( ) code_uri = upload_training_code( entry_point=entry_point, - source_dir=None, sagemaker_session=self.sagemaker_session, s3_uri_prefix=f"{self.cloud_output_path}/code/{job_name}/source", ) - stopping_condition: Dict[str, Any] = {"MaxRuntimeInSeconds": timeout} - if use_spot_instances: - stopping_condition["MaxWaitTimeInSeconds"] = max_wait or timeout request: Dict[str, Any] = { "TrainingJobName": job_name, "RoleArn": self.role_arn, "AlgorithmSpecification": { "TrainingImage": resolve_image_uri( - image_uri, framework_version, py_version, self._region, "training", instance_type + custom_image_uri, framework_version, py_version, self._region, "training", instance_type ), "TrainingInputMode": "File", }, @@ -379,20 +349,10 @@ def fit( "InstanceCount": instance_count, "VolumeSizeInGB": volume_size, }, - "StoppingCondition": stopping_condition, + "StoppingCondition": {"MaxRuntimeInSeconds": timeout}, "ProfilerConfig": {"DisableProfiler": True}, "Tags": self._resolve_tags(extra_tags), } - if environment: - request["Environment"] = dict(environment) - if use_spot_instances: - request["EnableManagedSpotTraining"] = True - if self.config.vpc_config is not None: - request["VpcConfig"] = _vpc_request(self.config.vpc_config) - if self.config.output_kms_key is not None: - request["OutputDataConfig"]["KmsKeyId"] = self.config.output_kms_key - if self.config.volume_kms_key is not None: - request["ResourceConfig"]["VolumeKmsKeyId"] = self.config.volume_kms_key request = deep_merge(request, overrides.get("create_training_job", {})) self._fit_job = SageMakerFitJob(session=self.sagemaker_session) @@ -425,8 +385,6 @@ def _create_model( "ExecutionRoleArn": self.role_arn, "Tags": tags, } - if self.config.vpc_config is not None: - request["VpcConfig"] = _vpc_request(self.config.vpc_config) request = deep_merge(request, overrides.get("create_model", {})) logger.log(20, "Creating inference model...") self.sagemaker_session.sagemaker_client.create_model(**request) @@ -449,20 +407,8 @@ def _prepare_model_data( entry_point=entry_point, repacked_model_uri=repacked_model_uri, sagemaker_session=self.sagemaker_session, - kms_key=self.config.output_kms_key, ) - @staticmethod - def _model_server_environment(environment: Optional[Dict[str, str]]) -> Dict[str, str]: - environment = dict(environment or {}) - if SAGEMAKER_MODEL_SERVER_WORKERS in environment and int(environment[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: - environment[SAGEMAKER_MODEL_SERVER_WORKERS] = "1" - return environment - def deploy( self, predictor_path: Optional[str] = None, @@ -470,17 +416,16 @@ def deploy( framework_version: str = "latest", instance_type: Optional[str] = "ml.m5.2xlarge", initial_instance_count: int = 1, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, volume_size: Optional[int] = None, wait: bool = True, - environment: Optional[Dict[str, str]] = 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, repack: bool = True, - extra_tags: Optional[Dict[str, str]] = None, + extra_tags: Optional[List[Dict[str, str]]] = None, ) -> None: """ Deploy a predictor as a SageMaker endpoint, which can be used to do real-time inference later. @@ -499,12 +444,12 @@ def deploy( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. instance_type: str, default = 'ml.m5.2xlarge' Instance to be deployed for the endpoint initial_instance_count: int, default = 1, Initial number of instances to be deployed for the endpoint - image_uri: Optional[str], default = None, + custom_image_uri: Optional[str], default = None, Custom image to use to deploy endpoint with. If not specified, with use official DLC image: https://aws.github.io/deep-learning-containers/reference/available_images/#autogluon @@ -514,8 +459,6 @@ 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the inference container. 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"``, @@ -554,9 +497,9 @@ def deploy( endpoint_name = unique_name_from_base(self.resource_prefix) # Resolve container image - if image_uri: + if custom_image_uri: framework_version, py_version = None, None - logger.log(20, f"Deploying with image_uri=={image_uri}") + logger.log(20, f"Deploying with custom_image_uri=={custom_image_uri}") else: framework_version, py_version = parse_framework_version( framework_version, "inference", minimum_version="0.6.0" @@ -602,7 +545,7 @@ def deploy( repacked_model_uri=f"{self.cloud_output_path}/endpoints/{endpoint_name}/model/model.tar.gz", ) - container_environment = self._model_server_environment(environment) + container_environment = {SAGEMAKER_MODEL_SERVER_WORKERS: "1"} if fm_serve_config is not None: container_environment["AG_FM_SERVE_CONFIG"] = json.dumps(fm_serve_config) if inference_mode == "serverless": @@ -616,7 +559,7 @@ def deploy( model_name=unique_name_from_base(endpoint_name), model_data=model_data, image_uri=resolve_image_uri( - image_uri, framework_version, py_version, self._region, "inference", instance_type + custom_image_uri, framework_version, py_version, self._region, "inference", instance_type ), entry_point=entry_point, environment=container_environment, @@ -630,7 +573,9 @@ def deploy( variant["InitialInstanceCount"] = initial_instance_count if volume_size: variant["VolumeSizeInGB"] = volume_size - inference_ami_version = infer_sagemaker_ami_version(image_uri, instance_type, image_scope="inference") + 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 elif inference_mode == "serverless": @@ -644,8 +589,6 @@ def deploy( "ProductionVariants": [variant], "Tags": tags, } - if self.config.volume_kms_key is not None and inference_mode == "realtime": - endpoint_config_request["KmsKeyId"] = self.config.volume_kms_key endpoint_config_request = deep_merge(endpoint_config_request, overrides.get("create_endpoint_config", {})) endpoint_request = deep_merge( { @@ -665,11 +608,16 @@ def deploy( 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/ and upload it.""" + """Create a minimal model.tar.gz containing the serve script + serving_utils/ under code/.""" + tarball_dir = tempfile.mkdtemp(prefix="ag_serve_") - tarball_path = create_serve_script_tarball(serve_script_path, tarball_dir) + tarball_path = os.path.join(tarball_dir, "model.tar.gz") + with tarfile.open(tarball_path, "w:gz") as tar: + tar.add(serve_script_path, arcname=f"code/{os.path.basename(serve_script_path)}") + tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils") s3_key = f"endpoints/{endpoint_name}/model/model.tar.gz" - return self._upload_predictor(tarball_path, s3_key) + s3_path = self._upload_predictor(tarball_path, s3_key) + return s3_path def cleanup_deployment(self) -> None: """ @@ -834,12 +782,11 @@ def predict( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, download: bool = True, persist: bool = True, save_path: Optional[str] = None, - environment: Optional[Dict[str, str]] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ @@ -864,16 +811,14 @@ def predict( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None - Name of the launched transform job. + Name of the launched training job. If None, AutoGluon Cloud creates one with a predictor- or model-specific prefix. instance_count: int, default = 1, Number of instances used to do batch transform. instance_type: str, default = 'ml.m5.2xlarge' Instance to be used for batch transform. - image_uri: Optional[str], default = None - Custom inference container image. If None, the official AutoGluon DLC is used. 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. @@ -888,8 +833,6 @@ def predict( 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the inference container. 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"``. @@ -908,12 +851,11 @@ def predict( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, download=download, persist=persist, save_path=save_path, - environment=environment, backend_overrides=backend_overrides, original_features=self.original_features, ) @@ -925,17 +867,69 @@ def predict_proba( test_data: Union[str, pd.DataFrame], test_data_image_column: Optional[str] = None, include_predict: bool = True, - **kwargs, + predictor_path: Optional[str] = None, + framework_version: str = "latest", + job_name: Optional[str] = None, + instance_type: str = "ml.m5.2xlarge", + instance_count: int = 1, + custom_image_uri: Optional[str] = None, + wait: bool = True, + download: bool = True, + persist: bool = True, + save_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 probabilities using SageMaker batch transform. - Accepts the same arguments as :meth:`predict`. + 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 creates a SageMaker model with the trained predictor and runs a transform job with it. Parameters ---------- + test_data: Union(str, pandas.DataFrame) + The test data to be inferenced. Can be a pandas.DataFrame, or a local path to a csv. + test_data_image_column: str, default = None + If test_data involves image modality, you must specify the column name corresponding to image paths. + The path MUST be an abspath include_predict: bool, default = True Whether to include predict result along with predict_proba results. This flag can save you time from making two calls to get both the prediction and the probability as batch inference involves noticeable overhead. + predictor_path: str + Path to the predictor tarball you want to use to predict. + Path can be both a local path or a S3 location. + If None, will use the most recent trained predictor trained with `fit()`. + framework_version: str, default = `latest` + Inference container version of autogluon. + If `latest`, will use the latest available container version. + If provided a specific version, will use this version. + If `custom_image_uri` is set, this argument will be ignored. + job_name: str, default = None + Name of the launched training job. + If None, AutoGluon Cloud creates one with a predictor- or model-specific prefix. + instance_count: int, default = 1, + Number of instances used to do batch transform. + instance_type: str, default = 'ml.m5.2xlarge' + Instance to be used for batch transform. + 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. + 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 ------- @@ -948,8 +942,18 @@ def predict_proba( pred, pred_proba = self._predict( test_data=test_data, test_data_image_column=test_data_image_column, + predictor_path=predictor_path, + framework_version=framework_version, + job_name=job_name, + instance_type=instance_type, + instance_count=instance_count, + custom_image_uri=custom_image_uri, + wait=wait, + download=download, + persist=persist, + save_path=save_path, + backend_overrides=backend_overrides, original_features=self.original_features, - **kwargs, ) if include_predict: @@ -1001,21 +1005,7 @@ def download_predict_results(self, job_name: Optional[str] = None, save_path: Op def get_fit_predict_results(self) -> pd.DataFrame: """Read predictions produced by a completed ``fit_predict`` job from S3.""" - return self._load_fit_predict_results(self._fit_job) - - def get_prediction_future( - self, *, result_transform: Optional[Callable[[pd.DataFrame], Any]] = None - ) -> JobPredictionFuture: - """Bind waiting and result loading to the current job, independent of later submissions.""" - job = self._fit_job - if job.job_name is None: - raise ValueError("No prediction job found. Submit a prediction first.") - load_results = partial(self._load_fit_predict_results, job) - result_loader = load_results if result_transform is None else lambda: result_transform(load_results()) - return JobPredictionFuture(job=job, result_loader=result_loader) - - def _load_fit_predict_results(self, job: SageMakerFitJob) -> pd.DataFrame: - ag_args = self._download_ag_args_from_job(job) + ag_args = self._download_ag_args_from_job() predictions_path = ag_args.get("predictions_path") assert predictions_path is not None, "No fit_predict job found. Call `fit_predict()` first." bucket, key = s3_path_to_bucket_prefix(predictions_path) @@ -1024,16 +1014,16 @@ def _load_fit_predict_results(self, job: SageMakerFitJob) -> pd.DataFrame: self.sagemaker_session.s3_client.download_file(bucket, key, local_path) return load_pd.load(local_path) - def _download_ag_args_from_job(self, job: Optional[SageMakerFitJob] = None) -> Dict[str, Any]: + def _download_ag_args_from_job(self) -> Dict[str, Any]: """Fetch and parse the ``ag_args.json`` that was uploaded as the ``ag_args`` channel. Each training job carries the exact config it was launched with as an input channel, making this the authoritative source — independent of local-disk lifetime. """ - job = self._fit_job if job is None else job - job_name = job.job_name + job_name = self._fit_job.job_name assert job_name is not None, "No fit job found. Call `fit()` / `fit_predict()` first." - channels = job.get_input_channels() + 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, ( f"Training job {job_name!r} has no `ag_args` input channel — cannot recover predictions_path." @@ -1080,7 +1070,6 @@ def _find_common_path_and_replace_image_column(self, data, image_column): def _upload_fit_artifact( self, - job_name: str, data_channels, label, ag_args, @@ -1088,7 +1077,7 @@ def _upload_fit_artifact( image_column=None, ): cloud_bucket, cloud_key_prefix = s3_path_to_bucket_prefix(self.cloud_output_path) - util_key_prefix = f"{cloud_key_prefix}/{job_name}/utils" + util_key_prefix = cloud_key_prefix + "/utils" # Image-column mode: rewrite image paths to be container-relative; common image directories # are zipped and uploaded as separate train_images / tune_images channels below. @@ -1228,12 +1217,11 @@ def _predict( job_name=None, instance_type="ml.m5.2xlarge", instance_count=1, - image_uri=None, + custom_image_uri=None, wait=True, download=True, persist=True, save_path=None, - environment=None, backend_overrides=None, split_pred_proba=True, original_features=None, @@ -1249,9 +1237,9 @@ def _predict( predictor_path = self._fit_job.get_output_path() assert predictor_path, "No cloud trained model found." - if image_uri: + if custom_image_uri: framework_version, py_version = None, None - logger.log(20, f"Predicting with image_uri=={image_uri}") + logger.log(20, f"Predicting with custom_image_uri=={custom_image_uri}") else: framework_version, py_version = parse_framework_version( framework_version, "inference", minimum_version="0.6.0" @@ -1337,10 +1325,10 @@ def _predict( model_name=job_name, model_data=model_data, image_uri=resolve_image_uri( - image_uri, framework_version, py_version, self._region, "inference", instance_type + custom_image_uri, framework_version, py_version, self._region, "inference", instance_type ), entry_point=entry_point, - environment=dict(environment or {}), + environment={}, tags=tags, overrides=overrides, ) @@ -1355,13 +1343,9 @@ def _predict( 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(image_uri, instance_type, image_scope="transform") + 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 - if self.config.output_kms_key is not None: - transform_output["KmsKeyId"] = self.config.output_kms_key - if self.config.volume_kms_key is not None: - transform_resources["VolumeKmsKeyId"] = self.config.volume_kms_key request = { "TransformJobName": job_name, "ModelName": model_name, @@ -1411,7 +1395,7 @@ def __getstate__(self) -> Dict[str, Any]: def __setstate__(self, state): """Custom implementation of the unpickle process""" self.__dict__.update(state) - self.sagemaker_session = setup_sagemaker_session(region=self.config.region) + self.sagemaker_session = setup_sagemaker_session() self._region = self.sagemaker_session.boto_region_name self._fit_job.session = self.sagemaker_session for job in self._batch_transform_jobs.values(): diff --git a/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py b/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py index 4b7b94dd..5dbd4359 100644 --- a/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py +++ b/src/autogluon/cloud/backend/timeseries_sagemaker_backend.py @@ -1,6 +1,6 @@ import logging import os -from typing import Any, Dict, Optional, Union +from typing import Any, Dict, List, Optional, Union import pandas as pd @@ -24,15 +24,22 @@ def fit( data_channels: Dict[str, Optional[Union[str, pd.DataFrame]]], id_column: str, timestamp_column: str, + framework_version: str = "latest", + job_name: Optional[str] = None, + instance_type: str = "ml.m5.2xlarge", + instance_count: int = 1, volume_size: int = 100, + custom_image_uri: Optional[str] = None, + wait: bool = True, + backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, extra_ag_args: Optional[Dict[str, Any]] = None, - **kwargs, + extra_tags: Optional[List[Dict[str, str]]] = None, ) -> None: """Fit a TimeSeriesPredictor in SageMaker. ``id_column`` / ``timestamp_column`` are forwarded to the training script via ``ag_args.json``. ``known_covariates`` (if present in ``data_channels``) is only honored when - ``extra_ag_args["predict_after_fit"]`` is True. Other arguments are forwarded to ``SagemakerBackend.fit()``. + ``extra_ag_args["predict_after_fit"]`` is True. """ extra_ag_args = {**(extra_ag_args or {}), "id_column": id_column, "timestamp_column": timestamp_column} if data_channels.get("known_covariates") is not None and not extra_ag_args.get("predict_after_fit", False): @@ -48,9 +55,16 @@ def fit( predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, data_channels=data_channels, + framework_version=framework_version, + job_name=job_name, + instance_type=instance_type, + instance_count=instance_count, volume_size=volume_size, + custom_image_uri=custom_image_uri, + wait=wait, + backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, - **kwargs, + extra_tags=extra_tags, ) def predict_real_time( diff --git a/src/autogluon/cloud/config.py b/src/autogluon/cloud/config.py index d8fb48f9..31ca9cbf 100644 --- a/src/autogluon/cloud/config.py +++ b/src/autogluon/cloud/config.py @@ -1,4 +1,4 @@ -"""Backend settings and persistent resource identifiers for AutoGluon-Cloud. +"""Persistent config for AutoGluon-Cloud. Stores resource identifiers (region, stack name, bucket, IAM role ARN) at ``~/.autogluon/cloud.yaml`` so users don't need to re-specify them every @@ -20,52 +20,13 @@ import stat from dataclasses import asdict, dataclass, field from pathlib import Path -from typing import ClassVar, Dict, List, Optional +from typing import Dict, Optional import yaml CONFIG_DIR_ENV = "AG_CONFIG_DIR" -@dataclass(kw_only=True) -class SageMakerConfig: - """Reusable SageMaker settings for predictors and foundation models. - - Pass this as ``backend=`` to a cloud predictor or foundation model. Each - object creates its own backend and jobs; sharing this config does not share - execution state. Resource sizes and other operation settings remain named - arguments to ``fit()``, ``predict()`` and ``deploy()``. - - Parameters - ---------- - region - AWS region. If omitted, use the region in ``~/.autogluon/cloud.yaml``, - then the boto3 default region. - role_arn - SageMaker execution role ARN. If omitted, use the saved role, then the - role of the current AWS identity. - vpc_config - Networking for training jobs and models, as - ``{"subnets": [...], "security_group_ids": [...]}``. - output_kms_key - KMS key for training artifacts, batch transform outputs, and repacked - or cached model artifacts in S3. - volume_kms_key - KMS key for training, batch transform and realtime endpoint storage - volumes. Leave unset for instance types with local NVMe storage. - tags - Tags added to every SageMaker resource created by this backend. - """ - - name: ClassVar[str] = "sagemaker" - region: Optional[str] = None - role_arn: Optional[str] = None - vpc_config: Optional[Dict[str, List[str]]] = None - output_kms_key: Optional[str] = None - volume_kms_key: Optional[str] = None - tags: Dict[str, str] = field(default_factory=dict) - - def get_config_dir() -> Path: override = os.environ.get(CONFIG_DIR_ENV) if override: diff --git a/src/autogluon/cloud/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index ad4a88fb..198afd68 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -192,13 +192,6 @@ def _get_hyperparameters(self): return self._describe().get("HyperParameters") return None - def get_input_channels(self) -> Dict[str, str]: - """Map each input channel name of the training job to its S3 URI.""" - return { - channel["ChannelName"]: channel["DataSource"]["S3DataSource"]["S3Uri"] - for channel in self._describe()["InputDataConfig"] - } - def run( self, training_job_request: Dict[str, Any], diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index ae925ede..f2ac48bb 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -7,7 +7,6 @@ import tarfile import tempfile from abc import abstractmethod -from functools import partial from pathlib import Path from typing import Any, Dict, List, Literal, Optional, Tuple, Union @@ -19,7 +18,6 @@ from ..backend.backend_factory import BackendFactory from ..backend.constant import SAGEMAKER, TABULAR_SAGEMAKER, TIMESERIES_SAGEMAKER -from ..config import SageMakerConfig from ..endpoint.prediction_future import JobPredictionFuture from ..endpoint.tabular_endpoint import TabularEndpoint from ..endpoint.timeseries_endpoint import TimeSeriesEndpoint @@ -83,9 +81,10 @@ def __init__( model_id: str, *, cloud_output_path: Optional[str] = None, + role: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, model_artifact_uri: Optional[str] = None, - backend: Union[str, SageMakerConfig] = SAGEMAKER, + backend: Literal["sagemaker"] = "sagemaker", ): """ Parameters @@ -103,35 +102,38 @@ def __init__( * ``None`` (default) — use the bucket saved in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` / :func:`autogluon.cloud.register`) and append a timestamped subfolder. Raises if no bucket is configured. + 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 the role of the current AWS identity. hyperparameters Default hyperparameters applied to inference and (when supported) training. model_artifact_uri S3 URI of a pre-bundled ``model.tar.gz`` produced by :meth:`cache_model_artifact`. When set, deploys skip the runtime HuggingFace download and load weights from the bundled artifact. backend - Backend name or reusable :class:`~autogluon.cloud.SageMakerConfig` with region, execution role, - networking, encryption and tags. ``"sagemaker"`` uses default settings. + Cloud backend to use. """ - backend_config = BackendFactory.resolve_config(backend) - backend_name = self._backend_map.get(backend_config.name) - if backend_name is None: - raise ValueError( - f"Backend {backend_config.name!r} is not supported for {self.__class__.__name__}. " - f"Available: {list(self._backend_map.keys())}" - ) self.model_id = model_id self.model_artifact_uri = model_artifact_uri - self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend_config.name) + self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend) self._config = get_model_config(model_id) self._hyperparameter_overrides = hyperparameters or {} self._tmpdir = tempfile.TemporaryDirectory(prefix="ag_fm_") + + backend_name = self._backend_map.get(backend) + if backend_name is None: + raise ValueError( + f"Backend '{backend}' is not supported for {self.__class__.__name__}. " + f"Available: {list(self._backend_map.keys())}" + ) self._backend = BackendFactory.get_backend( backend=backend_name, local_output_path=self._tmpdir.name, cloud_output_path=self.cloud_output_path, predictor_type=self._predictor_type, resource_prefix=f"ag-cloud-{self.model_id}", - config=backend_config, + role=role, ) def _get_hyperparameters( @@ -188,11 +190,11 @@ def _deploy_backend( endpoint_name: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, framework_version: str = "latest", - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - **kwargs, + **backend_kwargs, ) -> None: """Shared deploy logic. Subclasses call this then wrap the endpoint.""" if inference_mode == "serverless" and instance_type is not None: @@ -224,15 +226,15 @@ def _deploy_backend( endpoint_name=endpoint_name, framework_version=framework_version, instance_type=instance_type, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, entry_point=self._serve_script_path, fm_serve_config=fm_serve_config, inference_mode=inference_mode, inference_config=inference_config, repack=False, - extra_tags={"autogluon-cloud-model-id": self.model_id}, - **kwargs, + extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], + **backend_kwargs, ) assert self._backend.endpoint_name is not None @@ -339,14 +341,11 @@ def cache_model_artifact(self, cache_path: str, *, overwrite: bool = False) -> S tar.add(serve_script, arcname=f"code/{serve_script.name}") tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils") logger.info(f"Uploading to {cache_key}") - extra_args = {"Metadata": {_AG_CLOUD_VERSION_METADATA_KEY: __version__}} - if self._backend.config.output_kms_key is not None: - extra_args.update(ServerSideEncryption="aws:kms", SSEKMSKeyId=self._backend.config.output_kms_key) s3.upload_file( str(tarball), bucket, key, - ExtraArgs=extra_args, + ExtraArgs={"Metadata": {_AG_CLOUD_VERSION_METADATA_KEY: __version__}}, ) return self.__class__( @@ -354,11 +353,11 @@ def cache_model_artifact(self, cache_path: str, *, overwrite: bool = False) -> S hyperparameters=self._hyperparameter_overrides or None, model_artifact_uri=cache_key, cloud_output_path=self.cloud_output_path, - backend=self._backend.config, + role=self._backend.role_arn, ) def to_dict(self) -> Dict[str, Any]: - """Serialize the model identity. Runtime context (``backend``, ``cloud_output_path``) is excluded so configs can + """Serialize the model identity. Runtime context (``role``, ``cloud_output_path``) is excluded so configs can be shared across users.""" out: Dict[str, Any] = {"model_id": self.model_id} if self._hyperparameter_overrides: @@ -373,7 +372,7 @@ def to_json(self) -> str: @classmethod def from_dict(cls, config: Dict[str, Any], **runtime_context: Any) -> Self: - """Restore from :meth:`to_dict` output. Pass ``backend`` / ``cloud_output_path`` as ``runtime_context``.""" + """Restore from :meth:`to_dict` output. Pass ``role`` / ``cloud_output_path`` as ``runtime_context``.""" return cls(**config, **runtime_context) @classmethod @@ -416,13 +415,11 @@ def deploy( endpoint_name: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, framework_version: str = "latest", - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - environment: Optional[Dict[str, str]] = None, - backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, - **kwargs, + **backend_kwargs, ) -> TimeSeriesEndpoint: """ Deploy model to an inference endpoint. @@ -438,7 +435,7 @@ def deploy( Model hyperparameters for inference. Overrides values passed to the constructor. framework_version Container framework version. If 'latest', uses the most recent available. - image_uri + custom_image_uri Custom Docker image URI for the inference container. wait Whether to block until the endpoint is ready. @@ -447,27 +444,20 @@ def deploy( (no instance management, scales to zero). inference_config Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). - environment - Environment variables set in the inference container. - backend_overrides - Raw SageMaker request fields deep-merged over the requests built by AutoGluon-Cloud. Valid keys: - ``"create_model"``, ``"production_variant"``, ``"create_endpoint_config"``, ``"create_endpoint"``. See - :meth:`autogluon.cloud.TabularCloudPredictor.deploy`. - **kwargs - Additional deployment arguments (``initial_instance_count``, ``volume_size``). + **backend_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, endpoint_name=endpoint_name, hyperparameters=hyperparameters, framework_version=framework_version, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, inference_mode=inference_mode, inference_config=inference_config, - environment=environment, - backend_overrides=backend_overrides, - **kwargs, + **backend_kwargs, ) return TimeSeriesEndpoint( endpoint_name=self._backend.endpoint_name, @@ -512,10 +502,9 @@ def predict( hyperparameters: Optional[Dict[str, Any]] = None, instance_type: Optional[str] = None, framework_version: str = "latest", - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, - backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, - **kwargs, + **backend_kwargs, ) -> Union[pd.DataFrame, JobPredictionFuture]: """ Run batch prediction for time series. @@ -555,18 +544,15 @@ def predict( Instance type for the prediction job. If None, uses registry default. framework_version Container framework version. - image_uri + custom_image_uri Custom Docker image URI for the container. wait If True, block and return a DataFrame. If False, return a :class:`JobPredictionFuture` immediately — call ``.result()`` on it later to retrieve the DataFrame, or ``.status()`` to check progress. - backend_overrides - Raw provider request fields. This prediction path uses a SageMaker training job; - the valid key is ``"create_training_job"``. See :meth:`autogluon.cloud.TabularCloudPredictor.fit`. - **kwargs - Additional job arguments accepted by :meth:`autogluon.cloud.TimeSeriesCloudPredictor.fit` (e.g. - ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``). + **backend_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 ------- @@ -601,16 +587,18 @@ def predict( timestamp_column=timestamp_column, framework_version=framework_version, instance_type=instance_type, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, - backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, - extra_tags={"autogluon-cloud-model-id": self.model_id}, - **kwargs, + extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], + **backend_kwargs, ) if not wait: - return self._backend.get_prediction_future() + return JobPredictionFuture( + job=self._backend._fit_job, + result_loader=self._backend.get_fit_predict_results, + ) return self._backend.get_fit_predict_results() @@ -641,13 +629,11 @@ def deploy( endpoint_name: Optional[str] = None, hyperparameters: Optional[Dict[str, Any]] = None, framework_version: str = "latest", - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, inference_mode: Literal["realtime"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - environment: Optional[Dict[str, str]] = None, - backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, - **kwargs, + **backend_kwargs, ) -> TabularEndpoint: """Deploy the tabular foundation model to an inference endpoint. @@ -672,12 +658,10 @@ def deploy( endpoint_name=endpoint_name, hyperparameters=hyperparameters, framework_version=framework_version, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, inference_mode="realtime", - environment=environment, - backend_overrides=backend_overrides, - **kwargs, + **backend_kwargs, ) return TabularEndpoint( endpoint_name=self._backend.endpoint_name, @@ -697,15 +681,9 @@ def _build_predictor_fit_args(self, hyperparameters: Optional[Dict[str, Any]] = def _load_results( self, *, include_predict: bool, predict_only: bool = False - ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]: - raw = self._backend.get_fit_predict_results() - return self._format_results(raw, include_predict=include_predict, predict_only=predict_only) - - @staticmethod - def _format_results( - raw: pd.DataFrame, *, include_predict: bool, predict_only: bool = False ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series]]: # The training container writes [pred, _proba...]; regression has only the pred column. + raw = self._backend.get_fit_predict_results() pred, pred_proba = split_pred_and_pred_proba(raw) if pred_proba is None: # regression: proba mirrors pred, matching TabularPredictor.predict_proba pred_proba = pred @@ -727,10 +705,9 @@ def predict( hyperparameters: Optional[Dict[str, Any]] = None, instance_type: Optional[str] = None, framework_version: str = "latest", - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, - backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, - **kwargs, + **backend_kwargs, ) -> Union[pd.Series, JobPredictionFuture]: """ Run batch prediction for tabular tasks. @@ -756,17 +733,14 @@ def predict( Instance type for the prediction job. If None, uses registry default. framework_version Container framework version. - image_uri + custom_image_uri Custom Docker image URI for the container. wait If True, block and return the predictions. If False, return a :class:`JobPredictionFuture` immediately — call ``.result()`` on it later to retrieve the predictions. - backend_overrides - Raw provider request fields. This prediction path uses a SageMaker training job; - the valid key is ``"create_training_job"``. See :meth:`autogluon.cloud.TabularCloudPredictor.fit`. - **kwargs - Additional job arguments accepted by :meth:`autogluon.cloud.TabularCloudPredictor.fit` (e.g. - ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``). + **backend_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 ------- @@ -782,14 +756,14 @@ def predict( hyperparameters=hyperparameters, instance_type=instance_type, framework_version=framework_version, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, - backend_overrides=backend_overrides, - **kwargs, + **backend_kwargs, ) if not wait: - return self._backend.get_prediction_future( - result_transform=partial(self._format_results, include_predict=True, predict_only=True), + return JobPredictionFuture( + job=self._backend._fit_job, + result_loader=lambda: self._load_results(include_predict=True, predict_only=True), ) pred, _ = result return pred @@ -806,10 +780,9 @@ def predict_proba( hyperparameters: Optional[Dict[str, Any]] = None, instance_type: Optional[str] = None, framework_version: str = "latest", - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, - backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, - **kwargs, + **backend_kwargs, ) -> Union[Tuple[pd.Series, Union[pd.DataFrame, pd.Series]], Union[pd.DataFrame, pd.Series], JobPredictionFuture]: """ Run batch prediction returning class probabilities. @@ -837,15 +810,13 @@ def predict_proba( Instance type for the prediction job. If None, uses registry default. framework_version Container framework version. - image_uri + custom_image_uri Custom Docker image URI for the container. wait If True, block and return the result. If False, return a :class:`JobPredictionFuture` immediately. - backend_overrides - Raw provider request fields under ``"create_training_job"``. See :meth:`predict`. - **kwargs - Additional job arguments accepted by :meth:`autogluon.cloud.TabularCloudPredictor.fit` (e.g. - ``job_name``, ``volume_size``, ``environment``, ``use_spot_instances``). + **backend_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 ------- @@ -864,7 +835,7 @@ def predict_proba( extra_ag_args: Dict[str, Any] = {"predict_after_fit": True, "save_predictor": False} if predictions_path is not None: extra_ag_args["predictions_path"] = predictions_path - kwargs["leaderboard"] = False + backend_kwargs["leaderboard"] = False self._backend.fit( predictor_init_args=self._build_predictor_init_args(label=label), @@ -872,16 +843,16 @@ def predict_proba( data_channels={"train_data": train_data, "tuning_data": tuning_data, "test_data": test_data}, framework_version=framework_version, instance_type=instance_type, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, - backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, - extra_tags={"autogluon-cloud-model-id": self.model_id}, - **kwargs, + extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], + **backend_kwargs, ) if not wait: - return self._backend.get_prediction_future( - result_transform=partial(self._format_results, include_predict=include_predict), + return JobPredictionFuture( + job=self._backend._fit_job, + result_loader=lambda: self._load_results(include_predict=include_predict), ) return self._load_results(include_predict=include_predict) diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index 06827d17..f7390816 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -21,7 +21,6 @@ from ..backend.backend import Backend from ..backend.backend_factory import BackendFactory from ..backend.constant import SAGEMAKER -from ..config import SageMakerConfig from ..utils.aws_utils import resolve_cloud_output_path from ..utils.sagemaker_api import reject_legacy_kwargs from ..utils.utils import safe_unpack_archive @@ -37,7 +36,8 @@ def __init__( self, local_output_path: Optional[str] = None, cloud_output_path: Optional[str] = None, - backend: Union[str, SageMakerConfig] = SAGEMAKER, + backend: str = SAGEMAKER, + role: Optional[str] = None, verbosity: int = 2, ) -> None: """ @@ -60,10 +60,13 @@ def __init__( * ``None`` (default) — use the bucket saved in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` / :func:`autogluon.cloud.register`) and append a timestamped subfolder. Raises if no bucket is configured. - backend: Union[str, SageMakerConfig], default = "sagemaker" - Backend name or reusable :class:`~autogluon.cloud.SageMakerConfig` with region, execution role, - networking, encryption and tags. ``"sagemaker"`` uses default settings. - Only single instance training is supported. + backend: str, default = "sagemaker" + The backend to use. Currently only "sagemaker" is supported. + SageMaker backend supports training, deploying and batch inference on Amazon SageMaker. Only single instance training is supported. + 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 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). @@ -73,17 +76,18 @@ def __init__( self.verbosity = verbosity cloud_logger = logging.getLogger("autogluon.cloud") set_logger_verbosity(self.verbosity, logger=cloud_logger) - config = BackendFactory.resolve_config(backend) - if config.name not in self.backend_map: - raise ValueError(f"Unsupported backend {config.name!r}. Supported backends: {sorted(self.backend_map)}.") self.local_output_path = self._setup_local_output_path(local_output_path) - self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=config.name) + if backend in ("ray", "ray_aws"): + raise ValueError("The Ray backend was removed in AutoGluon-Cloud v0.7.0. Use backend='sagemaker' instead.") + if backend not in self.backend_map: + raise ValueError(f"Unsupported backend {backend!r}. Supported backends: {sorted(self.backend_map)}.") + self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend) self.backend: Backend = BackendFactory.get_backend( - backend=self.backend_map[config.name], + backend=self.backend_map[backend], local_output_path=self.local_output_path, cloud_output_path=self.cloud_output_path, predictor_type=self.predictor_type, - config=config, + role=role, ) @property @@ -170,12 +174,9 @@ def fit( instance_type: str = "ml.m5.2xlarge", instance_count: Union[int, str] = "auto", volume_size: int = 256, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, timeout: int = 24 * 60 * 60, wait: bool = True, - environment: Optional[Dict[str, str]] = None, - use_spot_instances: bool = False, - max_wait: Optional[int] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, **kwargs, ) -> CloudPredictor: @@ -199,7 +200,7 @@ def fit( Training container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with a predictor-specific prefix. @@ -211,26 +212,18 @@ def fit( volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training. Must be large enough to store training data if File Mode is used (which is the default). - image_uri: Optional[str], default = None - Custom training container image. If None, the official AutoGluon DLC for ``framework_version`` is used. timeout: int, default = 24*60*60 Timeout in seconds for training. This timeout doesn't include time for pre-processing or launching up the training job. wait: bool, default = True 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the training container. - use_spot_instances: bool, default = False - Whether to train on managed spot instances. - max_wait: Optional[int], default = None - Maximum seconds to wait for spot capacity plus training time. Defaults to ``timeout``. Requires - ``use_spot_instances=True``. 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. @@ -271,12 +264,9 @@ def fit( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - image_uri=image_uri, + custom_image_uri=custom_image_uri, timeout=timeout, wait=wait, - environment=environment, - use_spot_instances=use_spot_instances, - max_wait=max_wait, backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) @@ -385,12 +375,11 @@ def deploy( framework_version: str = "latest", instance_type: Optional[str] = None, initial_instance_count: int = 1, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, volume_size: Optional[int] = None, wait: bool = True, inference_mode: Literal["realtime", "serverless"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, - environment: Optional[Dict[str, str]] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> None: """ @@ -409,14 +398,14 @@ def deploy( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. instance_type: Optional[str], default = None Instance to be deployed for the endpoint. Defaults to ``ml.m5.2xlarge``. Must be ``None`` when ``inference_mode="serverless"``. initial_instance_count: int, default = 1, Initial number of instances to be deployed for the endpoint. Ignored when ``inference_mode="serverless"``. - image_uri: Optional[str], default = None, + custom_image_uri: Optional[str], default = None, Custom image to use to deploy endpoint with. If not specified, with use official DLC image: https://github.com/aws/deep-learning-containers/blob/master/available_images.md#autogluon-inference-containers @@ -431,14 +420,13 @@ def deploy( (no instance management, scales to zero). inference_config: Optional[Dict[str, Any]], default = None Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``). - environment: Optional[Dict[str, str]], default = None - Environment variables set in the inference container. 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. """ if inference_mode == "serverless" and instance_type is not None: raise ValueError("`instance_type` must not be set when `inference_mode='serverless'`.") @@ -450,12 +438,11 @@ def deploy( framework_version=framework_version, instance_type=instance_type, initial_instance_count=initial_instance_count, - image_uri=image_uri, + custom_image_uri=custom_image_uri, volume_size=volume_size, wait=wait, inference_mode=inference_mode, inference_config=inference_config, - environment=environment, backend_overrides=backend_overrides, ) @@ -564,12 +551,11 @@ def predict( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, download: bool = True, persist: bool = True, save_path: Optional[str] = None, - environment: Optional[Dict[str, str]] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ @@ -592,9 +578,9 @@ def predict( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None - Name of the launched batch transform job. + Name of the launched training job. If None, CloudPredictor creates one with a predictor-specific prefix. instance_count: int, default = 1, Number of instances used to do batch transform. @@ -614,13 +600,12 @@ def predict( 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the inference container. 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. Returns ------- @@ -636,12 +621,11 @@ def predict( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, download=download, persist=persist, save_path=save_path, - environment=environment, backend_overrides=backend_overrides, ) @@ -656,12 +640,11 @@ def predict_proba( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, download: bool = True, persist: bool = True, save_path: Optional[str] = None, - environment: Optional[Dict[str, 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]]]: """ @@ -687,9 +670,9 @@ def predict_proba( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None - Name of the launched batch transform job. + Name of the launched training job. If None, CloudPredictor creates one with a predictor-specific prefix. instance_count: int, default = 1, Number of instances used to do batch transform. @@ -709,13 +692,12 @@ def predict_proba( 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the inference container. 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. Returns ------- @@ -734,12 +716,11 @@ def predict_proba( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, download=download, persist=persist, save_path=save_path, - environment=environment, backend_overrides=backend_overrides, ) diff --git a/src/autogluon/cloud/predictor/tabular_cloud_predictor.py b/src/autogluon/cloud/predictor/tabular_cloud_predictor.py index 15cdf736..c9a613ed 100644 --- a/src/autogluon/cloud/predictor/tabular_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/tabular_cloud_predictor.py @@ -51,12 +51,9 @@ def fit_predict( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 256, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, predictions_path: Optional[str] = None, - environment: Optional[Dict[str, str]] = None, - use_spot_instances: bool = False, - max_wait: Optional[int] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ @@ -83,7 +80,7 @@ def fit_predict( Whether to include the leaderboard in the output artifact. framework_version: str, default = `latest` Training container version of autogluon. If `latest`, will use the latest available container version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-tabular``. instance_type: str, default = 'ml.m5.2xlarge' @@ -92,7 +89,7 @@ def fit_predict( Number of instances used to fit the predictor. volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training. - image_uri: Optional[str], default = None + custom_image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True Whether the call should wait until the job completes. @@ -100,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``. - environment, use_spot_instances, max_wait, backend_overrides: - Same as in :meth:`fit`. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields, same as in :meth:`fit`. Returns ------- @@ -121,12 +118,9 @@ def fit_predict( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, predictions_path=predictions_path, - environment=environment, - use_spot_instances=use_spot_instances, - max_wait=max_wait, backend_overrides=backend_overrides, ) if result is None: # wait=False @@ -149,12 +143,9 @@ def fit_predict_proba( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 256, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, predictions_path: Optional[str] = None, - environment: Optional[Dict[str, str]] = None, - use_spot_instances: bool = False, - max_wait: Optional[int] = 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]]]: """ @@ -179,7 +170,7 @@ def fit_predict_proba( leaderboard: bool, default = True Whether to include the leaderboard in the output artifact. framework_version: str, default = `latest` - Training container version of autogluon. If `image_uri` is set, this argument is ignored. + Training container version of autogluon. If `custom_image_uri` is set, this argument is ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-tabular``. instance_type: str, default = 'ml.m5.2xlarge' @@ -188,15 +179,15 @@ def fit_predict_proba( Number of instances used to fit the predictor. volume_size: int, default = 256 Size in GB of the EBS volume to use for storing input data during training. - image_uri: Optional[str], default = None + custom_image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True Whether the call should wait until the job completes. predictions_path: Optional[str] S3 URL where predictions will be written by the training container. Defaults to ``{cloud_output_path}/{job_name}/predictions.csv``. - environment, use_spot_instances, max_wait, backend_overrides: - Same as in :meth:`fit`. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields, same as in :meth:`fit`. Returns ------- @@ -220,11 +211,8 @@ def fit_predict_proba( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, - environment=environment, - use_spot_instances=use_spot_instances, - max_wait=max_wait, backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) diff --git a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py index 15c5c6bb..bc39ac9d 100644 --- a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py @@ -51,13 +51,10 @@ def fit( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 100, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, - known_covariates: Optional[Union[str, Path, pd.DataFrame]] = None, - environment: Optional[Dict[str, str]] = None, - use_spot_instances: bool = False, - max_wait: Optional[int] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, + known_covariates: Optional[Union[str, Path, pd.DataFrame]] = None, **kwargs, ) -> TimeSeriesCloudPredictor: """ @@ -94,7 +91,7 @@ def fit( Training container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-timeseries``. @@ -105,21 +102,14 @@ def fit( volume_size: int, default = 100 Size in GB of the EBS volume to use for storing input data during training. Must be large enough to store training data if File Mode is used (which is the default). - image_uri: Optional[str], default = None + custom_image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True 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. - environment: Optional[Dict[str, str]], default = None - Environment variables set in the training container. - use_spot_instances: bool, default = False - Whether to train on managed spot instances. - max_wait: Optional[int], default = None - Maximum seconds to wait for spot capacity plus training time. Requires ``use_spot_instances=True``. backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None - Raw ``CreateTrainingJob`` request fields under the ``"create_training_job"`` key. See - :meth:`TabularCloudPredictor.fit` for details. + Raw SageMaker request fields under ``"create_training_job"``. See :meth:`TabularCloudPredictor.fit`. Returns ------- @@ -128,11 +118,10 @@ def fit( assert not self.backend.is_fit, ( "Predictor is already fit! To fit additional models, create a new `CloudPredictor`" ) - # `extra_ag_args` is an internal channel for `fit_predict`, not part of the public signature. + # `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, @@ -160,11 +149,8 @@ def fit( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, - environment=environment, - use_spot_instances=use_spot_instances, - max_wait=max_wait, backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) @@ -232,12 +218,11 @@ def predict( job_name: Optional[str] = None, instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, download: bool = True, persist: bool = True, save_path: Optional[str] = None, - environment: Optional[Dict[str, str]] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ @@ -266,7 +251,7 @@ def predict( Inference container version of autogluon. If `latest`, will use the latest available container version. If provided a specific version, will use this version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-timeseries``. @@ -277,9 +262,7 @@ 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. - image_uri: Optional[str], default = None - Custom inference container image. If set, ``framework_version`` is ignored. - download, persist, save_path, environment, backend_overrides: + download, persist, save_path, backend_overrides: Same as in :meth:`TabularCloudPredictor.predict`. """ return self.backend.predict( @@ -291,12 +274,11 @@ def predict( job_name=job_name, instance_type=instance_type, instance_count=instance_count, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, download=download, persist=persist, save_path=save_path, - environment=environment, backend_overrides=backend_overrides, ) @@ -326,11 +308,8 @@ def fit_predict( instance_type: str = "ml.m5.2xlarge", instance_count: int = 1, volume_size: int = 100, - image_uri: Optional[str] = None, + custom_image_uri: Optional[str] = None, wait: bool = True, - environment: Optional[Dict[str, str]] = None, - use_spot_instances: bool = False, - max_wait: Optional[int] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ @@ -371,7 +350,7 @@ def fit_predict( names ``item_id`` and ``timestamp``, regardless of the ``id_column`` / ``timestamp_column`` passed in. framework_version: str, default = `latest` Training container version of autogluon. If `latest`, will use the latest available container version. - If `image_uri` is set, this argument will be ignored. + If `custom_image_uri` is set, this argument will be ignored. job_name: str, default = None Name of the launched training job. If None, CloudPredictor creates one with prefix ``ag-cloud-timeseries``. instance_type: str, default = 'ml.m5.2xlarge' @@ -380,12 +359,12 @@ def fit_predict( Number of instances used to fit the predictor. volume_size: int, default = 100 Size in GB of the EBS volume to use for storing input data during training. - image_uri: Optional[str], default = None + custom_image_uri: Optional[str], default = None Custom container image URI. If set, ``framework_version`` is ignored. wait: bool, default = True Whether the call should wait until the job completes. - environment, use_spot_instances, max_wait, backend_overrides: - Same as in :meth:`fit`. + backend_overrides: Optional[Dict[str, Dict[str, Any]]], default = None + Raw SageMaker request fields, same as in :meth:`fit`. Returns ------- @@ -409,11 +388,8 @@ def fit_predict( instance_type=instance_type, instance_count=instance_count, volume_size=volume_size, - image_uri=image_uri, + custom_image_uri=custom_image_uri, wait=wait, - environment=environment, - use_spot_instances=use_spot_instances, - max_wait=max_wait, backend_overrides=backend_overrides, extra_ag_args=extra_ag_args, ) diff --git a/src/autogluon/cloud/scripts/sagemaker_scripts/train.py b/src/autogluon/cloud/scripts/sagemaker_scripts/train.py index bdba3ea5..f4efaac3 100644 --- a/src/autogluon/cloud/scripts/sagemaker_scripts/train.py +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/train.py @@ -81,7 +81,7 @@ def prepare_data(data_file, predictor_type, ag_args, static_features_df=None): print(f"Args: {args}") - # See SageMaker-specific environment variables: https://github.com/aws/sagemaker-training-toolkit/blob/master/ENVIRONMENT_VARIABLES.md + # See SageMaker-specific environment variables: https://sagemaker.readthedocs.io/en/v2/overview.html#prepare-a-training-script os.makedirs(args.output_data_dir, mode=0o777, exist_ok=True) ag_args_file = get_input_path(args.ag_args) diff --git a/src/autogluon/cloud/utils/ag_sagemaker.py b/src/autogluon/cloud/utils/ag_sagemaker.py index 3c6e672e..c1e8aeba 100644 --- a/src/autogluon/cloud/utils/ag_sagemaker.py +++ b/src/autogluon/cloud/utils/ag_sagemaker.py @@ -42,19 +42,15 @@ def resolve_image_uri( ) -def upload_training_code(entry_point: str, source_dir: Optional[str], sagemaker_session, s3_uri_prefix: str) -> str: - """Bundle the training entry point (or ``source_dir`` containing it) as ``sourcedir.tar.gz`` and upload it. +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. 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: - if source_dir: - for name in os.listdir(source_dir): - tar.add(os.path.join(source_dir, name), arcname=name) - else: - tar.add(entry_point, arcname=os.path.basename(entry_point)) + 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) @@ -128,14 +124,3 @@ def repack_model_with_serving_code( 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 - - -def create_serve_script_tarball(entry_point: str, output_dir: str) -> str: - """Create a minimal ``model.tar.gz`` that only contains the serving code under ``code/``.""" - from ..scripts import ScriptManager # deferred: importing scripts pulls in the backend package - - tarball_path = os.path.join(output_dir, "model.tar.gz") - with tarfile.open(tarball_path, "w:gz") as tar: - tar.add(entry_point, arcname=f"code/{os.path.basename(entry_point)}") - tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils") - return tarball_path diff --git a/src/autogluon/cloud/utils/aws_utils.py b/src/autogluon/cloud/utils/aws_utils.py index 3cdf42a4..79649e52 100644 --- a/src/autogluon/cloud/utils/aws_utils.py +++ b/src/autogluon/cloud/utils/aws_utils.py @@ -106,7 +106,7 @@ def get_execution_role(session: Optional[AwsSession] = None) -> str: if match is None: raise ValueError( f"Cannot infer a SageMaker execution role from the current AWS identity {caller_arn}. Pass " - "`backend=SageMakerConfig(role_arn=)`, or run `autogluon.cloud.bootstrap()` / `register()` once to " + "`role=`, or run `autogluon.cloud.bootstrap()` / `register()` once to " "persist a role." ) partition, account, role_name = match.groups() @@ -118,7 +118,7 @@ def get_execution_role(session: Optional[AwsSession] = None) -> str: 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_arn` if this is wrong." + f"Could not look up role {role_name!r} in IAM, using {role_arn}. Pass `role=` if this is wrong." ) return role_arn @@ -225,13 +225,12 @@ def setup_sagemaker_session( connect_timeout: int = 60, read_timeout: int = 60, retries: Optional[dict] = None, - region: Optional[str] = None, **kwargs, ) -> AwsSession: """ Setup an :class:`AwsSession` with a given configuration - Region resolution (only when ``boto_session`` is not provided): use ``region``, then read from + 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, shared config). Raises if no region can be resolved at all. @@ -240,9 +239,6 @@ def setup_sagemaker_session( boto_session Pre-built ``boto3.Session`` to wrap. If provided, region resolution is skipped and the session is used as-is. - region - Explicit AWS region. Takes precedence over saved configuration and the boto3 default region. - Ignored when ``boto_session`` is provided. config A botocore.Config object providing the intended configuration https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html @@ -273,7 +269,7 @@ def setup_sagemaker_session( retries = {"max_attempts": 20} config = Config(connect_timeout=connect_timeout, read_timeout=read_timeout, retries=retries, **kwargs) if boto_session is None: - boto_session = boto3.Session(region_name=region or _resolve_sagemaker_region()) + boto_session = boto3.Session(region_name=_resolve_sagemaker_region()) if boto_session.region_name is None: raise ValueError( "AWS region could not be resolved. Set it in `~/.autogluon/cloud.yaml` (e.g. via " diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index 6be4e3a1..ca3aedf9 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -16,11 +16,10 @@ BATCH_PREDICT_OVERRIDE_KEYS = ("create_model", "create_transform_job") _REMOVED_KWARGS = { - "backend_kwargs": "named arguments, `backend=SageMakerConfig(...)`, or `backend_overrides`", - "custom_image_uri": "`image_uri`", + "backend_kwargs": "`backend_overrides` (and the `download` / `persist` / `save_path` arguments of `predict()`)", "autogluon_sagemaker_estimator_kwargs": "`backend_overrides={'create_training_job': ...}`", "fit_kwargs": "`backend_overrides={'create_training_job': ...}`", - "model_kwargs": "`environment` or `backend_overrides={'create_model': ...}`", + "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': ...}`", diff --git a/src/autogluon/cloud/utils/tag_utils.py b/src/autogluon/cloud/utils/tag_utils.py index abe5a313..3fe697b7 100644 --- a/src/autogluon/cloud/utils/tag_utils.py +++ b/src/autogluon/cloud/utils/tag_utils.py @@ -10,19 +10,18 @@ def build_tags( module: str, - extra_tags: Optional[Dict[str, str]] = None, - user_tags: Optional[Dict[str, str]] = None, -) -> Dict[str, str]: - """Final tags for a SageMaker resource: defaults + extras + user, with user winning on key collision. + extra_tags: Optional[List[Dict[str, str]]] = None, + user_tags: Optional[List[Dict[str, str]]] = None, +) -> List[Dict[str, str]]: + """Final tag list for a SageMaker resource: defaults + extras + user, with user winning on key collision. Defaults are skipped entirely when ``AG_CLOUD_DISABLE_DEFAULT_TAGS`` is truthy, so customers in tag-restricted AWS orgs can opt out without losing other functionality. """ if os.environ.get(DISABLE_DEFAULT_TAGS_ENV, "").lower() in ("1", "true", "yes"): - return dict(user_tags or {}) - return {"autogluon-cloud-module": module, **(extra_tags or {}), **(user_tags or {})} - - -def to_request_tags(tags: Dict[str, str]) -> List[Dict[str, str]]: - """Convert ``{key: value}`` tags to the list format of SageMaker API requests.""" - return [{"Key": key, "Value": value} for key, value in tags.items()] + return list(user_tags or []) + base = [{"Key": "autogluon-cloud-module", "Value": module}] + list(extra_tags or []) + if not user_tags: + return base + user_keys = {t["Key"] for t in user_tags} + return [t for t in base if t["Key"] not in user_keys] + list(user_tags) diff --git a/tests/conftest.py b/tests/conftest.py index 9d22c5f0..d79441ad 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -134,7 +134,7 @@ def tag_resources_with_ci_run(): build_tags = sagemaker_backend.build_tags def build_tags_with_ci_run(*args, **kwargs): - return {**build_tags(*args, **kwargs), CI_RUN_TAG: run_id} + return build_tags(*args, **kwargs) + [{"Key": CI_RUN_TAG, "Value": run_id}] with pytest.MonkeyPatch.context() as mp: mp.setattr(sagemaker_backend, "build_tags", build_tags_with_ci_run) diff --git a/tests/unittests/general/test_aws_session.py b/tests/unittests/general/test_aws_session.py index 224f1df3..829c1d1f 100644 --- a/tests/unittests/general/test_aws_session.py +++ b/tests/unittests/general/test_aws_session.py @@ -124,7 +124,7 @@ def test_execution_role_without_iam_access_falls_back_to_sts_arn(role_name, expe def test_execution_role_rejects_iam_users(): session, _ = _session_with_caller("arn:aws:iam::123456789012:user/alice") - with pytest.raises(ValueError, match="SageMakerConfig\\(role_arn="): + with pytest.raises(ValueError, match="role="): get_execution_role(session) diff --git a/tests/unittests/general/test_aws_utils.py b/tests/unittests/general/test_aws_utils.py index 1c3e94f1..a73cf9aa 100644 --- a/tests/unittests/general/test_aws_utils.py +++ b/tests/unittests/general/test_aws_utils.py @@ -9,7 +9,7 @@ CloudConfig, save_config, ) -from autogluon.cloud.utils.aws_utils import resolve_cloud_output_path, resolve_execution_role, setup_sagemaker_session +from autogluon.cloud.utils.aws_utils import resolve_cloud_output_path, resolve_execution_role @pytest.fixture(autouse=True) @@ -90,18 +90,6 @@ def test_role_fallback_uses_the_backend_session(): get_role.assert_called_once_with(session) -@pytest.mark.parametrize("explicit_region", [None, "eu-west-1"]) -def test_session_region_uses_explicit_setting_before_saved_config(explicit_region): - import boto3 - - _save_role_in_config("sagemaker", "role") - with mock.patch("autogluon.cloud.utils.aws_utils.boto3.Session", wraps=boto3.Session) as session_cls: - session = setup_sagemaker_session(region=explicit_region) - expected_region = explicit_region or "us-east-1" - session_cls.assert_called_once_with(region_name=expected_region) - assert session.boto_region_name == expected_region - - def _save_bucket_in_config(backend_name: str, bucket: str) -> None: save_config( CloudConfig( diff --git a/tests/unittests/general/test_backend_config.py b/tests/unittests/general/test_backend_config.py deleted file mode 100644 index e7d42ac3..00000000 --- a/tests/unittests/general/test_backend_config.py +++ /dev/null @@ -1,209 +0,0 @@ -"""Shared backend configuration and job-bound foundation-model results, without AWS calls.""" - -import pickle -from pathlib import Path -from unittest import mock - -import pandas as pd -import pytest - -from autogluon.cloud import FoundationModel, SageMakerConfig, TabularCloudPredictor, TimeSeriesCloudPredictor -from autogluon.cloud.backend.backend_factory import BackendFactory -from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend -from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend -from autogluon.cloud.backend.timeseries_sagemaker_backend import TimeSeriesSagemakerBackend -from autogluon.cloud.job.sagemaker_job import SageMakerFitJob - -SB = "autogluon.cloud.backend.sagemaker_backend" - - -@pytest.fixture(autouse=True) -def stub_aws(monkeypatch): - session_factory = mock.Mock( - side_effect=lambda *, region=None: mock.MagicMock(boto_region_name=region or "us-east-1") - ) - monkeypatch.setattr(f"{SB}.setup_sagemaker_session", session_factory) - monkeypatch.setattr( - f"{SB}.resolve_execution_role", - lambda role, backend_name, session=None: role or "arn:aws:iam::0:role/default", - ) - monkeypatch.setattr("autogluon.cloud.utils.aws_utils._s3_prefix_has_objects", lambda *_: False) - return session_factory - - -@pytest.mark.parametrize( - "predictor_cls,model_id", - [(TabularCloudPredictor, "mitra-classifier"), (TimeSeriesCloudPredictor, "chronos-2")], -) -def test_shared_config_creates_independent_predictor_and_model_backends(tmp_path, predictor_cls, model_id): - config = SageMakerConfig( - region="eu-west-1", - role_arn="arn:aws:iam::0:role/custom", - vpc_config={"subnets": ["subnet-1"], "security_group_ids": ["sg-1"]}, - output_kms_key="output-key", - tags={"team": "forecasting"}, - ) - predictor = predictor_cls(backend=config, cloud_output_path="s3://b/predictor", local_output_path=str(tmp_path)) - model = FoundationModel(model_id, backend=config, cloud_output_path="s3://b/model") - - assert predictor.backend.config == model._backend.config == config - assert predictor.backend is not model._backend - assert predictor.backend._fit_job is not model._backend._fit_job - assert predictor.backend.sagemaker_session is not model._backend.sagemaker_session - - predictor.backend.config.tags["team"] = "other" - predictor.backend.config.vpc_config["subnets"].append("subnet-2") - predictor.backend.attach_endpoint("endpoint") - assert config.tags == model._backend.config.tags == {"team": "forecasting"} - assert ( - config.vpc_config - == model._backend.config.vpc_config - == { - "subnets": ["subnet-1"], - "security_group_ids": ["sg-1"], - } - ) - assert model._backend.endpoint_name is None - - -def test_default_backend_name_uses_the_same_config_resolver(tmp_path): - predictor = TabularCloudPredictor(local_output_path=str(tmp_path), cloud_output_path="s3://b/run") - model = FoundationModel("chronos-2", cloud_output_path="s3://b/model") - assert predictor.backend.config == model._backend.config - assert predictor.backend.config.role_arn == "arn:aws:iam::0:role/default" - assert predictor.backend.config.region == "us-east-1" - - -@pytest.mark.parametrize("backend", ["unknown", "ray", "ray_aws"]) -def test_invalid_backend_rejected_by_both_entry_points(tmp_path, backend): - with pytest.raises(ValueError): - TabularCloudPredictor(backend=backend, local_output_path=str(tmp_path)) - with pytest.raises(ValueError): - FoundationModel("chronos-2", backend=backend) - - -def test_backend_resolver_rejects_untyped_dict(): - with pytest.raises(TypeError, match="SageMakerConfig"): - BackendFactory.resolve_config({"region": "us-east-1"}) - - -def test_predictor_reload_keeps_the_original_region(tmp_path, stub_aws): - predictor = TabularCloudPredictor( - backend=SageMakerConfig(region="eu-west-1"), - cloud_output_path="s3://b/run", - local_output_path=str(tmp_path), - ) - restored = pickle.loads(pickle.dumps(predictor)) - assert restored.backend.config == predictor.backend.config - stub_aws.assert_called_with(region="eu-west-1") - - -def test_real_training_submissions_keep_job_objects_and_inputs_separate(tmp_path, monkeypatch): - """Later predictions must not overwrite the first job's handle, config or data channels.""" - backend = TabularSagemakerBackend( - local_output_path=str(tmp_path), cloud_output_path="s3://b/run", predictor_type="tabular" - ) - uploads = {} - - def upload(path, bucket, key_prefix): - uri = f"s3://{bucket}/{key_prefix}/{Path(path).name}" - if Path(path).is_file(): - uploads[uri] = Path(path).read_bytes() - return uri - - def run(job, training_job_request, framework_version, wait): - job._job_name = training_job_request["TrainingJobName"] - job.request = training_job_request - - backend.sagemaker_session.upload_data.side_effect = upload - monkeypatch.setattr(SageMakerFitJob, "run", run) - monkeypatch.setattr(SageMakerFitJob, "_get_job_status", lambda self: "Completed") - monkeypatch.setattr(SB + ".upload_training_code", lambda **kwargs: "s3://b/code") - monkeypatch.setattr(backend, "_load_fit_predict_results", lambda job: pd.DataFrame({"job": [job.job_name]})) - futures, requests = [], [] - for name, value in [("first", 1), ("second", 2)]: - backend.fit( - predictor_init_args={"label": "y"}, - predictor_fit_args={}, - data_channels={"train_data": pd.DataFrame({"x": [value], "y": [0]})}, - job_name=name, - image_uri="example.com/autogluon:train", - wait=False, - extra_ag_args={"predict_after_fit": True, "save_predictor": False}, - ) - futures.append(backend.get_prediction_future()) - requests.append(backend._fit_job.request) - assert [future.job_name for future in futures] == ["first", "second"] - assert [future.result()["job"].iloc[0] for future in futures] == ["first", "second"] - channels = [ - { - channel["ChannelName"]: channel["DataSource"]["S3DataSource"]["S3Uri"] - for channel in request["InputDataConfig"] - } - for request in requests - ] - assert channels[0]["ag_args"] != channels[1]["ag_args"] - assert channels[0]["train_data"] != channels[1]["train_data"] - assert b"/first/predictions.csv" in uploads[channels[0]["ag_args"]] - assert b"/second/predictions.csv" in uploads[channels[1]["ag_args"]] - assert uploads[channels[0]["train_data"]] != uploads[channels[1]["train_data"]] - - -@pytest.mark.parametrize( - "model_id,operation,include_predict", - [ - ("chronos-2", "predict", None), - ("mitra-classifier", "predict", None), - ("mitra-classifier", "predict_proba", True), - ("mitra-classifier", "predict_proba", False), - ], -) -def test_async_model_results_stay_bound_to_the_submitted_job(monkeypatch, model_id, operation, include_predict): - model = FoundationModel(model_id, cloud_output_path="s3://b/model") - frames = [ - pd.DataFrame({"target": ["a"], "a_proba": [0.8], "b_proba": [0.2]}), - pd.DataFrame({"target": ["b"], "a_proba": [0.1], "b_proba": [0.9]}), - ] - jobs = [] - - def submit(self, **kwargs): - assert kwargs["backend_overrides"] == {"create_training_job": {"RetryStrategy": {}}} - job = mock.Mock(job_name=f"job-{len(jobs)}", completed=True) - job.frame = frames[len(jobs)] - jobs.append(job) - self._fit_job = job - - monkeypatch.setattr(TabularSagemakerBackend, "fit", submit) - monkeypatch.setattr(TimeSeriesSagemakerBackend, "fit", submit) - monkeypatch.setattr(SagemakerBackend, "_load_fit_predict_results", lambda self, job: job.frame) - kwargs = {"wait": False, "backend_overrides": {"create_training_job": {"RetryStrategy": {}}}} - if model_id == "chronos-2": - kwargs["data"] = pd.DataFrame({"target": [1.0]}) - else: - kwargs.update( - train_data=pd.DataFrame({"feature": [1], "target": ["a"]}), - test_data=pd.DataFrame({"feature": [2]}), - label="target", - ) - if include_predict is not None: - kwargs["include_predict"] = include_predict - - first = getattr(model, operation)(**kwargs) - second = getattr(model, operation)(**kwargs) - assert first.job_name == "job-0" - assert second.job_name == "job-1" - first_result, second_result = first.result(), second.result() - if model_id == "chronos-2": - pd.testing.assert_frame_equal(first_result, frames[0]) - pd.testing.assert_frame_equal(second_result, frames[1]) - elif operation == "predict": - assert first_result.tolist() == ["a"] - assert second_result.tolist() == ["b"] - else: - if include_predict: - first_pred, first_result = first_result - second_pred, second_result = second_result - assert first_pred.tolist() == ["a"] - assert second_pred.tolist() == ["b"] - assert first_result["a"].tolist() == [0.8] - assert second_result["a"].tolist() == [0.1] diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 8409ecf5..eb5142e9 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -7,7 +7,6 @@ import pandas as pd import pytest -from autogluon.cloud import SageMakerConfig from autogluon.cloud.model import FoundationModel @@ -20,7 +19,7 @@ def _stub_aws(monkeypatch): ) monkeypatch.setattr( "autogluon.cloud.backend.backend_factory.BackendFactory.get_backend", - lambda **kwargs: mock.MagicMock(role_arn="arn:aws:iam::0:role/stub", config=kwargs["config"]), + lambda **kwargs: mock.MagicMock(role_arn="arn:aws:iam::0:role/stub"), ) @@ -47,7 +46,7 @@ def test_to_dict_excludes_runtime_context(): fm = FoundationModel( "chronos-2", cloud_output_path="s3://my-bucket/runs/", - backend=SageMakerConfig(role_arn="arn:aws:iam::0:role/runtime"), + role="arn:aws:iam::0:role/runtime", ) d = fm.to_dict() assert "role" not in d @@ -137,7 +136,7 @@ def test_deploy_passes_artifact_uri_and_overrides_model_path_to_container_dir(): assert call.kwargs["repack"] is False serve_cfg = call.kwargs["fm_serve_config"] assert serve_cfg["hyperparameters"]["model_path"] == "/opt/ml/model/weights" - assert call.kwargs["extra_tags"] == {"autogluon-cloud-model-id": "chronos-2"} + assert {"Key": "autogluon-cloud-model-id", "Value": "chronos-2"} in call.kwargs["extra_tags"] def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): @@ -150,7 +149,7 @@ def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): assert call.kwargs["repack"] is False serve_cfg = call.kwargs["fm_serve_config"] assert serve_cfg["hyperparameters"]["model_path"] == "autogluon/chronos-2" - assert call.kwargs["extra_tags"] == {"autogluon-cloud-model-id": "chronos-2"} + assert {"Key": "autogluon-cloud-model-id", "Value": "chronos-2"} in call.kwargs["extra_tags"] def test_tabular_deploy_uses_tabular_fm_handler_and_returns_tabular_endpoint(): @@ -214,8 +213,7 @@ def test_cache_model_artifact_uploads_with_version_metadata(monkeypatch): """On cache miss, upload_file runs with the version metadata key — that's the cache-invalidation contract.""" from autogluon.cloud.version import __version__ - backend_config = SageMakerConfig(region="eu-west-1", output_kms_key="output-key", tags={"team": "ts"}) - fm = FoundationModel("chronos-2", cloud_output_path="s3://b", backend=backend_config) + fm = FoundationModel("chronos-2", cloud_output_path="s3://b") s3 = mock.MagicMock() fm._backend.sagemaker_session.boto_session.client.return_value = s3 monkeypatch.setattr("autogluon.cloud.model.foundation_model._s3_head_or_none", lambda *_: None) @@ -228,12 +226,9 @@ def test_cache_model_artifact_uploads_with_version_metadata(monkeypatch): new_fm = fm.cache_model_artifact("s3://b/cache") assert new_fm.model_artifact_uri == "s3://b/cache/chronos-2/model.tar.gz" - assert new_fm._backend.config == fm._backend.config == backend_config s3.upload_file.assert_called_once() metadata = s3.upload_file.call_args.kwargs["ExtraArgs"]["Metadata"] assert metadata == {"autogluon-cloud-version": __version__} - assert s3.upload_file.call_args.kwargs["ExtraArgs"]["SSEKMSKeyId"] == "output-key" - assert s3.upload_file.call_args.kwargs["ExtraArgs"]["ServerSideEncryption"] == "aws:kms" def test_cache_model_artifact_raises_on_stale_version_without_overwrite(): diff --git a/tests/unittests/general/test_inference_modes.py b/tests/unittests/general/test_inference_modes.py index 1d51f700..66d5538e 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -4,7 +4,6 @@ import pytest -from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.sagemaker_backend import SagemakerBackend GPU_IMAGE_URI = "123456789012.dkr.ecr.us-east-1.amazonaws.com/autogluon:1.6-cu133-amzn2023" @@ -27,10 +26,8 @@ def deploy_requests(assert_valid_request): ) backend._fit_job = None # deploy a serve-script tarball, not a fit-job artifact - def run(backend_config=None, **kwargs): + def run(**kwargs): backend.endpoint_name = None # allow re-deploy across cases - if backend_config is not None: - backend.config = backend_config backend.deploy(endpoint_name="ep", entry_point="stub.py", **kwargs) client = backend.sagemaker_session.sagemaker_client requests = { @@ -70,7 +67,7 @@ def test_when_deployed_then_model_endpoint_config_and_endpoint_are_linked(deploy def test_when_cuda_13_custom_image_then_inference_ami_is_inferred(deploy_requests): - variant = _variant(deploy_requests(instance_type="ml.g4dn.xlarge", image_uri=GPU_IMAGE_URI)) + 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" @@ -78,7 +75,7 @@ def test_when_inference_ami_is_overridden_then_override_wins(deploy_requests): variant = _variant( deploy_requests( instance_type="ml.g4dn.xlarge", - image_uri=GPU_IMAGE_URI, + custom_image_uri=GPU_IMAGE_URI, backend_overrides={"production_variant": {"InferenceAmiVersion": "custom-ami"}}, ) ) @@ -119,8 +116,11 @@ def test_when_inference_mode_is_unknown_then_value_error_is_raised(deploy_reques deploy_requests(inference_mode="batch") -def test_when_environment_given_then_it_reaches_the_container(deploy_requests): - requests = deploy_requests(instance_type="ml.m5.xlarge", environment={"FOO": "bar"}) +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" @@ -129,19 +129,3 @@ def test_when_environment_given_then_it_reaches_the_container(deploy_requests): 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": {}}) - - -@pytest.mark.parametrize( - ("inference_mode", "volume_key"), - [("realtime", None), ("realtime", "volume-key"), ("serverless", "volume-key")], -) -def test_endpoint_volume_encryption_is_independent_of_output_encryption(deploy_requests, inference_mode, volume_key): - requests = deploy_requests( - inference_mode=inference_mode, - backend_config=SageMakerConfig(output_kms_key="output-key", volume_kms_key=volume_key), - ) - config = requests["endpoint_config"] - if inference_mode == "realtime" and volume_key is not None: - assert config["KmsKeyId"] == volume_key - else: - assert "KmsKeyId" not in config diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index 59970477..b6a89717 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -3,7 +3,6 @@ import pandas as pd import pytest -from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend from autogluon.cloud.utils.dlc_utils import infer_sagemaker_ami_version @@ -72,7 +71,6 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value( local_output_path="/tmp/test", cloud_output_path="s3://bucket/run", predictor_type="tabular", - config=SageMakerConfig(output_kms_key="output-key"), ) backend._fit_job = mock.MagicMock() backend._predict( @@ -80,7 +78,7 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value( predictor_path="s3://bucket/model.tar.gz", job_name="job", instance_type="ml.g4dn.xlarge", - image_uri=GPU_IMAGE_URI, + custom_image_uri=GPU_IMAGE_URI, wait=False, download=False, persist=False, @@ -91,5 +89,3 @@ def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value( assert_valid_request("CreateTransformJob", request) assert request["TransformResources"]["TransformAmiVersion"] == expected assert request["TransformResources"]["InstanceType"] == "ml.g4dn.xlarge" - assert request["TransformOutput"]["KmsKeyId"] == "output-key" - assert "VolumeKmsKeyId" not in request["TransformResources"] diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index 490f245b..c93ac817 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -3,7 +3,6 @@ import pandas as pd import pytest -from autogluon.cloud import SageMakerConfig from autogluon.cloud.backend.tabular_sagemaker_backend import TabularSagemakerBackend from autogluon.cloud.utils.sagemaker_api import ( check_override_keys, @@ -34,9 +33,7 @@ def test_reject_legacy_kwargs_points_to_replacement(): def fit(**kwargs): return kwargs - assert fit(image_uri="x") == {"image_uri": "x"} - with pytest.raises(TypeError, match="Use `image_uri` instead"): - fit(custom_image_uri="x") + assert fit(custom_image_uri="x") == {"custom_image_uri": "x"} with pytest.raises(TypeError, match="backend_overrides"): fit(backend_kwargs={}) @@ -82,19 +79,18 @@ def fit_request(tmp_path, assert_valid_request): ), ): - def run(backend_config=None, **fit_kwargs): + def run(**fit_kwargs): backend = TabularSagemakerBackend( local_output_path=str(tmp_path), cloud_output_path="s3://bucket/run", predictor_type="tabular", - config=backend_config, ) backend.fit( predictor_init_args={"label": "y"}, predictor_fit_args={}, data_channels={"train_data": pd.DataFrame({"x": [1], "y": [0]})}, job_name="job", - image_uri="example.com/autogluon:train", + custom_image_uri="example.com/autogluon:train", **fit_kwargs, ) request = fit_job_cls.return_value.run.call_args.kwargs["training_job_request"] @@ -116,55 +112,18 @@ def test_fit_builds_script_mode_training_job(fit_request): assert request["StoppingCondition"] == {"MaxRuntimeInSeconds": 3600} assert request["OutputDataConfig"] == {"S3OutputPath": "s3://bucket/run/model"} assert {"Key": "autogluon-cloud-module", "Value": "tabular"} in request["Tags"] - assert "VpcConfig" not in request -def test_fit_applies_infra_settings_spot_and_overrides(fit_request): +def test_fit_applies_overrides(fit_request): request = fit_request( - backend_config=SageMakerConfig( - vpc_config={"subnets": ["s-1"], "security_group_ids": ["sg-1"]}, - output_kms_key="output-key", - volume_kms_key="volume-key", - tags={"team": "ts"}, - ), - timeout=3600, - environment={"FOO": "bar"}, - use_spot_instances=True, backend_overrides={"create_training_job": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}, ) - assert request["VpcConfig"] == {"Subnets": ["s-1"], "SecurityGroupIds": ["sg-1"]} - assert request["OutputDataConfig"]["KmsKeyId"] == "output-key" - assert request["ResourceConfig"]["VolumeKmsKeyId"] == "volume-key" - assert {"Key": "team", "Value": "ts"} in request["Tags"] - assert request["Environment"] == {"FOO": "bar"} - assert request["EnableManagedSpotTraining"] is True - assert request["StoppingCondition"]["MaxWaitTimeInSeconds"] == 3600 assert request["RetryStrategy"] == {"MaximumRetryAttempts": 2} -@pytest.mark.parametrize( - "vpc_config", - [{"subnets": ["s-1"]}, {"subnets": ["s-1"], "security_group_ids": ["sg-1"], "security_groups": ["sg-2"]}], -) -def test_fit_rejects_malformed_vpc_config(fit_request, vpc_config): - with pytest.raises(ValueError, match="security_group"): - fit_request(backend_config=SageMakerConfig(vpc_config=vpc_config)) - - -def test_output_encryption_does_not_set_volume_key_on_nvme_instance(fit_request): - request = fit_request( - backend_config=SageMakerConfig(output_kms_key="output-key"), - instance_type="ml.g5.xlarge", - ) - assert request["OutputDataConfig"]["KmsKeyId"] == "output-key" - assert "VolumeKmsKeyId" not in request["ResourceConfig"] - - -def test_fit_rejects_local_mode_and_max_wait_without_spot(fit_request): +def test_fit_rejects_local_mode(fit_request): with pytest.raises(ValueError, match="local mode"): fit_request(instance_type="local") - with pytest.raises(ValueError, match="use_spot_instances"): - fit_request(max_wait=100) def test_misspelled_override_field_fails_request_validation(fit_request): diff --git a/tests/unittests/general/test_tabular_foundation_model_predict.py b/tests/unittests/general/test_tabular_foundation_model_predict.py index 74966fd3..e6ee928b 100644 --- a/tests/unittests/general/test_tabular_foundation_model_predict.py +++ b/tests/unittests/general/test_tabular_foundation_model_predict.py @@ -33,21 +33,10 @@ def _stub_aws(monkeypatch): "autogluon.cloud.model.foundation_model.resolve_cloud_output_path", lambda path, backend_name: path or "s3://stub/output", ) - - def make_backend(**kwargs): - backend = mock.MagicMock(role_arn="arn:aws:iam::0:role/stub", config=kwargs["config"]) - - def get_prediction_future(*, result_transform=None): - def load(): - raw = backend.get_fit_predict_results() - return raw if result_transform is None else result_transform(raw) - - return JobPredictionFuture(job=backend._fit_job, result_loader=load) - - backend.get_prediction_future.side_effect = get_prediction_future - return backend - - monkeypatch.setattr("autogluon.cloud.backend.backend_factory.BackendFactory.get_backend", make_backend) + monkeypatch.setattr( + "autogluon.cloud.backend.backend_factory.BackendFactory.get_backend", + lambda **kwargs: mock.MagicMock(role_arn="arn:aws:iam::0:role/stub"), + ) def _make_fm(model_id="mitra-classifier", result=CLASSIFICATION_FRAME): diff --git a/tests/unittests/general/test_tags.py b/tests/unittests/general/test_tags.py index 85192e9f..f52777a7 100644 --- a/tests/unittests/general/test_tags.py +++ b/tests/unittests/general/test_tags.py @@ -2,36 +2,40 @@ import pytest -from autogluon.cloud.utils.tag_utils import DISABLE_DEFAULT_TAGS_ENV, build_tags, to_request_tags +from autogluon.cloud.utils.tag_utils import DISABLE_DEFAULT_TAGS_ENV, build_tags def test_when_no_extras_or_user_then_only_module_tag_is_returned(): - assert build_tags("timeseries") == {"autogluon-cloud-module": "timeseries"} + assert build_tags("timeseries") == [{"Key": "autogluon-cloud-module", "Value": "timeseries"}] -def test_when_extra_tags_provided_then_added_to_module_tag(): - tags = build_tags("timeseries", extra_tags={"autogluon-cloud-model-id": "chronos-2"}) - assert tags == {"autogluon-cloud-module": "timeseries", "autogluon-cloud-model-id": "chronos-2"} +def test_when_extra_tags_provided_then_appended_after_module(): + tags = build_tags("timeseries", extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": "chronos-2"}]) + assert tags == [ + {"Key": "autogluon-cloud-module", "Value": "timeseries"}, + {"Key": "autogluon-cloud-model-id", "Value": "chronos-2"}, + ] def test_when_user_tag_collides_with_default_then_user_wins(): - tags = build_tags("timeseries", user_tags={"autogluon-cloud-module": "override"}) - assert tags == {"autogluon-cloud-module": "override"} + tags = build_tags("timeseries", user_tags=[{"Key": "autogluon-cloud-module", "Value": "override"}]) + assert tags == [{"Key": "autogluon-cloud-module", "Value": "override"}] -def test_when_user_tags_unique_then_added_to_defaults(): - tags = build_tags("tabular", user_tags={"Owner": "team"}) - assert tags == {"autogluon-cloud-module": "tabular", "Owner": "team"} +def test_when_user_tags_unique_then_appended_after_defaults(): + tags = build_tags("tabular", user_tags=[{"Key": "Owner", "Value": "team"}]) + assert tags == [ + {"Key": "autogluon-cloud-module", "Value": "tabular"}, + {"Key": "Owner", "Value": "team"}, + ] @pytest.mark.parametrize("value", ["1", "true", "True", "yes"]) def test_when_disable_env_var_set_then_defaults_and_extras_are_skipped(monkeypatch, value): """Extras are AG-cloud defaults too — opt-out drops them along with module.""" monkeypatch.setenv(DISABLE_DEFAULT_TAGS_ENV, value) - assert build_tags("timeseries") == {} - assert build_tags("timeseries", extra_tags={"autogluon-cloud-model-id": "chronos-2"}) == {} - assert build_tags("timeseries", user_tags={"Owner": "team"}) == {"Owner": "team"} - - -def test_to_request_tags_uses_api_field_names(): - assert to_request_tags({"Owner": "team"}) == [{"Key": "Owner", "Value": "team"}] + assert build_tags("timeseries") == [] + assert build_tags("timeseries", extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": "chronos-2"}]) == [] + assert build_tags("timeseries", user_tags=[{"Key": "Owner", "Value": "team"}]) == [ + {"Key": "Owner", "Value": "team"} + ] diff --git a/tests/unittests/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index 5d54f800..be021305 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -47,7 +47,7 @@ def test_tabular_train(test_helper, framework_version, shared_training_job_name) predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), job_name=shared_training_job_name, ) info = predictor.info() @@ -74,7 +74,7 @@ def test_tabular_endpoint_lifecycle(test_helper, framework_version, shared_train predictor.deploy( framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=predictor.endpoint_name)["EndpointArn"] test_helper.assert_ag_cloud_tags(endpoint_arn, module="tabular") @@ -106,7 +106,7 @@ def test_tabular_batch_predict(test_helper, framework_version, shared_training_j pred, pred_proba = predictor.predict_proba( _TEST_DATA, framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) assert isinstance(pred, pd.Series) assert isinstance(pred_proba, pd.DataFrame) @@ -129,7 +129,7 @@ def test_tabular_deploy_trained_artifact(test_helper, framework_version, shared_ predictor.deploy( predictor_path=artifact_path, framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) test_helper.test_endpoint(predictor, _TEST_DATA) predictor.cleanup_deployment() @@ -152,7 +152,7 @@ def test_tabular_predict_trained_artifact(test_helper, framework_version, shared _TEST_DATA, predictor_path=artifact_path, framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), ) assert isinstance(pred, pd.Series) assert isinstance(pred_proba, pd.DataFrame) @@ -180,7 +180,7 @@ def test_tabular_foundation_model_predict(test_helper, framework_version): label="class", include_predict=True, framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), predictions_path=predictions_path, ) @@ -227,7 +227,7 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): "mitra-classifier", cloud_output_path=(f"s3://autogluon-cloud-ci/test-tabular-fm-deploy/{framework_version}/{timestamp}"), ) - endpoint = model.deploy(image_uri=inference_custom_image_uri) + endpoint = model.deploy(custom_image_uri=inference_custom_image_uri) try: endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=endpoint.endpoint_name)[ "EndpointArn" diff --git a/tests/unittests/timeseries/test_timeseries.py b/tests/unittests/timeseries/test_timeseries.py index 48ea755d..8e42598f 100644 --- a/tests/unittests/timeseries/test_timeseries.py +++ b/tests/unittests/timeseries/test_timeseries.py @@ -59,7 +59,7 @@ def retail_sales_dataset(): def _deploy_kwargs(test_helper, framework_version: str) -> dict: return { "framework_version": framework_version, - "image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + "custom_image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), } @@ -68,7 +68,7 @@ def _predict_kwargs(test_helper, framework_version: str, ds: dict) -> dict: "static_features": ds["static_features"], "known_covariates": ds["known_covariates"], "framework_version": framework_version, - "image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), + "custom_image_uri": test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False), } @@ -102,7 +102,7 @@ def test_timeseries_train(test_helper, framework_version, shared_training_job_na timestamp_column=ds["timestamp_column"], static_features=ds["static_features"], framework_version=framework_version, - image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), + custom_image_uri=test_helper.get_custom_image_uri(framework_version, type="training", gpu=False), job_name=shared_training_job_name, ) info = predictor.info() @@ -255,7 +255,7 @@ def test_timeseries_fit_predict_chronos( id_column=ds["id_column"], timestamp_column=ds["timestamp_column"], framework_version=framework_version, - image_uri=training_custom_image_uri, + custom_image_uri=training_custom_image_uri, predictions_path=predictions_path, ) @@ -340,7 +340,7 @@ def test_foundation_model_cache_artifact_then_deploy_serverless(test_helper, fra assert cached_model.model_artifact_uri.startswith("s3://") endpoint = cached_model.deploy( - image_uri=inference_custom_image_uri, + custom_image_uri=inference_custom_image_uri, inference_mode="serverless", inference_config={"memory_size_in_mb": 6144}, ) @@ -466,9 +466,9 @@ def test_timeseries_endpoint_payload_formats(test_helper, framework_version, pla predictor_init_args=dict(target="target", prediction_length=_PLAIN_PREDICTION_LENGTH), predictor_fit_args=dict(presets="medium_quality", time_limit=60), framework_version=framework_version, - image_uri=training_custom_image_uri, + custom_image_uri=training_custom_image_uri, ) - cloud_predictor.deploy(framework_version=framework_version, image_uri=inference_custom_image_uri) + cloud_predictor.deploy(framework_version=framework_version, custom_image_uri=inference_custom_image_uri) try: format_pairs = list( itertools.product( @@ -512,7 +512,7 @@ def test_foundation_model_deploy(test_helper, framework_version, retail_sales_da "chronos-bolt-tiny", cloud_output_path=f"s3://autogluon-cloud-ci/test-fm-deploy-{device}/{framework_version}/{timestamp}", ) - endpoint = model.deploy(image_uri=inference_custom_image_uri, **deploy_kwargs) + endpoint = model.deploy(custom_image_uri=inference_custom_image_uri, **deploy_kwargs) try: endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=endpoint.endpoint_name)[ "EndpointArn" From 2c60d6af557800313839be108e2cf95a8d374014 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 12:11:06 +0000 Subject: [PATCH 09/16] Align predict() result handling with fit_predict() - Replace predict()/predict_proba()'s download / persist / save_path with predictions_path, the S3 prefix the batch transform job writes to. With wait=True results are always loaded and returned, as in fit_predict(). - Inline resolve_image_uri: use custom_image_uri or retrieve_image_uri(). --- .../cloud/backend/sagemaker_backend.py | 126 ++++++------------ .../cloud/predictor/cloud_predictor.py | 52 ++------ .../predictor/timeseries_cloud_predictor.py | 10 +- src/autogluon/cloud/utils/ag_sagemaker.py | 21 --- src/autogluon/cloud/utils/sagemaker_api.py | 2 +- tests/unittests/general/test_sagemaker_ami.py | 88 +++++++----- 6 files changed, 118 insertions(+), 181 deletions(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 37a73682..7dfdeb74 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -17,7 +17,6 @@ from ..scripts import ScriptManager from ..utils.ag_sagemaker import ( repack_model_with_serving_code, - resolve_image_uri, script_mode_environment, staged_serving_code, training_script_hyperparameters, @@ -26,7 +25,7 @@ from ..utils.aws_utils import resolve_execution_role, setup_sagemaker_session 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 +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, @@ -334,9 +333,8 @@ def fit( "TrainingJobName": job_name, "RoleArn": self.role_arn, "AlgorithmSpecification": { - "TrainingImage": resolve_image_uri( - custom_image_uri, framework_version, py_version, self._region, "training", instance_type - ), + "TrainingImage": custom_image_uri + or retrieve_image_uri(framework_version, self._region, "training", instance_type, py_version), "TrainingInputMode": "File", }, "HyperParameters": training_script_hyperparameters( @@ -558,9 +556,8 @@ def deploy( model_name = self._create_model( model_name=unique_name_from_base(endpoint_name), model_data=model_data, - image_uri=resolve_image_uri( - custom_image_uri, framework_version, py_version, self._region, "inference", instance_type - ), + image_uri=custom_image_uri + or retrieve_image_uri(framework_version, self._region, "inference", instance_type, py_version), entry_point=entry_point, environment=container_environment, tags=tags, @@ -784,9 +781,7 @@ 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, + predictions_path: Optional[str] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ @@ -822,17 +817,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. - 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. + 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"``. @@ -840,8 +827,8 @@ def predict( 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, @@ -853,9 +840,7 @@ def predict( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, + predictions_path=predictions_path, backend_overrides=backend_overrides, original_features=self.original_features, ) @@ -874,9 +859,7 @@ 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, + 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]]]: """ @@ -915,17 +898,9 @@ 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. + 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"``. @@ -934,8 +909,8 @@ def predict_proba( 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. """ @@ -949,9 +924,7 @@ def predict_proba( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, + predictions_path=predictions_path, backend_overrides=backend_overrides, original_features=self.original_features, ) @@ -1219,9 +1192,7 @@ def _predict( instance_count=1, custom_image_uri=None, wait=True, - download=True, - persist=True, - save_path=None, + predictions_path=None, backend_overrides=None, split_pred_proba=True, original_features=None, @@ -1233,6 +1204,8 @@ def _predict( ): _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." @@ -1302,31 +1275,12 @@ def _predict( repacked_model_uri=f"{output_path}/model/model.tar.gz", ) - 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 - tags = self._resolve_tags() model_name = self._create_model( model_name=job_name, model_data=model_data, - image_uri=resolve_image_uri( - custom_image_uri, framework_version, py_version, self._region, "inference", instance_type - ), + image_uri=custom_image_uri + or retrieve_image_uri(framework_version, self._region, "inference", instance_type, py_version), entry_point=entry_point, environment={}, tags=tags, @@ -1339,7 +1293,10 @@ def _predict( } if split_type is not None: transform_input["SplitType"] = split_type - transform_output: Dict[str, Any] = {"S3OutputPath": output_path + "/results", "Accept": accept} + 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} @@ -1366,22 +1323,23 @@ def _predict( self._batch_transform_jobs[job_name] = batch_transform_job pred, pred_proba = None, None - if download: - results_path = self.download_predict_results(save_path=save_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}") + 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 diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index f7390816..08e54562 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -553,9 +553,7 @@ 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, + predictions_path: Optional[str] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.Series]: """ @@ -589,17 +587,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. - 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. + 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: @@ -610,8 +600,8 @@ def predict( 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 """ return self.backend.predict( test_data=test_data, @@ -623,9 +613,7 @@ def predict( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, + predictions_path=predictions_path, backend_overrides=backend_overrides, ) @@ -642,9 +630,7 @@ 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, + 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]]]: """ @@ -681,17 +667,9 @@ 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. + 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: @@ -702,8 +680,8 @@ def predict_proba( 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. """ @@ -718,9 +696,7 @@ def predict_proba( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, + predictions_path=predictions_path, backend_overrides=backend_overrides, ) diff --git a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py index bc39ac9d..012f0a50 100644 --- a/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py +++ b/src/autogluon/cloud/predictor/timeseries_cloud_predictor.py @@ -220,9 +220,7 @@ 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, + predictions_path: Optional[str] = None, backend_overrides: Optional[Dict[str, Dict[str, Any]]] = None, ) -> Optional[pd.DataFrame]: """ @@ -262,7 +260,7 @@ 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, persist, save_path, backend_overrides: + predictions_path, backend_overrides: Same as in :meth:`TabularCloudPredictor.predict`. """ return self.backend.predict( @@ -276,9 +274,7 @@ def predict( instance_count=instance_count, custom_image_uri=custom_image_uri, wait=wait, - download=download, - persist=persist, - save_path=save_path, + predictions_path=predictions_path, backend_overrides=backend_overrides, ) diff --git a/src/autogluon/cloud/utils/ag_sagemaker.py b/src/autogluon/cloud/utils/ag_sagemaker.py index c1e8aeba..619a8384 100644 --- a/src/autogluon/cloud/utils/ag_sagemaker.py +++ b/src/autogluon/cloud/utils/ag_sagemaker.py @@ -16,32 +16,11 @@ from autogluon.common.utils.s3_utils import s3_path_to_bucket_prefix -from .dlc_utils import retrieve_image_uri from .utils import safe_unpack_archive SOURCE_DIR_TARBALL_NAME = "sourcedir.tar.gz" -def resolve_image_uri( - image_uri: Optional[str], - framework_version: Optional[str], - py_version: Optional[str], - region: str, - image_scope: str, - instance_type: str, -) -> str: - """Return ``image_uri`` if set, otherwise the official AutoGluon DLC for the given version and instance.""" - if image_uri: - return image_uri - return retrieve_image_uri( - framework_version=framework_version, - region=region, - image_scope=image_scope, - instance_type=instance_type, - py_version=py_version, - ) - - 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. diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index ca3aedf9..e3c7e4c1 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -16,7 +16,7 @@ BATCH_PREDICT_OVERRIDE_KEYS = ("create_model", "create_transform_job") _REMOVED_KWARGS = { - "backend_kwargs": "`backend_overrides` (and the `download` / `persist` / `save_path` arguments of `predict()`)", + "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': ...}`", diff --git a/tests/unittests/general/test_sagemaker_ami.py b/tests/unittests/general/test_sagemaker_ami.py index b6a89717..e8c8b056 100644 --- a/tests/unittests/general/test_sagemaker_ami.py +++ b/tests/unittests/general/test_sagemaker_ami.py @@ -48,6 +48,42 @@ def test_infer_realtime_ami_ignores_unsupported_or_already_compatible_instance_f 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( ("backend_overrides", "expected"), [ @@ -56,36 +92,28 @@ def test_infer_realtime_ami_ignores_unsupported_or_already_compatible_instance_f ], ) def test_batch_transform_job_sets_inferred_ami_without_overriding_user_value( - backend_overrides, expected, assert_valid_request + transform_request, backend_overrides, expected ): - 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", - instance_type="ml.g4dn.xlarge", - custom_image_uri=GPU_IMAGE_URI, - wait=False, - download=False, - persist=False, - backend_overrides=backend_overrides, - ) - - request = job_cls.return_value.run.call_args.kwargs["transform_job_request"] - assert_valid_request("CreateTransformJob", request) + 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") From 1d51cdde68f2f906e202ab7158b144c7bcee261c Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 12:11:41 +0000 Subject: [PATCH 10/16] Let retrieve_image_uri pass through a custom image URI --- src/autogluon/cloud/backend/sagemaker_backend.py | 15 +++++++++------ src/autogluon/cloud/utils/dlc_utils.py | 6 ++++-- 2 files changed, 13 insertions(+), 8 deletions(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 7dfdeb74..e290000d 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -333,8 +333,9 @@ def fit( "TrainingJobName": job_name, "RoleArn": self.role_arn, "AlgorithmSpecification": { - "TrainingImage": custom_image_uri - or retrieve_image_uri(framework_version, self._region, "training", instance_type, py_version), + "TrainingImage": retrieve_image_uri( + framework_version, self._region, "training", instance_type, py_version, custom_image_uri + ), "TrainingInputMode": "File", }, "HyperParameters": training_script_hyperparameters( @@ -556,8 +557,9 @@ def deploy( model_name = self._create_model( model_name=unique_name_from_base(endpoint_name), model_data=model_data, - image_uri=custom_image_uri - or retrieve_image_uri(framework_version, self._region, "inference", instance_type, py_version), + image_uri=retrieve_image_uri( + framework_version, self._region, "inference", instance_type, py_version, custom_image_uri + ), entry_point=entry_point, environment=container_environment, tags=tags, @@ -1279,8 +1281,9 @@ def _predict( model_name = self._create_model( model_name=job_name, model_data=model_data, - image_uri=custom_image_uri - or retrieve_image_uri(framework_version, self._region, "inference", instance_type, py_version), + image_uri=retrieve_image_uri( + framework_version, self._region, "inference", instance_type, py_version, custom_image_uri + ), entry_point=entry_point, environment={}, tags=tags, 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] From 4a62b4b9ba3aa211b2acc7007dca7f56ed9111c7 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 12:36:48 +0000 Subject: [PATCH 11/16] Fix resource tracking with backend_overrides and restore SDK behaviors - Clean up the model and endpoint config when deploy() fails partway, and give each deploy a unique endpoint config name so a leftover config can't block a redeploy. - Batch transform: delete only the model this call created, and register the job under its actual (possibly overridden) name. - Pass training code as a `code` input channel, as SDK v2 did, so training works with EnableNetworkIsolation. - attach_job() raises again when the training job did not complete. --- .../cloud/backend/sagemaker_backend.py | 28 ++++++++++++++----- src/autogluon/cloud/job/sagemaker_job.py | 6 ++-- .../cloud/predictor/cloud_predictor.py | 9 ++++-- .../unittests/general/test_inference_modes.py | 12 ++++++++ tests/unittests/general/test_sagemaker_api.py | 8 +++--- 5 files changed, 46 insertions(+), 17 deletions(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index e290000d..bc4c0975 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -16,6 +16,7 @@ from ..job import SageMakerBatchTransformationJob, SageMakerFitJob from ..scripts import ScriptManager from ..utils.ag_sagemaker import ( + SOURCE_DIR_TARBALL_NAME, repack_model_with_serving_code, script_mode_environment, staged_serving_code, @@ -338,10 +339,15 @@ def fit( ), "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=code_uri, job_name=job_name, region=self._region + 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.items()], + "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, @@ -584,7 +590,7 @@ def deploy( variant = deep_merge(variant, overrides.get("production_variant", {})) endpoint_config_request: Dict[str, Any] = { - "EndpointConfigName": endpoint_name, + "EndpointConfigName": unique_name_from_base(endpoint_name), "ProductionVariants": [variant], "Tags": tags, } @@ -600,8 +606,16 @@ def deploy( logger.log(20, f"Deploying model to the endpoint (inference_mode={inference_mode})") client = self.sagemaker_session.sagemaker_client - client.create_endpoint_config(**endpoint_config_request) - client.create_endpoint(**endpoint_request) + try: + client.create_endpoint_config(**endpoint_config_request) + try: + client.create_endpoint(**endpoint_request) + except Exception: + client.delete_endpoint_config(EndpointConfigName=endpoint_config_request["EndpointConfigName"]) + raise + except Exception: + 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) @@ -1322,8 +1336,8 @@ def _predict( request = deep_merge(request, overrides.get("create_transform_job", {})) batch_transform_job = SageMakerBatchTransformationJob(session=self.sagemaker_session) - batch_transform_job.run(transform_job_request=request, wait=wait) - 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 wait: diff --git a/src/autogluon/cloud/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index 198afd68..d1b71bde 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -161,7 +161,7 @@ def attach(cls, job_name, session=None): # FIXME: find a way to recover framework version obj = cls(session=session) obj._job_name = job_name - obj.wait(logs=True) + obj._wait_until_completed() return obj @property @@ -247,15 +247,15 @@ def _delete_model(self, model_name: str) -> None: def run( self, transform_job_request: Dict[str, Any], + model_name: str, wait: bool, ): """Create the transform job from a ``CreateTransformJob`` request. - The SageMaker model referenced by the request is deleted once the job finishes (``wait=True``) or fails to + ``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"] - model_name = transform_job_request["ModelName"] try: logger.log(20, "Transforming") self.session.sagemaker_client.create_transform_job(**transform_job_request) diff --git a/src/autogluon/cloud/predictor/cloud_predictor.py b/src/autogluon/cloud/predictor/cloud_predictor.py index 08e54562..eaaee208 100644 --- a/src/autogluon/cloud/predictor/cloud_predictor.py +++ b/src/autogluon/cloud/predictor/cloud_predictor.py @@ -426,7 +426,8 @@ def deploy( ``"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. + 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'`.") @@ -595,7 +596,8 @@ def predict( 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. + Nested dicts merge recursively; other values, including lists, replace the generated ones. Only + resources created by AutoGluon-Cloud are cleaned up. Returns ------- @@ -675,7 +677,8 @@ def predict_proba( 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. + Nested dicts merge recursively; other values, including lists, replace the generated ones. Only + resources created by AutoGluon-Cloud are cleaned up. Returns ------- diff --git a/tests/unittests/general/test_inference_modes.py b/tests/unittests/general/test_inference_modes.py index 66d5538e..0337f026 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -129,3 +129,15 @@ def test_when_container_environment_overridden_then_it_merges_with_defaults(depl 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") + config_name = client.create_endpoint_config.call_args.kwargs["EndpointConfigName"] + assert config_name.startswith("ep-") # unique per deploy, so a leftover config can't block a redeploy + client.delete_endpoint_config.assert_called_once_with(EndpointConfigName=config_name) + 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_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index c93ac817..b1a30756 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -105,10 +105,10 @@ def test_fit_builds_script_mode_training_job(fit_request): 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"] == ( - '"s3://bucket/run/code/job/source/sourcedir.tar.gz"' - ) - assert request["InputDataConfig"][0]["ChannelName"] == "train_data" + 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"] From dbfa59cf69c2b05f79900b4491196a36b1ea6c4d Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 12:40:00 +0000 Subject: [PATCH 12/16] Reject custom entry point scripts with an explicit error --- src/autogluon/cloud/utils/sagemaker_api.py | 18 +++++++++++++++--- tests/unittests/general/test_sagemaker_api.py | 2 ++ 2 files changed, 17 insertions(+), 3 deletions(-) diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index e3c7e4c1..2f22fa20 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -31,16 +31,28 @@ def reject_legacy_kwargs(func): @functools.wraps(func) def wrapper(*args, **kwargs): - for name in kwargs: - if name in _REMOVED_KWARGS: + for name, value in kwargs.items(): + if name not in _REMOVED_KWARGS: + continue + if _sets_custom_entry_point(value): raise TypeError( - f"`{name}` was removed from {func.__qualname__}(). Use {_REMOVED_KWARGS[name]} instead." + 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 {}) diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index b1a30756..5f27a50d 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -36,6 +36,8 @@ def fit(**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(): From 4382ea7023a739d8c9024b1283db6d7ff7837c3d Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 12:42:04 +0000 Subject: [PATCH 13/16] Reject overrides of resource references; harden deploy rollback - backend_overrides can't set fields that link created resources (ProductionVariants / their ModelName, EndpointConfigName on the endpoint, ModelName on the transform job), so cleanup only deletes our own resources. - Validate inference_mode before creating anything. - Rollback deletes log failures instead of masking the original error. --- .../cloud/backend/sagemaker_backend.py | 13 ++++++----- src/autogluon/cloud/job/sagemaker_job.py | 3 ++- src/autogluon/cloud/utils/sagemaker_api.py | 22 ++++++++++++++++++- tests/unittests/general/test_sagemaker_api.py | 5 +++++ 4 files changed, 36 insertions(+), 7 deletions(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index bc4c0975..ff5b5621 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -35,6 +35,7 @@ check_override_keys, deep_merge, delete_endpoint, + delete_quietly, invoke_endpoint, ) from ..utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer @@ -487,6 +488,8 @@ def deploy( 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": @@ -583,10 +586,8 @@ def deploy( ) if inference_ami_version is not None: variant["InferenceAmiVersion"] = inference_ami_version - elif inference_mode == "serverless": - variant["ServerlessConfig"] = serverless_config 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] = { @@ -611,10 +612,12 @@ def deploy( try: client.create_endpoint(**endpoint_request) except Exception: - client.delete_endpoint_config(EndpointConfigName=endpoint_config_request["EndpointConfigName"]) + delete_quietly( + client.delete_endpoint_config, EndpointConfigName=endpoint_config_request["EndpointConfigName"] + ) raise except Exception: - client.delete_model(ModelName=model_name) + delete_quietly(client.delete_model, ModelName=model_name) raise self.endpoint_name = endpoint_request["EndpointName"] if wait: diff --git a/src/autogluon/cloud/job/sagemaker_job.py b/src/autogluon/cloud/job/sagemaker_job.py index d1b71bde..97deaad3 100644 --- a/src/autogluon/cloud/job/sagemaker_job.py +++ b/src/autogluon/cloud/job/sagemaker_job.py @@ -5,6 +5,7 @@ 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__) @@ -264,7 +265,7 @@ def run( self._wait_until_completed() logger.log(20, "Transform done") except Exception as e: - self._delete_model(model_name) + delete_quietly(self.session.sagemaker_client.delete_model, ModelName=model_name) raise e input_uri = transform_job_request["TransformInput"]["DataSource"]["S3DataSource"]["S3Uri"] diff --git a/src/autogluon/cloud/utils/sagemaker_api.py b/src/autogluon/cloud/utils/sagemaker_api.py index 2f22fa20..24f3f142 100644 --- a/src/autogluon/cloud/utils/sagemaker_api.py +++ b/src/autogluon/cloud/utils/sagemaker_api.py @@ -3,7 +3,7 @@ import copy import functools import logging -from typing import Any, Dict, Iterable, Mapping, Optional +from typing import Any, Callable, Dict, Iterable, Mapping, Optional from .aws_utils import AwsSession @@ -14,6 +14,14 @@ 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)", @@ -59,9 +67,21 @@ def check_override_keys(overrides: Optional[Mapping[str, Any]], allowed_keys: It 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)) diff --git a/tests/unittests/general/test_sagemaker_api.py b/tests/unittests/general/test_sagemaker_api.py index 5f27a50d..1933a065 100644 --- a/tests/unittests/general/test_sagemaker_api.py +++ b/tests/unittests/general/test_sagemaker_api.py @@ -28,6 +28,11 @@ def test_check_override_keys_rejects_unknown_keys(): 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): From d56f415c5e94effacbd04afc6522ce4c6b19b048 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 12:44:01 +0000 Subject: [PATCH 14/16] Drop attach_endpoint type check --- src/autogluon/cloud/backend/sagemaker_backend.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index ff5b5621..3a64a728 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -655,8 +655,6 @@ def attach_endpoint(self, endpoint: str) -> 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 not isinstance(endpoint, str): - raise ValueError(f"Please provide the endpoint name as a string, got {type(endpoint).__name__}.") self.endpoint_name = endpoint def detach_endpoint(self) -> str: From ce4f8166ce612e44b09a71f130a482b7301707d7 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 13:04:55 +0000 Subject: [PATCH 15/16] Name the endpoint config after the endpoint, as SDK v2 did --- src/autogluon/cloud/backend/sagemaker_backend.py | 2 +- tests/unittests/general/test_inference_modes.py | 4 +--- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 3a64a728..385804e2 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -591,7 +591,7 @@ def deploy( variant = deep_merge(variant, overrides.get("production_variant", {})) endpoint_config_request: Dict[str, Any] = { - "EndpointConfigName": unique_name_from_base(endpoint_name), + "EndpointConfigName": endpoint_name, "ProductionVariants": [variant], "Tags": tags, } diff --git a/tests/unittests/general/test_inference_modes.py b/tests/unittests/general/test_inference_modes.py index 0337f026..a6a351e4 100644 --- a/tests/unittests/general/test_inference_modes.py +++ b/tests/unittests/general/test_inference_modes.py @@ -136,8 +136,6 @@ def test_when_endpoint_creation_fails_then_model_and_config_are_deleted(deploy_r client.create_endpoint.side_effect = RuntimeError("boom") with pytest.raises(RuntimeError, match="boom"): deploy_requests(instance_type="ml.m5.xlarge") - config_name = client.create_endpoint_config.call_args.kwargs["EndpointConfigName"] - assert config_name.startswith("ep-") # unique per deploy, so a leftover config can't block a redeploy - client.delete_endpoint_config.assert_called_once_with(EndpointConfigName=config_name) + 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 From 025258cb9ec4ef047f498a0cd7c5e6d1fda8b658 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 2 Oct 2026 13:06:08 +0000 Subject: [PATCH 16/16] Name the endpoint's model after the endpoint, like batch predict does --- src/autogluon/cloud/backend/sagemaker_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 385804e2..c2d1902c 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -564,7 +564,7 @@ def deploy( tags = self._resolve_tags(extra_tags) model_name = self._create_model( - model_name=unique_name_from_base(endpoint_name), + 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