mlflow--mlflow
1794 行
72 KiB
Python
1794 行
72 KiB
Python
import logging
|
|
import threading
|
|
import urllib
|
|
import uuid
|
|
from typing import Any
|
|
|
|
import sqlalchemy
|
|
from sqlalchemy import select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from mlflow.entities.model_registry.model_version_stages import (
|
|
ALL_STAGES,
|
|
DEFAULT_STAGES_FOR_GET_LATEST_VERSIONS,
|
|
STAGE_ARCHIVED,
|
|
STAGE_DELETED_INTERNAL,
|
|
get_canonical_stage,
|
|
)
|
|
from mlflow.entities.model_registry.prompt_version import IS_PROMPT_TAG_KEY
|
|
from mlflow.entities.webhook import Webhook, WebhookEvent, WebhookStatus
|
|
from mlflow.exceptions import MlflowException
|
|
from mlflow.prompt.registry_utils import handle_resource_already_exist_error, has_prompt_tag
|
|
from mlflow.protos.databricks_pb2 import (
|
|
INVALID_PARAMETER_VALUE,
|
|
INVALID_STATE,
|
|
RESOURCE_ALREADY_EXISTS,
|
|
RESOURCE_DOES_NOT_EXIST,
|
|
)
|
|
from mlflow.store.artifact.utils.models import _parse_model_uri
|
|
from mlflow.store.db.utils import (
|
|
_all_tables_exist,
|
|
_get_managed_session_maker,
|
|
_get_routing_session_maker,
|
|
_initialize_tables,
|
|
create_sqlalchemy_engine_with_retry,
|
|
)
|
|
from mlflow.store.entities.paged_list import PagedList
|
|
from mlflow.store.model_registry import (
|
|
SEARCH_MODEL_VERSION_MAX_RESULTS_DEFAULT,
|
|
SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD,
|
|
SEARCH_REGISTERED_MODEL_MAX_RESULTS_DEFAULT,
|
|
SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD,
|
|
)
|
|
from mlflow.store.model_registry.abstract_store import AbstractStore
|
|
from mlflow.store.model_registry.dbmodels.models import (
|
|
SqlModelVersion,
|
|
SqlModelVersionTag,
|
|
SqlRegisteredModel,
|
|
SqlRegisteredModelAlias,
|
|
SqlRegisteredModelTag,
|
|
SqlWebhook,
|
|
SqlWebhookEvent,
|
|
)
|
|
from mlflow.tracking.client import MlflowClient
|
|
from mlflow.utils.search_utils import SearchModelUtils, SearchModelVersionUtils, SearchUtils
|
|
from mlflow.utils.time import get_current_time_millis
|
|
from mlflow.utils.uri import extract_db_type_from_uri
|
|
from mlflow.utils.validation import (
|
|
_REGISTERED_MODEL_ALIAS_LATEST,
|
|
_validate_model_alias_name,
|
|
_validate_model_alias_name_reserved,
|
|
_validate_model_name,
|
|
_validate_model_renaming,
|
|
_validate_model_version,
|
|
_validate_model_version_tag,
|
|
_validate_registered_model_tag,
|
|
_validate_tag_name,
|
|
_validate_webhook_events,
|
|
_validate_webhook_name,
|
|
_validate_webhook_url,
|
|
)
|
|
from mlflow.utils.workspace_utils import DEFAULT_WORKSPACE_NAME
|
|
|
|
_logger = logging.getLogger(__name__)
|
|
|
|
# Models that carry a ``workspace`` column and must be filtered
|
|
# by the active workspace in every query.
|
|
_WORKSPACE_MODELS = (
|
|
SqlRegisteredModel,
|
|
SqlModelVersion,
|
|
SqlWebhook,
|
|
SqlRegisteredModelTag,
|
|
SqlModelVersionTag,
|
|
SqlRegisteredModelAlias,
|
|
)
|
|
|
|
# For each database table, fetch its columns and define an appropriate attribute for each column
|
|
# on the table's associated object representation (Mapper). This is necessary to ensure that
|
|
# columns defined via backreference are available as Mapper instance attributes (e.g.,
|
|
# ``SqlRegisteredModel.model_versions``). For more information, see
|
|
# https://docs.sqlalchemy.org/en/latest/orm/mapping_api.html#sqlalchemy.orm.configure_mappers
|
|
# and https://docs.sqlalchemy.org/en/latest/orm/mapping_api.html#sqlalchemy.orm.mapper.Mapper
|
|
sqlalchemy.orm.configure_mappers()
|
|
|
|
|
|
class SqlAlchemyStore(AbstractStore):
|
|
"""
|
|
This entity may change or be removed in a future release without warning.
|
|
SQLAlchemy compliant backend store for tracking meta data for MLflow entities. MLflow
|
|
supports the database dialects ``mysql``, ``mssql``, ``sqlite``, and ``postgresql``.
|
|
As specified in the
|
|
`SQLAlchemy docs <https://docs.sqlalchemy.org/en/latest/core/engines.html#database-urls>`_ ,
|
|
the database URI is expected in the format
|
|
``<dialect>+<driver>://<username>:<password>@<host>:<port>/<database>``. If you do not
|
|
specify a driver, SQLAlchemy uses a dialect's default driver.
|
|
|
|
This store interacts with SQL store using SQLAlchemy abstractions defined for MLflow entities.
|
|
:py:class:`mlflow.store.model_registry.models.RegisteredModel` and
|
|
:py:class:`mlflow.store.model_registry.models.ModelVersion`
|
|
"""
|
|
|
|
CREATE_MODEL_VERSION_RETRIES = 3
|
|
|
|
# Class-level cache for SQLAlchemy engines to prevent connection pool leaks
|
|
# when multiple store instances are created with the same database URI.
|
|
_engine_map: dict[str, sqlalchemy.engine.Engine] = {}
|
|
_engine_map_lock = threading.Lock()
|
|
|
|
@classmethod
|
|
def _get_or_create_engine(cls, db_uri: str) -> sqlalchemy.engine.Engine:
|
|
"""Get a cached engine or create a new one for the given database URI."""
|
|
if db_uri not in cls._engine_map:
|
|
with cls._engine_map_lock:
|
|
if db_uri not in cls._engine_map:
|
|
cls._engine_map[db_uri] = create_sqlalchemy_engine_with_retry(db_uri)
|
|
return cls._engine_map[db_uri]
|
|
|
|
def __init__(self, db_uri, read_db_uri=None):
|
|
"""
|
|
Create a database backed store.
|
|
|
|
Args:
|
|
db_uri: The SQLAlchemy database URI string to connect to the database. See
|
|
the `SQLAlchemy docs
|
|
<https://docs.sqlalchemy.org/en/latest/core/engines.html#database-urls>`_
|
|
for format specifications. MLflow supports the dialects ``mysql``,
|
|
``mssql``, ``sqlite``, and ``postgresql``.
|
|
read_db_uri: Optional SQLAlchemy database URI for a read replica. When provided,
|
|
read operations are routed to this URI while write operations use ``db_uri``.
|
|
If not provided, all operations use ``db_uri``.
|
|
"""
|
|
super().__init__()
|
|
self.db_uri = db_uri
|
|
self.db_type = extract_db_type_from_uri(db_uri)
|
|
self.engine = self._get_or_create_engine(db_uri)
|
|
if not _all_tables_exist(self.engine):
|
|
_initialize_tables(self.engine)
|
|
# Verify that all model registry tables exist.
|
|
SqlAlchemyStore._verify_registry_tables_exist(self.engine)
|
|
|
|
# Set up read replica engine if provided
|
|
if read_db_uri and read_db_uri != db_uri:
|
|
self.read_engine = self._get_or_create_engine(read_db_uri)
|
|
WriteSessionMaker = sqlalchemy.orm.sessionmaker(bind=self.engine)
|
|
ReadSessionMaker = sqlalchemy.orm.sessionmaker(bind=self.read_engine)
|
|
self.ManagedSessionMaker = _get_routing_session_maker(
|
|
WriteSessionMaker, ReadSessionMaker, self.db_type
|
|
)
|
|
else:
|
|
if read_db_uri and read_db_uri == db_uri:
|
|
_logger.warning(
|
|
"read_db_uri is the same as the primary db_uri; "
|
|
"read replica routing will not be enabled. "
|
|
"This is likely a configuration mistake."
|
|
)
|
|
self.read_engine = None
|
|
SessionMaker = sqlalchemy.orm.sessionmaker(bind=self.engine)
|
|
self.ManagedSessionMaker = _get_managed_session_maker(SessionMaker, self.db_type)
|
|
# TODO: verify schema here once we add logic to initialize the registry tables if they
|
|
# don't exist (schema verification will fail in tests otherwise)
|
|
# mlflow.store.db.utils._verify_schema(self.engine)
|
|
self._initialize_store_state()
|
|
|
|
@property
|
|
def supports_workspaces(self) -> bool:
|
|
"""Indicates whether this store supports workspace isolation."""
|
|
return False
|
|
|
|
def _get_active_workspace(self) -> str:
|
|
"""
|
|
Get the active workspace name.
|
|
|
|
In single-tenant mode, always returns DEFAULT_WORKSPACE_NAME.
|
|
Workspace-aware subclasses override this to enforce isolation.
|
|
"""
|
|
return DEFAULT_WORKSPACE_NAME
|
|
|
|
def _get_query(self, session, model):
|
|
"""
|
|
Return a query for ``model``.
|
|
Always filter on workspace for relevant models, to benefit from DB index.
|
|
"""
|
|
query = session.query(model)
|
|
if model in _WORKSPACE_MODELS:
|
|
query = query.filter(model.workspace == self._get_active_workspace())
|
|
return query
|
|
|
|
def _with_workspace_field(self, instance):
|
|
"""
|
|
Allow subclasses to populate model fields (e.g., workspace metadata) on ORM instances.
|
|
"""
|
|
if hasattr(instance, "workspace") and getattr(instance, "workspace", None) is None:
|
|
instance.workspace = DEFAULT_WORKSPACE_NAME
|
|
return instance
|
|
|
|
def _initialize_store_state(self):
|
|
"""
|
|
Initialize store state after construction.
|
|
|
|
In single-tenant mode, validates no registry objects exist outside the default workspace.
|
|
"""
|
|
with self.ManagedSessionMaker() as session:
|
|
exists_non_default_rm = (
|
|
session
|
|
.query(SqlRegisteredModel.name)
|
|
.filter(SqlRegisteredModel.workspace.isnot(None))
|
|
.filter(SqlRegisteredModel.workspace != DEFAULT_WORKSPACE_NAME)
|
|
.first()
|
|
is not None
|
|
)
|
|
|
|
if exists_non_default_rm:
|
|
name, workspace = (
|
|
session
|
|
.query(SqlRegisteredModel.name, SqlRegisteredModel.workspace)
|
|
.filter(SqlRegisteredModel.workspace != DEFAULT_WORKSPACE_NAME)
|
|
.first()
|
|
)
|
|
raise MlflowException(
|
|
"Cannot disable workspaces because registered models exist outside the default "
|
|
f"workspace (for example, model '{name}' in workspace '{workspace}'). "
|
|
"Either remove those models or re-enable workspaces before starting.",
|
|
error_code=INVALID_STATE,
|
|
)
|
|
|
|
exists_non_default_webhook = (
|
|
session
|
|
.query(SqlWebhook.webhook_id)
|
|
.filter(SqlWebhook.workspace.isnot(None))
|
|
.filter(SqlWebhook.workspace != DEFAULT_WORKSPACE_NAME)
|
|
.first()
|
|
is not None
|
|
)
|
|
|
|
if exists_non_default_webhook:
|
|
webhook_id, workspace = (
|
|
session
|
|
.query(SqlWebhook.webhook_id, SqlWebhook.workspace)
|
|
.filter(SqlWebhook.workspace != DEFAULT_WORKSPACE_NAME)
|
|
.first()
|
|
)
|
|
raise MlflowException(
|
|
"Cannot disable workspaces because webhooks exist outside the default "
|
|
f"workspace (for example, webhook '{webhook_id}' in workspace '{workspace}'). "
|
|
"Either remove those webhooks or re-enable workspaces before starting.",
|
|
error_code=INVALID_STATE,
|
|
)
|
|
|
|
def _get_dialect(self):
|
|
return self.engine.dialect.name
|
|
|
|
def _dispose_engine(self):
|
|
self.engine.dispose()
|
|
|
|
@staticmethod
|
|
def _verify_registry_tables_exist(engine):
|
|
# Verify that all tables have been created.
|
|
inspected_tables = set(sqlalchemy.inspect(engine).get_table_names())
|
|
expected_tables = [
|
|
SqlRegisteredModel.__tablename__,
|
|
SqlModelVersion.__tablename__,
|
|
SqlWebhook.__tablename__,
|
|
SqlWebhookEvent.__tablename__,
|
|
]
|
|
if any(table not in inspected_tables for table in expected_tables):
|
|
# TODO: Replace the MlflowException with the following line once it's possible to run
|
|
# the registry against a different DB than the tracking server:
|
|
# mlflow.store.db.utils._initialize_tables(self.engine)
|
|
raise MlflowException("Database migration in unexpected state. Run manual upgrade.")
|
|
|
|
@staticmethod
|
|
def _get_eager_registered_model_query_options():
|
|
"""
|
|
A list of SQLAlchemy query options that can be used to eagerly
|
|
load the following registered model attributes
|
|
when fetching a registered model: ``registered_model_tags`` and
|
|
``registered_model_aliases``.
|
|
"""
|
|
# Use a subquery load rather than a joined load in order to minimize the memory overhead
|
|
# of the eager loading procedure. For more information about relationship loading
|
|
# techniques, see https://docs.sqlalchemy.org/en/13/orm/
|
|
# loading_relationships.html#relationship-loading-techniques
|
|
return [
|
|
sqlalchemy.orm.subqueryload(SqlRegisteredModel.registered_model_tags),
|
|
sqlalchemy.orm.subqueryload(SqlRegisteredModel.registered_model_aliases),
|
|
]
|
|
|
|
def _get_latest_versions_for_models(
|
|
self, session, model_names: list[str]
|
|
) -> dict[str, list[SqlModelVersion]]:
|
|
"""
|
|
Batch-fetch the latest model version per stage for multiple registered models.
|
|
|
|
Uses a SQL window function to compute the latest version per (name, stage) directly
|
|
in the database, avoiding N+1 queries and Python-side iteration through all versions.
|
|
"""
|
|
if not model_names:
|
|
return {}
|
|
|
|
workspace_clauses = self._get_workspace_clauses(SqlModelVersion)
|
|
|
|
row_num = (
|
|
sqlalchemy.func
|
|
.row_number()
|
|
.over(
|
|
partition_by=[
|
|
SqlModelVersion.workspace,
|
|
SqlModelVersion.name,
|
|
SqlModelVersion.current_stage,
|
|
],
|
|
order_by=SqlModelVersion.version.desc(),
|
|
)
|
|
.label("rn")
|
|
)
|
|
|
|
subquery = (
|
|
select(SqlModelVersion, row_num)
|
|
.where(
|
|
*workspace_clauses,
|
|
SqlModelVersion.name.in_(model_names),
|
|
SqlModelVersion.current_stage != STAGE_DELETED_INTERNAL,
|
|
)
|
|
.subquery()
|
|
)
|
|
|
|
query = (
|
|
select(SqlModelVersion)
|
|
.where(*workspace_clauses)
|
|
.join(
|
|
subquery,
|
|
sqlalchemy.and_(
|
|
SqlModelVersion.workspace == subquery.c.workspace,
|
|
SqlModelVersion.name == subquery.c.name,
|
|
SqlModelVersion.version == subquery.c.version,
|
|
),
|
|
)
|
|
.where(subquery.c.rn == 1)
|
|
.options(*self._get_eager_model_version_query_options())
|
|
)
|
|
|
|
latest_versions = session.execute(query).scalars().all()
|
|
|
|
result: dict[str, list[SqlModelVersion]] = {name: [] for name in model_names}
|
|
for mv in latest_versions:
|
|
result[mv.name].append(mv)
|
|
|
|
return result
|
|
|
|
@staticmethod
|
|
def _get_eager_model_version_query_options():
|
|
"""
|
|
A list of SQLAlchemy query options that can be used to eagerly
|
|
load the following model version attributes
|
|
when fetching a model version: ``model_version_tags``.
|
|
"""
|
|
# Use a subquery load rather than a joined load in order to minimize the memory overhead
|
|
# of the eager loading procedure. For more information about relationship loading
|
|
# techniques, see https://docs.sqlalchemy.org/en/13/orm/
|
|
# loading_relationships.html#relationship-loading-techniques
|
|
return [sqlalchemy.orm.subqueryload(SqlModelVersion.model_version_tags)]
|
|
|
|
def _get_workspace_clauses(self, model):
|
|
"""
|
|
Return workspace filter clauses for the model.
|
|
Always filter on workspace for relevant models, to benefit from DB index.
|
|
"""
|
|
if model in _WORKSPACE_MODELS:
|
|
return [model.workspace == self._get_active_workspace()]
|
|
return []
|
|
|
|
def create_registered_model(self, name, tags=None, description=None, deployment_job_id=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 version.
|
|
deployment_job_id: Optional deployment job ID.
|
|
|
|
Returns:
|
|
A single object of :py:class:`mlflow.entities.model_registry.RegisteredModel`
|
|
created in the backend.
|
|
"""
|
|
_validate_model_name(name)
|
|
for tag in tags or []:
|
|
_validate_registered_model_tag(tag.key, tag.value)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
try:
|
|
creation_time = get_current_time_millis()
|
|
registered_model = self._with_workspace_field(
|
|
SqlRegisteredModel(
|
|
name=name,
|
|
creation_time=creation_time,
|
|
last_updated_time=creation_time,
|
|
description=description,
|
|
)
|
|
)
|
|
tags_dict = {}
|
|
for tag in tags or []:
|
|
tags_dict[tag.key] = tag.value
|
|
registered_model.registered_model_tags = [
|
|
self._with_workspace_field(
|
|
SqlRegisteredModelTag(name=name, key=key, value=value)
|
|
)
|
|
for key, value in tags_dict.items()
|
|
]
|
|
session.add(registered_model)
|
|
session.flush()
|
|
return registered_model.to_mlflow_entity()
|
|
except sqlalchemy.exc.IntegrityError:
|
|
existing_model = self.get_registered_model(name)
|
|
handle_resource_already_exist_error(
|
|
name, has_prompt_tag(existing_model._tags), has_prompt_tag(tags)
|
|
)
|
|
|
|
def _get_registered_model(self, session, name, eager=False):
|
|
"""
|
|
Args:
|
|
eager: If ``True``, eagerly loads the registered model's tags. If ``False``, these
|
|
attributes are not eagerly loaded and will be loaded when their corresponding object
|
|
properties are accessed from the resulting ``SqlRegisteredModel`` object.
|
|
"""
|
|
_validate_model_name(name)
|
|
query = self._get_query(session, SqlRegisteredModel)
|
|
if eager:
|
|
query = query.options(*self._get_eager_registered_model_query_options())
|
|
rms = query.filter(SqlRegisteredModel.name == name).all()
|
|
if len(rms) == 0:
|
|
raise MlflowException(
|
|
f"Registered Model with name={name} not found", RESOURCE_DOES_NOT_EXIST
|
|
)
|
|
if len(rms) > 1:
|
|
raise MlflowException(
|
|
f"Expected only 1 registered model with name={name}. Found {len(rms)}.",
|
|
INVALID_STATE,
|
|
)
|
|
return rms[0]
|
|
|
|
def update_registered_model(self, name, description, deployment_job_id=None):
|
|
"""
|
|
Update description of the registered model.
|
|
|
|
Args:
|
|
name: Registered model name.
|
|
description: New description.
|
|
deployment_job_id: Optional deployment job ID.
|
|
|
|
Returns:
|
|
A single updated :py:class:`mlflow.entities.model_registry.RegisteredModel` object.
|
|
|
|
"""
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
sql_registered_model = self._get_registered_model(session, name)
|
|
updated_time = get_current_time_millis()
|
|
sql_registered_model.description = description
|
|
sql_registered_model.last_updated_time = updated_time
|
|
session.add(sql_registered_model)
|
|
session.flush()
|
|
return sql_registered_model.to_mlflow_entity()
|
|
|
|
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.
|
|
|
|
"""
|
|
_validate_model_renaming(new_name)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
sql_registered_model = self._get_registered_model(session, name)
|
|
try:
|
|
updated_time = get_current_time_millis()
|
|
sql_registered_model.name = new_name
|
|
for sql_model_version in sql_registered_model.model_versions:
|
|
sql_model_version.name = new_name
|
|
sql_model_version.last_updated_time = updated_time
|
|
sql_registered_model.last_updated_time = updated_time
|
|
session.add_all([sql_registered_model] + sql_registered_model.model_versions)
|
|
session.flush()
|
|
return sql_registered_model.to_mlflow_entity()
|
|
except sqlalchemy.exc.IntegrityError as e:
|
|
raise MlflowException(
|
|
f"Registered Model (name={new_name}) already exists. Error: {e}",
|
|
RESOURCE_ALREADY_EXISTS,
|
|
)
|
|
|
|
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
|
|
"""
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
sql_registered_model = self._get_registered_model(session, name)
|
|
session.delete(sql_registered_model)
|
|
|
|
def _compute_next_token(self, max_results_for_query, current_size, offset, max_results):
|
|
next_token = None
|
|
if max_results_for_query == current_size:
|
|
final_offset = offset + max_results
|
|
next_token = SearchUtils.create_page_token(final_offset)
|
|
return next_token
|
|
|
|
def search_registered_models(
|
|
self,
|
|
filter_string=None,
|
|
max_results=SEARCH_REGISTERED_MODEL_MAX_RESULTS_DEFAULT,
|
|
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.
|
|
"""
|
|
if max_results > SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD:
|
|
raise MlflowException(
|
|
"Invalid value for request parameter max_results. It must be at most "
|
|
f"{SEARCH_REGISTERED_MODEL_MAX_RESULTS_THRESHOLD}, but got value {max_results}",
|
|
INVALID_PARAMETER_VALUE,
|
|
)
|
|
|
|
parsed_filters = SearchModelUtils.parse_search_filter(filter_string)
|
|
parsed_orderby = self._parse_search_registered_models_order_by(order_by)
|
|
offset = SearchUtils.parse_start_offset_from_page_token(page_token)
|
|
# we query for max_results + 1 items to check whether there is another page to return.
|
|
# this remediates having to make another query which returns no items.
|
|
max_results_for_query = max_results + 1
|
|
|
|
with self.ManagedSessionMaker() as session:
|
|
filter_query = self._get_search_registered_model_filter_query(
|
|
session, parsed_filters, self.engine.dialect.name
|
|
)
|
|
query = (
|
|
filter_query
|
|
.options(*self._get_eager_registered_model_query_options())
|
|
.order_by(*parsed_orderby)
|
|
.limit(max_results_for_query)
|
|
)
|
|
if page_token:
|
|
query = query.offset(offset)
|
|
sql_registered_models = session.execute(query).scalars().all()
|
|
next_page_token = self._compute_next_token(
|
|
max_results_for_query, len(sql_registered_models), offset, max_results
|
|
)
|
|
|
|
# Batch-fetch latest versions for all models to avoid N+1 queries
|
|
model_names = [rm.name for rm in sql_registered_models[:max_results]]
|
|
latest_versions_map = self._get_latest_versions_for_models(session, model_names)
|
|
|
|
rm_entities = [
|
|
rm.to_mlflow_entity(preloaded_latest_versions=latest_versions_map.get(rm.name))
|
|
for rm in sql_registered_models[:max_results]
|
|
]
|
|
return PagedList(rm_entities, next_page_token)
|
|
|
|
def _get_search_registered_model_filter_query(self, session, parsed_filters, dialect):
|
|
attribute_filters = []
|
|
tag_filters = {}
|
|
tag_where_clauses = self._get_workspace_clauses(SqlRegisteredModelTag)
|
|
for f in parsed_filters:
|
|
type_ = f["type"]
|
|
key = f["key"]
|
|
comparator = f["comparator"]
|
|
value = f["value"]
|
|
if type_ == "attribute":
|
|
if key != "name":
|
|
raise MlflowException(
|
|
f"Invalid attribute name: {key}", error_code=INVALID_PARAMETER_VALUE
|
|
)
|
|
if comparator not in ("=", "!=", "LIKE", "ILIKE"):
|
|
raise MlflowException(
|
|
f"Invalid comparator for attribute: {comparator}",
|
|
error_code=INVALID_PARAMETER_VALUE,
|
|
)
|
|
attr = getattr(SqlRegisteredModel, key)
|
|
attr_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)(attr, value)
|
|
attribute_filters.append(attr_filter)
|
|
elif type_ == "tag":
|
|
if comparator not in ("=", "!=", "LIKE", "ILIKE"):
|
|
raise MlflowException.invalid_parameter_value(
|
|
f"Invalid comparator for tag: {comparator}"
|
|
)
|
|
if key not in tag_filters:
|
|
key_filter = SearchUtils.get_sql_comparison_func("=", dialect)(
|
|
SqlRegisteredModelTag.key, key
|
|
)
|
|
tag_filters[key] = [key_filter]
|
|
tag_filters[key].extend(tag_where_clauses)
|
|
|
|
val_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)(
|
|
SqlRegisteredModelTag.value, value
|
|
)
|
|
tag_filters[key].append(val_filter)
|
|
else:
|
|
raise MlflowException(
|
|
f"Invalid token type: {type_}", error_code=INVALID_PARAMETER_VALUE
|
|
)
|
|
|
|
attribute_filters.extend(self._get_workspace_clauses(SqlRegisteredModel))
|
|
|
|
rm_query = select(SqlRegisteredModel).filter(*attribute_filters)
|
|
|
|
if not self._is_querying_prompt(parsed_filters):
|
|
rm_query = self._update_query_to_exclude_prompts(
|
|
rm_query,
|
|
tag_filters,
|
|
dialect,
|
|
SqlRegisteredModel,
|
|
SqlRegisteredModelTag,
|
|
)
|
|
|
|
if tag_filters:
|
|
sql_tag_filters = (sqlalchemy.and_(*x) for x in tag_filters.values())
|
|
tag_filter_query = (
|
|
select(SqlRegisteredModelTag.workspace, SqlRegisteredModelTag.name)
|
|
.filter(sqlalchemy.or_(*sql_tag_filters))
|
|
.group_by(SqlRegisteredModelTag.workspace, SqlRegisteredModelTag.name)
|
|
.having(sqlalchemy.func.count(sqlalchemy.literal(1)) == len(tag_filters))
|
|
.subquery()
|
|
)
|
|
|
|
return rm_query.join(
|
|
tag_filter_query,
|
|
sqlalchemy.and_(
|
|
SqlRegisteredModel.workspace == tag_filter_query.c.workspace,
|
|
SqlRegisteredModel.name == tag_filter_query.c.name,
|
|
),
|
|
)
|
|
else:
|
|
return rm_query
|
|
|
|
def _get_search_model_versions_filter_clauses(self, parsed_filters, dialect):
|
|
attribute_filters = []
|
|
tag_filters = {}
|
|
tag_where_clauses = self._get_workspace_clauses(SqlModelVersionTag)
|
|
for f in parsed_filters:
|
|
type_ = f["type"]
|
|
key = f["key"]
|
|
comparator = f["comparator"]
|
|
value = f["value"]
|
|
if type_ == "attribute":
|
|
if key not in SearchModelVersionUtils.VALID_SEARCH_ATTRIBUTE_KEYS:
|
|
raise MlflowException(
|
|
f"Invalid attribute name: {key}", error_code=INVALID_PARAMETER_VALUE
|
|
)
|
|
if key in SearchModelVersionUtils.NUMERIC_ATTRIBUTES:
|
|
if (
|
|
comparator
|
|
not in SearchModelVersionUtils.VALID_NUMERIC_ATTRIBUTE_COMPARATORS
|
|
):
|
|
raise MlflowException(
|
|
f"Invalid comparator for attribute {key}: {comparator}",
|
|
error_code=INVALID_PARAMETER_VALUE,
|
|
)
|
|
elif (
|
|
comparator not in SearchModelVersionUtils.VALID_STRING_ATTRIBUTE_COMPARATORS
|
|
or (comparator == "IN" and key != "run_id")
|
|
):
|
|
raise MlflowException(
|
|
f"Invalid comparator for attribute: {comparator}",
|
|
error_code=INVALID_PARAMETER_VALUE,
|
|
)
|
|
if key == "source_path":
|
|
key_name = "source"
|
|
elif key == "version_number":
|
|
key_name = "version"
|
|
else:
|
|
key_name = key
|
|
attr = getattr(SqlModelVersion, key_name)
|
|
if comparator == "IN":
|
|
# Note: Here the run_id values in databases contain only lower case letters,
|
|
# so we already filter out comparison values containing upper case letters
|
|
# in `SearchModelUtils._get_value`. This addresses MySQL IN clause case
|
|
# in-sensitive issue.
|
|
val_filter = attr.in_(value)
|
|
else:
|
|
val_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)(
|
|
attr, value
|
|
)
|
|
attribute_filters.append(val_filter)
|
|
elif type_ == "tag":
|
|
if comparator not in ("=", "!=", "LIKE", "ILIKE"):
|
|
raise MlflowException.invalid_parameter_value(
|
|
f"Invalid comparator for tag: {comparator}",
|
|
)
|
|
if key not in tag_filters:
|
|
key_filter = SearchUtils.get_sql_comparison_func("=", dialect)(
|
|
SqlModelVersionTag.key, key
|
|
)
|
|
tag_filters[key] = [key_filter]
|
|
tag_filters[key].extend(tag_where_clauses)
|
|
|
|
val_filter = SearchUtils.get_sql_comparison_func(comparator, dialect)(
|
|
SqlModelVersionTag.value, value
|
|
)
|
|
tag_filters[key].append(val_filter)
|
|
else:
|
|
raise MlflowException(
|
|
f"Invalid token type: {type_}", error_code=INVALID_PARAMETER_VALUE
|
|
)
|
|
|
|
attribute_filters.extend(self._get_workspace_clauses(SqlModelVersion))
|
|
|
|
mv_query = select(SqlModelVersion).filter(*attribute_filters)
|
|
|
|
if not self._is_querying_prompt(parsed_filters):
|
|
mv_query = self._update_query_to_exclude_prompts(
|
|
mv_query,
|
|
tag_filters,
|
|
dialect,
|
|
SqlModelVersion,
|
|
SqlModelVersionTag,
|
|
)
|
|
|
|
if tag_filters:
|
|
sql_tag_filters = (sqlalchemy.and_(*x) for x in tag_filters.values())
|
|
tag_filter_query = (
|
|
select(
|
|
SqlModelVersionTag.workspace,
|
|
SqlModelVersionTag.name,
|
|
SqlModelVersionTag.version,
|
|
)
|
|
.filter(sqlalchemy.or_(*sql_tag_filters))
|
|
.group_by(
|
|
SqlModelVersionTag.workspace,
|
|
SqlModelVersionTag.name,
|
|
SqlModelVersionTag.version,
|
|
)
|
|
.having(sqlalchemy.func.count(sqlalchemy.literal(1)) == len(tag_filters))
|
|
.subquery()
|
|
)
|
|
return mv_query.join(
|
|
tag_filter_query,
|
|
sqlalchemy.and_(
|
|
SqlModelVersion.workspace == tag_filter_query.c.workspace,
|
|
SqlModelVersion.name == tag_filter_query.c.name,
|
|
SqlModelVersion.version == tag_filter_query.c.version,
|
|
),
|
|
)
|
|
else:
|
|
return mv_query
|
|
|
|
def _update_query_to_exclude_prompts(
|
|
self,
|
|
query: Any,
|
|
tag_filters: dict[str, list[Any]],
|
|
dialect: str,
|
|
main_db_model: SqlModelVersion | SqlRegisteredModel,
|
|
tag_db_model: SqlModelVersionTag | SqlRegisteredModelTag,
|
|
):
|
|
"""
|
|
Update query to exclude all prompt rows and return only normal model or model versions.
|
|
|
|
Prompts and normal models are distinguished by the `mlflow.prompt.is_prompt` tag.
|
|
The search API should only return normal models by default. However, simply filtering
|
|
rows using the tag like this does not work because models do not have the prompt tag.
|
|
|
|
tags.`mlflow.prompt.is_prompt` != 'true'
|
|
tags.`mlflow.prompt.is_prompt` = 'false'
|
|
|
|
To workaround this, we need to use a subquery to get all prompt rows and then use an
|
|
anti-join for excluding prompts.
|
|
"""
|
|
# If the tag filter contains the prompt tag, remove it
|
|
tag_filters.pop(IS_PROMPT_TAG_KEY, [])
|
|
|
|
# Filter to get all prompt rows
|
|
equal = SearchUtils.get_sql_comparison_func("=", dialect)
|
|
prompts_subquery = (
|
|
select(tag_db_model.workspace, tag_db_model.name)
|
|
.filter(
|
|
equal(tag_db_model.key, IS_PROMPT_TAG_KEY),
|
|
equal(tag_db_model.value, "true"),
|
|
*self._get_workspace_clauses(tag_db_model),
|
|
)
|
|
.group_by(tag_db_model.workspace, tag_db_model.name)
|
|
.subquery()
|
|
)
|
|
return query.join(
|
|
prompts_subquery,
|
|
sqlalchemy.and_(
|
|
main_db_model.workspace == prompts_subquery.c.workspace,
|
|
main_db_model.name == prompts_subquery.c.name,
|
|
),
|
|
isouter=True,
|
|
).filter(prompts_subquery.c.name.is_(None))
|
|
|
|
@classmethod
|
|
def _is_querying_prompt(cls, parsed_filters: list[dict[str, Any]]) -> bool:
|
|
for f in parsed_filters:
|
|
if f["type"] != "tag" or f["key"] != IS_PROMPT_TAG_KEY:
|
|
continue
|
|
|
|
return (f["comparator"] == "=" and f["value"].lower() == "true") or (
|
|
f["comparator"] == "!=" and f["value"].lower() == "false"
|
|
)
|
|
|
|
# Query should return only normal models by default
|
|
return False
|
|
|
|
@classmethod
|
|
def _parse_search_registered_models_order_by(cls, order_by_list):
|
|
"""Sorts a set of registered models based on their natural ordering and an overriding set
|
|
of order_bys. Registered models are naturally ordered first by name ascending.
|
|
"""
|
|
clauses = []
|
|
observed_order_by_clauses = set()
|
|
if order_by_list:
|
|
for order_by_clause in order_by_list:
|
|
(
|
|
attribute_token,
|
|
ascending,
|
|
) = SearchUtils.parse_order_by_for_search_registered_models(order_by_clause)
|
|
if attribute_token == SqlRegisteredModel.name.key:
|
|
field = SqlRegisteredModel.name
|
|
elif attribute_token in SearchUtils.VALID_TIMESTAMP_ORDER_BY_KEYS:
|
|
field = SqlRegisteredModel.last_updated_time
|
|
else:
|
|
raise MlflowException(
|
|
f"Invalid order by key '{attribute_token}' specified."
|
|
+ "Valid keys are "
|
|
+ f"'{SearchUtils.RECOMMENDED_ORDER_BY_KEYS_REGISTERED_MODELS}'",
|
|
error_code=INVALID_PARAMETER_VALUE,
|
|
)
|
|
if field.key in observed_order_by_clauses:
|
|
raise MlflowException(f"`order_by` contains duplicate fields: {order_by_list}")
|
|
observed_order_by_clauses.add(field.key)
|
|
if ascending:
|
|
clauses.append(field.asc())
|
|
else:
|
|
clauses.append(field.desc())
|
|
|
|
if SqlRegisteredModel.name.key not in observed_order_by_clauses:
|
|
clauses.append(SqlRegisteredModel.name.asc())
|
|
return clauses
|
|
|
|
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.
|
|
"""
|
|
with self.ManagedSessionMaker() as session:
|
|
return self._get_registered_model(session, name, eager=True).to_mlflow_entity()
|
|
|
|
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.
|
|
|
|
"""
|
|
with self.ManagedSessionMaker() as session:
|
|
sql_registered_model = self._get_registered_model(session, name)
|
|
# Convert to RegisteredModel entity first and then extract latest_versions
|
|
latest_versions = sql_registered_model.to_mlflow_entity().latest_versions
|
|
if stages is None or len(stages) == 0:
|
|
expected_stages = {get_canonical_stage(stage) for stage in ALL_STAGES}
|
|
else:
|
|
expected_stages = {get_canonical_stage(stage) for stage in stages}
|
|
mvs = [mv for mv in latest_versions if mv.current_stage in expected_stages]
|
|
|
|
# Populate aliases for each model version
|
|
for mv in mvs:
|
|
model_aliases = sql_registered_model.registered_model_aliases
|
|
mv.aliases = [alias.alias for alias in model_aliases if alias.version == mv.version]
|
|
|
|
return mvs
|
|
|
|
def _get_registered_model_tag(self, session, name, key):
|
|
tags = (
|
|
self
|
|
._get_query(session, SqlRegisteredModelTag)
|
|
.filter(
|
|
SqlRegisteredModelTag.name == name,
|
|
SqlRegisteredModelTag.key == key,
|
|
)
|
|
.all()
|
|
)
|
|
if len(tags) == 0:
|
|
return None
|
|
if len(tags) > 1:
|
|
raise MlflowException(
|
|
f"Expected only 1 registered model tag with name={name}, key={key}. "
|
|
f"Found {len(tags)}.",
|
|
INVALID_STATE,
|
|
)
|
|
return tags[0]
|
|
|
|
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
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_registered_model_tag(tag.key, tag.value)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
# check if registered model exists
|
|
sql_registered_model = self._get_registered_model(session, name)
|
|
session.merge(
|
|
SqlRegisteredModelTag(
|
|
workspace=sql_registered_model.workspace,
|
|
name=name,
|
|
key=tag.key,
|
|
value=tag.value,
|
|
)
|
|
)
|
|
|
|
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
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_tag_name(key)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
# check if registered model exists
|
|
self._get_registered_model(session, name)
|
|
existing_tag = self._get_registered_model_tag(session, name, key)
|
|
if existing_tag is not None:
|
|
session.delete(existing_tag)
|
|
|
|
# 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,
|
|
model_id: str | None = 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.
|
|
model_id: The ID of the model (from an Experiment) that is being promoted to a
|
|
registered model version, if applicable.
|
|
|
|
Returns:
|
|
A single object of :py:class:`mlflow.entities.model_registry.ModelVersion`
|
|
created in the backend.
|
|
|
|
"""
|
|
|
|
_validate_model_name(name)
|
|
for tag in tags or []:
|
|
_validate_model_version_tag(tag.key, tag.value)
|
|
storage_location = source
|
|
if urllib.parse.urlparse(source).scheme == "models":
|
|
parsed_model_uri = _parse_model_uri(source)
|
|
try:
|
|
if parsed_model_uri.model_id is not None:
|
|
# TODO: Propagate tracking URI to file sqlalchemy directly, rather than relying
|
|
# on global URI (individual MlflowClient instances may have different tracking
|
|
# URIs)
|
|
model = MlflowClient().get_logged_model(parsed_model_uri.model_id)
|
|
storage_location = model.artifact_location
|
|
run_id = run_id or model.source_run_id
|
|
else:
|
|
storage_location = self.get_model_version_download_uri(
|
|
parsed_model_uri.name, parsed_model_uri.version
|
|
)
|
|
except Exception as e:
|
|
raise MlflowException(
|
|
f"Unable to fetch model from model URI source artifact location '{source}'."
|
|
f"Error: {e}"
|
|
) from e
|
|
|
|
if not run_id and model_id:
|
|
model = MlflowClient().get_logged_model(model_id)
|
|
run_id = model.source_run_id
|
|
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
creation_time = get_current_time_millis()
|
|
for attempt in range(self.CREATE_MODEL_VERSION_RETRIES):
|
|
try:
|
|
sql_registered_model = self._get_registered_model(session, name)
|
|
sql_registered_model.last_updated_time = creation_time
|
|
max_version = (
|
|
session
|
|
.query(sqlalchemy.func.max(SqlModelVersion.version))
|
|
.filter(
|
|
SqlModelVersion.name == name,
|
|
*self._get_workspace_clauses(SqlModelVersion),
|
|
)
|
|
.scalar()
|
|
)
|
|
version = (max_version or 0) + 1
|
|
model_version = self._with_workspace_field(
|
|
SqlModelVersion(
|
|
name=name,
|
|
version=version,
|
|
creation_time=creation_time,
|
|
last_updated_time=creation_time,
|
|
source=source,
|
|
storage_location=storage_location,
|
|
run_id=run_id,
|
|
run_link=run_link,
|
|
description=description,
|
|
)
|
|
)
|
|
tags_dict = {}
|
|
for tag in tags or []:
|
|
tags_dict[tag.key] = tag.value
|
|
model_version.model_version_tags = [
|
|
self._with_workspace_field(
|
|
SqlModelVersionTag(name=name, version=version, key=key, value=value)
|
|
)
|
|
for key, value in tags_dict.items()
|
|
]
|
|
session.add_all([sql_registered_model, model_version])
|
|
session.flush()
|
|
return self._populate_model_version_aliases(
|
|
session, name, model_version.to_mlflow_entity()
|
|
)
|
|
except sqlalchemy.exc.IntegrityError:
|
|
session.rollback()
|
|
more_retries = self.CREATE_MODEL_VERSION_RETRIES - attempt - 1
|
|
_logger.info(
|
|
"Model Version creation error (name=%s) Retrying %s more time%s.",
|
|
name,
|
|
str(more_retries),
|
|
"s" if more_retries > 1 else "",
|
|
)
|
|
raise MlflowException(
|
|
f"Model Version creation error (name={name}). Giving up after "
|
|
f"{self.CREATE_MODEL_VERSION_RETRIES} attempts."
|
|
)
|
|
|
|
def _populate_model_version_aliases(self, session, name, version):
|
|
model_aliases = self._get_registered_model(session, name).registered_model_aliases
|
|
version.aliases = [
|
|
alias.alias for alias in model_aliases if alias.version == version.version
|
|
]
|
|
return version
|
|
|
|
def _get_model_version_from_db(self, session, name, version, conditions, query_options=None):
|
|
if query_options is None:
|
|
query_options = []
|
|
versions = (
|
|
self
|
|
._get_query(session, SqlModelVersion)
|
|
.options(*query_options)
|
|
.filter(*conditions)
|
|
.all()
|
|
)
|
|
|
|
if len(versions) == 0:
|
|
raise MlflowException(
|
|
f"Model Version (name={name}, version={version}) not found",
|
|
RESOURCE_DOES_NOT_EXIST,
|
|
)
|
|
if len(versions) > 1:
|
|
raise MlflowException(
|
|
f"Expected only 1 model version with (name={name}, version={version}). "
|
|
f"Found {len(versions)}.",
|
|
INVALID_STATE,
|
|
)
|
|
return versions[0]
|
|
|
|
def _get_sql_model_version(self, session, name, version, eager=False):
|
|
"""
|
|
Args:
|
|
eager: If ``True``, eagerly loads the model version's tags.
|
|
If ``False``, these attributes are not eagerly loaded and
|
|
will be loaded when their corresponding object properties
|
|
are accessed from the resulting ``SqlModelVersion`` object.
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_model_version(version)
|
|
query_options = self._get_eager_model_version_query_options() if eager else []
|
|
conditions = [
|
|
SqlModelVersion.name == name,
|
|
SqlModelVersion.version == version,
|
|
SqlModelVersion.current_stage != STAGE_DELETED_INTERNAL,
|
|
]
|
|
return self._get_model_version_from_db(session, name, version, conditions, query_options)
|
|
|
|
def _get_sql_model_version_including_deleted(self, name, version):
|
|
"""
|
|
Private method to retrieve model versions including those that are internally deleted.
|
|
Used in tests to verify redaction behavior on deletion.
|
|
|
|
Args:
|
|
name: Registered model name.
|
|
version: Registered model version.
|
|
|
|
Returns:
|
|
A single :py:class:`mlflow.entities.model_registry.ModelVersion` object.
|
|
"""
|
|
with self.ManagedSessionMaker() as session:
|
|
conditions = [
|
|
SqlModelVersion.name == name,
|
|
SqlModelVersion.version == version,
|
|
]
|
|
sql_model_version = self._get_model_version_from_db(session, name, version, conditions)
|
|
return self._populate_model_version_aliases(
|
|
session, name, sql_model_version.to_mlflow_entity()
|
|
)
|
|
|
|
def update_model_version(self, name, version, description=None):
|
|
"""
|
|
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.
|
|
|
|
"""
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
updated_time = get_current_time_millis()
|
|
sql_model_version = self._get_sql_model_version(session, name=name, version=version)
|
|
sql_model_version.description = description
|
|
sql_model_version.last_updated_time = updated_time
|
|
session.add(sql_model_version)
|
|
return self._populate_model_version_aliases(
|
|
session, name, sql_model_version.to_mlflow_entity()
|
|
)
|
|
|
|
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.
|
|
|
|
"""
|
|
is_active_stage = get_canonical_stage(stage) in DEFAULT_STAGES_FOR_GET_LATEST_VERSIONS
|
|
if archive_existing_versions and not is_active_stage:
|
|
msg_tpl = (
|
|
"Model version transition cannot archive existing model versions "
|
|
"because '{}' is not an Active stage. Valid stages are {}"
|
|
)
|
|
raise MlflowException(msg_tpl.format(stage, DEFAULT_STAGES_FOR_GET_LATEST_VERSIONS))
|
|
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
last_updated_time = get_current_time_millis()
|
|
|
|
model_versions = []
|
|
if archive_existing_versions:
|
|
conditions = [
|
|
SqlModelVersion.name == name,
|
|
SqlModelVersion.version != version,
|
|
SqlModelVersion.current_stage == get_canonical_stage(stage),
|
|
]
|
|
model_versions = self._get_query(session, SqlModelVersion).filter(*conditions).all()
|
|
for mv in model_versions:
|
|
mv.current_stage = STAGE_ARCHIVED
|
|
mv.last_updated_time = last_updated_time
|
|
|
|
sql_model_version = self._get_sql_model_version(
|
|
session=session, name=name, version=version
|
|
)
|
|
sql_model_version.current_stage = get_canonical_stage(stage)
|
|
sql_model_version.last_updated_time = last_updated_time
|
|
sql_registered_model = sql_model_version.registered_model
|
|
sql_registered_model.last_updated_time = last_updated_time
|
|
session.add_all([*model_versions, sql_model_version, sql_registered_model])
|
|
return self._populate_model_version_aliases(
|
|
session, name, sql_model_version.to_mlflow_entity()
|
|
)
|
|
|
|
def delete_model_version(self, name, version):
|
|
"""
|
|
Delete model version in backend.
|
|
|
|
Args:
|
|
name: Registered model name.
|
|
version: Registered model version.
|
|
|
|
Returns:
|
|
None
|
|
"""
|
|
# currently delete model version still keeps the tags associated with the version
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
updated_time = get_current_time_millis()
|
|
sql_model_version = self._get_sql_model_version(session, name, version)
|
|
sql_registered_model = sql_model_version.registered_model
|
|
sql_registered_model.last_updated_time = updated_time
|
|
aliases = sql_registered_model.registered_model_aliases
|
|
for alias in aliases:
|
|
if alias.version == version:
|
|
session.delete(alias)
|
|
sql_model_version.current_stage = STAGE_DELETED_INTERNAL
|
|
sql_model_version.last_updated_time = updated_time
|
|
sql_model_version.description = None
|
|
sql_model_version.user_id = None
|
|
sql_model_version.source = "REDACTED-SOURCE-PATH"
|
|
sql_model_version.run_id = "REDACTED-RUN-ID"
|
|
sql_model_version.run_link = "REDACTED-RUN-LINK"
|
|
sql_model_version.status_message = None
|
|
session.add_all([sql_registered_model, sql_model_version])
|
|
|
|
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.
|
|
"""
|
|
with self.ManagedSessionMaker() as session:
|
|
sql_model_version = self._get_sql_model_version(session, name, version, eager=True)
|
|
return self._populate_model_version_aliases(
|
|
session, name, sql_model_version.to_mlflow_entity()
|
|
)
|
|
|
|
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.
|
|
"""
|
|
with self.ManagedSessionMaker() as session:
|
|
sql_model_version = self._get_sql_model_version(session, name, version)
|
|
return sql_model_version.storage_location or sql_model_version.source
|
|
|
|
def search_model_versions(
|
|
self,
|
|
filter_string=None,
|
|
max_results=SEARCH_MODEL_VERSION_MAX_RESULTS_DEFAULT,
|
|
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.
|
|
|
|
"""
|
|
if not isinstance(max_results, int) or max_results < 1:
|
|
raise MlflowException(
|
|
"Invalid value for max_results. It must be a positive integer,"
|
|
f" but got {max_results}",
|
|
INVALID_PARAMETER_VALUE,
|
|
)
|
|
|
|
if max_results > SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD:
|
|
raise MlflowException(
|
|
"Invalid value for request parameter max_results. It must be at most "
|
|
f"{SEARCH_MODEL_VERSION_MAX_RESULTS_THRESHOLD}, but got value {max_results}",
|
|
INVALID_PARAMETER_VALUE,
|
|
)
|
|
|
|
parsed_filters = SearchModelVersionUtils.parse_search_filter(filter_string)
|
|
|
|
filter_query = self._get_search_model_versions_filter_clauses(
|
|
parsed_filters, self.engine.dialect.name
|
|
)
|
|
|
|
parsed_orderby = self._parse_search_model_versions_order_by(
|
|
order_by or ["last_updated_timestamp DESC", "name ASC", "version_number DESC"]
|
|
)
|
|
offset = SearchUtils.parse_start_offset_from_page_token(page_token)
|
|
# we query for max_results + 1 items to check whether there is another page to return.
|
|
# this remediates having to make another query which returns no items.
|
|
max_results_for_query = max_results + 1
|
|
|
|
with self.ManagedSessionMaker() as session:
|
|
query = (
|
|
filter_query
|
|
.options(*self._get_eager_model_version_query_options())
|
|
.filter(SqlModelVersion.current_stage != STAGE_DELETED_INTERNAL)
|
|
.order_by(*parsed_orderby)
|
|
.limit(max_results_for_query)
|
|
)
|
|
if page_token:
|
|
query = query.offset(offset)
|
|
sql_model_versions = session.execute(query).scalars().all()
|
|
next_page_token = self._compute_next_token(
|
|
max_results_for_query, len(sql_model_versions), offset, max_results
|
|
)
|
|
model_versions = [mv.to_mlflow_entity() for mv in sql_model_versions][:max_results]
|
|
return PagedList(model_versions, next_page_token)
|
|
|
|
@classmethod
|
|
def _parse_search_model_versions_order_by(cls, order_by_list):
|
|
"""Sorts a set of model versions based on their natural ordering and an overriding set
|
|
of order_bys. Model versions are naturally ordered first by name ascending, then by
|
|
version ascending.
|
|
"""
|
|
clauses = []
|
|
observed_order_by_clauses = set()
|
|
if order_by_list:
|
|
for order_by_clause in order_by_list:
|
|
(
|
|
_,
|
|
key,
|
|
ascending,
|
|
) = SearchModelVersionUtils.parse_order_by_for_search_model_versions(
|
|
order_by_clause
|
|
)
|
|
if key not in SearchModelVersionUtils.VALID_ORDER_BY_ATTRIBUTE_KEYS:
|
|
raise MlflowException(
|
|
f"Invalid order by key '{key}' specified. "
|
|
"Valid keys are "
|
|
f"{SearchModelVersionUtils.VALID_ORDER_BY_ATTRIBUTE_KEYS}",
|
|
error_code=INVALID_PARAMETER_VALUE,
|
|
)
|
|
else:
|
|
if key == "version_number":
|
|
field = SqlModelVersion.version
|
|
elif key == "creation_timestamp":
|
|
field = SqlModelVersion.creation_time
|
|
elif key == "last_updated_timestamp":
|
|
field = SqlModelVersion.last_updated_time
|
|
else:
|
|
field = getattr(SqlModelVersion, key)
|
|
if field.key in observed_order_by_clauses:
|
|
raise MlflowException(f"`order_by` contains duplicate fields: {order_by_list}")
|
|
observed_order_by_clauses.add(field.key)
|
|
if ascending:
|
|
clauses.append(field.asc())
|
|
else:
|
|
clauses.append(field.desc())
|
|
|
|
if SqlModelVersion.name.key not in observed_order_by_clauses:
|
|
clauses.append(SqlModelVersion.name.asc())
|
|
if SqlModelVersion.version.key not in observed_order_by_clauses:
|
|
clauses.append(SqlModelVersion.version.desc())
|
|
return clauses
|
|
|
|
def _get_model_version_tag(self, session, name, version, key):
|
|
tags = (
|
|
self
|
|
._get_query(session, SqlModelVersionTag)
|
|
.filter(
|
|
SqlModelVersionTag.name == name,
|
|
SqlModelVersionTag.version == version,
|
|
SqlModelVersionTag.key == key,
|
|
)
|
|
.all()
|
|
)
|
|
if len(tags) == 0:
|
|
return None
|
|
if len(tags) > 1:
|
|
raise MlflowException(
|
|
f"Expected only 1 model version tag with name={name}, version={version}, "
|
|
f"key={key}. Found {len(tags)}.",
|
|
INVALID_STATE,
|
|
)
|
|
return tags[0]
|
|
|
|
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
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_model_version(version)
|
|
_validate_model_version_tag(tag.key, tag.value)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
# check if model version exists
|
|
sql_model_version = self._get_sql_model_version(session, name, version)
|
|
session.merge(
|
|
SqlModelVersionTag(
|
|
workspace=sql_model_version.workspace,
|
|
name=name,
|
|
version=version,
|
|
key=tag.key,
|
|
value=tag.value,
|
|
)
|
|
)
|
|
|
|
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
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_model_version(version)
|
|
_validate_tag_name(key)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
# check if model version exists
|
|
self._get_sql_model_version(session, name, version)
|
|
existing_tag = self._get_model_version_tag(session, name, version, key)
|
|
if existing_tag is not None:
|
|
session.delete(existing_tag)
|
|
|
|
def _get_registered_model_alias(self, session, name, alias):
|
|
return (
|
|
self
|
|
._get_query(session, SqlRegisteredModelAlias)
|
|
.filter(
|
|
SqlRegisteredModelAlias.name == name,
|
|
SqlRegisteredModelAlias.alias == alias,
|
|
)
|
|
.first()
|
|
)
|
|
|
|
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
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_model_alias_name(alias)
|
|
_validate_model_alias_name_reserved(alias)
|
|
_validate_model_version(version)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
# check if model version exists
|
|
sql_model_version = self._get_sql_model_version(session, name, version)
|
|
session.merge(
|
|
SqlRegisteredModelAlias(
|
|
workspace=sql_model_version.workspace,
|
|
name=name,
|
|
alias=alias,
|
|
version=version,
|
|
)
|
|
)
|
|
|
|
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
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_model_alias_name(alias)
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
# check if registered model exists
|
|
self._get_registered_model(session, name)
|
|
existing_alias = self._get_registered_model_alias(session, name, alias)
|
|
if existing_alias is not None:
|
|
session.delete(existing_alias)
|
|
|
|
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.
|
|
"""
|
|
_validate_model_name(name)
|
|
_validate_model_alias_name(alias)
|
|
|
|
if alias.lower() == _REGISTERED_MODEL_ALIAS_LATEST:
|
|
if versions := self.get_latest_versions(name):
|
|
return versions[0]
|
|
else:
|
|
raise MlflowException(
|
|
f"Latest version not found for model {name}.", RESOURCE_DOES_NOT_EXIST
|
|
)
|
|
|
|
with self.ManagedSessionMaker() as session:
|
|
# check if registered model exists
|
|
self._get_registered_model(session, name)
|
|
existing_alias = self._get_registered_model_alias(session, name, alias)
|
|
if existing_alias is not None:
|
|
sql_model_version = self._get_sql_model_version(
|
|
session, existing_alias.name, existing_alias.version
|
|
)
|
|
return self._populate_model_version_aliases(
|
|
session, name, sql_model_version.to_mlflow_entity()
|
|
)
|
|
else:
|
|
raise MlflowException(
|
|
f"Registered model alias {alias} not found.", INVALID_PARAMETER_VALUE
|
|
)
|
|
|
|
def _await_model_version_creation(self, mv, await_creation_for):
|
|
"""
|
|
Does not wait for the model version to become READY as a successful creation will
|
|
immediately place the model version in a READY state.
|
|
"""
|
|
|
|
# Webhook CRUD operations
|
|
def create_webhook(
|
|
self,
|
|
name: str,
|
|
url: str,
|
|
events: list[WebhookEvent],
|
|
description: str | None = None,
|
|
secret: str | None = None,
|
|
status: WebhookStatus | None = None,
|
|
) -> Webhook:
|
|
_validate_webhook_name(name)
|
|
_validate_webhook_url(url)
|
|
_validate_webhook_events(events)
|
|
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
webhook_id = str(uuid.uuid4())
|
|
creation_time = get_current_time_millis()
|
|
webhook = self._with_workspace_field(
|
|
SqlWebhook(
|
|
webhook_id=webhook_id,
|
|
name=name,
|
|
url=url,
|
|
description=description,
|
|
secret=secret,
|
|
status=(status or WebhookStatus.ACTIVE).value,
|
|
creation_timestamp=creation_time,
|
|
last_updated_timestamp=creation_time,
|
|
)
|
|
)
|
|
session.add(webhook)
|
|
session.add_all(
|
|
SqlWebhookEvent(
|
|
webhook_id=webhook_id,
|
|
entity=e.entity.value,
|
|
action=e.action.value,
|
|
)
|
|
for e in events
|
|
)
|
|
session.flush()
|
|
return webhook.to_mlflow_entity()
|
|
|
|
def get_webhook(self, webhook_id: str) -> Webhook:
|
|
with self.ManagedSessionMaker() as session:
|
|
webhook = self._get_webhook_by_id(session, webhook_id)
|
|
return webhook.to_mlflow_entity()
|
|
|
|
def list_webhooks(
|
|
self,
|
|
max_results: int | None = None,
|
|
page_token: str | None = None,
|
|
) -> PagedList[Webhook]:
|
|
max_results = max_results or 100
|
|
if max_results < 1 or max_results > 1000:
|
|
raise MlflowException(
|
|
"max_results must be between 1 and 1000.", INVALID_PARAMETER_VALUE
|
|
)
|
|
|
|
offset = SearchUtils.parse_start_offset_from_page_token(page_token)
|
|
with self.ManagedSessionMaker() as session:
|
|
query = (
|
|
self
|
|
._get_query(session, SqlWebhook)
|
|
.filter(SqlWebhook.deleted_timestamp.is_(None))
|
|
.order_by(SqlWebhook.creation_timestamp.desc())
|
|
.limit(max_results + 1)
|
|
)
|
|
|
|
if page_token:
|
|
query = query.offset(offset)
|
|
|
|
webhooks = query.all()
|
|
|
|
# Check if there's a next page
|
|
has_next_page = len(webhooks) > max_results
|
|
next_page_token = None
|
|
if has_next_page:
|
|
webhooks = webhooks[:max_results]
|
|
next_page_token = SearchUtils.create_page_token(offset + max_results)
|
|
|
|
return PagedList([w.to_mlflow_entity() for w in webhooks], next_page_token)
|
|
|
|
def list_webhooks_by_event(
|
|
self,
|
|
event: WebhookEvent,
|
|
max_results: int | None = None,
|
|
page_token: str | None = None,
|
|
) -> PagedList[Webhook]:
|
|
max_results = max_results or 100
|
|
if max_results < 1 or max_results > 1000:
|
|
raise MlflowException(
|
|
"max_results must be between 1 and 1000.", INVALID_PARAMETER_VALUE
|
|
)
|
|
|
|
offset = SearchUtils.parse_start_offset_from_page_token(page_token)
|
|
with self.ManagedSessionMaker() as session:
|
|
# Query webhooks that have the specific event in their related webhook_events
|
|
query = (
|
|
self
|
|
._get_query(session, SqlWebhook)
|
|
.join(SqlWebhookEvent)
|
|
.filter(SqlWebhook.deleted_timestamp.is_(None))
|
|
.filter(SqlWebhookEvent.entity == event.entity.value)
|
|
.filter(SqlWebhookEvent.action == event.action.value)
|
|
.order_by(SqlWebhook.creation_timestamp.desc())
|
|
.limit(max_results + 1)
|
|
)
|
|
|
|
if page_token:
|
|
query = query.offset(offset)
|
|
|
|
webhooks = query.all()
|
|
|
|
# Check if there's a next page
|
|
has_next_page = len(webhooks) > max_results
|
|
next_page_token = None
|
|
if has_next_page:
|
|
webhooks = webhooks[:max_results]
|
|
next_page_token = SearchUtils.create_page_token(offset + max_results)
|
|
|
|
return PagedList([w.to_mlflow_entity() for w in webhooks], next_page_token)
|
|
|
|
def update_webhook(
|
|
self,
|
|
webhook_id: str,
|
|
name: str | None = None,
|
|
description: str | None = None,
|
|
url: str | None = None,
|
|
events: list[WebhookEvent] | None = None,
|
|
secret: str | None = None,
|
|
status: WebhookStatus | None = None,
|
|
) -> Webhook:
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
webhook = self._get_webhook_by_id(session, webhook_id)
|
|
|
|
# Update fields if provided
|
|
if name is not None:
|
|
_validate_webhook_name(name)
|
|
webhook.name = name
|
|
if url is not None:
|
|
_validate_webhook_url(url)
|
|
webhook.url = url
|
|
if events is not None:
|
|
_validate_webhook_events(events)
|
|
# Delete existing webhook events
|
|
session.query(SqlWebhookEvent).filter(
|
|
SqlWebhookEvent.webhook_id == webhook_id
|
|
).delete()
|
|
# Create new webhook events
|
|
session.add_all(
|
|
SqlWebhookEvent(
|
|
webhook_id=webhook_id, entity=e.entity.value, action=e.action.value
|
|
)
|
|
for e in events
|
|
)
|
|
if description is not None:
|
|
webhook.description = description
|
|
if secret is not None:
|
|
webhook.secret = secret
|
|
if status is not None:
|
|
webhook.status = status.value
|
|
|
|
webhook.last_updated_timestamp = get_current_time_millis()
|
|
|
|
session.add(webhook)
|
|
session.flush()
|
|
|
|
return webhook.to_mlflow_entity()
|
|
|
|
def delete_webhook(self, webhook_id: str) -> None:
|
|
with self.ManagedSessionMaker(read_only=False) as session:
|
|
webhook = self._get_webhook_by_id(session, webhook_id)
|
|
|
|
# Soft delete by setting deleted_timestamp
|
|
webhook.deleted_timestamp = get_current_time_millis()
|
|
webhook.last_updated_timestamp = webhook.deleted_timestamp
|
|
|
|
session.add(webhook)
|
|
session.flush()
|
|
|
|
# Helper methods for webhooks
|
|
def _get_webhook_by_id(self, session: Session, webhook_id: str) -> SqlWebhook:
|
|
if webhook := (
|
|
self
|
|
._get_query(session, SqlWebhook)
|
|
.filter(
|
|
SqlWebhook.webhook_id == webhook_id,
|
|
SqlWebhook.deleted_timestamp.is_(None),
|
|
)
|
|
.first()
|
|
):
|
|
return webhook
|
|
|
|
raise MlflowException(f"Webhook with ID {webhook_id} not found.", RESOURCE_DOES_NOT_EXIST)
|