{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"7b768955-39ac-4632-b50b-362a3c0e4236","cell_type":"markdown","source":"# Pneumonia Detection - Kermany Dataset\n## Algorithme : Ensemble(ResNet50 + GoogLeNet) - Multi-channel Preprocessing - Train/Val Split\n","metadata":{}},{"id":"5c5074b5-c0b8-467f-ad82-9329aea59d57","cell_type":"code","source":"# -- Cellule 1 : Installation\n!pip install albumentations scikit-learn tqdm pandas opencv-python-headless -q\nprint(\"OK deps\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:20.175972Z","iopub.execute_input":"2026-03-08T21:36:20.176751Z","iopub.status.idle":"2026-03-08T21:36:23.465947Z","shell.execute_reply.started":"2026-03-08T21:36:20.176689Z","shell.execute_reply":"2026-03-08T21:36:23.465131Z"}},"outputs":[],"execution_count":null},{"id":"c6ccd9f9-aaf3-42fc-9016-c655f6ae521d","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 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 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-03-08T21:36:23.467718Z","iopub.execute_input":"2026-03-08T21:36:23.467977Z","iopub.status.idle":"2026-03-08T21:36:23.475754Z","shell.execute_reply.started":"2026-03-08T21:36:23.467951Z","shell.execute_reply":"2026-03-08T21:36:23.475021Z"}},"outputs":[],"execution_count":null},{"id":"81d94ab0-7a0d-45aa-bf31-fc46974ceda7","cell_type":"code","source":"# -- Cellule 3 : Telechargement du Dataset\nimport zipfile\n\nDATA_DIR = Path('data/chest_xray')\nif not DATA_DIR.exists():\n    print('Downloading from Kaggle...')\n    os.system('kaggle datasets download -d paultimothymooney/chest-xray-pneumonia')\n    with zipfile.ZipFile('chest-xray-pneumonia.zip', 'r') as z:\n        z.extractall(DATA_DIR)\n    os.remove('chest-xray-pneumonia.zip')\n    print('Done!')\nelse:\n    print('Dataset already present')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:23.476688Z","iopub.execute_input":"2026-03-08T21:36:23.477011Z","iopub.status.idle":"2026-03-08T21:36:23.493001Z","shell.execute_reply.started":"2026-03-08T21:36:23.476991Z","shell.execute_reply":"2026-03-08T21:36:23.492433Z"}},"outputs":[],"execution_count":null},{"id":"29bc8c63-5d6a-40d5-acc0-407318bee671","cell_type":"code","source":"# -- Cellule 4 : Configuration\nSEED         = 42\nDATA_DIR     = Path('data/chest_xray/chest_xray')\n\nIMG_SIZE     = 224\nBATCH_SIZE   = 32\nEPOCHS       = 25\nLR           = 3e-4\nWEIGHT_DECAY = 1e-2\nVAL_SPLIT    = 0.2        # 80% train / 20% val\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 split : {int((1-VAL_SPLIT)*100)}% / {int(VAL_SPLIT*100)}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:23.493945Z","iopub.execute_input":"2026-03-08T21:36:23.494232Z","iopub.status.idle":"2026-03-08T21:36:23.508696Z","shell.execute_reply.started":"2026-03-08T21:36:23.494204Z","shell.execute_reply":"2026-03-08T21:36:23.508115Z"}},"outputs":[],"execution_count":null},{"id":"7b0dd002-aecb-4a94-9109-7228cc2ae40d","cell_type":"code","source":"# -- Cellule 5 : Multi-Channel Preprocessing\n#\n# Chaque image radiologique est transformee en un tenseur 3 canaux :\n#   Canal 0 : Image originale normalisee (information globale)\n#   Canal 1 : CLAHE (Contrast Limited Adaptive Histogram Equalization)\n#             -> rehausse les details locaux (opacites pulmonaires)\n#   Canal 2 : Gradient de Sobel (magnitude des contours)\n#             -> met en evidence les bordures des infiltrats\n\ndef multichannel_preprocess(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    # Canal 0 : image grise normalisee\n    ch0 = gray.astype(np.float32) / 255.0\n\n    # Canal 1 : CLAHE\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n    ch1   = clahe.apply(gray).astype(np.float32) / 255.0\n\n    # Canal 2 : Gradient Sobel\n    sobelx = cv2.Sobel(gray, cv2.CV_64F, 1, 0, ksize=3)\n    sobely = cv2.Sobel(gray, cv2.CV_64F, 0, 1, ksize=3)\n    mag    = np.sqrt(sobelx**2 + sobely**2)\n    if mag.max() > 0:\n        mag = mag / mag.max()\n    ch2 = mag.astype(np.float32)\n\n    return np.stack([ch0, ch1, ch2], axis=-1)  # [H, W, 3]\n\n\ndef visualize_multichannel(img_path: str):\n    img = cv2.imread(img_path)\n    mc  = multichannel_preprocess(img, IMG_SIZE)\n    orig = cv2.cvtColor(cv2.resize(img, (IMG_SIZE, IMG_SIZE)), cv2.COLOR_BGR2RGB)\n\n    fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n    images = [orig, mc[:,:,0], mc[:,:,1], mc[:,:,2]]\n    titles = ['Original', 'Canal 0\\n(Normalise)', 'Canal 1\\n(CLAHE)', 'Canal 2\\n(Sobel)']\n    cmaps  = [None, 'gray', 'gray', 'hot']\n    for ax, im, t, c in zip(axes, images, titles, cmaps):\n        ax.imshow(im, cmap=c)\n        ax.set_title(t, fontsize=11, fontweight='bold')\n        ax.axis('off')\n    plt.suptitle('Multi-Channel Preprocessing', fontsize=13, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\n\n\nprint('Multi-Channel Preprocessing defini')\nprint('   Canal 0 : Image normalisee')\nprint('   Canal 1 : CLAHE (contraste local adaptatif)')\nprint('   Canal 2 : Gradient Sobel (contours/infiltrats)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:23.510537Z","iopub.execute_input":"2026-03-08T21:36:23.510790Z","iopub.status.idle":"2026-03-08T21:36:23.532764Z","shell.execute_reply.started":"2026-03-08T21:36:23.510771Z","shell.execute_reply":"2026-03-08T21:36:23.532101Z"}},"outputs":[],"execution_count":null},{"id":"bc3363e3-50f4-4a7c-a96c-9127a0203283","cell_type":"code","source":"# -- Cellule 6 : Augmentations\n# IMPORTANT : max_pixel_value=1.0 car nos images multi-canal sont\n# en float32 [0,1] et non uint8 [0,255].\n# Par defaut albumentations suppose max_pixel_value=255 ce qui fausse\n# completement la normalisation sur des float32 -> collapse du modele.\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.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1,\n                           rotate_limit=10, p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.OneOf([\n            A.GaussNoise(var_limit=(0.001, 0.01)),\n            A.GaussianBlur(blur_limit=(3, 5)),\n        ], p=0.3),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.4),\n        A.CoarseDropout(max_holes=6, max_height=IMG_SIZE//16,\n                        max_width=IMG_SIZE//16, p=0.3),\n        A.Normalize(mean=MEAN, std=STD, max_pixel_value=1.0),  # fix float32\n        ToTensorV2(),\n    ])\n\ndef get_val_transforms():\n    return A.Compose([\n        A.Normalize(mean=MEAN, std=STD, max_pixel_value=1.0),  # fix float32\n        ToTensorV2(),\n    ])\n\nprint('Augmentations definies (max_pixel_value=1.0 pour float32)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:23.534386Z","iopub.execute_input":"2026-03-08T21:36:23.534667Z","iopub.status.idle":"2026-03-08T21:36:23.550028Z","shell.execute_reply.started":"2026-03-08T21:36:23.534648Z","shell.execute_reply":"2026-03-08T21:36:23.549404Z"}},"outputs":[],"execution_count":null},{"id":"3bb5c34d-8c07-4de5-8848-9e069075f17e","cell_type":"code","source":"# -- Cellule 7 : Dataset Multi-Channel + Train/Val Split\nclass KermanyMultiChannelDataset(Dataset):\n    CLASSES = {'NORMAL': 0, 'PNEUMONIA': 1}\n\n    def __init__(self, samples, transform=None):\n        self.samples   = samples\n        self.transform = transform\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 = cv2.imread(path)\n        mc = multichannel_preprocess(img_bgr, IMG_SIZE)  # [H,W,3] float32\n        if self.transform:\n            mc = self.transform(image=mc)['image']\n        return mc, label\n\n\n# Charge tous les echantillons (train + val + test officiel -> on refait le split)\nall_samples = []\nfor split in ['train', 'val', 'test']:\n    for cls_name, label in KermanyMultiChannelDataset.CLASSES.items():\n        cls_dir = DATA_DIR / split / cls_name\n        if cls_dir.exists():\n            for p in cls_dir.glob('*'):\n                if p.suffix.lower() in ('.jpg', '.jpeg', '.png'):\n                    all_samples.append((str(p), label))\n\nall_labels = [l for _, l in all_samples]\nn_normal   = sum(1 for l in all_labels if l == 0)\nn_pneumo   = sum(1 for l in all_labels if l == 1)\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# -- Train / Val Split stratifie --\ntrain_samples, val_samples = train_test_split(\n    all_samples,\n    test_size    = VAL_SPLIT,\n    random_state = SEED,\n    stratify     = all_labels,\n)\n\ntrain_ds = KermanyMultiChannelDataset(train_samples, get_train_transforms())\nval_ds   = KermanyMultiChannelDataset(val_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)\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)\nprint(f'Train : {len(train_samples)} ({n_tn} Normal, {n_tp} Pneumonia)')\nprint(f'Val   : {len(val_samples)}   ({n_vn} Normal, {n_vp} Pneumonia)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:23.550966Z","iopub.execute_input":"2026-03-08T21:36:23.551252Z","iopub.status.idle":"2026-03-08T21:36:23.603969Z","shell.execute_reply.started":"2026-03-08T21:36:23.551222Z","shell.execute_reply":"2026-03-08T21:36:23.603457Z"}},"outputs":[],"execution_count":null},{"id":"896ba35f-4283-48da-84af-38870a7da1e0","cell_type":"code","source":"# -- Cellule 8 : Visualisation Multi-Channel\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'Exemple : {label_name}')\n    visualize_multichannel(example_path)\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 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]]\n\naxes[0].bar(['NORMAL', 'PNEUMONIA'], counts_train, color=['steelblue', 'tomato'])\naxes[0].set_title('Distribution Train', fontweight='bold')\naxes[0].set_ylabel('Images')\nfor i, v in enumerate(counts_train):\n    axes[0].text(i, v + 20, str(v), ha='center', fontweight='bold')\n\naxes[1].bar(['NORMAL', 'PNEUMONIA'], counts_val, color=['steelblue', 'tomato'])\naxes[1].set_title('Distribution Val', fontweight='bold')\nfor i, v in enumerate(counts_val):\n    axes[1].text(i, v + 5, str(v), ha='center', fontweight='bold')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:23.604727Z","iopub.execute_input":"2026-03-08T21:36:23.604988Z","iopub.status.idle":"2026-03-08T21:36:24.633666Z","shell.execute_reply.started":"2026-03-08T21:36:23.604969Z","shell.execute_reply":"2026-03-08T21:36:24.633074Z"}},"outputs":[],"execution_count":null},{"id":"2513039c-2a0f-451a-8fbf-40670fde7930","cell_type":"code","source":"# -- Cellule 9 : Modeles ResNet50 + GoogLeNet\ndef build_resnet50(num_classes=2):\n    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\n    in_f  = model.fc.in_features  # 2048\n    model.fc = nn.Sequential(\n        nn.Dropout(0.4),\n        nn.Linear(in_f, 512),\n        nn.GELU(),\n        nn.Dropout(0.2),\n        nn.Linear(512, num_classes),\n    )\n    return model\n\n\ndef build_googlenet(num_classes=2):\n    # torchvision force aux_logits=True au chargement des poids pretrained.\n    # Solution : charger avec aux_logits=True, puis desactiver apres.\n    model = models.googlenet(\n        weights    = models.GoogLeNet_Weights.IMAGENET1K_V1,\n        aux_logits = True,   # requis pour charger le checkpoint\n    )\n    # Desactive les branches auxiliaires : on ne les utilise pas\n    model.aux_logits = False\n    model.aux1       = None\n    model.aux2       = None\n\n    in_f  = model.fc.in_features  # 1024\n    model.fc = nn.Sequential(\n        nn.Dropout(0.4),\n        nn.Linear(in_f, 512),\n        nn.GELU(),\n        nn.Dropout(0.2),\n        nn.Linear(512, num_classes),\n    )\n    return model\n\n\nresnet    = build_resnet50(NUM_CLASSES).to(DEVICE)\ngooglenet = build_googlenet(NUM_CLASSES).to(DEVICE)\n\nn_r = sum(p.numel() for p in resnet.parameters()    if p.requires_grad)\nn_g = sum(p.numel() for p in googlenet.parameters() if p.requires_grad)\nprint(f'ResNet-50  : {n_r/1e6:.1f}M params')\nprint(f'GoogLeNet  : {n_g/1e6:.1f}M params')\nprint(f'Input      : [B, 3, {IMG_SIZE}, {IMG_SIZE}]  (multi-canal)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:24.634530Z","iopub.execute_input":"2026-03-08T21:36:24.634838Z","iopub.status.idle":"2026-03-08T21:36:25.588418Z","shell.execute_reply.started":"2026-03-08T21:36:24.634815Z","shell.execute_reply":"2026-03-08T21:36:25.587763Z"}},"outputs":[],"execution_count":null},{"id":"a4a97a0c-626b-46af-a33f-d1b910855abb","cell_type":"code","source":"# -- Cellule 10 : Ensemble Model\nclass EnsembleModel(nn.Module):\n    def __init__(self, resnet, googlenet, fusion='mean', num_classes=2):\n        super().__init__()\n        self.resnet    = resnet\n        self.googlenet = googlenet\n        self.fusion    = fusion\n        if fusion == 'learned':\n            self.meta = nn.Sequential(\n                nn.Linear(num_classes * 2, 64),\n                nn.ReLU(),\n                nn.Linear(64, num_classes),\n            )\n\n    def forward(self, x):\n        l_r = self.resnet(x)\n        l_g = self.googlenet(x)\n        if self.fusion == 'mean':\n            return (l_r + l_g) / 2.0\n        elif self.fusion == 'learned':\n            return self.meta(torch.cat([l_r, l_g], dim=1))\n\n\nensemble = EnsembleModel(resnet, googlenet, fusion='mean').to(DEVICE)\nprint(f'Ensemble(ResNet50 + GoogLeNet) cree')\nprint(f'   Fusion : mean')\nprint(f'   Total  : {(n_r + n_g)/1e6:.1f}M params')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:25.589370Z","iopub.execute_input":"2026-03-08T21:36:25.589892Z","iopub.status.idle":"2026-03-08T21:36:25.600308Z","shell.execute_reply.started":"2026-03-08T21:36:25.589868Z","shell.execute_reply":"2026-03-08T21:36:25.599734Z"}},"outputs":[],"execution_count":null},{"id":"236ac9f9-94bf-479b-b328-ce328845efca","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.05)\n\n# Optimiseurs separes : chaque backbone a son propre LR\noptimizer_r = AdamW(resnet.parameters(),    lr=LR, weight_decay=WEIGHT_DECAY)\noptimizer_g = AdamW(googlenet.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n\nscheduler_r = CosineAnnealingLR(optimizer_r, T_max=EPOCHS, eta_min=1e-6)\nscheduler_g = CosineAnnealingLR(optimizer_g, T_max=EPOCHS, eta_min=1e-6)\n\nscaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n\nprint(f'CrossEntropyLoss (label_smoothing=0.05, weighted)')\nprint(f'AdamW x2  (lr={LR}, wd={WEIGHT_DECAY})')\nprint(f'CosineAnnealingLR (T_max={EPOCHS})')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:25.601135Z","iopub.execute_input":"2026-03-08T21:36:25.601378Z","iopub.status.idle":"2026-03-08T21:36:25.620955Z","shell.execute_reply.started":"2026-03-08T21:36:25.601358Z","shell.execute_reply":"2026-03-08T21:36:25.620311Z"}},"outputs":[],"execution_count":null},{"id":"02891c2f-15ec-4f53-9561-6ffcf19c3a78","cell_type":"code","source":"# -- Cellule 12 : Fonctions Train / Validate\ndef train_epoch(resnet, googlenet, loader, opt_r, opt_g, criterion, scaler):\n    resnet.train()\n    googlenet.train()\n    total_loss = 0.0\n    preds_all, labels_all = [], []\n\n    pbar = tqdm(loader, desc='  [TRAIN]', leave=False)\n    for imgs, labels in pbar:\n        imgs   = imgs.float().to(DEVICE, non_blocking=True)\n        labels = labels.long().to(DEVICE, non_blocking=True)\n\n        opt_r.zero_grad(set_to_none=True)\n        opt_g.zero_grad(set_to_none=True)\n\n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits_r = resnet(imgs)\n            logits_g = googlenet(imgs)\n            logits_e = (logits_r + logits_g) / 2.0\n            # Loss combineee : individuelle + ensemble\n            loss = (0.3 * criterion(logits_r, labels)\n                  + 0.3 * criterion(logits_g, labels)\n                  + 0.4 * criterion(logits_e, labels))\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(opt_r)\n        scaler.unscale_(opt_g)\n        nn.utils.clip_grad_norm_(resnet.parameters(),    GRAD_CLIP)\n        nn.utils.clip_grad_norm_(googlenet.parameters(), GRAD_CLIP)\n        scaler.step(opt_r)\n        scaler.step(opt_g)\n        scaler.update()\n\n        total_loss += loss.item()\n        preds_all.append(logits_e.argmax(1).detach())\n        labels_all.append(labels.detach())\n        pbar.set_postfix(loss=f'{loss.item():.4f}')\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(resnet, googlenet, loader, criterion):\n    resnet.eval()\n    googlenet.eval()\n    total_loss, preds_all, labels_all, probs_all = 0.0, [], [], []\n\n    for imgs, labels in loader:\n        imgs   = imgs.float().to(DEVICE)\n        labels = labels.long().to(DEVICE)\n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits_r = resnet(imgs)\n            logits_g = googlenet(imgs)\n            logits_e = (logits_r + logits_g) / 2.0\n            loss     = criterion(logits_e, labels)\n        probs = F.softmax(logits_e, dim=1)[:, 1]\n        total_loss += loss.item()\n        preds_all.extend(logits_e.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    try:    auc = roc_auc_score(labels_all, probs_all)\n    except: auc = 0.0\n    return total_loss / len(loader), acc, f1, auc\n\n\nprint('train_epoch / validate definis')\nprint('   Loss = 0.3*L(ResNet) + 0.3*L(GoogLeNet) + 0.4*L(Ensemble)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:25.621822Z","iopub.execute_input":"2026-03-08T21:36:25.622214Z","iopub.status.idle":"2026-03-08T21:36:25.634514Z","shell.execute_reply.started":"2026-03-08T21:36:25.622171Z","shell.execute_reply":"2026-03-08T21:36:25.633936Z"}},"outputs":[],"execution_count":null},{"id":"ce5c6af2-9490-4e6e-ae3e-4edf44a4deab","cell_type":"code","source":"# -- Cellule 13 : Boucle d'entrainement avec Progressive Unfreezing\nos.makedirs('checkpoints', exist_ok=True)\n\nWARMUP_EPOCHS = 3    # epochs backbone gele\nPATIENCE      = 6    # early stopping\n\nbest_acc, best_state_r, best_state_g = 0.0, None, None\nhistory = []\npatience_ctr = 0\n\nprint(f'Entrainement : {EPOCHS} epochs  (warmup={WARMUP_EPOCHS}, patience={PATIENCE})')\nprint(f'   Backbone 1 : ResNet-50   ({n_r/1e6:.1f}M params)')\nprint(f'   Backbone 2 : GoogLeNet   ({n_g/1e6:.1f}M params)')\nprint(f'   Input      : Multi-channel [Normalise | CLAHE | Sobel]')\nprint(f'   Split      : Train {len(train_samples)} / Val {len(val_samples)}')\nprint('=' * 65)\n\n# -- Phase 1 : Warmup — gele les backbones, entraine seulement les tetes\ndef freeze_backbone(model):\n    for name, p in model.named_parameters():\n        if 'fc' not in name:\n            p.requires_grad = False\n\ndef unfreeze_all(model):\n    for p in model.parameters():\n        p.requires_grad = True\n\nfreeze_backbone(resnet)\nfreeze_backbone(googlenet)\nprint(f'Backbone gele pour {WARMUP_EPOCHS} epochs de warmup')\n\nfor epoch in range(EPOCHS):\n    # Degele backbone apres le warmup\n    if epoch == WARMUP_EPOCHS:\n        unfreeze_all(resnet)\n        unfreeze_all(googlenet)\n        # LR reduit pour le fine-tuning complet\n        for opt in [optimizer_r, optimizer_g]:\n            for g in opt.param_groups:\n                g['lr'] = LR * 0.1\n        print(f'Backbone degele. LR -> {LR*0.1:.1e}')\n\n    tr_loss, tr_acc = train_epoch(\n        resnet, googlenet, train_dl, optimizer_r, optimizer_g, criterion, scaler)\n    vl_loss, vl_acc, vl_f1, vl_auc = validate(\n        resnet, googlenet, val_dl, criterion)\n\n    if epoch >= WARMUP_EPOCHS:\n        scheduler_r.step()\n        scheduler_g.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    gap  = vl_acc * 100 - SOTA_ACC\n    tag  = '🏆' if gap > 0 else '  '\n    phase = 'WU' if epoch < WARMUP_EPOCHS else 'FT'\n    lr   = optimizer_r.param_groups[0]['lr']\n    print(f'{tag} [{phase}] 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={lr:.2e}  (delta={gap:+.2f}%)')\n\n    if vl_acc > best_acc:\n        best_acc     = vl_acc\n        best_state_r = {k: v.clone() for k, v in resnet.state_dict().items()}\n        best_state_g = {k: v.clone() for k, v in googlenet.state_dict().items()}\n        torch.save(best_state_r, 'checkpoints/resnet_best.pth')\n        torch.save(best_state_g, 'checkpoints/googlenet_best.pth')\n        patience_ctr = 0\n        print(f'     Nouveau record ! Acc={best_acc*100:.2f}%')\n    else:\n        if epoch >= WARMUP_EPOCHS:  # patience seulement apres warmup\n            patience_ctr += 1\n            if patience_ctr >= PATIENCE:\n                print(f'Early stopping (patience={PATIENCE})')\n                break\n\nresnet.load_state_dict(best_state_r)\ngooglenet.load_state_dict(best_state_g)\nprint(f'Entrainement termine. Meilleure Val Acc = {best_acc*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:36:25.635199Z","iopub.execute_input":"2026-03-08T21:36:25.635901Z","iopub.status.idle":"2026-03-08T21:48:52.070010Z","shell.execute_reply.started":"2026-03-08T21:36:25.635874Z","shell.execute_reply":"2026-03-08T21:48:52.069195Z"}},"outputs":[],"execution_count":null},{"id":"d967f127-3bbf-4805-bb8c-862462a25d75","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('Ensemble(ResNet50 + GoogLeNet) — Multi-Channel Preprocessing',\n             fontsize=13, fontweight='bold')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:48:52.072603Z","iopub.execute_input":"2026-03-08T21:48:52.072830Z","iopub.status.idle":"2026-03-08T21:48:52.555495Z","shell.execute_reply.started":"2026-03-08T21:48:52.072807Z","shell.execute_reply":"2026-03-08T21:48:52.554796Z"}},"outputs":[],"execution_count":null},{"id":"b4b8f880-ae52-42d0-af42-a29c0515ef2a","cell_type":"code","source":"# -- Cellule 15 : Evaluation Finale (Val Set)\nresnet.eval(); googlenet.eval()\nall_preds, all_probs, all_true = [], [], []\n\nwith torch.no_grad():\n    for imgs, labels in tqdm(val_dl, desc='Eval finale'):\n        imgs = imgs.float().to(DEVICE)\n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            logits_r = resnet(imgs)\n            logits_g = googlenet(imgs)\n            logits_e = (logits_r + logits_g) / 2.0\n        probs = F.softmax(logits_e, dim=1)[:, 1]\n        all_preds.extend(logits_e.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')\nauc  = roc_auc_score(all_true, all_probs)\ngap  = acc * 100 - SOTA_ACC\n\nprint(f'RESULTATS FINAUX')\nprint(f'Ensemble : ResNet50 + GoogLeNet')\nprint(f'Preprocessing : Multi-channel [Normalise | CLAHE | Sobel]')\nprint(f'Split : {int((1-VAL_SPLIT)*100)}% Train / {int(VAL_SPLIT*100)}% Val')\nprint(f'  Accuracy  : {acc*100:.2f}%   (SOTA : {SOTA_ACC}%)')\nprint(f'  F1-Score  : {f1*100:.2f}%')\nprint(f'  Precision : {prec*100:.2f}%')\nprint(f'  Rappel    : {rec*100:.2f}%')\nprint(f'  AUC-ROC   : {auc*100:.2f}%')\nif gap > 0:\n    print(f'  SOTA BATTU ! delta = +{gap:.2f}%')\nelse:\n    print(f'  Gap vs SOTA : {gap:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:48:52.556299Z","iopub.execute_input":"2026-03-08T21:48:52.556619Z","iopub.status.idle":"2026-03-08T21:48:59.365572Z","shell.execute_reply.started":"2026-03-08T21:48:52.556595Z","shell.execute_reply":"2026-03-08T21:48:59.364871Z"}},"outputs":[],"execution_count":null},{"id":"0dc12529-0fc0-44fe-919b-926e9691f0d2","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'],\n            yticklabels=['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'Ensemble 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, target_names=['NORMAL', 'PNEUMONIA']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:48:59.366835Z","iopub.execute_input":"2026-03-08T21:48:59.367101Z","iopub.status.idle":"2026-03-08T21:48:59.726622Z","shell.execute_reply.started":"2026-03-08T21:48:59.367076Z","shell.execute_reply":"2026-03-08T21:48:59.725805Z"}},"outputs":[],"execution_count":null},{"id":"6158c6b4-4bca-4177-976d-bebbfe228ed6","cell_type":"code","source":"# -- Cellule 17 : Comparaison ResNet vs GoogLeNet vs Ensemble\nresnet.eval(); googlenet.eval()\n\nres_indiv = {'ResNet50': [], 'GoogLeNet': [], 'Ensemble': []}\ntrue_labels = []\n\nwith torch.no_grad():\n    for imgs, labels in val_dl:\n        imgs = imgs.float().to(DEVICE)\n        with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=USE_AMP):\n            lr_ = resnet(imgs)\n            lg_ = googlenet(imgs)\n            le_ = (lr_ + lg_) / 2.0\n        res_indiv['ResNet50'].extend(lr_.argmax(1).cpu().numpy())\n        res_indiv['GoogLeNet'].extend(lg_.argmax(1).cpu().numpy())\n        res_indiv['Ensemble'].extend(le_.argmax(1).cpu().numpy())\n        true_labels.extend(labels.numpy())\n\nprint(f'{\"Model\":<12} {\"Accuracy\":>10} {\"F1\":>10} {\"Precision\":>12} {\"Recall\":>10}')\nprint('-' * 58)\nfor name, preds in res_indiv.items():\n    a = accuracy_score(true_labels, preds)\n    f = f1_score(true_labels, preds, average='weighted')\n    p = precision_score(true_labels, preds, average='weighted', zero_division=0)\n    r = recall_score(true_labels, preds, average='weighted')\n    tag = ' <- best' if name == 'Ensemble' else ''\n    print(f'{name:<12} {a*100:>9.2f}% {f*100:>9.2f}% {p*100:>11.2f}% {r*100:>9.2f}%{tag}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:48:59.727714Z","iopub.execute_input":"2026-03-08T21:48:59.727999Z","iopub.status.idle":"2026-03-08T21:49:06.467064Z","shell.execute_reply.started":"2026-03-08T21:48:59.727978Z","shell.execute_reply":"2026-03-08T21:49:06.466173Z"}},"outputs":[],"execution_count":null},{"id":"493015c4-1863-4db8-82b2-aa2c78382898","cell_type":"code","source":"# -- Cellule 18 : Sauvegarde finale\ntorch.save({\n    'resnet_state'   : resnet.state_dict(),\n    'googlenet_state': googlenet.state_dict(),\n    'accuracy'       : acc,\n    'f1'             : f1,\n    'auc'            : auc,\n    'img_size'       : IMG_SIZE,\n    'preprocessing'  : 'multichannel [original, clahe, sobel]',\n    'split'          : f'{int((1-VAL_SPLIT)*100)}/{int(VAL_SPLIT*100)} train/val',\n    'fusion'         : 'mean',\n}, 'checkpoints/ensemble_final.pth')\n\nprint('Modele sauvegarde -> checkpoints/ensemble_final.pth')\nprint(f'   Accuracy : {acc*100:.2f}%')\nprint(f'   AUC-ROC  : {auc*100:.2f}%')\nprint(f'   F1-Score : {f1*100:.2f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T21:49:06.468440Z","iopub.execute_input":"2026-03-08T21:49:06.468733Z","iopub.status.idle":"2026-03-08T21:49:06.694395Z","shell.execute_reply.started":"2026-03-08T21:49:06.468707Z","shell.execute_reply":"2026-03-08T21:49:06.693715Z"}},"outputs":[],"execution_count":null}]}