From f1c492b8baaf3b62e39bdf89fdb1dba82927eea8 Mon Sep 17 00:00:00 2001 From: luandalmazo Date: Tue, 18 Aug 2026 13:31:17 -0300 Subject: [PATCH] add load_model api method --- datamint/api/endpoints/model_types.py | 6 +++++ docs/source/client_api_content.rst | 2 ++ .../01_deploy_registered_model.ipynb | 25 +++++++++++++++++++ .../02_deploy_external_model.ipynb | 4 +-- 4 files changed, 35 insertions(+), 2 deletions(-) diff --git a/datamint/api/endpoints/model_types.py b/datamint/api/endpoints/model_types.py index cb43ea3a..48b3a623 100644 --- a/datamint/api/endpoints/model_types.py +++ b/datamint/api/endpoints/model_types.py @@ -13,6 +13,7 @@ if TYPE_CHECKING: from datamint.entities.project import Project + from datamint.mlflow.flavors.model import DatamintModel from .models_api import ModelsApi @@ -119,6 +120,11 @@ def get_metrics(self) -> dict[str, float]: def is_deployed(self) -> bool: return self._api._deploy_api.image_exists(self.name) + def load_model(self, device: str | None = None) -> 'DatamintModel': + """Load this model version via the ``datamint`` MLflow flavor.""" + from datamint.mlflow.flavors import load_model + return load_model(f"models:/{self.name}/{self.version}", device=device) + @dataclass class Model: diff --git a/docs/source/client_api_content.rst b/docs/source/client_api_content.rst index 495fb151..69aeff58 100644 --- a/docs/source/client_api_content.rst +++ b/docs/source/client_api_content.rst @@ -568,6 +568,8 @@ a Datamint :mod:`~datamint.lightning.trainers`), rather than raising. ``Model.get_supported_modes()``/``get_metrics()`` are shortcuts that delegate to the latest version when you don't need a specific one. +``latest.load_model()`` loads the model itself, ready for local inference. + Find which projects a model belongs to ++++++++++++++++++++++++++++++++++++++ diff --git a/notebooks/05_deployment/01_deploy_registered_model.ipynb b/notebooks/05_deployment/01_deploy_registered_model.ipynb index fd165f2f..9f79db69 100644 --- a/notebooks/05_deployment/01_deploy_registered_model.ipynb +++ b/notebooks/05_deployment/01_deploy_registered_model.ipynb @@ -87,6 +87,31 @@ " # If you need aliases, use ``get_model_version(mv.name, mv.version)`` method instead." ] }, + { + "cell_type": "markdown", + "id": "ca5f9fad", + "metadata": {}, + "source": [ + "## Loading a Registered Model Locally\n", + "\n", + "Before deploying, you can load a model version directly in Python via `ModelVersion.load_model()` (a shortcut over `datamint.mlflow.flavors.load_model`)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "cd55407f", + "metadata": {}, + "outputs": [], + "source": [ + "# Load a model version locally, without deploying it\n", + "model = api.models.get_by_name(all_registered_models[1].name)\n", + "latest = model.get_latest_version()\n", + "\n", + "loaded_model = latest.load_model()\n", + "print(loaded_model)" + ] + }, { "cell_type": "markdown", "id": "9da152da", diff --git a/notebooks/05_deployment/02_deploy_external_model.ipynb b/notebooks/05_deployment/02_deploy_external_model.ipynb index 5c8dc216..22002e32 100644 --- a/notebooks/05_deployment/02_deploy_external_model.ipynb +++ b/notebooks/05_deployment/02_deploy_external_model.ipynb @@ -424,8 +424,8 @@ "source": [ "from datamint.mlflow import flavors as datamint_flavor\n", "\n", - "model_uri = f\"models:/{MODEL_NAME}@champion\" # or f\"models:/{MODEL_NAME}/latest\" or f\"models:/{MODEL_NAME}/1\"\n", - "loaded_model = datamint_flavor.load_model(model_uri)\n", + "model = api.models.get_by_name(MODEL_NAME)\n", + "loaded_model = model.get_latest_version(alias='champion').load_model()\n", "\n", "print(f\"Loaded model type : {type(loaded_model).__name__}\")\n", "print(f\"Task type : {loaded_model.task_type}\")\n",