Files
zenml/venv/lib/python3.9/site-packages/mlflow/recipes/steps/register.py
Christian Mantha 2ca0b9ef7c star
2026-03-02 19:10:52 -05:00

206 lines
7.4 KiB
Python

import logging
from pathlib import Path
from typing import Any
import mlflow
from mlflow.entities import SourceType
from mlflow.exceptions import INVALID_PARAMETER_VALUE, MlflowException
from mlflow.recipes.artifacts import ModelVersionArtifact, RegisteredModelVersionInfo
from mlflow.recipes.cards import BaseCard
from mlflow.recipes.step import BaseStep, StepClass
from mlflow.recipes.steps.train import TrainStep
from mlflow.recipes.utils.execution import get_step_output_path
from mlflow.recipes.utils.tracking import (
TrackingConfig,
apply_recipe_tracking_config,
get_recipe_tracking_config,
)
from mlflow.tracking._model_registry import DEFAULT_AWAIT_MAX_SLEEP_SECONDS
from mlflow.utils.databricks_utils import (
get_databricks_env_vars,
get_databricks_model_version_url,
get_databricks_run_url,
)
from mlflow.utils.mlflow_tags import MLFLOW_RECIPE_TEMPLATE_NAME, MLFLOW_SOURCE_TYPE
_logger = logging.getLogger(__name__)
_REGISTERED_MV_INFO_FILE = "registered_model_version.json"
class RegisterStep(BaseStep):
def __init__(self, step_config: dict[str, Any], recipe_root: str):
super().__init__(step_config, recipe_root)
self.tracking_config = TrackingConfig.from_dict(self.step_config)
def _validate_and_apply_step_config(self):
self.num_dropped_rows = None
self.model_uri = None
self.model_details = None
self.version = None
self.register_model_name = self.step_config.get("model_name")
if self.register_model_name is None:
raise MlflowException(
"Missing 'model_name' config in register step config.",
error_code=INVALID_PARAMETER_VALUE,
)
self.allow_non_validated_model = self.step_config.get("allow_non_validated_model", False)
self.registry_uri = self.step_config.get("registry_uri", None)
def _run(self, output_directory):
apply_recipe_tracking_config(self.tracking_config)
run_id_path = get_step_output_path(
recipe_root_path=self.recipe_root,
step_name="train",
relative_path="run_id",
)
run_id = Path(run_id_path).read_text()
model_validation_path = get_step_output_path(
recipe_root_path=self.recipe_root,
step_name="evaluate",
relative_path="model_validation_status",
)
model_validation = Path(model_validation_path).read_text()
artifact_path = "train/model"
tags = {
MLFLOW_SOURCE_TYPE: SourceType.to_string(SourceType.RECIPE),
MLFLOW_RECIPE_TEMPLATE_NAME: self.step_config["recipe"],
}
self.model_uri = f"runs:/{run_id}/{artifact_path}"
if model_validation == "VALIDATED" or (
model_validation == "UNKNOWN" and self.allow_non_validated_model
):
if self.registry_uri:
mlflow.set_registry_uri(self.registry_uri)
self.model_details = mlflow.register_model(
model_uri=self.model_uri,
name=self.register_model_name,
tags=tags,
await_registration_for=DEFAULT_AWAIT_MAX_SLEEP_SECONDS,
)
self.version = self.model_details.version
registered_model_info = RegisteredModelVersionInfo(
name=self.register_model_name, version=self.version
)
registered_model_info.to_json(
path=str(Path(output_directory) / _REGISTERED_MV_INFO_FILE)
)
else:
raise MlflowException(
f"Model registration on {self.model_uri} failed because it "
"is not validated. Bypass by setting allow_non_validated_model to True. "
)
card = self._build_card(run_id)
card.save_as_html(output_directory)
self._log_step_card(run_id, self.name)
return card
def _build_card(self, run_id: str) -> BaseCard:
card = BaseCard(self.recipe_name, self.name)
card_tab = card.add_tab(
"Run Summary",
"{{ MODEL_NAME }}"
+ "{{ MODEL_VERSION }}"
+ "{{ MODEL_SOURCE_URI }}"
+ "{{ ALERTS }}"
+ "{{ EXE_DURATION }}"
+ "{{ LAST_UPDATE_TIME }}",
)
if self.version is not None:
model_version_url = get_databricks_model_version_url(
registry_uri=mlflow.get_registry_uri(),
name=self.register_model_name,
version=self.version,
)
if model_version_url is not None:
card_tab.add_html(
"MODEL_NAME",
(
f"<b>Model Name:</b> <a href={model_version_url}>"
f"{self.register_model_name}</a><br><br>"
),
)
card_tab.add_html(
"MODEL_VERSION",
(
f"<b>Model Version</b> <a href={model_version_url}>"
f"{self.version}</a><br><br>"
),
)
else:
card_tab.add_markdown(
"MODEL_NAME",
f"**Model Name:** `{self.register_model_name}`",
)
card_tab.add_markdown(
"MODEL_VERSION",
f"**Model Version:** `{self.version}`",
)
model_source_url = get_databricks_run_url(
tracking_uri=mlflow.get_tracking_uri(),
run_id=run_id,
artifact_path=f"train/{TrainStep.MODEL_ARTIFACT_RELATIVE_PATH}",
)
if self.model_uri is not None and model_source_url is not None:
card_tab.add_html(
"MODEL_SOURCE_URI",
f"<b>Model Source URI</b> <a href={model_source_url}>{self.model_uri}</a>",
)
elif self.model_uri is not None:
card_tab.add_markdown(
"MODEL_SOURCE_URI",
f"**Model Source URI:** `{self.model_uri}`",
)
return card
@classmethod
def from_recipe_config(cls, recipe_config, recipe_root):
step_config = {}
if recipe_config.get("steps", {}).get("register") is not None:
step_config.update(recipe_config.get("steps", {}).get("register"))
step_config["recipe"] = recipe_config.get("recipe")
if recipe_config.get("model_registry", {}).get("registry_uri") is not None:
step_config["registry_uri"] = recipe_config.get("model_registry", {}).get(
"registry_uri"
)
if recipe_config.get("model_registry", {}).get("model_name") is not None:
step_config["model_name"] = recipe_config.get("model_registry", {}).get("model_name")
step_config.update(
get_recipe_tracking_config(
recipe_root_path=recipe_root,
recipe_config=recipe_config,
).to_dict()
)
return cls(step_config, recipe_root)
@property
def name(self):
return "register"
@property
def environment(self):
return get_databricks_env_vars(tracking_uri=self.tracking_config.tracking_uri)
def get_artifacts(self):
return [
ModelVersionArtifact(
"registered_model_version",
self.recipe_root,
self.name,
self.tracking_config.tracking_uri,
)
]
def step_class(self):
return StepClass.TRAINING