TP3 Parties 2-4 : API REST FastAPI (health, predict, batch, erreurs 404)

This commit is contained in:
Johan LEROY
2026-07-22 11:22:50 +02:00
parent 9a52395c91
commit b0bd6cdb07
6 changed files with 262 additions and 0 deletions

0
lab/serving/__init__.py Normal file
View File

99
lab/serving/api.py Normal file
View 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
View 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
View 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
View 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
View 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