cyrille-elie commited on
Commit
50e58b9
·
1 Parent(s): d1ca45e

Refactor: Improved comments and code structure + test file errors fix for /src/data_processing and /src/modeling forlders + config.py update and model update

Browse files
models/attrition_logistic_regression.joblib CHANGED
Binary files a/models/attrition_logistic_regression.joblib and b/models/attrition_logistic_regression.joblib differ
 
src/config.py CHANGED
@@ -18,7 +18,7 @@ PROCESSED_DATA_PATH = DATA_DIR / "processed" / "data_rh_final.csv"
18
  MODEL_NAME = "attrition_logistic_regression.joblib"
19
  MODEL_PATH = MODELS_DIR / MODEL_NAME
20
 
21
- TARGET_VARIABLE = 'a_quitte_l_entreprise_numeric'
22
 
23
  DB_USER = os.getenv("POSTGRES_USER", "default_user")
24
  DB_PASSWORD = os.getenv("POSTGRES_PASSWORD", "default_password")
@@ -32,12 +32,12 @@ API_VERSION = "0.1.0"
32
 
33
  # --- AJOUT DES MAPPINGS ET CATÉGORIES ICI ---
34
  BINARY_FEATURES_MAPPING = {
35
- 'genre': {'M': 0, 'F': 1}, # Adaptez si 'Masculin'/'Féminin' etc.
36
- 'heure_supplementaires': {'Non': 0, 'Oui': 1} # Adaptez 'Non'/'Oui' si nécessaire.
37
  }
38
 
39
  ORDINAL_FEATURES_CATEGORIES = {
40
- 'frequence_deplacement': ["Aucun", "Occasionnel", "Frequent"],
41
  # Ajoutez vos autres colonnes ordinales et leurs catégories ordonnées ici
42
  }
43
- # --- FIN AJOUT ---
 
18
  MODEL_NAME = "attrition_logistic_regression.joblib"
19
  MODEL_PATH = MODELS_DIR / MODEL_NAME
20
 
21
+ TARGET_VARIABLE = "a_quitte_l_entreprise_numeric"
22
 
23
  DB_USER = os.getenv("POSTGRES_USER", "default_user")
24
  DB_PASSWORD = os.getenv("POSTGRES_PASSWORD", "default_password")
 
32
 
33
  # --- AJOUT DES MAPPINGS ET CATÉGORIES ICI ---
34
  BINARY_FEATURES_MAPPING = {
35
+ "genre": {"M": 0, "F": 1}, # Adaptez si 'Masculin'/'Féminin' etc.
36
+ "heure_supplementaires": {"Non": 0, "Oui": 1}, # Adaptez 'Non'/'Oui' si nécessaire.
37
  }
38
 
39
  ORDINAL_FEATURES_CATEGORIES = {
40
+ "frequence_deplacement": ["Aucun", "Occasionnel", "Frequent"],
41
  # Ajoutez vos autres colonnes ordinales et leurs catégories ordonnées ici
42
  }
43
+ # --- FIN AJOUT ---
src/data_processing/load_data.py CHANGED
@@ -1,136 +1,232 @@
 
 
 
 
 
 
 
 
 
1
  import pandas as pd
2
- # Supprimé: from sqlalchemy import create_engine # Nous utiliserons la session/engine de database_setup
3
- from src.database.database_setup import SessionLocal, engine # Importer SessionLocal et engine
4
- from src.database.models import Employee # Importer votre modèle Employee
 
5
  from src import config
6
  import logging
7
 
8
- logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
 
 
 
9
  logger = logging.getLogger(__name__)
10
 
11
- def load_data_from_csv(path=config.PROCESSED_DATA_PATH):
12
- """Charge les données depuis un fichier CSV (gardé pour référence ou fallback)."""
 
 
 
 
 
 
 
 
 
 
 
 
 
13
  try:
14
  logger.info(f"Chargement des données CSV depuis {path}...")
15
  df = pd.read_csv(path)
16
- logger.info("Données CSV chargées avec succès.")
17
  return df
18
  except FileNotFoundError:
19
  logger.error(f"Fichier CSV non trouvé : {path}")
20
  return None
 
 
 
 
21
 
22
  def load_data_from_postgres() -> pd.DataFrame:
23
- """Charge les données de la table 'employees' depuis PostgreSQL dans un DataFrame."""
24
- db = SessionLocal()
 
 
 
 
 
 
 
 
 
25
  try:
26
- logger.info("Chargement des données depuis la table 'employees' de PostgreSQL...")
27
- # Construire la requête pour sélectionner toutes les colonnes de la table Employee
28
- query = db.query(Employee)
29
- # Exécuter la requête et la charger dans un DataFrame Pandas
30
- df = pd.read_sql_query(query.statement, db.bind) # Utiliser db.bind pour la connexion
31
-
 
32
  if df.empty:
33
- logger.warning("Aucune donnée trouvée dans la table 'employees'. Le DataFrame est vide.")
 
 
34
  else:
35
  logger.info(f"{len(df)} lignes chargées depuis la table 'employees'.")
36
  return df
37
  except Exception as e:
38
- logger.error(f"Erreur lors du chargement des données depuis PostgreSQL : {e}", exc_info=True)
39
- return pd.DataFrame() # Retourner un DataFrame vide en cas d'erreur
 
 
 
40
  finally:
41
  db.close()
 
42
 
43
- def get_data(source: str = "postgres") -> pd.DataFrame: # Changement de la source par défaut
 
44
  """
45
- Fonction principale pour charger les données.
46
- 'source' peut être "postgres" ou "csv".
 
 
 
 
 
 
 
 
 
 
 
47
  """
48
  logger.info(f"Récupération des données depuis la source : {source}")
49
  if source == "postgres":
50
  return load_data_from_postgres()
51
  elif source == "csv":
52
- # Vous pourriez vouloir charger le CSV fusionné et nettoyé si vous l'aviez sauvegardé,
53
- # ou relancer le processus de fusion si vous voulez repartir des bruts.
54
- # Pour l'instant, on garde la logique précédente si 'csv' est appelé.
55
- # Si vous avez un fichier CSV qui correspond à ce qui est dans la BDD:
56
- # return load_data_from_csv(config.PROCESSED_DATA_PATH)
57
- # Sinon, si vous voulez repartir des bruts (ce qui est moins pertinent maintenant):
58
- from .load_data import load_and_merge_csvs # Assurez-vous que cette fonction existe toujours
59
- logger.warning("Chargement depuis les CSV bruts via load_and_merge_csvs. Cette source est moins recommandée maintenant.")
60
- return load_and_merge_csvs() # Si load_and_merge_csvs est toujours et fait le travail
 
 
61
  else:
62
- logger.error(f"Source de données non reconnue : {source}. Utilisation de PostgreSQL par défaut.")
 
 
63
  return load_data_from_postgres()
64
 
65
 
66
- # La fonction load_and_merge_csvs() peut être gardée si vous en avez encore besoin pour
67
- # peupler la base ou pour d'autres usages, sinon elle pourrait être enlevée à terme.
68
- # Assurez-vous qu'elle est toujours présente si get_data(source="csv") l'appelle.
69
- # Pour l'instant, nous la laissons pour ne pas casser le script de peuplement si vous l'avez utilisé.
70
- def load_and_merge_csvs():
71
- """Charge les 3 fichiers CSV bruts, prépare la clé de jointure et les fusionne."""
 
 
 
 
 
 
 
 
 
72
  try:
73
- logger.info("Chargement des fichiers CSV bruts...")
74
  df_sirh = pd.read_csv(config.RAW_SIRH_PATH)
75
  df_eval = pd.read_csv(config.RAW_EVAL_PATH)
76
  df_sondage = pd.read_csv(config.RAW_SONDAGE_PATH)
77
 
78
- logger.info("Préparation de la clé de jointure dans df_eval...")
79
- if 'eval_number' in df_eval.columns:
80
- # S'assurer que eval_number est traité comme string avant split
81
- df_eval['id_employee_str_temp'] = df_eval['eval_number'].astype(str).str.split('_').str[1]
82
-
83
- # Tenter une conversion numérique pour valider le format si nécessaire,
84
- # mais la clé finale pour la jointure sera une string.
85
- numeric_ids_for_check = pd.to_numeric(df_eval['id_employee_str_temp'], errors='coerce')
 
86
  nan_count = numeric_ids_for_check.isnull().sum()
87
  if nan_count > 0:
88
- # Le message exact attendu par votre test est important ici
89
- logger.warning(f"{nan_count} 'eval_number' n'ont pas pu être convertis en 'id_employee' valides.") # Message original
90
-
91
- # La clé de jointure finale doit être une string
92
- df_eval['id_employee'] = df_eval['id_employee_str_temp'].astype(str) # Assurer le type string
93
- df_eval = df_eval.drop(columns=['id_employee_str_temp'])
94
  logger.info("'id_employee' créé et formaté en string dans df_eval.")
95
  else:
96
- logger.error("La colonne 'eval_number' est introuvable dans df_eval.")
97
  return None
98
 
99
- # S'assurer que les id_employee dans les autres tables sont aussi des strings
100
- for df_temp, name in [(df_sirh, 'df_sirh'), (df_sondage, 'df_sondage')]:
101
- if 'id_employee' in df_temp.columns:
102
- df_temp['id_employee'] = df_temp['id_employee'].astype(str) # Assurer le type string
103
- else:
104
- logger.error(f"La colonne 'id_employee' est introuvable dans {name}.")
 
 
 
 
 
 
105
  return None
106
-
107
- logger.info("Fusion des DataFrames...")
108
- df_merged = pd.merge(df_sirh, df_eval, on='id_employee', how='left')
109
- df_merged = pd.merge(df_merged, df_sondage, on='id_employee', how='left')
 
 
 
 
 
 
 
 
 
 
 
110
 
111
- logger.info(f"Données fusionnées : {df_merged.shape}")
112
  return df_merged
113
 
114
  except FileNotFoundError as e:
115
  logger.error(f"Erreur de chargement CSV : Fichier non trouvé - {e}")
116
  return None
117
  except Exception as e:
118
- logger.error(f"Erreur inattendue lors du chargement/fusion CSV : {e}", exc_info=True)
 
 
119
  return None
120
 
121
 
122
- # Pour tester ce module : poetry run python -m src.data_processing.load_data
123
  if __name__ == "__main__":
124
- logger.info("--- Test du chargement depuis PostgreSQL ---")
125
  data_from_db = get_data(source="postgres")
126
  if data_from_db is not None and not data_from_db.empty:
 
127
  print(data_from_db.head())
128
  print(f"\nDimensions des données depuis DB : {data_from_db.shape}")
129
- print(f"\nTypes de données depuis DB :\n{data_from_db.dtypes}")
130
  else:
131
- print("Aucune donnée chargée depuis la base ou une erreur s'est produite.")
132
 
133
- # logger.info("\n--- Test du chargement depuis CSV (si toujours pertinent) ---")
134
- # data_from_csv = get_data(source="csv") # Nécessiterait les CSV bruts
135
- # if data_from_csv is not None:
136
- # print(data_from_csv.head())
 
 
 
 
 
 
1
+ """
2
+ Module pour le chargement et la préparation initiale des données.
3
+
4
+ Fonctions pour charger les données brutes depuis des fichiers CSV,
5
+ les fusionner, et préparer les clés de jointure. Fournit également des fonctions
6
+ pour charger les données depuis une base de données PostgreSQL une fois peuplée.
7
+ La fonction principale `get_data` sert d'interface pour obtenir les données
8
+ pour le reste de l'application, typiquement pour l'entraînement du modèle.
9
+ """
10
  import pandas as pd
11
+ from sqlalchemy.orm import Session # Importé pour être utilisé dans load_data_from_postgres
12
+
13
+ from src.database.database_setup import SessionLocal, engine # engine est utilisé implicitement par read_sql_query via db.bind
14
+ from src.database.models import Employee
15
  from src import config
16
  import logging
17
 
18
+ # Configuration du logging pour ce module
19
+ logging.basicConfig(
20
+ level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
21
+ )
22
  logger = logging.getLogger(__name__)
23
 
24
+
25
+ def load_data_from_csv(path: str = config.PROCESSED_DATA_PATH) -> pd.DataFrame | None:
26
+ """
27
+ Charge les données depuis un unique fichier CSV spécifié.
28
+
29
+ Utilisé comme fonction de secours ou pour charger un dataset déjà traité et sauvegardé en CSV.
30
+
31
+ Args:
32
+ path (str, optional): Chemin vers le fichier CSV.
33
+ Par défaut, utilise config.PROCESSED_DATA_PATH.
34
+
35
+ Returns:
36
+ pd.DataFrame | None: DataFrame Pandas contenant les données chargées,
37
+ ou None si le fichier n'est pas trouvé.
38
+ """
39
  try:
40
  logger.info(f"Chargement des données CSV depuis {path}...")
41
  df = pd.read_csv(path)
42
+ logger.info(f"Données CSV chargées avec succès depuis {path}: {df.shape[0]} lignes.")
43
  return df
44
  except FileNotFoundError:
45
  logger.error(f"Fichier CSV non trouvé : {path}")
46
  return None
47
+ except Exception as e:
48
+ logger.error(f"Erreur inattendue lors du chargement du CSV {path}: {e}", exc_info=True)
49
+ return None
50
+
51
 
52
  def load_data_from_postgres() -> pd.DataFrame:
53
+ """
54
+ Charge l'intégralité des données de la table 'employees' depuis PostgreSQL
55
+ dans un DataFrame Pandas.
56
+
57
+ Utilise SQLAlchemy pour interagir avec la base de données configurée.
58
+
59
+ Returns:
60
+ pd.DataFrame: DataFrame Pandas contenant les données de la table 'employees'.
61
+ Retourne un DataFrame vide en cas d'erreur ou si la table est vide.
62
+ """
63
+ db: Session = SessionLocal()
64
  try:
65
+ logger.info(
66
+ "Chargement des données depuis la table 'employees' de PostgreSQL..."
67
+ )
68
+ query = db.query(Employee) # Construit une requête pour sélectionner toutes les colonnes de Employee
69
+ # Exécute la requête et charge les résultats dans un DataFrame Pandas
70
+ df = pd.read_sql_query(sql=query.statement, con=db.bind)
71
+
72
  if df.empty:
73
+ logger.warning(
74
+ "Aucune donnée trouvée dans la table 'employees'. Le DataFrame est vide."
75
+ )
76
  else:
77
  logger.info(f"{len(df)} lignes chargées depuis la table 'employees'.")
78
  return df
79
  except Exception as e:
80
+ logger.error(
81
+ f"Erreur lors du chargement des données depuis PostgreSQL : {e}",
82
+ exc_info=True,
83
+ )
84
+ return pd.DataFrame() # Retourner un DataFrame vide en cas d'erreur
85
  finally:
86
  db.close()
87
+ logger.debug("Session PostgreSQL fermée pour load_data_from_postgres.")
88
 
89
+
90
+ def get_data(source: str = "postgres") -> pd.DataFrame | None:
91
  """
92
+ Fonction principale pour obtenir le jeu de données pour l'application.
93
+
94
+ Sert d'interface pour charger les données soit depuis PostgreSQL (par défaut),
95
+ soit depuis des fichiers CSV bruts (via `load_and_merge_csvs`).
96
+
97
+ Args:
98
+ source (str, optional): La source des données. Peut être "postgres" ou "csv".
99
+ Par défaut à "postgres".
100
+
101
+ Returns:
102
+ pd.DataFrame | None: DataFrame Pandas contenant les données, ou None si une
103
+ erreur majeure de chargement se produit (ex: CSV introuvables).
104
+ Peut retourner un DataFrame vide si la source est vide (ex: table PG vide).
105
  """
106
  logger.info(f"Récupération des données depuis la source : {source}")
107
  if source == "postgres":
108
  return load_data_from_postgres()
109
  elif source == "csv":
110
+ logger.warning(
111
+ "Chargement depuis les CSV bruts via load_and_merge_csvs. "
112
+ "Cette source est principalement pour le peuplement initial de la BDD."
113
+ )
114
+ # L'import local est inhabituel. S'il est utilisé, il devrait être en haut du fichier.
115
+ # Cependant, si load_and_merge_csvs est spécifique à ce cas, le garder ici
116
+ # évite un import circulaire si load_and_merge_csvs devait importer qqch de ce module.
117
+ # Pour l'instant, nous le laissons, mais c'est un point d'attention.
118
+ # from .load_data import load_and_merge_csvs # Ceci créerait une récursion.
119
+ # Il faut s'assurer que load_and_merge_csvs est bien défini dans ce même fichier.
120
+ return load_and_merge_csvs()
121
  else:
122
+ logger.error(
123
+ f"Source de données non reconnue : {source}. Tentative avec PostgreSQL par défaut."
124
+ )
125
  return load_data_from_postgres()
126
 
127
 
128
+ def load_and_merge_csvs() -> pd.DataFrame | None:
129
+ """
130
+ Charge les données depuis les trois fichiers CSV bruts (`extrait_sirh.csv`,
131
+ `extrait_eval.csv`, `extrait_sondage.csv`), prépare les clés de jointure `id_employee`
132
+ (notamment à partir de `eval_number` et `code_sondage`), et fusionne les DataFrames.
133
+
134
+ Cette fonction est principalement utilisée pour le peuplement initial de la base de données.
135
+ Elle s'assure que les `id_employee` sont traités comme des chaînes de caractères pour
136
+ des fusions cohérentes.
137
+
138
+ Returns:
139
+ pd.DataFrame | None: Un DataFrame fusionné contenant les données des trois sources,
140
+ ou None si une erreur critique se produit (ex: fichier introuvable,
141
+ colonne clé de jointure manquante).
142
+ """
143
  try:
144
+ logger.info("Chargement des fichiers CSV bruts pour fusion...")
145
  df_sirh = pd.read_csv(config.RAW_SIRH_PATH)
146
  df_eval = pd.read_csv(config.RAW_EVAL_PATH)
147
  df_sondage = pd.read_csv(config.RAW_SONDAGE_PATH)
148
 
149
+ # Préparation de df_eval
150
+ logger.info("Préparation de la clé de jointure 'id_employee' dans df_eval à partir de 'eval_number'...")
151
+ if "eval_number" in df_eval.columns:
152
+ df_eval["id_employee_str_temp"] = (
153
+ df_eval["eval_number"].astype(str).str.split("_", n=1).str.get(1) # Prend tout après le premier '_'
154
+ )
155
+ numeric_ids_for_check = pd.to_numeric(
156
+ df_eval["id_employee_str_temp"], errors="coerce"
157
+ )
158
  nan_count = numeric_ids_for_check.isnull().sum()
159
  if nan_count > 0:
160
+ logger.warning(
161
+ f"{nan_count} 'eval_number' (après extraction) n'ont pas pu être convertis "
162
+ f"en id_employee numériques valides et sont devenus NaN."
163
+ )
164
+ df_eval["id_employee"] = df_eval["id_employee_str_temp"].astype(str) # Clé finale en string
165
+ df_eval = df_eval.drop(columns=["id_employee_str_temp"])
166
  logger.info("'id_employee' créé et formaté en string dans df_eval.")
167
  else:
168
+ logger.error("La colonne 'eval_number' est introuvable dans df_eval. Impossible de créer 'id_employee'.")
169
  return None
170
 
171
+ # Préparation de df_sondage
172
+ logger.info("Préparation de la clé de jointure 'id_employee' dans df_sondage à partir de 'code_sondage'...")
173
+ if "code_sondage" in df_sondage.columns: # Supposons que code_sondage est directement l'id_employee
174
+ # Si code_sondage nécessite une transformation similaire à eval_number, appliquez-la ici.
175
+ # Pour cet exemple, on suppose que code_sondage EST l'id_employee.
176
+ df_sondage["id_employee"] = df_sondage["code_sondage"].astype(str)
177
+ # Si vous aviez 'id_employee_str_temp' ici aussi, n'oubliez pas de le drop.
178
+ logger.info("'id_employee' (depuis code_sondage) formaté en string dans df_sondage.")
179
+ else:
180
+ # Si 'code_sondage' n'existe pas, on pourrait vérifier si 'id_employee' existe déjà
181
+ if "id_employee" not in df_sondage.columns:
182
+ logger.error("Ni 'code_sondage' ni 'id_employee' trouvés dans df_sondage.")
183
  return None
184
+ else: # id_employee existe déjà, on s'assure juste du type
185
+ df_sondage["id_employee"] = df_sondage["id_employee"].astype(str)
186
+ logger.info("'id_employee' existant dans df_sondage formaté en string.")
187
+
188
+
189
+ # Standardisation du type de 'id_employee' dans df_sirh
190
+ if "id_employee" in df_sirh.columns:
191
+ df_sirh["id_employee"] = df_sirh["id_employee"].astype(str)
192
+ else:
193
+ logger.error("La colonne 'id_employee' est introuvable dans df_sirh.")
194
+ return None
195
+
196
+ logger.info("Fusion des DataFrames (df_sirh <- df_eval <- df_sondage)...")
197
+ df_merged = pd.merge(df_sirh, df_eval, on="id_employee", how="left")
198
+ df_merged = pd.merge(df_merged, df_sondage, on="id_employee", how="left")
199
 
200
+ logger.info(f"Données fusionnées avec succès : {df_merged.shape[0]} lignes, {df_merged.shape[1]} colonnes.")
201
  return df_merged
202
 
203
  except FileNotFoundError as e:
204
  logger.error(f"Erreur de chargement CSV : Fichier non trouvé - {e}")
205
  return None
206
  except Exception as e:
207
+ logger.error(
208
+ f"Erreur inattendue lors du chargement/fusion CSV : {e}", exc_info=True
209
+ )
210
  return None
211
 
212
 
 
213
  if __name__ == "__main__":
214
+ logger.info("--- Test du chargement des données depuis PostgreSQL ---")
215
  data_from_db = get_data(source="postgres")
216
  if data_from_db is not None and not data_from_db.empty:
217
+ print("Premières lignes de data_from_db:")
218
  print(data_from_db.head())
219
  print(f"\nDimensions des données depuis DB : {data_from_db.shape}")
220
+ # print(f"\nTypes de données depuis DB :\n{data_from_db.dtypes}") # Peut être verbeux
221
  else:
222
+ print("Aucune donnée chargée depuis la base ou une erreur s'est produite lors du chargement.")
223
 
224
+ # Décommentez pour tester le chargement et la fusion des CSV bruts
225
+ # logger.info("\n--- Test du chargement et de la fusion des CSV bruts ---")
226
+ # data_from_csv_merge = get_data(source="csv")
227
+ # if data_from_csv_merge is not None and not data_from_csv_merge.empty:
228
+ # print("\nPremières lignes de data_from_csv_merge:")
229
+ # print(data_from_csv_merge.head())
230
+ # print(f"\nDimensions des données fusionnées depuis CSV : {data_from_csv_merge.shape}")
231
+ # else:
232
+ # print("Aucune donnée chargée depuis les CSV ou une erreur s'est produite lors de la fusion.")
src/data_processing/preprocess.py CHANGED
@@ -1,101 +1,192 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  import pandas as pd
2
  from sklearn.preprocessing import StandardScaler, OneHotEncoder, OrdinalEncoder
3
  from sklearn.compose import ColumnTransformer
4
  from sklearn.pipeline import Pipeline
5
  from sklearn.impute import SimpleImputer
6
- from src import config
7
  import logging
8
- import numpy as np
9
 
 
10
  logging.basicConfig(
11
- level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
12
  )
13
  logger = logging.getLogger(__name__)
14
 
15
 
16
  def map_binary_features(df: pd.DataFrame, binary_cols_map: dict) -> pd.DataFrame:
17
- """Mappe les colonnes binaires spécifiées en 0 et 1."""
18
- df = df.copy()
19
- logger.info("Mappage des features binaires...")
 
 
 
 
 
 
 
 
 
 
 
20
  for col, mapping in binary_cols_map.items():
21
  if col in df.columns:
22
- # S'assurer que les valeurs sont bien des strings avant de mapper si nécessaire
 
 
23
  df[col] = df[col].astype(str).map(mapping)
24
- logger.info(f"Colonne '{col}' mappée en binaire.")
25
- # Vérifier si des NaN ont été introduits (signe d'un problème de mapping)
 
 
 
 
 
 
 
26
  nan_count = df[col].isnull().sum()
27
  if nan_count > 0:
28
  logger.warning(
29
- f"{nan_count} valeurs dans '{col}' n'ont pas pu être mappées et sont devenues NaN."
 
30
  )
31
- # Option: Imputer ici ou laisser SimpleImputer le gérer plus tard
32
- # df[col] = df[col].fillna(df[col].mode()[0] if not df[col].mode().empty else 0)
33
  else:
34
- logger.warning(f"Colonne binaire '{col}' non trouvée.")
35
  return df
36
 
37
 
38
  def clean_data(df: pd.DataFrame) -> pd.DataFrame:
39
- """Applique les étapes de nettoyage et de conversion de type."""
40
- logger.info("Nettoyage des données...")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
41
  df = df.copy()
42
 
 
43
  if "a_quitte_l_entreprise" in df.columns:
44
  df[config.TARGET_VARIABLE] = df["a_quitte_l_entreprise"].map(
45
  {"Oui": 1, "Non": 0}
46
  )
47
- logger.info(f"Colonne '{config.TARGET_VARIABLE}' créée.")
 
 
48
  else:
49
- logger.error("La colonne 'a_quitte_l_entreprise' est manquante.")
50
- raise ValueError("Colonne cible manquante.")
51
-
 
 
 
 
 
 
 
 
 
 
 
 
52
  cols_to_drop = [
53
- "id_employee",
54
- "nombre_heures_travailless",
55
  "eval_number",
56
  "nombre_employee_sous_responsabilite",
57
  "code_sondage",
58
  "ayant_enfants",
59
- "a_quitte_l_entreprise",
60
- "annee_experience_totale",
61
- "niveau_hierarchique_poste",
62
- "annees_dans_le_poste_actuel",
63
- "annes_sous_responsable_actuel",
64
  ]
65
- df = df.drop(
66
- columns=[col for col in cols_to_drop if col in df.columns], errors="ignore"
67
- )
68
- logger.info(f"Colonnes supprimées (si présentes) : {cols_to_drop}")
 
 
 
69
 
 
70
  initial_rows = len(df)
71
  df = df.drop_duplicates()
72
  if len(df) < initial_rows:
73
  logger.info(f"{initial_rows - len(df)} lignes dupliquées supprimées.")
74
 
75
- col_aug = "augementation_salaire_precedente" # Mettez le nom exact
 
76
  if col_aug in df.columns:
77
- logger.info(f"Conversion de '{col_aug}' en numérique...")
78
- df[col_aug] = df[col_aug].astype(str).str.replace(" %", "", regex=False)
79
  df[col_aug] = pd.to_numeric(df[col_aug], errors="coerce")
80
  nan_count = df[col_aug].isnull().sum()
81
  if nan_count > 0:
82
- logger.warning(f"{nan_count} valeurs dans '{col_aug}' sont devenues NaN.")
83
- logger.info(f"'{col_aug}' convertie avec succès.")
 
 
 
84
  else:
85
  logger.warning(
86
- f"La colonne '{col_aug}' n'a pas été trouvée pour la conversion."
87
  )
88
 
89
- logger.info("Nettoyage terminé.")
90
  return df
91
 
92
 
93
  def create_features(df: pd.DataFrame) -> pd.DataFrame:
94
- """Crée de nouvelles features."""
95
- logger.info("Création de features...")
 
 
 
 
 
 
 
 
96
  df = df.copy()
97
- # --- AJOUTEZ VOS FEATURES ICI ---
98
- logger.info("Création de features terminée.")
 
 
 
 
 
 
 
 
 
 
 
99
  return df
100
 
101
 
@@ -103,10 +194,32 @@ def build_preprocessor(
103
  numerical_cols: list,
104
  onehot_cols: list,
105
  ordinal_cols: list,
106
- ordinal_categories: dict,
107
  ) -> ColumnTransformer:
108
- """Construit la pipeline de preprocessing Sklearn."""
109
- logger.info("Construction du préprocesseur Sklearn...")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
110
  transformers_list = []
111
 
112
  if numerical_cols:
@@ -117,6 +230,7 @@ def build_preprocessor(
117
  ]
118
  )
119
  transformers_list.append(("num", numerical_transformer, numerical_cols))
 
120
 
121
  if onehot_cols:
122
  onehot_transformer = Pipeline(
@@ -131,12 +245,16 @@ def build_preprocessor(
131
  ]
132
  )
133
  transformers_list.append(("onehot", onehot_transformer, onehot_cols))
 
 
134
 
135
  if ordinal_cols:
136
  for col in ordinal_cols:
137
- if col not in ordinal_categories:
138
- raise ValueError(f"Les catégories pour '{col}' ne sont pas définies.")
139
- categories_for_col = [ordinal_categories[col]]
 
 
140
  ordinal_transformer_col = Pipeline(
141
  steps=[
142
  ("imputer", SimpleImputer(strategy="most_frequent")),
@@ -145,55 +263,94 @@ def build_preprocessor(
145
  OrdinalEncoder(
146
  categories=categories_for_col,
147
  handle_unknown="use_encoded_value",
148
- unknown_value=-1,
 
149
  ),
150
  ),
151
  ]
152
  )
153
  transformers_list.append((f"ord_{col}", ordinal_transformer_col, [col]))
 
 
154
 
155
  if not transformers_list:
156
- logger.warning("Aucune colonne spécifiée. Préprocesseur vide.")
 
157
  return ColumnTransformer(transformers=[], remainder="passthrough")
158
 
159
  preprocessor = ColumnTransformer(transformers=transformers_list, remainder="drop")
160
- logger.info("Préprocesseur Sklearn construit.")
161
  return preprocessor
162
 
163
 
164
  def run_preprocessing_pipeline(
165
  df: pd.DataFrame,
166
- binary_cols_map: dict = None,
167
- ordinal_cols_categories: dict = None,
168
  preprocessor: ColumnTransformer = None,
169
  fit: bool = False,
170
  ):
171
- """Exécute le pipeline de preprocessing complet."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
172
  df_clean = clean_data(df)
173
 
174
- if binary_cols_map:
175
- df_clean = map_binary_features(df_clean, binary_cols_map)
 
 
 
 
176
 
177
- df_featured = create_features(df_clean)
178
 
179
  if config.TARGET_VARIABLE not in df_featured.columns:
180
  raise ValueError(
181
- f"La colonne cible '{config.TARGET_VARIABLE}' n'est pas présente."
182
  )
183
 
184
  y = df_featured[config.TARGET_VARIABLE]
185
  X = df_featured.drop(config.TARGET_VARIABLE, axis=1)
186
 
187
- ordinal_to_encode = (
188
- list(ordinal_cols_categories.keys()) if ordinal_cols_categories else []
189
- )
190
 
 
 
191
  potential_numerical_cols = X.select_dtypes(include=[np.number]).columns.tolist()
192
- # Les colonnes numériques incluent maintenant les binaires mappées.
193
- # On enlève les ordinales, car elles seront traitées séparément (même si numériques).
194
  numerical_to_scale = [
195
  col for col in potential_numerical_cols if col not in ordinal_to_encode
196
  ]
 
 
197
 
198
  onehot_to_encode = [
199
  col
@@ -201,85 +358,104 @@ def run_preprocessing_pipeline(
201
  if col not in ordinal_to_encode
202
  ]
203
 
 
204
  for col in ordinal_to_encode:
205
  if col not in X.columns:
206
- raise ValueError(f"La colonne ordinale '{col}' spécifiée n'existe pas.")
 
 
207
 
208
- logger.info(f"Colonnes numériques à scaler : {numerical_to_scale}")
209
- logger.info(f"Colonnes catégorielles pour OneHotEncoding : {onehot_to_encode}")
210
- logger.info(f"Colonnes catégorielles pour OrdinalEncoding : {ordinal_to_encode}")
211
 
212
  if fit:
 
213
  processor_instance = build_preprocessor(
214
  numerical_to_scale,
215
  onehot_to_encode,
216
  ordinal_to_encode,
217
- ordinal_cols_categories or {},
218
  )
219
- logger.info("Ajustement (fit) et transformation des données...")
220
  X_processed = processor_instance.fit_transform(X)
221
  try:
222
  feature_names = processor_instance.get_feature_names_out()
223
  except Exception as e:
224
  logger.warning(
225
- f"Impossible d'obtenir les noms de features: {e}. Noms génériques utilisés."
226
  )
227
  feature_names = [f"feature_{i}" for i in range(X_processed.shape[1])]
228
- logger.info("Transformation terminée.")
229
  return (
230
  pd.DataFrame(X_processed, columns=feature_names, index=X.index),
231
  y,
232
  processor_instance,
233
  )
234
- else:
235
  if preprocessor:
236
- logger.info("Transformation des données (sans ajustement)...")
237
  X_processed = preprocessor.transform(X)
238
  try:
239
  feature_names = preprocessor.get_feature_names_out()
240
  except Exception as e:
241
  logger.warning(
242
- f"Impossible d'obtenir les noms de features: {e}. Noms génériques utilisés."
243
  )
244
  feature_names = [f"feature_{i}" for i in range(X_processed.shape[1])]
245
- logger.info("Transformation terminée.")
246
  return pd.DataFrame(X_processed, columns=feature_names, index=X.index), y
247
  else:
 
248
  raise ValueError("Un preprocessor doit être fourni si fit=False.")
249
 
250
 
251
- # Pour tester : poetry run python -m src.data_processing.preprocess
252
  if __name__ == "__main__":
253
- from .load_data import load_and_merge_csvs
254
-
255
- df_raw = load_and_merge_csvs()
256
-
257
- if df_raw is not None:
258
-
259
- try:
 
 
 
 
260
  X_p, y_p, proc = run_preprocessing_pipeline(
261
  df_raw,
262
- binary_cols_map=config.BINARY_FEATURES_MAPPING,
263
- ordinal_cols_categories=config.ORDINAL_FEATURES_CATEGORIES,
264
  fit=True,
265
  )
266
- print("\n--- Preprocessing Réussi ---")
267
  print("Shape X_processed:", X_p.shape)
268
  print("X_processed (head):")
269
  print(X_p.head())
270
-
271
- # Vérifier que les colonnes binaires sont bien numériques et scalées
272
- genre_col = [col for col in X_p.columns if "genre" in col]
273
- hs_col = [col for col in X_p.columns if "heure_supplementaires" in col]
274
-
275
- if genre_col:
276
- print(f"\nColonne 'genre' traitée (scalée) : {genre_col[0]}")
277
- if hs_col:
278
- print(f"Colonne 'heure_supplementaires' traitée (scalée) : {hs_col[0]}")
279
-
280
- except Exception as e:
281
- print("\n--- Erreur lors du preprocessing ---")
282
- print(e)
283
- import traceback
284
-
285
- traceback.print_exc()
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Dans src/data_processing/preprocess.py
2
+ """
3
+ Module de prétraitement des données pour le projet d'attrition RH.
4
+
5
+ Ce module contient les fonctions nécessaires pour :
6
+ - Mapper les features binaires.
7
+ - Nettoyer les données brutes (conversion de la cible, suppression de colonnes,
8
+ gestion des doublons, conversion de types spécifiques comme les pourcentages).
9
+ - Créer de nouvelles features (feature engineering) - actuellement un placeholder.
10
+ - Construire une pipeline de prétraitement Scikit-learn (ColumnTransformer) pour
11
+ l'imputation, la mise à l'échelle des numériques, et l'encodage des catégorielles
12
+ (OneHot et Ordinal).
13
+ - Exécuter l'ensemble de ce pipeline de preprocessing.
14
+ """
15
  import pandas as pd
16
  from sklearn.preprocessing import StandardScaler, OneHotEncoder, OrdinalEncoder
17
  from sklearn.compose import ColumnTransformer
18
  from sklearn.pipeline import Pipeline
19
  from sklearn.impute import SimpleImputer
20
+ from src import config # Pour config.TARGET_VARIABLE, config.BINARY_FEATURES_MAPPING, etc.
21
  import logging
22
+ import numpy as np # Pour np.number lors de la sélection de dtypes
23
 
24
+ # Configuration du logging pour ce module
25
  logging.basicConfig(
26
+ level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
27
  )
28
  logger = logging.getLogger(__name__)
29
 
30
 
31
  def map_binary_features(df: pd.DataFrame, binary_cols_map: dict) -> pd.DataFrame:
32
+ """
33
+ Mappe les valeurs des colonnes catégorielles binaires spécifiées en 0 et 1.
34
+
35
+ Args:
36
+ df (pd.DataFrame): DataFrame d'entrée.
37
+ binary_cols_map (dict): Dictionnaire où les clés sont les noms de colonnes
38
+ et les valeurs sont des dictionnaires de mapping
39
+ (ex: {'Oui': 1, 'Non': 0}).
40
+
41
+ Returns:
42
+ pd.DataFrame: DataFrame avec les colonnes binaires mappées.
43
+ """
44
+ df = df.copy() # Évite les SettingWithCopyWarning
45
+ logger.info("Application du mappage sur les features binaires...")
46
  for col, mapping in binary_cols_map.items():
47
  if col in df.columns:
48
+ # Assurer que la colonne est de type string avant d'appliquer .map
49
+ # Cela évite des erreurs si une colonne est déjà numérique (ex: 0/1)
50
+ # ou contient des types mixtes.
51
  df[col] = df[col].astype(str).map(mapping)
52
+ # Après .map, les valeurs non trouvées dans le mapping deviennent NaN.
53
+ # Il est important que ces NaN soient de type float pour que SimpleImputer (numérique) fonctionne.
54
+ # Si la colonne originale était object et contenait des strings et des NaN après map,
55
+ # SimpleImputer(strategy='most_frequent') fonctionnerait mais SimpleImputer(strategy='median') échouerait.
56
+ # Le plus sûr est de convertir en type numérique si on s'attend à 0 et 1.
57
+ df[col] = pd.to_numeric(df[col], errors='coerce') # Convertit en float, les erreurs (anciens NaN) restent NaN
58
+
59
+ logger.info(f"Colonne '{col}' mappée en binaire numérique.")
60
+
61
  nan_count = df[col].isnull().sum()
62
  if nan_count > 0:
63
  logger.warning(
64
+ f"{nan_count} valeur(s) dans la colonne '{col}' n'ont pas pu être mappée(s) "
65
+ f"correctement et sont devenue(s) NaN (ou étaient déjà NaN)."
66
  )
 
 
67
  else:
68
+ logger.warning(f"Colonne binaire '{col}' spécifiée pour mappage mais non trouvée dans le DataFrame.")
69
  return df
70
 
71
 
72
  def clean_data(df: pd.DataFrame) -> pd.DataFrame:
73
+ """
74
+ Applique les étapes de nettoyage de base et les conversions de types spécifiques.
75
+
76
+ - Convertit la variable cible textuelle en format numérique.
77
+ - Supprime les colonnes jugées inutiles ou redondantes.
78
+ - Supprime les lignes dupliquées.
79
+ - Convertit la colonne 'augementation_salaire_precedente' (texte en "XX %") en numérique.
80
+
81
+ Args:
82
+ df (pd.DataFrame): DataFrame d'entrée brut ou fusionné.
83
+
84
+ Returns:
85
+ pd.DataFrame: DataFrame nettoyé.
86
+
87
+ Raises:
88
+ ValueError: Si la colonne cible 'a_quitte_l_entreprise' est manquante.
89
+ """
90
+ logger.info("Début du nettoyage des données...")
91
  df = df.copy()
92
 
93
+ # Conversion de la variable cible
94
  if "a_quitte_l_entreprise" in df.columns:
95
  df[config.TARGET_VARIABLE] = df["a_quitte_l_entreprise"].map(
96
  {"Oui": 1, "Non": 0}
97
  )
98
+ # Convertir en type entier nullable pour gérer les NaN si map échoue pour certaines valeurs
99
+ df[config.TARGET_VARIABLE] = pd.to_numeric(df[config.TARGET_VARIABLE], errors='coerce').astype('Int64')
100
+ logger.info(f"Colonne cible '{config.TARGET_VARIABLE}' créée et convertie en Int64.")
101
  else:
102
+ # Pour la prédiction, cette colonne ne sera pas présente.
103
+ # On ne devrait pas lever d'erreur si on est en mode prédiction.
104
+ # Cette fonction est appelée par run_preprocessing_pipeline qui sépare X et y *après*.
105
+ # Et aussi par populate_db qui a besoin de la cible.
106
+ # Pour predict.py, on ajoute une colonne factice 'a_quitte_l_entreprise'.
107
+ # Donc, cette condition est surtout pour la robustesse lors de l'entraînement.
108
+ logger.warning("La colonne 'a_quitte_l_entreprise' est manquante. Si c'est pour la prédiction, c'est normal.")
109
+ # Si cette fonction est appelée par un flux où la cible DOIT être là (comme populate_db ou train_model avant split X/y)
110
+ # alors une erreur est appropriée.
111
+ # Pour l'instant, on logue un warning, le check de présence de TARGET_VARIABLE se fera plus tard.
112
+
113
+
114
+ # Liste des colonnes à supprimer (adaptez si nécessaire)
115
+ # 'id_employee' est conservé pour le moment, car il est utilisé par populate_db
116
+ # et sera explicitement retiré de X dans train_model.py avant l'entraînement.
117
  cols_to_drop = [
118
+ "nombre_heures_travailless", # Semble être une coquille (s en trop)
 
119
  "eval_number",
120
  "nombre_employee_sous_responsabilite",
121
  "code_sondage",
122
  "ayant_enfants",
123
+ "a_quitte_l_entreprise", # Version texte de la cible, une fois la version numérique créée
124
+ "annee_experience_totale", # Supposée redondante ou moins pertinente
125
+ "niveau_hierarchique_poste", # Supposée redondante ou moins pertinente
126
+ "annees_dans_le_poste_actuel", # Supposée redondante ou moins pertinente
127
+ "annes_sous_responsable_actuel", # Coquille "annes" -> "annees" ?
128
  ]
129
+ actual_cols_to_drop = [col for col in cols_to_drop if col in df.columns]
130
+ if actual_cols_to_drop:
131
+ df = df.drop(columns=actual_cols_to_drop, errors="ignore")
132
+ logger.info(f"Colonnes supprimées : {actual_cols_to_drop}")
133
+ else:
134
+ logger.info("Aucune colonne de la liste `cols_to_drop` n'a été trouvée pour suppression.")
135
+
136
 
137
+ # Suppression des doublons
138
  initial_rows = len(df)
139
  df = df.drop_duplicates()
140
  if len(df) < initial_rows:
141
  logger.info(f"{initial_rows - len(df)} lignes dupliquées supprimées.")
142
 
143
+ # Conversion de 'augementation_salaire_precedente'
144
+ col_aug = "augementation_salaire_precedente"
145
  if col_aug in df.columns:
146
+ logger.info(f"Conversion de la colonne '{col_aug}' en numérique...")
147
+ df[col_aug] = df[col_aug].astype(str).str.replace(" %", "", regex=False).str.strip()
148
  df[col_aug] = pd.to_numeric(df[col_aug], errors="coerce")
149
  nan_count = df[col_aug].isnull().sum()
150
  if nan_count > 0:
151
+ logger.warning(
152
+ f"{nan_count} valeur(s) dans '{col_aug}' n'ont pas pu être converties "
153
+ f"en numérique et sont devenue(s) NaN."
154
+ )
155
+ logger.info(f"Colonne '{col_aug}' convertie avec succès en type numérique.")
156
  else:
157
  logger.warning(
158
+ f"La colonne '{col_aug}' n'a pas été trouvée pour la conversion numérique."
159
  )
160
 
161
+ logger.info("Nettoyage des données terminé.")
162
  return df
163
 
164
 
165
  def create_features(df: pd.DataFrame) -> pd.DataFrame:
166
+ """
167
+ Crée de nouvelles features (ingénierie des features) à partir des colonnes existantes.
168
+
169
+ Args:
170
+ df (pd.DataFrame): DataFrame d'entrée (généralement après nettoyage et mappage binaire).
171
+
172
+ Returns:
173
+ pd.DataFrame: DataFrame avec les nouvelles features ajoutées.
174
+ """
175
+ logger.info("Début de la création de features...")
176
  df = df.copy()
177
+
178
+ # --- EXEMPLE DE FEATURE A AJOUTER ---
179
+ # Si vous voulez ajouter une feature 'satisfaction_moyenne' basée sur les colonnes Q_...
180
+ # sondage_cols = [col for col in df.columns if col.startswith('Q_') and col not in ['Q_satisfaction_generale']]
181
+ # if sondage_cols:
182
+ # df['satisfaction_moyenne_autres_sondages'] = df[sondage_cols].mean(axis=1, skipna=True)
183
+ # logger.info("Feature 'satisfaction_moyenne_autres_sondages' créée.")
184
+ # else:
185
+ # logger.info("Aucune colonne de sondage (Q_...) trouvée pour calculer 'satisfaction_moyenne_autres_sondages'.")
186
+
187
+ # Pour l'instant, cette fonction ne fait rien de plus.
188
+ # Vous pouvez ajouter ici la logique de création de features que vous aviez dans vos notebooks.
189
+ logger.info("Création de features terminée (actuellement, pas de nouvelles features implémentées).")
190
  return df
191
 
192
 
 
194
  numerical_cols: list,
195
  onehot_cols: list,
196
  ordinal_cols: list,
197
+ ordinal_categories_map: dict, # Renommé pour correspondre à l'usage dans la fct
198
  ) -> ColumnTransformer:
199
+ """
200
+ Construit et retourne un objet ColumnTransformer de Scikit-learn pour le prétraitement.
201
+
202
+ Le ColumnTransformer applique :
203
+ - Imputation par la médiane puis StandardScaler aux colonnes numériques.
204
+ - Imputation par la valeur la plus fréquente puis OneHotEncoder aux colonnes catégorielles nominales.
205
+ - Imputation par la valeur la plus fréquente puis OrdinalEncoder aux colonnes catégorielles ordinales.
206
+
207
+ Args:
208
+ numerical_cols (list): Liste des noms des colonnes numériques.
209
+ onehot_cols (list): Liste des noms des colonnes catégorielles à encoder en One-Hot.
210
+ ordinal_cols (list): Liste des noms des colonnes catégorielles ordinales.
211
+ ordinal_categories_map (dict): Dictionnaire spécifiant l'ordre des catégories
212
+ pour chaque colonne ordinale.
213
+ Format: {'nom_col_ord': ['cat1', 'cat2', ...]}
214
+
215
+ Returns:
216
+ ColumnTransformer: Objet ColumnTransformer configuré mais non ajusté.
217
+
218
+ Raises:
219
+ ValueError: Si les catégories pour une colonne ordinale ne sont pas définies
220
+ dans `ordinal_categories_map`.
221
+ """
222
+ logger.info("Construction du préprocesseur Sklearn (ColumnTransformer)...")
223
  transformers_list = []
224
 
225
  if numerical_cols:
 
230
  ]
231
  )
232
  transformers_list.append(("num", numerical_transformer, numerical_cols))
233
+ logger.info(f"Transformateur numérique configuré pour les colonnes : {numerical_cols}")
234
 
235
  if onehot_cols:
236
  onehot_transformer = Pipeline(
 
245
  ]
246
  )
247
  transformers_list.append(("onehot", onehot_transformer, onehot_cols))
248
+ logger.info(f"Transformateur OneHot configuré pour les colonnes : {onehot_cols}")
249
+
250
 
251
  if ordinal_cols:
252
  for col in ordinal_cols:
253
+ if col not in ordinal_categories_map: # Utilisation de ordinal_categories_map
254
+ raise ValueError(
255
+ f"Les catégories pour la colonne ordinale '{col}' ne sont pas définies dans ordinal_categories_map."
256
+ )
257
+ categories_for_col = [ordinal_categories_map[col]] # Utilisation de ordinal_categories_map
258
  ordinal_transformer_col = Pipeline(
259
  steps=[
260
  ("imputer", SimpleImputer(strategy="most_frequent")),
 
263
  OrdinalEncoder(
264
  categories=categories_for_col,
265
  handle_unknown="use_encoded_value",
266
+ unknown_value= np.nan # Utiliser np.nan, qui sera géré par l'imputer numérique si la colonne devient numérique
267
+ # ou -1 si vous préférez et que l'imputer numérique n'est pas appliqué ensuite
268
  ),
269
  ),
270
  ]
271
  )
272
  transformers_list.append((f"ord_{col}", ordinal_transformer_col, [col]))
273
+ logger.info(f"Transformateurs ordinaux configurés pour les colonnes : {ordinal_cols}")
274
+
275
 
276
  if not transformers_list:
277
+ logger.warning("Aucune colonne spécifiée pour le preprocessing. Le ColumnTransformer sera vide et utilisera remainder='passthrough'.")
278
+ # Retourner un transformateur qui ne fait rien mais ne plante pas.
279
  return ColumnTransformer(transformers=[], remainder="passthrough")
280
 
281
  preprocessor = ColumnTransformer(transformers=transformers_list, remainder="drop")
282
+ logger.info("Préprocesseur Sklearn (ColumnTransformer) construit avec succès.")
283
  return preprocessor
284
 
285
 
286
  def run_preprocessing_pipeline(
287
  df: pd.DataFrame,
288
+ binary_cols_map: dict = None, # Doit venir de config.BINARY_FEATURES_MAPPING
289
+ ordinal_cols_categories_map: dict = None, # Doit venir de config.ORDINAL_FEATURES_CATEGORIES. Renommé pour cohérence.
290
  preprocessor: ColumnTransformer = None,
291
  fit: bool = False,
292
  ):
293
+ """
294
+ Exécute le pipeline de preprocessing complet sur le DataFrame fourni.
295
+
296
+ Orchestre les étapes de nettoyage, mappage binaire, création de features,
297
+ et application (ajustement ou transformation) du ColumnTransformer.
298
+
299
+ Args:
300
+ df (pd.DataFrame): DataFrame d'entrée brut.
301
+ binary_cols_map (dict, optional): Mapping pour les features binaires.
302
+ Utilise config.BINARY_FEATURES_MAPPING si non fourni (dans l'appelant).
303
+ ordinal_cols_categories_map (dict, optional): Catégories pour les features ordinales.
304
+ Utilise config.ORDINAL_FEATURES_CATEGORIES si non fourni (dans l'appelant).
305
+ preprocessor (ColumnTransformer, optional): Un ColumnTransformer pré-ajusté.
306
+ Requis si `fit` est False.
307
+ fit (bool, optional): Si True, le préprocesseur est ajusté (`fit_transform`) sur les données.
308
+ Si False, le `preprocessor` fourni est utilisé pour transformer (`transform`) les données.
309
+ Par défaut à False.
310
+
311
+ Returns:
312
+ Tuple: Contenant selon le mode `fit`:
313
+ Si `fit` est True: (X_processed_df, y_series, fitted_processor_instance)
314
+ Si `fit` est False: (X_processed_df, y_series)
315
+ Où X_processed_df est un DataFrame et y_series est une Series.
316
+
317
+ Raises:
318
+ ValueError: Si la colonne cible est manquante après les premières étapes,
319
+ ou si `fit` est False et aucun `preprocessor` n'est fourni,
320
+ ou si une colonne ordinale spécifiée n'existe pas dans le DataFrame.
321
+ """
322
+ logger.info(f"Exécution du pipeline de preprocessing (fit={fit})...")
323
  df_clean = clean_data(df)
324
 
325
+ # Utiliser les mappings passés en argument, ou ceux de config si non fournis par l'appelant
326
+ current_binary_map = binary_cols_map if binary_cols_map is not None else config.BINARY_FEATURES_MAPPING
327
+ if current_binary_map: # Vérifier si le mapping est non vide/None
328
+ df_mapped = map_binary_features(df_clean, current_binary_map)
329
+ else:
330
+ df_mapped = df_clean.copy() # Pas de mappage binaire à faire
331
 
332
+ df_featured = create_features(df_mapped)
333
 
334
  if config.TARGET_VARIABLE not in df_featured.columns:
335
  raise ValueError(
336
+ f"La colonne cible '{config.TARGET_VARIABLE}' n'est pas présente après clean/map/feature."
337
  )
338
 
339
  y = df_featured[config.TARGET_VARIABLE]
340
  X = df_featured.drop(config.TARGET_VARIABLE, axis=1)
341
 
342
+ # Utiliser les catégories ordinales passées en argument, ou celles de config
343
+ current_ordinal_map = ordinal_cols_categories_map if ordinal_cols_categories_map is not None else config.ORDINAL_FEATURES_CATEGORIES
344
+ ordinal_to_encode = list(current_ordinal_map.keys()) if current_ordinal_map else []
345
 
346
+
347
+ # Identification des types de colonnes pour le ColumnTransformer
348
  potential_numerical_cols = X.select_dtypes(include=[np.number]).columns.tolist()
 
 
349
  numerical_to_scale = [
350
  col for col in potential_numerical_cols if col not in ordinal_to_encode
351
  ]
352
+ # Les colonnes binaires mappées en 0/1 sont numériques et seront scalées par défaut ici.
353
+ # Si on veut les exclure du scaling, il faudrait les retirer de numerical_to_scale.
354
 
355
  onehot_to_encode = [
356
  col
 
358
  if col not in ordinal_to_encode
359
  ]
360
 
361
+ # Vérifier que les colonnes ordinales existent bien dans X avant de les passer à build_preprocessor
362
  for col in ordinal_to_encode:
363
  if col not in X.columns:
364
+ # Cela peut arriver si une colonne listée dans ORDINAL_FEATURES_CATEGORIES a été droppée
365
+ # ou n'était pas dans le df initial.
366
+ raise ValueError(f"La colonne ordinale '{col}' spécifiée dans ordinal_cols_categories_map n'existe pas dans le DataFrame X.")
367
 
368
+ logger.info(f"Colonnes identifiées pour la mise à l'échelle (numériques) : {numerical_to_scale}")
369
+ logger.info(f"Colonnes identifiées pour OneHotEncoding : {onehot_to_encode}")
370
+ logger.info(f"Colonnes identifiées pour OrdinalEncoding : {ordinal_to_encode}")
371
 
372
  if fit:
373
+ # Construire le préprocesseur avec les catégories ordinales actuelles
374
  processor_instance = build_preprocessor(
375
  numerical_to_scale,
376
  onehot_to_encode,
377
  ordinal_to_encode,
378
+ current_ordinal_map, # Passer le dictionnaire complet
379
  )
380
+ logger.info("Ajustement (fit) et transformation des données X...")
381
  X_processed = processor_instance.fit_transform(X)
382
  try:
383
  feature_names = processor_instance.get_feature_names_out()
384
  except Exception as e:
385
  logger.warning(
386
+ f"Impossible d'obtenir les noms de features via get_feature_names_out(): {e}. Noms génériques utilisés."
387
  )
388
  feature_names = [f"feature_{i}" for i in range(X_processed.shape[1])]
389
+ logger.info("Transformation X terminée.")
390
  return (
391
  pd.DataFrame(X_processed, columns=feature_names, index=X.index),
392
  y,
393
  processor_instance,
394
  )
395
+ else: # Mode transform uniquement
396
  if preprocessor:
397
+ logger.info("Transformation des données X (sans ajustement) en utilisant le préprocesseur fourni...")
398
  X_processed = preprocessor.transform(X)
399
  try:
400
  feature_names = preprocessor.get_feature_names_out()
401
  except Exception as e:
402
  logger.warning(
403
+ f"Impossible d'obtenir les noms de features via get_feature_names_out(): {e}. Noms génériques utilisés."
404
  )
405
  feature_names = [f"feature_{i}" for i in range(X_processed.shape[1])]
406
+ logger.info("Transformation X terminée.")
407
  return pd.DataFrame(X_processed, columns=feature_names, index=X.index), y
408
  else:
409
+ logger.error("Erreur: `fit` est False mais aucun `preprocessor` n'a été fourni.")
410
  raise ValueError("Un preprocessor doit être fourni si fit=False.")
411
 
412
 
413
+ # Bloc de test pour une exécution directe (poetry run python -m src.data_processing.preprocess)
414
  if __name__ == "__main__":
415
+ logger.info("--- Démarrage du test direct du module preprocess.py ---")
416
+ # Importer load_and_merge_csvs ici pour éviter l'import circulaire au niveau du module
417
+ # si load_data devait importer qqch de preprocess (ce qui n'est pas le cas actuellement)
418
+ try:
419
+ from src.data_processing.load_data import load_and_merge_csvs
420
+ df_raw = load_and_merge_csvs()
421
+
422
+ if df_raw is not None and not df_raw.empty:
423
+ logger.info(f"Données brutes chargées pour le test: {df_raw.shape[0]} lignes.")
424
+ # Test du pipeline de preprocessing complet en mode fit=True
425
+ # Les mappings et catégories sont pris depuis config.py par défaut dans run_preprocessing_pipeline
426
  X_p, y_p, proc = run_preprocessing_pipeline(
427
  df_raw,
428
+ # binary_cols_map et ordinal_cols_categories sont pris de config par défaut
 
429
  fit=True,
430
  )
431
+ print("\n--- Preprocessing Réussi (fit=True) ---")
432
  print("Shape X_processed:", X_p.shape)
433
  print("X_processed (head):")
434
  print(X_p.head())
435
+ print("\ny (head):")
436
+ print(y_p.head())
437
+ print("\nColonnes après preprocessing:", X_p.columns.tolist())
438
+
439
+ # Exemple de test en mode transform=True avec le processeur fitté
440
+ # On prend un échantillon des données brutes pour simuler de nouvelles données
441
+ if len(df_raw) > 5:
442
+ df_sample_for_transform = df_raw.sample(n=min(5, len(df_raw)), random_state=42)
443
+ logger.info(f"\nTest du mode transform sur {len(df_sample_for_transform)} échantillons...")
444
+ X_t, y_t = run_preprocessing_pipeline(
445
+ df_sample_for_transform,
446
+ preprocessor=proc, # Utiliser le processeur fitté
447
+ fit=False
448
+ )
449
+ print("\n--- Preprocessing Réussi (fit=False) ---")
450
+ print("Shape X_transformed:", X_t.shape)
451
+ print("X_transformed (head):")
452
+ print(X_t.head())
453
+ else:
454
+ logger.error("Impossible de charger les données brutes pour le test du module preprocess.")
455
+
456
+ except Exception as e_main:
457
+ print(f"\n--- Erreur lors de l'exécution du test de preprocess.py ---")
458
+ print(e_main)
459
+ import traceback
460
+ traceback.print_exc()
461
+ logger.info("--- Fin du test direct du module preprocess.py ---")
src/modeling/predict.py CHANGED
@@ -1,141 +1,201 @@
 
 
 
 
 
 
 
 
 
1
  import pandas as pd
2
  from joblib import load
3
  import logging
4
- from src import config
5
- from typing import List, Dict, Union
6
 
7
- # Importer les fonctions de preprocessing nécessaires !
 
8
  from src.data_processing.preprocess import (
9
  clean_data,
10
  map_binary_features,
11
  create_features,
12
  )
13
 
 
14
  logging.basicConfig(
15
- level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
16
  )
17
  logger = logging.getLogger(__name__)
18
 
19
- _pipeline = None
 
20
 
21
- # --- FIN DÉFINITION ---
22
 
 
 
 
 
 
 
23
 
24
- def load_prediction_pipeline():
25
- """Charge la pipeline complète (preprocessor + modèle) depuis le disque."""
 
 
26
  global _pipeline
27
- if _pipeline is None:
28
  if not config.MODEL_PATH.exists():
29
- logger.error(f"Fichier pipeline non trouvé : {config.MODEL_PATH}")
30
  return None
31
  try:
32
- logger.info(f"Chargement de la pipeline depuis {config.MODEL_PATH}...")
33
  _pipeline = load(config.MODEL_PATH)
34
- logger.info("Pipeline chargée avec succès.")
35
  except Exception as e:
36
  logger.error(
37
- f"Erreur lors du chargement de la pipeline : {e}", exc_info=True
38
  )
39
- _pipeline = None
40
  return None
41
  return _pipeline
42
 
43
 
44
- def predict_attrition(input_data: pd.DataFrame) -> Union[List[Dict], Dict]:
45
  """
46
- Fait une prédiction sur de nouvelles données brutes (DataFrame).
47
- Applique les transformations manuelles PUIS la pipeline sauvegardée.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
48
  """
49
  pipeline = load_prediction_pipeline()
50
  if pipeline is None:
51
- return {"error": "Pipeline non chargée. Veuillez entraîner le modèle."}
 
52
 
53
  try:
54
- logger.info(f"Prédiction sur {len(input_data)} enregistrement(s)...")
 
 
 
55
 
56
  # --- ÉTAPE 1 : Appliquer les transformations manuelles ---
57
- # Note: clean_data attend 'a_quitte_l_entreprise', on doit l'ajouter
58
- # temporairement si elle n'y est pas, ou modifier clean_data.
59
- # Pour la prédiction, il est plus simple de ne pas l'exiger.
60
- # Modifions légèrement l'appel ou clean_data.
61
- # Ici, on suppose que clean_data peut fonctionner sans la cible,
62
- # ou on la modifie pour qu'elle puisse.
63
- # Pour l'instant, on ajoute une colonne 'bidon' pour que ça passe.
64
- if "a_quitte_l_entreprise" not in input_data.columns:
65
- input_data["a_quitte_l_entreprise"] = "Non" # Valeur arbitraire
66
-
67
- df_cleaned = clean_data(input_data)
68
  df_mapped = map_binary_features(df_cleaned, config.BINARY_FEATURES_MAPPING)
69
  df_featured = create_features(df_mapped)
70
 
71
- # Enlever la colonne cible si on l'a ajoutée ou si elle était là
72
- X_predict = df_featured.drop(config.TARGET_VARIABLE, axis=1, errors="ignore")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
73
 
74
- logger.info("Transformations manuelles appliquées.")
75
 
76
  # --- ÉTAPE 2 : Utiliser la pipeline complète pour prédire ---
77
- # La pipeline va maintenant appliquer le ColumnTransformer (scaling, OHE, etc.)
78
- # ET ensuite le classifier.
79
  probabilities = pipeline.predict_proba(X_predict)[:, 1]
80
  predictions = pipeline.predict(X_predict)
 
81
 
82
  results = []
83
- for i, index in enumerate(X_predict.index):
84
  results.append(
85
  {
86
- "id_employe": index,
87
  "probabilite_depart": float(probabilities[i]),
88
  "prediction_depart": int(predictions[i]),
89
  }
90
  )
91
- logger.info("Prédiction terminée.")
92
  return results
93
 
94
  except Exception as e:
95
- logger.error(f"Erreur lors de la prédiction : {e}", exc_info=True)
96
- return {"error": f"Erreur de prédiction: {e}"}
97
 
98
 
99
- # Pour tester : poetry run python -m src.modeling.predict
100
  if __name__ == "__main__":
101
- # Créez un exemple de données BRUTES (comme si elles venaient d'un nouveau formulaire ou BDD)
102
- # Doit contenir TOUTES les colonnes présentes dans les données avant le preprocessing.
103
- # Utilisez les mêmes noms que dans vos fichiers CSV initiaux.
104
- sample_data = pd.DataFrame(
105
- {
106
- "age": 45,
107
- "genre": "M",
108
- "revenu_mensuel": 4850,
109
- "statut_marital": "Célibataire",
110
- "departement": "Commercial",
111
- "poste": "Cadre Commercial",
112
- "nombre_experiences_precedentes": 8,
113
- "annees_dans_l_entreprise": 5,
114
- "satisfaction_employee_environnement": 4,
115
- "note_evaluation_precedente": 3,
116
- "satisfaction_employee_nature_travail": 3,
117
- "satisfaction_employee_equipe": 3,
118
- "satisfaction_employee_equilibre_pro_perso": 3,
119
- "note_evaluation_actuelle": 3,
120
- "heure_supplementaires": "Non",
121
- "augementation_salaire_precedente": 15,
122
- "nombre_participation_pee": 0,
123
- "nb_formations_suivies": 3,
124
- "distance_domicile_travail": 20,
125
- "niveau_education": 3,
126
- "domaine_etude": "Infra & Cloud",
127
- "frequence_deplacement": "Occasionnel",
128
- "annees_depuis_la_derniere_promotion": 0,
129
- },
130
- index=["EMP_TEST_1"],
131
- ) # Donner des index pour l'ID
132
-
133
- # Vérifier si le modèle existe avant de tester
134
  if config.MODEL_PATH.exists():
135
- predictions_output = predict_attrition(sample_data)
 
136
  print("\n--- Résultat de la Prédiction Test ---")
137
- print(predictions_output)
 
 
 
 
138
  else:
139
  print(
140
- f"\nModèle non trouvé à {config.MODEL_PATH}. Veuillez l'entraîner d'abord."
141
  )
 
 
1
+ """
2
+ Module pour charger la pipeline de Machine Learning entraînée et effectuer des prédictions.
3
+
4
+ Fonctions principales :
5
+ - `load_prediction_pipeline()`: Charge la pipeline Scikit-learn sauvegardée (modèle + préprocesseur).
6
+ - `predict_attrition()`: Prend des données brutes d'employés en entrée (DataFrame),
7
+ applique les transformations de preprocessing initiales, puis utilise la pipeline
8
+ chargée pour prédire la probabilité d'attrition et la classe de départ.
9
+ """
10
  import pandas as pd
11
  from joblib import load
12
  import logging
13
+ from typing import List, Dict, Union, Any # Ajout de Any
 
14
 
15
+ from src import config # Pour MODEL_PATH et BINARY_FEATURES_MAPPING
16
+ # Importer les fonctions de preprocessing nécessaires
17
  from src.data_processing.preprocess import (
18
  clean_data,
19
  map_binary_features,
20
  create_features,
21
  )
22
 
23
+ # Configuration du logging pour ce module
24
  logging.basicConfig(
25
+ level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
26
  )
27
  logger = logging.getLogger(__name__)
28
 
29
+ # Variable globale pour stocker la pipeline chargée et éviter de la recharger à chaque appel
30
+ _pipeline: Any = None # Utilisation de Any car le type exact de la pipeline peut être complexe
31
 
 
32
 
33
+ def load_prediction_pipeline() -> Any | None:
34
+ """
35
+ Charge la pipeline de prédiction complète (préprocesseur + modèle) depuis le disque.
36
+
37
+ La pipeline est chargée une seule fois et stockée dans une variable globale `_pipeline`
38
+ pour optimiser les appels suivants.
39
 
40
+ Returns:
41
+ Pipeline | None: L'objet pipeline Scikit-learn chargé, ou None si une erreur
42
+ se produit (ex: fichier non trouvé, erreur de chargement).
43
+ """
44
  global _pipeline
45
+ if _pipeline is None: # Charger seulement si pas déjà en mémoire
46
  if not config.MODEL_PATH.exists():
47
+ logger.error(f"Fichier pipeline non trouvé à l'emplacement configuré : {config.MODEL_PATH}")
48
  return None
49
  try:
50
+ logger.info(f"Chargement de la pipeline de prédiction depuis : {config.MODEL_PATH}...")
51
  _pipeline = load(config.MODEL_PATH)
52
+ logger.info("Pipeline de prédiction chargée avec succès.")
53
  except Exception as e:
54
  logger.error(
55
+ f"Erreur critique lors du chargement de la pipeline : {e}", exc_info=True
56
  )
57
+ _pipeline = None # S'assurer que _pipeline est None en cas d'échec
58
  return None
59
  return _pipeline
60
 
61
 
62
+ def predict_attrition(input_data: pd.DataFrame) -> Union[List[Dict[str, Any]], Dict[str, str]]:
63
  """
64
+ Prédit le risque d'attrition pour les employés fournis en entrée.
65
+
66
+ Cette fonction prend un DataFrame de données brutes, applique les étapes
67
+ de preprocessing initiales (nettoyage, mappage binaire, création de features)
68
+ puis utilise la pipeline Scikit-learn entraînée et chargée pour effectuer
69
+ les prédictions.
70
+
71
+ Args:
72
+ input_data (pd.DataFrame): DataFrame contenant les données brutes des employés
73
+ pour lesquels faire une prédiction. Les colonnes doivent
74
+ correspondre aux attentes de `clean_data` et
75
+ `map_binary_features` (format brut).
76
+
77
+ Returns:
78
+ Union[List[Dict[str, Any]], Dict[str, str]]:
79
+ - Une liste de dictionnaires si la prédiction réussit, chaque dictionnaire
80
+ contenant 'id_employe', 'probabilite_depart', 'prediction_depart'.
81
+ - Un dictionnaire d'erreur `{"error": "message"}` si la pipeline n'est pas
82
+ chargée ou si une autre erreur majeure se produit.
83
  """
84
  pipeline = load_prediction_pipeline()
85
  if pipeline is None:
86
+ # Le logger dans load_prediction_pipeline aura déjà enregistré l'erreur spécifique.
87
+ return {"error": "Pipeline de prédiction non chargée. Le modèle n'a peut-être pas été entraîné ou le fichier est manquant."}
88
 
89
  try:
90
+ logger.info(f"Début de la prédiction pour {len(input_data)} enregistrement(s)...")
91
+
92
+ # Travailler sur une copie pour éviter de modifier le DataFrame original
93
+ data_to_process = input_data.copy()
94
 
95
  # --- ÉTAPE 1 : Appliquer les transformations manuelles ---
96
+ # clean_data s'attend à la colonne 'a_quitte_l_entreprise' pour créer config.TARGET_VARIABLE.
97
+ # Pour la prédiction, cette colonne n'existe pas dans les données d'entrée.
98
+ # Nous ajoutons une colonne factice 'a_quitte_l_entreprise' pour que clean_data
99
+ # puisse s'exécuter sans erreur. Cette colonne sera ensuite supprimée.
100
+ if "a_quitte_l_entreprise" not in data_to_process.columns:
101
+ data_to_process["a_quitte_l_entreprise"] = "Non" # Valeur factice, n'influence pas les features
102
+ logger.debug("Colonne 'a_quitte_l_entreprise' factice ajoutée pour le preprocessing.")
103
+
104
+ df_cleaned = clean_data(data_to_process)
105
+ # Utiliser le mapping centralisé depuis config.py
 
106
  df_mapped = map_binary_features(df_cleaned, config.BINARY_FEATURES_MAPPING)
107
  df_featured = create_features(df_mapped)
108
 
109
+ # Préparer X_predict : supprimer la variable cible (même factice) et toute autre
110
+ # colonne qui ne fait pas partie des features attendues par la pipeline.
111
+ # La pipeline a été entraînée sur un X qui n'incluait pas config.TARGET_VARIABLE.
112
+ cols_to_drop_for_predict = [config.TARGET_VARIABLE]
113
+ # Si d'autres colonnes comme 'id_employee' sont présentes dans df_featured mais que la pipeline
114
+ # ne les attend pas comme features (ce qui est le cas si elles ont été droppées avant le fit
115
+ # du preprocessor dans train_model.py), il faudrait les lister ici aussi.
116
+ # Cependant, le ColumnTransformer avec remainder='drop' devrait ignorer les colonnes inconnues.
117
+ # Mais c'est plus propre de présenter à la pipeline exactement ce qu'elle a vu à l'entraînement.
118
+ # X_predict = df_featured.drop(columns=[col for col in cols_to_drop_for_predict if col in df_featured.columns], errors="ignore")
119
+
120
+ # En se basant sur train_model.py, X_predict ne doit pas contenir TARGET_VARIABLE.
121
+ # Les autres colonnes non-features (comme id_employee) ont été retirées AVANT
122
+ # la définition des listes de colonnes pour build_preprocessor, donc le preprocessor ne les attend pas.
123
+ if config.TARGET_VARIABLE in df_featured.columns:
124
+ X_predict = df_featured.drop(config.TARGET_VARIABLE, axis=1)
125
+ else:
126
+ X_predict = df_featured # Si TARGET_VARIABLE n'a pas été créé (ex: clean_data modifié)
127
+
128
+ logger.info(f"Transformations manuelles appliquées. Shape de X_predict: {X_predict.shape}")
129
+ # logger.debug(f"Colonnes de X_predict avant la pipeline : {X_predict.columns.tolist()}")
130
 
 
131
 
132
  # --- ÉTAPE 2 : Utiliser la pipeline complète pour prédire ---
133
+ logger.info("Application de la pipeline Scikit-learn pour la prédiction...")
 
134
  probabilities = pipeline.predict_proba(X_predict)[:, 1]
135
  predictions = pipeline.predict(X_predict)
136
+ logger.info("Calcul des probabilités et des classes terminé.")
137
 
138
  results = []
139
+ for i, index_val in enumerate(X_predict.index): # Renommé index en index_val pour éviter conflit
140
  results.append(
141
  {
142
+ "id_employe": index_val, # Utilise l'index du DataFrame d'entrée
143
  "probabilite_depart": float(probabilities[i]),
144
  "prediction_depart": int(predictions[i]),
145
  }
146
  )
147
+ logger.info("Prédictions formatées avec succès.")
148
  return results
149
 
150
  except Exception as e:
151
+ logger.error(f"Erreur critique lors du processus de prédiction : {e}", exc_info=True)
152
+ return {"error": f"Erreur de prédiction : {str(e)}"}
153
 
154
 
155
+ # Pour tester ce module : poetry run python -m src.modeling.predict
156
  if __name__ == "__main__":
157
+ logger.info("--- Test direct du module predict.py ---")
158
+ # Exemple de données brutes (doit correspondre au schéma EmployeeInput de l'API)
159
+ sample_data_dict = {
160
+ # Assurez-vous que toutes les clés correspondent à EmployeeInput et aux attentes de clean_data
161
+ "age": [45, 30],
162
+ "genre": ["M", "F"], # Utiliser les valeurs textuelles brutes
163
+ "revenu_mensuel": [4850, 6000],
164
+ "statut_marital": ["Célibataire", "Marié(e)"],
165
+ "departement": ["Commercial", "R&D"],
166
+ "poste": ["Cadre Commercial", "Ingenieur"],
167
+ "nombre_experiences_precedentes": [8, 5],
168
+ "annees_dans_l_entreprise": [5, 2],
169
+ "satisfaction_employee_environnement": [4,3],
170
+ "note_evaluation_precedente": [3,4],
171
+ "satisfaction_employee_nature_travail": [3,2],
172
+ "satisfaction_employee_equipe": [3,4],
173
+ "satisfaction_employee_equilibre_pro_perso": [3,2],
174
+ "note_evaluation_actuelle": [3,4],
175
+ "heure_supplementaires": ["Non", "Oui"], # Valeurs textuelles brutes
176
+ "augementation_salaire_precedente": ["15 %", "10 %"], # Format texte brut
177
+ "nombre_participation_pee": [0,1],
178
+ "nb_formations_suivies": [3,1],
179
+ "distance_domicile_travail": [20,5],
180
+ "niveau_education": ["Bac+3", "Bac+5"], # Exemple, adaptez à vos données
181
+ "domaine_etude": ["Infra & Cloud", "Developpement"],
182
+ "frequence_deplacement": ["Occasionnel", "Frequent"], # Exemple, adaptez
183
+ "annees_depuis_la_derniere_promotion": [0,1],
184
+ # Pas besoin de 'a_quitte_l_entreprise' ici, la fonction predict_attrition l'ajoute temporairement
185
+ }
186
+ sample_df = pd.DataFrame(sample_data_dict, index=["EMP_TEST_A", "EMP_TEST_B"])
187
+
 
 
188
  if config.MODEL_PATH.exists():
189
+ logger.info("Modèle trouvé, lancement de la prédiction test...")
190
+ predictions_output = predict_attrition(sample_df)
191
  print("\n--- Résultat de la Prédiction Test ---")
192
+ if isinstance(predictions_output, dict) and "error" in predictions_output:
193
+ print(f"Erreur: {predictions_output['error']}")
194
+ else:
195
+ for res in predictions_output:
196
+ print(res)
197
  else:
198
  print(
199
+ f"\nModèle non trouvé à {config.MODEL_PATH}. Veuillez l'entraîner d'abord avec train_model.py."
200
  )
201
+ logger.info("--- Fin du test direct du module predict.py ---")
src/modeling/train_model.py CHANGED
@@ -1,9 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  import numpy as np
2
- from joblib import dump
3
  import logging
4
 
5
  from sklearn.model_selection import train_test_split
6
- from sklearn.linear_model import LogisticRegression # Ou RandomForestClassifier, etc.
7
  from sklearn.pipeline import Pipeline
8
  from sklearn.metrics import (
9
  classification_report,
@@ -11,140 +28,173 @@ from sklearn.metrics import (
11
  confusion_matrix,
12
  )
13
 
 
14
  from src.data_processing.load_data import get_data
15
- # from src.data_processing.load_data import load_and_merge_csvs
16
  from src.data_processing.preprocess import (
17
- clean_data,
18
- map_binary_features,
19
- create_features,
20
- build_preprocessor,
21
  )
22
- from src import config
23
 
 
24
  logging.basicConfig(
25
- level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
26
  )
27
  logger = logging.getLogger(__name__)
28
 
29
 
30
  def train_and_evaluate_pipeline():
31
  """
32
- Orchestre le chargement, la préparation, l'entraînement,
33
- l'évaluation et la sauvegarde de la pipeline ML.
 
 
 
 
 
 
 
 
 
 
 
34
  """
35
- logger.info(">>> Début du processus d'entraînement et d'évaluation <<<")
36
 
37
  # --- 1. Charger les données ---
38
- # df_raw = load_and_merge_csvs()
39
  df_loaded = get_data(source="postgres")
40
  if df_loaded is None or df_loaded.empty:
41
- logger.error("Arrêt : Impossible de charger les données ou DataFrame vide.")
42
  return
43
 
44
- # --- 3. Appliquer les transformations "sûres" (avant split) --- PAS BESOIN SI DEJA FAIT DANS populate_db
45
- # df_cleaned = clean_data(df_raw)
46
- # df_mapped = map_binary_features(df_cleaned, config.BINARY_FEATURES_MAPPING)
47
- df_featured = create_features(df_loaded) # Assurez-vous que df_featured est bien défini
 
48
 
49
- if df_featured.empty: # Vérifier après create_features (ou sur df_loaded si create_features ne fait rien)
50
- logger.error("Arrêt : DataFrame vide après chargement/création de features.")
51
  return
52
 
53
- df_for_training = df_featured.copy()
54
 
55
- if config.TARGET_VARIABLE not in df_for_training.columns: # CETTE VERIFICATION EST MAINTENANT APRES LE CHECK EMPTY
56
- raise ValueError(f"La colonne cible '{config.TARGET_VARIABLE}' n'est pas présente.")
 
 
 
 
 
57
 
58
  y = df_for_training[config.TARGET_VARIABLE]
59
- X = df_for_training.drop(config.TARGET_VARIABLE, axis=1)
60
-
61
- # Colonnes à exclure de X car ce ne sont pas des features pour le modèle
62
- # mais des identifiants ou des métadonnées de la BDD
63
  cols_to_exclude_from_features = [
64
- config.TARGET_VARIABLE,
65
- 'a_quitte_l_entreprise', # La version texte de la cible, si elle est encore là
66
- 'id_employee',
67
- 'date_creation_enregistrement', # Si elle vient de la BDD
68
- 'date_derniere_modification' # Si elle vient de la BDD
69
- # Ajoutez toute autre colonne qui ne doit pas être une feature
70
  ]
71
-
72
- X = df_for_training.drop(columns=[col for col in cols_to_exclude_from_features if col in df_for_training.columns], errors='ignore')
73
- logger.info(f"Colonnes dans X avant train_test_split (après exclusion): {X.columns.tolist()}")
74
 
75
- # --- 5. Séparer en Train / Test ---
76
- logger.info("Séparation des données en ensembles d'entraînement et de test...")
 
 
 
 
 
 
 
 
 
77
  X_train, X_test, y_train, y_test = train_test_split(
78
  X,
79
  y,
80
- test_size=0.20,
81
- random_state=42,
82
- stratify=y, # Stratify est important pour les classes déséquilibrées
83
  )
84
- logger.info(f"Taille Train: {X_train.shape}, Taille Test: {X_test.shape}")
 
85
 
86
- # --- 6. Identifier les types de colonnes (basé sur X_train !) ---
87
- # C'est important de le faire sur X_train pour éviter toute fuite,
88
- # même si dans ce cas, c'est juste pour obtenir les listes de noms.
89
  potential_numerical_cols = X_train.select_dtypes(
90
  include=[np.number]
91
  ).columns.tolist()
92
- ordinal_to_encode = list(config.ORDINAL_FEATURES_CATEGORIES.keys())
 
 
93
  numerical_to_scale = [
94
  col for col in potential_numerical_cols if col not in ordinal_to_encode
95
  ]
96
- onehot_to_encode = X_train.select_dtypes(
97
- include=["object", "category"]
98
- ).columns.tolist()
 
 
 
 
 
99
 
100
- # --- 7. Construire le préprocesseur (non ajusté) ---
 
101
  preprocessor = build_preprocessor(
102
  numerical_cols=numerical_to_scale,
103
  onehot_cols=onehot_to_encode,
104
  ordinal_cols=ordinal_to_encode,
105
- ordinal_categories=config.ORDINAL_FEATURES_CATEGORIES,
106
  )
107
 
108
- # --- 8. Définir le modèle ---
109
- # Ajoutez ici vos hyperparamètres optimisés si vous en avez.
110
- # L'utilisation de class_weight='balanced' est une première approche
111
- # pour gérer le déséquilibre. Vous pourriez intégrer SMOTE dans la pipeline
112
- # si nécessaire (en utilisant imblearn.pipeline.Pipeline).
113
  classifier = LogisticRegression(
114
  random_state=42, class_weight="balanced", max_iter=1000
115
  )
116
 
117
- # --- 9. Créer la Pipeline Complète ---
 
118
  full_pipeline = Pipeline(
119
  steps=[("preprocessor", preprocessor), ("classifier", classifier)]
120
  )
121
  logger.info("Pipeline complète créée.")
122
 
123
- # --- 10. Entraîner la Pipeline Complète (sur X_train, y_train) ---
124
- # C'est ici que le 'fit' du preprocessor ET du classifier a lieu,
125
- # UNIQUEMENT sur les données d'entraînement.
126
- logger.info("Entraînement de la pipeline...")
127
  full_pipeline.fit(X_train, y_train)
128
- logger.info("Entraînement terminé.")
129
 
130
- # --- 11. Évaluer la Pipeline (sur X_test, y_test) ---
131
- logger.info("\n--- Évaluation sur le jeu de Test ---")
132
  y_pred = full_pipeline.predict(X_test)
 
133
 
 
134
  print("\nMatrice de Confusion :\n", confusion_matrix(y_test, y_pred))
135
  print(
136
  "\nRapport de Classification :\n",
137
  classification_report(y_test, y_pred, target_names=["Non", "Oui"]),
138
  )
139
 
140
- f2_scorer = fbeta_score(y_test, y_pred, beta=2)
141
- print(f"\nF2-Score (privilégie le Rappel) : {f2_scorer:.4f}")
 
 
 
 
 
 
 
 
 
 
142
 
143
- # --- 12. Sauvegarder la Pipeline Ajustée ---
144
- dump(full_pipeline, config.MODEL_PATH)
145
- logger.info(f"Pipeline complète et ajustée sauvegardée dans {config.MODEL_PATH}")
146
- logger.info(">>> Fin du processus <<<")
147
 
148
 
149
  if __name__ == "__main__":
150
- train_and_evaluate_pipeline()
 
1
+ """
2
+ Module d'entraînement et d'évaluation du modèle de prédiction d'attrition.
3
+
4
+ Ce script orchestre le pipeline complet de Machine Learning :
5
+ 1. Chargement des données prétraitées depuis la source configurée (PostgreSQL).
6
+ 2. (Optionnel) Création de features supplémentaires.
7
+ 3. Séparation des données en ensembles d'entraînement et de test.
8
+ 4. Identification des types de colonnes pour le preprocessing.
9
+ 5. Construction d'une pipeline Scikit-learn incluant :
10
+ - Un préprocesseur (ColumnTransformer) pour imputer, scaler (numériques)
11
+ et encoder (catégorielles OneHot et Ordinal).
12
+ - Un classifieur (actuellement LogisticRegression).
13
+ 6. Entraînement de la pipeline complète sur les données d'entraînement.
14
+ 7. Évaluation du modèle sur les données de test (Matrice de confusion, rapport de
15
+ classification, F2-score).
16
+ 8. Sauvegarde de la pipeline entraînée pour une utilisation ultérieure (prédictions).
17
+ """
18
  import numpy as np
19
+ from joblib import dump # Pour sauvegarder la pipeline
20
  import logging
21
 
22
  from sklearn.model_selection import train_test_split
23
+ from sklearn.linear_model import LogisticRegression
24
  from sklearn.pipeline import Pipeline
25
  from sklearn.metrics import (
26
  classification_report,
 
28
  confusion_matrix,
29
  )
30
 
31
+ # Fonctions et configurations importées des autres modules du projet
32
  from src.data_processing.load_data import get_data
 
33
  from src.data_processing.preprocess import (
34
+ create_features, # Utilisé pour la création de features à la volée
35
+ build_preprocessor, # Pour construire le ColumnTransformer
 
 
36
  )
37
+ from src import config # Pour TARGET_VARIABLE, ORDINAL_FEATURES_CATEGORIES, MODEL_PATH
38
 
39
+ # Configuration du logging pour ce module
40
  logging.basicConfig(
41
+ level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
42
  )
43
  logger = logging.getLogger(__name__)
44
 
45
 
46
  def train_and_evaluate_pipeline():
47
  """
48
+ Orchestre le chargement des données, la préparation, l'entraînement d'un modèle
49
+ de classification, son évaluation, et la sauvegarde de la pipeline entraînée.
50
+
51
+ Cette fonction est le point d'entrée principal pour le processus d'entraînement.
52
+ Elle utilise les configurations définies dans `config.py` et les fonctions
53
+ de `load_data.py` et `preprocess.py`.
54
+
55
+ Le modèle actuel est une Régression Logistique, et la métrique principale
56
+ d'évaluation est le F2-score pour la classe positive (départ d'employé).
57
+
58
+ Raises:
59
+ ValueError: Si la colonne cible n'est pas trouvée dans les données
60
+ après les étapes initiales de chargement et de feature engineering.
61
  """
62
+ logger.info(">>> Début du processus d'entraînement et d'évaluation du modèle <<<")
63
 
64
  # --- 1. Charger les données ---
65
+ logger.info("Étape 1: Chargement des données depuis la source configurée (PostgreSQL).")
66
  df_loaded = get_data(source="postgres")
67
  if df_loaded is None or df_loaded.empty:
68
+ logger.error("Arrêt : Impossible de charger les données ou DataFrame vide depuis la source.")
69
  return
70
 
71
+ # --- 2. Création de Features (si applicable) ---
72
+ # Cette étape applique la fonction create_features pour toute ingénierie de features
73
+ # qui n'aurait pas été faite avant le stockage en base, ou qui est spécifique à l'entraînement.
74
+ logger.info("Étape 2: Application de la création de features (si définie).")
75
+ df_featured = create_features(df_loaded)
76
 
77
+ if df_featured.empty:
78
+ logger.error("Arrêt : DataFrame vide après l'étape de création de features.")
79
  return
80
 
81
+ df_for_training = df_featured.copy() # Utiliser une copie pour les modifications suivantes
82
 
83
+ # --- 3. Vérification et Séparation de la Cible (y) et des Features (X) ---
84
+ logger.info("Étape 3: Séparation des features (X) et de la variable cible (y).")
85
+ if config.TARGET_VARIABLE not in df_for_training.columns:
86
+ logger.error(f"La colonne cible '{config.TARGET_VARIABLE}' est introuvable dans le DataFrame.")
87
+ raise ValueError(
88
+ f"La colonne cible '{config.TARGET_VARIABLE}' n'est pas présente dans les données chargées."
89
+ )
90
 
91
  y = df_for_training[config.TARGET_VARIABLE]
92
+
93
+ # Exclure la cible et les colonnes non-feature de X
 
 
94
  cols_to_exclude_from_features = [
95
+ config.TARGET_VARIABLE,
96
+ "a_quitte_l_entreprise", # Version texte originale de la cible, si présente
97
+ "id_employee", # Identifiant, non utilisé comme feature
98
+ "date_creation_enregistrement",
99
+ "date_derniere_modification"
100
+ # Ajoutez d'autres colonnes à exclure si nécessaire
101
  ]
 
 
 
102
 
103
+ X = df_for_training.drop(
104
+ columns=[col for col in cols_to_exclude_from_features if col in df_for_training.columns],
105
+ errors="ignore" # ignore si une colonne de la liste n'est pas trouvée dans df_for_training
106
+ )
107
+ logger.info(
108
+ f"Nombre de features initiales dans X avant train_test_split : {len(X.columns)}"
109
+ )
110
+ # logger.debug(f"Colonnes dans X : {X.columns.tolist()}") # Décommentez pour débogage détaillé
111
+
112
+ # --- 4. Séparation en Ensembles d'Entraînement et de Test ---
113
+ logger.info("Étape 4: Séparation des données en ensembles d'entraînement et de test (80%/20%).")
114
  X_train, X_test, y_train, y_test = train_test_split(
115
  X,
116
  y,
117
+ test_size=0.20, # 20% des données pour le test
118
+ random_state=42, # Pour la reproductibilité
119
+ stratify=y, # Maintient la proportion des classes dans les deux ensembles
120
  )
121
+ logger.info(f"Dimensions de l'ensemble d'entraînement - X_train: {X_train.shape}, y_train: {y_train.shape}")
122
+ logger.info(f"Dimensions de l'ensemble de test - X_test: {X_test.shape}, y_test: {y_test.shape}")
123
 
124
+ # --- 5. Identification des Types de Colonnes pour le Préprocessing (basé sur X_train) ---
125
+ logger.info("Étape 5: Identification des types de colonnes pour le préprocessing (sur X_train).")
 
126
  potential_numerical_cols = X_train.select_dtypes(
127
  include=[np.number]
128
  ).columns.tolist()
129
+ ordinal_to_encode = list(config.ORDINAL_FEATURES_CATEGORIES.keys()) # Colonnes définies comme ordinales
130
+
131
+ # Les colonnes numériques à scaler sont celles qui sont numériques et non ordinales
132
  numerical_to_scale = [
133
  col for col in potential_numerical_cols if col not in ordinal_to_encode
134
  ]
135
+ # Les colonnes pour OneHotEncoding sont celles de type object/category qui ne sont pas ordinales
136
+ onehot_to_encode = [
137
+ col for col in X_train.select_dtypes(include=["object", "category"]).columns.tolist()
138
+ if col not in ordinal_to_encode
139
+ ]
140
+ logger.info(f"Colonnes numériques à scaler : {numerical_to_scale}")
141
+ logger.info(f"Colonnes pour OneHotEncoding : {onehot_to_encode}")
142
+ logger.info(f"Colonnes pour OrdinalEncoding : {ordinal_to_encode}")
143
 
144
+ # --- 6. Construction du Préprocesseur ---
145
+ logger.info("Étape 6: Construction du préprocesseur Scikit-learn.")
146
  preprocessor = build_preprocessor(
147
  numerical_cols=numerical_to_scale,
148
  onehot_cols=onehot_to_encode,
149
  ordinal_cols=ordinal_to_encode,
150
+ ordinal_categories_map=config.ORDINAL_FEATURES_CATEGORIES, # Le nom du paramètre a été harmonisé
151
  )
152
 
153
+ # --- 7. Définition du Modèle de Classification ---
154
+ logger.info("Étape 7: Définition du modèle (LogisticRegression).")
 
 
 
155
  classifier = LogisticRegression(
156
  random_state=42, class_weight="balanced", max_iter=1000
157
  )
158
 
159
+ # --- 8. Création de la Pipeline Complète (Préprocesseur + Classifieur) ---
160
+ logger.info("Étape 8: Création de la pipeline Scikit-learn complète.")
161
  full_pipeline = Pipeline(
162
  steps=[("preprocessor", preprocessor), ("classifier", classifier)]
163
  )
164
  logger.info("Pipeline complète créée.")
165
 
166
+ # --- 9. Entraînement de la Pipeline Complète ---
167
+ logger.info("Étape 9: Entraînement de la pipeline sur l'ensemble d'entraînement...")
 
 
168
  full_pipeline.fit(X_train, y_train)
169
+ logger.info("Entraînement de la pipeline terminé.")
170
 
171
+ # --- 10. Évaluation de la Pipeline sur l'Ensemble de Test ---
172
+ logger.info("Étape 10: Évaluation de la pipeline sur l'ensemble de test...")
173
  y_pred = full_pipeline.predict(X_test)
174
+ # y_pred_proba = full_pipeline.predict_proba(X_test)[:, 1] # Gardé pour référence, mais non utilisé actuellement
175
 
176
+ logger.info("\n--- Résultats de l'Évaluation sur le Jeu de Test ---")
177
  print("\nMatrice de Confusion :\n", confusion_matrix(y_test, y_pred))
178
  print(
179
  "\nRapport de Classification :\n",
180
  classification_report(y_test, y_pred, target_names=["Non", "Oui"]),
181
  )
182
 
183
+ f2_scorer = fbeta_score(y_test, y_pred, beta=2, zero_division=0) # Ajout de zero_division pour éviter warning/erreur
184
+ print(f"\nF2-Score (privilégie le Rappel pour la classe 'Oui') : {f2_scorer:.4f}")
185
+
186
+ # --- 11. Sauvegarde de la Pipeline Ajustée ---
187
+ logger.info("Étape 11: Sauvegarde de la pipeline entraînée...")
188
+ try:
189
+ config.MODELS_DIR.mkdir(parents=True, exist_ok=True) # S'assurer que le dossier existe
190
+ dump(full_pipeline, config.MODEL_PATH)
191
+ logger.info(f"Pipeline complète et ajustée sauvegardée dans : {config.MODEL_PATH}")
192
+ except Exception as e_save:
193
+ logger.error(f"Erreur lors de la sauvegarde de la pipeline : {e_save}", exc_info=True)
194
+
195
 
196
+ logger.info(">>> Fin du processus d'entraînement et d'évaluation <<<")
 
 
 
197
 
198
 
199
  if __name__ == "__main__":
200
+ train_and_evaluate_pipeline()
tests/unit/test_load_data.py CHANGED
@@ -1,75 +1,88 @@
1
  import pandas as pd
2
- import pytest
3
  from unittest.mock import patch, MagicMock, call
4
 
5
  from src.data_processing.load_data import (
6
  load_data_from_csv,
7
  load_data_from_postgres,
8
  get_data,
9
- load_and_merge_csvs
10
  )
11
- from src import config as app_config # Pour config.PROCESSED_DATA_PATH
 
12
  # Supposons que Employee est défini dans models.py pour le test de load_data_from_postgres
13
  # from src.database.models import Employee # Nécessaire si on mock db.query(Employee) spécifiquement
14
 
15
- @patch('src.data_processing.load_data.pd.read_csv')
16
- @patch('src.data_processing.load_data.logger')
 
17
  def test_load_data_from_csv_success(mock_logger, mock_read_csv):
18
  """Teste le chargement réussi depuis un CSV."""
19
- sample_df = pd.DataFrame({'col1': [1], 'col2': ['a']})
20
  mock_read_csv.return_value = sample_df
21
-
22
- df = load_data_from_csv("dummy_path.csv")
23
-
24
- mock_read_csv.assert_called_once_with("dummy_path.csv")
25
- mock_logger.info.assert_any_call("Chargement des données CSV depuis dummy_path.csv...")
26
- mock_logger.info.assert_any_call("Données CSV chargées avec succès.")
 
27
  pd.testing.assert_frame_equal(df, sample_df)
28
 
29
- @patch('src.data_processing.load_data.pd.read_csv')
30
- @patch('src.data_processing.load_data.logger')
 
31
  def test_load_data_from_csv_file_not_found(mock_logger, mock_read_csv):
32
  """Teste la gestion de FileNotFoundError pour load_data_from_csv."""
33
  mock_read_csv.side_effect = FileNotFoundError("File not found")
34
-
35
  df = load_data_from_csv("non_existent.csv")
36
-
37
  assert df is None
38
- mock_logger.error.assert_called_once_with("Fichier CSV non trouvé : non_existent.csv")
 
 
39
 
40
- @patch('src.data_processing.load_data.SessionLocal') # Mocker la classe SessionLocal
41
- @patch('src.data_processing.load_data.pd.read_sql_query')
42
- @patch('src.data_processing.load_data.logger')
 
43
  def test_load_data_from_postgres_success(mock_logger, mock_read_sql, mock_SessionLocal):
44
  """Teste le chargement réussi depuis PostgreSQL."""
45
- sample_df = pd.DataFrame({'id_employee': ['E_1'], 'age': [30]})
46
  mock_read_sql.return_value = sample_df
47
-
48
  # Configurer le mock de la session et de la query
49
  mock_db_session = MagicMock()
50
- mock_SessionLocal.return_value = mock_db_session # Quand SessionLocal() est appelé, il retourne notre mock
51
-
 
 
52
  # Si votre query est db.query(Employee), vous pouvez mocker Employee
53
  # from src.database.models import Employee (déjà importé si décommenté plus haut)
54
  # mock_query_obj = MagicMock()
55
  # mock_db_session.query.return_value = mock_query_obj
56
 
57
  df = load_data_from_postgres()
58
-
59
- mock_SessionLocal.assert_called_once() # Vérifie que SessionLocal() a été instancié
60
  # mock_db_session.query.assert_called_once_with(Employee) # Si vous mockez Employee
61
  # mock_read_sql.assert_called_once_with(mock_query_obj.statement, mock_db_session.bind)
62
  # Version plus simple si on ne mocke pas la query en détail :
63
  assert mock_read_sql.call_count == 1
64
 
65
- mock_logger.info.assert_any_call("Chargement des données depuis la table 'employees' de PostgreSQL...")
66
- mock_logger.info.assert_any_call(f"{len(sample_df)} lignes chargées depuis la table 'employees'.")
 
 
 
 
67
  pd.testing.assert_frame_equal(df, sample_df)
68
- mock_db_session.close.assert_called_once() # Vérifie que la session est fermée
 
69
 
70
- @patch('src.data_processing.load_data.SessionLocal')
71
- @patch('src.data_processing.load_data.pd.read_sql_query')
72
- @patch('src.data_processing.load_data.logger')
73
  def test_load_data_from_postgres_empty(mock_logger, mock_read_sql, mock_SessionLocal):
74
  """Teste le chargement depuis PostgreSQL quand la table est vide."""
75
  empty_df = pd.DataFrame()
@@ -78,31 +91,39 @@ def test_load_data_from_postgres_empty(mock_logger, mock_read_sql, mock_SessionL
78
  mock_SessionLocal.return_value = mock_db_session
79
 
80
  df = load_data_from_postgres()
81
-
82
  pd.testing.assert_frame_equal(df, empty_df)
83
- mock_logger.warning.assert_called_once_with("Aucune donnée trouvée dans la table 'employees'. Le DataFrame est vide.")
 
 
84
  mock_db_session.close.assert_called_once()
85
 
86
- @patch('src.data_processing.load_data.SessionLocal')
87
- @patch('src.data_processing.load_data.pd.read_sql_query')
88
- @patch('src.data_processing.load_data.logger')
89
- def test_load_data_from_postgres_exception(mock_logger, mock_read_sql, mock_SessionLocal):
 
 
 
90
  """Teste la gestion d'exception lors du chargement depuis PostgreSQL."""
91
  mock_read_sql.side_effect = Exception("DB Connection Error")
92
  mock_db_session = MagicMock()
93
  mock_SessionLocal.return_value = mock_db_session
94
 
95
  df = load_data_from_postgres()
96
-
97
- assert df.empty # Doit retourner un DataFrame vide
98
  mock_logger.error.assert_called_once_with(
99
  "Erreur lors du chargement des données depuis PostgreSQL : DB Connection Error",
100
- exc_info=True
101
  )
102
  mock_db_session.close.assert_called_once()
103
 
104
- @patch('src.data_processing.load_data.load_data_from_postgres')
105
- @patch('src.data_processing.load_data.load_and_merge_csvs') # Si get_data(source='csv') l'appelle
 
 
 
106
  # Ou @patch('src.data_processing.load_data.load_data_from_csv') si c'est ce qu'elle appelle
107
  def test_get_data_source_postgres(mock_load_csv_or_merge, mock_load_postgres):
108
  """Teste get_data avec source='postgres'."""
@@ -112,8 +133,9 @@ def test_get_data_source_postgres(mock_load_csv_or_merge, mock_load_postgres):
112
  mock_load_csv_or_merge.assert_not_called()
113
  assert result == "data_from_db"
114
 
115
- @patch('src.data_processing.load_data.load_data_from_postgres')
116
- @patch('src.data_processing.load_data.load_and_merge_csvs') # Adaptez ce mock
 
117
  def test_get_data_source_csv(mock_load_csv_or_merge, mock_load_postgres):
118
  """Teste get_data avec source='csv'."""
119
  mock_load_csv_or_merge.return_value = "data_from_csv"
@@ -122,34 +144,43 @@ def test_get_data_source_csv(mock_load_csv_or_merge, mock_load_postgres):
122
  mock_load_postgres.assert_not_called()
123
  assert result == "data_from_csv"
124
 
125
- @patch('src.data_processing.load_data.load_data_from_postgres')
126
- @patch('src.data_processing.load_data.logger')
 
127
  def test_get_data_invalid_source(mock_logger, mock_load_postgres):
128
  """Teste get_data avec une source invalide (doit utiliser postgres par défaut)."""
129
  mock_load_postgres.return_value = "data_from_db_default"
130
  result = get_data(source="invalid_source")
131
- mock_load_postgres.assert_called_once() # Appelée car c'est le défaut
132
- mock_logger.error.assert_any_call("Source de données non reconnue : invalid_source. Utilisation de PostgreSQL par défaut.")
 
 
133
  assert result == "data_from_db_default"
134
-
135
- @patch('src.data_processing.load_data.logger') # Mocker le logger en premier (ordre des décorateurs inversé)
136
- @patch('src.data_processing.load_data.pd.read_csv')
 
 
 
137
  def test_load_and_merge_csvs_success(mock_read_csv, mock_logger):
138
  """Teste le chargement et la fusion réussis des 3 CSV."""
139
  # 1. Préparer les DataFrames de simulation retournés par read_csv
140
- df_sirh_sample = pd.DataFrame({
141
- 'id_employee': ['1', '2', '3'],
142
- 'sirh_feature': ['A1', 'A2', 'A3']
143
- })
144
  # df_eval avec des eval_number qui seront transformés en id_employee '1', '2', 'SPECIAL'
145
- df_eval_sample = pd.DataFrame({
146
- 'eval_number': ['E_1', 'E_2', 'E_SPECIAL'],
147
- 'eval_feature': ['EvalX', 'EvalY', 'EvalZ']
148
- })
149
- df_sondage_sample = pd.DataFrame({
150
- 'id_employee': ['1', '3'], # ID '2' manquant, 'SPECIAL' non présent
151
- 'sondage_feature': ['SondA', 'SondC']
152
- })
 
 
 
 
153
 
154
  # Configurer le side_effect de mock_read_csv pour retourner ces DataFrames dans l'ordre
155
  mock_read_csv.side_effect = [df_sirh_sample, df_eval_sample, df_sondage_sample]
@@ -162,107 +193,132 @@ def test_load_and_merge_csvs_success(mock_read_csv, mock_logger):
162
  expected_calls = [
163
  call(app_config.RAW_SIRH_PATH),
164
  call(app_config.RAW_EVAL_PATH),
165
- call(app_config.RAW_SONDAGE_PATH)
166
  ]
167
  mock_read_csv.assert_has_calls(expected_calls, any_order=False)
168
  assert mock_read_csv.call_count == 3
169
 
170
  # Vérifier les logs importants
171
- mock_logger.info.assert_any_call("Chargement des fichiers CSV bruts...")
172
- mock_logger.info.assert_any_call("Préparation de la clé de jointure dans df_eval...")
 
173
  # mock_logger.info.assert_any_call("'id_employee' créé dans df_eval.") # log absent du fichier load_data.py
174
- mock_logger.info.assert_any_call("Fusion des DataFrames...")
175
  # mock_logger.info.assert_any_call(f"Données fusionnées : {merged_df.shape}") # Le shape peut varier
176
 
177
  # Vérifier le DataFrame fusionné
178
  assert merged_df is not None
179
- assert list(merged_df.columns) == ['id_employee', 'sirh_feature', 'eval_number', 'eval_feature', 'sondage_feature']
180
- assert len(merged_df) == 3 # Car left merge depuis df_sirh_sample
 
 
 
 
 
 
181
 
182
  # Vérifier le contenu (basé sur df_sirh en tant que table de gauche pour les merges)
183
  # Ligne pour id_employee '1'
184
- row_1 = merged_df[merged_df['id_employee'] == '1'].iloc[0]
185
- assert row_1['sirh_feature'] == 'A1'
186
- assert row_1['eval_feature'] == 'EvalX'
187
- assert row_1['sondage_feature'] == 'SondA'
188
 
189
  # Ligne pour id_employee '2'
190
- row_2 = merged_df[merged_df['id_employee'] == '2'].iloc[0]
191
- assert row_2['sirh_feature'] == 'A2'
192
- assert row_2['eval_feature'] == 'EvalY'
193
- assert pd.isna(row_2['sondage_feature']) # Pas de '2' dans df_sondage_sample
194
 
195
  # Ligne pour id_employee '3'
196
- row_3 = merged_df[merged_df['id_employee'] == '3'].iloc[0]
197
- assert row_3['sirh_feature'] == 'A3'
198
- assert pd.isna(row_3['eval_feature']) # 'E_3' n'est pas dans df_eval_sample, df_eval a 'E_SPECIAL' qui donne 'SPECIAL'
199
- assert row_3['sondage_feature'] == 'SondC'
200
-
 
 
201
  # Vérifier le type de la colonne 'id_employee'
202
- assert merged_df['id_employee'].dtype == 'object'
203
 
204
 
205
- @patch('src.data_processing.load_data.logger')
206
- @patch('src.data_processing.load_data.pd.read_csv')
207
  def test_load_and_merge_csvs_sirh_file_not_found(mock_read_csv, mock_logger):
208
  """Teste FileNotFoundError pour le fichier SIRH."""
209
  mock_read_csv.side_effect = FileNotFoundError("SIRH file missing")
210
-
211
  result_df = load_and_merge_csvs()
212
-
213
  assert result_df is None
214
- mock_logger.error.assert_any_call("Erreur de chargement CSV : Fichier non trouvé - SIRH file missing")
 
 
215
 
216
- @patch('src.data_processing.load_data.logger')
217
- @patch('src.data_processing.load_data.pd.read_csv')
 
218
  def test_load_and_merge_csvs_eval_file_not_found(mock_read_csv, mock_logger):
219
  """Teste FileNotFoundError pour le fichier EVAL."""
220
- df_sirh_sample = pd.DataFrame({'id_employee': ['1']})
221
  mock_read_csv.side_effect = [df_sirh_sample, FileNotFoundError("EVAL file missing")]
222
-
223
  result_df = load_and_merge_csvs()
224
-
225
  assert result_df is None
226
- mock_logger.error.assert_any_call("Erreur de chargement CSV : Fichier non trouvé - EVAL file missing")
 
 
227
 
228
 
229
- @patch('src.data_processing.load_data.logger')
230
- @patch('src.data_processing.load_data.pd.read_csv')
231
  def test_load_and_merge_csvs_missing_eval_number_col(mock_read_csv, mock_logger):
232
  """Teste le cas où 'eval_number' manque dans df_eval."""
233
- df_sirh_sample = pd.DataFrame({'id_employee': ['1']})
234
- df_eval_no_eval_number = pd.DataFrame({'autre_col': ['X']}) # Pas de eval_number
235
- df_sondage_sample = pd.DataFrame({'id_employee': ['1']})
236
- mock_read_csv.side_effect = [df_sirh_sample, df_eval_no_eval_number, df_sondage_sample]
237
-
 
 
 
 
238
  result_df = load_and_merge_csvs()
239
-
240
  assert result_df is None
241
- mock_logger.error.assert_any_call("La colonne 'eval_number' est introuvable dans df_eval.")
242
 
243
 
244
- @patch('src.data_processing.load_data.logger')
245
- @patch('src.data_processing.load_data.pd.read_csv')
246
  def test_load_and_merge_csvs_missing_id_employee_in_sirh(mock_read_csv, mock_logger):
247
  """Teste le cas où 'id_employee' manque dans df_sirh."""
248
- df_sirh_no_id = pd.DataFrame({'autre_col_sirh': ['Y']}) # Pas de id_employee
249
- df_eval_sample = pd.DataFrame({'eval_number': ['E_1']})
250
- df_sondage_sample = pd.DataFrame({'id_employee': ['1']})
251
  mock_read_csv.side_effect = [df_sirh_no_id, df_eval_sample, df_sondage_sample]
252
 
253
  result_df = load_and_merge_csvs()
254
 
255
  assert result_df is None
256
- mock_logger.error.assert_any_call("La colonne 'id_employee' est introuvable dans df_sirh.")
 
 
257
 
258
 
259
- @patch('src.data_processing.load_data.logger')
260
- @patch('src.data_processing.load_data.pd.read_csv')
261
  def test_load_and_merge_csvs_eval_number_conversion_warning(mock_read_csv, mock_logger):
262
  """Teste le warning pour les eval_number non convertibles."""
263
- df_sirh_sample = pd.DataFrame({'id_employee': ['1', 'bad_id_format']})
264
- df_eval_sample = pd.DataFrame({'eval_number': ['E_1', 'E_OK', 'WRONG_FORMAT', 'E_WRONG_TOO'], 'eval_feature': ['x1', 'x2', 'x3', 'x4']})
265
- df_sondage_sample = pd.DataFrame({'id_employee': ['1', 'OK']})
 
 
 
 
 
266
  mock_read_csv.side_effect = [df_sirh_sample, df_eval_sample, df_sondage_sample]
267
 
268
  merged_df = load_and_merge_csvs()
@@ -272,10 +328,17 @@ def test_load_and_merge_csvs_eval_number_conversion_warning(mock_read_csv, mock_
272
  # La logique de split('_').str[1] va lever une erreur ou retourner NaN si pas de '_'.
273
  # pd.to_numeric(..., errors='coerce') transformera ces erreurs en NaN.
274
  # Ici, 'WRONG_FORMAT' et 'E_WRONG_TOO' devraient produire des NaN pour 'id_employee' dans df_eval.
275
- mock_logger.warning.assert_any_call("3 'eval_number' n'ont pas pu être convertis en 'id_employee' valides.")
 
276
  # mock_logger.warning.assert_any_call("3 'eval_number' n'ont pas pu être convertis en 'id_employee' valides.")
277
  assert merged_df is not None
278
  # La ligne avec id_employee '1' devrait avoir eval_feature 'x1'
279
- assert merged_df.loc[merged_df['id_employee'] == '1', 'eval_feature'].iloc[0] == 'x1'
 
 
280
  # La ligne avec id_employee 'bad_id_format' ne trouvera pas de correspondance dans df_eval et aura NaN pour eval_feature
281
- assert pd.isna(merged_df.loc[merged_df['id_employee'] == 'bad_id_format', 'eval_feature'].iloc[0])
 
 
 
 
 
1
  import pandas as pd
 
2
  from unittest.mock import patch, MagicMock, call
3
 
4
  from src.data_processing.load_data import (
5
  load_data_from_csv,
6
  load_data_from_postgres,
7
  get_data,
8
+ load_and_merge_csvs,
9
  )
10
+ from src import config as app_config # Pour config.PROCESSED_DATA_PATH
11
+
12
  # Supposons que Employee est défini dans models.py pour le test de load_data_from_postgres
13
  # from src.database.models import Employee # Nécessaire si on mock db.query(Employee) spécifiquement
14
 
15
+
16
+ @patch("src.data_processing.load_data.pd.read_csv")
17
+ @patch("src.data_processing.load_data.logger")
18
  def test_load_data_from_csv_success(mock_logger, mock_read_csv):
19
  """Teste le chargement réussi depuis un CSV."""
20
+ sample_df = pd.DataFrame({"col1": [1], "col2": ["a"]})
21
  mock_read_csv.return_value = sample_df
22
+
23
+ dummy_path = "dummy_path.csv"
24
+ df = load_data_from_csv(dummy_path) # path est dummy_path
25
+
26
+ mock_read_csv.assert_called_once_with(dummy_path)
27
+ mock_logger.info.assert_any_call(f"Chargement des données CSV depuis {dummy_path}...")
28
+ mock_logger.info.assert_any_call(f"Données CSV chargées avec succès depuis {dummy_path}: {len(sample_df)} lignes.")
29
  pd.testing.assert_frame_equal(df, sample_df)
30
 
31
+
32
+ @patch("src.data_processing.load_data.pd.read_csv")
33
+ @patch("src.data_processing.load_data.logger")
34
  def test_load_data_from_csv_file_not_found(mock_logger, mock_read_csv):
35
  """Teste la gestion de FileNotFoundError pour load_data_from_csv."""
36
  mock_read_csv.side_effect = FileNotFoundError("File not found")
37
+
38
  df = load_data_from_csv("non_existent.csv")
39
+
40
  assert df is None
41
+ mock_logger.error.assert_called_once_with(
42
+ "Fichier CSV non trouvé : non_existent.csv"
43
+ )
44
 
45
+
46
+ @patch("src.data_processing.load_data.SessionLocal") # Mocker la classe SessionLocal
47
+ @patch("src.data_processing.load_data.pd.read_sql_query")
48
+ @patch("src.data_processing.load_data.logger")
49
  def test_load_data_from_postgres_success(mock_logger, mock_read_sql, mock_SessionLocal):
50
  """Teste le chargement réussi depuis PostgreSQL."""
51
+ sample_df = pd.DataFrame({"id_employee": ["E_1"], "age": [30]})
52
  mock_read_sql.return_value = sample_df
53
+
54
  # Configurer le mock de la session et de la query
55
  mock_db_session = MagicMock()
56
+ mock_SessionLocal.return_value = (
57
+ mock_db_session # Quand SessionLocal() est appelé, il retourne notre mock
58
+ )
59
+
60
  # Si votre query est db.query(Employee), vous pouvez mocker Employee
61
  # from src.database.models import Employee (déjà importé si décommenté plus haut)
62
  # mock_query_obj = MagicMock()
63
  # mock_db_session.query.return_value = mock_query_obj
64
 
65
  df = load_data_from_postgres()
66
+
67
+ mock_SessionLocal.assert_called_once() # Vérifie que SessionLocal() a été instancié
68
  # mock_db_session.query.assert_called_once_with(Employee) # Si vous mockez Employee
69
  # mock_read_sql.assert_called_once_with(mock_query_obj.statement, mock_db_session.bind)
70
  # Version plus simple si on ne mocke pas la query en détail :
71
  assert mock_read_sql.call_count == 1
72
 
73
+ mock_logger.info.assert_any_call(
74
+ "Chargement des données depuis la table 'employees' de PostgreSQL..."
75
+ )
76
+ mock_logger.info.assert_any_call(
77
+ f"{len(sample_df)} lignes chargées depuis la table 'employees'."
78
+ )
79
  pd.testing.assert_frame_equal(df, sample_df)
80
+ mock_db_session.close.assert_called_once() # Vérifie que la session est fermée
81
+
82
 
83
+ @patch("src.data_processing.load_data.SessionLocal")
84
+ @patch("src.data_processing.load_data.pd.read_sql_query")
85
+ @patch("src.data_processing.load_data.logger")
86
  def test_load_data_from_postgres_empty(mock_logger, mock_read_sql, mock_SessionLocal):
87
  """Teste le chargement depuis PostgreSQL quand la table est vide."""
88
  empty_df = pd.DataFrame()
 
91
  mock_SessionLocal.return_value = mock_db_session
92
 
93
  df = load_data_from_postgres()
94
+
95
  pd.testing.assert_frame_equal(df, empty_df)
96
+ mock_logger.warning.assert_called_once_with(
97
+ "Aucune donnée trouvée dans la table 'employees'. Le DataFrame est vide."
98
+ )
99
  mock_db_session.close.assert_called_once()
100
 
101
+
102
+ @patch("src.data_processing.load_data.SessionLocal")
103
+ @patch("src.data_processing.load_data.pd.read_sql_query")
104
+ @patch("src.data_processing.load_data.logger")
105
+ def test_load_data_from_postgres_exception(
106
+ mock_logger, mock_read_sql, mock_SessionLocal
107
+ ):
108
  """Teste la gestion d'exception lors du chargement depuis PostgreSQL."""
109
  mock_read_sql.side_effect = Exception("DB Connection Error")
110
  mock_db_session = MagicMock()
111
  mock_SessionLocal.return_value = mock_db_session
112
 
113
  df = load_data_from_postgres()
114
+
115
+ assert df.empty # Doit retourner un DataFrame vide
116
  mock_logger.error.assert_called_once_with(
117
  "Erreur lors du chargement des données depuis PostgreSQL : DB Connection Error",
118
+ exc_info=True,
119
  )
120
  mock_db_session.close.assert_called_once()
121
 
122
+
123
+ @patch("src.data_processing.load_data.load_data_from_postgres")
124
+ @patch(
125
+ "src.data_processing.load_data.load_and_merge_csvs"
126
+ ) # Si get_data(source='csv') l'appelle
127
  # Ou @patch('src.data_processing.load_data.load_data_from_csv') si c'est ce qu'elle appelle
128
  def test_get_data_source_postgres(mock_load_csv_or_merge, mock_load_postgres):
129
  """Teste get_data avec source='postgres'."""
 
133
  mock_load_csv_or_merge.assert_not_called()
134
  assert result == "data_from_db"
135
 
136
+
137
+ @patch("src.data_processing.load_data.load_data_from_postgres")
138
+ @patch("src.data_processing.load_data.load_and_merge_csvs") # Adaptez ce mock
139
  def test_get_data_source_csv(mock_load_csv_or_merge, mock_load_postgres):
140
  """Teste get_data avec source='csv'."""
141
  mock_load_csv_or_merge.return_value = "data_from_csv"
 
144
  mock_load_postgres.assert_not_called()
145
  assert result == "data_from_csv"
146
 
147
+
148
+ @patch("src.data_processing.load_data.load_data_from_postgres")
149
+ @patch("src.data_processing.load_data.logger")
150
  def test_get_data_invalid_source(mock_logger, mock_load_postgres):
151
  """Teste get_data avec une source invalide (doit utiliser postgres par défaut)."""
152
  mock_load_postgres.return_value = "data_from_db_default"
153
  result = get_data(source="invalid_source")
154
+ mock_load_postgres.assert_called_once() # Appelée car c'est le défaut
155
+ source_name = "invalid_source"
156
+ expected_log_msg = f"Source de données non reconnue : {source_name}. Tentative avec PostgreSQL par défaut."
157
+ mock_logger.error.assert_any_call(expected_log_msg)
158
  assert result == "data_from_db_default"
159
+
160
+
161
+ @patch(
162
+ "src.data_processing.load_data.logger"
163
+ ) # Mocker le logger en premier (ordre des décorateurs inversé)
164
+ @patch("src.data_processing.load_data.pd.read_csv")
165
  def test_load_and_merge_csvs_success(mock_read_csv, mock_logger):
166
  """Teste le chargement et la fusion réussis des 3 CSV."""
167
  # 1. Préparer les DataFrames de simulation retournés par read_csv
168
+ df_sirh_sample = pd.DataFrame(
169
+ {"id_employee": ["1", "2", "3"], "sirh_feature": ["A1", "A2", "A3"]}
170
+ )
 
171
  # df_eval avec des eval_number qui seront transformés en id_employee '1', '2', 'SPECIAL'
172
+ df_eval_sample = pd.DataFrame(
173
+ {
174
+ "eval_number": ["E_1", "E_2", "E_SPECIAL"],
175
+ "eval_feature": ["EvalX", "EvalY", "EvalZ"],
176
+ }
177
+ )
178
+ df_sondage_sample = pd.DataFrame(
179
+ {
180
+ "id_employee": ["1", "3"], # ID '2' manquant, 'SPECIAL' non présent
181
+ "sondage_feature": ["SondA", "SondC"],
182
+ }
183
+ )
184
 
185
  # Configurer le side_effect de mock_read_csv pour retourner ces DataFrames dans l'ordre
186
  mock_read_csv.side_effect = [df_sirh_sample, df_eval_sample, df_sondage_sample]
 
193
  expected_calls = [
194
  call(app_config.RAW_SIRH_PATH),
195
  call(app_config.RAW_EVAL_PATH),
196
+ call(app_config.RAW_SONDAGE_PATH),
197
  ]
198
  mock_read_csv.assert_has_calls(expected_calls, any_order=False)
199
  assert mock_read_csv.call_count == 3
200
 
201
  # Vérifier les logs importants
202
+ mock_logger.info.assert_any_call("Chargement des fichiers CSV bruts pour fusion...")
203
+ mock_logger.info.assert_any_call("Préparation de la clé de jointure 'id_employee' dans df_eval à partir de 'eval_number'...")
204
+ mock_logger.info.assert_any_call("Préparation de la clé de jointure 'id_employee' dans df_sondage à partir de 'code_sondage'...")
205
  # mock_logger.info.assert_any_call("'id_employee' créé dans df_eval.") # log absent du fichier load_data.py
206
+ mock_logger.info.assert_any_call("Fusion des DataFrames (df_sirh <- df_eval <- df_sondage)...")
207
  # mock_logger.info.assert_any_call(f"Données fusionnées : {merged_df.shape}") # Le shape peut varier
208
 
209
  # Vérifier le DataFrame fusionné
210
  assert merged_df is not None
211
+ assert list(merged_df.columns) == [
212
+ "id_employee",
213
+ "sirh_feature",
214
+ "eval_number",
215
+ "eval_feature",
216
+ "sondage_feature",
217
+ ]
218
+ assert len(merged_df) == 3 # Car left merge depuis df_sirh_sample
219
 
220
  # Vérifier le contenu (basé sur df_sirh en tant que table de gauche pour les merges)
221
  # Ligne pour id_employee '1'
222
+ row_1 = merged_df[merged_df["id_employee"] == "1"].iloc[0]
223
+ assert row_1["sirh_feature"] == "A1"
224
+ assert row_1["eval_feature"] == "EvalX"
225
+ assert row_1["sondage_feature"] == "SondA"
226
 
227
  # Ligne pour id_employee '2'
228
+ row_2 = merged_df[merged_df["id_employee"] == "2"].iloc[0]
229
+ assert row_2["sirh_feature"] == "A2"
230
+ assert row_2["eval_feature"] == "EvalY"
231
+ assert pd.isna(row_2["sondage_feature"]) # Pas de '2' dans df_sondage_sample
232
 
233
  # Ligne pour id_employee '3'
234
+ row_3 = merged_df[merged_df["id_employee"] == "3"].iloc[0]
235
+ assert row_3["sirh_feature"] == "A3"
236
+ assert pd.isna(
237
+ row_3["eval_feature"]
238
+ ) # 'E_3' n'est pas dans df_eval_sample, df_eval a 'E_SPECIAL' qui donne 'SPECIAL'
239
+ assert row_3["sondage_feature"] == "SondC"
240
+
241
  # Vérifier le type de la colonne 'id_employee'
242
+ assert merged_df["id_employee"].dtype == "object"
243
 
244
 
245
+ @patch("src.data_processing.load_data.logger")
246
+ @patch("src.data_processing.load_data.pd.read_csv")
247
  def test_load_and_merge_csvs_sirh_file_not_found(mock_read_csv, mock_logger):
248
  """Teste FileNotFoundError pour le fichier SIRH."""
249
  mock_read_csv.side_effect = FileNotFoundError("SIRH file missing")
250
+
251
  result_df = load_and_merge_csvs()
252
+
253
  assert result_df is None
254
+ mock_logger.error.assert_any_call(
255
+ "Erreur de chargement CSV : Fichier non trouvé - SIRH file missing"
256
+ )
257
 
258
+
259
+ @patch("src.data_processing.load_data.logger")
260
+ @patch("src.data_processing.load_data.pd.read_csv")
261
  def test_load_and_merge_csvs_eval_file_not_found(mock_read_csv, mock_logger):
262
  """Teste FileNotFoundError pour le fichier EVAL."""
263
+ df_sirh_sample = pd.DataFrame({"id_employee": ["1"]})
264
  mock_read_csv.side_effect = [df_sirh_sample, FileNotFoundError("EVAL file missing")]
265
+
266
  result_df = load_and_merge_csvs()
267
+
268
  assert result_df is None
269
+ mock_logger.error.assert_any_call(
270
+ "Erreur de chargement CSV : Fichier non trouvé - EVAL file missing"
271
+ )
272
 
273
 
274
+ @patch("src.data_processing.load_data.logger")
275
+ @patch("src.data_processing.load_data.pd.read_csv")
276
  def test_load_and_merge_csvs_missing_eval_number_col(mock_read_csv, mock_logger):
277
  """Teste le cas où 'eval_number' manque dans df_eval."""
278
+ df_sirh_sample = pd.DataFrame({"id_employee": ["1"]})
279
+ df_eval_no_eval_number = pd.DataFrame({"autre_col": ["X"]}) # Pas de eval_number
280
+ df_sondage_sample = pd.DataFrame({"id_employee": ["1"]})
281
+ mock_read_csv.side_effect = [
282
+ df_sirh_sample,
283
+ df_eval_no_eval_number,
284
+ df_sondage_sample,
285
+ ]
286
+
287
  result_df = load_and_merge_csvs()
288
+
289
  assert result_df is None
290
+ mock_logger.error.assert_any_call("La colonne 'eval_number' est introuvable dans df_eval. Impossible de créer 'id_employee'.")
291
 
292
 
293
+ @patch("src.data_processing.load_data.logger")
294
+ @patch("src.data_processing.load_data.pd.read_csv")
295
  def test_load_and_merge_csvs_missing_id_employee_in_sirh(mock_read_csv, mock_logger):
296
  """Teste le cas où 'id_employee' manque dans df_sirh."""
297
+ df_sirh_no_id = pd.DataFrame({"autre_col_sirh": ["Y"]}) # Pas de id_employee
298
+ df_eval_sample = pd.DataFrame({"eval_number": ["E_1"]})
299
+ df_sondage_sample = pd.DataFrame({"id_employee": ["1"]})
300
  mock_read_csv.side_effect = [df_sirh_no_id, df_eval_sample, df_sondage_sample]
301
 
302
  result_df = load_and_merge_csvs()
303
 
304
  assert result_df is None
305
+ mock_logger.error.assert_any_call(
306
+ "La colonne 'id_employee' est introuvable dans df_sirh."
307
+ )
308
 
309
 
310
+ @patch("src.data_processing.load_data.logger")
311
+ @patch("src.data_processing.load_data.pd.read_csv")
312
  def test_load_and_merge_csvs_eval_number_conversion_warning(mock_read_csv, mock_logger):
313
  """Teste le warning pour les eval_number non convertibles."""
314
+ df_sirh_sample = pd.DataFrame({"id_employee": ["1", "bad_id_format"]})
315
+ df_eval_sample = pd.DataFrame(
316
+ {
317
+ "eval_number": ["E_1", "E_OK", "WRONG_FORMAT", "E_WRONG_TOO"],
318
+ "eval_feature": ["x1", "x2", "x3", "x4"],
319
+ }
320
+ )
321
+ df_sondage_sample = pd.DataFrame({"id_employee": ["1", "OK"]})
322
  mock_read_csv.side_effect = [df_sirh_sample, df_eval_sample, df_sondage_sample]
323
 
324
  merged_df = load_and_merge_csvs()
 
328
  # La logique de split('_').str[1] va lever une erreur ou retourner NaN si pas de '_'.
329
  # pd.to_numeric(..., errors='coerce') transformera ces erreurs en NaN.
330
  # Ici, 'WRONG_FORMAT' et 'E_WRONG_TOO' devraient produire des NaN pour 'id_employee' dans df_eval.
331
+ expected_warning_msg = "3 'eval_number' (après extraction) n'ont pas pu être convertis en id_employee numériques valides et sont devenus NaN."
332
+ mock_logger.warning.assert_any_call(expected_warning_msg)
333
  # mock_logger.warning.assert_any_call("3 'eval_number' n'ont pas pu être convertis en 'id_employee' valides.")
334
  assert merged_df is not None
335
  # La ligne avec id_employee '1' devrait avoir eval_feature 'x1'
336
+ assert (
337
+ merged_df.loc[merged_df["id_employee"] == "1", "eval_feature"].iloc[0] == "x1"
338
+ )
339
  # La ligne avec id_employee 'bad_id_format' ne trouvera pas de correspondance dans df_eval et aura NaN pour eval_feature
340
+ assert pd.isna(
341
+ merged_df.loc[merged_df["id_employee"] == "bad_id_format", "eval_feature"].iloc[
342
+ 0
343
+ ]
344
+ )
tests/unit/test_predict.py CHANGED
@@ -1,95 +1,123 @@
1
- import pytest
2
  from unittest.mock import patch, MagicMock
3
  import numpy as np
4
  import pandas as pd
5
  from src.modeling.predict import load_prediction_pipeline, predict_attrition
6
- from src import config # Pour config.MODEL_PATH
7
 
8
- @patch('src.modeling.predict.load') # Mocker joblib.load
9
- @patch('src.modeling.predict.config.MODEL_PATH')
10
- @patch('src.modeling.predict.logger')
11
- def test_load_prediction_pipeline_success(mock_logger, mock_model_path, mock_joblib_load):
 
 
 
12
  """Teste le chargement réussi de la pipeline."""
13
  mock_pipeline_obj = MagicMock()
14
  mock_joblib_load.return_value = mock_pipeline_obj
15
  mock_model_path.exists.return_value = True
16
-
17
  # Réinitialiser la variable globale _pipeline pour ce test
18
  from src.modeling import predict
19
- predict._pipeline = None
20
 
21
  pipeline = load_prediction_pipeline()
22
 
23
  mock_model_path.exists.assert_called_once()
24
  mock_joblib_load.assert_called_once_with(mock_model_path)
25
- mock_logger.info.assert_any_call(f"Chargement de la pipeline depuis {mock_model_path}...")
26
- mock_logger.info.assert_any_call("Pipeline chargée avec succès.")
 
 
 
 
27
  assert pipeline == mock_pipeline_obj
28
 
29
- @patch('src.modeling.predict.config.MODEL_PATH')
30
- @patch('src.modeling.predict.logger')
 
31
  def test_load_prediction_pipeline_file_not_found(mock_logger, mock_model_path):
32
  """Teste le cas où le fichier modèle n'est pas trouvé."""
33
  mock_model_path.exists.return_value = False
34
-
35
  from src.modeling import predict
36
  predict._pipeline = None
37
 
38
  pipeline = load_prediction_pipeline()
39
 
 
40
  mock_model_path.exists.assert_called_once()
41
- mock_logger.error.assert_any_call(f"Fichier pipeline non trouvé : {mock_model_path}")
 
 
 
42
  assert pipeline is None
43
 
44
- @patch('src.modeling.predict.load')
45
- @patch('src.modeling.predict.config.MODEL_PATH')
46
- @patch('src.modeling.predict.logger')
47
- def test_load_prediction_pipeline_load_exception(mock_logger, mock_model_path, mock_joblib_load):
 
 
 
48
  """Teste une exception lors du chargement du modèle."""
49
  mock_model_path.exists.return_value = True
50
- mock_joblib_load.side_effect = Exception("Load error")
 
 
51
 
52
  from src.modeling import predict
 
53
  predict._pipeline = None
54
 
55
  pipeline = load_prediction_pipeline()
56
 
57
- mock_logger.error.assert_any_call("Erreur lors du chargement de la pipeline : Load error", exc_info=True)
 
58
  assert pipeline is None
59
-
60
- @patch('src.modeling.predict.create_features')
61
- @patch('src.modeling.predict.map_binary_features')
62
- @patch('src.modeling.predict.clean_data')
63
- @patch('src.modeling.predict.load_prediction_pipeline')
64
- @patch('src.modeling.predict.logger')
 
65
  def test_predict_attrition_success_path(
66
- mock_logger, mock_load_pipeline, mock_clean_data,
67
- mock_map_binary, mock_create_features
 
 
 
68
  ):
69
  """Teste le chemin de succès de predict_attrition avec des mocks."""
70
  # 1. Configurer les mocks
71
  mock_pipeline_instance = MagicMock()
72
  # MODIFIÉ : Retourner des tableaux NumPy et pour 1 seul échantillon
73
- mock_pipeline_instance.predict_proba.return_value = np.array([[0.3, 0.7]]) # Pour 1 échantillon
74
- mock_pipeline_instance.predict.return_value = np.array([1]) # Pour 1 échantillon
 
 
75
  mock_load_pipeline.return_value = mock_pipeline_instance
76
 
77
  # Simuler les DataFrames retournés par chaque étape de preprocessing
78
  # df_raw_input a 1 ligne, donc toutes les sorties mockées doivent aussi correspondre à 1 ligne
79
- df_raw_input = pd.DataFrame({'feature1': ['A'], 'feature2': [10]}, index=['EMP_X'])
80
-
81
  # La fonction clean_data dans predict.py ajoute temporairement 'a_quitte_l_entreprise'
82
  df_cleaned_for_predict = df_raw_input.copy()
83
  # Supposons que config.TARGET_VARIABLE est 'a_quitte_l_entreprise_numeric'
84
- df_cleaned_for_predict[config.TARGET_VARIABLE] = 0 # Ajout factice pour le mock
85
- df_cleaned_for_predict['feature1_clean'] = 'A_clean' # Simuler le nettoyage
86
  mock_clean_data.return_value = df_cleaned_for_predict
87
 
88
  # df_mapped_mock et df_featured_mock doivent aussi avoir 1 ligne
89
- df_mapped_mock = pd.DataFrame({'feature1_map': [0], 'feature2_map': [10]}, index=['EMP_X'])
 
 
90
  mock_map_binary.return_value = df_mapped_mock
91
-
92
- df_featured_mock = pd.DataFrame({'feature1_final': [0], 'feature2_final': [100]}, index=['EMP_X'])
 
 
93
  mock_create_features.return_value = df_featured_mock
94
 
95
  # 2. Appeler la fonction
@@ -99,23 +127,27 @@ def test_predict_attrition_success_path(
99
  # 3. Assertions
100
  mock_load_pipeline.assert_called_once()
101
  mock_clean_data.assert_called_once()
102
- mock_map_binary.assert_called_once_with(mock_clean_data.return_value, config.BINARY_FEATURES_MAPPING)
 
 
103
  mock_create_features.assert_called_once_with(mock_map_binary.return_value)
104
 
105
  # X_predict_expected est dérivé de df_featured_mock (1 ligne)
106
- X_predict_expected = df_featured_mock.drop(columns=[config.TARGET_VARIABLE], errors='ignore')
107
-
108
  # Vérifier que les méthodes de la pipeline sont appelées
109
  # L'argument passé à predict_proba sera X_predict_expected
110
  # Vous pouvez utiliser pd.testing.assert_frame_equal si vous voulez être très précis sur l'argument
111
  # mock_pipeline_instance.predict_proba.assert_called_once_with(X_predict_expected) # Ne fonctionnera pas directement avec assert_frame_equal ici
112
-
113
  # Vérification des appels (plus simple et souvent suffisant)
114
  assert mock_pipeline_instance.predict_proba.call_count == 1
115
- assert mock_pipeline_instance.predict.call_count == 1 # Devrait passer maintenant
116
 
117
  assert len(results) == 1
118
- assert results[0]['id_employe'] == 'EMP_X'
119
- assert results[0]['probabilite_depart'] == 0.7 # Correspond à np.array([[0.3, 0.7]])[:, 1][0]
120
- assert results[0]['prediction_depart'] == 1 # Correspond à np.array([1])[0]
121
- mock_logger.info.assert_any_call("Prédiction terminée.")
 
 
 
 
1
  from unittest.mock import patch, MagicMock
2
  import numpy as np
3
  import pandas as pd
4
  from src.modeling.predict import load_prediction_pipeline, predict_attrition
5
+ from src import config # Pour config.MODEL_PATH
6
 
7
+
8
+ @patch("src.modeling.predict.load") # Mocker joblib.load
9
+ @patch("src.modeling.predict.config.MODEL_PATH")
10
+ @patch("src.modeling.predict.logger")
11
+ def test_load_prediction_pipeline_success(
12
+ mock_logger, mock_model_path, mock_joblib_load
13
+ ):
14
  """Teste le chargement réussi de la pipeline."""
15
  mock_pipeline_obj = MagicMock()
16
  mock_joblib_load.return_value = mock_pipeline_obj
17
  mock_model_path.exists.return_value = True
18
+
19
  # Réinitialiser la variable globale _pipeline pour ce test
20
  from src.modeling import predict
21
+ predict._pipeline = None
22
 
23
  pipeline = load_prediction_pipeline()
24
 
25
  mock_model_path.exists.assert_called_once()
26
  mock_joblib_load.assert_called_once_with(mock_model_path)
27
+
28
+ # CORRECTION ICI :
29
+ expected_log_msg_loading = f"Chargement de la pipeline de prédiction depuis : {mock_model_path}..."
30
+ mock_logger.info.assert_any_call(expected_log_msg_loading)
31
+
32
+ mock_logger.info.assert_any_call("Pipeline de prédiction chargée avec succès.") # Vérifiez aussi ce message
33
  assert pipeline == mock_pipeline_obj
34
 
35
+
36
+ @patch("src.modeling.predict.config.MODEL_PATH")
37
+ @patch("src.modeling.predict.logger")
38
  def test_load_prediction_pipeline_file_not_found(mock_logger, mock_model_path):
39
  """Teste le cas où le fichier modèle n'est pas trouvé."""
40
  mock_model_path.exists.return_value = False
41
+
42
  from src.modeling import predict
43
  predict._pipeline = None
44
 
45
  pipeline = load_prediction_pipeline()
46
 
47
+ # Assertions
48
  mock_model_path.exists.assert_called_once()
49
+
50
+ # CORRECTION ICI :
51
+ expected_log_msg = f"Fichier pipeline non trouvé à l'emplacement configuré : {mock_model_path}"
52
+ mock_logger.error.assert_any_call(expected_log_msg)
53
  assert pipeline is None
54
 
55
+
56
+ @patch("src.modeling.predict.load")
57
+ @patch("src.modeling.predict.config.MODEL_PATH")
58
+ @patch("src.modeling.predict.logger")
59
+ def test_load_prediction_pipeline_load_exception(
60
+ mock_logger, mock_model_path, mock_joblib_load
61
+ ):
62
  """Teste une exception lors du chargement du modèle."""
63
  mock_model_path.exists.return_value = True
64
+ simulated_exception = Exception("Load error")
65
+ mock_joblib_load.side_effect = simulated_exception
66
+ mock_model_path.exists.return_value = True
67
 
68
  from src.modeling import predict
69
+
70
  predict._pipeline = None
71
 
72
  pipeline = load_prediction_pipeline()
73
 
74
+ expected_log_msg = f"Erreur critique lors du chargement de la pipeline : {simulated_exception}"
75
+ mock_logger.error.assert_any_call(expected_log_msg, exc_info=True)
76
  assert pipeline is None
77
+
78
+
79
+ @patch("src.modeling.predict.create_features")
80
+ @patch("src.modeling.predict.map_binary_features")
81
+ @patch("src.modeling.predict.clean_data")
82
+ @patch("src.modeling.predict.load_prediction_pipeline")
83
+ @patch("src.modeling.predict.logger")
84
  def test_predict_attrition_success_path(
85
+ mock_logger,
86
+ mock_load_pipeline,
87
+ mock_clean_data,
88
+ mock_map_binary,
89
+ mock_create_features,
90
  ):
91
  """Teste le chemin de succès de predict_attrition avec des mocks."""
92
  # 1. Configurer les mocks
93
  mock_pipeline_instance = MagicMock()
94
  # MODIFIÉ : Retourner des tableaux NumPy et pour 1 seul échantillon
95
+ mock_pipeline_instance.predict_proba.return_value = np.array(
96
+ [[0.3, 0.7]]
97
+ ) # Pour 1 échantillon
98
+ mock_pipeline_instance.predict.return_value = np.array([1]) # Pour 1 échantillon
99
  mock_load_pipeline.return_value = mock_pipeline_instance
100
 
101
  # Simuler les DataFrames retournés par chaque étape de preprocessing
102
  # df_raw_input a 1 ligne, donc toutes les sorties mockées doivent aussi correspondre à 1 ligne
103
+ df_raw_input = pd.DataFrame({"feature1": ["A"], "feature2": [10]}, index=["EMP_X"])
104
+
105
  # La fonction clean_data dans predict.py ajoute temporairement 'a_quitte_l_entreprise'
106
  df_cleaned_for_predict = df_raw_input.copy()
107
  # Supposons que config.TARGET_VARIABLE est 'a_quitte_l_entreprise_numeric'
108
+ df_cleaned_for_predict[config.TARGET_VARIABLE] = 0 # Ajout factice pour le mock
109
+ df_cleaned_for_predict["feature1_clean"] = "A_clean" # Simuler le nettoyage
110
  mock_clean_data.return_value = df_cleaned_for_predict
111
 
112
  # df_mapped_mock et df_featured_mock doivent aussi avoir 1 ligne
113
+ df_mapped_mock = pd.DataFrame(
114
+ {"feature1_map": [0], "feature2_map": [10]}, index=["EMP_X"]
115
+ )
116
  mock_map_binary.return_value = df_mapped_mock
117
+
118
+ df_featured_mock = pd.DataFrame(
119
+ {"feature1_final": [0], "feature2_final": [100]}, index=["EMP_X"]
120
+ )
121
  mock_create_features.return_value = df_featured_mock
122
 
123
  # 2. Appeler la fonction
 
127
  # 3. Assertions
128
  mock_load_pipeline.assert_called_once()
129
  mock_clean_data.assert_called_once()
130
+ mock_map_binary.assert_called_once_with(
131
+ mock_clean_data.return_value, config.BINARY_FEATURES_MAPPING
132
+ )
133
  mock_create_features.assert_called_once_with(mock_map_binary.return_value)
134
 
135
  # X_predict_expected est dérivé de df_featured_mock (1 ligne)
136
+ # X_predict_expected = df_featured_mock.drop(columns=[config.TARGET_VARIABLE], errors="ignore")
137
+
138
  # Vérifier que les méthodes de la pipeline sont appelées
139
  # L'argument passé à predict_proba sera X_predict_expected
140
  # Vous pouvez utiliser pd.testing.assert_frame_equal si vous voulez être très précis sur l'argument
141
  # mock_pipeline_instance.predict_proba.assert_called_once_with(X_predict_expected) # Ne fonctionnera pas directement avec assert_frame_equal ici
142
+
143
  # Vérification des appels (plus simple et souvent suffisant)
144
  assert mock_pipeline_instance.predict_proba.call_count == 1
145
+ assert mock_pipeline_instance.predict.call_count == 1 # Devrait passer maintenant
146
 
147
  assert len(results) == 1
148
+ assert results[0]["id_employe"] == "EMP_X"
149
+ assert (
150
+ results[0]["probabilite_depart"] == 0.7
151
+ ) # Correspond à np.array([[0.3, 0.7]])[:, 1][0]
152
+ assert results[0]["prediction_depart"] == 1 # Correspond à np.array([1])[0]
153
+ mock_logger.info.assert_any_call("Prédictions formatées avec succès.")
tests/unit/test_preprocess.py CHANGED
@@ -4,9 +4,15 @@ from sklearn.compose import ColumnTransformer
4
  from sklearn.pipeline import Pipeline
5
  from sklearn.preprocessing import StandardScaler, OneHotEncoder, OrdinalEncoder
6
  from sklearn.impute import SimpleImputer
7
- from src.data_processing.preprocess import map_binary_features, clean_data, build_preprocessor, run_preprocessing_pipeline
 
 
 
 
 
8
  from src import config # Pour config.TARGET_VARIABLE
9
 
 
10
  # --- Tests existants pour map_binary_features (assurez-vous qu'ils sont robustes) ---
11
  def test_map_binary_features_simple_case():
12
  """Teste le mappage binaire simple."""
@@ -40,6 +46,7 @@ def test_map_binary_features_unmapped_values():
40
  assert df_output["genre"].iloc[1] == 1
41
  assert pd.isna(df_output["genre"].iloc[2]) # Les valeurs non mappées deviennent NaN
42
 
 
43
  @pytest.fixture
44
  def sample_df_for_clean():
45
  """Crée un DataFrame d'exemple pour tester clean_data."""
@@ -62,9 +69,8 @@ def test_clean_data_target_conversion(sample_df_for_clean):
62
  """Vérifie la conversion correcte de la variable cible."""
63
  df_cleaned = clean_data(sample_df_for_clean)
64
  assert config.TARGET_VARIABLE in df_cleaned.columns
65
- assert df_cleaned[config.TARGET_VARIABLE].equals(
66
- pd.Series([1, 0, 1, 0], name=config.TARGET_VARIABLE)
67
- )
68
 
69
 
70
  def test_clean_data_column_dropping(sample_df_for_clean):
@@ -91,112 +97,127 @@ def test_clean_data_percentage_conversion(sample_df_for_clean):
91
  ) # "Pas d'augmentation" devient NaN à cause de errors='coerce'
92
 
93
 
94
- def test_clean_data_with_missing_target_column():
95
- """Vérifie que clean_data lève une erreur si la colonne cible est manquante."""
96
- data = {"employee_id": [1], "eval_number": ["E_1"]}
97
- df_no_target = pd.DataFrame(data)
98
- with pytest.raises(ValueError, match="Colonne cible manquante"):
99
- clean_data(df_no_target)
100
-
101
  def test_build_preprocessor_all_types():
102
  """Teste build_preprocessor avec tous les types de colonnes."""
103
- numerical_cols = ['age', 'salaire']
104
- onehot_cols = ['departement', 'poste']
105
- ordinal_cols = ['frequence_deplacement']
106
- ordinal_categories = {
107
- 'frequence_deplacement': ['Bas', 'Moyen', 'Haut']
108
- }
109
 
110
- preprocessor = build_preprocessor(numerical_cols, onehot_cols, ordinal_cols, ordinal_categories)
 
 
111
 
112
  assert isinstance(preprocessor, ColumnTransformer)
113
- assert len(preprocessor.transformers) == 3 # Un pour num, un pour onehot, un pour ord_frequence_deplacement
 
 
114
 
115
  # Vérifier le transformateur numérique
116
- num_transformer_tuple = next(t for t in preprocessor.transformers if t[0] == 'num')
117
  assert isinstance(num_transformer_tuple[1], Pipeline)
118
  assert len(num_transformer_tuple[1].steps) == 2
119
  assert isinstance(num_transformer_tuple[1].steps[0][1], SimpleImputer)
120
- assert num_transformer_tuple[1].steps[0][1].strategy == 'median'
121
  assert isinstance(num_transformer_tuple[1].steps[1][1], StandardScaler)
122
  assert num_transformer_tuple[2] == numerical_cols
123
 
124
  # Vérifier le transformateur one-hot
125
- onehot_transformer_tuple = next(t for t in preprocessor.transformers if t[0] == 'onehot')
 
 
126
  assert isinstance(onehot_transformer_tuple[1], Pipeline)
127
  assert len(onehot_transformer_tuple[1].steps) == 2
128
  assert isinstance(onehot_transformer_tuple[1].steps[0][1], SimpleImputer)
129
- assert onehot_transformer_tuple[1].steps[0][1].strategy == 'most_frequent'
130
  assert isinstance(onehot_transformer_tuple[1].steps[1][1], OneHotEncoder)
131
- assert onehot_transformer_tuple[1].steps[1][1].drop == 'first'
132
- assert onehot_transformer_tuple[1].steps[1][1].handle_unknown == 'ignore'
133
  assert onehot_transformer_tuple[2] == onehot_cols
134
-
135
  # Vérifier le transformateur ordinal pour 'frequence_deplacement'
136
- ord_transformer_tuple = next(t for t in preprocessor.transformers if t[0] == 'ord_frequence_deplacement')
 
 
137
  assert isinstance(ord_transformer_tuple[1], Pipeline)
138
  assert len(ord_transformer_tuple[1].steps) == 2
139
  assert isinstance(ord_transformer_tuple[1].steps[0][1], SimpleImputer)
140
  assert isinstance(ord_transformer_tuple[1].steps[1][1], OrdinalEncoder)
141
- assert ord_transformer_tuple[1].steps[1][1].categories == [['Bas', 'Moyen', 'Haut']]
142
- assert ord_transformer_tuple[2] == ['frequence_deplacement']
143
 
144
 
145
  def test_build_preprocessor_only_numerical():
146
  """Teste build_preprocessor avec seulement des colonnes numériques."""
147
- numerical_cols = ['age', 'salaire']
148
  preprocessor = build_preprocessor(numerical_cols, [], [], {})
149
  assert isinstance(preprocessor, ColumnTransformer)
150
  assert len(preprocessor.transformers) == 1
151
- assert preprocessor.transformers[0][0] == 'num'
152
  assert preprocessor.transformers[0][2] == numerical_cols
153
 
 
154
  def test_build_preprocessor_ordinal_col_missing_categories():
155
  """Teste que build_preprocessor lève une erreur si une catégorie ordinale est manquante."""
156
  numerical_cols = []
157
  onehot_cols = []
158
- ordinal_cols = ['niveau_satisfaction'] # Cette colonne n'est pas dans ordinal_categories
159
- with pytest.raises(ValueError, match="Les catégories pour 'niveau_satisfaction' ne sont pas définies."):
160
- build_preprocessor(numerical_cols, onehot_cols, ordinal_cols, config.ORDINAL_FEATURES_CATEGORIES)
 
 
 
 
161
 
162
  def test_build_preprocessor_no_cols():
163
  """Teste build_preprocessor quand aucune colonne n'est spécifiée."""
164
  preprocessor = build_preprocessor([], [], [], {})
165
  assert isinstance(preprocessor, ColumnTransformer)
166
- assert len(preprocessor.transformers) == 0 # Devrait retourner un ColumnTransformer vide
167
- assert preprocessor.remainder == 'passthrough' # Selon votre implémentation actuelle
168
- # (j'avais mis 'drop', si c'est passthrough, adaptez le test ou la fonction)
 
 
 
 
169
  # Ma version de build_preprocessor avait 'remainder="drop"' et retournait un CT vide.
170
  # Si vous voulez qu'il soit 'passthrough' dans ce cas, modifiez la fonction ou le test.
171
  # Pour 'remainder="drop"', len(preprocessor.transformers) == 0 est correct.
172
-
 
173
  @pytest.fixture
174
  def sample_raw_df_for_pipeline():
175
  """DataFrame brut d'exemple pour tester le pipeline complet."""
176
  data = {
177
- 'employee_id': ['E_1', 'E_2', 'E_3', 'E_4'],
178
- 'a_quitte_l_entreprise': ['Oui', 'Non', 'Oui', 'Non'],
179
- 'genre': ['M', 'F', 'M', 'F'],
180
- 'heure_supplementaires': ['Oui', 'Non', 'Non', 'Oui'],
181
- 'frequence_deplacement': ['Occasionnel', 'Aucun', 'Frequent', 'Occasionnel'],
182
- 'augmentation_salaire_precedente': ["10 %", "5 %", "12 %", "8 %"],
183
- 'age': [30, 45, 22, 50],
184
- 'salaire_mensuel_brut': [5000, 8000, 3500, 9000],
185
- 'departement': ['Ventes', 'R&D', 'Ventes', 'Marketing'],
186
  # Ajoutez d'autres colonnes pour couvrir tous vos types
187
- config.TARGET_VARIABLE: [1,0,1,0] # La fonction clean_data va écraser ça, mais pour la cohérence
 
 
 
 
 
188
  }
189
  return pd.DataFrame(data)
190
 
 
191
  def test_run_preprocessing_pipeline_fit_true(sample_raw_df_for_pipeline):
192
  """Teste run_preprocessing_pipeline en mode fit=True."""
193
  df_in = sample_raw_df_for_pipeline.copy()
194
-
195
  X_processed, y_processed, fitted_processor = run_preprocessing_pipeline(
196
  df_in,
197
  binary_cols_map=config.BINARY_FEATURES_MAPPING,
198
- ordinal_cols_categories=config.ORDINAL_FEATURES_CATEGORIES,
199
- fit=True
200
  )
201
 
202
  assert isinstance(X_processed, pd.DataFrame)
@@ -204,20 +225,26 @@ def test_run_preprocessing_pipeline_fit_true(sample_raw_df_for_pipeline):
204
  assert isinstance(fitted_processor, ColumnTransformer)
205
  assert X_processed.shape[0] == len(df_in)
206
  assert y_processed.shape[0] == len(df_in)
207
-
208
  # Vérifier que les colonnes attendues (après OHE, etc.) sont là
209
  # Ceci dépendra de vos colonnes exactes et de get_feature_names_out
210
  # Par exemple, on s'attend à voir des colonnes issues du OHE de 'departement'
211
- assert any(col.startswith('onehot__departement_') for col in X_processed.columns)
212
- assert 'num__age' in X_processed.columns
213
- assert f'ord_frequence_deplacement__{config.ORDINAL_FEATURES_CATEGORIES["frequence_deplacement"][0]}' not in X_processed.columns # OrdinalEncoder ne préfixe pas comme ça
214
- assert 'ord_frequence_deplacement__frequence_deplacement' in X_processed.columns # ou juste 'frequence_deplacement' si remainder='passthrough' et ord pas dans CT
 
 
 
 
 
215
 
216
  # Vérifier que le processeur est "fitté" (ex: StandardScaler a mean_ et scale_)
217
  # Accéder au transformateur numérique dans le ColumnTransformer
218
- num_pipeline = fitted_processor.named_transformers_['num']
219
- scaler = num_pipeline.named_steps['scaler']
220
- assert hasattr(scaler, 'mean_') and scaler.mean_ is not None
 
221
 
222
  @pytest.mark.filterwarnings("ignore:Found unknown categories.*:UserWarning")
223
  def test_run_preprocessing_pipeline_fit_false(sample_raw_df_for_pipeline):
@@ -228,16 +255,16 @@ def test_run_preprocessing_pipeline_fit_false(sample_raw_df_for_pipeline):
228
  _, _, fitted_processor = run_preprocessing_pipeline(
229
  df_train,
230
  binary_cols_map=config.BINARY_FEATURES_MAPPING,
231
- ordinal_cols_categories=config.ORDINAL_FEATURES_CATEGORIES,
232
- fit=True
233
  )
234
-
235
  X_test_processed, y_test_processed = run_preprocessing_pipeline(
236
- df_test,
237
- binary_cols_map=config.BINARY_FEATURES_MAPPING, # Toujours nécessaire pour map_binary_features
238
- ordinal_cols_categories=config.ORDINAL_FEATURES_CATEGORIES, # Toujours nécessaire pour build_preprocessor si pas fitté
239
- preprocessor=fitted_processor,
240
- fit=False
241
  )
242
 
243
  assert isinstance(X_test_processed, pd.DataFrame)
@@ -245,9 +272,22 @@ def test_run_preprocessing_pipeline_fit_false(sample_raw_df_for_pipeline):
245
  # Le nombre de colonnes doit être le même que celui obtenu après fit_transform sur le train
246
  # Pour cela, il faudrait stocker X_train_processed.shape[1] du test précédent
247
 
248
- def test_run_preprocessing_pipeline_fit_false_no_processor():
249
- """Teste que fit=False sans preprocessor lève une ValueError."""
250
- df_in = pd.DataFrame({'a_quitte_l_entreprise': ['Oui']}) # Minimal DF
251
-
252
- with pytest.raises(ValueError, match="Un preprocessor doit être fourni si fit=False."):
253
- run_preprocessing_pipeline(df_in, fit=False)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4
  from sklearn.pipeline import Pipeline
5
  from sklearn.preprocessing import StandardScaler, OneHotEncoder, OrdinalEncoder
6
  from sklearn.impute import SimpleImputer
7
+ from src.data_processing.preprocess import (
8
+ map_binary_features,
9
+ clean_data,
10
+ build_preprocessor,
11
+ run_preprocessing_pipeline,
12
+ )
13
  from src import config # Pour config.TARGET_VARIABLE
14
 
15
+
16
  # --- Tests existants pour map_binary_features (assurez-vous qu'ils sont robustes) ---
17
  def test_map_binary_features_simple_case():
18
  """Teste le mappage binaire simple."""
 
46
  assert df_output["genre"].iloc[1] == 1
47
  assert pd.isna(df_output["genre"].iloc[2]) # Les valeurs non mappées deviennent NaN
48
 
49
+
50
  @pytest.fixture
51
  def sample_df_for_clean():
52
  """Crée un DataFrame d'exemple pour tester clean_data."""
 
69
  """Vérifie la conversion correcte de la variable cible."""
70
  df_cleaned = clean_data(sample_df_for_clean)
71
  assert config.TARGET_VARIABLE in df_cleaned.columns
72
+ expected_series = pd.Series([1, 0, 1, 0], name=config.TARGET_VARIABLE, dtype='Int64') # Spécifier dtype
73
+ assert df_cleaned[config.TARGET_VARIABLE].equals(expected_series)
 
74
 
75
 
76
  def test_clean_data_column_dropping(sample_df_for_clean):
 
97
  ) # "Pas d'augmentation" devient NaN à cause de errors='coerce'
98
 
99
 
 
 
 
 
 
 
 
100
  def test_build_preprocessor_all_types():
101
  """Teste build_preprocessor avec tous les types de colonnes."""
102
+ numerical_cols = ["age", "salaire"]
103
+ onehot_cols = ["departement", "poste"]
104
+ ordinal_cols = ["frequence_deplacement"]
105
+ ordinal_categories = {"frequence_deplacement": ["Bas", "Moyen", "Haut"]}
 
 
106
 
107
+ preprocessor = build_preprocessor(
108
+ numerical_cols, onehot_cols, ordinal_cols, ordinal_categories
109
+ )
110
 
111
  assert isinstance(preprocessor, ColumnTransformer)
112
+ assert (
113
+ len(preprocessor.transformers) == 3
114
+ ) # Un pour num, un pour onehot, un pour ord_frequence_deplacement
115
 
116
  # Vérifier le transformateur numérique
117
+ num_transformer_tuple = next(t for t in preprocessor.transformers if t[0] == "num")
118
  assert isinstance(num_transformer_tuple[1], Pipeline)
119
  assert len(num_transformer_tuple[1].steps) == 2
120
  assert isinstance(num_transformer_tuple[1].steps[0][1], SimpleImputer)
121
+ assert num_transformer_tuple[1].steps[0][1].strategy == "median"
122
  assert isinstance(num_transformer_tuple[1].steps[1][1], StandardScaler)
123
  assert num_transformer_tuple[2] == numerical_cols
124
 
125
  # Vérifier le transformateur one-hot
126
+ onehot_transformer_tuple = next(
127
+ t for t in preprocessor.transformers if t[0] == "onehot"
128
+ )
129
  assert isinstance(onehot_transformer_tuple[1], Pipeline)
130
  assert len(onehot_transformer_tuple[1].steps) == 2
131
  assert isinstance(onehot_transformer_tuple[1].steps[0][1], SimpleImputer)
132
+ assert onehot_transformer_tuple[1].steps[0][1].strategy == "most_frequent"
133
  assert isinstance(onehot_transformer_tuple[1].steps[1][1], OneHotEncoder)
134
+ assert onehot_transformer_tuple[1].steps[1][1].drop == "first"
135
+ assert onehot_transformer_tuple[1].steps[1][1].handle_unknown == "ignore"
136
  assert onehot_transformer_tuple[2] == onehot_cols
137
+
138
  # Vérifier le transformateur ordinal pour 'frequence_deplacement'
139
+ ord_transformer_tuple = next(
140
+ t for t in preprocessor.transformers if t[0] == "ord_frequence_deplacement"
141
+ )
142
  assert isinstance(ord_transformer_tuple[1], Pipeline)
143
  assert len(ord_transformer_tuple[1].steps) == 2
144
  assert isinstance(ord_transformer_tuple[1].steps[0][1], SimpleImputer)
145
  assert isinstance(ord_transformer_tuple[1].steps[1][1], OrdinalEncoder)
146
+ assert ord_transformer_tuple[1].steps[1][1].categories == [["Bas", "Moyen", "Haut"]]
147
+ assert ord_transformer_tuple[2] == ["frequence_deplacement"]
148
 
149
 
150
  def test_build_preprocessor_only_numerical():
151
  """Teste build_preprocessor avec seulement des colonnes numériques."""
152
+ numerical_cols = ["age", "salaire"]
153
  preprocessor = build_preprocessor(numerical_cols, [], [], {})
154
  assert isinstance(preprocessor, ColumnTransformer)
155
  assert len(preprocessor.transformers) == 1
156
+ assert preprocessor.transformers[0][0] == "num"
157
  assert preprocessor.transformers[0][2] == numerical_cols
158
 
159
+
160
  def test_build_preprocessor_ordinal_col_missing_categories():
161
  """Teste que build_preprocessor lève une erreur si une catégorie ordinale est manquante."""
162
  numerical_cols = []
163
  onehot_cols = []
164
+ ordinal_cols = [
165
+ "niveau_satisfaction"
166
+ ] # Cette colonne n'est pas dans ordinal_categories
167
+ expected_error_message = "Les catégories pour la colonne ordinale 'niveau_satisfaction' ne sont pas définies dans ordinal_categories_map."
168
+ with pytest.raises(ValueError, match=expected_error_message):
169
+ build_preprocessor(numerical_cols, onehot_cols, ordinal_cols, config.ORDINAL_FEATURES_CATEGORIES) # Utilisez un nom de variable clair pour le dict de catégories
170
+
171
 
172
  def test_build_preprocessor_no_cols():
173
  """Teste build_preprocessor quand aucune colonne n'est spécifiée."""
174
  preprocessor = build_preprocessor([], [], [], {})
175
  assert isinstance(preprocessor, ColumnTransformer)
176
+ assert (
177
+ len(preprocessor.transformers) == 0
178
+ ) # Devrait retourner un ColumnTransformer vide
179
+ assert (
180
+ preprocessor.remainder == "passthrough"
181
+ ) # Selon votre implémentation actuelle
182
+ # (j'avais mis 'drop', si c'est passthrough, adaptez le test ou la fonction)
183
  # Ma version de build_preprocessor avait 'remainder="drop"' et retournait un CT vide.
184
  # Si vous voulez qu'il soit 'passthrough' dans ce cas, modifiez la fonction ou le test.
185
  # Pour 'remainder="drop"', len(preprocessor.transformers) == 0 est correct.
186
+
187
+
188
  @pytest.fixture
189
  def sample_raw_df_for_pipeline():
190
  """DataFrame brut d'exemple pour tester le pipeline complet."""
191
  data = {
192
+ "employee_id": ["E_1", "E_2", "E_3", "E_4"],
193
+ "a_quitte_l_entreprise": ["Oui", "Non", "Oui", "Non"],
194
+ "genre": ["M", "F", "M", "F"],
195
+ "heure_supplementaires": ["Oui", "Non", "Non", "Oui"],
196
+ "frequence_deplacement": ["Occasionnel", "Aucun", "Frequent", "Occasionnel"],
197
+ "augmentation_salaire_precedente": ["10 %", "5 %", "12 %", "8 %"],
198
+ "age": [30, 45, 22, 50],
199
+ "salaire_mensuel_brut": [5000, 8000, 3500, 9000],
200
+ "departement": ["Ventes", "R&D", "Ventes", "Marketing"],
201
  # Ajoutez d'autres colonnes pour couvrir tous vos types
202
+ config.TARGET_VARIABLE: [
203
+ 1,
204
+ 0,
205
+ 1,
206
+ 0,
207
+ ], # La fonction clean_data va écraser ça, mais pour la cohérence
208
  }
209
  return pd.DataFrame(data)
210
 
211
+
212
  def test_run_preprocessing_pipeline_fit_true(sample_raw_df_for_pipeline):
213
  """Teste run_preprocessing_pipeline en mode fit=True."""
214
  df_in = sample_raw_df_for_pipeline.copy()
215
+
216
  X_processed, y_processed, fitted_processor = run_preprocessing_pipeline(
217
  df_in,
218
  binary_cols_map=config.BINARY_FEATURES_MAPPING,
219
+ ordinal_cols_categories_map=config.ORDINAL_FEATURES_CATEGORIES,
220
+ fit=True,
221
  )
222
 
223
  assert isinstance(X_processed, pd.DataFrame)
 
225
  assert isinstance(fitted_processor, ColumnTransformer)
226
  assert X_processed.shape[0] == len(df_in)
227
  assert y_processed.shape[0] == len(df_in)
228
+
229
  # Vérifier que les colonnes attendues (après OHE, etc.) sont là
230
  # Ceci dépendra de vos colonnes exactes et de get_feature_names_out
231
  # Par exemple, on s'attend à voir des colonnes issues du OHE de 'departement'
232
+ assert any(col.startswith("onehot__departement_") for col in X_processed.columns)
233
+ assert "num__age" in X_processed.columns
234
+ assert (
235
+ f'ord_frequence_deplacement__{config.ORDINAL_FEATURES_CATEGORIES["frequence_deplacement"][0]}'
236
+ not in X_processed.columns
237
+ ) # OrdinalEncoder ne préfixe pas comme ça
238
+ assert (
239
+ "ord_frequence_deplacement__frequence_deplacement" in X_processed.columns
240
+ ) # ou juste 'frequence_deplacement' si remainder='passthrough' et ord pas dans CT
241
 
242
  # Vérifier que le processeur est "fitté" (ex: StandardScaler a mean_ et scale_)
243
  # Accéder au transformateur numérique dans le ColumnTransformer
244
+ num_pipeline = fitted_processor.named_transformers_["num"]
245
+ scaler = num_pipeline.named_steps["scaler"]
246
+ assert hasattr(scaler, "mean_") and scaler.mean_ is not None
247
+
248
 
249
  @pytest.mark.filterwarnings("ignore:Found unknown categories.*:UserWarning")
250
  def test_run_preprocessing_pipeline_fit_false(sample_raw_df_for_pipeline):
 
255
  _, _, fitted_processor = run_preprocessing_pipeline(
256
  df_train,
257
  binary_cols_map=config.BINARY_FEATURES_MAPPING,
258
+ ordinal_cols_categories_map=config.ORDINAL_FEATURES_CATEGORIES,
259
+ fit=True,
260
  )
261
+
262
  X_test_processed, y_test_processed = run_preprocessing_pipeline(
263
+ df_test,
264
+ binary_cols_map=config.BINARY_FEATURES_MAPPING, # Toujours nécessaire pour map_binary_features
265
+ ordinal_cols_categories_map=config.ORDINAL_FEATURES_CATEGORIES, # Toujours nécessaire pour build_preprocessor si pas fitté
266
+ preprocessor=fitted_processor,
267
+ fit=False,
268
  )
269
 
270
  assert isinstance(X_test_processed, pd.DataFrame)
 
272
  # Le nombre de colonnes doit être le même que celui obtenu après fit_transform sur le train
273
  # Pour cela, il faudrait stocker X_train_processed.shape[1] du test précédent
274
 
275
+
276
+ def test_run_preprocessing_pipeline_fit_false_no_processor(sample_raw_df_for_pipeline): # Utilisez la fixture
277
+ """Teste que fit=False sans preprocessor lève une ValueError pour cette raison spécifique."""
278
+ # Utilisez un DataFrame d'entrée qui est suffisamment complet pour passer les étapes
279
+ # de clean_data, map_binary_features, create_features et l'identification des colonnes.
280
+ # La fixture sample_raw_df_for_pipeline devrait convenir si elle est bien définie.
281
+ df_in = sample_raw_df_for_pipeline.copy()
282
+
283
+ with pytest.raises(
284
+ ValueError, match="Un preprocessor doit être fourni si fit=False."
285
+ ):
286
+ # Appelez avec les mappings pour éviter des erreurs avant le check final
287
+ run_preprocessing_pipeline(
288
+ df_in,
289
+ binary_cols_map=config.BINARY_FEATURES_MAPPING,
290
+ ordinal_cols_categories_map=config.ORDINAL_FEATURES_CATEGORIES,
291
+ preprocessor=None, # Explicitement None
292
+ fit=False
293
+ )
tests/unit/test_train_model.py CHANGED
@@ -1,50 +1,117 @@
1
- import pytest
2
- from unittest.mock import patch, MagicMock, ANY
3
  import pandas as pd
4
  import numpy as np
5
  from src.modeling.train_model import train_and_evaluate_pipeline
6
- from src import config # Pour les mappings et la cible
7
 
8
- @patch('src.modeling.train_model.dump') # joblib.dump
9
- @patch('src.modeling.train_model.confusion_matrix', return_value=np.array([[0,0],[0,0]]))
10
- @patch('src.modeling.train_model.fbeta_score', return_value=0.5)
11
- @patch('src.modeling.train_model.classification_report', return_value="Mocked Report")
12
- @patch('src.modeling.train_model.Pipeline') # sklearn.pipeline.Pipeline
13
- @patch('src.modeling.train_model.LogisticRegression') # sklearn.linear_model.LogisticRegression
14
- @patch('src.modeling.train_model.build_preprocessor')
15
- @patch('src.modeling.train_model.create_features') # de preprocess
16
- @patch('src.modeling.train_model.map_binary_features') # de preprocess
17
- @patch('src.modeling.train_model.clean_data') # de preprocess
18
- @patch('src.modeling.train_model.get_data') # de load_data
19
- @patch('src.modeling.train_model.logger')
 
 
 
20
  def test_train_and_evaluate_pipeline_success_path(
21
- mock_logger, mock_get_data, mock_clean_data, mock_map_binary, mock_create_features,
22
- mock_build_preprocessor, mock_LogisticRegression, mock_Pipeline,
23
- mock_classification_report, mock_fbeta_score, mock_confusion_matrix, mock_joblib_dump
 
 
 
 
 
 
 
24
  ):
25
  """Teste le chemin principal de train_and_evaluate_pipeline avec des mocks."""
26
  # 1. Configurer les mocks
27
  # Mock get_data pour retourner des données avec ASSEZ DE LIGNES
28
- sample_df_from_db = pd.DataFrame({
29
- 'id_employee': [str(i) for i in range(1, 11)], # 10 lignes
30
- config.TARGET_VARIABLE: [0, 1, 0, 1, 0, 1, 0, 1, 0, 0], # 4 de classe 1, 6 de classe 0
31
- 'genre': [0, 1, 0, 1, 0, 1, 0, 1, 0, 1], # Alterner pour avoir des valeurs
32
- 'heure_supplementaires': [0, 1, 0, 0, 1, 1, 0, 0, 1, 0],
33
- 'frequence_deplacement': ['Non-Travel', 'Travel_Rarely', 'Travel_Frequently', 'Non-Travel', 'Travel_Rarely',
34
- 'Non-Travel', 'Travel_Rarely', 'Travel_Frequently', 'Non-Travel', 'Travel_Rarely'],
35
- 'age': [30, 45, 22, 50, 38, 29, 54, 33, 41, 47],
36
- 'salaire_mensuel_brut': [5000, 8000, 3500, 9000, 6000, 4800, 9500, 5500, 7000, 8200.0],
37
- 'department': ['Ventes', 'R&D', 'Ventes', 'Marketing', 'R&D',
38
- 'Ventes', 'R&D', 'Ventes', 'Marketing', 'R&D'],
39
- 'augmentation_salaire_precedente': [10.0, 15.0, 5.0, 12.0, 8.0, 9.0, 11.0, 6.0, 13.0, 7.0]
40
- })
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
41
  mock_get_data.return_value = sample_df_from_db
42
 
43
  df_after_features = sample_df_from_db.copy()
44
  mock_create_features.return_value = df_after_features
45
-
46
  mock_preprocessor_instance = MagicMock()
47
- mock_build_preprocessor.return_value = (mock_preprocessor_instance, ['mock_col1', 'mock_col2']) # Exemple
 
 
 
48
 
49
  mock_classifier_instance = MagicMock()
50
  mock_LogisticRegression.return_value = mock_classifier_instance
@@ -53,8 +120,10 @@ def test_train_and_evaluate_pipeline_success_path(
53
  # Avec 10 lignes et test_size=0.2, X_test aura 2 lignes (10 * 0.2 = 2)
54
  # X_train aura 8 lignes.
55
  # Les mocks pour predict et predict_proba doivent donc retourner pour 2 échantillons
56
- mock_pipeline_instance.predict.return_value = np.array([0, 1])
57
- mock_pipeline_instance.predict_proba.return_value = np.array([[0.9, 0.1], [0.4, 0.6]])
 
 
58
  mock_Pipeline.return_value = mock_pipeline_instance
59
 
60
  # Appeler la fonction
@@ -62,17 +131,18 @@ def test_train_and_evaluate_pipeline_success_path(
62
 
63
  # 3. Assertions
64
  mock_get_data.assert_called_once_with(source="postgres")
65
-
66
- # Si vous avez bien enlevé les appels à clean_data et map_binary_features de train_model.py :
67
- mock_clean_data.assert_not_called()
68
- mock_map_binary.assert_not_called()
69
  # Si create_features est toujours appelé :
70
- mock_create_features.assert_called_once_with(sample_df_from_db) # Ou df_for_training avant le copy
71
-
 
 
72
  mock_build_preprocessor.assert_called_once()
73
- mock_LogisticRegression.assert_called_once_with(random_state=42, class_weight='balanced', max_iter=1000)
 
 
74
  # Par exemple, les appels à predict et predict_proba se feront sur X_test qui a 2 lignes
75
- mock_pipeline_instance.predict.assert_called_once()
76
  # On peut vérifier la forme de l'argument avec lequel predict a été appelé
77
  # args_predict, _ = mock_pipeline_instance.predict.call_args
78
  # assert args_predict[0].shape[0] == 2 # X_test doit avoir 2 lignes
@@ -83,18 +153,28 @@ def test_train_and_evaluate_pipeline_success_path(
83
  mock_joblib_dump.assert_called_once_with(mock_pipeline_instance, config.MODEL_PATH)
84
 
85
 
86
- @patch('src.modeling.train_model.get_data', return_value=None) # Simule un échec de chargement
87
- @patch('src.modeling.train_model.logger')
88
- def test_train_and_evaluate_pipeline_get_data_returns_none(mock_logger, mock_get_data_none):
 
 
 
 
89
  """Teste le cas où get_data retourne None."""
90
  train_and_evaluate_pipeline()
91
- mock_logger.error.assert_called_once_with("Arrêt : Impossible de charger les données ou DataFrame vide.")
 
 
92
 
93
 
94
- @patch('src.modeling.train_model.get_data')
95
- @patch('src.modeling.train_model.logger')
96
- def test_train_and_evaluate_pipeline_get_data_returns_empty_df(mock_logger, mock_get_data_empty):
 
 
97
  """Teste le cas où get_data retourne un DataFrame vide."""
98
- mock_get_data_empty.return_value = pd.DataFrame() # df vide
99
- train_and_evaluate_pipeline() # Devrait maintenant appeler logger.error et return
100
- mock_logger.error.assert_called_once_with("Arrêt : Impossible de charger les données ou DataFrame vide.")
 
 
 
1
+ from unittest.mock import patch, MagicMock
 
2
  import pandas as pd
3
  import numpy as np
4
  from src.modeling.train_model import train_and_evaluate_pipeline
5
+ from src import config # Pour les mappings et la cible
6
 
7
+
8
+ @patch("src.modeling.train_model.dump") # joblib.dump
9
+ @patch(
10
+ "src.modeling.train_model.confusion_matrix", return_value=np.array([[0, 0], [0, 0]])
11
+ )
12
+ @patch("src.modeling.train_model.fbeta_score", return_value=0.5)
13
+ @patch("src.modeling.train_model.classification_report", return_value="Mocked Report")
14
+ @patch("src.modeling.train_model.Pipeline") # sklearn.pipeline.Pipeline
15
+ @patch(
16
+ "src.modeling.train_model.LogisticRegression"
17
+ ) # sklearn.linear_model.LogisticRegression
18
+ @patch("src.modeling.train_model.build_preprocessor")
19
+ @patch("src.modeling.train_model.create_features") # de preprocess
20
+ @patch("src.modeling.train_model.get_data") # de load_data
21
+ @patch("src.modeling.train_model.logger")
22
  def test_train_and_evaluate_pipeline_success_path(
23
+ mock_logger,
24
+ mock_get_data,
25
+ mock_create_features,
26
+ mock_build_preprocessor,
27
+ mock_LogisticRegression,
28
+ mock_Pipeline,
29
+ mock_classification_report,
30
+ mock_fbeta_score,
31
+ mock_confusion_matrix,
32
+ mock_joblib_dump,
33
  ):
34
  """Teste le chemin principal de train_and_evaluate_pipeline avec des mocks."""
35
  # 1. Configurer les mocks
36
  # Mock get_data pour retourner des données avec ASSEZ DE LIGNES
37
+ sample_df_from_db = pd.DataFrame(
38
+ {
39
+ "id_employee": [str(i) for i in range(1, 11)], # 10 lignes
40
+ config.TARGET_VARIABLE: [
41
+ 0,
42
+ 1,
43
+ 0,
44
+ 1,
45
+ 0,
46
+ 1,
47
+ 0,
48
+ 1,
49
+ 0,
50
+ 0,
51
+ ], # 4 de classe 1, 6 de classe 0
52
+ "genre": [0, 1, 0, 1, 0, 1, 0, 1, 0, 1], # Alterner pour avoir des valeurs
53
+ "heure_supplementaires": [0, 1, 0, 0, 1, 1, 0, 0, 1, 0],
54
+ "frequence_deplacement": [
55
+ "Non-Travel",
56
+ "Travel_Rarely",
57
+ "Travel_Frequently",
58
+ "Non-Travel",
59
+ "Travel_Rarely",
60
+ "Non-Travel",
61
+ "Travel_Rarely",
62
+ "Travel_Frequently",
63
+ "Non-Travel",
64
+ "Travel_Rarely",
65
+ ],
66
+ "age": [30, 45, 22, 50, 38, 29, 54, 33, 41, 47],
67
+ "salaire_mensuel_brut": [
68
+ 5000,
69
+ 8000,
70
+ 3500,
71
+ 9000,
72
+ 6000,
73
+ 4800,
74
+ 9500,
75
+ 5500,
76
+ 7000,
77
+ 8200.0,
78
+ ],
79
+ "department": [
80
+ "Ventes",
81
+ "R&D",
82
+ "Ventes",
83
+ "Marketing",
84
+ "R&D",
85
+ "Ventes",
86
+ "R&D",
87
+ "Ventes",
88
+ "Marketing",
89
+ "R&D",
90
+ ],
91
+ "augmentation_salaire_precedente": [
92
+ 10.0,
93
+ 15.0,
94
+ 5.0,
95
+ 12.0,
96
+ 8.0,
97
+ 9.0,
98
+ 11.0,
99
+ 6.0,
100
+ 13.0,
101
+ 7.0,
102
+ ],
103
+ }
104
+ )
105
  mock_get_data.return_value = sample_df_from_db
106
 
107
  df_after_features = sample_df_from_db.copy()
108
  mock_create_features.return_value = df_after_features
109
+
110
  mock_preprocessor_instance = MagicMock()
111
+ mock_build_preprocessor.return_value = (
112
+ mock_preprocessor_instance,
113
+ ["mock_col1", "mock_col2"],
114
+ ) # Exemple
115
 
116
  mock_classifier_instance = MagicMock()
117
  mock_LogisticRegression.return_value = mock_classifier_instance
 
120
  # Avec 10 lignes et test_size=0.2, X_test aura 2 lignes (10 * 0.2 = 2)
121
  # X_train aura 8 lignes.
122
  # Les mocks pour predict et predict_proba doivent donc retourner pour 2 échantillons
123
+ mock_pipeline_instance.predict.return_value = np.array([0, 1])
124
+ mock_pipeline_instance.predict_proba.return_value = np.array(
125
+ [[0.9, 0.1], [0.4, 0.6]]
126
+ )
127
  mock_Pipeline.return_value = mock_pipeline_instance
128
 
129
  # Appeler la fonction
 
131
 
132
  # 3. Assertions
133
  mock_get_data.assert_called_once_with(source="postgres")
134
+
 
 
 
135
  # Si create_features est toujours appelé :
136
+ mock_create_features.assert_called_once_with(
137
+ sample_df_from_db
138
+ ) # Ou df_for_training avant le copy
139
+
140
  mock_build_preprocessor.assert_called_once()
141
+ mock_LogisticRegression.assert_called_once_with(
142
+ random_state=42, class_weight="balanced", max_iter=1000
143
+ )
144
  # Par exemple, les appels à predict et predict_proba se feront sur X_test qui a 2 lignes
145
+ mock_pipeline_instance.predict.assert_called_once()
146
  # On peut vérifier la forme de l'argument avec lequel predict a été appelé
147
  # args_predict, _ = mock_pipeline_instance.predict.call_args
148
  # assert args_predict[0].shape[0] == 2 # X_test doit avoir 2 lignes
 
153
  mock_joblib_dump.assert_called_once_with(mock_pipeline_instance, config.MODEL_PATH)
154
 
155
 
156
+ @patch(
157
+ "src.modeling.train_model.get_data", return_value=None
158
+ ) # Simule un échec de chargement
159
+ @patch("src.modeling.train_model.logger")
160
+ def test_train_and_evaluate_pipeline_get_data_returns_none(
161
+ mock_logger, mock_get_data_none
162
+ ):
163
  """Teste le cas où get_data retourne None."""
164
  train_and_evaluate_pipeline()
165
+ mock_logger.error.assert_called_once_with(
166
+ "Arrêt : Impossible de charger les données ou DataFrame vide depuis la source."
167
+ )
168
 
169
 
170
+ @patch("src.modeling.train_model.get_data")
171
+ @patch("src.modeling.train_model.logger")
172
+ def test_train_and_evaluate_pipeline_get_data_returns_empty_df(
173
+ mock_logger, mock_get_data_empty
174
+ ):
175
  """Teste le cas où get_data retourne un DataFrame vide."""
176
+ mock_get_data_empty.return_value = pd.DataFrame() # df vide
177
+ train_and_evaluate_pipeline() # Devrait maintenant appeler logger.error et return
178
+ mock_logger.error.assert_called_once_with(
179
+ "Arrêt : Impossible de charger les données ou DataFrame vide depuis la source."
180
+ )