{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# -- Cellule 1 : Installation\n!pip install albumentations scikit-learn tqdm pandas opencv-python-headless pydicom -q\n!pip install snntorch -q\nprint(\"OK deps\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:13.790931Z","iopub.execute_input":"2026-05-10T09:48:13.791680Z","iopub.status.idle":"2026-05-10T09:48:22.766674Z","shell.execute_reply.started":"2026-05-10T09:48:13.791626Z","shell.execute_reply":"2026-05-10T09:48:22.765748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 2 : Imports\nimport os, random, warnings\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\nwarnings.filterwarnings('ignore')\n \nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport torchvision.models as models\n \nimport snntorch as snn\nfrom snntorch import surrogate\n \nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n \nfrom sklearn.metrics import (accuracy_score, f1_score, precision_score,\n                              recall_score, roc_auc_score, confusion_matrix,\n                              classification_report, roc_curve)\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight\n \nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'PyTorch : {torch.__version__}')\nprint(f'Device  : {DEVICE}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:22.768760Z","iopub.execute_input":"2026-05-10T09:48:22.769563Z","iopub.status.idle":"2026-05-10T09:48:39.603538Z","shell.execute_reply.started":"2026-05-10T09:48:22.769510Z","shell.execute_reply":"2026-05-10T09:48:39.602818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\n# Liste les dossiers disponibles dans /kaggle/input pour trouver le dataset RSNA\nprint(\"Datasets disponibles dans /kaggle/input :\")\nprint(os.listdir('/kaggle/input'))\n\n# Recherche automatique du dossier RSNA Pneumonia Detection Challenge\nrsna_candidates = [p for p in Path('/kaggle/input').glob('*') if 'rsna' in p.name.lower() and 'pneumonia' in p.name.lower()]\nprint(\"\\nCandidats RSNA trouvés :\")\nfor p in rsna_candidates:\n    print(\" -\", p)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:39.604492Z","iopub.execute_input":"2026-05-10T09:48:39.605051Z","iopub.status.idle":"2026-05-10T09:48:39.611129Z","shell.execute_reply.started":"2026-05-10T09:48:39.605024Z","shell.execute_reply":"2026-05-10T09:48:39.610362Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cellule 3 : Dataset Kaggle déjà ajouté via \"Add Data\"\n# Dataset : RSNA Pneumonia Detection Challenge\n# Lien : https://www.kaggle.com/competitions/rsna-pneumonia-detection-challenge\n\nfrom pathlib import Path\nimport os\n\n# Chemin Kaggle du dataset RSNA. Sur Kaggle, il est généralement ici :\nDATASET_ROOT = Path(\"/kaggle/input/competitions/rsna-pneumonia-detection-challenge\")\n\n# Fallback si le nom exact du dossier est légèrement différent\nif not DATASET_ROOT.exists():\n    candidates = [p for p in Path('/kaggle/input').glob('*') if 'rsna' in p.name.lower() and 'pneumonia' in p.name.lower()]\n    if len(candidates) > 0:\n        DATASET_ROOT = candidates[0]\n\nprint(\"Dataset root :\", DATASET_ROOT)\nprint(\"Contenu du dossier dataset :\")\nprint(os.listdir(DATASET_ROOT))\n\nprint(\"\\nChemins utiles :\")\nprint(\"Train images :\", DATASET_ROOT / \"stage_2_train_images\")\nprint(\"Labels CSV   :\", DATASET_ROOT / \"stage_2_train_labels.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:39.613197Z","iopub.execute_input":"2026-05-10T09:48:39.613754Z","iopub.status.idle":"2026-05-10T09:48:39.684186Z","shell.execute_reply.started":"2026-05-10T09:48:39.613715Z","shell.execute_reply":"2026-05-10T09:48:39.683438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 4 : Configuration\nSEED         = 42\nDATA_DIR     = DATASET_ROOT  # RSNA Pneumonia Detection Challenge\n \nIMG_SIZE     = 224\nBATCH_SIZE   = 32\nEPOCHS       = 5\nLR           = 3e-4\nWEIGHT_DECAY = 1e-2\nVAL_SPLIT    = 0.1  # 10% val\nTEST_SPLIT   = 0.1  # 10% test\nGRAD_CLIP    = 1.0\nNUM_WORKERS  = 2\nUSE_AMP      = torch.cuda.is_available()\nNUM_CLASSES  = 2\nSOTA_ACC     = 98.81\n \n# =====================================================================\n# MULTI-CHANNEL PREPROCESSING : nombre de canaux = 6\n#   ch0 : Grayscale normalisé        [0,1]\n#   ch1 : CLAHE (contraste local)    [0,1]\n#   ch2 : Canny edges normalisé      [0,1]\n#   ch3 : Gradient Sobel normalisé   [0,1]\n#   ch4 : LBP (texture locale)       [0,1]\n#   ch5 : Top-Hat morphologique      [0,1]\n# =====================================================================\nNUM_CHANNELS = 6\n \nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark     = False\n \nprint('Config chargee')\nprint(f'   IMG_SIZE={IMG_SIZE}  BATCH={BATCH_SIZE}  EPOCHS={EPOCHS}')\nprint(f'   Train/Val/Test : {int((1-(VAL_SPLIT+TEST_SPLIT))*100)}% / {int(VAL_SPLIT*100)}% / {int(TEST_SPLIT*100)}%')\nprint(f'   Canaux : {NUM_CHANNELS} (Grayscale, CLAHE, Canny, Sobel, LBP, Top-Hat)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:39.685200Z","iopub.execute_input":"2026-05-10T09:48:39.685608Z","iopub.status.idle":"2026-05-10T09:48:39.712958Z","shell.execute_reply.started":"2026-05-10T09:48:39.685576Z","shell.execute_reply":"2026-05-10T09:48:39.712304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 5 : ============================================================\n#                MULTI-CHANNEL PREPROCESSING\n# ===========================================================================\nSPIKE_STEPS = 8\n \n \ndef _safe_norm(arr: np.ndarray) -> np.ndarray:\n    \"\"\"Normalise un tableau [H,W] float32 dans [0,1].\"\"\"\n    mn, mx = arr.min(), arr.max()\n    if mx - mn < 1e-6:\n        return np.zeros_like(arr, dtype=np.float32)\n    return ((arr - mn) / (mx - mn)).astype(np.float32)\n \n \ndef compute_lbp(gray: np.ndarray, radius: int = 1, n_points: int = 8) -> np.ndarray:\n    \"\"\"\n    LBP uniforme simplifié calculé avec OpenCV (sans scikit-image).\n    gray : uint8 [H,W]\n    return : float32 [H,W] dans [0,1]\n    \"\"\"\n    H, W = gray.shape\n    lbp = np.zeros((H, W), dtype=np.float32)\n    offsets = [(int(round(radius * np.sin(2 * np.pi * i / n_points))),\n                int(round(radius * np.cos(2 * np.pi * i / n_points))))\n               for i in range(n_points)]\n    gray_f = gray.astype(np.float32)\n    for bit, (dy, dx) in enumerate(offsets):\n        shifted = np.roll(np.roll(gray_f, dy, axis=0), dx, axis=1)\n        lbp += ((gray_f >= shifted).astype(np.float32)) * (2 ** bit)\n    return _safe_norm(lbp)\n \n \ndef preprocess_multichannel(img_bgr: np.ndarray, img_size: int = 224) -> np.ndarray:\n    \"\"\"\n    Construit un tenseur [H, W, NUM_CHANNELS] à partir d'une image BGR.\n \n    Canaux :\n        0 - Grayscale normalisé\n        1 - CLAHE (Contrast Limited Adaptive Histogram Equalization)\n        2 - Canny edges (normalisé)\n        3 - Gradient Sobel magnitude (normalisé)\n        4 - LBP — Local Binary Pattern (texture)\n        5 - Top-Hat morphologique (détection de petites structures claires)\n    \"\"\"\n    # --- Resize ---\n    img = cv2.resize(img_bgr, (img_size, img_size))\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)   # uint8 [H,W]\n \n    # --- Canal 0 : Grayscale normalisé [0,1] ---\n    ch0_gray = gray.astype(np.float32) / 255.0\n \n    # --- Canal 1 : CLAHE ---\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    ch1_clahe = clahe.apply(gray).astype(np.float32) / 255.0\n \n    # --- Canal 2 : Canny edges ---\n    edges = cv2.Canny(gray, threshold1=50, threshold2=150)\n    ch2_canny = edges.astype(np.float32) / 255.0\n \n    # --- Canal 3 : Gradient Sobel magnitude ---\n    sobelx = cv2.Sobel(gray, cv2.CV_32F, 1, 0, ksize=3)\n    sobely = cv2.Sobel(gray, cv2.CV_32F, 0, 1, ksize=3)\n    sobel_mag = np.sqrt(sobelx ** 2 + sobely ** 2)\n    ch3_sobel = _safe_norm(sobel_mag)\n \n    # --- Canal 4 : LBP (texture locale) ---\n    ch4_lbp = compute_lbp(gray, radius=1, n_points=8)\n \n    # --- Canal 5 : Top-Hat morphologique ---\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))\n    tophat = cv2.morphologyEx(gray, cv2.MORPH_TOPHAT, kernel)\n    ch5_tophat = tophat.astype(np.float32) / 255.0\n \n    # --- Empilement [H,W,6] ---\n    multichannel = np.stack(\n        [ch0_gray, ch1_clahe, ch2_canny, ch3_sobel, ch4_lbp, ch5_tophat],\n        axis=-1\n    )\n    return multichannel.astype(np.float32)\n \n \ndef rate_encode_spikes(x: torch.Tensor, num_steps: int = SPIKE_STEPS) -> torch.Tensor:\n    \"\"\"\n    x : [C, H, W] dans [0,1]\n    return : [T, C, H, W]  — encodage par taux de décharge (Bernoulli)\n    \"\"\"\n    return torch.stack([torch.bernoulli(x) for _ in range(num_steps)], dim=0)\n \n \ndef read_rsna_image_bgr(img_path: str) -> np.ndarray:\n    \"\"\"Lit une image RSNA DICOM (.dcm) ou une image classique et retourne BGR uint8.\"\"\"\n    if str(img_path).lower().endswith('.dcm'):\n        ds = pydicom.dcmread(img_path)\n        img = ds.pixel_array.astype(np.float32)\n        img = img - img.min()\n        if img.max() > 0:\n            img = img / img.max()\n        img_uint8 = (img * 255).astype(np.uint8)\n        return cv2.cvtColor(img_uint8, cv2.COLOR_GRAY2BGR)\n\n    img_bgr = cv2.imread(img_path)\n    if img_bgr is None:\n        raise ValueError(f\"Image non lue correctement : {img_path}\")\n    return img_bgr\n\n\ndef visualize_multichannel(img_path: str):\n    \"\"\"Visualise les 6 canaux multi-channel pour une image donnée.\"\"\"\n    img = read_rsna_image_bgr(img_path)\n    mc = preprocess_multichannel(img, IMG_SIZE)\n \n    channel_names = [\n        'ch0 — Grayscale norm.',\n        'ch1 — CLAHE',\n        'ch2 — Canny edges',\n        'ch3 — Sobel gradient',\n        'ch4 — LBP texture',\n        'ch5 — Top-Hat morph.',\n    ]\n \n    fig, axes = plt.subplots(1, NUM_CHANNELS + 1, figsize=(20, 3))\n \n    # Image originale\n    axes[0].imshow(cv2.cvtColor(\n        cv2.resize(img, (IMG_SIZE, IMG_SIZE)), cv2.COLOR_BGR2RGB))\n    axes[0].set_title('Original (RGB)', fontsize=9)\n \n    for i in range(NUM_CHANNELS):\n        axes[i + 1].imshow(mc[:, :, i], cmap='gray', vmin=0, vmax=1)\n        axes[i + 1].set_title(channel_names[i], fontsize=8)\n \n    for ax in axes:\n        ax.axis('off')\n \n    plt.suptitle('Multi-Channel Preprocessing', fontweight='bold', fontsize=12)\n    plt.tight_layout()\n    plt.show()\n \n \nprint(f'Preprocessing multi-canal defini : {NUM_CHANNELS} canaux')\nprint('  ch0=Grayscale  ch1=CLAHE  ch2=Canny  ch3=Sobel  ch4=LBP  ch5=Top-Hat')\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:39.714046Z","iopub.execute_input":"2026-05-10T09:48:39.714361Z","iopub.status.idle":"2026-05-10T09:48:39.731748Z","shell.execute_reply.started":"2026-05-10T09:48:39.714331Z","shell.execute_reply":"2026-05-10T09:48:39.731163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 6 : Augmentations (adaptées aux 6 canaux) ----------------------\n#   mean/std pour 6 canaux identiques (0.5 par canal)\nMEAN_MC = [0.5] * NUM_CHANNELS\nSTD_MC  = [0.5] * NUM_CHANNELS\n \n \ndef get_train_transforms():\n    return A.Compose([\n        A.Resize(IMG_SIZE, IMG_SIZE),\n        A.HorizontalFlip(p=0.5),\n        A.ShiftScaleRotate(\n            shift_limit=0.05,\n            scale_limit=0.08,\n            rotate_limit=10,\n            p=0.5\n        ),\n        A.RandomBrightnessContrast(\n            brightness_limit=0.1,\n            contrast_limit=0.1,\n            p=0.3\n        ),\n        A.GaussNoise(std_range=(0.01, 0.03), p=0.2),\n        A.Normalize(mean=MEAN_MC, std=STD_MC, max_pixel_value=1.0),\n        ToTensorV2(),\n    ])\n \n \ndef get_val_transforms():\n    return A.Compose([\n        A.Resize(IMG_SIZE, IMG_SIZE),\n        A.Normalize(mean=MEAN_MC, std=STD_MC, max_pixel_value=1.0),\n        ToTensorV2(),\n    ])\n \n \nprint(f'Augmentations definies pour {NUM_CHANNELS} canaux')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:39.732559Z","iopub.execute_input":"2026-05-10T09:48:39.732870Z","iopub.status.idle":"2026-05-10T09:48:39.754264Z","shell.execute_reply.started":"2026-05-10T09:48:39.732849Z","shell.execute_reply":"2026-05-10T09:48:39.753503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 7 : Dataset Multi-Channel + Split 80/10/10 + Spike Encoding ----\n \nclass KermanySpikeDataset(Dataset):\n    # On garde le même nom de classe pour ne pas changer le reste du code\n    # Normal / No Pneumonia = 0\n    # Pneumonia = 1\n    CLASSES = {\n        'Normal': 0,\n        'Pneumonia': 1\n    }\n\n    def __init__(self, samples, transform=None, spike_steps=SPIKE_STEPS):\n        self.samples     = samples\n        self.transform   = transform\n        self.spike_steps = spike_steps\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        path, label = self.samples[idx]\n\n        # RSNA fournit des images DICOM (.dcm). On les convertit en BGR uint8\n        # pour garder exactement le même preprocessing multi-canal en aval.\n        img_bgr = read_rsna_image_bgr(path)\n\n        # ---- Multi-Channel Preprocessing ----\n        x = preprocess_multichannel(img_bgr, IMG_SIZE)   # [H, W, 6]\n\n        if self.transform:\n            x = self.transform(image=x)['image']         # [6, H, W]\n\n        # Ramener dans [0,1] pour spike encoding\n        x = (x * 0.5) + 0.5\n        x = torch.clamp(x, 0.0, 1.0)\n\n        # ---- Spike Encoding ----\n        spikes = rate_encode_spikes(x, self.spike_steps)  # [T, 6, H, W]\n\n        return spikes, label\n \n \n# --- Chargement des images depuis RSNA Pneumonia Detection Challenge ---\n# Structure Kaggle :\n# DATA_DIR/\n#   stage_2_train_images/*.dcm\n#   stage_2_train_labels.csv\n#\n# Remarque : un patient positif peut apparaître plusieurs fois dans le CSV\n# s'il possède plusieurs bounding boxes. Pour une classification binaire,\n# on garde une seule ligne par patientId avec Target=max(Target).\n\nlabels_csv = DATA_DIR / \"stage_2_train_labels.csv\"\ntrain_img_dir = DATA_DIR / \"stage_2_train_images\"\n\nif not labels_csv.exists():\n    raise FileNotFoundError(f\"CSV introuvable : {labels_csv}\")\nif not train_img_dir.exists():\n    raise FileNotFoundError(f\"Dossier images introuvable : {train_img_dir}\")\n\nlabels_df = pd.read_csv(labels_csv)\nlabels_df = labels_df.groupby('patientId', as_index=False)['Target'].max()\n\nall_samples = []\nmissing_files = 0\n\nfor _, row in labels_df.iterrows():\n    patient_id = row['patientId']\n    label = int(row['Target'])\n    img_path = train_img_dir / f\"{patient_id}.dcm\"\n\n    if img_path.exists():\n        all_samples.append((str(img_path), label))\n    else:\n        missing_files += 1\n\nall_labels = [l for _, l in all_samples]\n\nn_normal = sum(1 for l in all_labels if l == 0)\nn_pneumo = sum(1 for l in all_labels if l == 1)\n\nprint(f\"\\nTotal : {len(all_samples)} images\")\nprint(f\"  Normal / No Pneumonia : {n_normal}\")\nprint(f\"  Pneumonia             : {n_pneumo}\")\nprint(f\"  Fichiers manquants    : {missing_files}\")\n\nif len(all_samples) == 0:\n    raise ValueError(\"Aucune image trouvée. Vérifie le chemin DATA_DIR et le dataset RSNA ajouté à Kaggle.\")\n\nprint(f'Total : {len(all_samples)} images')\nprint(f'  Normal    : {n_normal} ({n_normal/len(all_samples)*100:.1f}%)')\nprint(f'  Pneumonia : {n_pneumo} ({n_pneumo/len(all_samples)*100:.1f}%)')\n \n# --- Split stratifié 80 / 10 / 10 ---\n# Étape 1 : Train (80%) + Temp (20%)\ntrain_samples, temp_samples, train_labels, temp_labels = train_test_split(\n    all_samples, all_labels,\n    test_size=0.2,\n    random_state=SEED,\n    stratify=all_labels\n)\n \n# Étape 2 : Temp → Val (10%) + Test (10%)\nval_samples, test_samples, val_labels, test_labels = train_test_split(\n    temp_samples, temp_labels,\n    test_size=0.5,\n    random_state=SEED,\n    stratify=temp_labels\n)\n \ntrain_ds = KermanySpikeDataset(train_samples, get_train_transforms())\nval_ds   = KermanySpikeDataset(val_samples,   get_val_transforms())\ntest_ds  = KermanySpikeDataset(test_samples,  get_val_transforms())\n \ntrain_dl = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                      num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\nval_dl   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=NUM_WORKERS, pin_memory=True)\ntest_dl  = DataLoader(test_ds,  batch_size=BATCH_SIZE, shuffle=False,\n                      num_workers=NUM_WORKERS, pin_memory=True)\n \nn_tn = sum(1 for _, l in train_samples if l == 0)\nn_tp = sum(1 for _, l in train_samples if l == 1)\nn_vn = sum(1 for _, l in val_samples   if l == 0)\nn_vp = sum(1 for _, l in val_samples   if l == 1)\nn_tsn = sum(1 for _, l in test_samples if l == 0)\nn_tsp = sum(1 for _, l in test_samples if l == 1)\nprint(f'Train : {len(train_samples)} ({n_tn} Normal, {n_tp} Pneumonia)')\nprint(f'Val   : {len(val_samples)}   ({n_vn} Normal, {n_vp} Pneumonia)')\nprint(f'Test  : {len(test_samples)}  ({n_tsn} Normal, {n_tsp} Pneumonia)')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:48:39.755310Z","iopub.execute_input":"2026-05-10T09:48:39.755589Z","iopub.status.idle":"2026-05-10T09:50:04.540238Z","shell.execute_reply.started":"2026-05-10T09:48:39.755558Z","shell.execute_reply":"2026-05-10T09:50:04.539551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 8 : Visualisation Multi-Channel + Distribution -----------------\nfor label_name, label_id in [('Normal', 0), ('Pneumonia', 1)]:\n    example_path = next(p for p, l in all_samples if l == label_id)\n    print(f'\\nVisualisation Multi-Channel — {label_name}')\n    visualize_multichannel(example_path)\n \nfig, axes = plt.subplots(1, 3, figsize=(16, 4))\ncounts_train = [sum(1 for _, l in train_samples if l == c) for c in [0, 1]]\ncounts_val   = [sum(1 for _, l in val_samples   if l == c) for c in [0, 1]]\ncounts_test  = [sum(1 for _, l in test_samples  if l == c) for c in [0, 1]]\n \nfor ax, counts, title in zip(\n        axes,\n        [counts_train, counts_val, counts_test],\n        ['Distribution Train (80%)', 'Distribution Val (10%)', 'Distribution Test (10%)']\n):\n    ax.bar(['Normal', 'Pneumonia'], counts, color=['steelblue', 'tomato'])\n    ax.set_title(title, fontweight='bold')\n    ax.set_ylabel('Images')\n    for i, v in enumerate(counts):\n        ax.text(i, v + 10, str(v), ha='center', fontweight='bold')\n \nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:50:04.541239Z","iopub.execute_input":"2026-05-10T09:50:04.541630Z","iopub.status.idle":"2026-05-10T09:50:05.912130Z","shell.execute_reply.started":"2026-05-10T09:50:04.541606Z","shell.execute_reply":"2026-05-10T09:50:05.911506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 9 : CNN + GRU + SNN + Attention-Guided (adapté 6 canaux) -------\n \nclass SpatialAttention(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels, in_channels // 2, kernel_size=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(in_channels // 2, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n \n    def forward(self, x):\n        attn = self.conv(x)        # [B,1,H,W]\n        return x * attn, attn\n \n \nclass CNNEncoder(nn.Module):\n    \"\"\"\n    CNNEncoder adapté pour NUM_CHANNELS canaux d'entrée (6 au lieu de 3).\n    La première couche Conv2d accepte in_channels=NUM_CHANNELS.\n    \"\"\"\n    def __init__(self, in_channels=NUM_CHANNELS):\n        super().__init__()\n        self.features = nn.Sequential(\n            nn.Conv2d(in_channels, 32, 3, padding=1),   # 6 → 32\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n \n            nn.Conv2d(32, 64, 3, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n \n            nn.Conv2d(64, 128, 3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n        )\n        self.attn = SpatialAttention(128)\n        self.pool = nn.AdaptiveAvgPool2d((1, 1))\n \n    def forward(self, x):\n        x = self.features(x)\n        x, attn = self.attn(x)\n        x = self.pool(x).flatten(1)   # [B,128]\n        return x, attn\n \n \nclass CNN_GRU_SNN_Attention(nn.Module):\n    def __init__(self, num_classes=2, spike_steps=SPIKE_STEPS,\n                 hidden_size=128, in_channels=NUM_CHANNELS):\n        super().__init__()\n        self.spike_steps = spike_steps\n        self.encoder = CNNEncoder(in_channels=in_channels)\n \n        self.gru = nn.GRU(\n            input_size=128,\n            hidden_size=hidden_size,\n            num_layers=1,\n            batch_first=True,\n            bidirectional=True\n        )\n \n        self.fc1  = nn.Linear(hidden_size * 2, 128)\n        self.lif1 = snn.Leaky(beta=0.9, spike_grad=surrogate.fast_sigmoid())\n        self.fc2  = nn.Linear(128, num_classes)\n \n    def forward(self, x):\n        # x : [B, T, C, H, W]   avec C=NUM_CHANNELS\n        B, T, C, H, W = x.shape\n        seq_feats, attn_maps = [], []\n \n        for t in range(T):\n            feat_t, attn_t = self.encoder(x[:, t])\n            seq_feats.append(feat_t)\n            attn_maps.append(attn_t)\n \n        seq_feats = torch.stack(seq_feats, dim=1)    # [B, T, 128]\n        gru_out, _ = self.gru(seq_feats)             # [B, T, 2H]\n        temporal_feat = gru_out[:, -1, :]            # [B, 2H]\n \n        cur  = self.fc1(temporal_feat)\n        mem1 = self.lif1.init_leaky()\n        spk1, mem1 = self.lif1(cur, mem1)\n        logits = self.fc2(mem1)\n        return logits\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:50:05.914176Z","iopub.execute_input":"2026-05-10T09:50:05.914504Z","iopub.status.idle":"2026-05-10T09:50:05.926232Z","shell.execute_reply.started":"2026-05-10T09:50:05.914479Z","shell.execute_reply":"2026-05-10T09:50:05.925631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 10 : Instanciation du modele -----------------------------------\nmodel = CNN_GRU_SNN_Attention(\n    num_classes=NUM_CLASSES,\n    spike_steps=SPIKE_STEPS,\n    hidden_size=128,\n    in_channels=NUM_CHANNELS\n).to(DEVICE)\n \nn_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'CNN + GRU + SNN + Attention : {n_params/1e6:.2f}M params')\nprint(f'Input attendu : [B, T={SPIKE_STEPS}, C={NUM_CHANNELS}, {IMG_SIZE}, {IMG_SIZE}]')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:50:05.926968Z","iopub.execute_input":"2026-05-10T09:50:05.927430Z","iopub.status.idle":"2026-05-10T09:50:06.524960Z","shell.execute_reply.started":"2026-05-10T09:50:05.927401Z","shell.execute_reply":"2026-05-10T09:50:06.524309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 11 : Loss + Optimizer + Scheduler ------------------------------\ncw = compute_class_weight('balanced', classes=np.unique(all_labels), y=all_labels)\nclass_weights = torch.tensor(cw, dtype=torch.float32).to(DEVICE)\nprint(f'Poids de classe : Normal={cw[0]:.3f}, Pneumonia={cw[1]:.3f}')\n \ncriterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=0.02)\noptimizer = AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)\nscaler    = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n \nprint('CrossEntropyLoss + AdamW + CosineAnnealingLR')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:50:06.525951Z","iopub.execute_input":"2026-05-10T09:50:06.526267Z","iopub.status.idle":"2026-05-10T09:50:06.549175Z","shell.execute_reply.started":"2026-05-10T09:50:06.526216Z","shell.execute_reply":"2026-05-10T09:50:06.548604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, loader, optimizer, criterion, scaler):\n    model.train()\n    total_loss = 0.0\n    preds_all, labels_all = [], []\n \n    pbar = tqdm(loader, desc='  [TRAIN]', leave=False)\n    for spikes, labels in pbar:\n        spikes = spikes.float().to(DEVICE, non_blocking=True)\n        labels = labels.long().to(DEVICE, non_blocking=True)\n \n        optimizer.zero_grad(set_to_none=True)\n \n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits = model(spikes)\n            loss   = criterion(logits, labels)\n \n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n \n        total_loss += loss.item()\n        preds_all.append(logits.argmax(1).detach())\n        labels_all.append(labels.detach())\n \n    preds_all  = torch.cat(preds_all).cpu().numpy()\n    labels_all = torch.cat(labels_all).cpu().numpy()\n    return total_loss / len(loader), accuracy_score(labels_all, preds_all)\n \n \n@torch.no_grad()\ndef validate(model, loader, criterion):\n    model.eval()\n    total_loss, preds_all, labels_all, probs_all = 0.0, [], [], []\n \n    for spikes, labels in loader:\n        spikes = spikes.float().to(DEVICE)\n        labels = labels.long().to(DEVICE)\n \n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits = model(spikes)\n            loss   = criterion(logits, labels)\n \n        probs = F.softmax(logits, dim=1)[:, 1]\n        total_loss += loss.item()\n        preds_all.extend(logits.argmax(1).cpu().numpy())\n        labels_all.extend(labels.cpu().numpy())\n        probs_all.extend(probs.cpu().numpy())\n \n    acc = accuracy_score(labels_all, preds_all)\n    f1  = f1_score(labels_all, preds_all, average='weighted')\n    auc = roc_auc_score(labels_all, probs_all)\n    return total_loss / len(loader), acc, f1, auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:50:06.550007Z","iopub.execute_input":"2026-05-10T09:50:06.550313Z","iopub.status.idle":"2026-05-10T09:50:06.559432Z","shell.execute_reply.started":"2026-05-10T09:50:06.550291Z","shell.execute_reply":"2026-05-10T09:50:06.558608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 13 : Boucle d'entrainement -------------------------------------\nos.makedirs('checkpoints', exist_ok=True)\n \nPATIENCE = 6\nbest_acc, best_state = 0.0, None\nhistory = []\npatience_ctr = 0\n \nprint(f'Entrainement : {EPOCHS} epochs  (patience={PATIENCE})')\nprint(f'   Modele       : CNN + GRU + SNN + Attention')\nprint(f'   Preprocessing: Multi-Channel ({NUM_CHANNELS} canaux)')\nprint(f'   Split        : Train {len(train_samples)} / Val {len(val_samples)} / Test {len(test_samples)}')\nprint('=' * 70)\n \nfor epoch in range(EPOCHS):\n    tr_loss, tr_acc = train_epoch(model, train_dl, optimizer, criterion, scaler)\n    vl_loss, vl_acc, vl_f1, vl_auc = validate(model, val_dl, criterion)\n    scheduler.step()\n \n    history.append({\n        'epoch':   epoch + 1,\n        'tr_loss': tr_loss, 'tr_acc': tr_acc,\n        'vl_loss': vl_loss, 'vl_acc': vl_acc,\n        'vl_f1':   vl_f1,   'vl_auc': vl_auc,\n    })\n \n    print(f'Epoch {epoch+1:02d}/{EPOCHS} | '\n          f'Tr={tr_acc:.4f}  Val={vl_acc:.4f}  '\n          f'F1={vl_f1:.4f}  AUC={vl_auc:.4f}  '\n          f'LR={optimizer.param_groups[0][\"lr\"]:.2e}')\n \n    if vl_acc > best_acc:\n        best_acc   = vl_acc\n        best_state = {k: v.clone() for k, v in model.state_dict().items()}\n        torch.save(best_state, 'checkpoints/cnn_gru_snn_attention_multichannel_best.pth')\n        patience_ctr = 0\n        print(f'   Nouveau record ! Acc={best_acc*100:.2f}%')\n    else:\n        patience_ctr += 1\n        if patience_ctr >= PATIENCE:\n            print(f'Early stopping (patience={PATIENCE})')\n            break\n \nmodel.load_state_dict(best_state)\nprint(f'Entrainement termine. Meilleure Val Acc = {best_acc*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T09:50:06.560226Z","iopub.execute_input":"2026-05-10T09:50:06.560604Z","iopub.status.idle":"2026-05-10T10:54:31.917040Z","shell.execute_reply.started":"2026-05-10T09:50:06.560583Z","shell.execute_reply":"2026-05-10T10:54:31.915959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 14 : Courbes d'apprentissage -----------------------------------\nhist_df = pd.DataFrame(history)\nfig, axes = plt.subplots(1, 3, figsize=(20, 5))\n \naxes[0].plot(hist_df.epoch, hist_df.tr_loss, label='Train', color='steelblue')\naxes[0].plot(hist_df.epoch, hist_df.vl_loss, label='Val',   color='tomato')\naxes[0].set_title('Loss'); axes[0].legend(); axes[0].set_xlabel('Epoch')\n \naxes[1].plot(hist_df.epoch, hist_df.tr_acc * 100, label='Train', color='steelblue')\naxes[1].plot(hist_df.epoch, hist_df.vl_acc * 100, label='Val',   color='tomato')\naxes[1].axhline(y=SOTA_ACC, color='green', linestyle='--', linewidth=2,\n                label=f'SOTA {SOTA_ACC}%')\naxes[1].set_title('Accuracy (%)'); axes[1].legend(); axes[1].set_xlabel('Epoch')\n \naxes[2].plot(hist_df.epoch, hist_df.vl_f1,  label='F1',  color='purple')\naxes[2].plot(hist_df.epoch, hist_df.vl_auc, label='AUC', color='orange')\naxes[2].set_title('F1 & AUC'); axes[2].legend(); axes[2].set_xlabel('Epoch')\n \nplt.suptitle('CNN+GRU+SNN+Attention — Multi-Channel Preprocessing (6 canaux)',\n             fontsize=13, fontweight='bold')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T10:54:31.918675Z","iopub.execute_input":"2026-05-10T10:54:31.919313Z","iopub.status.idle":"2026-05-10T10:54:32.416012Z","shell.execute_reply.started":"2026-05-10T10:54:31.919244Z","shell.execute_reply":"2026-05-10T10:54:32.415008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 15 : Evaluation Finale (Test Set) ------------------------------\nmodel.eval()\nall_preds, all_probs, all_true = [], [], []\n \nwith torch.no_grad():\n    for spikes, labels in tqdm(test_dl, desc='Eval finale'):\n        spikes = spikes.float().to(DEVICE)\n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits = model(spikes)\n        probs = F.softmax(logits, dim=1)[:, 1]\n        all_preds.extend(logits.argmax(1).cpu().numpy())\n        all_probs.extend(probs.cpu().numpy())\n        all_true.extend(labels.numpy())\n \nacc  = accuracy_score(all_true, all_preds)\nf1   = f1_score(all_true, all_preds, average='weighted')\nprec = precision_score(all_true, all_preds, average='weighted', zero_division=0)\nrec  = recall_score(all_true, all_preds, average='weighted', zero_division=0)\nauc  = roc_auc_score(all_true, all_probs)\n \nprint('RESULTATS FINAUX')\nprint('Modele : CNN + GRU + SNN + Attention-Guided')\nprint(f'Preprocessing : Multi-Channel ({NUM_CHANNELS} canaux) + Spike Encoding')\nprint(f'Split : 80% Train / 10% Val / 10% Test')\nprint(f'  Accuracy  : {acc*100:.2f}%')\nprint(f'  F1-Score  : {f1*100:.2f}%')\nprint(f'  Precision : {prec*100:.2f}%')\nprint(f'  Recall    : {rec*100:.2f}%')\nprint(f'  AUC-ROC   : {auc*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T10:54:32.417155Z","iopub.execute_input":"2026-05-10T10:54:32.417560Z","iopub.status.idle":"2026-05-10T10:56:01.380031Z","shell.execute_reply.started":"2026-05-10T10:54:32.417524Z","shell.execute_reply":"2026-05-10T10:56:01.379135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 16 : Matrice de Confusion + Courbe ROC -------------------------\ncm = confusion_matrix(all_true, all_preds)\n \nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n \nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=axes[0],\n            xticklabels=['Normal', 'Pneumonia'],\nyticklabels=['Normal', 'Pneumonia'])\naxes[0].set_ylabel('Vrai label')\naxes[0].set_xlabel('Prediction')\naxes[0].set_title(f'Matrice de Confusion\\nAcc={acc*100:.2f}%  AUC={auc*100:.2f}%',\n                  fontweight='bold')\n \nfpr, tpr, _ = roc_curve(all_true, all_probs)\naxes[1].plot(fpr, tpr, color='darkorange', lw=2,\n             label=f'CNN+GRU+SNN ROC (AUC={auc:.4f})')\naxes[1].plot([0, 1], [0, 1], 'navy', linestyle='--', lw=1)\naxes[1].set_xlabel('False Positive Rate')\naxes[1].set_ylabel('True Positive Rate')\naxes[1].set_title('Courbe ROC', fontweight='bold')\naxes[1].legend(loc='lower right')\n \nplt.tight_layout()\nplt.show()\nprint(classification_report(all_true, all_preds, target_names=['NORMAL', 'PNEUMONIA']))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T10:56:01.381428Z","iopub.execute_input":"2026-05-10T10:56:01.381764Z","iopub.status.idle":"2026-05-10T10:56:02.026612Z","shell.execute_reply.started":"2026-05-10T10:56:01.381736Z","shell.execute_reply":"2026-05-10T10:56:02.025979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 17 : Tableau récapitulatif final --------------------------------\nmodel.eval()\nall_preds, all_probs, true_labels = [], [], []\n \nwith torch.no_grad():\n    for spikes, labels in test_dl:\n        spikes = spikes.float().to(DEVICE)\n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits = model(spikes)\n        probs = F.softmax(logits, dim=1)[:, 1]\n        all_preds.extend(logits.argmax(1).cpu().numpy())\n        all_probs.extend(probs.cpu().numpy())\n        true_labels.extend(labels.numpy())\n \nacc  = accuracy_score(true_labels, all_preds)\nf1   = f1_score(true_labels, all_preds, average='weighted')\nprec = precision_score(true_labels, all_preds, average='weighted', zero_division=0)\nrec  = recall_score(true_labels, all_preds, average='weighted', zero_division=0)\nauc  = roc_auc_score(true_labels, all_probs)\n \nprint(f'{\"Model\":<35} {\"Accuracy\":>10} {\"F1\":>10} {\"Precision\":>12} {\"Recall\":>10} {\"AUC\":>10}')\nprint('-' * 90)\nprint(f'{\"CNN+GRU+SNN+Attn (6-ch preproc)\":<35} '\n      f'{acc*100:>9.2f}% {f1*100:>9.2f}% {prec*100:>11.2f}% {rec*100:>9.2f}% {auc*100:>9.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T10:56:02.027780Z","iopub.execute_input":"2026-05-10T10:56:02.028125Z","iopub.status.idle":"2026-05-10T10:57:22.080234Z","shell.execute_reply.started":"2026-05-10T10:56:02.028088Z","shell.execute_reply":"2026-05-10T10:57:22.079205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 18 : Sauvegarde finale -----------------------------------------\ntorch.save({\n    'model_state'    : model.state_dict(),\n    'accuracy'       : acc,\n    'precision'      : prec,\n    'recall'         : rec,\n    'f1'             : f1,\n    'auc'            : auc,\n    'img_size'       : IMG_SIZE,\n    'spike_steps'    : SPIKE_STEPS,\n    'num_channels'   : NUM_CHANNELS,\n    'preprocessing'  : 'multi-channel (Grayscale+CLAHE+Canny+Sobel+LBP+TopHat) + spike encoding',\n    'channels_desc'  : ['Grayscale', 'CLAHE', 'Canny edges', 'Sobel gradient', 'LBP', 'Top-Hat'],\n    'split'          : '80/10/10 train/val/test',\n    'architecture'   : 'CNN + GRU + SNN + Attention-Guided',\n    'num_classes'    : NUM_CLASSES,\n}, 'checkpoints/cnn_gru_snn_attention_multichannel_final.pth')\n \nprint('Modele sauvegarde -> checkpoints/cnn_gru_snn_attention_multichannel_final.pth')\nprint(f'   Preprocessing : {NUM_CHANNELS} canaux (Grayscale, CLAHE, Canny, Sobel, LBP, Top-Hat)')\nprint(f'   Accuracy  : {acc*100:.2f}%')\nprint(f'   Precision : {prec*100:.2f}%')\nprint(f'   Recall    : {rec*100:.2f}%')\nprint(f'   AUC-ROC   : {auc*100:.2f}%')\nprint(f'   F1-Score  : {f1*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T10:57:22.081725Z","iopub.execute_input":"2026-05-10T10:57:22.082026Z","iopub.status.idle":"2026-05-10T10:57:22.096365Z","shell.execute_reply.started":"2026-05-10T10:57:22.081997Z","shell.execute_reply":"2026-05-10T10:57:22.095547Z"}},"outputs":[],"execution_count":null}]}