gdleds commited on
Commit
3083676
ยท
1 Parent(s): c0fbae9
Files changed (1) hide show
  1. app.py +20 -23
app.py CHANGED
@@ -45,6 +45,15 @@ from sksurv.metrics import concordance_index_censored
45
  from sksurv.linear_model import CoxnetSurvivalAnalysis, CoxPHSurvivalAnalysis
46
  from sksurv.preprocessing import OneHotEncoder
47
  from sksurv.util import Surv
 
 
 
 
 
 
 
 
 
48
 
49
 
50
  from sksurv.ensemble import GradientBoostingSurvivalAnalysis
@@ -99,26 +108,7 @@ def load_df_merge():
99
  url = 'https://fireprojectbislead.s3.us-east-1.amazonaws.com/dataset/historique_incendies_avec_coordonnees.csv'
100
  return pd.read_csv(url, sep=';', encoding='utf-8')
101
  #------------------------------------------------------- ----------------Notre produit#_________________________________________________
102
- import streamlit as st
103
- import pandas as pd
104
- import numpy as np
105
- import plotly.express as px
106
- import warnings
107
- from sklearn.exceptions import UndefinedMetricWarning
108
- from sklearn import set_config
109
- from sklearn.model_selection import train_test_split
110
- from sklearn.pipeline import Pipeline
111
- from sklearn.preprocessing import StandardScaler
112
- from sklearn.impute import SimpleImputer
113
- from xgboost import XGBRegressor, DMatrix, train as xgb_train
114
- from lifelines import CoxPHFitter
115
- from sksurv.util import Surv
116
- from sksurv.metrics import concordance_index_censored
117
- from mlflow import sklearn as mlflow_sklearn
118
- import mlflow.sklearn
119
 
120
- warnings.filterwarnings("ignore", category=UndefinedMetricWarning)
121
- set_config(display="text")
122
 
123
  # โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
124
  # 1) FONCTION DE CHARGEMENT DU CSV BRUT
@@ -133,12 +123,17 @@ def load_raw_data() -> pd.DataFrame:
133
  # 2) FONCTION Dโ€™ENTRAรŽNEMENT + PRร‰DICTIONS
134
  # โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
135
 
136
-
137
- mlflow.set_tracking_uri(os.getenv("BACKEND_STORE_URI")) # URI NeonDB
138
  os.environ["MLFLOW_DEFAULT_ARTIFACT_ROOT"] = os.getenv("MLFLOW_DEFAULT_ARTIFACT_ROOT") # S3
139
  os.environ["AWS_ACCESS_KEY_ID"] = os.getenv("AWS_ACCESS_KEY_ID")
140
  os.environ["AWS_SECRET_ACCESS_KEY"] = os.getenv("AWS_SECRET_ACCESS_KEY")
141
 
 
 
 
 
 
 
 
142
 
143
  # โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
144
  # 2) FONCTION DE PREDICTIONS (sans entraรฎnement)
@@ -168,8 +163,10 @@ def train_model_and_predict(df_raw: pd.DataFrame) -> pd.DataFrame:
168
  features = [f for f in features if f in df.columns]
169
 
170
  # c) Chargement du modรจle depuis MLflow
171
- model_uri = "models:/fire_survival/Production" # ou /1 pour la v1
172
- pipe = mlflow.sklearn.load_model(model_uri)
 
 
173
 
174
  # d) Estimation baseline hazard avec Cox factice (comme avant)
175
  y_struct = Surv.from_dataframe("event", "duration", df)
 
45
  from sksurv.linear_model import CoxnetSurvivalAnalysis, CoxPHSurvivalAnalysis
46
  from sksurv.preprocessing import OneHotEncoder
47
  from sksurv.util import Surv
48
+ from xgboost import XGBRegressor, DMatrix, train as xgb_train
49
+ from lifelines import CoxPHFitter
50
+ from mlflow import sklearn as mlflow_sklearn
51
+ import mlflow.sklearn
52
+ import boto3
53
+ import io
54
+
55
+ warnings.filterwarnings("ignore", category=UndefinedMetricWarning)
56
+ set_config(display="text")
57
 
58
 
59
  from sksurv.ensemble import GradientBoostingSurvivalAnalysis
 
108
  url = 'https://fireprojectbislead.s3.us-east-1.amazonaws.com/dataset/historique_incendies_avec_coordonnees.csv'
109
  return pd.read_csv(url, sep=';', encoding='utf-8')
110
  #------------------------------------------------------- ----------------Notre produit#_________________________________________________
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
111
 
 
 
112
 
113
  # โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
114
  # 1) FONCTION DE CHARGEMENT DU CSV BRUT
 
123
  # 2) FONCTION Dโ€™ENTRAรŽNEMENT + PRร‰DICTIONS
124
  # โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
125
 
 
 
126
  os.environ["MLFLOW_DEFAULT_ARTIFACT_ROOT"] = os.getenv("MLFLOW_DEFAULT_ARTIFACT_ROOT") # S3
127
  os.environ["AWS_ACCESS_KEY_ID"] = os.getenv("AWS_ACCESS_KEY_ID")
128
  os.environ["AWS_SECRET_ACCESS_KEY"] = os.getenv("AWS_SECRET_ACCESS_KEY")
129
 
130
+ def load_model_from_s3(bucket: str, key: str):
131
+ s3 = boto3.client("s3")
132
+ buffer = io.BytesIO()
133
+ s3.download_fileobj(bucket, key, buffer)
134
+ buffer.seek(0)
135
+ model = joblib.load(buffer)
136
+ return model
137
 
138
  # โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
139
  # 2) FONCTION DE PREDICTIONS (sans entraรฎnement)
 
163
  features = [f for f in features if f in df.columns]
164
 
165
  # c) Chargement du modรจle depuis MLflow
166
+ pipe = load_model_from_s3(
167
+ bucket=os.getenv("S3_BUCKET"),
168
+ key="mlflow/models/xgboost_survivalCOX_model_2f26eb52269844a2b01f13974185f102.joblib" # ton chemin exacts
169
+ )
170
 
171
  # d) Estimation baseline hazard avec Cox factice (comme avant)
172
  y_struct = Surv.from_dataframe("event", "duration", df)