Djohell commited on
Commit
e36a5c5
·
verified ·
1 Parent(s): 5899ebd

Upload app.py

Browse files
Files changed (1) hide show
  1. app.py +45 -30
app.py CHANGED
@@ -8,12 +8,11 @@ from fastapi import FastAPI, HTTPException, Body
8
  from fastapi.responses import RedirectResponse
9
  from dotenv import load_dotenv
10
 
11
- # On importe tout, y compris FEATURES pour le debug
12
  from processing import prepare_input, calculate_survival_risk, map_statut_expert, get_sigma, FEATURES
13
 
14
  # --- 1. CONFIGURATION MLFLOW ---
15
  load_dotenv()
16
-
17
  mlflow.set_tracking_uri(os.getenv("MLFLOW_TRACKING_URI"))
18
 
19
  RUN_ID = "674d07aab0b0493a838310da47c71a95"
@@ -21,11 +20,12 @@ MODEL_URI = f"runs:/{RUN_ID}/model"
21
 
22
  # --- 2. INITIALISATION DE L'API ---
23
  app = FastAPI(
24
- title="Business Risk API - Test MLflow Remote",
25
- description="API de simulation utilisant un modèle stocké sur un serveur MLflow distant.",
26
- version="3.5.0"
27
  )
28
 
 
29
  model = None
30
  SIGMA = None
31
 
@@ -34,8 +34,6 @@ async def load_model():
34
  global model, SIGMA
35
  try:
36
  print(f"🚀 Connexion à MLflow : {os.getenv('MLFLOW_TRACKING_URI')}")
37
-
38
- # Chargement du modèle
39
  loaded_model = mlflow.xgboost.load_model(MODEL_URI)
40
 
41
  if isinstance(loaded_model, xgb.Booster):
@@ -46,7 +44,7 @@ async def load_model():
46
  SIGMA = get_sigma(model)
47
  print(f"✅ Modèle chargé avec succès (Sigma: {round(SIGMA, 4)})")
48
  except Exception as e:
49
- print(f"❌ Erreur lors du chargement : {e}")
50
 
51
  # --- 3. ROUTES ---
52
 
@@ -59,26 +57,38 @@ def health():
59
  return {
60
  "status": "online",
61
  "model_loaded": model is not None,
62
- "run_id": RUN_ID
 
63
  }
64
 
65
- @app.post("/predict")
66
- async def predict(data: dict):
67
- try:
68
- if model is None:
69
- raise HTTPException(status_code=503, detail="Modèle non chargé")
 
 
 
 
 
 
 
 
 
 
 
70
 
71
- # 1. Préparation
 
72
  dmatrix = prepare_input(data)
73
 
74
- # 2. Prédiction
75
  mu = float(model.predict(dmatrix)[0])
76
- s = get_sigma(model)
77
 
78
- # 3. Risques
79
- p1 = calculate_survival_risk(mu, 1, s)
80
- p2 = calculate_survival_risk(mu, 2, s)
81
- p3 = calculate_survival_risk(mu, 3, s)
82
 
83
  return {
84
  "diagnostic": {
@@ -90,22 +100,27 @@ async def predict(data: dict):
90
  "2_ans": f"{p2}%",
91
  "3_ans": f"{p3}%"
92
  },
 
 
 
 
 
93
  "debug_internal": {
94
  "features_count": len(FEATURES),
95
- "first_feature": FEATURES[0] if FEATURES else "None",
96
- "input_received": {
97
- "age": data.get("age_estime"),
98
- "dep": data.get("code_departement")
99
- }
100
  },
101
  "metadonnees": {
102
- "run_id": os.urandom(8).hex(),
103
- "sigma_utilise": s
 
104
  }
105
  }
 
106
  except Exception as e:
107
- # Correction de la parenthèse ici
108
- return {"error": str(e)}
 
 
109
 
110
  if __name__ == "__main__":
111
  import uvicorn
 
8
  from fastapi.responses import RedirectResponse
9
  from dotenv import load_dotenv
10
 
11
+ # On importe les fonctions et la constante FEATURES depuis processing
12
  from processing import prepare_input, calculate_survival_risk, map_statut_expert, get_sigma, FEATURES
13
 
14
  # --- 1. CONFIGURATION MLFLOW ---
15
  load_dotenv()
 
16
  mlflow.set_tracking_uri(os.getenv("MLFLOW_TRACKING_URI"))
17
 
18
  RUN_ID = "674d07aab0b0493a838310da47c71a95"
 
20
 
21
  # --- 2. INITIALISATION DE L'API ---
22
  app = FastAPI(
23
+ title="Business Risk API",
24
+ description="API de prédiction du risque de fermeture des entreprises via modèle AFT.",
25
+ version="3.6.0"
26
  )
27
 
28
+ # Variables globales
29
  model = None
30
  SIGMA = None
31
 
 
34
  global model, SIGMA
35
  try:
36
  print(f"🚀 Connexion à MLflow : {os.getenv('MLFLOW_TRACKING_URI')}")
 
 
37
  loaded_model = mlflow.xgboost.load_model(MODEL_URI)
38
 
39
  if isinstance(loaded_model, xgb.Booster):
 
44
  SIGMA = get_sigma(model)
45
  print(f"✅ Modèle chargé avec succès (Sigma: {round(SIGMA, 4)})")
46
  except Exception as e:
47
+ print(f"❌ Erreur lors du chargement du modèle : {e}")
48
 
49
  # --- 3. ROUTES ---
50
 
 
57
  return {
58
  "status": "online",
59
  "model_loaded": model is not None,
60
+ "run_id": RUN_ID,
61
+ "features_synced": len(FEATURES) > 0
62
  }
63
 
64
+ @app.post("/predict", tags=["Prédiction"])
65
+ async def predict(
66
+ data: dict = Body(..., example={
67
+ "age_estime": 0.5,
68
+ "Tranche_effectif_num": 0,
69
+ "code_departement": "75",
70
+ "code_ape": "56",
71
+ "categorie_juridique": "5499",
72
+ "is_ess": 0
73
+ })
74
+ ):
75
+ """
76
+ Simule le risque de fermeture d'une entreprise à 1, 2 et 3 ans.
77
+ """
78
+ if model is None:
79
+ raise HTTPException(status_code=503, detail="Modèle non disponible")
80
 
81
+ try:
82
+ # 1. Préparation des données (Utilise le mapping S3)
83
  dmatrix = prepare_input(data)
84
 
85
+ # 2. Inférence (Score MU)
86
  mu = float(model.predict(dmatrix)[0])
 
87
 
88
+ # 3. Calcul des probabilités avec le Sigma extrait du modèle
89
+ p1 = calculate_survival_risk(mu, 1, SIGMA)
90
+ p2 = calculate_survival_risk(mu, 2, SIGMA)
91
+ p3 = calculate_survival_risk(mu, 3, SIGMA)
92
 
93
  return {
94
  "diagnostic": {
 
100
  "2_ans": f"{p2}%",
101
  "3_ans": f"{p3}%"
102
  },
103
+ "entrees_recues": {
104
+ "age_saisi": data.get("age_estime"),
105
+ "division_ape": data.get("code_ape"),
106
+ "departement": data.get("code_departement")
107
+ },
108
  "debug_internal": {
109
  "features_count": len(FEATURES),
110
+ "first_feature": FEATURES[0] if FEATURES else "None"
 
 
 
 
111
  },
112
  "metadonnees": {
113
+ "run_id": RUN_ID,
114
+ "sigma_utilise": round(SIGMA, 6),
115
+ "api_version": "3.6.0"
116
  }
117
  }
118
+
119
  except Exception as e:
120
+ raise HTTPException(
121
+ status_code=500,
122
+ detail=f"Erreur lors de la prédiction : {str(e)}"
123
+ )
124
 
125
  if __name__ == "__main__":
126
  import uvicorn