This commit is contained in:
Christian Mantha
2026-03-02 19:10:52 -05:00
commit 2ca0b9ef7c
28907 changed files with 5233713 additions and 0 deletions

View File

@@ -0,0 +1,10 @@
# Path to default location for backend when using local FileStore.
DEFAULT_LOCAL_FILE_AND_ARTIFACT_PATH = "./mlruns"
SEARCH_REGISTERED_MODEL_MAX_RESULTS_DEFAULT = 100
SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD = 1000
# Some backends have a low maximum results threshold; for example, Databricks only allows
# `max_results` request parameter values up to 10,000. Accordingly, **be very careful** when
# increasing this default maximum results value to avoid breaking compatibility with such backends
SEARCH_MODEL_VERSION_MAX_RESULTS_DEFAULT = 10000
SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD = 200_000

View File

@@ -0,0 +1,454 @@
import logging
from abc import ABCMeta, abstractmethod
from time import sleep, time
from mlflow.entities.model_registry import ModelVersionTag
from mlflow.entities.model_registry.model_version_status import ModelVersionStatus
from mlflow.exceptions import MlflowException
from mlflow.prompt.registry_utils import has_prompt_tag
from mlflow.protos.databricks_pb2 import RESOURCE_ALREADY_EXISTS, ErrorCode
from mlflow.utils.annotations import developer_stable
from mlflow.utils.logging_utils import eprint
_logger = logging.getLogger(__name__)
AWAIT_MODEL_VERSION_CREATE_SLEEP_INTERVAL_SECONDS = 3
@developer_stable
class AbstractStore:
"""
Abstract class that defines API interfaces for storing Model Registry metadata.
"""
__metaclass__ = ABCMeta
def __init__(self, store_uri=None, tracking_uri=None):
"""
Empty constructor. This is deliberately not marked as abstract, else every derived class
would be forced to create one.
Args:
store_uri: The model registry store URI.
tracking_uri: URI of the current MLflow tracking server, used to perform operations
like fetching source run metadata or downloading source run artifacts
to support subsequently uploading them to the model registry storage
location.
"""
# CRUD API for RegisteredModel objects
@abstractmethod
def create_registered_model(self, name, tags=None, description=None):
"""
Create a new registered model in backend store.
Args:
name: Name of the new model. This is expected to be unique in the backend store.
tags: A list of :py:class:`mlflow.entities.model_registry.RegisteredModelTag`
instances associated with this registered model.
description: Description of the model.
Returns:
A single object of :py:class:`mlflow.entities.model_registry.RegisteredModel`
created in the backend.
"""
@abstractmethod
def update_registered_model(self, name, description):
"""
Update description of the registered model.
Args:
name: Registered model name.
description: New description.
Returns:
A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
"""
@abstractmethod
def rename_registered_model(self, name, new_name):
"""
Rename the registered model.
Args:
name: Registered model name.
new_name: New proposed name.
Returns:
A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
"""
@abstractmethod
def delete_registered_model(self, name):
"""
Delete the registered model.
Backend raises exception if a registered model with given name does not exist.
Args:
name: Registered model name.
Returns:
None
"""
@abstractmethod
def search_registered_models(
self, filter_string=None, max_results=None, order_by=None, page_token=None
):
"""
Search for registered models in backend that satisfy the filter criteria.
Args:
filter_string: Filter query string, defaults to searching all registered models.
max_results: Maximum number of registered models desired.
order_by: List of column names with ASC|DESC annotation, to be used for ordering
matching search results.
page_token: Token specifying the next page of results. It should be obtained from
a ``search_registered_models`` call.
Returns:
A PagedList of :py:class:`mlflow.entities.model_registry.RegisteredModel` objects
that satisfy the search expressions. The pagination token for the next page can be
obtained via the ``token`` attribute of the object.
"""
@abstractmethod
def get_registered_model(self, name):
"""
Get registered model instance by name.
Args:
name: Registered model name.
Returns:
A single :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
"""
@abstractmethod
def get_latest_versions(self, name, stages=None):
"""
Latest version models for each requested stage. If no ``stages`` argument is provided,
returns the latest version for each stage.
Args:
name: Registered model name.
stages: List of desired stages. If input list is None, return latest versions for
each stage.
Returns:
List of :py:class:`mlflow.entities.model_registry.ModelVersion` objects.
"""
@abstractmethod
def set_registered_model_tag(self, name, tag):
"""
Set a tag for the registered model.
Args:
name: Registered model name.
tag: :py:class:`mlflow.entities.model_registry.RegisteredModelTag` instance to log.
Returns:
None
"""
@abstractmethod
def delete_registered_model_tag(self, name, key):
"""
Delete a tag associated with the registered model.
Args:
name: Registered model name.
key: Registered model tag key.
Returns:
None
"""
# CRUD API for ModelVersion objects
@abstractmethod
def create_model_version(
self,
name,
source,
run_id=None,
tags=None,
run_link=None,
description=None,
local_model_path=None,
):
"""
Create a new model version from given source and run ID.
Args:
name: Registered model name.
source: URI indicating the location of the model artifacts.
run_id: Run ID from MLflow tracking server that generated the model.
tags: A list of :py:class:`mlflow.entities.model_registry.ModelVersionTag`
instances associated with this model version.
run_link: Link to the run from an MLflow tracking server that generated this model.
description: Description of the version.
local_model_path: Local path to the MLflow model, if it's already accessible
on the local filesystem. Can be used by AbstractStores that
upload model version files to the model registry to avoid
a redundant download from the source location when logging
and registering a model via a single
mlflow.<flavor>.log_model(..., registered_model_name) call
Returns:
A single object of :py:class:`mlflow.entities.model_registry.ModelVersion`
created in the backend.
"""
@abstractmethod
def update_model_version(self, name, version, description):
"""
Update metadata associated with a model version in backend.
Args:
name: Registered model name.
version: Registered model version.
description: New model description.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
@abstractmethod
def transition_model_version_stage(self, name, version, stage, archive_existing_versions):
"""
Update model version stage.
Args:
name: Registered model name.
version: Registered model version.
stage: New desired stage for this model version.
archive_existing_versions: If this flag is set to ``True``, all existing model
versions in the stage will be automatically moved to the "archived" stage. Only
valid when ``stage`` is ``"staging"`` or ``"production"`` otherwise an error will
be raised.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
@abstractmethod
def delete_model_version(self, name, version):
"""
Delete model model version in backend.
Args:
name: Registered model name.
version: Registered model version.
Returns:
None
"""
@abstractmethod
def get_model_version(self, name, version):
"""
Get the model version instance by name and version.
Args:
name: Registered model name.
version: Registered model version.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
@abstractmethod
def get_model_version_download_uri(self, name, version):
"""
Get the download location in Model Registry for this model version.
NOTE: For first version of Model Registry, since the models are not copied over to another
location, download URI points to input source path.
Args:
name: Registered model name.
version: Registered model version.
Returns:
A single URI location that allows reads for downloading.
"""
@abstractmethod
def search_model_versions(
self, filter_string=None, max_results=None, order_by=None, page_token=None
):
"""
Search for model versions in backend that satisfy the filter criteria.
Args:
filter_string: A filter string expression. Currently supports a single filter
condition either name of model like ``name = 'model_name'`` or
``run_id = '...'``.
max_results: Maximum number of model versions desired.
order_by: List of column names with ASC|DESC annotation, to be used for ordering
matching search results.
page_token: Token specifying the next page of results. It should be obtained from
a ``search_model_versions`` call.
Returns:
A PagedList of :py:class:`mlflow.entities.model_registry.ModelVersion`
objects that satisfy the search expressions. The pagination token for the next
page can be obtained via the ``token`` attribute of the object.
"""
@abstractmethod
def set_model_version_tag(self, name, version, tag):
"""
Set a tag for the model version.
Args:
name: Registered model name.
version: Registered model version.
tag: :py:class:`mlflow.entities.model_registry.ModelVersionTag` instance to log.
Returns:
None
"""
@abstractmethod
def delete_model_version_tag(self, name, version, key):
"""
Delete a tag associated with the model version.
Args:
name: Registered model name.
version: Registered model version.
key: Tag key.
Returns:
None
"""
@abstractmethod
def set_registered_model_alias(self, name, alias, version):
"""
Set a registered model alias pointing to a model version.
Args:
name: Registered model name.
alias: Name of the alias.
version: Registered model version number.
Returns:
None
"""
@abstractmethod
def delete_registered_model_alias(self, name, alias):
"""
Delete an alias associated with a registered model.
Args:
name: Registered model name.
alias: Name of the alias.
Returns:
None
"""
@abstractmethod
def get_model_version_by_alias(self, name, alias):
"""
Get the model version instance by name and alias.
Args:
name: Registered model name.
alias: Name of the alias.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
def copy_model_version(self, src_mv, dst_name):
"""
Copy a model version from one registered model to another as a new model version.
Args:
src_mv: A :py:class:`mlflow.entities.model_registry.ModelVersion` object representing
the source model version.
dst_name: The name of the registered model to copy the model version to. If a
registered model with this name does not exist, it will be created.
Returns:
Single :py:class:`mlflow.entities.model_registry.ModelVersion` object representing
the cloned model version.
"""
try:
create_model_response = self.create_registered_model(dst_name)
eprint(f"Successfully registered model '{create_model_response.name}'.")
except MlflowException as e:
if e.error_code != ErrorCode.Name(RESOURCE_ALREADY_EXISTS):
raise
eprint(
f"Registered model '{dst_name}' already exists."
f" Creating a new version of this model..."
)
try:
mv_copy = self.create_model_version(
name=dst_name,
source=f"models:/{src_mv.name}/{src_mv.version}",
run_id=src_mv.run_id,
tags=[ModelVersionTag(k, v) for k, v in src_mv.tags.items()],
run_link=src_mv.run_link,
description=src_mv.description,
)
eprint(
f"Copied version '{src_mv.version}' of model '{src_mv.name}'"
f" to version '{mv_copy.version}' of model '{mv_copy.name}'."
)
except MlflowException as e:
raise MlflowException(
f"Failed to create model version copy. The current model registry backend "
f"may not yet support model version URI sources.\nError: {e}"
) from e
return mv_copy
def _await_model_version_creation(self, mv, await_creation_for):
"""
Await for model version to become ready after creation.
Args:
mv: A :py:class:`mlflow.entities.model_registry.ModelVersion` object.
await_creation_for: Number of seconds to wait for the model version to finish being
created and is in ``READY`` status.
"""
self._await_model_version_creation_impl(mv, await_creation_for)
def _await_model_version_creation_impl(self, mv, await_creation_for, hint=""):
entity_type = "Prompt" if has_prompt_tag(mv.tags) else "Model"
_logger.info(
f"Waiting up to {await_creation_for} seconds for {entity_type.lower()} version to "
f"finish creation. {entity_type} name: {mv.name}, version {mv.version}",
)
max_time = time() + await_creation_for
pending_status = ModelVersionStatus.to_string(ModelVersionStatus.PENDING_REGISTRATION)
while mv.status == pending_status:
if time() > max_time:
raise MlflowException(
f"Exceeded max wait time for model name: {mv.name} version: {mv.version} "
f"to become READY. Status: {mv.status} Wait Time: {await_creation_for}"
f".{hint}"
)
mv = self.get_model_version(mv.name, mv.version)
if mv.status != pending_status:
break
sleep(AWAIT_MODEL_VERSION_CREATE_SLEEP_INTERVAL_SECONDS)
if mv.status != ModelVersionStatus.to_string(ModelVersionStatus.READY):
raise MlflowException(
f"{entity_type} version creation failed for {entity_type.lower()} name: {mv.name} "
f"version: {mv.version} with status: {mv.status} and message: {mv.status_message}"
)

View File

@@ -0,0 +1,46 @@
from abc import ABCMeta, abstractmethod
from mlflow.store.model_registry.abstract_store import AbstractStore
from mlflow.utils.annotations import experimental
from mlflow.utils.rest_utils import (
call_endpoint,
call_endpoints,
)
@experimental
class BaseRestStore(AbstractStore):
"""
Base class client for a remote model registry server accessed via REST API calls
"""
__metaclass__ = ABCMeta
def __init__(self, get_host_creds):
super().__init__()
self.get_host_creds = get_host_creds
@abstractmethod
def _get_all_endpoints_from_method(self, method):
pass
@abstractmethod
def _get_endpoint_from_method(self, method):
pass
@abstractmethod
def _get_response_from_method(self, method):
pass
def _call_endpoint(self, api, json_body, call_all_endpoints=False, extra_headers=None):
response_proto = self._get_response_from_method(api)
if call_all_endpoints:
endpoints = self._get_all_endpoints_from_method(api)
return call_endpoints(
self.get_host_creds(), endpoints, json_body, response_proto, extra_headers
)
else:
endpoint, method = self._get_endpoint_from_method(api)
return call_endpoint(
self.get_host_creds(), endpoint, method, json_body, response_proto, extra_headers
)

View File

@@ -0,0 +1,37 @@
from mlflow.exceptions import MlflowException
from mlflow.store.model_registry.rest_store import RestStore
def _raise_unsupported_method(method, message=None):
messages = [
f"Method '{method}' is unsupported for models in the Workspace Model Registry. "
f"Upgrade to Models in Unity Catalog to access the latest features. You can configure "
f"the MLflow Python client to access models in Unity Catalog by running "
f"mlflow.set_registry_uri('databricks-uc') before accessing models.",
]
if message is not None:
messages.append(message)
raise MlflowException(" ".join(messages))
class DatabricksWorkspaceModelRegistryRestStore(RestStore):
def set_registered_model_alias(self, name, alias, version):
_raise_unsupported_method(method="set_registered_model_alias")
def delete_registered_model_alias(self, name, alias):
_raise_unsupported_method(method="delete_registered_model_alias")
def get_model_version_by_alias(self, name, alias):
_raise_unsupported_method(
method="get_model_version_by_alias",
message="If attempting to load a model version by alias via a URI of the form "
"'models:/model_name@alias_name', configure the MLflow client to target Unity Catalog "
"and try again.",
)
def _await_model_version_creation(self, mv, await_creation_for):
uc_hint = (
" For faster model version creation, use Models in Unity Catalog "
"(https://docs.databricks.com/en/machine-learning/manage-model-lifecycle/index.html)."
)
self._await_model_version_creation_impl(mv, await_creation_for, hint=uc_hint)

View File

@@ -0,0 +1,205 @@
from sqlalchemy import (
BigInteger,
Column,
ForeignKey,
ForeignKeyConstraint,
Integer,
PrimaryKeyConstraint,
String,
)
from sqlalchemy.orm import backref, relationship
from mlflow.entities.model_registry import (
ModelVersion,
ModelVersionTag,
RegisteredModel,
RegisteredModelAlias,
RegisteredModelTag,
)
from mlflow.entities.model_registry.model_version_stages import STAGE_DELETED_INTERNAL, STAGE_NONE
from mlflow.entities.model_registry.model_version_status import ModelVersionStatus
from mlflow.store.db.base_sql_model import Base
from mlflow.utils.time import get_current_time_millis
class SqlRegisteredModel(Base):
__tablename__ = "registered_models"
name = Column(String(256), unique=True, nullable=False)
creation_time = Column(BigInteger, default=get_current_time_millis)
last_updated_time = Column(BigInteger, nullable=True, default=None)
description = Column(String(5000), nullable=True)
__table_args__ = (PrimaryKeyConstraint("name", name="registered_model_pk"),)
def __repr__(self):
return (
f"<SqlRegisteredModel ({self.name}, {self.description}, "
f"{self.creation_time}, {self.last_updated_time})>"
)
def to_mlflow_entity(self):
# SqlRegisteredModel has backref to all "model_versions". Filter latest for each stage.
latest_versions = {}
for mv in self.model_versions:
stage = mv.current_stage
if stage != STAGE_DELETED_INTERNAL and (
stage not in latest_versions or latest_versions[stage].version < mv.version
):
latest_versions[stage] = mv
return RegisteredModel(
self.name,
self.creation_time,
self.last_updated_time,
self.description,
[mvd.to_mlflow_entity() for mvd in latest_versions.values()],
[tag.to_mlflow_entity() for tag in self.registered_model_tags],
[alias.to_mlflow_entity() for alias in self.registered_model_aliases],
)
class SqlModelVersion(Base):
__tablename__ = "model_versions"
name = Column(String(256), ForeignKey("registered_models.name", onupdate="cascade"))
version = Column(Integer, nullable=False)
creation_time = Column(BigInteger, default=get_current_time_millis)
last_updated_time = Column(BigInteger, nullable=True, default=None)
description = Column(String(5000), nullable=True)
user_id = Column(String(256), nullable=True, default=None)
current_stage = Column(String(20), default=STAGE_NONE)
source = Column(String(500), nullable=True, default=None)
storage_location = Column(String(500), nullable=True, default=None)
run_id = Column(String(32), nullable=True, default=None)
run_link = Column(String(500), nullable=True, default=None)
status = Column(String(20), default=ModelVersionStatus.to_string(ModelVersionStatus.READY))
status_message = Column(String(500), nullable=True, default=None)
# linked entities
registered_model = relationship(
"SqlRegisteredModel", backref=backref("model_versions", cascade="all")
)
__table_args__ = (PrimaryKeyConstraint("name", "version", name="model_version_pk"),)
# entity mappers
def to_mlflow_entity(self):
return ModelVersion(
self.name,
self.version,
self.creation_time,
self.last_updated_time,
self.description,
self.user_id,
self.current_stage,
self.source,
self.run_id,
self.status,
self.status_message,
[tag.to_mlflow_entity() for tag in self.model_version_tags],
self.run_link,
[],
)
class SqlRegisteredModelTag(Base):
__tablename__ = "registered_model_tags"
name = Column(String(256), ForeignKey("registered_models.name", onupdate="cascade"))
key = Column(String(250), nullable=False)
value = Column(String(5000), nullable=True)
# linked entities
registered_model = relationship(
"SqlRegisteredModel", backref=backref("registered_model_tags", cascade="all")
)
__table_args__ = (PrimaryKeyConstraint("key", "name", name="registered_model_tag_pk"),)
def __repr__(self):
return f"<SqlRegisteredModelTag ({self.name}, {self.key}, {self.value})>"
# entity mappers
def to_mlflow_entity(self):
return RegisteredModelTag(self.key, self.value)
class SqlModelVersionTag(Base):
__tablename__ = "model_version_tags"
name = Column(String(256))
version = Column(Integer)
key = Column(String(250), nullable=False)
value = Column(String(5000), nullable=True)
# linked entities
model_version = relationship(
"SqlModelVersion",
foreign_keys=[name, version],
backref=backref("model_version_tags", cascade="all"),
)
__table_args__ = (
PrimaryKeyConstraint("key", "name", "version", name="model_version_tag_pk"),
ForeignKeyConstraint(
("name", "version"),
("model_versions.name", "model_versions.version"),
onupdate="cascade",
),
)
def __repr__(self):
return f"<SqlModelVersionTag ({self.name}, {self.version}, {self.key}, {self.value})>"
# entity mappers
def to_mlflow_entity(self):
return ModelVersionTag(self.key, self.value)
class SqlRegisteredModelAlias(Base):
__tablename__ = "registered_model_aliases"
name = Column(
String(256),
ForeignKey(
"registered_models.name",
onupdate="cascade",
ondelete="cascade",
name="registered_model_alias_name_fkey",
),
)
alias = Column(String(256), nullable=False)
version = Column(Integer, nullable=False)
# linked entities
registered_model = relationship(
"SqlRegisteredModel", backref=backref("registered_model_aliases", cascade="all")
)
__table_args__ = (PrimaryKeyConstraint("name", "alias", name="registered_model_alias_pk"),)
def __repr__(self):
return f"<SqlRegisteredModelAlias ({self.name}, {self.alias}, {self.version})>"
# entity mappers
def to_mlflow_entity(self):
return RegisteredModelAlias(self.alias, self.version)

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,474 @@
import logging
from mlflow.entities.model_registry import ModelVersion, RegisteredModel
from mlflow.protos.model_registry_pb2 import (
CreateModelVersion,
CreateRegisteredModel,
DeleteModelVersion,
DeleteModelVersionTag,
DeleteRegisteredModel,
DeleteRegisteredModelAlias,
DeleteRegisteredModelTag,
GetLatestVersions,
GetModelVersion,
GetModelVersionByAlias,
GetModelVersionDownloadUri,
GetRegisteredModel,
ModelRegistryService,
RenameRegisteredModel,
SearchModelVersions,
SearchRegisteredModels,
SetModelVersionTag,
SetRegisteredModelAlias,
SetRegisteredModelTag,
TransitionModelVersionStage,
UpdateModelVersion,
UpdateRegisteredModel,
)
from mlflow.store.entities.paged_list import PagedList
from mlflow.store.model_registry.base_rest_store import BaseRestStore
from mlflow.utils.proto_json_utils import message_to_json
from mlflow.utils.rest_utils import (
_REST_API_PATH_PREFIX,
extract_all_api_info_for_service,
extract_api_info_for_service,
)
_METHOD_TO_INFO = extract_api_info_for_service(ModelRegistryService, _REST_API_PATH_PREFIX)
_METHOD_TO_ALL_INFO = extract_all_api_info_for_service(ModelRegistryService, _REST_API_PATH_PREFIX)
_logger = logging.getLogger(__name__)
class RestStore(BaseRestStore):
"""
Client for a remote model registry server accessed via REST API calls
Args:
get_host_creds: Method to be invoked prior to every REST request to get the
:py:class:`mlflow.rest_utils.MlflowHostCreds` for the request. Note that this
is a function so that we can obtain fresh credentials in the case of expiry.
"""
def _get_response_from_method(self, method):
return method.Response()
def _get_endpoint_from_method(self, method):
return _METHOD_TO_INFO[method]
def _get_all_endpoints_from_method(self, method):
return _METHOD_TO_ALL_INFO[method]
# CRUD API for RegisteredModel objects
def create_registered_model(self, name, tags=None, description=None):
"""
Create a new registered model in backend store.
Args:
name: Name of the new model. This is expected to be unique in the backend store.
tags: A list of :py:class:`mlflow.entities.model_registry.RegisteredModelTag`
instances associated with this registered model.
description: Description of the model.
Returns:
A single object of :py:class:`mlflow.entities.model_registry.RegisteredModel`
created in the backend.
"""
proto_tags = [tag.to_proto() for tag in tags or []]
req_body = message_to_json(
CreateRegisteredModel(name=name, tags=proto_tags, description=description)
)
response_proto = self._call_endpoint(CreateRegisteredModel, req_body)
return RegisteredModel.from_proto(response_proto.registered_model)
def update_registered_model(self, name, description):
"""
Update description of the registered model.
Args:
name: Registered model name.
description: New description.
Returns:
A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
"""
req_body = message_to_json(UpdateRegisteredModel(name=name, description=description))
response_proto = self._call_endpoint(UpdateRegisteredModel, req_body)
return RegisteredModel.from_proto(response_proto.registered_model)
def rename_registered_model(self, name, new_name):
"""
Rename the registered model.
Args:
name: Registered model name.
new_name: New proposed name.
Returns:
A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
"""
req_body = message_to_json(RenameRegisteredModel(name=name, new_name=new_name))
response_proto = self._call_endpoint(RenameRegisteredModel, req_body)
return RegisteredModel.from_proto(response_proto.registered_model)
def delete_registered_model(self, name):
"""
Delete the registered model.
Backend raises exception if a registered model with given name does not exist.
Args:
name: Registered model name.
Returns:
None
"""
req_body = message_to_json(DeleteRegisteredModel(name=name))
self._call_endpoint(DeleteRegisteredModel, req_body)
def search_registered_models(
self, filter_string=None, max_results=None, order_by=None, page_token=None
):
"""
Search for registered models in backend that satisfy the filter criteria.
Args:
filter_string: Filter query string, defaults to searching all registered models.
max_results: Maximum number of registered models desired.
order_by: List of column names with ASC|DESC annotation, to be used for ordering
matching search results.
page_token: Token specifying the next page of results. It should be obtained from
a ``search_registered_models`` call.
Returns:
A PagedList of :py:class:`mlflow.entities.model_registry.RegisteredModel` objects
that satisfy the search expressions. The pagination token for the next page can be
obtained via the ``token`` attribute of the object.
"""
req_body = message_to_json(
SearchRegisteredModels(
filter=filter_string,
max_results=max_results,
order_by=order_by,
page_token=page_token,
)
)
response_proto = self._call_endpoint(SearchRegisteredModels, req_body)
registered_models = [
RegisteredModel.from_proto(registered_model)
for registered_model in response_proto.registered_models
]
return PagedList(registered_models, response_proto.next_page_token)
def get_registered_model(self, name):
"""
Get registered model instance by name.
Args:
name: Registered model name.
Returns:
A single :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
"""
req_body = message_to_json(GetRegisteredModel(name=name))
response_proto = self._call_endpoint(GetRegisteredModel, req_body)
return RegisteredModel.from_proto(response_proto.registered_model)
def get_latest_versions(self, name, stages=None):
"""
Latest version models for each requested stage. If no ``stages`` argument is provided,
returns the latest version for each stage.
Args:
name: Registered model name.
stages: List of desired stages. If input list is None, return latest versions for
each stage.
Returns:
List of :py:class:`mlflow.entities.model_registry.ModelVersion` objects.
"""
req_body = message_to_json(GetLatestVersions(name=name, stages=stages))
response_proto = self._call_endpoint(GetLatestVersions, req_body, call_all_endpoints=True)
return [
ModelVersion.from_proto(model_version)
for model_version in response_proto.model_versions
]
def set_registered_model_tag(self, name, tag):
"""
Set a tag for the registered model.
Args:
name: Registered model name.
tag: :py:class:`mlflow.entities.model_registry.RegisteredModelTag` instance to log.
Returns:
None
"""
req_body = message_to_json(SetRegisteredModelTag(name=name, key=tag.key, value=tag.value))
self._call_endpoint(SetRegisteredModelTag, req_body)
def delete_registered_model_tag(self, name, key):
"""
Delete a tag associated with the registered model.
Args:
name: Registered model name.
key: Registered model tag key.
Returns:
None
"""
req_body = message_to_json(DeleteRegisteredModelTag(name=name, key=key))
self._call_endpoint(DeleteRegisteredModelTag, req_body)
# CRUD API for ModelVersion objects
def create_model_version(
self,
name,
source,
run_id=None,
tags=None,
run_link=None,
description=None,
local_model_path=None,
):
"""
Create a new model version from given source and run ID.
Args:
name: Registered model name.
source: URI indicating the location of the model artifacts.
run_id: Run ID from MLflow tracking server that generated the model.
tags: A list of :py:class:`mlflow.entities.model_registry.ModelVersionTag`
instances associated with this model version.
run_link: Link to the run from an MLflow tracking server that generated this model.
description: Description of the version.
local_model_path: Unused.
Returns:
A single object of :py:class:`mlflow.entities.model_registry.ModelVersion`
created in the backend.
"""
proto_tags = [tag.to_proto() for tag in tags or []]
req_body = message_to_json(
CreateModelVersion(
name=name,
source=source,
run_id=run_id,
run_link=run_link,
tags=proto_tags,
description=description,
)
)
response_proto = self._call_endpoint(CreateModelVersion, req_body)
return ModelVersion.from_proto(response_proto.model_version)
def transition_model_version_stage(self, name, version, stage, archive_existing_versions):
"""
Update model version stage.
Args:
name: Registered model name.
version: Registered model version.
stage: New desired stage for this model version.
archive_existing_versions: If this flag is set to ``True``, all existing model
versions in the stage will be automatically moved to the "archived" stage. Only
valid when ``stage`` is ``"staging"`` or ``"production"`` otherwise an error will
be raised.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
req_body = message_to_json(
TransitionModelVersionStage(
name=name,
version=str(version),
stage=stage,
archive_existing_versions=archive_existing_versions,
)
)
response_proto = self._call_endpoint(TransitionModelVersionStage, req_body)
return ModelVersion.from_proto(response_proto.model_version)
def update_model_version(self, name, version, description):
"""
Update metadata associated with a model version in backend.
Args:
name: Registered model name.
version: Registered model version.
description: New model description.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
req_body = message_to_json(
UpdateModelVersion(name=name, version=str(version), description=description)
)
response_proto = self._call_endpoint(UpdateModelVersion, req_body)
return ModelVersion.from_proto(response_proto.model_version)
def delete_model_version(self, name, version):
"""
Delete model version in backend.
Args:
name: Registered model name.
version: Registered model version.
Returns:
None
"""
req_body = message_to_json(DeleteModelVersion(name=name, version=str(version)))
self._call_endpoint(DeleteModelVersion, req_body)
def get_model_version(self, name, version):
"""
Get the model version instance by name and version.
Args:
name: Registered model name.
version: Registered model version.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
req_body = message_to_json(GetModelVersion(name=name, version=str(version)))
response_proto = self._call_endpoint(GetModelVersion, req_body)
return ModelVersion.from_proto(response_proto.model_version)
def get_model_version_download_uri(self, name, version):
"""
Get the download location in Model Registry for this model version.
NOTE: For first version of Model Registry, since the models are not copied over to another
location, download URI points to input source path.
Args:
name: Registered model name.
version: Registered model version.
Returns:
A single URI location that allows reads for downloading.
"""
req_body = message_to_json(GetModelVersionDownloadUri(name=name, version=str(version)))
response_proto = self._call_endpoint(GetModelVersionDownloadUri, req_body)
return response_proto.artifact_uri
def search_model_versions(
self, filter_string=None, max_results=None, order_by=None, page_token=None
):
"""
Search for model versions in backend that satisfy the filter criteria.
Args:
filter_string: A filter string expression. Currently supports a single filter
condition either name of model like ``name = 'model_name'`` or
``run_id = '...'``.
max_results: Maximum number of model versions desired.
order_by: List of column names with ASC|DESC annotation, to be used for ordering
matching search results.
page_token: Token specifying the next page of results. It should be obtained from
a ``search_model_versions`` call.
Returns:
A PagedList of :py:class:`mlflow.entities.model_registry.ModelVersion`
objects that satisfy the search expressions. The pagination token for the next
page can be obtained via the ``token`` attribute of the object.
"""
req_body = message_to_json(
SearchModelVersions(
filter=filter_string,
max_results=max_results,
order_by=order_by,
page_token=page_token,
)
)
response_proto = self._call_endpoint(SearchModelVersions, req_body)
model_versions = [ModelVersion.from_proto(mvd) for mvd in response_proto.model_versions]
return PagedList(model_versions, response_proto.next_page_token)
def set_model_version_tag(self, name, version, tag):
"""
Set a tag for the model version.
Args:
name: Registered model name.
version: Registered model version.
tag: :py:class:`mlflow.entities.model_registry.ModelVersionTag` instance to log.
Returns:
None
"""
req_body = message_to_json(
SetModelVersionTag(name=name, version=version, key=tag.key, value=tag.value)
)
self._call_endpoint(SetModelVersionTag, req_body)
def delete_model_version_tag(self, name, version, key):
"""
Delete a tag associated with the model version.
Args:
name: Registered model name.
version: Registered model version.
key: Tag key.
Returns:
None
"""
req_body = message_to_json(DeleteModelVersionTag(name=name, version=version, key=key))
self._call_endpoint(DeleteModelVersionTag, req_body)
def set_registered_model_alias(self, name, alias, version):
"""
Set a registered model alias pointing to a model version.
Args:
name: Registered model name.
alias: Name of the alias.
version: Registered model version number.
Returns:
None
"""
req_body = message_to_json(
SetRegisteredModelAlias(name=name, alias=alias, version=str(version))
)
self._call_endpoint(SetRegisteredModelAlias, req_body)
def delete_registered_model_alias(self, name, alias):
"""
Delete an alias associated with a registered model.
Args:
name: Registered model name.
alias: Name of the alias.
Returns:
None
"""
req_body = message_to_json(DeleteRegisteredModelAlias(name=name, alias=alias))
self._call_endpoint(DeleteRegisteredModelAlias, req_body)
def get_model_version_by_alias(self, name, alias):
"""
Get the model version instance by name and alias.
Args:
name: Registered model name.
alias: Name of the alias.
Returns:
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
"""
req_body = message_to_json(GetModelVersionByAlias(name=name, alias=alias))
response_proto = self._call_endpoint(GetModelVersionByAlias, req_body)
return ModelVersion.from_proto(response_proto.model_version)

File diff suppressed because it is too large Load Diff