predict¶
- TabularFoundationModel.predict(test_data: str | Path | DataFrame, train_data: str | Path | DataFrame, label: str, *, predictions_path: str | None = None, hyperparameters: dict[str, Any] | None = None, instance_type: str | None = None, framework_version: str = '1.6', custom_image_uri: str | None = None, wait: bool = True, **backend_kwargs) Series | JobPredictionFuture[source]¶
Run batch prediction for tabular tasks.
For tabular foundation models (e.g., Mitra),
train_dataprovides the few-shot context andtest_datacontains the rows to predict on.- Parameters:
test_data (str | Path | pd.DataFrame) – Data to predict on. Must contain every feature column present in
train_dataexceptlabel.train_data (str | Path | pd.DataFrame) – Labeled few-shot context for the foundation model, as a
pd.DataFrameor 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
JobPredictionFutureimmediately — 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.Seriesifwait=True; aJobPredictionFutureotherwise.
SageMaker API
CreateTrainingJob: runs the prediction as a training job (not a batch transform job) on
instance_type. Predictions are written topredictions_path.