Employee-Churn-Prediction / src /db /create_db.py
Alexis-Ravet's picture
Upload folder using huggingface_hub
b666236 verified
Raw
History Blame Contribute Delete
5.31 kB
"""
Script de création de la base de données PostgreSQL et d'insertion du dataset.
"""
from pathlib import Path
import pandas as pd
from sqlalchemy.exc import IntegrityError
from src.db.database import SessionLocal, engine
from src.db.models import Base, Employee
def racine_projet() -> Path:
"""Remonte l'arborescence jusqu'à trouver la racine du projet."""
courant = Path(__file__).resolve()
for parent in courant.parents:
if (parent / "pyproject.toml").exists():
return parent
raise FileNotFoundError("Racine du projet non trouvée")
PROJECT_ROOT = racine_projet()
RAW_DIR = PROJECT_ROOT / "data" / "raw"
PROCESSED_DIR = PROJECT_ROOT / "data" / "processed"
CSV_SIRH = RAW_DIR / "extrait_sirh.csv"
CSV_EVAL = RAW_DIR / "extrait_eval.csv"
CSV_SONDAGE = RAW_DIR / "extrait_sondage.csv"
CSV_EMPLOYES = PROCESSED_DIR / "employees.csv"
COLONNES_EMPLOYEES = [
"id_employee",
"age",
"genre",
"revenu_mensuel",
"statut_marital",
"departement",
"poste",
"annee_experience_totale",
"annees_dans_l_entreprise",
"satisfaction_employee_environnement",
"note_evaluation_precedente",
"satisfaction_employee_nature_travail",
"satisfaction_employee_equipe",
"satisfaction_employee_equilibre_pro_perso",
"note_evaluation_actuelle",
"heure_supplementaires",
"augementation_salaire_precedente",
"nombre_participation_pee",
"nb_formations_suivies",
"distance_domicile_travail",
"niveau_education",
"frequence_deplacement",
"annees_depuis_la_derniere_promotion",
"a_quitte_l_entreprise",
]
def fusionner_csv() -> pd.DataFrame:
"""
Charge les 3 CSV et les fusionne sur la colonne id_employee.
Reproduit les étapes du notebook :
1. Renomme eval_number en id_employee dans df_eval
2. Retire le préfixe 'E_' et convertit en int
3. Renomme code_sondage en id_employee dans df_sondage
4. Inner merge des 3 DataFrames
Returns:
DataFrame fusionné (32 colonnes, 1470 lignes)
"""
df_sirh = pd.read_csv(CSV_SIRH)
df_eval = pd.read_csv(CSV_EVAL)
df_sondage = pd.read_csv(CSV_SONDAGE)
df_eval = df_eval.rename(columns={"eval_number": "id_employee"})
df_eval["id_employee"] = df_eval["id_employee"].str[2:].astype("int64")
df_sondage = df_sondage.rename(columns={"code_sondage": "id_employee"})
df_central = pd.merge(df_sirh, df_eval, on="id_employee", how="inner")
df_central = pd.merge(df_central, df_sondage, on="id_employee", how="inner")
print(
f"DataFrame fusionné : {df_central.shape[0]} lignes, {df_central.shape[1]} colonnes"
)
return df_central
def nettoyer_dataframe(df_central: pd.DataFrame) -> pd.DataFrame:
"""
Nettoie le DataFrame fusionné pour ne garder que les colonnes utilisées par le modèle.
Étapes :
1. Retire le '%' de augementation_salaire_precedente et convertit en int
2. Sélectionne uniquement les 23 colonnes du modèle
Args:
df_central: DataFrame fusionné (32 colonnes)
Returns:
DataFrame nettoyé (23 colonnes)
"""
df_central["augementation_salaire_precedente"] = (
df_central["augementation_salaire_precedente"].str[:-2].astype("int64")
)
df_employees = df_central.loc[:, COLONNES_EMPLOYEES].copy()
print(
f"DataFrame nettoyé : {df_employees.shape[0]} lignes, {df_employees.shape[1]} colonnes"
)
return df_employees
def sauvegarder_csv(df_employees: pd.DataFrame) -> None:
"""
Sauvegarde le DataFrame nettoyé au format CSV.
Args:
df_employees: DataFrame nettoyé (23 colonnes)
"""
PROCESSED_DIR.mkdir(parents=True, exist_ok=True)
df_employees.to_csv(CSV_EMPLOYES, index=False)
print(f"CSV sauvegardé : {CSV_EMPLOYES}")
def creer_tables() -> None:
"""Crée les tables dans PostgreSQL à partir des modèles ORM."""
Base.metadata.create_all(engine)
print("Tables créées avec succès")
def inserer_donnees(df_employees: pd.DataFrame) -> None:
"""
Insère les données du DataFrame dans la table employees.
Args:
df_employees: DataFrame nettoyé (23 colonnes, 1470 lignes)
"""
session = SessionLocal()
try:
employees = [Employee(**row.to_dict()) for _, row in df_employees.iterrows()]
session.add_all(employees)
session.commit()
print(f"{len(employees)} employés insérés avec succès")
except IntegrityError as e:
session.rollback()
print(f"Erreur d'intégrité : {e}")
print("Les données existent peut-être déjà. Vider la table avant de réinsérer.")
except Exception as e:
session.rollback()
print(f"Erreur lors de l'insertion : {e}")
finally:
session.close()
def main() -> None:
"""Point d'entrée principal du script."""
# 1. Chargement et fusion des CSV
df_central = fusionner_csv()
# 2. Nettoyage du DataFrame"
df_employees = nettoyer_dataframe(df_central)
# 3. Sauvegarde du CSV nettoyé
sauvegarder_csv(df_employees)
# 4. Création des tables
creer_tables()
# 5. Insertion des données
inserer_donnees(df_employees)
print("Base de données créée avec succès")
if __name__ == "__main__":
main()