TP3 Parties 2-4 : API REST FastAPI (health, predict, batch, erreurs 404)
This commit is contained in:
43
lab/serving/registry.py
Normal file
43
lab/serving/registry.py
Normal file
@@ -0,0 +1,43 @@
|
||||
"""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))
|
||||
Reference in New Issue
Block a user