{"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":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":13479521,"sourceType":"datasetVersion","datasetId":8557821},{"sourceId":14682499,"sourceType":"datasetVersion","datasetId":9379797}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Lister TOUS les datasets disponibles dans /kaggle/input/\nprint(\"📁 Datasets disponibles dans /kaggle/input/ :\")\n!ls /kaggle/input/\n\nprint(\"\\n🔍 Contenu de chaque dataset :\")\nfor dirname in os.listdir('/kaggle/input'):\n    print(f\"\\n→ {dirname} :\")\n    !ls /kaggle/input/{dirname}/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:05:28.121728Z","iopub.execute_input":"2026-02-09T14:05:28.122249Z","iopub.status.idle":"2026-02-09T14:05:28.387727Z","shell.execute_reply.started":"2026-02-09T14:05:28.122212Z","shell.execute_reply":"2026-02-09T14:05:28.387072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Imports nécessaires\nimport torch\nimport torch.nn as nn\nfrom torchvision import models, transforms\nimport timm\nimport pandas as pd\nfrom pathlib import Path\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import classification_report\nimport numpy as np\nimport os\n\n# Configuration du device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"✅ Device : {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:07:03.996478Z","iopub.execute_input":"2026-02-09T14:07:03.996875Z","iopub.status.idle":"2026-02-09T14:07:14.816986Z","shell.execute_reply.started":"2026-02-09T14:07:03.996839Z","shell.execute_reply":"2026-02-09T14:07:14.816272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HybridResNet50_DeitBasic(nn.Module):\n    def __init__(self, num_classes=5):\n        super().__init__()\n        # 1. ResNet50 (features locaux)\n        self.resnet = models.resnet50(weights='IMAGENET1K_V1')\n        self.resnet.fc = nn.Identity()  # [B, 2048]\n        \n        # 2. DeiT-Basic (contexte global)\n        self.deit = timm.create_model('deit_base_patch16_224', pretrained=True, num_classes=0)  # [B, 768]\n        \n        # 3. Fusion + tête de classification\n        self.fusion = nn.Linear(2048 + 768, 512)\n        self.head = nn.Sequential(\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    \n    def forward(self, x):\n        res_feat = self.resnet(x)      # [B, 2048]\n        deit_feat = self.deit(x)       # [B, 768]\n        fused = torch.cat([res_feat, deit_feat], dim=1)  # [B, 2816]\n        return self.head(self.fusion(fused))  # [B, 5]\n\nprint(\"✅ Architecture du modèle définie\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:07:36.589832Z","iopub.execute_input":"2026-02-09T14:07:36.590326Z","iopub.status.idle":"2026-02-09T14:07:36.59719Z","shell.execute_reply.started":"2026-02-09T14:07:36.590299Z","shell.execute_reply":"2026-02-09T14:07:36.596511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Chemin EXACT vers votre modèle (d'après votre résultat)\nMODEL_PATH = '/kaggle/input/resnet50-deit-basic-best-pthyy/resnet50_deit_basic_best.pth'\n\nprint(f\"🔍 Chemin du modèle : {MODEL_PATH}\")\nprint(f\"✅ Fichier existe ? {os.path.exists(MODEL_PATH)}\")\n\n# Instanciation + chargement\nmodel = HybridResNet50_DeitBasic().to(device)\nmodel.load_state_dict(torch.load(MODEL_PATH, map_location=device))\nmodel.eval()\n\nprint(\"✅ Modèle chargé avec succès !\")\nprint(f\"✅ Nombre de paramètres : {sum(p.numel() for p in model.parameters()):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:07:55.461483Z","iopub.execute_input":"2026-02-09T14:07:55.462026Z","iopub.status.idle":"2026-02-09T14:08:04.565547Z","shell.execute_reply.started":"2026-02-09T14:07:55.461996Z","shell.execute_reply":"2026-02-09T14:08:04.564891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Classe Dataset (identique à votre docx)\nclass APTOSDataset(Dataset):\n    def __init__(self, img_dir, csv_path, transform=None):\n        self.img_dir = Path(img_dir)\n        self.transform = transform\n        df = pd.read_csv(csv_path)\n        self.samples = df[['id_code', 'diagnosis']].values.tolist()\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        img_name, label = self.samples[idx]\n        img_path = self.img_dir / f\"{img_name}.png\"\n        if not img_path.exists():\n            img_path = Path(\"/kaggle/input/aptos2019-blindness-detection/train_images\") / f\"{img_name}.png\"\n        image = Image.open(img_path).convert('RGB')\n        if self.transform:\n            image = self.transform(image)\n        return image, torch.tensor(label, dtype=torch.long)\n\n# Transforms de validation\nval_tf = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# Chemins des données\nVAL_DIR = \"/kaggle/input/aptos-2019/aptos2019/processed/validation\"\nCSV_PATH = \"/kaggle/input/aptos2019-blindness-detection/train.csv\"\n\n# Création du DataLoader\nval_ds = APTOSDataset(VAL_DIR, CSV_PATH, transform=val_tf)\nval_loader = DataLoader(val_ds, batch_size=8, shuffle=False, num_workers=2, pin_memory=True)\n\nprint(f\"✅ Dataset de validation chargé : {len(val_ds)} images\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:12:39.589043Z","iopub.execute_input":"2026-02-09T14:12:39.5894Z","iopub.status.idle":"2026-02-09T14:12:39.681304Z","shell.execute_reply.started":"2026-02-09T14:12:39.589371Z","shell.execute_reply":"2026-02-09T14:12:39.68072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"🚀 Évaluation du modèle sur le jeu de validation...\")\n\nall_preds, all_labels = [], []\n\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        imgs, labels = imgs.to(device), labels.to(device)\n        outputs = model(imgs)\n        preds = outputs.argmax(1)\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n# Rapport détaillé\nprint(\"\\n\" + \"=\"*70)\nprint(\"📊 RAPPORT DE CLASSIFICATION – Modèle Hybride ResNet50 + DeiT-Basic\")\nprint(\"=\"*70)\nprint(classification_report(\n    all_labels,\n    all_preds,\n    target_names=['0-No DR', '1-Mild', '2-Moderate', '3-Severe', '4-Proliferative'],\n    digits=4\n))\n\n# Précision globale\nfinal_acc = np.mean(np.array(all_preds) == np.array(all_labels))\nprint(f\"🎯 Précision globale : {final_acc:.4f} ({100 * final_acc:.2f}%)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:12:59.00665Z","iopub.execute_input":"2026-02-09T14:12:59.006996Z","iopub.status.idle":"2026-02-09T14:17:16.21465Z","shell.execute_reply.started":"2026-02-09T14:12:59.006969Z","shell.execute_reply":"2026-02-09T14:17:16.213898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision import models, transforms\nimport timm\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport cv2\nimport pandas as pd\nfrom pathlib import Path\nimport os\n\n# Vérifier les datasets disponibles\nprint(\"📁 Datasets disponibles dans /kaggle/input :\")\nfor d in os.listdir('/kaggle/input'):\n    print(f\"  - {d}\")\n\n# Vérifier que les datasets nécessaires sont présents\nrequired_datasets = ['aptos2019-blindness-detection', 'aptos-2019']\nfor ds in required_datasets:\n    if not any(ds.replace('-', '') in d.replace('-', '') for d in os.listdir('/kaggle/input')):\n        print(f\"⚠️  Dataset '{ds}' manquant ! Cliquez sur '+ Add Data' et ajoutez-le.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:45:01.236835Z","iopub.execute_input":"2026-02-09T14:45:01.237165Z","iopub.status.idle":"2026-02-09T14:45:01.969344Z","shell.execute_reply.started":"2026-02-09T14:45:01.23714Z","shell.execute_reply":"2026-02-09T14:45:01.968784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HybridResNet50_DeitBasic(nn.Module):\n    def __init__(self, num_classes=5):\n        super().__init__()\n        self.resnet = models.resnet50(weights='IMAGENET1K_V1')\n        self.resnet.fc = nn.Identity()\n        self.deit = timm.create_model('deit_base_patch16_224', pretrained=True, num_classes=0)\n        self.fusion = nn.Linear(2048 + 768, 512)\n        self.head = nn.Sequential(\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    \n    def forward(self, x):\n        res_feat = self.resnet(x)\n        deit_feat = self.deit(x)\n        fused = torch.cat([res_feat, deit_feat], dim=1)\n        return self.head(self.fusion(fused))\n\n# Chargement du modèle\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = HybridResNet50_DeitBasic().to(device)\n\nMODEL_PATH = '/kaggle/input/resnet50-deit-basic-best-pthyy/resnet50_deit_basic_best.pth'\nmodel.load_state_dict(torch.load(MODEL_PATH, map_location=device))\nmodel.eval()\nprint(\"✅ Modèle hybride chargé avec succès\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:45:23.954233Z","iopub.execute_input":"2026-02-09T14:45:23.954923Z","iopub.status.idle":"2026-02-09T14:45:26.249949Z","shell.execute_reply.started":"2026-02-09T14:45:23.954892Z","shell.execute_reply":"2026-02-09T14:45:26.249286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Chemins des datasets (adaptés aux noms réels dans Kaggle)\nAPTOS_CSV = \"/kaggle/input/aptos2019-blindness-detection/train.csv\"\nAPTOS_IMAGES = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\n\n# Charger le CSV pour obtenir les IDs existants\ndf = pd.read_csv(APTOS_CSV)\nprint(f\"✅ CSV chargé : {len(df)} images disponibles\")\n\n# Vérifier que le dossier d'images existe\nif not os.path.exists(APTOS_IMAGES):\n    print(f\"❌ Dossier images introuvable : {APTOS_IMAGES}\")\n    print(\"💡 Solutions :\")\n    print(\"1. Ajoutez le dataset 'aptos2019-blindness-detection'\")\n    print(\"2. Vérifiez le chemin exact avec : !ls /kaggle/input/\")\n    raise FileNotFoundError(\"Dossier images introuvable\")\n\n# Sélectionner une image existante (première image du CSV)\nsample_image_id = df.iloc[0]['id_code']\nsample_image_path = Path(APTOS_IMAGES) / f\"{sample_image_id}.png\"\n\nif not sample_image_path.exists():\n    # Essayer avec une autre image\n    for i in range(min(10, len(df))):\n        img_id = df.iloc[i]['id_code']\n        img_path = Path(APTOS_IMAGES) / f\"{img_id}.png\"\n        if img_path.exists():\n            sample_image_id = img_id\n            sample_image_path = img_path\n            break\n    else:\n        raise FileNotFoundError(\"Aucune image trouvée dans le dataset\")\n\ntrue_label = df[df['id_code'] == sample_image_id]['diagnosis'].values[0]\nprint(f\"✅ Image sélectionnée : {sample_image_id}.png (Classe: {true_label})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:45:46.391259Z","iopub.execute_input":"2026-02-09T14:45:46.392041Z","iopub.status.idle":"2026-02-09T14:45:46.408852Z","shell.execute_reply.started":"2026-02-09T14:45:46.39201Z","shell.execute_reply":"2026-02-09T14:45:46.4082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class IntegratedGradients:\n    def __init__(self, model):\n        self.model = model\n    \n    def generate_heatmap(self, input_tensor, target_class=None, steps=50):\n        # Baseline : image noire\n        baseline = torch.zeros_like(input_tensor).to(device)\n        \n        # Générer les inputs interpolés\n        scaled_inputs = []\n        for i in range(steps + 1):\n            alpha = float(i) / steps\n            scaled_input = baseline + alpha * (input_tensor - baseline)\n            scaled_inputs.append(scaled_input)\n        \n        scaled_inputs = torch.cat(scaled_inputs, dim=0).to(device)\n        \n        # Calculer les gradients\n        grads = []\n        for i in range(steps + 1):\n            inp = scaled_inputs[i:i+1].requires_grad_(True)\n            output = self.model(inp)\n            \n            if target_class is None:\n                target_class = output.argmax(dim=1).item()\n            \n            self.model.zero_grad()\n            score = output[0, target_class]\n            score.backward()\n            grads.append(inp.grad.detach().cpu())\n        \n        # Moyenne des gradients\n        avg_grads = torch.mean(torch.stack(grads), dim=0)\n        \n        # Integrated Gradients\n        integrated_grads = (input_tensor.cpu() - baseline.cpu()) * avg_grads\n        heatmap = torch.abs(integrated_grads).sum(dim=1).squeeze().numpy()\n        heatmap = np.maximum(heatmap, 0)\n        heatmap = (heatmap - heatmap.min()) / (heatmap.max() - heatmap.min() + 1e-8)\n        \n        return heatmap, target_class\n\nig = IntegratedGradients(model)\nprint(\"✅ Integrated Gradients configuré\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:46:13.150089Z","iopub.execute_input":"2026-02-09T14:46:13.150389Z","iopub.status.idle":"2026-02-09T14:46:13.161153Z","shell.execute_reply.started":"2026-02-09T14:46:13.150364Z","shell.execute_reply":"2026-02-09T14:46:13.160474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GradCAMppResNet:\n    def __init__(self, model):\n        self.model = model\n        self.gradients = None\n        self.activations = None\n        \n        # Hook sur la dernière couche de convolution de ResNet50\n        model.resnet.layer4[-1].register_forward_hook(self.save_activation)\n        model.resnet.layer4[-1].register_backward_hook(self.save_gradient)\n    \n    def save_activation(self, module, input, output):\n        self.activations = output.detach()\n    \n    def save_gradient(self, module, grad_input, grad_output):\n        self.gradients = grad_output[0].detach()\n    \n    def generate_heatmap(self, input_tensor, target_class):\n        output = self.model(input_tensor)\n        self.model.zero_grad()\n        output[0, target_class].backward()\n        \n        # Calcul Grad-CAM++ (formule complète)\n        alpha_numer = self.gradients.pow(2)\n        alpha_denom = self.gradients.pow(2).mul(2) + \\\n                      self.activations.mul(self.gradients.pow(3)).sum(dim=[2, 3], keepdim=True)\n        alpha_denom = torch.where(alpha_denom != 0.0, alpha_denom, torch.ones_like(alpha_denom))\n        alpha = alpha_numer / alpha_denom\n        \n        weights = torch.relu(self.gradients).mul(alpha).sum(dim=[2, 3])\n        weights = weights.unsqueeze(-1).unsqueeze(-1)\n        heatmap = torch.relu((weights * self.activations).sum(dim=1, keepdim=True))\n        heatmap = torch.nn.functional.interpolate(\n            heatmap, size=(224, 224), mode='bilinear', align_corners=False\n        )\n        heatmap = heatmap.squeeze().cpu().numpy()\n        heatmap = (heatmap - heatmap.min()) / (heatmap.max() - heatmap.min() + 1e-8)\n        return heatmap\n\ngradcampp = GradCAMppResNet(model)\nprint(\"✅ Grad-CAM++ configuré pour ResNet50\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:46:39.141261Z","iopub.execute_input":"2026-02-09T14:46:39.141557Z","iopub.status.idle":"2026-02-09T14:46:39.150356Z","shell.execute_reply.started":"2026-02-09T14:46:39.141533Z","shell.execute_reply":"2026-02-09T14:46:39.149676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Charger l'image sélectionnée\noriginal = Image.open(sample_image_path).convert('RGB')\noriginal_np = np.array(original.resize((224, 224)))\n\n# Prétraitement\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\ninput_tensor = transform(original).unsqueeze(0).to(device)\n\n# Prédiction\nwith torch.no_grad():\n    pred_class = model(input_tensor).argmax(1).item()\n\n# Générer les heatmaps\nheatmap_ig, _ = ig.generate_heatmap(input_tensor, target_class=pred_class)\nheatmap_gradcam = gradcampp.generate_heatmap(input_tensor, pred_class)\n\n# Visualisation\nfig, axes = plt.subplots(1, 3, figsize=(16, 5))\n\n# Image originale\naxes[0].imshow(original_np)\naxes[0].set_title(f\"Image originale\\nVraie: {true_label} | Prédite: {pred_class}\", \n                  fontsize=11, fontweight='bold', pad=10)\naxes[0].axis('off')\n\n# Integrated Gradients (sur modèle hybride COMPLET)\naxes[1].imshow(original_np)\naxes[1].imshow(heatmap_ig, cmap='jet', alpha=0.55)\naxes[1].set_title(\"Integrated Gradients\\n(Modèle hybride COMPLET)\", \n                  fontsize=11, fontweight='bold', color='green', pad=10)\naxes[1].axis('off')\n\n# Grad-CAM++ (ResNet50 uniquement)\naxes[2].imshow(original_np)\naxes[2].imshow(heatmap_gradcam, cmap='jet', alpha=0.55)\naxes[2].set_title(\"Grad-CAM++\\n(Branche ResNet50 uniquement)\", \n                  fontsize=11, fontweight='bold', color='blue', pad=10)\naxes[2].axis('off')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/xai_comparison_hybrid.png', dpi=150, bbox_inches='tight')\nprint(\"✅ Visualisation sauvegardée : /kaggle/working/xai_comparison_hybrid.png\")\nplt.show()\n\n# Afficher les classes pour référence\nprint(\"\\n📌 Légende des classes :\")\nprint(\"0 = No DR | 1 = Mild | 2 = Moderate | 3 = Severe | 4 = Proliferative\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T14:46:59.901432Z","iopub.execute_input":"2026-02-09T14:46:59.902351Z","iopub.status.idle":"2026-02-09T14:47:05.519193Z","shell.execute_reply.started":"2026-02-09T14:46:59.902305Z","shell.execute_reply":"2026-02-09T14:47:05.518269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch, torch.nn as nn\nfrom torchvision import models, transforms\nimport timm, numpy as np, pandas as pd, matplotlib.pyplot as plt, cv2\nfrom PIL import Image\nfrom pathlib import Path\nimport os\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Architecture hybride\nclass HybridResNet50_DeitBasic(nn.Module):\n    def __init__(self, num_classes=5):\n        super().__init__()\n        self.resnet = models.resnet50(weights='IMAGENET1K_V1')\n        self.resnet.fc = nn.Identity()\n        self.deit = timm.create_model('deit_base_patch16_224', pretrained=True, num_classes=0)\n        self.fusion = nn.Linear(2048 + 768, 512)\n        self.head = nn.Sequential(nn.ReLU(), nn.Dropout(0.3), nn.Linear(512, num_classes))\n    def forward(self, x):\n        return self.head(self.fusion(torch.cat([self.resnet(x), self.deit(x)], dim=1)))\n\n# Chargement du modèle\nmodel = HybridResNet50_DeitBasic().to(device)\nmodel.load_state_dict(torch.load('/kaggle/input/resnet50-deit-basic-best-pthyy/resnet50_deit_basic_best.pth', map_location=device))\nmodel.eval()\nprint(\"✅ Modèle hybride chargé | Device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T15:12:25.952282Z","iopub.execute_input":"2026-02-09T15:12:25.952946Z","iopub.status.idle":"2026-02-09T15:12:28.172845Z","shell.execute_reply.started":"2026-02-09T15:12:25.952916Z","shell.execute_reply":"2026-02-09T15:12:28.172193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Grad-CAM++ pour ResNet50\nclass GradCAMpp:\n    def __init__(self, model): \n        self.model = model\n        model.resnet.layer4[-1].register_forward_hook(lambda m,i,o: setattr(self, 'acts', o.detach()))\n        model.resnet.layer4[-1].register_backward_hook(lambda m,gi,go: setattr(self, 'grads', go[0].detach()))\n    def __call__(self, x, cls):\n        self.model.zero_grad()\n        self.model(x)[0, cls].backward()\n        a = torch.relu(self.grads)\n        alpha = self.grads.pow(2) / (2*self.grads.pow(2) + (self.acts * self.grads.pow(3)).sum([2,3], keepdim=True))\n        w = (alpha * a).sum([2,3], keepdim=True)\n        hm = torch.relu((w * self.acts).sum(1, keepdim=True))\n        hm = torch.nn.functional.interpolate(hm, (224,224), mode='bilinear')[0,0].cpu().numpy()\n        return (hm - hm.min()) / (hm.max() - hm.min() + 1e-8)\n\n# Attention Rollout pour DeiT\ndef attention_rollout(model, x):\n    with torch.no_grad():\n        tokens = model.deit.patch_embed(x)\n        tokens = torch.cat([model.deit.cls_token.expand(1,-1,-1), tokens], dim=1) + model.deit.pos_embed\n        attns = []\n        for blk in model.deit.blocks:\n            qkv = blk.attn.qkv(tokens).reshape(1, -1, 3, 12, 64).permute(2,0,3,1,4)\n            attn = (qkv[0] @ qkv[1].transpose(-2,-1)) * blk.attn.scale\n            attn = attn.softmax(dim=-1)\n            attns.append(attn[0, :, 0, 1:].mean(0))  # [CLS] → patches\n            tokens = blk(tokens)\n        rollout = attns[-1].reshape(14,14).cpu().numpy()\n        return cv2.resize(rollout, (224,224))\n\ngradcampp = GradCAMpp(model)\nprint(\"✅ XAI prêt : Grad-CAM++ (ResNet50) + Attention Rollout (DeiT)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T15:12:48.286375Z","iopub.execute_input":"2026-02-09T15:12:48.287238Z","iopub.status.idle":"2026-02-09T15:12:48.296602Z","shell.execute_reply.started":"2026-02-09T15:12:48.287207Z","shell.execute_reply":"2026-02-09T15:12:48.295985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Charger CSV et préparer transforms\ndf = pd.read_csv('/kaggle/input/aptos2019-blindness-detection/train.csv')\ntransform = transforms.Compose([transforms.Resize((224,224)), transforms.ToTensor(), \n                               transforms.Normalize([0.485,0.456,0.406], [0.229,0.224,0.225])])\n\n# Sélectionner 2 images par classe (10 images totales)\nrepresentative_imgs = []\nfor cls in range(5):\n    samples = df[df['diagnosis']==cls]['id_code'].values[:2]\n    for img_id in samples:\n        img_path = Path(f'/kaggle/input/aptos2019-blindness-detection/train_images/{img_id}.png')\n        if img_path.exists():\n            img = Image.open(img_path).convert('RGB')\n            tensor = transform(img).unsqueeze(0).to(device)\n            with torch.no_grad():\n                pred = model(tensor).argmax(1).item()\n            representative_imgs.append({\n                'id': img_id, 'true': cls, 'pred': pred, \n                'tensor': tensor, 'original': np.array(img.resize((224,224)))\n            })\n\nprint(f\"✅ {len(representative_imgs)} images sélectionnées (2 par classe)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T15:13:09.821375Z","iopub.execute_input":"2026-02-09T15:13:09.82198Z","iopub.status.idle":"2026-02-09T15:13:11.982843Z","shell.execute_reply.started":"2026-02-09T15:13:09.82195Z","shell.execute_reply":"2026-02-09T15:13:11.98195Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Créer grille 5 classes × 3 colonnes (originale | Grad-CAM++ | Attention Rollout)\nfig, axes = plt.subplots(5, 3, figsize=(15, 22))\nclass_names = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']\n\nfor i, sample in enumerate(representative_imgs):\n    row = i // 2  # 2 images par classe → 5 lignes\n    col_offset = (i % 2) * 3  # Chaque paire occupe 3 colonnes\n    \n    # Colonne 1: Image originale\n    axes[row, 0].imshow(sample['original'])\n    axes[row, 0].set_title(f\"Classe {sample['true']}: {class_names[sample['true']]}\\nPrédit: {class_names[sample['pred']]}\", \n                          fontsize=11, fontweight='bold')\n    axes[row, 0].axis('off')\n    \n    # Colonne 2: Grad-CAM++ (ResNet50)\n    hm_grad = gradcampp(sample['tensor'], sample['pred'])\n    axes[row, 1].imshow(sample['original'])\n    axes[row, 1].imshow(hm_grad, cmap='jet', alpha=0.55)\n    axes[row, 1].set_title(\"Grad-CAM++\\n(Détails fins)\", fontsize=10, color='darkred', fontweight='bold')\n    axes[row, 1].axis('off')\n    \n    # Colonne 3: Attention Rollout (DeiT)\n    hm_attn = attention_rollout(model, sample['tensor'])\n    axes[row, 2].imshow(sample['original'])\n    axes[row, 2].imshow(hm_attn, cmap='viridis', alpha=0.55)\n    axes[row, 2].set_title(\"Attention Rollout\\n(Contexte global)\", fontsize=10, color='darkblue', fontweight='bold')\n    axes[row, 2].axis('off')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/xai_grid.png', dpi=150, bbox_inches='tight')\nprint(\"✅ Visualisation sauvegardée: /kaggle/working/xai_grid.png\")\nplt.show()\n\n# Afficher métriques résumées\naccuracy = np.mean([s['true']==s['pred'] for s in representative_imgs])\nprint(\"\\n\" + \"=\"*60)\nprint(f\"📊 Résultats sur {len(representative_imgs)} images représentatives\")\nprint(\"=\"*60)\nprint(f\"🎯 Précision: {accuracy:.1%}\")\nprint(f\"✅ Bonnes prédictions: {sum(s['true']==s['pred'] for s in representative_imgs)}/{len(representative_imgs)}\")\nprint(\"\\n💡 Interprétation clinique:\")\nprint(\"   • Grad-CAM++ → Micro-anévrismes, hémorragies ponctuelles\")\nprint(\"   • Attention Rollout → Propagation des lésions, réseau vasculaire\")\nprint(\"   • SYNERGIE → Explication de la haute précision (98.91%) du modèle hybride\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T15:13:34.877691Z","iopub.execute_input":"2026-02-09T15:13:34.878036Z","iopub.status.idle":"2026-02-09T15:13:43.056674Z","shell.execute_reply.started":"2026-02-09T15:13:34.878009Z","shell.execute_reply":"2026-02-09T15:13:43.05567Z"}},"outputs":[],"execution_count":null}]}