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
6 changes: 6 additions & 0 deletions datamint/api/endpoints/model_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down
2 changes: 2 additions & 0 deletions docs/source/client_api_content.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
++++++++++++++++++++++++++++++++++++++

Expand Down
25 changes: 25 additions & 0 deletions notebooks/05_deployment/01_deploy_registered_model.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
4 changes: 2 additions & 2 deletions notebooks/05_deployment/02_deploy_external_model.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Loading