44 lines
1.4 KiB
Python
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))
|