diff --git a/docs/api.rst b/docs/api.rst index 5a1c9f05..f1d4483c 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -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 diff --git a/docs/api/tabular.rst b/docs/api/tabular.rst index 368b8f84..a6b2e1fb 100644 --- a/docs/api/tabular.rst +++ b/docs/api/tabular.rst @@ -8,3 +8,5 @@ Tabular :template: custom_class.rst TabularCloudPredictor + TabularFoundationModel + TabularEndpoint diff --git a/src/autogluon/cloud/__init__.py b/src/autogluon/cloud/__init__.py index 187838d5..343319ca 100644 --- a/src/autogluon/cloud/__init__.py +++ b/src/autogluon/cloud/__init__.py @@ -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 @@ -22,6 +23,7 @@ __all__ = [ "MultiModalCloudPredictor", "TabularCloudPredictor", + "TabularEndpoint", "TabularFoundationModel", "TimeSeriesCloudPredictor", "TimeSeriesEndpoint", diff --git a/src/autogluon/cloud/backend/sagemaker_backend.py b/src/autogluon/cloud/backend/sagemaker_backend.py index 05ee84b0..3c69ebf1 100644 --- a/src/autogluon/cloud/backend/sagemaker_backend.py +++ b/src/autogluon/cloud/backend/sagemaker_backend.py @@ -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: diff --git a/src/autogluon/cloud/endpoint/tabular_endpoint.py b/src/autogluon/cloud/endpoint/tabular_endpoint.py new file mode 100644 index 00000000..719fee7b --- /dev/null +++ b/src/autogluon/cloud/endpoint/tabular_endpoint.py @@ -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) diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 62a301f4..2534fa23 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -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 @@ -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", {}) @@ -608,9 +611,9 @@ 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} @@ -618,10 +621,52 @@ class TabularFoundationModel(FoundationModel): @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.""" @@ -778,6 +823,11 @@ 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 @@ -785,7 +835,7 @@ def predict_proba( 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, diff --git a/src/autogluon/cloud/model/registry.py b/src/autogluon/cloud/model/registry.py index a3b258b5..bcd3c16e 100644 --- a/src/autogluon/cloud/model/registry.py +++ b/src/autogluon/cloud/model/registry.py @@ -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", @@ -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", ), } diff --git a/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py new file mode 100644 index 00000000..d64af70d --- /dev/null +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py @@ -0,0 +1,117 @@ +"""Serve script for tabular foundation models (Mitra, etc.) on SageMaker endpoints. + +Each request contains labeled ``train_data`` used as in-context examples and ``data`` to score. +Configuration comes from the ``AG_FM_SERVE_CONFIG`` environment variable set during deployment. +""" + +import base64 +import copy +import json +import os +import tempfile +from io import BytesIO + +import pandas as pd +from huggingface_hub import snapshot_download + +from autogluon.tabular import TabularPredictor + +_FM_SERVE_CONFIG = json.loads(os.environ.get("AG_FM_SERVE_CONFIG", "{}")) +_SUPPORTED_INPUT_CONTENT_TYPES = {"application/x-autogluon"} + + +def model_fn(model_dir): + """Download remote weights during container startup and return model configuration.""" + model_config = copy.deepcopy(_FM_SERVE_CONFIG) + hyperparameters = model_config.get("hyperparameters", {}) + for source_key in ("hf_cls_model", "hf_reg_model", "hf_general_model", "hf_model"): + source = hyperparameters.get(source_key) + if source is not None and not os.path.isdir(source): + snapshot_download( + repo_id=source, + allow_patterns=["config.json", "model.safetensors"], + ) + break + return model_config + + +def _read_parquet(payload, key): + encoded = payload.get(key) + if encoded is None: + raise ValueError(f"Missing required field {key!r} in x-autogluon payload.") + return pd.read_parquet(BytesIO(base64.b64decode(encoded))) + + +def _parse_payload(request_body, input_content_type): + if input_content_type not in _SUPPORTED_INPUT_CONTENT_TYPES: + raise ValueError( + f"{input_content_type} input content type not supported. " + f"Supported: {sorted(_SUPPORTED_INPUT_CONTENT_TYPES)}" + ) + + payload = json.loads(request_body) + if payload.get("version") != 1: + raise ValueError(f"Unsupported x-autogluon payload version: {payload.get('version')}. Expected 1.") + + data = _read_parquet(payload, "data") + train_data = _read_parquet(payload, "train_data") + inference_kwargs = payload.get("inference_kwargs") or {} + return data, train_data, inference_kwargs + + +def _render_response(prediction, output_content_type): + if isinstance(prediction, pd.Series): + prediction = prediction.to_frame() + + output_content_type = output_content_type.lower() + if "application/x-parquet" in output_content_type: + prediction.columns = prediction.columns.astype(str) + return prediction.to_parquet(index=False), "application/x-parquet" + if "application/json" in output_content_type: + return prediction.to_json(orient="records"), "application/json" + if "text/csv" in output_content_type: + return prediction.to_csv(index=False), "text/csv" + raise ValueError(f"{output_content_type} content type not supported") + + +def transform_fn(model_config, request_body, input_content_type, output_content_type="application/json"): + """Fit a request-scoped TabularPredictor and score the request's prediction data.""" + data, train_data, inference_kwargs = _parse_payload(request_body, input_content_type) + inference_kwargs = dict(inference_kwargs) + label = inference_kwargs.pop("label", None) + prediction_only_kwargs = { + key: inference_kwargs.pop(key) for key in ("decision_threshold",) if key in inference_kwargs + } + if label is None: + raise ValueError("`inference_kwargs` must contain the training label column name under `label`.") + if label not in train_data.columns: + raise ValueError(f"Label column {label!r} is not present in `train_data`.") + + ag_model_key = model_config["ag_model_key"] + hyperparameters = model_config.get("hyperparameters", {}) + problem_type = model_config["problem_type"] + + with tempfile.TemporaryDirectory(prefix="ag_tabular_fm_") as temp_dir: + # Duplicate one tuning row so AutoGluon/Mitra do not hold out any rows from the prediction context. + tuning_data = train_data.iloc[[0]].copy() + predictor = TabularPredictor( + label=label, + problem_type=problem_type, + path=os.path.join(temp_dir, "predictor"), + ).fit( + train_data, + tuning_data=tuning_data, + hyperparameters={ag_model_key: hyperparameters}, + fit_weighted_ensemble=False, + ) + + if predictor.can_predict_proba: + pred_proba = predictor.predict_proba(data, as_pandas=True, **inference_kwargs) + pred = predictor.predict_from_proba(pred_proba, **prediction_only_kwargs) + pred_proba.columns = [f"{column}_proba" for column in pred_proba.columns] + pred.name = predictor.label + prediction = pd.concat([pred, pred_proba], axis=1) + else: + prediction = predictor.predict(data, as_pandas=True, **inference_kwargs, **prediction_only_kwargs) + + return _render_response(prediction, output_content_type) diff --git a/src/autogluon/cloud/scripts/script_manager.py b/src/autogluon/cloud/scripts/script_manager.py index 9dc93b31..7b8e2be1 100644 --- a/src/autogluon/cloud/scripts/script_manager.py +++ b/src/autogluon/cloud/scripts/script_manager.py @@ -11,6 +11,7 @@ class ScriptManager: RAY_SCRIPTS_PATH = os.path.join(SCRIPTS_PATH, "ray_scripts") SAGEMAKER_TRAIN_SCRIPT_PATH = os.path.join(SAGEMAKER_SCRIPTS_PATH, "train.py") SAGEMAKER_TABULAR_SERVE_SCRIPT_PATH = os.path.join(SAGEMAKER_SCRIPTS_PATH, "tabular_serve.py") + SAGEMAKER_TABULAR_FM_SERVE_SCRIPT_PATH = os.path.join(SAGEMAKER_SCRIPTS_PATH, "tabular_fm_serve.py") SAGEMAKER_MULTIMODAL_SERVE_SCRIPT_PATH = os.path.join(SAGEMAKER_SCRIPTS_PATH, "multimodal_serve.py") SAGEMAKER_TIMESERIES_SERVE_SCRIPT_PATH = os.path.join(SAGEMAKER_SCRIPTS_PATH, "timeseries_serve.py") SAGEMAKER_TIMESERIES_FM_SERVE_SCRIPT_PATH = os.path.join(SAGEMAKER_SCRIPTS_PATH, "timeseries_fm_serve.py") diff --git a/src/autogluon/cloud/utils/serializers.py b/src/autogluon/cloud/utils/serializers.py index 9a8a19f2..1f6633a1 100644 --- a/src/autogluon/cloud/utils/serializers.py +++ b/src/autogluon/cloud/utils/serializers.py @@ -29,6 +29,7 @@ class AutoGluonSerializationWrapper: data: pd.DataFrame inference_kwargs: Dict[str, Any] + train_data: Optional[pd.DataFrame] = field(default=None) static_features: Optional[pd.DataFrame] = field(default=None) known_covariates: Optional[pd.DataFrame] = field(default=None) @@ -64,6 +65,8 @@ def serialize(self, data: AutoGluonSerializationWrapper): "data": _dataframe_to_b64(data.data), "inference_kwargs": inference_kwargs, } + if data.train_data is not None: + package["train_data"] = _dataframe_to_b64(data.train_data) if data.static_features is not None: package["static_features"] = _dataframe_to_b64(data.static_features) if data.known_covariates is not None: diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 10ac8de4..204a9499 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -130,6 +130,44 @@ def test_deploy_without_artifact_passes_none_predictor_path_and_source_uri(): 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(): + fm = FoundationModel("mitra-classifier", cloud_output_path="s3://b") + fm._backend.endpoint = mock.MagicMock(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: + endpoint = fm.deploy() + + 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["fm_serve_config"] == { + "ag_model_key": "MITRA", + "hyperparameters": { + "fine_tune": False, + "hf_cls_model": "autogluon/mitra-classifier", + }, + "problem_type": "multiclass", + } + endpoint_cls.assert_called_once_with( + endpoint_name="mitra-endpoint", + session=mock.sentinel.boto_session, + ) + assert endpoint is endpoint_cls.return_value + + +def test_tabular_deploy_rejects_serverless_inference(): + fm = FoundationModel("mitra-classifier", cloud_output_path="s3://b") + + with pytest.raises(ValueError, match="only supports `inference_mode='realtime'`"): + fm.deploy(inference_mode="serverless") + + with pytest.raises(ValueError, match="`inference_config` is not supported"): + fm.deploy(inference_config={"memory_size_in_mb": 6144}) + + fm._backend.deploy.assert_not_called() + + def test_deploy_rejects_user_model_path_when_artifact_uri_set(): """User-supplied model_path is incoherent with model_artifact_uri (the bundled tarball dictates the in-container path). Raise rather than silently overwrite.""" diff --git a/tests/unittests/general/test_serializers.py b/tests/unittests/general/test_serializers.py index bea41832..510eb6d9 100644 --- a/tests/unittests/general/test_serializers.py +++ b/tests/unittests/general/test_serializers.py @@ -79,12 +79,22 @@ def test_when_all_fields_provided_then_payload_contains_version_and_data(df, sta pd.testing.assert_frame_equal(known_covariates, _decode_parquet(payload["known_covariates"])) +def test_when_train_data_provided_then_payload_contains_train_data(df): + train_data = pd.DataFrame({"a": [10, 20], "b": ["train-x", "train-y"], "label": [0, 1]}) + wrapper = AutoGluonSerializationWrapper(data=df, train_data=train_data, inference_kwargs={"label": "label"}) + payload = json.loads(AutoGluonSerializer().serialize(wrapper)) + + pd.testing.assert_frame_equal(df, _decode_parquet(payload["data"])) + pd.testing.assert_frame_equal(train_data, _decode_parquet(payload["train_data"])) + + def test_when_no_optional_fields_then_payload_omits_them(df): wrapper = AutoGluonSerializationWrapper(data=df, inference_kwargs={}) payload = json.loads(AutoGluonSerializer().serialize(wrapper)) assert "static_features" not in payload assert "known_covariates" not in payload + assert "train_data" not in payload pd.testing.assert_frame_equal(df, _decode_parquet(payload["data"])) diff --git a/tests/unittests/general/test_tabular_endpoint.py b/tests/unittests/general/test_tabular_endpoint.py new file mode 100644 index 00000000..baf934ca --- /dev/null +++ b/tests/unittests/general/test_tabular_endpoint.py @@ -0,0 +1,72 @@ +from unittest import mock + +import pandas as pd +import pytest + +from autogluon.cloud.endpoint.tabular_endpoint import TabularEndpoint +from autogluon.cloud.utils.serializers import AutoGluonSerializationWrapper + + +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 + + +def test_predict_sends_train_data_and_returns_prediction_series(): + 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) + + pred = endpoint.predict(data=data, train_data=train_data, label="label") + + assert pred.tolist() == ["a", "b"] + payload = endpoint._predictor.predict.call_args.args[0] + 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(): + response = pd.DataFrame({"label": ["a"], "a_proba": [0.7], "b_proba": [0.3]}) + endpoint = _make_endpoint(response) + train_data = pd.DataFrame({"feature": [0, 1], "label": ["a", "b"]}) + data = pd.DataFrame({"feature": [2]}) + + pred, proba = endpoint.predict_proba(data=data, train_data=train_data, label="label") + + assert pred.tolist() == ["a"] + assert proba.columns.tolist() == ["a", "b"] + 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]})) + train_data = pd.DataFrame({"feature": [0, 1], "target": [0.0, 1.0]}) + data = pd.DataFrame({"feature": [2, 3]}) + + pred, proba = endpoint.predict_proba(data=data, train_data=train_data, label="target") + + pd.testing.assert_series_equal(pred, proba) + + +def test_predict_validates_label_and_feature_columns(): + endpoint = _make_endpoint(pd.DataFrame()) + train_data = pd.DataFrame({"feature": [0], "label": ["a"]}) + + with pytest.raises(ValueError, match="Label column"): + endpoint.predict(data=pd.DataFrame({"feature": [1]}), train_data=train_data, label="missing") + + with pytest.raises(ValueError, match="missing 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() + + endpoint._predictor.delete_model.assert_called_once_with() + endpoint._predictor.delete_endpoint.assert_called_once_with(delete_endpoint_config=True) diff --git a/tests/unittests/general/test_tabular_foundation_model_predict.py b/tests/unittests/general/test_tabular_foundation_model_predict.py index 9cc9cec2..e848e93d 100644 --- a/tests/unittests/general/test_tabular_foundation_model_predict.py +++ b/tests/unittests/general/test_tabular_foundation_model_predict.py @@ -20,7 +20,9 @@ CLASSIFICATION_FRAME = pd.DataFrame({"class": ["a", "b"], "a_proba": [0.7, 0.3], "b_proba": [0.3, 0.7]}) REGRESSION_FRAME = pd.DataFrame({"target": [1.5, 2.5]}) -PREDICT_ARGS = dict(train_data="train.csv", test_data="test.csv", label="class") +TRAIN_DATA = pd.DataFrame({"feature": [1, 2], "class": ["a", "b"]}) +TEST_DATA = pd.DataFrame({"feature": [3, 4]}) +PREDICT_ARGS = dict(train_data=TRAIN_DATA, test_data=TEST_DATA, label="class") @pytest.fixture(autouse=True) @@ -52,9 +54,12 @@ def test_predict_returns_prediction_series(): def test_predict_launches_predict_after_fit_job(): fm = _make_fm() fm.predict(**PREDICT_ARGS) - extra_ag_args = fm._backend.fit.call_args.kwargs["extra_ag_args"] + fit_kwargs = fm._backend.fit.call_args.kwargs + extra_ag_args = fit_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 + pd.testing.assert_frame_equal(fit_kwargs["data_channels"]["train_data"], TRAIN_DATA) + pd.testing.assert_frame_equal(fit_kwargs["data_channels"]["tuning_data"], TRAIN_DATA.iloc[[0]]) @pytest.mark.parametrize( @@ -65,7 +70,11 @@ def test_predict_pins_problem_type_from_registry(model_id, expected_problem_type """The checkpoint's task is enforced via problem_type, not inferred from the label — otherwise a continuous label would silently route mitra-classifier to the default regressor.""" fm = _make_fm(model_id=model_id, result=REGRESSION_FRAME) - fm.predict(train_data="t.csv", test_data="s.csv", label="y") + fm.predict( + train_data=pd.DataFrame({"feature": [1, 2], "y": [0, 1]}), + test_data=TEST_DATA, + label="y", + ) init_args = fm._backend.fit.call_args.kwargs["predictor_init_args"] assert init_args["problem_type"] == expected_problem_type @@ -89,7 +98,11 @@ def test_predict_proba_include_predict_false_returns_only_proba(): def test_regression_proba_equals_pred(): fm = _make_fm(model_id="mitra-regressor", result=REGRESSION_FRAME) - pred, proba = fm.predict_proba(train_data="t.csv", test_data="s.csv", label="target") + pred, proba = fm.predict_proba( + train_data=pd.DataFrame({"feature": [1, 2], "target": [1.0, 2.0]}), + test_data=TEST_DATA, + label="target", + ) assert pred.tolist() == [1.5, 2.5] pd.testing.assert_series_equal(pred, proba) @@ -126,9 +139,10 @@ def test_fm_predict_matches_tcp_fit_predict(frame): """TabularFoundationModel.predict_proba and TabularCloudPredictor.fit_predict_proba must return the same (pred, proba) off the same backend frame — they wrap one shared result-loading mechanism.""" label = "class" if frame is CLASSIFICATION_FRAME else "target" + train_data = pd.DataFrame({"feature": [1, 2], label: [0, 1]}) fm = _make_fm(model_id="mitra-classifier", result=frame) - fm_pred, fm_proba = fm.predict_proba(train_data="t.csv", test_data="s.csv", label=label) + fm_pred, fm_proba = fm.predict_proba(train_data=train_data, test_data=TEST_DATA, label=label) tcp = _make_tcp(result=frame) tcp_pred, tcp_proba = tcp.fit_predict_proba( diff --git a/tests/unittests/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index 783a9eeb..04494ee1 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -46,3 +46,51 @@ def test_tabular_foundation_model_predict(test_helper, framework_version): head = boto3.client("s3").head_object(Bucket=bucket, Key=predictions_key) assert head["ContentLength"] > 0, "predictions file on S3 should not be empty" + + +def test_tabular_foundation_model_deploy(test_helper, framework_version): + """Test TabularFoundationModel deploy to a real-time CPU endpoint and predict.""" + import boto3 + + from autogluon.cloud.model import TabularFoundationModel + + train_data = "tabular_train.csv" + test_data = "tabular_test.csv" + timestamp = test_helper.get_utc_timestamp_now() + + with tempfile.TemporaryDirectory() as temp_dir: + os.chdir(temp_dir) + test_helper.prepare_data(train_data, test_data) + n_test_rows = len(pd.read_csv(test_data)) + + inference_custom_image_uri = test_helper.get_custom_image_uri(framework_version, type="inference", gpu=False) + + model = TabularFoundationModel( + "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) + try: + endpoint_arn = boto3.client("sagemaker").describe_endpoint(EndpointName=endpoint.endpoint_name)[ + "EndpointArn" + ] + test_helper.assert_ag_cloud_tags(endpoint_arn, module="tabular", model_id="mitra-classifier") + + pred_proba = endpoint.predict_proba( + data=test_data, + train_data=train_data, + label="class", + include_predict=False, + ) + assert isinstance(pred_proba, pd.DataFrame) + assert len(pred_proba) == n_test_rows + + pred = endpoint.predict( + data=test_data, + train_data=train_data, + label="class", + ) + assert isinstance(pred, pd.Series) + assert len(pred) == n_test_rows + finally: + endpoint.delete_endpoint()