From 6551d66c42aed8a044bc605cb983f7cedcb95ffa Mon Sep 17 00:00:00 2001 From: luandalmazo Date: Thu, 23 Jul 2026 13:39:45 -0300 Subject: [PATCH] add mlflow facade --- datamint/api/client.py | 2 + datamint/api/endpoints/annotations_api.py | 2 +- datamint/api/endpoints/model_types.py | 143 ++++++++++++++++++ datamint/api/endpoints/models_api.py | 106 +++++++++---- datamint/client_cmd_tools/datamint_upload.py | 2 +- .../lightning/callbacks/modelcheckpoint.py | 3 + datamint/mlflow/models/tags.py | 9 ++ 7 files changed, 237 insertions(+), 30 deletions(-) create mode 100644 datamint/api/endpoints/model_types.py create mode 100644 datamint/mlflow/models/tags.py diff --git a/datamint/api/client.py b/datamint/api/client.py index 8a84f2b6..cfd17abd 100644 --- a/datamint/api/client.py +++ b/datamint/api/client.py @@ -150,6 +150,8 @@ def _get_endpoint(self, name: str, is_mlflow: bool = False): kwargs: dict[str, Any] = {} if name == 'inference': kwargs['projects_api'] = self.projects + elif name == 'models': + kwargs['deploy_api'] = self.deploy endpoint = api_class(self.config, client=client, **kwargs) # Inject this API instance into the endpoint so it can inject into entities endpoint._api_instance = self diff --git a/datamint/api/endpoints/annotations_api.py b/datamint/api/endpoints/annotations_api.py index 669ac7ba..9671bfb2 100644 --- a/datamint/api/endpoints/annotations_api.py +++ b/datamint/api/endpoints/annotations_api.py @@ -586,7 +586,7 @@ def _check_model(self, ai_model_name: str | None) -> str | None: modelinfo = self._models_api.get_by_name(ai_model_name) if modelinfo is None: try: - available_models = [model['name'] for model in self._models_api.get_all()] + available_models = [model.name for model in self._models_api.get_all()] except Exception: _LOGGER.warning("Could not fetch available AI models from the server.") raise ItemNotFoundError('ai-model', diff --git a/datamint/api/endpoints/model_types.py b/datamint/api/endpoints/model_types.py new file mode 100644 index 00000000..99bd7486 --- /dev/null +++ b/datamint/api/endpoints/model_types.py @@ -0,0 +1,143 @@ +"""Thin wrapper objects over MLflow's model registry entities.""" +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import mlflow.models +from mlflow.entities.model_registry import ModelVersion as MlflowModelVersion +from mlflow.entities.model_registry import RegisteredModel as MlflowRegisteredModel + +from datamint.entities.annotations.annotation_spec import AnnotationSpec +from datamint.mlflow.flavors.datamint_flavor import FLAVOR_NAME +from datamint.mlflow.models.tags import DATAMINT_LOGGED_MODEL_ID_TAG + +if TYPE_CHECKING: + from .models_api import ModelsApi + + +@dataclass +class ModelVersion: + """A single version of a registered model.""" + + _raw: MlflowModelVersion + _api: 'ModelsApi' + + @property + def name(self) -> str: + return self._raw.name + + @property + def version(self) -> str: + return self._raw.version + + @property + def run_id(self) -> str | None: + return self._raw.run_id + + @property + def creation_timestamp(self) -> int: + return self._raw.creation_timestamp + + @property + def source(self) -> str | None: + return self._raw.source + + @property + def current_stage(self) -> str | None: + return self._raw.current_stage + + @property + def aliases(self) -> list[str]: + return self._raw.aliases + + @property + def tags(self) -> dict[str, str]: + return self._raw.tags + + def _flavor_data(self) -> dict: + if self.source is None: + return {} + model_info = mlflow.models.get_model_info(self.source) + return model_info.flavors.get(FLAVOR_NAME, {}) + + def get_supported_modes(self) -> list[str]: + """Prediction modes this model version supports (from the ``datamint`` flavor).""" + return self._flavor_data().get('supported_modes', []) + + def get_task_type(self) -> str | None: + """Task type this model version was trained for, or ``None`` if not recorded.""" + return self._flavor_data().get('task_type') + + def get_annotation_specs(self) -> list[AnnotationSpec] | None: + """Annotation specs this model version produces, or ``None`` if not recorded.""" + raw_specs = self._flavor_data().get('annotation_specs') + if not raw_specs: + return None + return [AnnotationSpec.create(**s) for s in raw_specs] + + def get_metrics(self) -> dict[str, float]: + """Training/test metrics logged for this version. + + Returns ``{}`` when this version has no ``DATAMINT_LOGGED_MODEL_ID_TAG`` + (e.g. an externally-registered model with no Datamint-trained run behind it). + """ + logged_model_id = self.tags.get(DATAMINT_LOGGED_MODEL_ID_TAG) + if logged_model_id is None: + return {} + logged_model = mlflow.get_logged_model(logged_model_id) + return {m.key: m.value for m in logged_model.metrics} + + def is_deployed(self) -> bool: + return self._api._deploy_api.image_exists(self.name) + + +@dataclass +class Model: + """A registered model: a named family of :class:`ModelVersion`.""" + + _raw: MlflowRegisteredModel + _api: 'ModelsApi' + + @property + def name(self) -> str: + return self._raw.name + + @property + def description(self) -> str | None: + return self._raw.description + + @property + def creation_timestamp(self) -> int: + return self._raw.creation_timestamp + + @property + def last_updated_timestamp(self) -> int: + return self._raw.last_updated_timestamp + + @property + def tags(self) -> dict[str, str]: + return self._raw.tags + + def get_versions(self) -> list[ModelVersion]: + raw_versions = self._api._mlflow_client.search_model_versions(f"name='{self.name}'") + return [ModelVersion(_raw=v, _api=self._api) for v in raw_versions] + + def get_latest_version(self, alias: str | None = None) -> ModelVersion | None: + """Most recently created version, or the version at *alias* if given.""" + if alias is not None: + raw = self._api._mlflow_client.get_model_version_by_alias(self.name, alias) + return ModelVersion(_raw=raw, _api=self._api) if raw else None + versions = self.get_versions() + if not versions: + return None + return max(versions, key=lambda v: int(v.version)) + + def get_supported_modes(self, version: ModelVersion | None = None) -> list[str]: + version = version or self.get_latest_version() + return version.get_supported_modes() if version else [] + + def get_metrics(self, version: ModelVersion | None = None) -> dict[str, float]: + version = version or self.get_latest_version() + return version.get_metrics() if version else {} + + def is_deployed(self) -> bool: + return self._api._deploy_api.image_exists(self.name) diff --git a/datamint/api/endpoints/models_api.py b/datamint/api/endpoints/models_api.py index 01a18c9e..b07aeb31 100644 --- a/datamint/api/endpoints/models_api.py +++ b/datamint/api/endpoints/models_api.py @@ -1,37 +1,87 @@ -"""Deprecated: Use MLFlow API instead.""" -from collections.abc import Sequence -from ..entity_base_api import BaseApi +"""API handler for the model registry, backed by MLflow.""" import httpx -from datamint.exceptions import EntityAlreadyExistsError +import mlflow.exceptions +import mlflow.tracking + +from ..entity_base_api import ApiConfig, BaseApi +from .deploy_model_api import DeployModelApi +from .model_types import Model class ModelsApi(BaseApi): - """API handler for project-related endpoints.""" + """API handler for the model registry. + + Wraps MLflow's model registry (registered models / model versions) behind + plain Python objects (:class:`~.model_types.Model`, :class:`~.model_types.ModelVersion`) + so callers never need to know MLflow's object model. + """ + + def __init__(self, + config: ApiConfig, + client: httpx.Client | None = None, + deploy_api: DeployModelApi | None = None) -> None: + super().__init__(config, client) + self._deploy_api = deploy_api or DeployModelApi(config, client=client) + + @property + def _mlflow_client(self) -> mlflow.tracking.MlflowClient: + import datamint.mlflow + return mlflow.tracking.MlflowClient() + + def get_list(self, + only_deployed: bool = False, + max_results: int | None = None) -> list[Model]: + """List registered models. - def create(self, - name: str) -> dict: - json = { - 'name': name - } + Args: + only_deployed: If ``True``, only return models with a deployed image. + max_results: Maximum number of models to return. If ``None``, all + registered models are returned (paginating through the registry). + """ + if max_results is not None: + raw_models = list(self._mlflow_client.search_registered_models(max_results=max_results)) + else: + raw_models = [] + page_token = None + while True: + page = self._mlflow_client.search_registered_models(page_token=page_token) + raw_models.extend(page) + page_token = page.token + if not page_token: + break + models = [Model(_raw=m, _api=self) for m in raw_models] + if only_deployed: + models = [m for m in models if m.is_deployed()] + return models + + def get_all(self, only_deployed: bool = False, max_results: int | None = None) -> list[Model]: + """Alias for :meth:`get_list`, kept for backwards compatibility with existing call sites.""" + return self.get_list(only_deployed=only_deployed, max_results=max_results) + + def get_by_name(self, name: str) -> Model | None: + """Get a registered model by name, or ``None`` if it does not exist.""" try: - response = self._make_request('POST', - 'ai-models', - json=json) - return response.json() - except httpx.HTTPStatusError as e: - if e.response.status_code == 409: - raise EntityAlreadyExistsError('ai-model', {'name': name}) + raw_model = self._mlflow_client.get_registered_model(name) + except mlflow.exceptions.MlflowException as e: + if e.error_code == 'RESOURCE_DOES_NOT_EXIST': + return None raise + return Model(_raw=raw_model, _api=self) + + def create(self, name: str, description: str | None = None, exists_ok: bool = True) -> Model: + """Create a new registered model. - def get_all(self) -> Sequence[dict]: - response = self._make_request('GET', - 'ai-models') - return response.json() - - def get_by_name(self, name: str) -> dict | None: - models = self.get_all() - for model in models: - if model['name'] == name: - return model - return None + Args: + name: Name of the model to register. + description: Optional description. + exists_ok: If ``True`` (default), return the existing model instead of + raising when a model with this name already exists. + """ + try: + raw_model = self._mlflow_client.create_registered_model(name, description=description) + except mlflow.exceptions.MlflowException as e: + if exists_ok and e.error_code == 'RESOURCE_ALREADY_EXISTS': + return self.get_by_name(name) + raise + return Model(_raw=raw_model, _api=self) diff --git a/datamint/client_cmd_tools/datamint_upload.py b/datamint/client_cmd_tools/datamint_upload.py index 779bcddb..b79b862c 100644 --- a/datamint/client_cmd_tools/datamint_upload.py +++ b/datamint/client_cmd_tools/datamint_upload.py @@ -794,7 +794,7 @@ def main(): model_info = api.models.get_by_name(args.ai_model) if model_info is None: available_models = api.models.get_all() - model_names = [model['name'] for model in available_models] + model_names = [model.name for model in available_models] _USER_LOGGER.error(f'❌ AI model "{args.ai_model}" not found. Available models: {model_names}') return print_input_summary(files_path, diff --git a/datamint/mlflow/lightning/callbacks/modelcheckpoint.py b/datamint/mlflow/lightning/callbacks/modelcheckpoint.py index 2cb641e1..64bcb0b0 100644 --- a/datamint/mlflow/lightning/callbacks/modelcheckpoint.py +++ b/datamint/mlflow/lightning/callbacks/modelcheckpoint.py @@ -9,6 +9,7 @@ from torch import nn import lightning.pytorch as L from datamint.mlflow.models import log_model_metadata, _get_MLFlowLogger +from datamint.mlflow.models.tags import DATAMINT_LOGGED_MODEL_ID_TAG from datamint.mlflow.env_utils import ensure_mlflow_configured import mlflow.models import mlflow.exceptions @@ -309,9 +310,11 @@ def register_model(self, trainer=None): return self.registered_model_info # mlflow_client = _get_MLFlowLogger(trainer)._mlflow_client + tags = {DATAMINT_LOGGED_MODEL_ID_TAG: self._last_model_id} if self._last_model_id else None self.registered_model_info = mlflow.register_model( model_uri=self._last_model_uri, name=self.model_name, + tags=tags, ) # Update the registered state hash after successful registration diff --git a/datamint/mlflow/models/tags.py b/datamint/mlflow/models/tags.py new file mode 100644 index 00000000..0fe1cf34 --- /dev/null +++ b/datamint/mlflow/models/tags.py @@ -0,0 +1,9 @@ +"""Shared MLflow tag keys used by Datamint's model registry integration.""" + +DATAMINT_LOGGED_MODEL_ID_TAG = 'datamint.logged_model_id' +"""ModelVersion tag holding the LoggedModel.model_id it was registered from. + +Stamped at registration time in ``_BaseMLFlowModelCheckpoint.register_model()``, +read back by ``ModelVersion.get_metrics()`` to find the LoggedModel that carries +the training/test metrics for this version. +"""