diff --git a/lab/serving/__init__.py b/lab/serving/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/lab/serving/api.py b/lab/serving/api.py new file mode 100644 index 0000000..1da2503 --- /dev/null +++ b/lab/serving/api.py @@ -0,0 +1,99 @@ +"""API REST de prediction de consommation electrique (FastAPI). + +Le service orchestre la chaine de prediction : + 1. recevoir la requete (client_id, date) ; + 2. recuperer les features du client (feature store simule) ; + 3. charger le modele promu (Model Registry, au demarrage) ; + 4. calculer la prediction et repondre en JSON. + +Lancement : uvicorn lab.serving.api:app --host 0.0.0.0 --port 8000 +Swagger UI : /docs +""" + +import logging +from contextlib import asynccontextmanager + +from fastapi import FastAPI, HTTPException + +from . import features, registry +from .schemas import ( + BatchPredictionRequest, + BatchPredictionResponse, + PredictionRequest, + PredictionResponse, +) + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +# Modele charge une fois au demarrage (couteux) et reutilise a chaque requete. +_state: dict[str, registry.LoadedModel] = {} + + +@asynccontextmanager +async def lifespan(app: FastAPI): + _state["model"] = registry.load_champion() + yield + _state.clear() + + +app = FastAPI( + title="Electricity Consumption Prediction API", + description="Expose le modele promu (MLflow Model Registry) via une API REST.", + version="1.0.0", + lifespan=lifespan, +) + + +def _get_model() -> registry.LoadedModel: + model = _state.get("model") + if model is None: # modele indisponible au demarrage + raise HTTPException(status_code=503, detail="Modele non charge.") + return model + + +def _predict(client_id: str, model: registry.LoadedModel) -> PredictionResponse: + feats = features.get_features(client_id) # peut lever UnknownClientError + prediction = model.predict_one(feats) + return PredictionResponse( + client_id=client_id, + prediction_kwh=prediction, + model_name=model.name, + model_version=model.version, + ) + + +@app.get("/health", summary="Verification de l'etat du service") +def health() -> dict[str, str]: + """Endpoint de sante : renvoie 200 si le service repond.""" + return {"status": "ok"} + + +@app.post("/predict", response_model=PredictionResponse, summary="Prediction unitaire") +def predict(request: PredictionRequest) -> PredictionResponse: + model = _get_model() + try: + return _predict(request.client_id, model) + except features.UnknownClientError: + # Partie 3 : client inconnu -> erreur cliente, pas un plantage du service. + raise HTTPException( + status_code=404, + detail=f"Client inconnu : aucune feature pour '{request.client_id}'.", + ) + + +@app.post( + "/predict/batch", + response_model=BatchPredictionResponse, + summary="Prediction en batch (plusieurs clients)", +) +def predict_batch(request: BatchPredictionRequest) -> BatchPredictionResponse: + model = _get_model() + predictions: list[PredictionResponse] = [] + unknown: list[str] = [] + for client_id in request.client_ids: + try: + predictions.append(_predict(client_id, model)) + except features.UnknownClientError: + unknown.append(client_id) + return BatchPredictionResponse(predictions=predictions, unknown_client_ids=unknown) diff --git a/lab/serving/features.py b/lab/serving/features.py new file mode 100644 index 0000000..1c85672 --- /dev/null +++ b/lab/serving/features.py @@ -0,0 +1,66 @@ +"""Recuperation des features de prediction. + +Dans un systeme reel, ces valeurs proviendraient d'un feature store / d'une base +alimentee par le pipeline de calcul de features (lags, moyennes glissantes) sur +l'historique de consommation. On separe volontairement cette etape du calcul de la +prediction : la source des features peut changer sans toucher au modele. + +Ici on SIMULE cette recuperation par un simple dictionnaire Python, comme demande +par l'enonce. Les valeurs sont des observations reelles (derniere ligne connue de +quelques clients dans data/test.parquet). +""" + +from .. import constants + + +class UnknownClientError(KeyError): + """Aucune feature disponible pour ce client (identifiant inconnu).""" + + +# feature store simule : client_id -> {feature: valeur} +FEATURE_STORE: dict[str, dict[str, float]] = { + "MT_124": { + "lag_1d": 107.656, + "lag_7d": 25.120, + "lag_30d": 70.574, + "lag_365d": 25.120, + "rolling_mean_7d": 65.870, + "rolling_mean_30d": 71.310, + }, + "MT_156": { + "lag_1d": 13.149, + "lag_7d": 13.929, + "lag_30d": 21.577, + "lag_365d": 8.935, + "rolling_mean_7d": 16.648, + "rolling_mean_30d": 19.720, + }, + "MT_158": { + "lag_1d": 30.739, + "lag_7d": 16.608, + "lag_30d": 34.094, + "lag_365d": 6.574, + "rolling_mean_7d": 19.067, + "rolling_mean_30d": 21.486, + }, + "MT_159": { + "lag_1d": 23.305, + "lag_7d": 24.741, + "lag_30d": 21.386, + "lag_365d": 5.333, + "rolling_mean_7d": 11.707, + "rolling_mean_30d": 13.619, + }, +} + + +def get_features(client_id: str) -> dict[str, float]: + """Renvoyer les features du client, ordonnees comme a l'entrainement du modele. + + Leve UnknownClientError si le client est inconnu. + """ + if client_id not in FEATURE_STORE: + raise UnknownClientError(client_id) + raw = FEATURE_STORE[client_id] + # On respecte l'ordre des colonnes attendu par le modele (SERVING_FEATURES). + return {feature: raw[feature] for feature in constants.SERVING_FEATURES} diff --git a/lab/serving/registry.py b/lab/serving/registry.py new file mode 100644 index 0000000..71e91b1 --- /dev/null +++ b/lab/serving/registry.py @@ -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:/@` : 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)) diff --git a/lab/serving/schemas.py b/lab/serving/schemas.py new file mode 100644 index 0000000..e53bc5d --- /dev/null +++ b/lab/serving/schemas.py @@ -0,0 +1,44 @@ +"""Schemas Pydantic du service de prediction (contrat d'entree/sortie de l'API).""" + +import datetime + +from pydantic import BaseModel, Field + + +class PredictionRequest(BaseModel): + """Ce que le client fournit : QUI et QUAND, pas les features (calculees cote serveur).""" + + client_id: str = Field( + ..., + description="Identifiant du client (ex. 'MT_124').", + examples=["MT_124"], + ) + date: datetime.date | None = Field( + default=None, + description="Date de la prediction (parametre de requete). Optionnelle.", + examples=["2015-01-01"], + ) + + +class BatchPredictionRequest(BaseModel): + """Prediction pour plusieurs clients en une seule requete (inference batch).""" + + client_ids: list[str] = Field( + ..., + description="Liste d'identifiants clients.", + examples=[["MT_124", "MT_156", "MT_158"]], + ) + date: datetime.date | None = None + + +class PredictionResponse(BaseModel): + client_id: str + prediction_kwh: float = Field(description="Consommation predite (kWh).") + model_name: str + model_version: str + + +class BatchPredictionResponse(BaseModel): + predictions: list[PredictionResponse] + # Clients ignores (features introuvables) : on ne fait pas echouer tout le lot. + unknown_client_ids: list[str] = Field(default_factory=list) diff --git a/serve.sh b/serve.sh new file mode 100755 index 0000000..c5abaa0 --- /dev/null +++ b/serve.sh @@ -0,0 +1,10 @@ +#!/usr/bin/env bash +# Lance l'API de prediction (TP03). A executer depuis la racine du depot (/home/user/tp). +# Charge .env (MLflow + creds S3 pour recuperer le modele promu depuis le Registry). +set -euo pipefail +cd "$(dirname "$0")" +set -a +# shellcheck disable=SC1091 +source .env +set +a +exec /opt/venvs/mlops/bin/uvicorn lab.serving.api:app --host 0.0.0.0 --port 8000