|
| 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) |
0 commit comments