Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions datamint/api/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion datamint/api/endpoints/annotations_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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',
Expand Down
143 changes: 143 additions & 0 deletions datamint/api/endpoints/model_types.py
Original file line number Diff line number Diff line change
@@ -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)
106 changes: 78 additions & 28 deletions datamint/api/endpoints/models_api.py
Original file line number Diff line number Diff line change
@@ -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)
2 changes: 1 addition & 1 deletion datamint/client_cmd_tools/datamint_upload.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
3 changes: 3 additions & 0 deletions datamint/mlflow/lightning/callbacks/modelcheckpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
9 changes: 9 additions & 0 deletions datamint/mlflow/models/tags.py
Original file line number Diff line number Diff line change
@@ -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.
"""
Loading