Download src/api/main.py from Alexis-Ravet/Employee-Churn-Prediction: direct link, hf CLI and curl.
- Browser
- Download file 5.41 kB
-
https://huggingface.co/spaces/Alexis-Ravet/Employee-Churn-Prediction/resolve/main/src/api/main.py
- Command line
-
hf download hf://spaces/Alexis-Ravet/Employee-Churn-Prediction/src/api/main.py
-
curl -L -o main.py https://huggingface.co/spaces/Alexis-Ravet/Employee-Churn-Prediction/resolve/main/src/api/main.py
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, | |
| ) | |
| 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) | |
| async def racine(): | |
| """Endpoint racine - informations générales.""" | |
| return { | |
| "message": "API Prédiction Churn Employé", | |
| "version": "1.0.0", | |
| "documentation": "/docs", | |
| "redoc": "/redoc", | |
| } | |
| async def health_check(): | |
| """Vérifie que l'API fonctionne.""" | |
| return {"status": "ok", "message": "L'API est opérationnelle"} | |
| 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) | |
| 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)", | |
| } | |
| 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) | |