File size: 5,413 Bytes
8200ee1 b666236 8200ee1 02f31d3 8200ee1 b301ad2 b666236 8200ee1 02f31d3 8200ee1 b666236 1d72461 b666236 8200ee1 b301ad2 8200ee1 02f31d3 8200ee1 b301ad2 8200ee1 02f31d3 8200ee1 02f31d3 8200ee1 b301ad2 02f31d3 8200ee1 b301ad2 02f31d3 8200ee1 b301ad2 8200ee1 b666236 8200ee1 b301ad2 02f31d3 8200ee1 b301ad2 02f31d3 8200ee1 02f31d3 8200ee1 02f31d3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 | """
API de prédiction du churn employé.
Contenu:
- Configuration de l'application FastAPI
- Fonction de prédiction
- Définition des endpoints
- Point d'entrée de l'application
"""
import logging
from fastapi import FastAPI, HTTPException
from fastapi.responses import Response
from src.api.schemas import DonneesEmploye, ResultatPrediction
from src.db.logger_db import (
log_api_operation,
log_prediction_input,
log_prediction_output,
)
from src.utils.model_loader import get_modele
from src.utils.transformer import transformer_donnees
# CONFIGURATION DU LOGGER
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# APPLICATION FASTAPI
tags_metadata = [
{
"name": "Service API",
"description": "Endpoint du service API pour l'inférence.",
},
{
"name": "Infrastructure",
"description": "Endpoints courants d'infrastructure pour l'observabilité (contrôle d'intégrité, métriques).",
},
]
app = FastAPI(
title="API Prédiction Churn Employé",
description="""
API qui prédit si un employé va quitter l'entreprise.
## Fonctionnalités
- Prédiction du churn (départ) d'un employé
- Validation automatique des données entrantes
- Documentation interactive Swagger UI (/docs)
## Données requises
L'API attend les données complètes d'un employé incluant :
- Informations personnelles (âge, genre, salaire, etc.)
- Expérience professionnelle
- Scores de satisfaction
- Informations sur le poste
""",
version="1.0.0",
docs_url="/docs",
redoc_url="/redoc",
openapi_tags=tags_metadata,
)
@app.get("/favicon.ico", include_in_schema=False)
async def favicon():
return Response(status_code=204)
# FONCTIONS DE PRÉDICTION
def predire_churn(donnees: dict) -> dict:
"""
Fonction principale de prédiction.
Étapes:
1. Transformer les données (encodage, etc.)
2. Appliquer le modèle
3. Retourner le résultat
Args:
donnees: Dict avec les données de l'employé
Returns:
Dict avec prédiction, probabilité et classe
Raises:
HTTPException: Si les features ne correspondent pas au modèle
"""
donnees_transformees = transformer_donnees(donnees)
modele = get_modele()
try:
prediction = modele.predict(donnees_transformees)[0]
probabilites = modele.predict_proba(donnees_transformees)[0]
except Exception as e:
raise HTTPException(
status_code=422,
detail=(
"Erreur de features: Les colonnes fournies ne correspondent pas "
f"aux features du modèle entraîné. Message original: {e}"
),
) from e
proba_depart = probabilites[1]
return {
"prediction": "Oui" if prediction == 1 else "Non",
"probabilite": round(float(proba_depart), 3),
"classe": int(prediction),
}
# ENDPOINTS (routes de l'API)
@app.get("/", tags=["Infrastructure"])
async def racine():
"""Endpoint racine - informations générales."""
return {
"message": "API Prédiction Churn Employé",
"version": "1.0.0",
"documentation": "/docs",
"redoc": "/redoc",
}
@app.get("/health", tags=["Infrastructure"])
async def health_check():
"""Vérifie que l'API fonctionne."""
return {"status": "ok", "message": "L'API est opérationnelle"}
@app.post("/predire", response_model=ResultatPrediction, tags=["Service API"])
def predire(employe: DonneesEmploye):
"""
Endpoint de prédiction du churn.
Args:
employe: Données de l'employé (validées par Pydantic)
Returns:
Prédiction avec probabilité de départ
Example:
```json
{
"prediction": "Oui",
"probabilite": 0.75,
"classe": 1
}
```
"""
donnees = employe.model_dump(exclude_none=True)
resultat = predire_churn(donnees)
# Logging dans la base de données (optionnel, ne bloque pas l'API)
input_id = log_prediction_input(donnees)
log_prediction_output(
input_id=input_id,
prediction=resultat["prediction"],
probabilite=resultat["probabilite"],
classe=resultat["classe"],
)
log_api_operation(
operation="PREDICT",
table_cible="prediction_inputs",
details=f"id_employee={donnees.get('id_employee')}, prediction={resultat['prediction']}",
statut="SUCCESS",
)
return ResultatPrediction(**resultat)
@app.get("/modele/info", tags=["Infrastructure"])
async def infos_modele():
"""Retourne des informations sur le modèle."""
return {
"type": "RandomForestClassifier",
"description": "Modèle de prédiction du churn employé",
"version": "1.0.0",
"nombre_features": 24,
"cible": "Prédiction du départ d'un employé (0=restera, 1=partira)",
}
@app.get("/modele/features", tags=["Infrastructure"])
async def liste_features():
"""Retourne la liste des features attendues par le modèle."""
from src.utils.transformer import get_liste_features
return {"features": get_liste_features(), "nombre": len(get_liste_features())}
# POINT D'ENTRÉE
if __name__ == "__main__":
import uvicorn
uvicorn.run("src.api.main:app", host="0.0.0.0", port=7860, reload=True)
|