"""
Module: Régression logistique
Catégorie : Apprentissage supervisé - Classification
Difficulté : Débutant

Généré depuis la plateforme ML Formation
"""

# Imports
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, mean_squared_error, r2_score

# Charger le dataset
df = pd.read_csv('binary_classification.csv')

# Explorer les données
# Type: Code exécutable
# =============================================================================
# ETAPE 1 : EXPLORATION DU DATASET DE CLASSIFICATION
# =============================================================================
# En classification, l'exploration des donnees est cruciale pour comprendre :
# 1) La distribution des classes (equilibre ou desequilibre ?)
# 2) La separabilite des classes (les features permettent-elles de distinguer ?)
# 3) Les caracteristiques de chaque classe

print("=" * 70)
print("EXPLORATION DU DATASET DE CLASSIFICATION BINAIRE")
print("=" * 70)
print()

# --- 1.1 Apercu des donnees ---
print("1. APERCU DES DONNEES")
print("-" * 40)
print("Chaque ligne represente un echantillon avec 2 caracteristiques (features)")
print("et une etiquette (label) indiquant sa classe (0 ou 1).")
print()
display(df.head(10), title="Apercu du dataset de classification")

# --- 1.2 Dimensions ---
n_samples, n_cols = df.shape
print()
print("2. DIMENSIONS DU DATASET")
print("-" * 40)
print(f"   Nombre d'echantillons : {n_samples}")
print(f"   Nombre de colonnes    : {n_cols} (2 features + 1 label)")
print()

# --- 1.3 Distribution des classes ---
print("3. DISTRIBUTION DES CLASSES")
print("-" * 40)
class_counts = df['label'].value_counts().sort_index()
total = len(df)

print("   Repartition des echantillons par classe :")
print()
for classe, count in class_counts.items():
    pct = count / total * 100
    bar = "█" * int(pct / 2)
    print(f"   Classe {classe} : {count:4d} echantillons ({pct:5.1f}%) {bar}")

print()

# Analyse de l'equilibre
ratio = class_counts.min() / class_counts.max()
if ratio > 0.8:
    equilibre = "EQUILIBRE"
    emoji = "✓"
    conseil = "Excellent ! Les metriques standard (accuracy) seront fiables."
elif ratio > 0.5:
    equilibre = "LEGEREMENT DESEQUILIBRE"
    emoji = "~"
    conseil = "Acceptable, mais surveillez precision et recall par classe."
else:
    equilibre = "DESEQUILIBRE"
    emoji = "⚠"
    conseil = "Attention ! L'accuracy peut etre trompeuse. Utilisez F1-score."

print(f"   DIAGNOSTIC : {equilibre} {emoji}")
print(f"   Ratio minoritaire/majoritaire : {ratio:.2f}")
print(f"   → {conseil}")
print()

# --- 1.4 Statistiques descriptives ---
print("4. STATISTIQUES DESCRIPTIVES")
print("-" * 40)
print("Voici les statistiques cles de chaque colonne :")
print()
display(df.describe().round(3), title="Statistiques descriptives")

# --- 1.5 Analyse de separabilite ---
print()
print("5. ANALYSE DE SEPARABILITE")
print("-" * 40)
print("   La separabilite mesure si les classes sont distinguables par les features.")
print()

for feature in ['feature1', 'feature2']:
    mean_0 = df[df['label'] == 0][feature].mean()
    mean_1 = df[df['label'] == 1][feature].mean()
    std_0 = df[df['label'] == 0][feature].std()
    std_1 = df[df['label'] == 1][feature].std()

    # Calcul d'une metrique de separabilite simplifiee
    pooled_std = np.sqrt((std_0**2 + std_1**2) / 2)
    if pooled_std > 0:
        separabilite = abs(mean_1 - mean_0) / pooled_std
    else:
        separabilite = 0

    print(f"   {feature}:")
    print(f"      Classe 0 : moyenne = {mean_0:.3f}, std = {std_0:.3f}")
    print(f"      Classe 1 : moyenne = {mean_1:.3f}, std = {std_1:.3f}")
    print(f"      Difference des moyennes : {abs(mean_1 - mean_0):.3f}")

    if separabilite > 1.5:
        verdict = "BONNE separabilite"
    elif separabilite > 0.8:
        verdict = "Separabilite MODEREE"
    else:
        verdict = "Separabilite FAIBLE"

    print(f"      → {verdict} (score: {separabilite:.2f})")
    print()

# --- 1.6 Conclusion ---
print("=" * 70)
print("CONCLUSION DE L'EXPLORATION")
print("=" * 70)
print()
print("   Ce que nous avons appris :")
print("   → Les deux classes sont-elles equilibrees ?")
print("   → Les features permettent-elles de distinguer les classes ?")
print("   → Y a-t-il des patterns visibles dans les statistiques ?")
print()
print("   Prochaine etape : Visualiser les donnees pour confirmer ces observations.")


# Visualiser les classes
# Type: Code exécutable
# =============================================================================
# ETAPE 2 : VISUALISATION DES CLASSES
# =============================================================================
# La visualisation est essentielle en classification pour :
# 1) Voir si les classes sont naturellement separables
# 2) Identifier la forme de la frontiere de decision necessaire
# 3) Detecter des outliers ou des chevauchements problematiques

print("=" * 70)
print("VISUALISATION DES DEUX CLASSES")
print("=" * 70)
print()
print("Question cle : Les deux classes sont-elles visuellement separables ?")
print("→ Si oui, la regression logistique (frontiere lineaire) conviendra.")
print("→ Si les classes sont melangees, un modele plus complexe sera necessaire.")
print()

# Separation des classes
class_0 = df[df['label'] == 0]
class_1 = df[df['label'] == 1]

# Creation du graphique
plt.figure(figsize=(10, 7))

plt.scatter(class_0['feature1'], class_0['feature2'],
            alpha=0.6, color='#9B7AC4', label='Classe 0', s=60,
            edgecolors='white', linewidth=0.5)
plt.scatter(class_1['feature1'], class_1['feature2'],
            alpha=0.6, color='#F7E64D', label='Classe 1', s=60,
            edgecolors='white', linewidth=0.5)

plt.xlabel('Feature 1', fontsize=12)
plt.ylabel('Feature 2', fontsize=12)
plt.title('Distribution des deux classes dans l\'espace des features', fontsize=14)
plt.legend(loc='best', fontsize=11)
plt.grid(True, alpha=0.3)

# Ajouter les centroïdes
c0_center = (class_0['feature1'].mean(), class_0['feature2'].mean())
c1_center = (class_1['feature1'].mean(), class_1['feature2'].mean())
plt.plot(*c0_center, 'o', color='#9B7AC4', markersize=15, markeredgecolor='black', markeredgewidth=2)
plt.plot(*c1_center, 'o', color='#F7E64D', markersize=15, markeredgecolor='black', markeredgewidth=2)

plt.tight_layout()
plt.show()

# --- Analyse visuelle ---
print()
print("ANALYSE VISUELLE")
print("-" * 40)
print()
print("   LEGENDE :")
print("   → Points violets : Classe 0")
print("   → Points jaunes  : Classe 1")
print("   → Grands cercles : Centres (moyennes) de chaque classe")
print()

# Calcul de la distance entre centroïdes
distance = np.sqrt((c1_center[0] - c0_center[0])**2 + (c1_center[1] - c0_center[1])**2)
print(f"   DISTANCE ENTRE LES CENTRES : {distance:.2f}")
print()

if distance > 2:
    print("   OBSERVATION : Les classes sont BIEN SEPAREES")
    print("   → Une frontiere lineaire devrait bien fonctionner.")
    print("   → La regression logistique est appropriee.")
elif distance > 1:
    print("   OBSERVATION : Les classes ont un CHEVAUCHEMENT PARTIEL")
    print("   → Quelques erreurs de classification sont attendues.")
    print("   → La regression logistique reste un bon choix de base.")
else:
    print("   OBSERVATION : Les classes sont FORTEMENT MELANGEES")
    print("   → La classification sera difficile.")
    print("   → Envisagez des modeles non-lineaires (SVM, arbres).")

print()
print("   QUESTION A SE POSER :")
print("   'Puis-je tracer une droite qui separe globalement les deux couleurs ?'")
print("   → Si oui, la regression logistique conviendra !")


# ----------------------------------------------------------------------
# La suite de ce module demande un compte.
# Les cellules de code restantes ne sont pas dans ce fichier.
# ----------------------------------------------------------------------
