DavidJyes commited on
Commit
fe22c25
·
verified ·
1 Parent(s): ca1611c

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +21 -28
app.py CHANGED
@@ -1,3 +1,4 @@
 
1
  from fastapi import FastAPI
2
  from fastapi.responses import RedirectResponse
3
  import joblib
@@ -6,8 +7,7 @@ import numpy as np
6
  from pydantic import BaseModel
7
  from sklearn.base import BaseEstimator, TransformerMixin
8
 
9
- # --- 1. DEFINITION DES CLASSES PERSONNALISÉES (CRUCIAL) ---
10
- # Cette classe doit être définie AVANT le chargement du modèle pour que joblib la reconnaisse
11
  class DFToHashedTokens(BaseEstimator, TransformerMixin):
12
  def __init__(self, columns=None):
13
  self.columns = columns
@@ -17,28 +17,34 @@ class DFToHashedTokens(BaseEstimator, TransformerMixin):
17
 
18
  def transform(self, X):
19
  X = X.copy()
20
- # Ici, remets la logique exacte que tu avais dans ton notebook
21
- # Par exemple, si tu transformais des colonnes en chaînes de caractères :
22
- for col in self.columns:
23
- X[col] = X[col].astype(str)
24
  return X
25
 
26
- # --- 2. CHARGEMENT DU MODÈLE ---
27
- # Maintenant que la classe est définie, joblib ne plantera plus
 
 
 
 
 
28
  try:
 
29
  model = joblib.load('fraud_model_hashing.pkl')
30
  print("✅ Modèle chargé avec succès")
31
  except Exception as e:
32
- print(f"❌ Erreur lors du chargement du modèle : {e}")
 
33
 
34
- # --- 3. CONFIGURATION API ---
35
  app = FastAPI(title="Fraud Detection API")
36
 
37
  @app.get("/", include_in_schema=False)
38
  def root():
39
  return RedirectResponse(url="/docs")
40
 
41
- # --- 4. SCHÉMA DES DONNÉES ---
42
  class Transaction(BaseModel):
43
  amt: float
44
  trans_date_trans_time: str
@@ -55,51 +61,38 @@ class Transaction(BaseModel):
55
  job: str
56
  cc_num: int
57
 
58
- # --- 5. LOGIQUE DE PRÉPARATION ---
59
  def prepare_input(data: dict):
60
  df = pd.DataFrame([data])
61
-
62
- # Dates
63
  dt = pd.to_datetime(df["trans_date_trans_time"])
64
  df["hour"] = dt.dt.hour
65
  df["day_of_week"] = dt.dt.dayofweek
66
  df["day"] = dt.dt.day
67
  df["month"] = dt.dt.month
68
 
69
- # Âge
70
  dob = pd.to_datetime(df["dob"])
71
  df["age"] = ((dt - dob).dt.days / 365.25).astype("float32")
72
 
73
- # Distance
74
  lat1, lon1 = np.radians(df["lat"]), np.radians(df["long"])
75
  lat2, lon2 = np.radians(df["merch_lat"]), np.radians(df["merch_long"])
76
  d = np.sin((lat2-lat1)/2)**2 + np.cos(lat1)*np.cos(lat2)*np.sin((lon2-lon1)/2)**2
77
  df["distance"] = (6371 * 2 * np.arcsin(np.sqrt(d))).astype("float32")
78
 
79
- # Agrégats par défaut (pour une transaction isolée via API)
80
  df["avg_amt"] = df["amt"]
81
  df["std_amt"] = 0.0
82
  df["nb_trans"] = 1.0
83
 
84
- # Ordre strict des colonnes
85
  cols = ['amt', 'hour', 'day_of_week', 'day', 'month', 'age', 'lat', 'long',
86
  'city_pop', 'distance', 'avg_amt', 'std_amt', 'nb_trans',
87
  'category', 'gender', 'state', 'merchant', 'job']
88
  return df[cols]
89
 
90
- # --- 6. ENDPOINT DE PRÉDICTION ---
91
  @app.post("/predict")
92
  def predict(data: Transaction):
 
 
93
  try:
94
  X = prepare_input(data.dict())
95
  prediction = model.predict(X)
96
- return {
97
- "is_fraud": int(prediction[0]),
98
- "status": "success"
99
- }
100
  except Exception as e:
101
- return {"status": "error", "message": str(e)}
102
-
103
- if __name__ == "__main__":
104
- import uvicorn
105
- uvicorn.run(app, host="0.0.0.0", port=7860) # Port par défaut Hugging Face
 
1
+ import sys
2
  from fastapi import FastAPI
3
  from fastapi.responses import RedirectResponse
4
  import joblib
 
7
  from pydantic import BaseModel
8
  from sklearn.base import BaseEstimator, TransformerMixin
9
 
10
+ # --- 1. DEFINITION DE LA CLASSE ---
 
11
  class DFToHashedTokens(BaseEstimator, TransformerMixin):
12
  def __init__(self, columns=None):
13
  self.columns = columns
 
17
 
18
  def transform(self, X):
19
  X = X.copy()
20
+ if self.columns:
21
+ for col in self.columns:
22
+ X[col] = X[col].astype(str)
 
23
  return X
24
 
25
+ # --- 2. L'ASTUCE ANTI-ERREUR (CRUCIAL) ---
26
+ # On injecte la classe dans le module '__main__' pour que joblib la trouve
27
+ import __main__
28
+ __main__.DFToHashedTokens = DFToHashedTokens
29
+
30
+ # --- 3. CHARGEMENT DU MODÈLE ---
31
+ # On le place dans une fonction ou on le charge après l'injection
32
  try:
33
+ # Assure-toi que le nom du fichier est exactement celui-ci sur Hugging Face
34
  model = joblib.load('fraud_model_hashing.pkl')
35
  print("✅ Modèle chargé avec succès")
36
  except Exception as e:
37
+ model = None
38
+ print(f"❌ Erreur lors du chargement : {e}")
39
 
40
+ # --- 4. CONFIGURATION API ---
41
  app = FastAPI(title="Fraud Detection API")
42
 
43
  @app.get("/", include_in_schema=False)
44
  def root():
45
  return RedirectResponse(url="/docs")
46
 
47
+ # --- 5. SCHÉMA ET PRÉPARATION ---
48
  class Transaction(BaseModel):
49
  amt: float
50
  trans_date_trans_time: str
 
61
  job: str
62
  cc_num: int
63
 
 
64
  def prepare_input(data: dict):
65
  df = pd.DataFrame([data])
 
 
66
  dt = pd.to_datetime(df["trans_date_trans_time"])
67
  df["hour"] = dt.dt.hour
68
  df["day_of_week"] = dt.dt.dayofweek
69
  df["day"] = dt.dt.day
70
  df["month"] = dt.dt.month
71
 
 
72
  dob = pd.to_datetime(df["dob"])
73
  df["age"] = ((dt - dob).dt.days / 365.25).astype("float32")
74
 
 
75
  lat1, lon1 = np.radians(df["lat"]), np.radians(df["long"])
76
  lat2, lon2 = np.radians(df["merch_lat"]), np.radians(df["merch_long"])
77
  d = np.sin((lat2-lat1)/2)**2 + np.cos(lat1)*np.cos(lat2)*np.sin((lon2-lon1)/2)**2
78
  df["distance"] = (6371 * 2 * np.arcsin(np.sqrt(d))).astype("float32")
79
 
 
80
  df["avg_amt"] = df["amt"]
81
  df["std_amt"] = 0.0
82
  df["nb_trans"] = 1.0
83
 
 
84
  cols = ['amt', 'hour', 'day_of_week', 'day', 'month', 'age', 'lat', 'long',
85
  'city_pop', 'distance', 'avg_amt', 'std_amt', 'nb_trans',
86
  'category', 'gender', 'state', 'merchant', 'job']
87
  return df[cols]
88
 
 
89
  @app.post("/predict")
90
  def predict(data: Transaction):
91
+ if model is None:
92
+ return {"status": "error", "message": "Model not loaded"}
93
  try:
94
  X = prepare_input(data.dict())
95
  prediction = model.predict(X)
96
+ return {"is_fraud": int(prediction[0])}
 
 
 
97
  except Exception as e:
98
+ return {"status": "error", "message": str(e)}