Alexis-Ravet's picture
Upload folder using huggingface_hub
1d72461 verified
Raw
History Blame Contribute Delete
5.41 kB
"""
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)