Skip to content

Commit 6551d66

Browse files
committed
add mlflow facade
1 parent 3285b6a commit 6551d66

7 files changed

Lines changed: 237 additions & 30 deletions

File tree

datamint/api/client.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -150,6 +150,8 @@ def _get_endpoint(self, name: str, is_mlflow: bool = False):
150150
kwargs: dict[str, Any] = {}
151151
if name == 'inference':
152152
kwargs['projects_api'] = self.projects
153+
elif name == 'models':
154+
kwargs['deploy_api'] = self.deploy
153155
endpoint = api_class(self.config, client=client, **kwargs)
154156
# Inject this API instance into the endpoint so it can inject into entities
155157
endpoint._api_instance = self

datamint/api/endpoints/annotations_api.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -586,7 +586,7 @@ def _check_model(self, ai_model_name: str | None) -> str | None:
586586
modelinfo = self._models_api.get_by_name(ai_model_name)
587587
if modelinfo is None:
588588
try:
589-
available_models = [model['name'] for model in self._models_api.get_all()]
589+
available_models = [model.name for model in self._models_api.get_all()]
590590
except Exception:
591591
_LOGGER.warning("Could not fetch available AI models from the server.")
592592
raise ItemNotFoundError('ai-model',
Lines changed: 143 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,143 @@
1+
"""Thin wrapper objects over MLflow's model registry entities."""
2+
from dataclasses import dataclass
3+
from typing import TYPE_CHECKING
4+
5+
import mlflow.models
6+
from mlflow.entities.model_registry import ModelVersion as MlflowModelVersion
7+
from mlflow.entities.model_registry import RegisteredModel as MlflowRegisteredModel
8+
9+
from datamint.entities.annotations.annotation_spec import AnnotationSpec
10+
from datamint.mlflow.flavors.datamint_flavor import FLAVOR_NAME
11+
from datamint.mlflow.models.tags import DATAMINT_LOGGED_MODEL_ID_TAG
12+
13+
if TYPE_CHECKING:
14+
from .models_api import ModelsApi
15+
16+
17+
@dataclass
18+
class ModelVersion:
19+
"""A single version of a registered model."""
20+
21+
_raw: MlflowModelVersion
22+
_api: 'ModelsApi'
23+
24+
@property
25+
def name(self) -> str:
26+
return self._raw.name
27+
28+
@property
29+
def version(self) -> str:
30+
return self._raw.version
31+
32+
@property
33+
def run_id(self) -> str | None:
34+
return self._raw.run_id
35+
36+
@property
37+
def creation_timestamp(self) -> int:
38+
return self._raw.creation_timestamp
39+
40+
@property
41+
def source(self) -> str | None:
42+
return self._raw.source
43+
44+
@property
45+
def current_stage(self) -> str | None:
46+
return self._raw.current_stage
47+
48+
@property
49+
def aliases(self) -> list[str]:
50+
return self._raw.aliases
51+
52+
@property
53+
def tags(self) -> dict[str, str]:
54+
return self._raw.tags
55+
56+
def _flavor_data(self) -> dict:
57+
if self.source is None:
58+
return {}
59+
model_info = mlflow.models.get_model_info(self.source)
60+
return model_info.flavors.get(FLAVOR_NAME, {})
61+
62+
def get_supported_modes(self) -> list[str]:
63+
"""Prediction modes this model version supports (from the ``datamint`` flavor)."""
64+
return self._flavor_data().get('supported_modes', [])
65+
66+
def get_task_type(self) -> str | None:
67+
"""Task type this model version was trained for, or ``None`` if not recorded."""
68+
return self._flavor_data().get('task_type')
69+
70+
def get_annotation_specs(self) -> list[AnnotationSpec] | None:
71+
"""Annotation specs this model version produces, or ``None`` if not recorded."""
72+
raw_specs = self._flavor_data().get('annotation_specs')
73+
if not raw_specs:
74+
return None
75+
return [AnnotationSpec.create(**s) for s in raw_specs]
76+
77+
def get_metrics(self) -> dict[str, float]:
78+
"""Training/test metrics logged for this version.
79+
80+
Returns ``{}`` when this version has no ``DATAMINT_LOGGED_MODEL_ID_TAG``
81+
(e.g. an externally-registered model with no Datamint-trained run behind it).
82+
"""
83+
logged_model_id = self.tags.get(DATAMINT_LOGGED_MODEL_ID_TAG)
84+
if logged_model_id is None:
85+
return {}
86+
logged_model = mlflow.get_logged_model(logged_model_id)
87+
return {m.key: m.value for m in logged_model.metrics}
88+
89+
def is_deployed(self) -> bool:
90+
return self._api._deploy_api.image_exists(self.name)
91+
92+
93+
@dataclass
94+
class Model:
95+
"""A registered model: a named family of :class:`ModelVersion`."""
96+
97+
_raw: MlflowRegisteredModel
98+
_api: 'ModelsApi'
99+
100+
@property
101+
def name(self) -> str:
102+
return self._raw.name
103+
104+
@property
105+
def description(self) -> str | None:
106+
return self._raw.description
107+
108+
@property
109+
def creation_timestamp(self) -> int:
110+
return self._raw.creation_timestamp
111+
112+
@property
113+
def last_updated_timestamp(self) -> int:
114+
return self._raw.last_updated_timestamp
115+
116+
@property
117+
def tags(self) -> dict[str, str]:
118+
return self._raw.tags
119+
120+
def get_versions(self) -> list[ModelVersion]:
121+
raw_versions = self._api._mlflow_client.search_model_versions(f"name='{self.name}'")
122+
return [ModelVersion(_raw=v, _api=self._api) for v in raw_versions]
123+
124+
def get_latest_version(self, alias: str | None = None) -> ModelVersion | None:
125+
"""Most recently created version, or the version at *alias* if given."""
126+
if alias is not None:
127+
raw = self._api._mlflow_client.get_model_version_by_alias(self.name, alias)
128+
return ModelVersion(_raw=raw, _api=self._api) if raw else None
129+
versions = self.get_versions()
130+
if not versions:
131+
return None
132+
return max(versions, key=lambda v: int(v.version))
133+
134+
def get_supported_modes(self, version: ModelVersion | None = None) -> list[str]:
135+
version = version or self.get_latest_version()
136+
return version.get_supported_modes() if version else []
137+
138+
def get_metrics(self, version: ModelVersion | None = None) -> dict[str, float]:
139+
version = version or self.get_latest_version()
140+
return version.get_metrics() if version else {}
141+
142+
def is_deployed(self) -> bool:
143+
return self._api._deploy_api.image_exists(self.name)
Lines changed: 78 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -1,37 +1,87 @@
1-
"""Deprecated: Use MLFlow API instead."""
2-
from collections.abc import Sequence
3-
from ..entity_base_api import BaseApi
1+
"""API handler for the model registry, backed by MLflow."""
42
import httpx
5-
from datamint.exceptions import EntityAlreadyExistsError
3+
import mlflow.exceptions
4+
import mlflow.tracking
5+
6+
from ..entity_base_api import ApiConfig, BaseApi
7+
from .deploy_model_api import DeployModelApi
8+
from .model_types import Model
69

710

811
class ModelsApi(BaseApi):
9-
"""API handler for project-related endpoints."""
12+
"""API handler for the model registry.
13+
14+
Wraps MLflow's model registry (registered models / model versions) behind
15+
plain Python objects (:class:`~.model_types.Model`, :class:`~.model_types.ModelVersion`)
16+
so callers never need to know MLflow's object model.
17+
"""
18+
19+
def __init__(self,
20+
config: ApiConfig,
21+
client: httpx.Client | None = None,
22+
deploy_api: DeployModelApi | None = None) -> None:
23+
super().__init__(config, client)
24+
self._deploy_api = deploy_api or DeployModelApi(config, client=client)
25+
26+
@property
27+
def _mlflow_client(self) -> mlflow.tracking.MlflowClient:
28+
import datamint.mlflow
29+
return mlflow.tracking.MlflowClient()
30+
31+
def get_list(self,
32+
only_deployed: bool = False,
33+
max_results: int | None = None) -> list[Model]:
34+
"""List registered models.
1035
11-
def create(self,
12-
name: str) -> dict:
13-
json = {
14-
'name': name
15-
}
36+
Args:
37+
only_deployed: If ``True``, only return models with a deployed image.
38+
max_results: Maximum number of models to return. If ``None``, all
39+
registered models are returned (paginating through the registry).
40+
"""
41+
if max_results is not None:
42+
raw_models = list(self._mlflow_client.search_registered_models(max_results=max_results))
43+
else:
44+
raw_models = []
45+
page_token = None
46+
while True:
47+
page = self._mlflow_client.search_registered_models(page_token=page_token)
48+
raw_models.extend(page)
49+
page_token = page.token
50+
if not page_token:
51+
break
1652

53+
models = [Model(_raw=m, _api=self) for m in raw_models]
54+
if only_deployed:
55+
models = [m for m in models if m.is_deployed()]
56+
return models
57+
58+
def get_all(self, only_deployed: bool = False, max_results: int | None = None) -> list[Model]:
59+
"""Alias for :meth:`get_list`, kept for backwards compatibility with existing call sites."""
60+
return self.get_list(only_deployed=only_deployed, max_results=max_results)
61+
62+
def get_by_name(self, name: str) -> Model | None:
63+
"""Get a registered model by name, or ``None`` if it does not exist."""
1764
try:
18-
response = self._make_request('POST',
19-
'ai-models',
20-
json=json)
21-
return response.json()
22-
except httpx.HTTPStatusError as e:
23-
if e.response.status_code == 409:
24-
raise EntityAlreadyExistsError('ai-model', {'name': name})
65+
raw_model = self._mlflow_client.get_registered_model(name)
66+
except mlflow.exceptions.MlflowException as e:
67+
if e.error_code == 'RESOURCE_DOES_NOT_EXIST':
68+
return None
2569
raise
70+
return Model(_raw=raw_model, _api=self)
71+
72+
def create(self, name: str, description: str | None = None, exists_ok: bool = True) -> Model:
73+
"""Create a new registered model.
2674
27-
def get_all(self) -> Sequence[dict]:
28-
response = self._make_request('GET',
29-
'ai-models')
30-
return response.json()
31-
32-
def get_by_name(self, name: str) -> dict | None:
33-
models = self.get_all()
34-
for model in models:
35-
if model['name'] == name:
36-
return model
37-
return None
75+
Args:
76+
name: Name of the model to register.
77+
description: Optional description.
78+
exists_ok: If ``True`` (default), return the existing model instead of
79+
raising when a model with this name already exists.
80+
"""
81+
try:
82+
raw_model = self._mlflow_client.create_registered_model(name, description=description)
83+
except mlflow.exceptions.MlflowException as e:
84+
if exists_ok and e.error_code == 'RESOURCE_ALREADY_EXISTS':
85+
return self.get_by_name(name)
86+
raise
87+
return Model(_raw=raw_model, _api=self)

datamint/client_cmd_tools/datamint_upload.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -794,7 +794,7 @@ def main():
794794
model_info = api.models.get_by_name(args.ai_model)
795795
if model_info is None:
796796
available_models = api.models.get_all()
797-
model_names = [model['name'] for model in available_models]
797+
model_names = [model.name for model in available_models]
798798
_USER_LOGGER.error(f'❌ AI model "{args.ai_model}" not found. Available models: {model_names}')
799799
return
800800
print_input_summary(files_path,

datamint/mlflow/lightning/callbacks/modelcheckpoint.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from torch import nn
1010
import lightning.pytorch as L
1111
from datamint.mlflow.models import log_model_metadata, _get_MLFlowLogger
12+
from datamint.mlflow.models.tags import DATAMINT_LOGGED_MODEL_ID_TAG
1213
from datamint.mlflow.env_utils import ensure_mlflow_configured
1314
import mlflow.models
1415
import mlflow.exceptions
@@ -309,9 +310,11 @@ def register_model(self, trainer=None):
309310
return self.registered_model_info
310311

311312
# mlflow_client = _get_MLFlowLogger(trainer)._mlflow_client
313+
tags = {DATAMINT_LOGGED_MODEL_ID_TAG: self._last_model_id} if self._last_model_id else None
312314
self.registered_model_info = mlflow.register_model(
313315
model_uri=self._last_model_uri,
314316
name=self.model_name,
317+
tags=tags,
315318
)
316319

317320
# Update the registered state hash after successful registration

datamint/mlflow/models/tags.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,9 @@
1+
"""Shared MLflow tag keys used by Datamint's model registry integration."""
2+
3+
DATAMINT_LOGGED_MODEL_ID_TAG = 'datamint.logged_model_id'
4+
"""ModelVersion tag holding the LoggedModel.model_id it was registered from.
5+
6+
Stamped at registration time in ``_BaseMLFlowModelCheckpoint.register_model()``,
7+
read back by ``ModelVersion.get_metrics()`` to find the LoggedModel that carries
8+
the training/test metrics for this version.
9+
"""

0 commit comments

Comments
 (0)