Source code for autogluon.cloud.config

"""Backend settings and persistent resource identifiers for AutoGluon-Cloud.

Stores resource identifiers (region, stack name, bucket, IAM role ARN) at
``~/.autogluon/cloud.yaml`` so users don't need to re-specify them every
session. The file contains only non-secret identifiers — no AWS credentials
are ever written to disk.

The file is keyed by backend name::

    sagemaker:
      region: us-east-1
      role_arn: arn:aws:iam::...:role/ag-cloud-sagemaker-execution-role
      bucket: ag-cloud-sagemaker-bucket-...
      stack_name: ag-cloud-sagemaker
"""

from __future__ import annotations

import os
import stat
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import ClassVar, Dict, List, Optional

import yaml

CONFIG_DIR_ENV = "AG_CONFIG_DIR"


[docs] @dataclass(kw_only=True) class SageMakerConfig: """Reusable SageMaker settings for predictors and foundation models. Pass this as ``backend=`` to a cloud predictor or foundation model. Each object creates its own backend and jobs; sharing this config does not share execution state. Resource sizes and other operation settings remain named arguments to ``fit()``, ``predict()`` and ``deploy()``. Parameters ---------- region AWS region. If omitted, use the region in ``~/.autogluon/cloud.yaml``, then the boto3 default region. role_arn SageMaker execution role ARN. If omitted, use the saved role, then the role of the current AWS identity. vpc_config Networking for training jobs and models, as ``{"subnets": [...], "security_group_ids": [...]}``. output_kms_key KMS key for training artifacts, batch transform outputs, and repacked or cached model artifacts in S3. volume_kms_key KMS key for training, batch transform and realtime endpoint storage volumes. Leave unset for instance types with local NVMe storage. tags Tags added to every SageMaker resource created by this backend. """ name: ClassVar[str] = "sagemaker" region: Optional[str] = None role_arn: Optional[str] = None vpc_config: Optional[Dict[str, List[str]]] = None output_kms_key: Optional[str] = None volume_kms_key: Optional[str] = None tags: Dict[str, str] = field(default_factory=dict)
def get_config_dir() -> Path: override = os.environ.get(CONFIG_DIR_ENV) if override: return Path(override).expanduser() return Path.home() / ".autogluon" def get_config_path() -> Path: return get_config_dir() / "cloud.yaml" @dataclass class BackendConfig: """Persisted identifiers for a single AutoGluon-Cloud backend.""" region: str role_arn: str bucket: str stack_name: Optional[str] = None @dataclass class CloudConfig: """Top-level config: maps backend name → BackendConfig.""" backends: Dict[str, BackendConfig] = field(default_factory=dict) def load_config() -> Optional[CloudConfig]: """Load the config file, or return None if it doesn't exist or is empty.""" path = get_config_path() if not path.exists(): return None with path.open("r") as f: raw = yaml.safe_load(f) or {} if not raw: return None backends = {name: BackendConfig(**data) for name, data in raw.items()} return CloudConfig(backends=backends) def save_config(config: CloudConfig) -> Path: """Persist config atomically with 0600 file perms.""" path = get_config_path() path.parent.mkdir(parents=True, exist_ok=True) payload = {name: asdict(b) for name, b in config.backends.items()} tmp = path.with_suffix(".yaml.tmp") with tmp.open("w") as f: yaml.safe_dump(payload, f, sort_keys=False) os.chmod(tmp, stat.S_IRUSR | stat.S_IWUSR) os.replace(tmp, path) return path def delete_config() -> bool: """Remove the config file. Returns True if a file was deleted.""" path = get_config_path() if not path.exists(): return False path.unlink() return True