DavidJyes commited on
Commit
0b59a26
·
verified ·
1 Parent(s): ead8493

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +352 -352
app.py CHANGED
@@ -1,352 +1,352 @@
1
- import streamlit as st
2
- import pandas as pd
3
- import pickle
4
- import matplotlib.pyplot as plt
5
- from dotenv import load_dotenv
6
- import os
7
-
8
- import numpy as np
9
- import shap
10
- from sklearn.linear_model import LinearRegression
11
- from langchain_mistralai import ChatMistralAI
12
- from langchain_core.output_parsers import StrOutputParser
13
-
14
- # partie initialisation
15
- model_paths = {
16
- "Deblais et Gravats": "model_paths/model_Deblais_gravats.pkl",
17
- "Dechets verts": "model_paths/model_Dechets_verts.pkl",
18
- "Encombrants": "model_paths/model_Encombrants.pkl",
19
- "Materiaux recyclables": "model_paths/model_Materiaux_recyclables.pkl"
20
- }
21
-
22
- col_mapping = {
23
- "Deblais et Gravats": "Deblais_gravats",
24
- "Dechets verts": "Dechets_verts",
25
- "Encombrants": "Encombrants",
26
- "Materiaux recyclables": "Materiaux_recyclables"
27
- }
28
-
29
- default_dept = "Aisne"
30
-
31
- valeurs_observees = []
32
- valeurs_predites = []
33
- labels = []
34
-
35
- categories = {
36
- "📊 Population": [
37
- "pop_globale",
38
- "tranche_age_0-24", "tranche_age_25-59", "tranche_age_60+",
39
- "csp1_agriculteurs", "csp2_artisans_commerçant_chef_entreprises",
40
- "csp3_cadres_professions_intellectuelles", "csp4_professions_intermediaires",
41
- "csp5_employes", "csp6_ouvriers", "csp7_retraites", "csp8_sans_activite",
42
- "densite"
43
- ],
44
- "🏭 Activite economique": [
45
- "nbre_entreprises", "nbre_entreprises_agricole", "nb_salaries_secteur_agricole",
46
- "nbre_entreprises_industrie", "nb_salaries_secteur_industrie",
47
- "nb_salaries_secteur_service", "nbre_entreprises_service"
48
- ]
49
- }
50
-
51
- # Pour la gestion automatique du run eval uniquement au demarrage et en cas de changement de departement
52
- # sinon necessaire de cliquer sur Lancer l'evaluation
53
- if "previous_dept" not in st.session_state:
54
- st.session_state["previous_dept"] = None
55
-
56
- if "auto_run_done" not in st.session_state:
57
- st.session_state["auto_run_done"] = False
58
-
59
- def run_eval(selected_dept, form_input):
60
- # Transformation en DataFrame
61
- input_df = pd.DataFrame([form_input])
62
- input_df_complete = row_default.to_frame().T.copy()
63
- for col in input_df.columns:
64
- if col in input_df_complete.columns:
65
- input_df_complete.at[input_df_complete.index[0], col] = input_df.at[0, col]
66
-
67
- # Verification des incoherences
68
- Liste_age = ["pop_globale", "tranche_age_0-24", "tranche_age_25-59", "tranche_age_60+"]
69
- Somme = 0
70
- if all(v in form_input for v in Liste_age):
71
- for elt in Liste_age:
72
- if elt != "pop_globale":
73
- Somme += form_input[elt]
74
- if abs(form_input["pop_globale"] - Somme) > 1:
75
- st.error(f"❌ Incoherence : Population globale = {form_input['pop_globale']} ne correspond pas à la somme des tranches d'âge : {Somme}")
76
-
77
- Liste_CSP = ["pop_globale","csp1_agriculteurs", "csp2_artisans_commercant_chef_entreprises",
78
- "csp3_cadres_professions_intellectuelles", "csp4_professions_intermediaires",
79
- "csp5_employes", "csp6_ouvriers", "csp7_retraites", "csp8_sans_activite"]
80
- Somme = 0
81
- if all(v in form_input for v in Liste_CSP):
82
- for elt in Liste_CSP:
83
- if elt != "pop_globale":
84
- Somme += form_input[elt]
85
- if abs(form_input["pop_globale"] - Somme) >1:
86
- st.error(f"❌ Incoherence : Population globale = {form_input['pop_globale']} ne correspond pas à la somme des CSP : {Somme}")
87
-
88
- Liste_Entreprise = ["nbre_entreprises", "nbre_entreprises_agricole",
89
- "nbre_entreprises_industrie", "nbre_entreprises_service"]
90
- Somme = 0
91
- if all(v in form_input for v in Liste_Entreprise):
92
- for elt in Liste_Entreprise:
93
- if elt != "nbre_entreprises":
94
- Somme += form_input[elt]
95
- if abs(form_input["nbre_entreprises"] - Somme) > 1:
96
- st.error(f"❌ Incoherence : Le nombre d'entreprises = {form_input['nbre_entreprises']} ne correspond pas à la somme des types d'entreprise : {Somme}")
97
-
98
-
99
- # on reinitialise pour que ça ne se lance pas automatiquement (lourd)
100
- evaluation = False
101
- for typologie, path in model_paths.items():
102
- try:
103
- if os.path.exists(path):
104
- with open(path, "rb") as f:
105
- model = pickle.load(f)
106
- else:
107
- st.error(f"❌ Modèle manquant : {path}")
108
- expected_cols = model.model.exog_names
109
- if "const" in expected_cols and "const" not in input_df_complete.columns:
110
- input_df_complete["const"] = 1.0
111
-
112
- prediction = max(0, model.predict(input_df_complete[expected_cols]).iloc[0])
113
- valeurs_predites.append(prediction)
114
- labels.append(typologie)
115
-
116
- filtered = observed_df[
117
- (observed_df["Departement"] == selected_dept) & (observed_df["annee"] == 2021)
118
- ]
119
-
120
- excel_col = col_mapping.get(typologie)
121
- if not filtered.empty and excel_col in filtered.columns:
122
- valeurs_observees.append(filtered[excel_col].values[0])
123
- else:
124
- valeurs_observees.append(0.0)
125
- except Exception as e:
126
- st.error(f"Erreur avec le modèle {typologie}")
127
- st.exception(e)
128
-
129
- # === Chargement des donnees
130
- df = pd.read_csv(
131
- r"/mnt/c/Users/david/Desktop/CDSD 2025/BLOC_6_PROJET_FINAL/geodechet/GeoDechet_streamlit/Data/df_dummies.csv").drop(columns=["Unnamed: 0"], errors="ignore")
132
- observed_df = pd.read_csv(r"/mnt/c/Users/david/Desktop/CDSD 2025/BLOC_6_PROJET_FINAL/geodechet/GeoDechet_streamlit/Data/data_wip.csv")
133
-
134
- # === Liste des departements
135
- departements = [col.replace("Departement_", "") for col in df.columns if col.startswith("Departement_")]
136
-
137
- # === Mise en page
138
- st.set_page_config(layout="wide")
139
- st.markdown("<h1 style='text-align: left;'>♻️ Simulateur de production de dechets par departement</h1>", unsafe_allow_html=True)
140
-
141
- st.markdown("<h3 style='text-align: left;'>📍 Choix du departement</h3>", unsafe_allow_html=True)
142
-
143
- st.markdown("<h3 style='text-align: left;'>📈 Comparaison entre valeurs observees et predites</h3>", unsafe_allow_html=True)
144
-
145
- st.markdown("<div style='text-align:left: 60px;'></div>", unsafe_allow_html=True)
146
-
147
- selected_dept = st.selectbox("Selectionner un departement", sorted(departements), index=sorted(departements).index(default_dept))
148
-
149
- row_default = df[df[f"Departement_{selected_dept}"] == 1].iloc[0]
150
- default_dict = row_default.to_dict()
151
-
152
- st.subheader("⚙️ Paramètres modifiables")
153
- form_input = {}
154
- for category_name, variables in categories.items():
155
- with st.expander(category_name, expanded=False):
156
- for i, var in enumerate(variables):
157
- if var in default_dict:
158
- if category_name == "📊 Population":
159
- default_value = int(round(float(default_dict[var])/100)*100)
160
- else:
161
- default_value = int(round(float(default_dict[var])/10)*10)
162
- val = st.number_input(
163
- f"✏️ {var}",
164
- min_value=0,
165
- value=default_value,
166
- step=1,
167
- format="%d",
168
- key=f"number_input_{selected_dept}_{var}"
169
- )
170
- form_input[var] = val
171
- st.markdown("<div style='margin-bottom: 10px;'></div>", unsafe_allow_html=True)
172
-
173
- if st.session_state["previous_dept"] != selected_dept or not st.session_state["auto_run_done"]:
174
- st.session_state["previous_dept"] = selected_dept
175
- st.session_state["auto_run_done"] = True
176
- run_eval(selected_dept, form_input)
177
-
178
-
179
- # with chart_col:
180
- st.markdown("<div style='text-align:center: 30px;'></div>", unsafe_allow_html=True)
181
-
182
- # btn_col = st.columns([3, 2, 3])[1]
183
- # with btn_col:
184
- evaluation = st.button("🔍 Lancer l'evaluation")
185
-
186
- st.markdown("<div style='text-align:center: 40px;'></div>", unsafe_allow_html=True)
187
-
188
- if evaluation:
189
- run_eval(selected_dept, form_input)
190
- st.session_state["auto_run_done"] = True
191
-
192
- if valeurs_observees and valeurs_predites:
193
- x = np.arange(len(labels))
194
- width = 0.4
195
- fig, ax = plt.subplots(figsize=(10, 6))
196
-
197
- bars1 = ax.bar(x - width / 2, valeurs_observees, width, label='Observé (2021)', color='steelblue')
198
- bar_colors = [(1, 0, 0, 0.6) if pred > obs else (0, 0.6, 0, 0.6)
199
- for pred, obs in zip(valeurs_predites, valeurs_observees)]
200
- bars2 = ax.bar(x + width / 2, valeurs_predites, width, label='Prevision', color=bar_colors)
201
-
202
- for i in range(len(labels)):
203
- ax.text(x[i] - width / 2, valeurs_observees[i] + max(valeurs_observees) * 0.01, f"{valeurs_observees[i]:,.0f}",
204
- ha='center', va='bottom', fontsize=9)
205
- ax.text(x[i] + width / 2, valeurs_predites[i] + max(valeurs_predites) * 0.01, f"{valeurs_predites[i]:,.0f}",
206
- ha='center', va='bottom', fontsize=9)
207
-
208
- ax.set_ylabel("Tonnes")
209
- ax.set_title("Comparaison Observe vs Predit")
210
- ax.set_xticks(x)
211
- ax.set_xticklabels(labels, rotation=45, ha='right')
212
- ax.legend()
213
- st.pyplot(fig)
214
-
215
- # === Graphiques SHAP ===
216
- st.markdown("---")
217
- st.subheader(f"📉 SHAP - Analyse des contributions pour le departement : {selected_dept}")
218
-
219
- # Menu deroulant
220
- selected_typologie = st.selectbox("Choisissez une typologie de dechets à analyser avec SHAP :", list(model_paths.keys()))
221
-
222
- # SHAP pour la typologie selectionnee
223
- typologie = selected_typologie
224
- path = model_paths[typologie]
225
-
226
- st.markdown(f"### 🔍 {typologie}")
227
-
228
- try:
229
- with open(path, "rb") as f:
230
- model_sm = pickle.load(f)
231
-
232
- used_features = model_sm.model.exog_names
233
- used_features_no_const = [f for f in used_features if f != "const"]
234
- X_used = df[used_features_no_const].copy()
235
-
236
- if "const" in used_features:
237
- X_used["const"] = 1.0
238
-
239
- intercept = model_sm.params['const'] if 'const' in model_sm.params else 0
240
- coefs = model_sm.params[used_features_no_const].values
241
-
242
- lr = LinearRegression()
243
- lr.intercept_ = intercept
244
- lr.coef_ = coefs
245
- lr.feature_names_in_ = np.array(used_features_no_const)
246
-
247
- X_used_corrected = X_used.reindex(columns=lr.feature_names_in_, fill_value=0)
248
-
249
- explainer = shap.Explainer(lr, X_used_corrected)
250
- shap_values = explainer(X_used_corrected)
251
-
252
- selected_index = df[df[f"Departement_{selected_dept}"] == 1].index[0]
253
-
254
- # === Première ligne : Waterfall + Beeswarm + Moyenne des contributions
255
- exclude_vars = [
256
- "Deblais_gravats", "Dechets_verts", "Encombrants", "Materiaux_recyclables"
257
- ]
258
- exclude_vars += [
259
- name for name in shap_values.feature_names
260
- if name.startswith(("Departement_", "Region_"))
261
- ]
262
-
263
- # Creation d’un masque pour filtrer les SHAP plots sans toucher à la prediction
264
- mask = np.array([name not in exclude_vars for name in shap_values.feature_names])
265
- filtered_shap = shap.Explanation(
266
- values=shap_values.values[:, mask],
267
- base_values=shap_values.base_values,
268
- data=shap_values.data[:, mask],
269
- feature_names=[name for name in shap_values.feature_names if name not in exclude_vars]
270
- )
271
-
272
- col1 = st.columns(1)[0]
273
-
274
- with col1:
275
- st.markdown("<h6 style='text-align: center;'>🩜 Waterfall</h6>", unsafe_allow_html=True)
276
- fig = plt.figure(figsize=(3, 2))
277
- shap.plots.waterfall(filtered_shap[selected_index], max_display=10, show=False)
278
- st.pyplot(fig, bbox_inches='tight', dpi=200, clear_figure=True)
279
-
280
-
281
-
282
-
283
-
284
- except Exception as e:
285
- st.error(f"Erreur dans le SHAP pour {typologie}")
286
- st.exception(e)
287
-
288
- # === 🧠 Explication automatique avec Mistral ===
289
-
290
-
291
-
292
- # 1. Recuperation des moyennes absolues des SHAP values
293
- shap_local = filtered_shap[selected_index]
294
-
295
- # 2. Creation d’un resume lisible des coefficients (tries par impact)
296
- sorted_indices = np.argsort(np.abs(shap_local.values))[::-1]
297
- top_n = 10
298
- list_coef = "\n".join([
299
- f"{shap_local.feature_names[i]}: {shap_local.values[i]:.2f}"
300
- for i in sorted_indices[:top_n]
301
- ])
302
-
303
- # 3. Prompt + contexte
304
- prompt_template = f"""
305
- Tu es un expert en data science et en statistique, specialise dans l'interpretation des resultats de modèles explicatifs à l'aide des coefficients de Shapley.
306
-
307
- Je vais te fournir les valeurs des coefficients de Shapley pour un modèle lineaire de regression, associes à chaque variable explicative.
308
-
309
- Ta mission :
310
- Redige un paragraphe clair et synthetique, de 1500 caractères maximum, interpretant le rôle des variables dans le modèle pour repondre au besoin de notre client qui sont les presidents des conseils departementaux. Tu dois identifier :
311
-
312
- - Les variables qui ont le plus d’impact positif ou negatif sur la variable cible.
313
- - Les grandes tendances demographiques ou economiques qui expliquent la production de dechets.
314
- - Une interpretation comprehensible par un public non-expert, mais avec une rigueur statistique.
315
-
316
- Voici les coefficients :
317
- {list_coef}
318
-
319
- Contexte :
320
- - Objectif : Comprendre comment les caracteristiques demographiques et economiques influencent la production des differents types de dechets en France.
321
- - Variable cible : {typologie}
322
- - Modèle utilisé : OLS de Statsmodel avec coefficients de Shapley.
323
- - Variables explicatives :
324
- - Secteurs d’activites : Nombre de salaries par secteur : Agricole, Service, Industrie.
325
- - Profils socioprofessionnels (CSP) :
326
- csp1_agriculteurs, csp2_artisans_commerçant_chef_entreprises, csp3_cadres_professions_intellectuelles, csp4_professions_intermediaires, csp5_employes, csp6_ouvriers, csp7_retraites, csp8_sans_activite.
327
- - Tranches d’âge : tranche_age_0-24, tranche_age_25-59, tranche_age_60+.
328
- - Autres variables :
329
- Densite de population, Population globale, Typologie d'entreprises
330
- - tu ne dois pas prendre en compte les {exclude_vars} dans ton analyse
331
- """
332
-
333
- # 4. Appel à l’API Mistral
334
- # Charge les variables d'environnement à partir du fichier .env
335
- load_dotenv()
336
- api_key = os.getenv("MISTRAL_API_KEY")
337
- #st.write(f"Clé API : `****-****-****-{api_key[-4:]}`")
338
- try:
339
- with st.spinner("🧠 Generation de l'interpretation avec Mistral..."):
340
- # model_llm = ChatMistralAI(model="mistral-large-latest", mistral_api_key=api_key)
341
- model_llm = ChatMistralAI(model="mistral-small-latest", mistral_api_key=api_key)
342
- parser = StrOutputParser()
343
- response = model_llm.invoke(prompt_template)
344
- explanation_text = parser.invoke(response)
345
-
346
- # 5. Affichage dans l’interface Streamlit
347
- st.markdown("#### 🤖 Interpretation automatique (LLM)")
348
- st.success(explanation_text)
349
-
350
- except Exception as e:
351
- st.error("Erreur lors de l'appel au LLM Mistral.")
352
- st.exception(e)
 
1
+ import streamlit as st
2
+ import pandas as pd
3
+ import pickle
4
+ import matplotlib.pyplot as plt
5
+ from dotenv import load_dotenv
6
+ import os
7
+
8
+ import numpy as np
9
+ import shap
10
+ from sklearn.linear_model import LinearRegression
11
+ from langchain_mistralai import ChatMistralAI
12
+ from langchain_core.output_parsers import StrOutputParser
13
+
14
+ # partie initialisation
15
+ model_paths = {
16
+ "Deblais et Gravats": "model_paths/model_Deblais_gravats.pkl",
17
+ "Dechets verts": "model_paths/model_Dechets_verts.pkl",
18
+ "Encombrants": "model_paths/model_Encombrants.pkl",
19
+ "Materiaux recyclables": "model_paths/model_Materiaux_recyclables.pkl"
20
+ }
21
+
22
+ col_mapping = {
23
+ "Deblais et Gravats": "Deblais_gravats",
24
+ "Dechets verts": "Dechets_verts",
25
+ "Encombrants": "Encombrants",
26
+ "Materiaux recyclables": "Materiaux_recyclables"
27
+ }
28
+
29
+ default_dept = "Aisne"
30
+
31
+ valeurs_observees = []
32
+ valeurs_predites = []
33
+ labels = []
34
+
35
+ categories = {
36
+ "📊 Population": [
37
+ "pop_globale",
38
+ "tranche_age_0-24", "tranche_age_25-59", "tranche_age_60+",
39
+ "csp1_agriculteurs", "csp2_artisans_commerçant_chef_entreprises",
40
+ "csp3_cadres_professions_intellectuelles", "csp4_professions_intermediaires",
41
+ "csp5_employes", "csp6_ouvriers", "csp7_retraites", "csp8_sans_activite",
42
+ "densite"
43
+ ],
44
+ "🏭 Activite economique": [
45
+ "nbre_entreprises", "nbre_entreprises_agricole", "nb_salaries_secteur_agricole",
46
+ "nbre_entreprises_industrie", "nb_salaries_secteur_industrie",
47
+ "nb_salaries_secteur_service", "nbre_entreprises_service"
48
+ ]
49
+ }
50
+
51
+ # Pour la gestion automatique du run eval uniquement au demarrage et en cas de changement de departement
52
+ # sinon necessaire de cliquer sur Lancer l'evaluation
53
+ if "previous_dept" not in st.session_state:
54
+ st.session_state["previous_dept"] = None
55
+
56
+ if "auto_run_done" not in st.session_state:
57
+ st.session_state["auto_run_done"] = False
58
+
59
+ def run_eval(selected_dept, form_input):
60
+ # Transformation en DataFrame
61
+ input_df = pd.DataFrame([form_input])
62
+ input_df_complete = row_default.to_frame().T.copy()
63
+ for col in input_df.columns:
64
+ if col in input_df_complete.columns:
65
+ input_df_complete.at[input_df_complete.index[0], col] = input_df.at[0, col]
66
+
67
+ # Verification des incoherences
68
+ Liste_age = ["pop_globale", "tranche_age_0-24", "tranche_age_25-59", "tranche_age_60+"]
69
+ Somme = 0
70
+ if all(v in form_input for v in Liste_age):
71
+ for elt in Liste_age:
72
+ if elt != "pop_globale":
73
+ Somme += form_input[elt]
74
+ if abs(form_input["pop_globale"] - Somme) > 1:
75
+ st.error(f"❌ Incoherence : Population globale = {form_input['pop_globale']} ne correspond pas à la somme des tranches d'âge : {Somme}")
76
+
77
+ Liste_CSP = ["pop_globale","csp1_agriculteurs", "csp2_artisans_commercant_chef_entreprises",
78
+ "csp3_cadres_professions_intellectuelles", "csp4_professions_intermediaires",
79
+ "csp5_employes", "csp6_ouvriers", "csp7_retraites", "csp8_sans_activite"]
80
+ Somme = 0
81
+ if all(v in form_input for v in Liste_CSP):
82
+ for elt in Liste_CSP:
83
+ if elt != "pop_globale":
84
+ Somme += form_input[elt]
85
+ if abs(form_input["pop_globale"] - Somme) >1:
86
+ st.error(f"❌ Incoherence : Population globale = {form_input['pop_globale']} ne correspond pas à la somme des CSP : {Somme}")
87
+
88
+ Liste_Entreprise = ["nbre_entreprises", "nbre_entreprises_agricole",
89
+ "nbre_entreprises_industrie", "nbre_entreprises_service"]
90
+ Somme = 0
91
+ if all(v in form_input for v in Liste_Entreprise):
92
+ for elt in Liste_Entreprise:
93
+ if elt != "nbre_entreprises":
94
+ Somme += form_input[elt]
95
+ if abs(form_input["nbre_entreprises"] - Somme) > 1:
96
+ st.error(f"❌ Incoherence : Le nombre d'entreprises = {form_input['nbre_entreprises']} ne correspond pas à la somme des types d'entreprise : {Somme}")
97
+
98
+
99
+ # on reinitialise pour que ça ne se lance pas automatiquement (lourd)
100
+ evaluation = False
101
+ for typologie, path in model_paths.items():
102
+ try:
103
+ if os.path.exists(path):
104
+ with open(path, "rb") as f:
105
+ model = pickle.load(f)
106
+ else:
107
+ st.error(f"❌ Modèle manquant : {path}")
108
+ expected_cols = model.model.exog_names
109
+ if "const" in expected_cols and "const" not in input_df_complete.columns:
110
+ input_df_complete["const"] = 1.0
111
+
112
+ prediction = max(0, model.predict(input_df_complete[expected_cols]).iloc[0])
113
+ valeurs_predites.append(prediction)
114
+ labels.append(typologie)
115
+
116
+ filtered = observed_df[
117
+ (observed_df["Departement"] == selected_dept) & (observed_df["annee"] == 2021)
118
+ ]
119
+
120
+ excel_col = col_mapping.get(typologie)
121
+ if not filtered.empty and excel_col in filtered.columns:
122
+ valeurs_observees.append(filtered[excel_col].values[0])
123
+ else:
124
+ valeurs_observees.append(0.0)
125
+ except Exception as e:
126
+ st.error(f"Erreur avec le modèle {typologie}")
127
+ st.exception(e)
128
+
129
+ # === Chargement des donnees
130
+ df = pd.read_csv(
131
+ Data/df_dummies.csv").drop(columns=["Unnamed: 0"], errors="ignore")
132
+ observed_df = pd.read_csv(Data/data_wip.csv")
133
+
134
+ # === Liste des departements
135
+ departements = [col.replace("Departement_", "") for col in df.columns if col.startswith("Departement_")]
136
+
137
+ # === Mise en page
138
+ st.set_page_config(layout="wide")
139
+ st.markdown("<h1 style='text-align: left;'>♻️ Simulateur de production de dechets par departement</h1>", unsafe_allow_html=True)
140
+
141
+ st.markdown("<h3 style='text-align: left;'>📍 Choix du departement</h3>", unsafe_allow_html=True)
142
+
143
+ st.markdown("<h3 style='text-align: left;'>📈 Comparaison entre valeurs observees et predites</h3>", unsafe_allow_html=True)
144
+
145
+ st.markdown("<div style='text-align:left: 60px;'></div>", unsafe_allow_html=True)
146
+
147
+ selected_dept = st.selectbox("Selectionner un departement", sorted(departements), index=sorted(departements).index(default_dept))
148
+
149
+ row_default = df[df[f"Departement_{selected_dept}"] == 1].iloc[0]
150
+ default_dict = row_default.to_dict()
151
+
152
+ st.subheader("⚙️ Paramètres modifiables")
153
+ form_input = {}
154
+ for category_name, variables in categories.items():
155
+ with st.expander(category_name, expanded=False):
156
+ for i, var in enumerate(variables):
157
+ if var in default_dict:
158
+ if category_name == "📊 Population":
159
+ default_value = int(round(float(default_dict[var])/100)*100)
160
+ else:
161
+ default_value = int(round(float(default_dict[var])/10)*10)
162
+ val = st.number_input(
163
+ f"✏️ {var}",
164
+ min_value=0,
165
+ value=default_value,
166
+ step=1,
167
+ format="%d",
168
+ key=f"number_input_{selected_dept}_{var}"
169
+ )
170
+ form_input[var] = val
171
+ st.markdown("<div style='margin-bottom: 10px;'></div>", unsafe_allow_html=True)
172
+
173
+ if st.session_state["previous_dept"] != selected_dept or not st.session_state["auto_run_done"]:
174
+ st.session_state["previous_dept"] = selected_dept
175
+ st.session_state["auto_run_done"] = True
176
+ run_eval(selected_dept, form_input)
177
+
178
+
179
+ # with chart_col:
180
+ st.markdown("<div style='text-align:center: 30px;'></div>", unsafe_allow_html=True)
181
+
182
+ # btn_col = st.columns([3, 2, 3])[1]
183
+ # with btn_col:
184
+ evaluation = st.button("🔍 Lancer l'evaluation")
185
+
186
+ st.markdown("<div style='text-align:center: 40px;'></div>", unsafe_allow_html=True)
187
+
188
+ if evaluation:
189
+ run_eval(selected_dept, form_input)
190
+ st.session_state["auto_run_done"] = True
191
+
192
+ if valeurs_observees and valeurs_predites:
193
+ x = np.arange(len(labels))
194
+ width = 0.4
195
+ fig, ax = plt.subplots(figsize=(10, 6))
196
+
197
+ bars1 = ax.bar(x - width / 2, valeurs_observees, width, label='Observé (2021)', color='steelblue')
198
+ bar_colors = [(1, 0, 0, 0.6) if pred > obs else (0, 0.6, 0, 0.6)
199
+ for pred, obs in zip(valeurs_predites, valeurs_observees)]
200
+ bars2 = ax.bar(x + width / 2, valeurs_predites, width, label='Prevision', color=bar_colors)
201
+
202
+ for i in range(len(labels)):
203
+ ax.text(x[i] - width / 2, valeurs_observees[i] + max(valeurs_observees) * 0.01, f"{valeurs_observees[i]:,.0f}",
204
+ ha='center', va='bottom', fontsize=9)
205
+ ax.text(x[i] + width / 2, valeurs_predites[i] + max(valeurs_predites) * 0.01, f"{valeurs_predites[i]:,.0f}",
206
+ ha='center', va='bottom', fontsize=9)
207
+
208
+ ax.set_ylabel("Tonnes")
209
+ ax.set_title("Comparaison Observe vs Predit")
210
+ ax.set_xticks(x)
211
+ ax.set_xticklabels(labels, rotation=45, ha='right')
212
+ ax.legend()
213
+ st.pyplot(fig)
214
+
215
+ # === Graphiques SHAP ===
216
+ st.markdown("---")
217
+ st.subheader(f"📉 SHAP - Analyse des contributions pour le departement : {selected_dept}")
218
+
219
+ # Menu deroulant
220
+ selected_typologie = st.selectbox("Choisissez une typologie de dechets à analyser avec SHAP :", list(model_paths.keys()))
221
+
222
+ # SHAP pour la typologie selectionnee
223
+ typologie = selected_typologie
224
+ path = model_paths[typologie]
225
+
226
+ st.markdown(f"### 🔍 {typologie}")
227
+
228
+ try:
229
+ with open(path, "rb") as f:
230
+ model_sm = pickle.load(f)
231
+
232
+ used_features = model_sm.model.exog_names
233
+ used_features_no_const = [f for f in used_features if f != "const"]
234
+ X_used = df[used_features_no_const].copy()
235
+
236
+ if "const" in used_features:
237
+ X_used["const"] = 1.0
238
+
239
+ intercept = model_sm.params['const'] if 'const' in model_sm.params else 0
240
+ coefs = model_sm.params[used_features_no_const].values
241
+
242
+ lr = LinearRegression()
243
+ lr.intercept_ = intercept
244
+ lr.coef_ = coefs
245
+ lr.feature_names_in_ = np.array(used_features_no_const)
246
+
247
+ X_used_corrected = X_used.reindex(columns=lr.feature_names_in_, fill_value=0)
248
+
249
+ explainer = shap.Explainer(lr, X_used_corrected)
250
+ shap_values = explainer(X_used_corrected)
251
+
252
+ selected_index = df[df[f"Departement_{selected_dept}"] == 1].index[0]
253
+
254
+ # === Première ligne : Waterfall + Beeswarm + Moyenne des contributions
255
+ exclude_vars = [
256
+ "Deblais_gravats", "Dechets_verts", "Encombrants", "Materiaux_recyclables"
257
+ ]
258
+ exclude_vars += [
259
+ name for name in shap_values.feature_names
260
+ if name.startswith(("Departement_", "Region_"))
261
+ ]
262
+
263
+ # Creation d’un masque pour filtrer les SHAP plots sans toucher à la prediction
264
+ mask = np.array([name not in exclude_vars for name in shap_values.feature_names])
265
+ filtered_shap = shap.Explanation(
266
+ values=shap_values.values[:, mask],
267
+ base_values=shap_values.base_values,
268
+ data=shap_values.data[:, mask],
269
+ feature_names=[name for name in shap_values.feature_names if name not in exclude_vars]
270
+ )
271
+
272
+ col1 = st.columns(1)[0]
273
+
274
+ with col1:
275
+ st.markdown("<h6 style='text-align: center;'>🩜 Waterfall</h6>", unsafe_allow_html=True)
276
+ fig = plt.figure(figsize=(3, 2))
277
+ shap.plots.waterfall(filtered_shap[selected_index], max_display=10, show=False)
278
+ st.pyplot(fig, bbox_inches='tight', dpi=200, clear_figure=True)
279
+
280
+
281
+
282
+
283
+
284
+ except Exception as e:
285
+ st.error(f"Erreur dans le SHAP pour {typologie}")
286
+ st.exception(e)
287
+
288
+ # === 🧠 Explication automatique avec Mistral ===
289
+
290
+
291
+
292
+ # 1. Recuperation des moyennes absolues des SHAP values
293
+ shap_local = filtered_shap[selected_index]
294
+
295
+ # 2. Creation d’un resume lisible des coefficients (tries par impact)
296
+ sorted_indices = np.argsort(np.abs(shap_local.values))[::-1]
297
+ top_n = 10
298
+ list_coef = "\n".join([
299
+ f"{shap_local.feature_names[i]}: {shap_local.values[i]:.2f}"
300
+ for i in sorted_indices[:top_n]
301
+ ])
302
+
303
+ # 3. Prompt + contexte
304
+ prompt_template = f"""
305
+ Tu es un expert en data science et en statistique, specialise dans l'interpretation des resultats de modèles explicatifs à l'aide des coefficients de Shapley.
306
+
307
+ Je vais te fournir les valeurs des coefficients de Shapley pour un modèle lineaire de regression, associes à chaque variable explicative.
308
+
309
+ Ta mission :
310
+ Redige un paragraphe clair et synthetique, de 1500 caractères maximum, interpretant le rôle des variables dans le modèle pour repondre au besoin de notre client qui sont les presidents des conseils departementaux. Tu dois identifier :
311
+
312
+ - Les variables qui ont le plus d’impact positif ou negatif sur la variable cible.
313
+ - Les grandes tendances demographiques ou economiques qui expliquent la production de dechets.
314
+ - Une interpretation comprehensible par un public non-expert, mais avec une rigueur statistique.
315
+
316
+ Voici les coefficients :
317
+ {list_coef}
318
+
319
+ Contexte :
320
+ - Objectif : Comprendre comment les caracteristiques demographiques et economiques influencent la production des differents types de dechets en France.
321
+ - Variable cible : {typologie}
322
+ - Modèle utilisé : OLS de Statsmodel avec coefficients de Shapley.
323
+ - Variables explicatives :
324
+ - Secteurs d’activites : Nombre de salaries par secteur : Agricole, Service, Industrie.
325
+ - Profils socioprofessionnels (CSP) :
326
+ csp1_agriculteurs, csp2_artisans_commerçant_chef_entreprises, csp3_cadres_professions_intellectuelles, csp4_professions_intermediaires, csp5_employes, csp6_ouvriers, csp7_retraites, csp8_sans_activite.
327
+ - Tranches d’âge : tranche_age_0-24, tranche_age_25-59, tranche_age_60+.
328
+ - Autres variables :
329
+ Densite de population, Population globale, Typologie d'entreprises
330
+ - tu ne dois pas prendre en compte les {exclude_vars} dans ton analyse
331
+ """
332
+
333
+ # 4. Appel à l’API Mistral
334
+ # Charge les variables d'environnement à partir du fichier .env
335
+ load_dotenv()
336
+ api_key = os.getenv("MISTRAL_API_KEY")
337
+ #st.write(f"Clé API : `****-****-****-{api_key[-4:]}`")
338
+ try:
339
+ with st.spinner("🧠 Generation de l'interpretation avec Mistral..."):
340
+ # model_llm = ChatMistralAI(model="mistral-large-latest", mistral_api_key=api_key)
341
+ model_llm = ChatMistralAI(model="mistral-small-latest", mistral_api_key=api_key)
342
+ parser = StrOutputParser()
343
+ response = model_llm.invoke(prompt_template)
344
+ explanation_text = parser.invoke(response)
345
+
346
+ # 5. Affichage dans l’interface Streamlit
347
+ st.markdown("#### 🤖 Interpretation automatique (LLM)")
348
+ st.success(explanation_text)
349
+
350
+ except Exception as e:
351
+ st.error("Erreur lors de l'appel au LLM Mistral.")
352
+ st.exception(e)