"""FoundationModel — predict with pretrained foundation models on AWS."""
from __future__ import annotations
import json
import logging
import tarfile
import tempfile
from abc import abstractmethod
from pathlib import Path
from typing import Any, Literal
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
from ..utils.constants import DEFAULT_FRAMEWORK_VERSION
from ..utils.sagemaker_api import reject_legacy_kwargs
from ..utils.utils import split_pred_and_pred_proba
from ..version import __version__
from .registry import get_model_config
logger = logging.getLogger(__name__)
# SageMaker extracts model.tar.gz to /opt/ml/model in the container.
_CONTAINER_WEIGHTS_DIR = "/opt/ml/model/weights"
_AG_CLOUD_VERSION_METADATA_KEY = "autogluon-cloud-version"
def _s3_head_or_none(s3_client: Any, bucket: str, key: str) -> dict[str, Any] | None:
"""Return ``head_object`` response if the key exists, ``None`` for 404. Other errors propagate."""
from botocore.exceptions import ClientError
try:
return s3_client.head_object(Bucket=bucket, Key=key)
except ClientError as e:
if e.response.get("Error", {}).get("Code") in ("404", "NoSuchKey", "NotFound"):
return None
raise
class FoundationModel:
"""
Pretrained foundation model inference on AWS.
Factory: ``FoundationModel(model_id, ...)`` dispatches on the model's task and returns the
appropriate task-specific subclass (:class:`TimeSeriesFoundationModel`, ``TabularFoundationModel``).
Most users instantiate the subclass directly instead.
Examples
--------
>>> model = FoundationModel("chronos-2") # returns a TimeSeriesFoundationModel
>>> predictions = model.predict(data, prediction_length=24)
"""
_backend_map: dict[str, str] = {}
_predictor_type: str
def __new__(cls, model_id: str, **kwargs) -> Self:
if cls is not FoundationModel:
return super().__new__(cls)
config = get_model_config(model_id)
problem_type = config.problem_type
if problem_type == "forecasting":
return super().__new__(TimeSeriesFoundationModel)
elif problem_type in ("multiclass", "regression"):
return super().__new__(TabularFoundationModel)
raise ValueError(f"Unsupported problem_type: {problem_type}")
def __init__(
self,
model_id: str,
*,
cloud_output_path: str | None = None,
role: str | None = None,
hyperparameters: dict[str, Any] | None = None,
model_artifact_uri: str | None = None,
backend: Literal["sagemaker"] = "sagemaker",
):
"""
Parameters
----------
model_id: str
ID of the foundation model from the model registry. See
`Available models <https://auto.gluon.ai/cloud/stable/tutorials/foundation-model-timeseries.html#available-models>`_
in the foundation model tutorial for the list of supported values.
cloud_output_path: str | None, default = None
S3 location where intermediate artifacts are stored. Accepts:
* ``s3://bucket`` — a unique timestamped subfolder ``ag-<timestamp>`` is appended.
* ``s3://bucket/prefix`` — used verbatim. Re-running with the same prefix will overwrite previously written
artifacts.
* ``None`` (default) — use the bucket saved in ``~/.autogluon/cloud.yaml`` (set by
:func:`autogluon.cloud.bootstrap` / :func:`autogluon.cloud.register`) and append a timestamped subfolder.
Raises if no bucket is configured.
role: str | None, default = None
ARN of the SageMaker execution role used to run training and inference jobs. If ``None``, falls back to
``role_arn`` in ``~/.autogluon/cloud.yaml`` (set by :func:`autogluon.cloud.bootstrap` /
:func:`autogluon.cloud.register`), and finally to the role of the current AWS identity.
hyperparameters: dict[str, Any] | None, default = None
Default hyperparameters applied to inference and (when supported) training.
model_artifact_uri: str | None, default = None
S3 URI of a pre-bundled ``model.tar.gz`` produced by :meth:`cache_model_artifact`. When set, deploys skip
the runtime HuggingFace download and load weights from the bundled artifact.
backend: Literal["sagemaker"], default = "sagemaker"
Cloud backend to use.
"""
self.model_id = model_id
self.model_artifact_uri = model_artifact_uri
self.cloud_output_path = resolve_cloud_output_path(cloud_output_path, backend_name=backend)
self._config = get_model_config(model_id)
self._hyperparameter_overrides = hyperparameters or {}
self._tmpdir = tempfile.TemporaryDirectory(prefix="ag_fm_")
backend_name = self._backend_map.get(backend)
if backend_name is None:
raise ValueError(
f"Backend '{backend}' is not supported for {self.__class__.__name__}. "
f"Available: {list(self._backend_map.keys())}"
)
self._backend = BackendFactory.get_backend(
backend=backend_name,
local_output_path=self._tmpdir.name,
cloud_output_path=self.cloud_output_path,
predictor_type=self._predictor_type,
resource_prefix=f"ag-cloud-{self.model_id}",
role=role,
)
def _get_hyperparameters(
self, context: Literal["inference", "training"], overrides: dict[str, Any] | None = None
) -> dict[str, Any]:
"""Merge registry defaults → constructor overrides → call-site overrides, defaulting the model's
weights-source hyperparameter (``model_source_hyperparameter``) to ``model_source_uri`` if not set."""
if context == "inference":
registry_defaults = self._config.inference_hyperparameters
else:
registry_defaults = self._config.training_hyperparameters
merged = registry_defaults | self._hyperparameter_overrides | (overrides or {})
if self._config.model_source_hyperparameter is not None:
merged.setdefault(self._config.model_source_hyperparameter, self._config.model_source_uri)
return merged
@abstractmethod
def _build_predictor_init_args(self, **user_kwargs) -> dict[str, Any]:
"""Build predictor_init_args dict from user-provided kwargs.
Subclasses override to map their public API kwargs (e.g., prediction_length,
target, known_covariates_names) to the dict that TimeSeriesPredictor/TabularPredictor expects.
"""
...
@abstractmethod
def _build_predictor_fit_args(self, hyperparameters: dict[str, Any] | None = None) -> dict[str, Any]:
"""Build predictor_fit_args dict. Subclasses override with task-specific logic."""
...
@property
@abstractmethod
def _serve_script_path(self) -> str:
"""Path to the serve script for this model type."""
...
@abstractmethod
def deploy(self, **kwargs):
"""Deploy model to a real-time endpoint.
Subclasses implement this and return a task-specific endpoint
(e.g., TimeSeriesEndpoint, TabularEndpoint).
"""
...
@abstractmethod
def predict(self, data: str | Path | pd.DataFrame, wait: bool = True, **kwargs) -> pd.DataFrame | None:
"""Subclasses override with task-specific signature."""
...
def _deploy_backend(
self,
instance_type: str | None = None,
endpoint_name: str | None = None,
hyperparameters: dict[str, Any] | None = None,
framework_version: str = DEFAULT_FRAMEWORK_VERSION,
custom_image_uri: str | None = None,
wait: bool = True,
inference_mode: Literal["realtime", "serverless"] = "realtime",
inference_config: dict[str, Any] | None = None,
**backend_kwargs,
) -> None:
"""Shared deploy logic. Subclasses call this then wrap the endpoint."""
if inference_mode == "serverless" and instance_type is not None:
raise ValueError("`instance_type` must not be set when `inference_mode='serverless'`.")
if instance_type is None and inference_mode == "realtime":
instance_type = self._config.deploy_instance_type
merged_hp = self._get_hyperparameters("inference", hyperparameters)
if self.model_artifact_uri is not None:
source_hp = self._config.model_source_hyperparameter
user_model_path = (hyperparameters or {}).get(source_hp) or self._hyperparameter_overrides.get(source_hp)
if user_model_path is not None:
raise ValueError(
f"Cannot set hyperparameters['{source_hp}'] when model_artifact_uri is in use — the bundled "
f"artifact determines the in-container weights path ({_CONTAINER_WEIGHTS_DIR}). Drop "
f"'{source_hp}', or call deploy() on a FoundationModel without model_artifact_uri."
)
merged_hp[source_hp] = _CONTAINER_WEIGHTS_DIR
fm_serve_config = {
"ag_model_key": self._config.ag_model_key,
"hyperparameters": merged_hp,
"problem_type": self._config.problem_type,
}
# FM deploys never repack: predictor_path is either None (script-only tarball is built locally) or a
# pre-bundled cache artifact that already contains the serve script.
self._backend.deploy(
predictor_path=self.model_artifact_uri,
endpoint_name=endpoint_name,
framework_version=framework_version,
instance_type=instance_type,
custom_image_uri=custom_image_uri,
wait=wait,
entry_point=self._serve_script_path,
fm_serve_config=fm_serve_config,
inference_mode=inference_mode,
inference_config=inference_config,
repack=False,
extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}],
**backend_kwargs,
)
assert self._backend.endpoint_name is not None
def fit(
self,
train_data: str | Path | pd.DataFrame,
output_path: str | None = None,
instance_type: str | None = None,
hyperparameters: dict[str, Any] | None = None,
wait: bool = True,
**kwargs,
) -> Self:
"""
Fine-tune the model. Returns a new FoundationModel pointing to the fine-tuned artifact.
Parameters
----------
train_data: str | Path | pd.DataFrame
Training data, as a ``pd.DataFrame`` or local/S3 path to a data file.
output_path: str | None, default = None
S3 path to store fine-tuned model.
If None, will auto-generate under cloud_output_path.
instance_type: str | None, default = None
Instance type for the training job.
If None, will use the default from the model registry.
hyperparameters: dict[str, Any] | None, default = None
Model hyperparameters for training. Overrides values passed to the constructor.
Available hyperparameters for each model are listed in the AutoGluon documentation.
wait: bool, default = True
If True, block until training completes.
Returns
-------
FoundationModel
New instance with hyperparameters pointing to the fine-tuned artifact.
:meta private:
"""
if not self._config.fine_tunable:
raise ValueError(f"Model '{self.model_id}' does not support fine-tuning.")
raise NotImplementedError
def cache_model_artifact(self, cache_path: str, *, overwrite: bool = False) -> Self:
"""
Download model weights from HuggingFace, bundle them with the FM serve script into a SageMaker-compatible
``model.tar.gz``, and upload to S3.
Lets :meth:`deploy` skip the runtime HuggingFace download — required for network-isolated endpoints (e.g.
SageMaker Serverless Inference). Returns a new :class:`FoundationModel` with ``model_artifact_uri`` set to the
uploaded tarball.
Destination key: ``{cache_path}/{model_id}/model.tar.gz``. If it already exists, upload is skipped unless
``overwrite=True``; a stale-cache mismatch between the bundled artifact's autogluon-cloud version and the
current version raises ``RuntimeError`` and prompts the caller to re-bundle.
Parameters
----------
cache_path: str
S3 prefix under which the artifact will be uploaded. Multiple foundation models can share one prefix.
overwrite: bool, default = False
If True, re-upload even when the destination key exists.
Returns
-------
FoundationModel
A new instance with ``model_artifact_uri`` populated. The original is unchanged.
"""
from huggingface_hub import snapshot_download
if not cache_path.startswith("s3://"):
raise ValueError(f"cache_path must be an s3:// URI, got: {cache_path!r}")
if self._config.model_source_hyperparameter is None:
raise ValueError(
f"Model '{self.model_id}' does not support cache_model_artifact: its weights are downloaded by "
f"AutoGluon at runtime and cannot be loaded from a bundled artifact."
)
source_uri = self._config.model_source_uri
cache_key = f"{cache_path.rstrip('/')}/{self.model_id}/model.tar.gz"
bucket, key = s3_path_to_bucket_prefix(cache_key)
s3 = self._backend.sagemaker_session.boto_session.client("s3")
head = None if overwrite else _s3_head_or_none(s3, bucket, key)
if head is not None:
cached_version = head["Metadata"].get(_AG_CLOUD_VERSION_METADATA_KEY)
if cached_version != __version__:
raise RuntimeError(
f"Cached artifact at {cache_key} was bundled with autogluon-cloud "
f"{cached_version!r}, current is {__version__!r}. "
f"Pass overwrite=True to re-bundle and re-upload."
)
logger.info(f"Cached artifact already exists at {cache_key}; skipping upload")
else:
with tempfile.TemporaryDirectory(prefix="ag_fm_cache_") as tmp:
tmp_path = Path(tmp)
weights_dir = tmp_path / "weights"
logger.info(f"Downloading {source_uri} from HuggingFace to {weights_dir}")
# trusted AG-owned repo, numeric-only outputs, no code-execution path
snapshot_download(repo_id=source_uri, local_dir=str(weights_dir)) # nosec B615
# Mirror the layout produced by SagemakerBackend._create_serve_script_tarball:
# entry-point script + serving_utils/ under code/, so the cached endpoint can
# `from serving_utils.timeseries import ...` exactly like a fresh deploy.
serve_script = Path(self._serve_script_path)
tarball = tmp_path / "model.tar.gz"
logger.info(f"Bundling weights + serve script into {tarball}")
with tarfile.open(tarball, "w:gz") as tar:
tar.add(weights_dir, arcname="weights")
tar.add(serve_script, arcname=f"code/{serve_script.name}")
tar.add(ScriptManager.SAGEMAKER_SERVING_UTILS_DIR, arcname="code/serving_utils")
logger.info(f"Uploading to {cache_key}")
s3.upload_file(
str(tarball),
bucket,
key,
ExtraArgs={"Metadata": {_AG_CLOUD_VERSION_METADATA_KEY: __version__}},
)
return self.__class__(
model_id=self.model_id,
hyperparameters=self._hyperparameter_overrides or None,
model_artifact_uri=cache_key,
cloud_output_path=self.cloud_output_path,
role=self._backend.role_arn,
)
def to_dict(self) -> dict[str, Any]:
"""Serialize the model identity. Runtime context (``role``, ``cloud_output_path``) is excluded so configs can
be shared across users."""
out: dict[str, Any] = {"model_id": self.model_id}
if self._hyperparameter_overrides:
out["hyperparameters"] = self._hyperparameter_overrides
if self.model_artifact_uri:
out["model_artifact_uri"] = self.model_artifact_uri
return out
def to_json(self) -> str:
"""Serialize :meth:`to_dict` output as a JSON string."""
return json.dumps(self.to_dict())
@classmethod
def from_dict(cls, config: dict[str, Any], **runtime_context: Any) -> Self:
"""Restore from :meth:`to_dict` output. Pass ``role`` / ``cloud_output_path`` as ``runtime_context``."""
return cls(**config, **runtime_context)
@classmethod
def from_json(cls, s: str, **runtime_context: Any) -> Self:
"""Restore from a :meth:`to_json` string."""
return cls.from_dict(json.loads(s), **runtime_context)
[docs]
class TimeSeriesFoundationModel(FoundationModel):
"""Pretrained time series foundation model for zero-shot forecasting on Amazon SageMaker.
Wraps pretrained models like `Chronos-2 <https://huggingface.co/autogluon/chronos-2>`_ and
Chronos-Bolt and runs prediction as a managed SageMaker job, with no training required. See
`the foundation model tutorial <https://auto.gluon.ai/cloud/stable/tutorials/foundation-model-timeseries.html>`_
for the supported ``model_id`` values and a full walkthrough.
Predictions can be produced in three modes:
* **Batch** — :meth:`predict` runs a one-off SageMaker training job and writes forecasts to S3.
Best for one-shot inference.
* **Real-time** — :meth:`deploy` provisions a real-time endpoint; call
:meth:`TimeSeriesEndpoint.predict` for low-latency inference, then
:meth:`TimeSeriesEndpoint.delete_endpoint` to tear it down.
* **Serverless** — :meth:`deploy` with ``inference_mode="serverless"`` provisions a SageMaker
Serverless Inference endpoint that scales to zero. Requires a cached model artifact (see
:meth:`cache_model_artifact`).
"""
_backend_map = {SAGEMAKER: TIMESERIES_SAGEMAKER}
_predictor_type = "timeseries"
@property
def _serve_script_path(self) -> str:
return ScriptManager.SAGEMAKER_TIMESERIES_FM_SERVE_SCRIPT_PATH
[docs]
@reject_legacy_kwargs
def deploy(
self,
instance_type: str | None = None,
endpoint_name: str | None = None,
hyperparameters: dict[str, Any] | None = None,
framework_version: str = DEFAULT_FRAMEWORK_VERSION,
custom_image_uri: str | None = None,
wait: bool = True,
inference_mode: Literal["realtime", "serverless"] = "realtime",
inference_config: dict[str, Any] | None = None,
**backend_kwargs,
) -> TimeSeriesEndpoint:
"""
Deploy model to an inference endpoint.
Parameters
----------
instance_type: str | None, default = None
Instance type for the endpoint. Defaults to the model registry value. Must be ``None``
when ``inference_mode="serverless"``.
endpoint_name: str | None, default = None
Custom endpoint name. If None, will auto-generate a unique name.
hyperparameters: dict[str, Any] | None, default = None
Model hyperparameters for inference. Overrides values passed to the constructor.
framework_version: str, default = "1.6"
AutoGluon version, e.g. "1.6". Uses the official AutoGluon DLC image for this version.
custom_image_uri: str | None, default = None
Custom Docker image URI for the inference container.
wait: bool, default = True
Whether to block until the endpoint is ready.
inference_mode: Literal["realtime", "serverless"], default = "realtime"
Endpoint type. ``"serverless"`` provisions a SageMaker Serverless Inference endpoint
(no instance management, scales to zero).
inference_config: dict[str, Any] | None, default = None
Serverless settings (``memory_size_in_mb``, ``max_concurrency``, ``provisioned_concurrency``).
**backend_kwargs: Any
Additional SageMaker arguments:
* ``initial_instance_count``: Number of instances for the endpoint. Defaults to 1. Ignored when
``inference_mode="serverless"``.
* ``volume_size``: Size in GB of the EBS volume to use for the endpoint. Ignored for GPU instances.
* ``backend_overrides``: raw SageMaker request fields for settings without a dedicated argument.
* Keys: request names from the *SageMaker API* section below.
* Values: request fields in PascalCase, as in the SageMaker API and boto3. Deep-merged over the request
built by AutoGluon-Cloud; lists and other non-dict values replace the generated ones.
* Example: ``{"ProductionVariant": {"ModelDataDownloadTimeoutInSeconds": 1200}}``
SageMaker API
-------------
* :sm-api:`CreateModel`: registers the model artifact and inference image as a SageMaker model.
* :sm-api:`CreateEndpointConfig`: defines the endpoint's single :sm-api:`ProductionVariant`: instance type and
count, or the serverless settings.
* :sm-api:`CreateEndpoint`: launches the endpoint.
The endpoint is billed until :meth:`TimeSeriesEndpoint.delete_endpoint` deletes it.
"""
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 TimeSeriesEndpoint(
endpoint_name=self._backend.endpoint_name,
session=self._backend.sagemaker_session.boto_session,
)
def _build_predictor_fit_args(self, hyperparameters: dict[str, Any] | None = None) -> dict[str, Any]:
merged_hp = self._get_hyperparameters("inference", hyperparameters)
return {
"hyperparameters": {self._config.ag_model_key: merged_hp},
"skip_model_selection": True,
}
def _build_predictor_init_args(
self,
target: str = "target",
prediction_length: int = 1,
quantile_levels: list[float] | None = None,
**kwargs,
) -> dict[str, Any]:
"""Map user kwargs to TimeSeriesPredictor init args."""
args: dict[str, Any] = {
"target": target,
"prediction_length": prediction_length,
}
if quantile_levels is not None:
args["quantile_levels"] = quantile_levels
return args
[docs]
@reject_legacy_kwargs
def predict(
self,
data: str | Path | pd.DataFrame,
target: str = "target",
id_column: str = "item_id",
timestamp_column: str = "timestamp",
known_covariates: str | Path | pd.DataFrame | None = None,
static_features: str | Path | pd.DataFrame | None = None,
prediction_length: int = 1,
quantile_levels: list[float] | None = None,
predictions_path: str | None = None,
hyperparameters: dict[str, Any] | None = None,
instance_type: str | None = None,
framework_version: str = DEFAULT_FRAMEWORK_VERSION,
custom_image_uri: str | None = None,
wait: bool = True,
**backend_kwargs,
) -> pd.DataFrame | JobPredictionFuture:
"""
Run batch prediction for time series.
Parameters
----------
data: str | Path | pd.DataFrame
Historical time series to forecast from, in long format, as a ``pd.DataFrame`` or local/S3 path to
a data file. See the `TimeSeriesPredictor docs <https://auto.gluon.ai/stable/api/autogluon.timeseries.TimeSeriesPredictor.html>`_
for the expected format.
target: str, default = "target"
Name of the column that contains the target values to forecast.
id_column: str, default = "item_id"
Name of the column with the unique identifier of each time series (item).
timestamp_column: str, default = "timestamp"
Name of the column with the observation timestamps.
known_covariates: str | Path | pd.DataFrame | None, default = None
Future values of the known covariates over the forecast horizon. Covariate column names are
inferred from the columns (excluding ``id_column`` and ``timestamp_column``).
static_features: str | Path | pd.DataFrame | None, default = None
Static (time-independent) features describing each individual time series.
prediction_length: int, default = 1
Forecast horizon: how many time steps into the future the model should predict.
quantile_levels: list[float] | None, default = None
List of increasing decimals between 0 and 1 specifying which quantiles to estimate. Defaults
to ``[0.1, 0.2, ..., 0.9]``.
predictions_path: str | None, default = None
S3 URL where predictions will be written by the prediction job (e.g.
``s3://my-bucket/runs/2024-05-01/predictions.csv``). The container's SageMaker execution
role must have ``s3:PutObject`` permission for this location. Defaults to
``{cloud_output_path}/{job_name}/predictions.csv``. Predictions use AutoGluon's canonical
column names ``item_id`` and ``timestamp``, regardless of the ``id_column`` /
``timestamp_column`` passed in.
hyperparameters: dict[str, Any] | None, default = None
Model hyperparameters for inference. Overrides values passed to the constructor.
instance_type: str | None, default = None
Instance type for the prediction job. If None, uses registry default.
framework_version: str, default = "1.6"
AutoGluon version, e.g. "1.6". Uses the official AutoGluon DLC image for this version.
custom_image_uri: str | None, default = None
Custom Docker image URI for the container.
wait: bool, default = True
If True, block and return a ``pd.DataFrame``. If False, return a
:class:`JobPredictionFuture` immediately — call ``.result()`` on it later to
retrieve the ``pd.DataFrame``, or ``.status()`` to check progress.
**backend_kwargs: Any
Additional SageMaker arguments:
* ``job_name``: Name of the training job that runs the prediction. Auto-generated if not set.
* ``volume_size``: Size in GB of the storage volume to use for the job. Defaults to 100.
* ``backend_overrides``: raw SageMaker request fields for settings without a dedicated argument.
* Keys: request names from the *SageMaker API* section below.
* Values: request fields in PascalCase, as in the SageMaker API and boto3. Deep-merged over the request
built by AutoGluon-Cloud; lists and other non-dict values replace the generated ones.
* Example: ``{"CreateTrainingJob": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}``
Returns
-------
pd.DataFrame | JobPredictionFuture
``pd.DataFrame`` if ``wait=True``; a :class:`JobPredictionFuture` otherwise.
SageMaker API
-------------
* :sm-api:`CreateTrainingJob`: runs the prediction as a training job (not a batch transform job) on
``instance_type``. Predictions are written to ``predictions_path``.
"""
if instance_type is None:
instance_type = self._config.predict_instance_type
predictor_init_args = self._build_predictor_init_args(
target=target,
prediction_length=prediction_length,
quantile_levels=quantile_levels,
)
predictor_fit_args = self._build_predictor_fit_args(hyperparameters)
data_channels = {
"train_data": data,
"known_covariates": known_covariates,
"static_features": static_features,
}
extra_ag_args: dict[str, Any] = {"predict_after_fit": True, "save_predictor": False}
if predictions_path is not None:
extra_ag_args["predictions_path"] = predictions_path
self._backend.fit(
predictor_init_args=predictor_init_args,
predictor_fit_args=predictor_fit_args,
data_channels=data_channels,
id_column=id_column,
timestamp_column=timestamp_column,
framework_version=framework_version,
instance_type=instance_type,
custom_image_uri=custom_image_uri,
wait=wait,
extra_ag_args=extra_ag_args,
extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}],
**backend_kwargs,
)
if not wait:
return JobPredictionFuture(
job=self._backend._fit_job,
result_loader=self._backend.get_fit_predict_results,
)
return self._backend.get_fit_predict_results()
[docs]
class TabularFoundationModel(FoundationModel):
"""Foundation model for tabular prediction on Amazon SageMaker.
Wraps pretrained tabular models like `Mitra <https://huggingface.co/autogluon/mitra-classifier>`_ and
runs prediction as a managed SageMaker job, with no training required. Each ``model_id`` targets a
single task:
* Classification: ``mitra-classifier``, ``tabicl-v2-classifier``, ``tabdpt-turbo-classifier``.
* Regression: ``mitra-regressor``, ``tabicl-v2-regressor``, ``tabdpt-turbo-regressor``,
``nori-regressor``, ``nori-30m-regressor``.
Predictions can be produced in batch mode with :meth:`predict` / :meth:`predict_proba`, or through a
real-time endpoint created with :meth:`deploy`. In both modes, labeled ``train_data`` provides the
in-context examples for each prediction.
"""
_backend_map = {SAGEMAKER: TABULAR_SAGEMAKER}
_predictor_type = "tabular"
@property
def _serve_script_path(self) -> str:
return ScriptManager.SAGEMAKER_TABULAR_FM_SERVE_SCRIPT_PATH
[docs]
@reject_legacy_kwargs
def deploy(
self,
instance_type: str | None = None,
endpoint_name: str | None = None,
hyperparameters: dict[str, Any] | None = None,
framework_version: str = DEFAULT_FRAMEWORK_VERSION,
custom_image_uri: str | None = None,
wait: bool = True,
inference_mode: Literal["realtime"] = "realtime",
inference_config: dict[str, Any] | None = 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.
Parameters
----------
instance_type: str | None, default = None
Instance type for the endpoint. Defaults to the model registry value.
endpoint_name: str | None, default = None
Custom endpoint name. If None, will auto-generate a unique name.
hyperparameters: dict[str, Any] | None, default = None
Model hyperparameters for inference. Overrides values passed to the constructor.
framework_version: str, default = "1.6"
AutoGluon version, e.g. "1.6". Uses the official AutoGluon DLC image for this version.
custom_image_uri: str | None, default = None
Custom Docker image URI for the inference container.
wait: bool, default = True
Whether to block until the endpoint is ready.
inference_mode: Literal["realtime"], default = "realtime"
Endpoint type. Only ``"realtime"`` is supported.
inference_config: dict[str, Any] | None, default = None
Not supported; must be None.
**backend_kwargs: Any
Additional SageMaker arguments:
* ``initial_instance_count``: Number of instances for the endpoint. Defaults to 1.
* ``volume_size``: Size in GB of the EBS volume to use for the endpoint. Ignored for GPU instances.
* ``backend_overrides``: raw SageMaker request fields for settings without a dedicated argument.
* Keys: request names from the *SageMaker API* section below.
* Values: request fields in PascalCase, as in the SageMaker API and boto3. Deep-merged over the request
built by AutoGluon-Cloud; lists and other non-dict values replace the generated ones.
* Example: ``{"ProductionVariant": {"ModelDataDownloadTimeoutInSeconds": 1200}}``
SageMaker API
-------------
* :sm-api:`CreateModel`: registers the model artifact and inference image as a SageMaker model.
* :sm-api:`CreateEndpointConfig`: defines the endpoint's single :sm-api:`ProductionVariant`: instance type and
count.
* :sm-api:`CreateEndpoint`: launches the endpoint.
The endpoint is billed until :meth:`TabularEndpoint.delete_endpoint` deletes it.
"""
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_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."""
return {"label": label, "problem_type": self._config.problem_type}
def _build_predictor_fit_args(self, hyperparameters: dict[str, Any] | None = None) -> dict[str, Any]:
merged_hp = self._get_hyperparameters("inference", hyperparameters)
return {
"hyperparameters": {self._config.ag_model_key: merged_hp},
"fit_weighted_ensemble": False,
}
def _load_results(
self, *, include_predict: bool, predict_only: bool = False
) -> tuple[pd.Series, pd.DataFrame | pd.Series] | pd.DataFrame | pd.Series:
# The training container writes [pred, <class>_proba...]; regression has only the pred column.
raw = self._backend.get_fit_predict_results()
pred, pred_proba = split_pred_and_pred_proba(raw)
if pred_proba is None: # regression: proba mirrors pred, matching TabularPredictor.predict_proba
pred_proba = pred
if predict_only:
return pred
elif include_predict:
return pred, pred_proba
else:
return pred_proba
[docs]
@reject_legacy_kwargs
def predict(
self,
test_data: str | Path | pd.DataFrame,
train_data: str | Path | pd.DataFrame,
label: str,
*,
predictions_path: str | None = None,
hyperparameters: dict[str, Any] | None = None,
instance_type: str | None = None,
framework_version: str = DEFAULT_FRAMEWORK_VERSION,
custom_image_uri: str | None = None,
wait: bool = True,
**backend_kwargs,
) -> pd.Series | JobPredictionFuture:
"""
Run batch prediction for tabular tasks.
For tabular foundation models (e.g., Mitra), ``train_data`` provides the few-shot context and
``test_data`` contains the rows to predict on.
Parameters
----------
test_data: str | Path | pd.DataFrame
Data to predict on. Must contain every feature column present in ``train_data`` except ``label``.
train_data: str | Path | pd.DataFrame
Labeled few-shot context for the foundation model, as a ``pd.DataFrame`` or local/S3 path to a data file.
label: str
Target column name in ``train_data``.
predictions_path: str | None, default = None
S3 URL where predictions will be written by the training container (e.g.
``s3://my-bucket/runs/2024-05-01/predictions.csv``). Defaults to
``{cloud_output_path}/{job_name}/predictions.csv``.
hyperparameters: dict[str, Any] | None, default = None
Model hyperparameters for inference. Overrides values passed to the constructor.
instance_type: str | None, default = None
Instance type for the prediction job. If None, uses registry default.
framework_version: str, default = "1.6"
AutoGluon version, e.g. "1.6". Uses the official AutoGluon DLC image for this version.
custom_image_uri: str | None, default = None
Custom Docker image URI for the container.
wait: bool, default = True
If True, block and return the predictions. If False, return a :class:`JobPredictionFuture`
immediately — call ``.result()`` on it later to retrieve the predictions.
**backend_kwargs: Any
Additional SageMaker arguments:
* ``job_name``: Name of the training job that runs the prediction. Auto-generated if not set.
* ``volume_size``: Size in GB of the storage volume to use for the job. Defaults to 100.
* ``backend_overrides``: raw SageMaker request fields for settings without a dedicated argument.
* Keys: request names from the *SageMaker API* section below.
* Values: request fields in PascalCase, as in the SageMaker API and boto3. Deep-merged over the request
built by AutoGluon-Cloud; lists and other non-dict values replace the generated ones.
* Example: ``{"CreateTrainingJob": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}``
Returns
-------
pd.Series | JobPredictionFuture
Predictions as a ``pd.Series`` if ``wait=True``; a :class:`JobPredictionFuture` otherwise.
SageMaker API
-------------
* :sm-api:`CreateTrainingJob`: runs the prediction as a training job (not a batch transform job) on
``instance_type``. Predictions are written to ``predictions_path``.
"""
result = self.predict_proba(
test_data,
train_data,
label=label,
include_predict=True,
predictions_path=predictions_path,
hyperparameters=hyperparameters,
instance_type=instance_type,
framework_version=framework_version,
custom_image_uri=custom_image_uri,
wait=wait,
**backend_kwargs,
)
if not wait:
return JobPredictionFuture(
job=self._backend._fit_job,
result_loader=lambda: self._load_results(include_predict=True, predict_only=True),
)
pred, _ = result
return pred
[docs]
@reject_legacy_kwargs
def predict_proba(
self,
test_data: str | Path | pd.DataFrame,
train_data: str | Path | pd.DataFrame,
label: str,
*,
include_predict: bool = True,
predictions_path: str | None = None,
hyperparameters: dict[str, Any] | None = None,
instance_type: str | None = None,
framework_version: str = DEFAULT_FRAMEWORK_VERSION,
custom_image_uri: str | None = None,
wait: bool = True,
**backend_kwargs,
) -> tuple[pd.Series, pd.DataFrame | pd.Series] | pd.DataFrame | pd.Series | JobPredictionFuture:
"""
Run batch prediction returning class probabilities.
For tabular foundation models (e.g., Mitra), ``train_data`` provides the few-shot context and
``test_data`` contains the rows to predict on. For regression the probabilities are identical to the
predictions.
Parameters
----------
test_data: str | Path | pd.DataFrame
Data to predict on. Must contain every feature column present in ``train_data`` except ``label``.
train_data: str | Path | pd.DataFrame
Labeled few-shot context for the foundation model, as a ``pd.DataFrame`` or local/S3 path to a data file.
label: str
Target column name in ``train_data``.
include_predict: bool, default = True
Whether to return the predictions along with the probabilities. Comes for free — the job always
computes both.
predictions_path: str | None, default = None
S3 URL where predictions will be written by the training container. Defaults to
``{cloud_output_path}/{job_name}/predictions.csv``.
hyperparameters: dict[str, Any] | None, default = None
Model hyperparameters for inference. Overrides values passed to the constructor.
instance_type: str | None, default = None
Instance type for the prediction job. If None, uses registry default.
framework_version: str, default = "1.6"
AutoGluon version, e.g. "1.6". Uses the official AutoGluon DLC image for this version.
custom_image_uri: str | None, default = None
Custom Docker image URI for the container.
wait: bool, default = True
If True, block and return the result. If False, return a :class:`JobPredictionFuture` immediately.
**backend_kwargs: Any
Additional SageMaker arguments:
* ``job_name``: Name of the training job that runs the prediction. Auto-generated if not set.
* ``volume_size``: Size in GB of the storage volume to use for the job. Defaults to 100.
* ``backend_overrides``: raw SageMaker request fields for settings without a dedicated argument.
* Keys: request names from the *SageMaker API* section below.
* Values: request fields in PascalCase, as in the SageMaker API and boto3. Deep-merged over the request
built by AutoGluon-Cloud; lists and other non-dict values replace the generated ones.
* Example: ``{"CreateTrainingJob": {"RetryStrategy": {"MaximumRetryAttempts": 2}}}``
Returns
-------
tuple[pd.Series, pd.DataFrame | pd.Series] | pd.DataFrame | pd.Series | JobPredictionFuture
If ``include_predict`` is True, returns ``(prediction, predict_probability)``; otherwise just
``predict_probability``. Returns a :class:`JobPredictionFuture` when ``wait=False``.
SageMaker API
-------------
* :sm-api:`CreateTrainingJob`: runs the prediction as a training job (not a batch transform job) on
``instance_type``. Predictions are written to ``predictions_path``.
"""
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 two tuning rows so AutoGluon does not hold out any rows from the prediction context. Two rather
# than one, since some models (e.g. Nori) return a 0-d array when predicting a single row.
tuning_data = train_data.iloc[:2].copy()
extra_ag_args: dict[str, Any] = {"predict_after_fit": True, "save_predictor": False}
if predictions_path is not None:
extra_ag_args["predictions_path"] = predictions_path
backend_kwargs["leaderboard"] = False
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, "tuning_data": tuning_data, "test_data": test_data},
framework_version=framework_version,
instance_type=instance_type,
custom_image_uri=custom_image_uri,
wait=wait,
extra_ag_args=extra_ag_args,
extra_tags=[{"Key": "autogluon-cloud-model-id", "Value": self.model_id}],
**backend_kwargs,
)
if not wait:
return JobPredictionFuture(
job=self._backend._fit_job,
result_loader=lambda: self._load_results(include_predict=include_predict),
)
return self._load_results(include_predict=include_predict)