{"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,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":" *Generated from: cnn-simple-covid-ds.ipynb Adapted for: RSNA Pneumonia Detection Challenge Dataset: ** https://www.kaggle.com/competitions/rsna-pneumonia-detection-challenge* \n\n* Algorithm: CNN + GRU + SNN + Attention-Guided\n* Split:     80% / 10% / 10% ","metadata":{}},{"cell_type":"code","source":" #-- Cellule 1 : Installation\n# pydicom ajouté pour lire les fichiers DICOM de l'RSNA\n!pip install albumentations scikit-learn tqdm pandas opencv-python-headless -q\n!pip install snntorch -q\n!pip install pydicom -q         \nprint(\"OK deps\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:44:18.622964Z","iopub.execute_input":"2026-05-09T12:44:18.623533Z","iopub.status.idle":"2026-05-09T12:44:30.942058Z","shell.execute_reply.started":"2026-05-09T12:44:18.623502Z","shell.execute_reply":"2026-05-09T12:44:30.940916Z"}},"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 matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\nwarnings.filterwarnings('ignore')\n \nimport pydicom                   # ← ajout RSNA\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-09T12:44:30.943788Z","iopub.execute_input":"2026-05-09T12:44:30.944634Z","iopub.status.idle":"2026-05-09T12:44:47.923219Z","shell.execute_reply.started":"2026-05-09T12:44:30.944600Z","shell.execute_reply":"2026-05-09T12:44:47.922517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 4 : Configuration  \nSEED         = 42\nIMG_SIZE     = 224\nBATCH_SIZE   = 32\nEPOCHS       = 5\nLR           = 3e-4\nWEIGHT_DECAY = 1e-2\nVAL_SPLIT    = 0.1\nTEST_SPLIT   = 0.1\nGRAD_CLIP    = 1.0\nNUM_WORKERS  = 2\nUSE_AMP      = torch.cuda.is_available()\nNUM_CLASSES  = 2\nSOTA_ACC     = 98.81\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 split : {int((1-(VAL_SPLIT+TEST_SPLIT))*100)}% / {int(VAL_SPLIT*100)}% / {int(TEST_SPLIT*100)}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:44:47.924037Z","iopub.execute_input":"2026-05-09T12:44:47.924511Z","iopub.status.idle":"2026-05-09T12:44:47.941050Z","shell.execute_reply.started":"2026-05-09T12:44:47.924487Z","shell.execute_reply":"2026-05-09T12:44:47.940291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 5bis : Détection du dataset RSNA\n# Structure RSNA sur Kaggle :\n#   /kaggle/input/competitions/rsna-pneumonia-detection-challenge/\n#       stage_2_train_images/   ← fichiers DICOM (.dcm)\n#       stage_2_train_labels.csv  ← colonnes : patientId, x, y, width, height, Target\n#           Target = 0  →  Normal\n#           Target = 1  →  Lung Opacity (Pneumonia)\n \ndef find_rsna_paths():\n    \"\"\"Détecte automatiquement les chemins RSNA quelle que soit la profondeur.\"\"\"\n    base = Path('/kaggle/input')\n    images_dir, labels_csv = None, None\n    for p in sorted(base.rglob('*')):\n        if p.is_dir() and 'train_images' in p.name and images_dir is None:\n            images_dir = p\n        if p.is_file() and 'train_labels' in p.name and p.suffix == '.csv':\n            labels_csv = p\n    return images_dir, labels_csv\n \nIMAGES_DIR, LABELS_CSV = find_rsna_paths()\n \nprint(f'IMAGES_DIR : {IMAGES_DIR}')\nprint(f'LABELS_CSV : {LABELS_CSV}')\n \nif IMAGES_DIR is None:\n    raise FileNotFoundError(\"Dossier stage_2_train_images introuvable. \"\n                            \"Vérifiez que le dataset RSNA est bien ajouté.\")\nif LABELS_CSV is None:\n    raise FileNotFoundError(\"stage_2_train_labels.csv introuvable.\")\n \n# Lecture et déduplication du CSV  (1 ligne par patient, Target = max des bboxs)\nlabels_raw = pd.read_csv(LABELS_CSV)\nprint(f'\\nCSV brut : {labels_raw.shape}')\nprint(labels_raw.head())\n \npatient_labels = (\n    labels_raw\n    .groupby('patientId', as_index=False)['Target']\n    .max()\n)\nprint(f'\\nPatients uniques : {len(patient_labels)}')\nprint(patient_labels['Target'].value_counts().rename({0: 'Normal', 1: 'Lung Opacity'}))\n \n# Classes RSNA  (mêmes indices que le dataset COVID original)\nCLASS_NAMES_RSNA = {0: 'Normal', 1: 'Lung_Opacity'}\n \n# ── Lecture DICOM → BGR numpy (même format qu'attendu par preprocess_base) ────\ndef read_dcm_as_bgr(dcm_path: str) -> np.ndarray:\n    \"\"\"\n    Lit un fichier DICOM et retourne une image BGR uint8 [H, W, 3].\n    Compatible avec preprocess_base() qui attend img_bgr.\n    \"\"\"\n    ds  = pydicom.dcmread(dcm_path)\n    arr = ds.pixel_array.astype(np.float32)\n    # Normalise vers [0, 255]\n    arr -= arr.min()\n    if arr.max() > 0:\n        arr /= arr.max()\n    arr = (arr * 255).astype(np.uint8)\n    # Grayscale → BGR 3 canaux (requis par cv2 + preprocess_base)\n    bgr = cv2.cvtColor(arr, cv2.COLOR_GRAY2BGR)\n    return bgr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:44:47.942601Z","iopub.execute_input":"2026-05-09T12:44:47.942911Z","iopub.status.idle":"2026-05-09T12:46:06.560564Z","shell.execute_reply.started":"2026-05-09T12:44:47.942889Z","shell.execute_reply":"2026-05-09T12:46:06.559913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 5 : Preprocessing (inchangé)\nSPIKE_STEPS = 8\n \ndef preprocess_base(img_bgr: np.ndarray, img_size: int = 224) -> np.ndarray:\n    img  = cv2.resize(img_bgr, (img_size, img_size))\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n \n    norm = gray.astype(np.float32) / 255.0\n    ch0  = norm\n    ch1  = cv2.equalizeHist(gray).astype(np.float32) / 255.0\n    blur = cv2.GaussianBlur(gray, (3, 3), 0).astype(np.float32) / 255.0\n \n    x = np.stack([ch0, ch1, blur], axis=-1)\n    return x.astype(np.float32)\n \ndef rate_encode_spikes(x: torch.Tensor, num_steps: int = SPIKE_STEPS) -> torch.Tensor:\n    spikes = []\n    for _ in range(num_steps):\n        spikes.append(torch.bernoulli(x))\n    return torch.stack(spikes, dim=0)\n \ndef visualize_preprocessing(img_bgr: np.ndarray):\n    \"\"\"Accepte directement un tableau BGR (plus de chemin fichier).\"\"\"\n    x = preprocess_base(img_bgr, IMG_SIZE)\n \n    fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n    axes[0].imshow(cv2.cvtColor(cv2.resize(img_bgr, (IMG_SIZE, IMG_SIZE)), cv2.COLOR_BGR2RGB))\n    axes[0].set_title('Original')\n    axes[1].imshow(x[:, :, 0], cmap='gray'); axes[1].set_title('Normalized')\n    axes[2].imshow(x[:, :, 1], cmap='gray'); axes[2].set_title('Equalized')\n    axes[3].imshow(x[:, :, 2], cmap='gray'); axes[3].set_title('Blurred')\n    for ax in axes: ax.axis('off')\n    plt.tight_layout(); plt.show()\n \nprint('Preprocessing défini : resize + normalization + spike encoding')\n \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:06.561599Z","iopub.execute_input":"2026-05-09T12:46:06.562043Z","iopub.status.idle":"2026-05-09T12:46:06.571965Z","shell.execute_reply.started":"2026-05-09T12:46:06.562018Z","shell.execute_reply":"2026-05-09T12:46:06.571134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 6 : Augmentations  (inchangées)\nMEAN = [0.5, 0.5, 0.5]\nSTD  = [0.5, 0.5, 0.5]\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(shift_limit=0.05, scale_limit=0.08, rotate_limit=10, p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),\n        A.GaussNoise(std_range=(0.01, 0.03), p=0.2),\n        A.Normalize(mean=MEAN, std=STD, max_pixel_value=1.0),\n        ToTensorV2(),\n    ])\n \ndef get_val_transforms():\n    return A.Compose([\n        A.Resize(IMG_SIZE, IMG_SIZE),\n        A.Normalize(mean=MEAN, std=STD, max_pixel_value=1.0),\n        ToTensorV2(),\n    ])\n \nprint('Augmentations définies : resize + normalization + augmentation')\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:06.572937Z","iopub.execute_input":"2026-05-09T12:46:06.573265Z","iopub.status.idle":"2026-05-09T12:46:06.595925Z","shell.execute_reply.started":"2026-05-09T12:46:06.573237Z","shell.execute_reply":"2026-05-09T12:46:06.595119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 7 : Dataset + Split 80/10/10 + Spike Encoding (Version RSNA)\nclass RSNASpikeDataset(Dataset):\n    \"\"\"\n    Même interface que COVIDSpikeDataset.\n    Lit les fichiers DICOM via read_dcm_as_bgr() puis applique\n    le même pipeline preprocess_base + spike encoding.\n    \"\"\"\n    CLASSES = {'Normal': 0, 'Lung_Opacity': 1}   # ← noms RSNA\n \n    def __init__(self, samples, transform=None, spike_steps=SPIKE_STEPS):\n        self.samples    = samples       # liste de (dcm_path_str, label_int)\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        img_bgr = read_dcm_as_bgr(path)          # ← lecture DICOM au lieu de cv2.imread\n        x = preprocess_base(img_bgr, IMG_SIZE)\n \n        if self.transform:\n            x = self.transform(image=x)['image']\n \n        x = (x * 0.5) + 0.5\n        x = torch.clamp(x, 0.0, 1.0)\n        spikes = rate_encode_spikes(x, self.spike_steps)\n        return spikes, label\n \n \n# Construction de all_samples depuis le CSV RSNA\nall_samples = []\nskipped     = 0\n \nfor _, row in patient_labels.iterrows():\n    pid    = row['patientId']\n    target = int(row['Target'])\n    dcm    = IMAGES_DIR / f'{pid}.dcm'\n \n    if dcm.exists():\n        all_samples.append((str(dcm), target))\n    else:\n        skipped += 1\n \nprint(f'\\nTotal images trouvées : {len(all_samples)}  (ignorées : {skipped})')\nprint(f\"  Normal       : {sum(1 for _, l in all_samples if l == 0)}\")\nprint(f\"  Lung Opacity : {sum(1 for _, l in all_samples if l == 1)}\")\n \nif len(all_samples) == 0:\n    raise FileNotFoundError('Aucun fichier DICOM trouvé dans ' + str(IMAGES_DIR))\nif len(set(l for _, l in all_samples)) < 2:\n    raise ValueError('Au moins 2 classes requises.')\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:06.596924Z","iopub.execute_input":"2026-05-09T12:46:06.597270Z","iopub.status.idle":"2026-05-09T12:46:42.098629Z","shell.execute_reply.started":"2026-05-09T12:46:06.597227Z","shell.execute_reply":"2026-05-09T12:46:42.097935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split Train / Val / Test  80 / 10 / 10  (inchangé)\nall_labels = [l for _, l in all_samples]\n \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 \nval_samples, test_samples, _, _ = train_test_split(\n    temp_samples, temp_labels,\n    test_size=0.5,\n    random_state=SEED,\n    stratify=temp_labels\n)\n \n# Créer les datasets  (RSNASpikeDataset à la place de COVIDSpikeDataset)\ntrain_ds = RSNASpikeDataset(train_samples, get_train_transforms())\nval_ds   = RSNASpikeDataset(val_samples,   get_val_transforms())\ntest_ds  = RSNASpikeDataset(test_samples,  get_val_transforms())\n \n# DataLoaders  (inchangés)\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_train_n = sum(1 for _, l in train_samples if l == 0)\nn_train_v = sum(1 for _, l in train_samples if l == 1)\nn_val_n   = sum(1 for _, l in val_samples   if l == 0)\nn_val_v   = sum(1 for _, l in val_samples   if l == 1)\nn_test_n  = sum(1 for _, l in test_samples  if l == 0)\nn_test_v  = sum(1 for _, l in test_samples  if l == 1)\n \nprint(f'\\nTrain : {len(train_samples)} (Normal:{n_train_n}, Lung Opacity:{n_train_v})')\nprint(f'Val   : {len(val_samples)}   (Normal:{n_val_n},   Lung Opacity:{n_val_v})')\nprint(f'Test  : {len(test_samples)}  (Normal:{n_test_n},  Lung Opacity:{n_test_v})')\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:42.099643Z","iopub.execute_input":"2026-05-09T12:46:42.099994Z","iopub.status.idle":"2026-05-09T12:46:42.143312Z","shell.execute_reply.started":"2026-05-09T12:46:42.099971Z","shell.execute_reply":"2026-05-09T12:46:42.142519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 8 : Visualisation Preprocessing + Distribution\nfor label_name, label_id in [('Normal', 0), ('Lung Opacity', 1)]:\n    example_path = next(p for p, l in all_samples if l == label_id)\n    print(f'Exemple : {label_name}')\n    visualize_preprocessing(read_dcm_as_bgr(example_path))   # ← lecture DICOM\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(axes,\n                              [counts_train, counts_val, counts_test],\n                              ['Distribution Train', 'Distribution Val', 'Distribution Test']):\n    ax.bar(['NORMAL', 'LUNG OPACITY'], 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 + max(counts) * 0.02, str(v), ha='center', fontweight='bold')\n \nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:42.144238Z","iopub.execute_input":"2026-05-09T12:46:42.144547Z","iopub.status.idle":"2026-05-09T12:46:43.458635Z","shell.execute_reply.started":"2026-05-09T12:46:42.144526Z","shell.execute_reply":"2026-05-09T12:46:43.458011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 9 : CNN + GRU + SNN + Attention-Guided  (inchangé)\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)\n        return x * attn, attn\n \nclass CNNEncoder(nn.Module):\n    def __init__(self, in_channels=3):\n        super().__init__()\n        self.features = nn.Sequential(\n            nn.Conv2d(in_channels, 32, 3, padding=1),\n            nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n \n            nn.Conv2d(32, 64, 3, padding=1),\n            nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n \n            nn.Conv2d(64, 128, 3, padding=1),\n            nn.BatchNorm2d(128), nn.ReLU(inplace=True), 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)\n        return x, attn\n \nclass CNN_GRU_SNN_Attention(nn.Module):\n    def __init__(self, num_classes=2, spike_steps=SPIKE_STEPS, hidden_size=128):\n        super().__init__()\n        self.spike_steps = spike_steps\n        self.encoder = CNNEncoder(in_channels=3)\n        self.gru = nn.GRU(input_size=128, hidden_size=hidden_size,\n                          num_layers=1, batch_first=True, bidirectional=True)\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        B, T, C, H, W = x.shape\n        seq_feats, attn_maps = [], []\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        seq_feats  = torch.stack(seq_feats, dim=1)\n        gru_out, _ = self.gru(seq_feats)\n        temporal_feat = gru_out[:, -1, :]\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:43.460684Z","iopub.execute_input":"2026-05-09T12:46:43.461021Z","iopub.status.idle":"2026-05-09T12:46:43.471423Z","shell.execute_reply.started":"2026-05-09T12:46:43.460997Z","shell.execute_reply":"2026-05-09T12:46:43.470859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 10 : Instanciation du modèle  (inchangé)\nmodel = CNN_GRU_SNN_Attention(\n    num_classes=NUM_CLASSES,\n    spike_steps=SPIKE_STEPS,\n    hidden_size=128\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, 3, {IMG_SIZE}, {IMG_SIZE}]')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:43.472300Z","iopub.execute_input":"2026-05-09T12:46:43.472634Z","iopub.status.idle":"2026-05-09T12:46:44.107693Z","shell.execute_reply.started":"2026-05-09T12:46:43.472598Z","shell.execute_reply":"2026-05-09T12:46:44.106977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 11 : Loss + Optimizer + Scheduler  (inchangé)\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}, Lung Opacity={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-09T12:46:44.108660Z","iopub.execute_input":"2026-05-09T12:46:44.109025Z","iopub.status.idle":"2026-05-09T12:46:44.133805Z","shell.execute_reply.started":"2026-05-09T12:46:44.109000Z","shell.execute_reply":"2026-05-09T12:46:44.133014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 12 : Fonctions Train / Validate  (inchangées)\ndef 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        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.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@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\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:44.134805Z","iopub.execute_input":"2026-05-09T12:46:44.135128Z","iopub.status.idle":"2026-05-09T12:46:44.144473Z","shell.execute_reply.started":"2026-05-09T12:46:44.135095Z","shell.execute_reply":"2026-05-09T12:46:44.143917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 13 : Boucle d'entraînement  (inchangée)\nos.makedirs('checkpoints', exist_ok=True)\n \nPATIENCE = 6\nbest_acc, best_state = 0.0, None\nhistory = []\npatience_ctr = 0\n \nprint(f'Entraînement : {EPOCHS} epochs  (patience={PATIENCE})')\nprint(f'   Modèle     : CNN + GRU + SNN + Attention')\nprint(f'   Input      : Spike encoding + resize + normalization + augmentation')\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({'epoch': epoch+1, 'tr_loss': tr_loss, 'tr_acc': tr_acc,\n                    'vl_loss': vl_loss, 'vl_acc': vl_acc, 'vl_f1': vl_f1, 'vl_auc': vl_auc})\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_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'Entraînement terminé. Meilleure Val Acc = {best_acc*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T12:46:44.145386Z","iopub.execute_input":"2026-05-09T12:46:44.145755Z","iopub.status.idle":"2026-05-09T13:26:49.765491Z","shell.execute_reply.started":"2026-05-09T12:46:44.145729Z","shell.execute_reply":"2026-05-09T13:26:49.764650Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 14 : Courbes d'apprentissage  (inchangées)\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 — RSNA Pneumonia Dataset',\n             fontsize=13, fontweight='bold')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T13:26:49.766914Z","iopub.execute_input":"2026-05-09T13:26:49.767527Z","iopub.status.idle":"2026-05-09T13:26:50.271908Z","shell.execute_reply.started":"2026-05-09T13:26:49.767496Z","shell.execute_reply":"2026-05-09T13:26:50.271216Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 15 : Évaluation Finale (Test Set)  (inchangée)\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('RÉSULTATS FINAUX')\nprint('Modèle  : CNN + GRU + SNN + Attention-Guided')\nprint('Dataset : RSNA Pneumonia Detection Challenge')\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-09T13:26:50.272900Z","iopub.execute_input":"2026-05-09T13:26:50.273247Z","iopub.status.idle":"2026-05-09T13:27:48.091863Z","shell.execute_reply.started":"2026-05-09T13:26:50.273223Z","shell.execute_reply":"2026-05-09T13:27:48.090992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 16 : Matrice de Confusion + Courbe ROC  (inchangées)\ncm = confusion_matrix(all_true, all_preds)\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', 'Lung Opacity'],\n            yticklabels=['Normal', 'Lung Opacity'])\naxes[0].set_ylabel('Vrai label')\naxes[0].set_xlabel('Prédiction')\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'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()\n \nprint(classification_report(all_true, all_preds,\n                            target_names=['NORMAL', 'LUNG OPACITY']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T13:27:48.093280Z","iopub.execute_input":"2026-05-09T13:27:48.093623Z","iopub.status.idle":"2026-05-09T13:27:48.429870Z","shell.execute_reply.started":"2026-05-09T13:27:48.093584Z","shell.execute_reply":"2026-05-09T13:27:48.429206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 17 : Tableau récapitulatif  (inchangé)\nprint(f'{\"Model\":<30} {\"Accuracy\":>10} {\"F1\":>10} {\"Precision\":>12} {\"Recall\":>10} {\"AUC\":>10}')\nprint('-' * 86)\nprint(f'{\"CNN+GRU+SNN+Attention\":<30} {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-09T13:27:48.430860Z","iopub.execute_input":"2026-05-09T13:27:48.431402Z","iopub.status.idle":"2026-05-09T13:27:48.435886Z","shell.execute_reply.started":"2026-05-09T13:27:48.431370Z","shell.execute_reply":"2026-05-09T13:27:48.435103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -- Cellule 18 : Sauvegarde finale  (inchangée)\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    'preprocessing' : 'spike encoding + resize + normalization + augmentation',\n    'split'         : '80/10/10 train/val/test',\n    'architecture'  : 'CNN + GRU + SNN + Attention-Guided',\n    'dataset'       : 'RSNA Pneumonia Detection Challenge',\n    'num_classes'   : NUM_CLASSES,\n}, 'checkpoints/cnn_gru_snn_attention_final.pth')\n \nprint('Modèle sauvegardé -> checkpoints/cnn_gru_snn_attention_final.pth')\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}%')\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-09T13:27:48.436816Z","iopub.execute_input":"2026-05-09T13:27:48.437122Z","iopub.status.idle":"2026-05-09T13:27:48.460177Z","shell.execute_reply.started":"2026-05-09T13:27:48.437091Z","shell.execute_reply":"2026-05-09T13:27:48.459315Z"}},"outputs":[],"execution_count":null}]}