Djohell commited on
Commit
a179cf2
·
1 Parent(s): 97f06af

maj app.py

Browse files
Files changed (1) hide show
  1. app.py +15 -16
app.py CHANGED
@@ -5,6 +5,7 @@ import xgboost as xgb
5
  import pandas as pd
6
  import mlflow.xgboost
7
  from fastapi import FastAPI, HTTPException, Body
 
8
  from dotenv import load_dotenv
9
  from processing import prepare_input, calculate_survival_risk, map_statut_expert, get_sigma
10
 
@@ -28,7 +29,6 @@ app = FastAPI(
28
  model = None
29
  SIGMA = None
30
 
31
- # Dans ton bloc startup, ajoute ces prints pour débugger :
32
  @app.on_event("startup")
33
  async def load_model():
34
  global model, SIGMA
@@ -38,9 +38,7 @@ async def load_model():
38
  # 1. On charge l'objet
39
  loaded_model = mlflow.xgboost.load_model(MODEL_URI)
40
 
41
- # 2. CORRECTION ICI :
42
- # Si c'est déjà un Booster, on l'utilise directement.
43
- # Si c'est un wrapper XGBModel, on appelle get_booster().
44
  if isinstance(loaded_model, xgb.Booster):
45
  model = loaded_model
46
  else:
@@ -55,15 +53,21 @@ async def load_model():
55
 
56
  # --- 3. ROUTES ---
57
 
58
- @app.get("/")
 
59
  def home():
 
 
 
 
 
60
  return {
61
  "status": "online",
62
- "model_source": "MLflow Remote",
63
  "run_id": RUN_ID
64
  }
65
 
66
- @app.post("/predict")
67
  async def predict(
68
  data: dict = Body(..., example={
69
  "age_estime": 4.5,
@@ -75,25 +79,19 @@ async def predict(
75
  })
76
  ):
77
  """
78
- Simule le risque de fermeture d'une entreprise.
79
-
80
- Exemple fourni :
81
- - Age : 4.5 ans
82
- - Effectif : Tranche 3
83
- - Localisation : Drôme (26)
84
- - Secteur : Construction (43)
85
  """
86
  if model is None:
87
  raise HTTPException(status_code=503, detail="Modèle non chargé")
88
 
89
  try:
90
- # 1. Préparation des données (Mapping APE 2 chiffres inclus)
91
  dmatrix = prepare_input(data)
92
 
93
  # 2. Inférence (Score MU)
94
  mu = float(model.predict(dmatrix)[0])
95
 
96
- # 3. Calcul des probabilités de fermeture aux horizons 1, 2 et 3 ans
97
  p1 = calculate_survival_risk(mu, 1, SIGMA)
98
  p2 = calculate_survival_risk(mu, 2, SIGMA)
99
  p3 = calculate_survival_risk(mu, 3, SIGMA)
@@ -126,4 +124,5 @@ async def predict(
126
 
127
  if __name__ == "__main__":
128
  import uvicorn
 
129
  uvicorn.run(app, host="0.0.0.0", port=7860)
 
5
  import pandas as pd
6
  import mlflow.xgboost
7
  from fastapi import FastAPI, HTTPException, Body
8
+ from fastapi.responses import RedirectResponse
9
  from dotenv import load_dotenv
10
  from processing import prepare_input, calculate_survival_risk, map_statut_expert, get_sigma
11
 
 
29
  model = None
30
  SIGMA = None
31
 
 
32
  @app.on_event("startup")
33
  async def load_model():
34
  global model, SIGMA
 
38
  # 1. On charge l'objet
39
  loaded_model = mlflow.xgboost.load_model(MODEL_URI)
40
 
41
+ # 2. Gestion du format Booster vs Wrapper
 
 
42
  if isinstance(loaded_model, xgb.Booster):
43
  model = loaded_model
44
  else:
 
53
 
54
  # --- 3. ROUTES ---
55
 
56
+ # Redirection automatique vers la doc Swagger à l'ouverture de l'URL
57
+ @app.get("/", include_in_schema=False)
58
  def home():
59
+ return RedirectResponse(url="/docs")
60
+
61
+ # Route de santé pour vérifier le statut sans redirection
62
+ @app.get("/health", tags=["Système"])
63
+ def health():
64
  return {
65
  "status": "online",
66
+ "model_loaded": model is not None,
67
  "run_id": RUN_ID
68
  }
69
 
70
+ @app.post("/predict", tags=["Prédiction"])
71
  async def predict(
72
  data: dict = Body(..., example={
73
  "age_estime": 4.5,
 
79
  })
80
  ):
81
  """
82
+ Simule le risque de fermeture d'une entreprise à 1, 2 et 3 ans.
 
 
 
 
 
 
83
  """
84
  if model is None:
85
  raise HTTPException(status_code=503, detail="Modèle non chargé")
86
 
87
  try:
88
+ # 1. Préparation des données
89
  dmatrix = prepare_input(data)
90
 
91
  # 2. Inférence (Score MU)
92
  mu = float(model.predict(dmatrix)[0])
93
 
94
+ # 3. Calcul des probabilités
95
  p1 = calculate_survival_risk(mu, 1, SIGMA)
96
  p2 = calculate_survival_risk(mu, 2, SIGMA)
97
  p3 = calculate_survival_risk(mu, 3, SIGMA)
 
124
 
125
  if __name__ == "__main__":
126
  import uvicorn
127
+ # Important : sur HF, le port par défaut attendu est 7860
128
  uvicorn.run(app, host="0.0.0.0", port=7860)