Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions docs/api.rst
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,20 @@ API

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
Expand Down
2 changes: 2 additions & 0 deletions docs/api/tabular.rst
Original file line number Diff line number Diff line change
Expand Up @@ -8,3 +8,5 @@ Tabular
:template: custom_class.rst

TabularCloudPredictor
TabularFoundationModel
TabularEndpoint
2 changes: 2 additions & 0 deletions src/autogluon/cloud/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from autogluon.common.utils.log_utils import _add_stream_handler

from .cloud_setup import bootstrap, register, status, teardown
from .endpoint.tabular_endpoint import TabularEndpoint
from .endpoint.timeseries_endpoint import TimeSeriesEndpoint
from .model.foundation_model import TabularFoundationModel, TimeSeriesFoundationModel
from .predictor import MultiModalCloudPredictor, TabularCloudPredictor, TimeSeriesCloudPredictor
Expand All @@ -22,6 +23,7 @@
__all__ = [
"MultiModalCloudPredictor",
"TabularCloudPredictor",
"TabularEndpoint",
"TabularFoundationModel",
"TimeSeriesCloudPredictor",
"TimeSeriesEndpoint",
Expand Down
2 changes: 1 addition & 1 deletion src/autogluon/cloud/backend/sagemaker_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -1410,7 +1410,7 @@ def __setstate__(self, 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.endpoiont = self.attach_endpoint(self._endpoint_saved)
self.attach_endpoint(self._endpoint_saved)
self._endpoint_saved = None
self._fit_job.session = self.sagemaker_session
for job in self._batch_transform_jobs:
Expand Down
128 changes: 128 additions & 0 deletions src/autogluon/cloud/endpoint/tabular_endpoint.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,128 @@
from pathlib import Path
from typing import Any, Dict, Optional, Tuple, Union

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.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer
from ..utils.utils import split_pred_and_pred_proba

DataInput = Union[str, Path, pd.DataFrame]
Prediction = Union[pd.DataFrame, pd.Series]


class TabularEndpoint:
"""High-level handle for an AutoGluon-Cloud tabular foundation-model endpoint."""

def __init__(self, endpoint_name: str, session: Optional[boto3.Session] = None):
"""
Parameters
----------
endpoint_name
Name of an existing SageMaker endpoint deployed through
:meth:`autogluon.cloud.TabularFoundationModel.deploy`.
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(),
)

@property
def endpoint_name(self) -> str:
return self._predictor.endpoint_name

@staticmethod
def _load_data(data: DataInput) -> pd.DataFrame:
if isinstance(data, (str, Path)):
return load_pd.load(str(data))
return data

def _predict(
self,
data: DataInput,
train_data: DataInput,
label: str,
inference_kwargs: Optional[Dict[str, Any]] = None,
) -> Tuple[pd.Series, Prediction]:
data = self._load_data(data)
train_data = self._load_data(train_data)

if label not in train_data.columns:
raise ValueError(f"Label column {label!r} is not present in `train_data`.")
feature_columns = [column for column in train_data.columns if column != label]
missing_columns = [column for column in feature_columns if column not in data.columns]
if missing_columns:
raise ValueError(f"`data` is missing feature columns present in `train_data`: {missing_columns}.")

payload = AutoGluonSerializationWrapper(
data=data,
train_data=train_data,
inference_kwargs={"label": label, **(inference_kwargs or {})},
)
raw = self._predictor.predict(payload, initial_args={"Accept": "application/x-parquet"})
pred, pred_proba = split_pred_and_pred_proba(raw)
if pred_proba is None:
pred_proba = pred
return pred, pred_proba

def predict(
self,
data: DataInput,
train_data: DataInput,
label: str,
**inference_kwargs: Any,
) -> pd.Series:
"""Fit the foundation model on ``train_data`` and predict ``data``.

The serialized request includes both ``train_data`` and ``data`` and must not exceed SageMaker's
6 MiB real-time invocation payload limit. Use
:meth:`autogluon.cloud.TabularFoundationModel.predict` for larger inputs.
"""
pred, _ = self._predict(
data=data,
train_data=train_data,
label=label,
inference_kwargs=inference_kwargs,
)
return pred

def predict_proba(
self,
data: DataInput,
train_data: DataInput,
label: str,
*,
include_predict: bool = True,
**inference_kwargs: Any,
) -> Union[Tuple[pd.Series, Prediction], Prediction]:
"""Fit the foundation model and return class probabilities.

For regression, the probability result is identical to the prediction.

The serialized request includes both ``train_data`` and ``data`` and must not exceed SageMaker's
6 MiB real-time invocation payload limit. Use
:meth:`autogluon.cloud.TabularFoundationModel.predict_proba` for larger inputs.
"""
pred, pred_proba = self._predict(
data=data,
train_data=train_data,
label=label,
inference_kwargs=inference_kwargs,
)
if include_predict:
return pred, pred_proba
return pred_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)
64 changes: 57 additions & 7 deletions src/autogluon/cloud/model/foundation_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,13 @@
import pandas as pd
from typing_extensions import Self

from autogluon.common.loaders import load_pd
from autogluon.common.utils.s3_utils import s3_path_to_bucket_prefix

from ..backend.backend_factory import BackendFactory
from ..backend.constant import SAGEMAKER, TABULAR_SAGEMAKER, TIMESERIES_SAGEMAKER
from ..endpoint.prediction_future import JobPredictionFuture
from ..endpoint.tabular_endpoint import TabularEndpoint
from ..endpoint.timeseries_endpoint import TimeSeriesEndpoint
from ..scripts.script_manager import ScriptManager
from ..utils.aws_utils import resolve_cloud_output_path
Expand Down Expand Up @@ -213,6 +215,7 @@ def _deploy_backend(
fm_serve_config = {
"ag_model_key": self._config.ag_model_key,
"hyperparameters": merged_hp,
"problem_type": self._config.problem_type,
}

model_kwargs = backend_kwargs.pop("model_kwargs", {})
Expand Down Expand Up @@ -608,20 +611,62 @@ class TabularFoundationModel(FoundationModel):
runs prediction as a managed SageMaker job, with no training required. Each ``model_id`` targets a
single task — ``mitra-classifier`` for classification and ``mitra-regressor`` for regression.

Predictions are produced in batch mode: :meth:`predict` (and :meth:`predict_proba`) runs a one-off
SageMaker training job where the labeled ``train_data`` provides the in-context examples and the
predictions for ``test_data`` are written to S3.
Predictions can be produced in batch mode with :meth:`predict` / :meth:`predict_proba`, or through a
real-time endpoint created with :meth:`deploy`. In both modes, labeled ``train_data`` provides the
in-context examples for each prediction.
"""

_backend_map = {SAGEMAKER: TABULAR_SAGEMAKER}
_predictor_type = "tabular"

@property
def _serve_script_path(self) -> str:
raise NotImplementedError("Tabular FM deploy is not yet supported")
return ScriptManager.SAGEMAKER_TABULAR_FM_SERVE_SCRIPT_PATH

def deploy(self, **kwargs):
raise NotImplementedError("Tabular FM deploy is not yet supported")
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,
wait: bool = True,
inference_mode: Literal["realtime"] = "realtime",
inference_config: Optional[Dict[str, Any]] = None,
**backend_kwargs,
) -> TabularEndpoint:
"""Deploy the tabular foundation model to an inference endpoint.

The returned endpoint accepts both labeled ``train_data`` and the rows to predict. It fits a
request-scoped :class:`TabularPredictor` before producing predictions.

Only real-time inference is supported. Tabular foundation models such as Mitra require a
provisioned instance and cannot be deployed with SageMaker Serverless Inference.
"""
if inference_mode != "realtime":
raise ValueError(
"TabularFoundationModel.deploy only supports `inference_mode='realtime'`; "
"SageMaker Serverless Inference does not provide sufficient resources for tabular foundation models."
)
if inference_config is not None:
raise ValueError(
"`inference_config` is not supported by TabularFoundationModel.deploy because tabular foundation "
"models do not support SageMaker Serverless Inference."
)
self._deploy_backend(
instance_type=instance_type,
endpoint_name=endpoint_name,
hyperparameters=hyperparameters,
framework_version=framework_version,
custom_image_uri=custom_image_uri,
wait=wait,
inference_mode="realtime",
**backend_kwargs,
)
return TabularEndpoint(
endpoint_name=self._backend.endpoint.endpoint_name,
session=self._backend.sagemaker_session.boto_session,
)

def _build_predictor_init_args(self, label: str = "target", **kwargs) -> Dict[str, Any]:
"""Map user kwargs to TabularPredictor init args."""
Expand Down Expand Up @@ -778,14 +823,19 @@ def predict_proba(
if instance_type is None:
instance_type = self._config.predict_instance_type

if isinstance(train_data, (str, Path)):
train_data = load_pd.load(str(train_data))
# Duplicate one tuning row so AutoGluon/Mitra do not hold out any rows from the prediction context.
tuning_data = train_data.iloc[[0]].copy()

extra_ag_args: Dict[str, Any] = {"predict_after_fit": True}
if predictions_path is not None:
extra_ag_args["predictions_path"] = predictions_path

self._backend.fit(
predictor_init_args=self._build_predictor_init_args(label=label),
predictor_fit_args=self._build_predictor_fit_args(hyperparameters),
data_channels={"train_data": train_data, "test_data": test_data},
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,
Expand Down
2 changes: 2 additions & 0 deletions src/autogluon/cloud/model/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@ class FoundationModelConfig:
model_source_hyperparameter="hf_cls_model",
inference_hyperparameters={"fine_tune": False},
predict_instance_type="ml.m5.4xlarge",
deploy_instance_type="ml.m5.4xlarge",
),
"mitra-regressor": FoundationModelConfig(
problem_type="regression",
Expand All @@ -68,6 +69,7 @@ class FoundationModelConfig:
model_source_hyperparameter="hf_reg_model",
inference_hyperparameters={"fine_tune": False},
predict_instance_type="ml.m5.4xlarge",
deploy_instance_type="ml.m5.4xlarge",
),
}

Expand Down
Loading
Loading