"""
Module: CNN (réseaux convolutifs)
Catégorie : Deep learning
Difficulté : Intermédiaire

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('digits_simple.csv')

# Explorer les données (chiffres)
# Type: Code exécutable
print("=" * 70)
print("       EXPLORATION DU DATASET DE CHIFFRES MANUSCRITS")
print("=" * 70)

print("""
Les CNN sont nes pour traiter des IMAGES. Nous allons travailler avec
un dataset classique: des chiffres manuscrits (version simplifiee).

POURQUOI CE DATASET EST PARFAIT POUR APPRENDRE LES CNN:
  • Images petites (4x4 = 16 pixels) → rapide a traiter
  • 10 classes (0-9) → classification multiclasse
  • Patterns visuels distincts → facile a interpreter
""")

print("\n" + "=" * 70)
print("1. APERCU DU DATASET")
print("=" * 70)
display(df.head(10), title="Premieres lignes du dataset de chiffres")

print("""
STRUCTURE DES DONNEES:
  • Chaque ligne = une image de chiffre
  • Colonnes pixel_* = valeurs des pixels (0=noir, 15=blanc)
  • Colonne 'label' = le chiffre represente (0-9)
""")

print("\n" + "=" * 70)
print("2. DIMENSIONS")
print("=" * 70)

# Separer features et labels
X = df.drop('label', axis=1).values
y = df['label'].values

n_samples, n_pixels = X.shape
img_size = int(np.sqrt(n_pixels))

print(f"""
DIMENSIONS DU DATASET:
  • Nombre d'images: {n_samples}
  • Pixels par image: {n_pixels}
  • Taille d'une image: {img_size}x{img_size} pixels

NOTE: En production, les images sont plus grandes:
  • MNIST: 28x28 = 784 pixels
  • CIFAR-10: 32x32x3 = 3 072 valeurs
  • ImageNet: 224x224x3 = 150 528 valeurs!
""")

print("\n" + "=" * 70)
print("3. DISTRIBUTION DES CLASSES")
print("=" * 70)

print("""
Combien d'exemples avons-nous pour chaque chiffre?
""")

unique, counts = np.unique(y, return_counts=True)
print("-" * 45)
print(f"{'Chiffre':<10} {'Nb images':<12} {'Distribution'}")
print("-" * 45)

for digit, count in zip(unique, counts):
    pct = count / len(y) * 100
    bar = "█" * int(pct / 2)
    print(f"{digit:<10} {count:<12} {pct:5.1f}% {bar}")

print("-" * 45)
print(f"{'TOTAL':<10} {len(y):<12} 100.0%")

# Verifier l'equilibre
balance_ratio = max(counts) / min(counts)
print(f"""

ANALYSE DE L'EQUILIBRE:
  • Classe la plus frequente: {max(counts)} images
  • Classe la moins frequente: {min(counts)} images
  • Ratio: {balance_ratio:.2f}
""")

if balance_ratio < 1.5:
    print("  → Dataset bien equilibre! Toutes les classes sont representees.")
else:
    print("  → Desequilibre detecte. Certaines classes sont sous-representees.")

print("\n" + "=" * 70)
print("4. STATISTIQUES DES PIXELS")
print("=" * 70)

print(f"""
VALEURS DES PIXELS:
  • Minimum: {X.min():.0f} (noir complet)
  • Maximum: {X.max():.0f} (blanc complet)
  • Moyenne: {X.mean():.2f}
  • Ecart-type: {X.std():.2f}

INTERPRETATION:
  Les valeurs sont entre 0 (noir) et ~16 (blanc).
  Les pixels forment des MOTIFS que le CNN va apprendre a reconnaitre.
""")

print("\n" + "=" * 70)
print("              PRET POUR LA VISUALISATION DES CHIFFRES!")
print("=" * 70)


# Visualiser les chiffres
# Type: Code exécutable
print("=" * 70)
print("       VISUALISATION DES CHIFFRES MANUSCRITS")
print("=" * 70)

print("""
Visualisons les donnees pour comprendre ce que le CNN doit apprendre.
Chaque chiffre a des caracteristiques visuelles distinctes.
""")

# Separer features et labels
X = df.drop('label', axis=1).values
y = df['label'].values
img_size = int(np.sqrt(X.shape[1]))

print("\n" + "=" * 70)
print("1. UN EXEMPLE DE CHAQUE CHIFFRE (0-9)")
print("=" * 70)

fig, axes = plt.subplots(2, 5, figsize=(14, 6))

for i, ax in enumerate(axes.flat):
    # Trouver un exemple du chiffre i
    idx = np.where(y == i)[0][0]
    image = X[idx].reshape(img_size, img_size)

    ax.imshow(image, cmap='gray', interpolation='nearest')
    ax.set_title(f'Chiffre: {i}', fontsize=12, fontweight='bold')
    ax.axis('off')

    # Ajouter une bordure coloree
    for spine in ax.spines.values():
        spine.set_edgecolor('#9B7AC4')
        spine.set_linewidth(2)

plt.suptitle('Exemples de chiffres manuscrits (4x4 pixels)', fontsize=14, fontweight='bold')
plt.tight_layout()
plt.show()

print("""
OBSERVATIONS:
  • Chaque chiffre a une forme caracteristique
  • Les pixels clairs (blancs) dessinent le chiffre
  • Les pixels sombres (noirs) forment le fond

CE QUE LE CNN DOIT APPRENDRE:
  Le CNN va detecter des MOTIFS (patterns) qui distinguent chaque chiffre:
  - Lignes verticales: 1, 4, 7
  - Boucles: 0, 6, 8, 9
  - Angles: 4, 7
  - Courbes: 2, 3, 5, 6, 9
""")

print("\n" + "=" * 70)
print("2. VARIABILITE D'UN MEME CHIFFRE")
print("=" * 70)

print("""
Un meme chiffre peut etre ecrit de differentes facons.
Le CNN doit etre ROBUSTE a ces variations!
""")

# Montrer plusieurs exemples du chiffre 3
target_digit = 3
indices = np.where(y == target_digit)[0][:8]

fig, axes = plt.subplots(1, 8, figsize=(14, 2.5))
fig.suptitle(f'8 facons differentes d\'ecrire le chiffre {target_digit}', fontsize=13, fontweight='bold')

for ax, idx in zip(axes, indices):
    image = X[idx].reshape(img_size, img_size)
    ax.imshow(image, cmap='gray', interpolation='nearest')
    ax.axis('off')

plt.tight_layout()
plt.show()

print("""
DEFI POUR LE CNN:
  Malgre ces differences, le CNN doit reconnaitre
  que tous ces exemples representent le meme chiffre!

C'est la que la CONVOLUTION aide:
  → Detecte les motifs LOCAUX (bords, courbes)
  → Invariante a la position (grace au pooling)
""")

print("\n" + "=" * 70)
print("3. VISUALISATION EN GRILLE DE PIXELS")
print("=" * 70)

print("""
Regardons un chiffre pixel par pixel pour comprendre
comment les valeurs numeriques forment l'image:
""")

# Prendre le premier exemple
image = X[0].reshape(img_size, img_size)
label = y[0]

fig, axes = plt.subplots(1, 2, figsize=(12, 5))

# Image avec valeurs
ax1 = axes[0]
im = ax1.imshow(image, cmap='gray', interpolation='nearest')
ax1.set_title(f'Chiffre {label} - Valeurs des pixels', fontsize=12, fontweight='bold')

# Afficher les valeurs
for i in range(img_size):
    for j in range(img_size):
        color = 'white' if image[i, j] < 8 else 'black'
        ax1.text(j, i, f'{image[i,j]:.0f}', ha='center', va='center',
                color=color, fontsize=10, fontweight='bold')
ax1.axis('off')

# Histogramme des valeurs
ax2 = axes[1]
ax2.hist(image.flatten(), bins=16, range=(0, 16), color='#9B7AC4', edgecolor='white', alpha=0.7)
ax2.set_xlabel('Valeur du pixel', fontsize=11)
ax2.set_ylabel('Frequence', fontsize=11)
ax2.set_title('Distribution des valeurs de pixels', fontsize=12, fontweight='bold')
ax2.grid(True, alpha=0.3)

plt.tight_layout()
plt.show()

print(f"""
LECTURE DU GRAPHIQUE:
  • Gauche: l'image avec la valeur numerique de chaque pixel
  • Droite: histogramme montrant la repartition des valeurs

Le chiffre est "dessine" par les pixels de haute valeur (clairs).
""")

print("\n" + "=" * 70)
print("              VISUALISATION TERMINEE!")
print("=" * 70)


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