TP3 Parties 2-4 : API REST FastAPI (health, predict, batch, erreurs 404)
This commit is contained in:
0
lab/serving/__init__.py
Normal file
0
lab/serving/__init__.py
Normal file
99
lab/serving/api.py
Normal file
99
lab/serving/api.py
Normal file
@@ -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)
|
||||||
66
lab/serving/features.py
Normal file
66
lab/serving/features.py
Normal file
@@ -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}
|
||||||
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))
|
||||||
44
lab/serving/schemas.py
Normal file
44
lab/serving/schemas.py
Normal file
@@ -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)
|
||||||
10
serve.sh
Executable file
10
serve.sh
Executable file
@@ -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
|
||||||
Reference in New Issue
Block a user