From caa2a99c6dac81ecb1e85c004fa6a172ab34e344 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 14:55:33 +0000 Subject: [PATCH 01/12] Deploy tabular foundation models to SageMaker endpoints --- docs/api.rst | 14 ++ docs/api/tabular.rst | 2 + src/autogluon/cloud/__init__.py | 2 + .../cloud/endpoint/tabular_endpoint.py | 119 +++++++++++++ src/autogluon/cloud/model/foundation_model.py | 44 ++++- src/autogluon/cloud/model/registry.py | 2 + .../sagemaker_scripts/tabular_fm_serve.py | 111 ++++++++++++ src/autogluon/cloud/scripts/script_manager.py | 1 + src/autogluon/cloud/utils/serializers.py | 3 + .../general/test_foundation_model.py | 26 +++ tests/unittests/general/test_serializers.py | 10 ++ .../general/test_tabular_endpoint.py | 72 ++++++++ .../general/test_tabular_fm_serve.py | 162 ++++++++++++++++++ 13 files changed, 562 insertions(+), 6 deletions(-) create mode 100644 src/autogluon/cloud/endpoint/tabular_endpoint.py create mode 100644 src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py create mode 100644 tests/unittests/general/test_tabular_endpoint.py create mode 100644 tests/unittests/general/test_tabular_fm_serve.py 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/endpoint/tabular_endpoint.py b/src/autogluon/cloud/endpoint/tabular_endpoint.py new file mode 100644 index 00000000..dcdeddc2 --- /dev/null +++ b/src/autogluon/cloud/endpoint/tabular_endpoint.py @@ -0,0 +1,119 @@ +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``.""" + 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. + """ + 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..32aad725 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -18,6 +18,7 @@ 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 +214,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 +610,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 +620,40 @@ 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", "serverless"] = "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. + """ + 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=inference_mode, + inference_config=inference_config, + **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.""" 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..a84a0ed3 --- /dev/null +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py @@ -0,0 +1,111 @@ +"""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) + 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: + predictor = TabularPredictor( + label=label, + problem_type=problem_type, + path=os.path.join(temp_dir, "predictor"), + ).fit( + train_data, + hyperparameters={ag_model_key: hyperparameters}, + fit_weighted_ensemble=False, + ) + + pred = predictor.predict(data, as_pandas=True, **inference_kwargs) + if predictor.can_predict_proba: + pred_proba = predictor.predict_proba(data, as_pandas=True, **inference_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 = pred + + 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..97fcf7a8 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -130,6 +130,32 @@ 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_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_fm_serve.py b/tests/unittests/general/test_tabular_fm_serve.py new file mode 100644 index 00000000..f398d58a --- /dev/null +++ b/tests/unittests/general/test_tabular_fm_serve.py @@ -0,0 +1,162 @@ +import importlib.util +import json +import sys +import types +from io import BytesIO +from pathlib import Path + +import pandas as pd + +from autogluon.cloud.utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer + +SERVE_SCRIPT = ( + Path(__file__).parents[3] + / "src" + / "autogluon" + / "cloud" + / "scripts" + / "sagemaker_scripts" + / "tabular_fm_serve.py" +) + + +class FakeTabularPredictor: + instances = [] + can_predict_proba = True + + def __init__(self, *, label, problem_type, path): + self.label = label + self.problem_type = problem_type + self.path = path + self.fit_data = None + self.fit_kwargs = None + self.predict_data = None + self.predict_kwargs = None + self.__class__.instances.append(self) + + def fit(self, train_data, **kwargs): + self.fit_data = train_data + self.fit_kwargs = kwargs + return self + + def predict(self, data, **kwargs): + self.predict_data = data + self.predict_kwargs = kwargs + return pd.Series(["a", "b"], name=self.label) + + def predict_proba(self, data, **kwargs): + return pd.DataFrame({"a": [0.8, 0.2], "b": [0.2, 0.8]}) + + +def _load_serve_module(monkeypatch): + FakeTabularPredictor.instances.clear() + tabular_module = types.ModuleType("autogluon.tabular") + tabular_module.TabularPredictor = FakeTabularPredictor + monkeypatch.setitem(sys.modules, "autogluon.tabular", tabular_module) + + spec = importlib.util.spec_from_file_location("tabular_fm_serve_under_test", SERVE_SCRIPT) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_model_fn_downloads_remote_weights_during_startup(monkeypatch): + config = { + "ag_model_key": "MITRA", + "hyperparameters": { + "fine_tune": False, + "hf_cls_model": "autogluon/mitra-classifier", + }, + "problem_type": "multiclass", + } + monkeypatch.setenv("AG_FM_SERVE_CONFIG", json.dumps(config)) + serve = _load_serve_module(monkeypatch) + snapshot_download = types.SimpleNamespace(calls=[]) + + def _snapshot_download(**kwargs): + snapshot_download.calls.append(kwargs) + return "/cache/snapshot" + + monkeypatch.setattr( + serve, + "snapshot_download", + _snapshot_download, + ) + + loaded_config = serve.model_fn("/opt/ml/model") + + assert loaded_config["hyperparameters"]["hf_cls_model"] == "autogluon/mitra-classifier" + assert snapshot_download.calls == [ + { + "repo_id": "autogluon/mitra-classifier", + "allow_patterns": ["config.json", "model.safetensors"], + } + ] + assert config["hyperparameters"]["hf_cls_model"] == "autogluon/mitra-classifier" + + +def test_transform_fits_request_train_data_before_predicting(monkeypatch): + serve = _load_serve_module(monkeypatch) + train_data = pd.DataFrame({"feature": [0, 1], "label": ["a", "b"]}) + data = pd.DataFrame({"feature": [2, 3]}) + payload = AutoGluonSerializer().serialize( + AutoGluonSerializationWrapper( + data=data, + train_data=train_data, + inference_kwargs={"label": "label"}, + ) + ) + config = { + "ag_model_key": "MITRA", + "hyperparameters": {"fine_tune": False, "hf_cls_model": "autogluon/mitra-classifier"}, + "problem_type": "multiclass", + } + + body, content_type = serve.transform_fn( + config, + payload, + "application/x-autogluon", + "application/x-parquet", + ) + + predictor = FakeTabularPredictor.instances[0] + pd.testing.assert_frame_equal(predictor.fit_data, train_data) + pd.testing.assert_frame_equal(predictor.predict_data, data) + assert predictor.problem_type == "multiclass" + assert predictor.fit_kwargs == { + "hyperparameters": { + "MITRA": {"fine_tune": False, "hf_cls_model": "autogluon/mitra-classifier"} + }, + "fit_weighted_ensemble": False, + } + assert content_type == "application/x-parquet" + result = pd.read_parquet(BytesIO(body)) + assert result.columns.tolist() == ["label", "a_proba", "b_proba"] + assert result["label"].tolist() == ["a", "b"] + + +def test_transform_requires_train_data_and_label(monkeypatch): + serve = _load_serve_module(monkeypatch) + data = pd.DataFrame({"feature": [2]}) + config = {"ag_model_key": "MITRA", "hyperparameters": {}, "problem_type": "multiclass"} + + payload_without_train = AutoGluonSerializer().serialize( + AutoGluonSerializationWrapper(data=data, inference_kwargs={"label": "label"}) + ) + try: + serve.transform_fn(config, payload_without_train, "application/x-autogluon") + except ValueError as error: + assert "train_data" in str(error) + else: + raise AssertionError("Expected missing train_data to raise") + + train_data = pd.DataFrame({"feature": [0], "label": ["a"]}) + payload_without_label = AutoGluonSerializer().serialize( + AutoGluonSerializationWrapper(data=data, train_data=train_data, inference_kwargs={}) + ) + try: + serve.transform_fn(config, payload_without_label, "application/x-autogluon") + except ValueError as error: + assert "label" in str(error) + else: + raise AssertionError("Expected missing label to raise") From c3a5f33681e94cd3987536683742ca6165741176 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:00:42 +0000 Subject: [PATCH 02/12] Fix endpoint reattachment after deserialization --- 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 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: From 72792c9d5a973ba4fbd7e028842658facdfd71af Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:01:51 +0000 Subject: [PATCH 03/12] Format tabular foundation model serving tests --- tests/unittests/general/test_tabular_fm_serve.py | 12 ++---------- 1 file changed, 2 insertions(+), 10 deletions(-) diff --git a/tests/unittests/general/test_tabular_fm_serve.py b/tests/unittests/general/test_tabular_fm_serve.py index f398d58a..80a67cca 100644 --- a/tests/unittests/general/test_tabular_fm_serve.py +++ b/tests/unittests/general/test_tabular_fm_serve.py @@ -10,13 +10,7 @@ from autogluon.cloud.utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer SERVE_SCRIPT = ( - Path(__file__).parents[3] - / "src" - / "autogluon" - / "cloud" - / "scripts" - / "sagemaker_scripts" - / "tabular_fm_serve.py" + Path(__file__).parents[3] / "src" / "autogluon" / "cloud" / "scripts" / "sagemaker_scripts" / "tabular_fm_serve.py" ) @@ -124,9 +118,7 @@ def test_transform_fits_request_train_data_before_predicting(monkeypatch): pd.testing.assert_frame_equal(predictor.predict_data, data) assert predictor.problem_type == "multiclass" assert predictor.fit_kwargs == { - "hyperparameters": { - "MITRA": {"fine_tune": False, "hf_cls_model": "autogluon/mitra-classifier"} - }, + "hyperparameters": {"MITRA": {"fine_tune": False, "hf_cls_model": "autogluon/mitra-classifier"}}, "fit_weighted_ensemble": False, } assert content_type == "application/x-parquet" From 9bc009fe273bfb4b5674328a6ddbf79bbc83f8f6 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:03:47 +0000 Subject: [PATCH 04/12] Simplify tabular foundation model deploy tests --- .../general/test_tabular_fm_serve.py | 154 ------------------ 1 file changed, 154 deletions(-) delete mode 100644 tests/unittests/general/test_tabular_fm_serve.py diff --git a/tests/unittests/general/test_tabular_fm_serve.py b/tests/unittests/general/test_tabular_fm_serve.py deleted file mode 100644 index 80a67cca..00000000 --- a/tests/unittests/general/test_tabular_fm_serve.py +++ /dev/null @@ -1,154 +0,0 @@ -import importlib.util -import json -import sys -import types -from io import BytesIO -from pathlib import Path - -import pandas as pd - -from autogluon.cloud.utils.serializers import AutoGluonSerializationWrapper, AutoGluonSerializer - -SERVE_SCRIPT = ( - Path(__file__).parents[3] / "src" / "autogluon" / "cloud" / "scripts" / "sagemaker_scripts" / "tabular_fm_serve.py" -) - - -class FakeTabularPredictor: - instances = [] - can_predict_proba = True - - def __init__(self, *, label, problem_type, path): - self.label = label - self.problem_type = problem_type - self.path = path - self.fit_data = None - self.fit_kwargs = None - self.predict_data = None - self.predict_kwargs = None - self.__class__.instances.append(self) - - def fit(self, train_data, **kwargs): - self.fit_data = train_data - self.fit_kwargs = kwargs - return self - - def predict(self, data, **kwargs): - self.predict_data = data - self.predict_kwargs = kwargs - return pd.Series(["a", "b"], name=self.label) - - def predict_proba(self, data, **kwargs): - return pd.DataFrame({"a": [0.8, 0.2], "b": [0.2, 0.8]}) - - -def _load_serve_module(monkeypatch): - FakeTabularPredictor.instances.clear() - tabular_module = types.ModuleType("autogluon.tabular") - tabular_module.TabularPredictor = FakeTabularPredictor - monkeypatch.setitem(sys.modules, "autogluon.tabular", tabular_module) - - spec = importlib.util.spec_from_file_location("tabular_fm_serve_under_test", SERVE_SCRIPT) - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module - - -def test_model_fn_downloads_remote_weights_during_startup(monkeypatch): - config = { - "ag_model_key": "MITRA", - "hyperparameters": { - "fine_tune": False, - "hf_cls_model": "autogluon/mitra-classifier", - }, - "problem_type": "multiclass", - } - monkeypatch.setenv("AG_FM_SERVE_CONFIG", json.dumps(config)) - serve = _load_serve_module(monkeypatch) - snapshot_download = types.SimpleNamespace(calls=[]) - - def _snapshot_download(**kwargs): - snapshot_download.calls.append(kwargs) - return "/cache/snapshot" - - monkeypatch.setattr( - serve, - "snapshot_download", - _snapshot_download, - ) - - loaded_config = serve.model_fn("/opt/ml/model") - - assert loaded_config["hyperparameters"]["hf_cls_model"] == "autogluon/mitra-classifier" - assert snapshot_download.calls == [ - { - "repo_id": "autogluon/mitra-classifier", - "allow_patterns": ["config.json", "model.safetensors"], - } - ] - assert config["hyperparameters"]["hf_cls_model"] == "autogluon/mitra-classifier" - - -def test_transform_fits_request_train_data_before_predicting(monkeypatch): - serve = _load_serve_module(monkeypatch) - train_data = pd.DataFrame({"feature": [0, 1], "label": ["a", "b"]}) - data = pd.DataFrame({"feature": [2, 3]}) - payload = AutoGluonSerializer().serialize( - AutoGluonSerializationWrapper( - data=data, - train_data=train_data, - inference_kwargs={"label": "label"}, - ) - ) - config = { - "ag_model_key": "MITRA", - "hyperparameters": {"fine_tune": False, "hf_cls_model": "autogluon/mitra-classifier"}, - "problem_type": "multiclass", - } - - body, content_type = serve.transform_fn( - config, - payload, - "application/x-autogluon", - "application/x-parquet", - ) - - predictor = FakeTabularPredictor.instances[0] - pd.testing.assert_frame_equal(predictor.fit_data, train_data) - pd.testing.assert_frame_equal(predictor.predict_data, data) - assert predictor.problem_type == "multiclass" - assert predictor.fit_kwargs == { - "hyperparameters": {"MITRA": {"fine_tune": False, "hf_cls_model": "autogluon/mitra-classifier"}}, - "fit_weighted_ensemble": False, - } - assert content_type == "application/x-parquet" - result = pd.read_parquet(BytesIO(body)) - assert result.columns.tolist() == ["label", "a_proba", "b_proba"] - assert result["label"].tolist() == ["a", "b"] - - -def test_transform_requires_train_data_and_label(monkeypatch): - serve = _load_serve_module(monkeypatch) - data = pd.DataFrame({"feature": [2]}) - config = {"ag_model_key": "MITRA", "hyperparameters": {}, "problem_type": "multiclass"} - - payload_without_train = AutoGluonSerializer().serialize( - AutoGluonSerializationWrapper(data=data, inference_kwargs={"label": "label"}) - ) - try: - serve.transform_fn(config, payload_without_train, "application/x-autogluon") - except ValueError as error: - assert "train_data" in str(error) - else: - raise AssertionError("Expected missing train_data to raise") - - train_data = pd.DataFrame({"feature": [0], "label": ["a"]}) - payload_without_label = AutoGluonSerializer().serialize( - AutoGluonSerializationWrapper(data=data, train_data=train_data, inference_kwargs={}) - ) - try: - serve.transform_fn(config, payload_without_label, "application/x-autogluon") - except ValueError as error: - assert "label" in str(error) - else: - raise AssertionError("Expected missing label to raise") From cde4f6a7f965ffd483987f0aa836848d439ddb35 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:06:09 +0000 Subject: [PATCH 05/12] Test tabular foundation model endpoint on SageMaker --- tests/unittests/tabular/test_tabular.py | 48 +++++++++++++++++++++++++ 1 file changed, 48 insertions(+) diff --git a/tests/unittests/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index 783a9eeb..fc6299cd 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 + + timestamp = test_helper.get_utc_timestamp_now() + train_data = pd.DataFrame( + { + "feature_a": list(range(40)), + "feature_b": [value % 3 for value in range(40)], + "class": ["negative"] * 20 + ["positive"] * 20, + } + ) + test_data = pd.DataFrame( + { + "feature_a": [5, 35], + "feature_b": [2, 2], + } + ) + + with tempfile.TemporaryDirectory() as temp_dir: + os.chdir(temp_dir) + 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) + 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") + + try: + pred, pred_proba = endpoint.predict_proba( + data=test_data, + train_data=train_data, + label="class", + ) + assert isinstance(pred, pd.Series) + assert len(pred) == len(test_data) + assert isinstance(pred_proba, pd.DataFrame) + assert len(pred_proba) == len(test_data) + assert set(pred_proba.columns) == {"negative", "positive"} + finally: + endpoint.delete_endpoint() From 4ac36a00a257f0ace914445673ca94c6527a5067 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:13:52 +0000 Subject: [PATCH 06/12] Reuse tabular test data for endpoint deployment --- tests/unittests/tabular/test_tabular.py | 23 +++++++---------------- 1 file changed, 7 insertions(+), 16 deletions(-) diff --git a/tests/unittests/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index fc6299cd..a7295299 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -54,23 +54,15 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): from autogluon.cloud.model import TabularFoundationModel + train_data = "tabular_train.csv" + test_data = "tabular_test.csv" timestamp = test_helper.get_utc_timestamp_now() - train_data = pd.DataFrame( - { - "feature_a": list(range(40)), - "feature_b": [value % 3 for value in range(40)], - "class": ["negative"] * 20 + ["positive"] * 20, - } - ) - test_data = pd.DataFrame( - { - "feature_a": [5, 35], - "feature_b": [2, 2], - } - ) 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( @@ -88,9 +80,8 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): label="class", ) assert isinstance(pred, pd.Series) - assert len(pred) == len(test_data) + assert len(pred) == n_test_rows assert isinstance(pred_proba, pd.DataFrame) - assert len(pred_proba) == len(test_data) - assert set(pred_proba.columns) == {"negative", "positive"} + assert len(pred_proba) == n_test_rows finally: endpoint.delete_endpoint() From 5793522cefe24797434ab639d6950eac25f633a9 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:15:58 +0000 Subject: [PATCH 07/12] Avoid duplicate tabular endpoint inference --- .../cloud/scripts/sagemaker_scripts/tabular_fm_serve.py | 7 +++++-- tests/unittests/tabular/test_tabular.py | 9 ++++++--- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py index a84a0ed3..940e9943 100644 --- a/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py @@ -79,6 +79,9 @@ def transform_fn(model_config, request_body, input_content_type, output_content_ 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: @@ -99,13 +102,13 @@ def transform_fn(model_config, request_body, input_content_type, output_content_ fit_weighted_ensemble=False, ) - pred = predictor.predict(data, as_pandas=True, **inference_kwargs) 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 = pred + prediction = predictor.predict(data, as_pandas=True, **inference_kwargs, **prediction_only_kwargs) return _render_response(prediction, output_content_type) diff --git a/tests/unittests/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index a7295299..f5ff81ff 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -70,14 +70,17 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): 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_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") - 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, pred_proba = endpoint.predict_proba( data=test_data, train_data=train_data, label="class", + decision_threshold=0.4, ) assert isinstance(pred, pd.Series) assert len(pred) == n_test_rows From 315ff67758b4292ae7a843bcad70a7560ac74f77 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:44:16 +0000 Subject: [PATCH 08/12] Restrict tabular foundation endpoints to realtime --- .../cloud/endpoint/tabular_endpoint.py | 11 ++++++++++- src/autogluon/cloud/model/foundation_model.py | 18 +++++++++++++++--- .../unittests/general/test_foundation_model.py | 12 ++++++++++++ tests/unittests/tabular/test_tabular.py | 1 - 4 files changed, 37 insertions(+), 5 deletions(-) diff --git a/src/autogluon/cloud/endpoint/tabular_endpoint.py b/src/autogluon/cloud/endpoint/tabular_endpoint.py index dcdeddc2..719fee7b 100644 --- a/src/autogluon/cloud/endpoint/tabular_endpoint.py +++ b/src/autogluon/cloud/endpoint/tabular_endpoint.py @@ -81,7 +81,12 @@ def predict( label: str, **inference_kwargs: Any, ) -> pd.Series: - """Fit the foundation model on ``train_data`` and predict ``data``.""" + """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, @@ -102,6 +107,10 @@ def predict_proba( """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, diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 32aad725..054272c6 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -630,7 +630,7 @@ def deploy( framework_version: str = "latest", custom_image_uri: Optional[str] = None, wait: bool = True, - inference_mode: Literal["realtime", "serverless"] = "realtime", + inference_mode: Literal["realtime"] = "realtime", inference_config: Optional[Dict[str, Any]] = None, **backend_kwargs, ) -> TabularEndpoint: @@ -638,7 +638,20 @@ def deploy( 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, @@ -646,8 +659,7 @@ def deploy( framework_version=framework_version, custom_image_uri=custom_image_uri, wait=wait, - inference_mode=inference_mode, - inference_config=inference_config, + inference_mode="realtime", **backend_kwargs, ) return TabularEndpoint( diff --git a/tests/unittests/general/test_foundation_model.py b/tests/unittests/general/test_foundation_model.py index 97fcf7a8..204a9499 100644 --- a/tests/unittests/general/test_foundation_model.py +++ b/tests/unittests/general/test_foundation_model.py @@ -156,6 +156,18 @@ def test_tabular_deploy_uses_tabular_fm_handler_and_returns_tabular_endpoint(): 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/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index f5ff81ff..bdebe8f9 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -80,7 +80,6 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): data=test_data, train_data=train_data, label="class", - decision_threshold=0.4, ) assert isinstance(pred, pd.Series) assert len(pred) == n_test_rows From 07a56f57e3107c8282089cc58b595faf268a5ef8 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:46:11 +0000 Subject: [PATCH 09/12] Exercise repeated tabular endpoint invocations --- tests/unittests/tabular/test_tabular.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/tests/unittests/tabular/test_tabular.py b/tests/unittests/tabular/test_tabular.py index bdebe8f9..04494ee1 100644 --- a/tests/unittests/tabular/test_tabular.py +++ b/tests/unittests/tabular/test_tabular.py @@ -76,14 +76,21 @@ def test_tabular_foundation_model_deploy(test_helper, framework_version): ] test_helper.assert_ag_cloud_tags(endpoint_arn, module="tabular", model_id="mitra-classifier") - pred, pred_proba = endpoint.predict_proba( + pred_proba = endpoint.predict_proba( data=test_data, train_data=train_data, label="class", + include_predict=False, ) - assert isinstance(pred, pd.Series) - assert len(pred) == n_test_rows 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() From 54b98e8ec474ca70ef0897b18f0aeb5b8c16bd26 Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:52:17 +0000 Subject: [PATCH 10/12] Use full context for tabular foundation inference --- .../cloud/backend/tabular_sagemaker_backend.py | 5 +++++ src/autogluon/cloud/model/foundation_model.py | 1 + .../scripts/sagemaker_scripts/tabular_fm_serve.py | 4 ++++ .../test_tabular_foundation_model_predict.py | 4 +++- .../unittests/tabular/test_tabular_fit_predict.py | 15 +++++++++++++++ 5 files changed, 28 insertions(+), 1 deletion(-) diff --git a/src/autogluon/cloud/backend/tabular_sagemaker_backend.py b/src/autogluon/cloud/backend/tabular_sagemaker_backend.py index 3aa2df3d..a7540df8 100644 --- a/src/autogluon/cloud/backend/tabular_sagemaker_backend.py +++ b/src/autogluon/cloud/backend/tabular_sagemaker_backend.py @@ -18,12 +18,17 @@ def fit( predictor_init_args: Dict[str, Any], predictor_fit_args: Dict[str, Any], data_channels: Dict[str, Optional[Union[str, pd.DataFrame]]], + use_full_train_data: bool = False, **kwargs, ) -> None: data_channels = self._validate_data_channels( data_channels=data_channels, predictor_init_args=predictor_init_args, ) + if use_full_train_data: + # An explicit tuning row prevents AutoGluon from removing a holdout from train_data. + # Keep the row in train_data as well so foundation models receive the full context. + data_channels["tuning_data"] = data_channels["train_data"].iloc[[0]].copy() super().fit( predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 054272c6..6fad291f 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -834,6 +834,7 @@ def predict_proba( instance_type=instance_type, custom_image_uri=custom_image_uri, wait=wait, + use_full_train_data=True, extra_ag_args=extra_ag_args, extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], **backend_kwargs, diff --git a/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py index 940e9943..af8616a0 100644 --- a/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py @@ -92,12 +92,16 @@ def transform_fn(model_config, request_body, input_content_type, output_content_ problem_type = model_config["problem_type"] with tempfile.TemporaryDirectory(prefix="ag_tabular_fm_") as temp_dir: + # Explicit tuning data prevents AutoGluon and Mitra from holding out context rows internally. + # Duplicate one row so every row in train_data remains available to Mitra during prediction. + 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, ) diff --git a/tests/unittests/general/test_tabular_foundation_model_predict.py b/tests/unittests/general/test_tabular_foundation_model_predict.py index 9cc9cec2..ac06269d 100644 --- a/tests/unittests/general/test_tabular_foundation_model_predict.py +++ b/tests/unittests/general/test_tabular_foundation_model_predict.py @@ -52,9 +52,11 @@ 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 + assert fit_kwargs["use_full_train_data"] is True @pytest.mark.parametrize( diff --git a/tests/unittests/tabular/test_tabular_fit_predict.py b/tests/unittests/tabular/test_tabular_fit_predict.py index 827e68f1..fc9ce2d7 100644 --- a/tests/unittests/tabular/test_tabular_fit_predict.py +++ b/tests/unittests/tabular/test_tabular_fit_predict.py @@ -99,6 +99,21 @@ def test_when_test_data_omits_label_column_then_validation_passes(backend): _backend_fit(backend, train, test) # does not raise +def test_when_full_train_data_requested_then_tuning_row_is_duplicated(backend): + train = pd.DataFrame({"x": [1, 2], "y": [0, 1]}) + + backend.fit( + predictor_init_args={"label": "y"}, + predictor_fit_args={}, + data_channels={"train_data": train}, + use_full_train_data=True, + ) + + data_channels = SagemakerBackend.fit.call_args.kwargs["data_channels"] + pd.testing.assert_frame_equal(data_channels["train_data"], train) + pd.testing.assert_frame_equal(data_channels["tuning_data"], train.iloc[[0]]) + + def test_when_test_data_missing_feature_columns_then_raises(backend): train = pd.DataFrame({"x": [1, 2], "z": [3, 4], "y": [0, 1]}) test = pd.DataFrame({"x": [5]}) # missing feature column `z` From 8eb978588204d1516f5f9d87b7eafcfebd48bf3d Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:53:16 +0000 Subject: [PATCH 11/12] Scope full-context workaround to endpoint serving --- .../cloud/backend/tabular_sagemaker_backend.py | 5 ----- src/autogluon/cloud/model/foundation_model.py | 1 - .../test_tabular_foundation_model_predict.py | 4 +--- .../unittests/tabular/test_tabular_fit_predict.py | 15 --------------- 4 files changed, 1 insertion(+), 24 deletions(-) diff --git a/src/autogluon/cloud/backend/tabular_sagemaker_backend.py b/src/autogluon/cloud/backend/tabular_sagemaker_backend.py index a7540df8..3aa2df3d 100644 --- a/src/autogluon/cloud/backend/tabular_sagemaker_backend.py +++ b/src/autogluon/cloud/backend/tabular_sagemaker_backend.py @@ -18,17 +18,12 @@ def fit( predictor_init_args: Dict[str, Any], predictor_fit_args: Dict[str, Any], data_channels: Dict[str, Optional[Union[str, pd.DataFrame]]], - use_full_train_data: bool = False, **kwargs, ) -> None: data_channels = self._validate_data_channels( data_channels=data_channels, predictor_init_args=predictor_init_args, ) - if use_full_train_data: - # An explicit tuning row prevents AutoGluon from removing a holdout from train_data. - # Keep the row in train_data as well so foundation models receive the full context. - data_channels["tuning_data"] = data_channels["train_data"].iloc[[0]].copy() super().fit( predictor_init_args=predictor_init_args, predictor_fit_args=predictor_fit_args, diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 6fad291f..054272c6 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -834,7 +834,6 @@ def predict_proba( instance_type=instance_type, custom_image_uri=custom_image_uri, wait=wait, - use_full_train_data=True, extra_ag_args=extra_ag_args, extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}], **backend_kwargs, diff --git a/tests/unittests/general/test_tabular_foundation_model_predict.py b/tests/unittests/general/test_tabular_foundation_model_predict.py index ac06269d..9cc9cec2 100644 --- a/tests/unittests/general/test_tabular_foundation_model_predict.py +++ b/tests/unittests/general/test_tabular_foundation_model_predict.py @@ -52,11 +52,9 @@ def test_predict_returns_prediction_series(): def test_predict_launches_predict_after_fit_job(): fm = _make_fm() fm.predict(**PREDICT_ARGS) - fit_kwargs = fm._backend.fit.call_args.kwargs - extra_ag_args = fit_kwargs["extra_ag_args"] + extra_ag_args = fm._backend.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 - assert fit_kwargs["use_full_train_data"] is True @pytest.mark.parametrize( diff --git a/tests/unittests/tabular/test_tabular_fit_predict.py b/tests/unittests/tabular/test_tabular_fit_predict.py index fc9ce2d7..827e68f1 100644 --- a/tests/unittests/tabular/test_tabular_fit_predict.py +++ b/tests/unittests/tabular/test_tabular_fit_predict.py @@ -99,21 +99,6 @@ def test_when_test_data_omits_label_column_then_validation_passes(backend): _backend_fit(backend, train, test) # does not raise -def test_when_full_train_data_requested_then_tuning_row_is_duplicated(backend): - train = pd.DataFrame({"x": [1, 2], "y": [0, 1]}) - - backend.fit( - predictor_init_args={"label": "y"}, - predictor_fit_args={}, - data_channels={"train_data": train}, - use_full_train_data=True, - ) - - data_channels = SagemakerBackend.fit.call_args.kwargs["data_channels"] - pd.testing.assert_frame_equal(data_channels["train_data"], train) - pd.testing.assert_frame_equal(data_channels["tuning_data"], train.iloc[[0]]) - - def test_when_test_data_missing_feature_columns_then_raises(backend): train = pd.DataFrame({"x": [1, 2], "z": [3, 4], "y": [0, 1]}) test = pd.DataFrame({"x": [5]}) # missing feature column `z` From 515232aa65398469299a0c766fa12df5b292ddfb Mon Sep 17 00:00:00 2001 From: Oleksandr Shchur Date: Fri, 18 Sep 2026 15:54:44 +0000 Subject: [PATCH 12/12] Provide explicit tuning data for tabular FM prediction --- src/autogluon/cloud/model/foundation_model.py | 8 ++++++- .../sagemaker_scripts/tabular_fm_serve.py | 3 +-- .../test_tabular_foundation_model_predict.py | 24 +++++++++++++++---- 3 files changed, 27 insertions(+), 8 deletions(-) diff --git a/src/autogluon/cloud/model/foundation_model.py b/src/autogluon/cloud/model/foundation_model.py index 054272c6..2534fa23 100644 --- a/src/autogluon/cloud/model/foundation_model.py +++ b/src/autogluon/cloud/model/foundation_model.py @@ -13,6 +13,7 @@ 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 @@ -822,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 @@ -829,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/scripts/sagemaker_scripts/tabular_fm_serve.py b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py index af8616a0..d64af70d 100644 --- a/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py +++ b/src/autogluon/cloud/scripts/sagemaker_scripts/tabular_fm_serve.py @@ -92,8 +92,7 @@ def transform_fn(model_config, request_body, input_content_type, output_content_ problem_type = model_config["problem_type"] with tempfile.TemporaryDirectory(prefix="ag_tabular_fm_") as temp_dir: - # Explicit tuning data prevents AutoGluon and Mitra from holding out context rows internally. - # Duplicate one row so every row in train_data remains available to Mitra during prediction. + # 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, 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(