Spaces:
Running
Running
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 +0 -0
- src/config.py +5 -5
- src/data_processing/load_data.py +169 -73
- src/data_processing/preprocess.py +277 -101
- src/modeling/predict.py +137 -77
- src/modeling/train_model.py +121 -71
- tests/unit/test_load_data.py +181 -118
- tests/unit/test_predict.py +78 -46
- tests/unit/test_preprocess.py +113 -73
- tests/unit/test_train_model.py +133 -53
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 =
|
| 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 |
-
|
| 36 |
-
|
| 37 |
}
|
| 38 |
|
| 39 |
ORDINAL_FEATURES_CATEGORIES = {
|
| 40 |
-
|
| 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 |
-
|
| 3 |
-
|
| 4 |
-
from src.database.
|
|
|
|
| 5 |
from src import config
|
| 6 |
import logging
|
| 7 |
|
| 8 |
-
|
|
|
|
|
|
|
|
|
|
| 9 |
logger = logging.getLogger(__name__)
|
| 10 |
|
| 11 |
-
|
| 12 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
"""
|
| 24 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 25 |
try:
|
| 26 |
-
logger.info(
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
#
|
| 30 |
-
|
| 31 |
-
|
|
|
|
| 32 |
if df.empty:
|
| 33 |
-
logger.warning(
|
|
|
|
|
|
|
| 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(
|
| 39 |
-
|
|
|
|
|
|
|
|
|
|
| 40 |
finally:
|
| 41 |
db.close()
|
|
|
|
| 42 |
|
| 43 |
-
|
|
|
|
| 44 |
"""
|
| 45 |
-
Fonction principale pour
|
| 46 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
#
|
| 57 |
-
#
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
|
|
|
|
|
|
| 61 |
else:
|
| 62 |
-
logger.error(
|
|
|
|
|
|
|
| 63 |
return load_data_from_postgres()
|
| 64 |
|
| 65 |
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
df_eval[
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
|
|
|
| 86 |
nan_count = numeric_ids_for_check.isnull().sum()
|
| 87 |
if nan_count > 0:
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
df_eval[
|
| 93 |
-
df_eval = df_eval.drop(columns=[
|
| 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 |
-
#
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
return None
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 |
-
#
|
| 134 |
-
#
|
| 135 |
-
#
|
| 136 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
"""
|
| 18 |
-
|
| 19 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
for col, mapping in binary_cols_map.items():
|
| 21 |
if col in df.columns:
|
| 22 |
-
#
|
|
|
|
|
|
|
| 23 |
df[col] = df[col].astype(str).map(mapping)
|
| 24 |
-
|
| 25 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
nan_count = df[col].isnull().sum()
|
| 27 |
if nan_count > 0:
|
| 28 |
logger.warning(
|
| 29 |
-
f"{nan_count}
|
|
|
|
| 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 |
-
"""
|
| 40 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 48 |
else:
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 52 |
cols_to_drop = [
|
| 53 |
-
"
|
| 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 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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(
|
| 83 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
"""
|
| 95 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
df = df.copy()
|
| 97 |
-
|
| 98 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 107 |
) -> ColumnTransformer:
|
| 108 |
-
"""
|
| 109 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 138 |
-
raise ValueError(
|
| 139 |
-
|
|
|
|
|
|
|
| 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=
|
|
|
|
| 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.
|
|
|
|
| 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 |
-
|
| 168 |
preprocessor: ColumnTransformer = None,
|
| 169 |
fit: bool = False,
|
| 170 |
):
|
| 171 |
-
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 172 |
df_clean = clean_data(df)
|
| 173 |
|
| 174 |
-
|
| 175 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 176 |
|
| 177 |
-
df_featured = create_features(
|
| 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 |
-
|
| 188 |
-
|
| 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 |
-
|
|
|
|
|
|
|
| 207 |
|
| 208 |
-
logger.info(f"Colonnes
|
| 209 |
-
logger.info(f"Colonnes
|
| 210 |
-
logger.info(f"Colonnes
|
| 211 |
|
| 212 |
if fit:
|
|
|
|
| 213 |
processor_instance = build_preprocessor(
|
| 214 |
numerical_to_scale,
|
| 215 |
onehot_to_encode,
|
| 216 |
ordinal_to_encode,
|
| 217 |
-
|
| 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 |
-
#
|
| 252 |
if __name__ == "__main__":
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 260 |
X_p, y_p, proc = run_preprocessing_pipeline(
|
| 261 |
df_raw,
|
| 262 |
-
binary_cols_map
|
| 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 |
-
|
| 272 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
| 282 |
-
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 5 |
-
from typing import List, Dict, Union
|
| 6 |
|
| 7 |
-
|
|
|
|
| 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 |
-
|
|
|
|
| 20 |
|
| 21 |
-
# --- FIN DÉFINITION ---
|
| 22 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 23 |
|
| 24 |
-
|
| 25 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 47 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 48 |
"""
|
| 49 |
pipeline = load_prediction_pipeline()
|
| 50 |
if pipeline is None:
|
| 51 |
-
|
|
|
|
| 52 |
|
| 53 |
try:
|
| 54 |
-
logger.info(f"
|
|
|
|
|
|
|
|
|
|
| 55 |
|
| 56 |
# --- ÉTAPE 1 : Appliquer les transformations manuelles ---
|
| 57 |
-
#
|
| 58 |
-
#
|
| 59 |
-
#
|
| 60 |
-
#
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 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 |
-
#
|
| 72 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 73 |
|
| 74 |
-
logger.info("Transformations manuelles appliquées.")
|
| 75 |
|
| 76 |
# --- ÉTAPE 2 : Utiliser la pipeline complète pour prédire ---
|
| 77 |
-
|
| 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,
|
| 84 |
results.append(
|
| 85 |
{
|
| 86 |
-
"id_employe":
|
| 87 |
"probabilite_depart": float(probabilities[i]),
|
| 88 |
"prediction_depart": int(predictions[i]),
|
| 89 |
}
|
| 90 |
)
|
| 91 |
-
logger.info("
|
| 92 |
return results
|
| 93 |
|
| 94 |
except Exception as e:
|
| 95 |
-
logger.error(f"Erreur lors
|
| 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 |
-
|
| 102 |
-
#
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 112 |
-
|
| 113 |
-
|
| 114 |
-
|
| 115 |
-
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
# Vérifier si le modèle existe avant de tester
|
| 134 |
if config.MODEL_PATH.exists():
|
| 135 |
-
|
|
|
|
| 136 |
print("\n--- Résultat de la Prédiction Test ---")
|
| 137 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 18 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
"""
|
| 35 |
-
logger.info(">>> Début du processus d'entraînement et d'évaluation <<<")
|
| 36 |
|
| 37 |
# --- 1. Charger les données ---
|
| 38 |
-
|
| 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 |
-
# ---
|
| 45 |
-
#
|
| 46 |
-
#
|
| 47 |
-
|
|
|
|
| 48 |
|
| 49 |
-
if df_featured.empty:
|
| 50 |
-
logger.error("Arrêt : DataFrame vide après
|
| 51 |
return
|
| 52 |
|
| 53 |
-
df_for_training = df_featured.copy()
|
| 54 |
|
| 55 |
-
|
| 56 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
|
| 58 |
y = df_for_training[config.TARGET_VARIABLE]
|
| 59 |
-
|
| 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 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
# Ajoutez
|
| 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 |
-
|
| 76 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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,
|
| 83 |
)
|
| 84 |
-
logger.info(f"
|
|
|
|
| 85 |
|
| 86 |
-
# ---
|
| 87 |
-
|
| 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 |
-
|
| 97 |
-
|
| 98 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
|
| 100 |
-
# ---
|
|
|
|
| 101 |
preprocessor = build_preprocessor(
|
| 102 |
numerical_cols=numerical_to_scale,
|
| 103 |
onehot_cols=onehot_to_encode,
|
| 104 |
ordinal_cols=ordinal_to_encode,
|
| 105 |
-
|
| 106 |
)
|
| 107 |
|
| 108 |
-
# ---
|
| 109 |
-
|
| 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 |
-
# ---
|
|
|
|
| 118 |
full_pipeline = Pipeline(
|
| 119 |
steps=[("preprocessor", preprocessor), ("classifier", classifier)]
|
| 120 |
)
|
| 121 |
logger.info("Pipeline complète créée.")
|
| 122 |
|
| 123 |
-
# ---
|
| 124 |
-
|
| 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 |
-
# ---
|
| 131 |
-
logger.info("
|
| 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 |
-
|
| 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
|
|
|
|
| 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(
|
|
|
|
| 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({
|
| 20 |
mock_read_csv.return_value = sample_df
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
mock_logger.info.assert_any_call("
|
|
|
|
| 27 |
pd.testing.assert_frame_equal(df, sample_df)
|
| 28 |
|
| 29 |
-
|
| 30 |
-
@patch(
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 39 |
|
| 40 |
-
|
| 41 |
-
@patch(
|
| 42 |
-
@patch(
|
|
|
|
| 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({
|
| 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 =
|
| 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()
|
| 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(
|
| 66 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 67 |
pd.testing.assert_frame_equal(df, sample_df)
|
| 68 |
-
mock_db_session.close.assert_called_once()
|
|
|
|
| 69 |
|
| 70 |
-
@patch(
|
| 71 |
-
@patch(
|
| 72 |
-
@patch(
|
| 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(
|
|
|
|
|
|
|
| 84 |
mock_db_session.close.assert_called_once()
|
| 85 |
|
| 86 |
-
|
| 87 |
-
@patch(
|
| 88 |
-
@patch(
|
| 89 |
-
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 105 |
-
@patch(
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 116 |
-
@patch(
|
|
|
|
| 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 |
-
|
| 126 |
-
@patch(
|
|
|
|
| 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()
|
| 132 |
-
|
|
|
|
|
|
|
| 133 |
assert result == "data_from_db_default"
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
@patch(
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 142 |
-
|
| 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 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 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) == [
|
| 180 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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[
|
| 185 |
-
assert row_1[
|
| 186 |
-
assert row_1[
|
| 187 |
-
assert row_1[
|
| 188 |
|
| 189 |
# Ligne pour id_employee '2'
|
| 190 |
-
row_2 = merged_df[merged_df[
|
| 191 |
-
assert row_2[
|
| 192 |
-
assert row_2[
|
| 193 |
-
assert pd.isna(row_2[
|
| 194 |
|
| 195 |
# Ligne pour id_employee '3'
|
| 196 |
-
row_3 = merged_df[merged_df[
|
| 197 |
-
assert row_3[
|
| 198 |
-
assert pd.isna(
|
| 199 |
-
|
| 200 |
-
|
|
|
|
|
|
|
| 201 |
# Vérifier le type de la colonne 'id_employee'
|
| 202 |
-
assert merged_df[
|
| 203 |
|
| 204 |
|
| 205 |
-
@patch(
|
| 206 |
-
@patch(
|
| 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(
|
|
|
|
|
|
|
| 215 |
|
| 216 |
-
|
| 217 |
-
@patch(
|
|
|
|
| 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({
|
| 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(
|
|
|
|
|
|
|
| 227 |
|
| 228 |
|
| 229 |
-
@patch(
|
| 230 |
-
@patch(
|
| 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({
|
| 234 |
-
df_eval_no_eval_number = pd.DataFrame({
|
| 235 |
-
df_sondage_sample = pd.DataFrame({
|
| 236 |
-
mock_read_csv.side_effect = [
|
| 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(
|
| 245 |
-
@patch(
|
| 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({
|
| 249 |
-
df_eval_sample = pd.DataFrame({
|
| 250 |
-
df_sondage_sample = pd.DataFrame({
|
| 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(
|
|
|
|
|
|
|
| 257 |
|
| 258 |
|
| 259 |
-
@patch(
|
| 260 |
-
@patch(
|
| 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({
|
| 264 |
-
df_eval_sample = pd.DataFrame(
|
| 265 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 7 |
|
| 8 |
-
|
| 9 |
-
@patch(
|
| 10 |
-
@patch(
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 26 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
assert pipeline == mock_pipeline_obj
|
| 28 |
|
| 29 |
-
|
| 30 |
-
@patch(
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 42 |
assert pipeline is None
|
| 43 |
|
| 44 |
-
|
| 45 |
-
@patch(
|
| 46 |
-
@patch(
|
| 47 |
-
|
|
|
|
|
|
|
|
|
|
| 48 |
"""Teste une exception lors du chargement du modèle."""
|
| 49 |
mock_model_path.exists.return_value = True
|
| 50 |
-
|
|
|
|
|
|
|
| 51 |
|
| 52 |
from src.modeling import predict
|
|
|
|
| 53 |
predict._pipeline = None
|
| 54 |
|
| 55 |
pipeline = load_prediction_pipeline()
|
| 56 |
|
| 57 |
-
|
|
|
|
| 58 |
assert pipeline is None
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
@patch(
|
| 62 |
-
@patch(
|
| 63 |
-
@patch(
|
| 64 |
-
@patch(
|
|
|
|
| 65 |
def test_predict_attrition_success_path(
|
| 66 |
-
mock_logger,
|
| 67 |
-
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 74 |
-
|
|
|
|
|
|
|
| 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({
|
| 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
|
| 85 |
-
df_cleaned_for_predict[
|
| 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(
|
|
|
|
|
|
|
| 90 |
mock_map_binary.return_value = df_mapped_mock
|
| 91 |
-
|
| 92 |
-
df_featured_mock = pd.DataFrame(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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=
|
| 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
|
| 116 |
|
| 117 |
assert len(results) == 1
|
| 118 |
-
assert results[0][
|
| 119 |
-
assert
|
| 120 |
-
|
| 121 |
-
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 66 |
-
|
| 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 = [
|
| 104 |
-
onehot_cols = [
|
| 105 |
-
ordinal_cols = [
|
| 106 |
-
ordinal_categories = {
|
| 107 |
-
'frequence_deplacement': ['Bas', 'Moyen', 'Haut']
|
| 108 |
-
}
|
| 109 |
|
| 110 |
-
preprocessor = build_preprocessor(
|
|
|
|
|
|
|
| 111 |
|
| 112 |
assert isinstance(preprocessor, ColumnTransformer)
|
| 113 |
-
assert
|
|
|
|
|
|
|
| 114 |
|
| 115 |
# Vérifier le transformateur numérique
|
| 116 |
-
num_transformer_tuple = next(t for t in preprocessor.transformers if t[0] ==
|
| 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 ==
|
| 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(
|
|
|
|
|
|
|
| 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 ==
|
| 130 |
assert isinstance(onehot_transformer_tuple[1].steps[1][1], OneHotEncoder)
|
| 131 |
-
assert onehot_transformer_tuple[1].steps[1][1].drop ==
|
| 132 |
-
assert onehot_transformer_tuple[1].steps[1][1].handle_unknown ==
|
| 133 |
assert onehot_transformer_tuple[2] == onehot_cols
|
| 134 |
-
|
| 135 |
# Vérifier le transformateur ordinal pour 'frequence_deplacement'
|
| 136 |
-
ord_transformer_tuple = next(
|
|
|
|
|
|
|
| 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 == [[
|
| 142 |
-
assert ord_transformer_tuple[2] == [
|
| 143 |
|
| 144 |
|
| 145 |
def test_build_preprocessor_only_numerical():
|
| 146 |
"""Teste build_preprocessor avec seulement des colonnes numériques."""
|
| 147 |
-
numerical_cols = [
|
| 148 |
preprocessor = build_preprocessor(numerical_cols, [], [], {})
|
| 149 |
assert isinstance(preprocessor, ColumnTransformer)
|
| 150 |
assert len(preprocessor.transformers) == 1
|
| 151 |
-
assert preprocessor.transformers[0][0] ==
|
| 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 = [
|
| 159 |
-
|
| 160 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 167 |
-
|
| 168 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
# Ajoutez d'autres colonnes pour couvrir tous vos types
|
| 187 |
-
config.TARGET_VARIABLE: [
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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(
|
| 212 |
-
assert
|
| 213 |
-
assert
|
| 214 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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_[
|
| 219 |
-
scaler = num_pipeline.named_steps[
|
| 220 |
-
assert hasattr(scaler,
|
|
|
|
| 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 |
-
|
| 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,
|
| 238 |
-
|
| 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 |
-
|
| 249 |
-
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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
|
| 7 |
|
| 8 |
-
|
| 9 |
-
@patch(
|
| 10 |
-
@patch(
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
@patch(
|
| 14 |
-
@patch(
|
| 15 |
-
@patch(
|
| 16 |
-
@patch(
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
@patch(
|
|
|
|
|
|
|
|
|
|
| 20 |
def test_train_and_evaluate_pipeline_success_path(
|
| 21 |
-
mock_logger,
|
| 22 |
-
|
| 23 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 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 = (
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
| 71 |
-
|
|
|
|
|
|
|
| 72 |
mock_build_preprocessor.assert_called_once()
|
| 73 |
-
mock_LogisticRegression.assert_called_once_with(
|
|
|
|
|
|
|
| 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(
|
| 87 |
-
|
| 88 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 89 |
"""Teste le cas où get_data retourne None."""
|
| 90 |
train_and_evaluate_pipeline()
|
| 91 |
-
mock_logger.error.assert_called_once_with(
|
|
|
|
|
|
|
| 92 |
|
| 93 |
|
| 94 |
-
@patch(
|
| 95 |
-
@patch(
|
| 96 |
-
def test_train_and_evaluate_pipeline_get_data_returns_empty_df(
|
|
|
|
|
|
|
| 97 |
"""Teste le cas où get_data retourne un DataFrame vide."""
|
| 98 |
-
mock_get_data_empty.return_value = pd.DataFrame()
|
| 99 |
-
train_and_evaluate_pipeline()
|
| 100 |
-
mock_logger.error.assert_called_once_with(
|
|
|
|
|
|
|
|
|
| 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 |
+
)
|