Files
ENI-ml-mlops/lab/serving/registry.py

44 lines
1.4 KiB
Python

"""Chargement du modele promu depuis le MLflow Model Registry (par alias)."""
import logging
from dataclasses import dataclass
import mlflow
import pandas as pd
from .. import constants
logger = logging.getLogger(__name__)
@dataclass
class LoadedModel:
"""Modele charge + metadonnees de version, garde en memoire par l'API."""
model: mlflow.pyfunc.PyFuncModel
name: str
version: str
def predict_one(self, features: dict[str, float]) -> float:
"""Prediction pour un jeu de features (1 ligne)."""
frame = pd.DataFrame([features], columns=constants.SERVING_FEATURES)
return float(self.model.predict(frame)[0])
def load_champion() -> LoadedModel:
"""Charger la version pointee par l'alias `MODEL_ALIAS` du modele enregistre.
On identifie le modele par `models:/<nom>@<alias>` : le meme code sert n'importe
quelle version promue, sans redeploiement, en changeant seulement l'alias cote MLflow.
"""
name = constants.REGISTERED_MODEL_NAME
alias = constants.MODEL_ALIAS
uri = f"models:/{name}@{alias}"
logger.info(f"Chargement du modele {uri}")
client = mlflow.MlflowClient()
version = client.get_model_version_by_alias(name=name, alias=alias)
model = mlflow.pyfunc.load_model(uri)
logger.info(f"Modele charge : {name} v{version.version}")
return LoadedModel(model=model, name=name, version=str(version.version))