fit_predict¶
- TabularCloudPredictor.fit_predict(train_data: str | Path | DataFrame, test_data: str | Path | DataFrame, *, predictor_init_args: dict[str, Any], predictor_fit_args: dict[str, Any] | None = None, leaderboard: bool = True, framework_version: str = '1.6', job_name: str | None = None, instance_type: str = 'ml.m5.2xlarge', instance_count: int = 1, volume_size: int = 100, custom_image_uri: str | None = None, wait: bool = True, predictions_path: str | None = None, backend_overrides: dict[str, dict[str, Any]] | None = None) Series | None[source]¶
Fit and predict in a single SageMaker training job.
Fits a
TabularPredictorontrain_dataand runs batch prediction ontest_datainside the same training container. This avoids the overhead of a separate batch-transform job (one cold start, one data upload, no predictor-tarball round-trip). The predictor is left fitted afterward, sodeploy()/predict()still work.- Parameters:
train_data (str | pathlib.Path | pd.DataFrame) – Training data, as a
pd.DataFrameor local/S3 path to a data file.test_data (str | pathlib.Path | pd.DataFrame) – Data to predict on, as a
pd.DataFrameor local/S3 path to a data file. Must contain every feature column present intrain_data(the label column is not required).predictor_init_args (dict) – Init args for the predictor.
predictor_fit_args (dict | None, default = None) – Additional fit args forwarded to
TabularPredictor.fit(). Must NOT containtrain_dataortuning_data.leaderboard (bool, default = True) – Whether to include the leaderboard in the output artifact.
framework_version (str, optional) – AutoGluon version, e.g. “1.6”. Training uses the official AutoGluon DLC image for this version. If custom_image_uri is set, this argument will be ignored.
job_name (str, default = None) – Name of the launched training job. If None, CloudPredictor creates one with prefix
ag-cloud-tabular.instance_type (str, default = 'ml.m5.2xlarge') – Instance type the predictor will be trained on with SageMaker.
instance_count (int, default = 1) – Number of instances used to fit the predictor.
volume_size (int, default = 100) – Size in GB of the EBS volume to use for storing input data during training.
custom_image_uri (str | None, default = None) – Custom container image URI. If set,
framework_versionis ignored.wait (bool, default = True) – Whether the call should wait until the job completes.
predictions_path (str | 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.backend_overrides (dict[str, dict[str, Any]] | None, default = None) –
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 | None – Predictions as a
pd.Series. ReturnsNonewhenwaitis False; fetch later viaget_fit_predict_results().
SageMaker API
CreateTrainingJob: trains the predictor and predicts in the same job on
instance_countxinstance_type. Predictions are written topredictions_path.