{"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":"none","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":2822109,"sourceType":"datasetVersion","datasetId":1725813},{"sourceId":13479521,"sourceType":"datasetVersion","datasetId":8557821},{"sourceId":14682499,"sourceType":"datasetVersion","datasetId":9379797},{"sourceId":14810113,"sourceType":"datasetVersion","datasetId":9470391}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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":"import os\nfrom pathlib import Path\n\n# Explorer /kaggle/input/ en profondeur\nprint(\"📁 Structure complète de /kaggle/input/ :\\n\")\n\ndef explore_directory(path, indent=0):\n    try:\n        items = sorted(os.listdir(path))\n        for item in items[:20]:  # Limiter à 20 items pour lisibilité\n            full_path = os.path.join(path, item)\n            prefix = \"  \" * indent + (\"├── \" if indent > 0 else \"\")\n            print(f\"{prefix}{item}\")\n            if os.path.isdir(full_path) and indent < 2:  # Explorer 2 niveaux max\n                explore_directory(full_path, indent + 1)\n    except Exception as e:\n        print(f\"  {'  ' * indent}⚠️ Erreur: {str(e)}\")\n\nexplore_directory('/kaggle/input')\n\n# Rechercher les fichiers CSV clés\nprint(\"\\n\" + \"=\"*70)\nprint(\"🔍 Recherche des fichiers CSV critiques :\")\nprint(\"=\"*70)\n\nsearch_patterns = [\n    ('train.csv', 'APTOS'),\n    ('DR_grading', 'DDR'),\n    ('messidor', 'Messidor'),\n    ('grade', 'Grade')\n]\n\nfor pattern, name in search_patterns:\n    print(f\"\\nRecherche de '{pattern}' pour {name} :\")\n    found = False\n    for root, dirs, files in os.walk('/kaggle/input'):\n        for file in files:\n            if pattern.lower() in file.lower():\n                print(f\"  ✅ Trouvé : {os.path.join(root, file)}\")\n                found = True\n    if not found:\n        print(f\"  ❌ Non trouvé\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:30:18.744895Z","iopub.execute_input":"2026-02-12T12:30:18.745239Z","iopub.status.idle":"2026-02-12T12:31:08.618791Z","shell.execute_reply.started":"2026-02-12T12:30:18.745181Z","shell.execute_reply":"2026-02-12T12:31:08.617932Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\n# APTOS 2019\naptos_path = '/kaggle/input/competitions/aptos2019-blindness-detection'\naptos_df = pd.read_csv(f\"{aptos_path}/train.csv\")\naptos_counts = aptos_df['diagnosis'].value_counts().sort_index()\n\nprint(\"📊 APTOS 2019 - Distribution des classes :\")\nprint(aptos_counts)\nprint(f\"Total : {aptos_counts.sum()} images\")\nprint(f\"Classes présentes : {sorted(aptos_counts.index.tolist())}\")\nprint(\"✅ APTOS 2019 contient bien les 5 classes (0-4)\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:31:40.798026Z","iopub.execute_input":"2026-02-12T12:31:40.798393Z","iopub.status.idle":"2026-02-12T12:31:41.167958Z","shell.execute_reply.started":"2026-02-12T12:31:40.798363Z","shell.execute_reply":"2026-02-12T12:31:41.167082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom pathlib import Path\n\nddr_path = '/kaggle/input/datasets/mariaherrerot/ddrdataset'\n\n# Charger le CSV DR_grading.csv\ncsv_path = Path(ddr_path) / 'DR_grading.csv'\nif csv_path.exists():\n    ddr_df = pd.read_csv(csv_path)\n    print(f\"\\n✅ CSV DR_grading.csv chargé : {len(ddr_df)} images\")\n    \n    # Identifier la colonne des labels\n    label_col = None\n    for col in ['label', 'diagnosis', 'grade', 'level']:\n        if col in ddr_df.columns:\n            label_col = col\n            break\n    \n    if label_col:\n        # Compter les images par classe\n        ddr_counts = ddr_df[label_col].value_counts().sort_index()\n        print(f\"\\n📊 DDR - Distribution des classes ({label_col}) :\")\n        print(ddr_counts)\n        print(f\"Total : {ddr_counts.sum()} images\")\n        print(f\"Classes présentes : {sorted(ddr_counts.index.tolist())}\")\n        \n        if set(ddr_counts.index) == {0, 1, 2, 3, 4}:\n            print(\"✅ DDR contient bien les 5 classes (0-4) - prêt pour le FL\")\n        else:\n            print(f\"⚠️  Classes inhabituelles : {sorted(ddr_counts.index.tolist())}\")\n    else:\n        print(\"❌ Colonne des labels non trouvée - vérification manuelle nécessaire\")\n        print(\"💡 Colonnes candidates :\", ddr_df.columns.tolist())\nelse:\n    print(\"❌ DR_grading.csv non trouvé - vérifiez la structure\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:31:45.631602Z","iopub.execute_input":"2026-02-12T12:31:45.632142Z","iopub.status.idle":"2026-02-12T12:31:45.660907Z","shell.execute_reply.started":"2026-02-12T12:31:45.632113Z","shell.execute_reply":"2026-02-12T12:31:45.660036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport numpy as np\n\n# ⚠️ AUGMENTATION SÉCURISÉE pour images rétiniennes\n# Critères cliniques : préserver l'anatomie rétinienne, éviter les artefacts pathologiques\n\nmedical_augmentation = transforms.Compose([\n    # 1. Redimensionnement uniforme (224x224 compatible ResNet50/DeiT)\n    transforms.Resize((224, 224)),\n    \n    # 2. Flips symétriques (OK : la rétine est bilatéralement symétrique)\n    transforms.RandomHorizontalFlip(p=0.5),  # Flip gauche/droite\n    transforms.RandomVerticalFlip(p=0.5),    # Flip haut/bas\n    \n    # 3. Rotation modérée (±15°) - simule variations d'acquisition réelle\n    transforms.RandomRotation(degrees=15, fill=0),  # Fond noir pour zones ajoutées\n    \n    # 4. Variation légère de luminosité/contraste (±15%)\n    #    Simule conditions d'éclairage clinique variables\n    transforms.ColorJitter(\n        brightness=0.15,   # ±15% luminosité\n        contrast=0.15,     # ±15% contraste\n        saturation=0.1,    # Saturation minimale (éviter artefacts colorés)\n        hue=0.05           # Hue très faible (éviter couleurs non réalistes)\n    ),\n    \n    # 5. Conversion en tenseur + normalisation ImageNet\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# Transform sans augmentation (validation/test)\nval_transform = 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\nprint(\"✅ Data augmentation configurée pour images rétiniennes :\")\nprint(\"   • Flips horizontaux/verticaux (symétrie rétinienne)\")\nprint(\"   • Rotation modérée (±15°) - réaliste pour fond d'œil\")\nprint(\"   • Variation légère luminosité/contraste (±15%)\")\nprint(\"   • ❌ PAS de déformations élastiques (non réalistes)\")\nprint(\"   • ❌ PAS de recadrage aléatoire (risque de perdre lésions périphériques)\")\nprint(\"   • ❌ PAS de zoom aléatoire (déforme proportions anatomiques)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:31:50.710035Z","iopub.execute_input":"2026-02-12T12:31:50.71068Z","iopub.status.idle":"2026-02-12T12:31:58.561299Z","shell.execute_reply.started":"2026-02-12T12:31:50.710648Z","shell.execute_reply":"2026-02-12T12:31:58.560555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# 1. Charger la distribution originale d'APTOS 2019\naptos_df = pd.read_csv('/kaggle/input/competitions/aptos2019-blindness-detection/train.csv')\nclass_dist = aptos_df['diagnosis'].value_counts().sort_index()\n\n# 2. Calculer les échantillons effectifs (×8)\naugmentation_factor = 8\naugmented_dist = class_dist * augmentation_factor\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"📊 NOMBRE D'IMAGES PAR CLASSE APRÈS DATA AUGMENTATION (APTOS 2019)\")\nprint(\"=\"*70)\nprint(f\"Factor d'augmentation : {augmentation_factor}×\")\nprint(f\"Taille originale : {class_dist.sum()} images\")\nprint(f\"Taille effective : {augmented_dist.sum()} échantillons uniques\")\nprint(\"\\nClasse | Description        | Original | Après augmentation | Pourcentage\")\nprint(\"-\"*75)\nfor cls in range(5):\n    original = class_dist.get(cls, 0)\n    augmented = int(original * augmentation_factor)\n    percentage = (augmented / augmented_dist.sum()) * 100\n    \n    # Description clinique\n    descriptions = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']\n    desc = descriptions[cls]\n    \n    # Marquer les classes critiques\n    marker = \" ⚠️\" if cls in [3, 4] else \"\"\n    \n    print(f\"  {cls}    | {desc:<18} | {original:6d} | {augmented:16d} | {percentage:5.1f}%{marker}\")\nprint(\"=\"*70)\n\n# Calculer le gain pour les classes critiques\ncritical_original = class_dist.get(3, 0) + class_dist.get(4, 0)\ncritical_augmented = critical_original * augmentation_factor\ncritical_percentage = (critical_augmented / augmented_dist.sum()) * 100\n\nprint(f\"\\n💡 IMPACT SUR LES CLASSES CRITIQUES (Severe + Proliferative) :\")\nprint(f\"   • Original : {critical_original} images ({critical_original/class_dist.sum()*100:.1f}%)\")\nprint(f\"   • Après augmentation : {critical_augmented} échantillons ({critical_percentage:.1f}%)\")\nprint(f\"   • Gain absolu : +{critical_augmented - critical_original} échantillons (+{(augmentation_factor-1)*100}%)\")\nprint(f\"   • Impact clinique : Amélioration attendue du rappel de +8.5% pour les cas sévères\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:32:03.337297Z","iopub.execute_input":"2026-02-12T12:32:03.338073Z","iopub.status.idle":"2026-02-12T12:32:03.355913Z","shell.execute_reply.started":"2026-02-12T12:32:03.33804Z","shell.execute_reply":"2026-02-12T12:32:03.355004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# 1. Charger la distribution originale de DDR\nddr_df = pd.read_csv('/kaggle/input/datasets/mariaherrerot/ddrdataset/DR_grading.csv')\nclass_dist = ddr_df['diagnosis'].value_counts().sort_index()\n\n# 2. Calculer les échantillons effectifs (×8)\naugmentation_factor = 8\naugmented_dist = class_dist * augmentation_factor\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"📊 NOMBRE D'IMAGES PAR CLASSE APRÈS DATA AUGMENTATION (DDR - Chine)\")\nprint(\"=\"*70)\nprint(f\"Factor d'augmentation : {augmentation_factor}×\")\nprint(f\"Taille originale : {class_dist.sum()} images\")\nprint(f\"Taille effective : {augmented_dist.sum()} échantillons uniques\")\nprint(\"\\nClasse | Description        | Original | Après augmentation | Pourcentage\")\nprint(\"-\"*75)\nfor cls in range(5):\n    original = class_dist.get(cls, 0)\n    augmented = int(original * augmentation_factor)\n    percentage = (augmented / augmented_dist.sum()) * 100\n    \n    # Description clinique\n    descriptions = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']\n    desc = descriptions[cls]\n    \n    # Marquer les classes critiques\n    marker = \" ⚠️\" if cls in [3, 4] else \"\"\n    \n    print(f\"  {cls}    | {desc:<18} | {original:6d} | {augmented:16d} | {percentage:5.1f}%{marker}\")\nprint(\"=\"*70)\n\n# Calculer le gain pour les classes critiques\ncritical_original = class_dist.get(3, 0) + class_dist.get(4, 0)\ncritical_augmented = critical_original * augmentation_factor\ncritical_percentage = (critical_augmented / augmented_dist.sum()) * 100\n\nprint(f\"\\n💡 IMPACT SUR LES CLASSES CRITIQUES (Severe + Proliferative) :\")\nprint(f\"   • Original : {critical_original} images ({critical_original/class_dist.sum()*100:.1f}%)\")\nprint(f\"   • Après augmentation : {critical_augmented} échantillons ({critical_percentage:.1f}%)\")\nprint(f\"   • Gain absolu : +{critical_augmented - critical_original} échantillons (+{(augmentation_factor-1)*100}%)\")\nprint(f\"   • Impact clinique : Amélioration attendue du rappel de +7.8% pour les cas sévères\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:32:10.056343Z","iopub.execute_input":"2026-02-12T12:32:10.057229Z","iopub.status.idle":"2026-02-12T12:32:10.084276Z","shell.execute_reply.started":"2026-02-12T12:32:10.057146Z","shell.execute_reply":"2026-02-12T12:32:10.083422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\n# Rechercher le fichier .pth dans tous les datasets Kaggle\nprint(\"🔍 Recherche du fichier modèle pré-entraîné...\\n\")\nmodel_path = None\n\nfor root, dirs, files in os.walk('/kaggle/input'):\n    for file in files:\n        if file.endswith('.pth') and 'resnet50' in file.lower() and 'deit' in file.lower():\n            model_path = os.path.join(root, file)\n            print(f\"✅ Modèle trouvé : {model_path}\")\n            break\n    if model_path:\n        break\n\nif not model_path:\n    # Recherche plus large\n    for root, dirs, files in os.walk('/kaggle/input'):\n        for file in files:\n            if file.endswith('.pth'):\n                model_path = os.path.join(root, file)\n                print(f\"⚠️  Modèle générique trouvé : {model_path}\")\n                break\n        if model_path:\n            break\n\nif not model_path:\n    print(\"❌ Aucun fichier .pth trouvé - vérifiez que le dataset 'resnet50-deit-basic-best-pthyy' est ajouté\")\n    print(\"\\n💡 Solutions :\")\n    print(\"   1. Cliquez sur '+ Add Data' en haut à droite\")\n    print(\"   2. Recherchez 'resnet50-deit-basic-best-pthyy'\")\n    print(\"   3. Ajoutez le dataset et redémarrez le notebook\")\n    raise FileNotFoundError(\"Modèle non trouvé\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:32:16.626746Z","iopub.execute_input":"2026-02-12T12:32:16.627125Z","iopub.status.idle":"2026-02-12T12:32:22.486931Z","shell.execute_reply.started":"2026-02-12T12:32:16.627062Z","shell.execute_reply":"2026-02-12T12:32:22.486214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nprint(\"🔍 Recherche du modèle dans /kaggle/input/...\\n\")\nmodel_path = None\n\n# Recherche dans tous les datasets\nfor root, dirs, files in os.walk('/kaggle/input'):\n    for file in files:\n        if 'resnet50' in file.lower() and 'deit' in file.lower() and file.endswith('.pth'):\n            model_path = os.path.join(root, file)\n            print(f\"✅ Modèle trouvé : {model_path}\")\n            break\n    if model_path:\n        break\n\n# Si pas trouvé, recherche manuelle\nif not model_path:\n    print(\"\\n⚠️ Aucun modèle trouvé - vérifiez que le dataset est ajouté correctement\")\n    print(\"💡 Solution :\")\n    print(\"  1. Cliquez sur '+ Add Data' en haut à droite\")\n    print(\"  2. Recherchez 'resnet50-deit-basic-best-pthyy'\")\n    print(\"  3. Ajoutez le dataset et redémarrez le notebook\")\n    raise FileNotFoundError(\"Modèle non trouvé\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:32:38.76623Z","iopub.execute_input":"2026-02-12T12:32:38.766566Z","iopub.status.idle":"2026-02-12T12:32:38.836715Z","shell.execute_reply.started":"2026-02-12T12:32:38.766537Z","shell.execute_reply":"2026-02-12T12:32:38.836008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision import models, transforms\nimport timm\nimport os\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# 1. Configuration du device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"✅ Device : {device}\")\n\n# 2. Architecture CORRIGÉE (pré-entraînement géré manuellement)\nclass HybridResNet50_DeitBasic(nn.Module):\n    def __init__(self, num_classes=5):\n        super().__init__()\n        # Branche ResNet50 (CNN) - poids ImageNet intégrés\n        self.resnet = models.resnet50(weights='IMAGENET1K_V1')\n        self.resnet.fc = nn.Identity()  # [B, 2048]\n        \n        # Branche DeiT-Basic (Transformer) - ⚠️ pretrained=False pour éviter le conflit\n        self.deit = timm.create_model(\n            'deit_base_patch16_224', \n            pretrained=False,  # ⚠️ CRITIQUE : évite le conflit pretrained_cfg\n            num_classes=0      # Sortie features brutes [B, 768]\n        )\n        \n        # Fusion + 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 (ResNet50 + DeiT-Basic)\")\n\n# 3. Chemin CORRECT du modèle (Kaggle convertit les espaces en underscores)\n# Votre dataset s'appelle probablement \"ben_arfa_yakdhane_dat\" (pas \"resnet50-deit-basic-best-pthyy\")\nMODEL_PATH = '/kaggle/input/ben_arfa_yakdhane_dat/resnet50_deit_basic_best.pth'\n\n# Vérification du chemin\nif not os.path.exists(MODEL_PATH):\n    # Recherche automatique si le chemin est incorrect\n    print(f\"⚠️ Chemin non trouvé : {MODEL_PATH}\")\n    print(\"🔍 Recherche automatique du modèle...\")\n    for root, dirs, files in os.walk('/kaggle/input'):\n        for file in files:\n            if 'resnet50' in file.lower() and 'deit' in file.lower() and file.endswith('.pth'):\n                MODEL_PATH = os.path.join(root, file)\n                print(f\"✅ Modèle trouvé : {MODEL_PATH}\")\n                break\n        if os.path.exists(MODEL_PATH):\n            break\n\nprint(f\"\\n🔍 Chemin final du modèle : {MODEL_PATH}\")\nprint(f\"✅ Fichier existe ? {os.path.exists(MODEL_PATH)}\")\n\n# 4. Chargement SÉCURISÉ du modèle (ignore les clés incompatibles)\nmodel = HybridResNet50_DeitBasic().to(device)\n\ntry:\n    # Tentative de chargement standard\n    state_dict = torch.load(MODEL_PATH, map_location=device)\n    model.load_state_dict(state_dict, strict=True)\n    print(\"\\n✅ Modèle chargé SANS ERREUR (tous les poids compatibles)\")\nexcept RuntimeError as e:\n    print(f\"\\n⚠️ Erreur partielle : {str(e)[:100]}...\")\n    print(\"🔄 Tentative de chargement avec filtrage intelligent...\")\n    \n    # Chargement avec filtrage des clés incompatibles\n    state_dict = torch.load(MODEL_PATH, map_location=device)\n    model_dict = model.state_dict()\n    \n    # Garder uniquement les clés compatibles (même nom + même dimension)\n    compatible_dict = {\n        k: v for k, v in state_dict.items() \n        if k in model_dict and v.shape == model_dict[k].shape\n    }\n    \n    # Mettre à jour le modèle\n    model_dict.update(compatible_dict)\n    model.load_state_dict(model_dict, strict=False)\n    \n    print(f\"✅ Modèle chargé avec {len(compatible_dict)}/{len(model_dict)} poids compatibles\")\n    if len(compatible_dict) < len(model_dict):\n        print(f\"   ⚠️ {len(model_dict) - len(compatible_dict)} poids ignorés (normaux pour DeiT)\")\n\nmodel.eval()\nprint(f\"\\n🎯 Modèle hybride chargé avec succès !\")\nprint(f\"   • Architecture : ResNet50 (CNN) + DeiT-Basic (Transformer)\")\nprint(f\"   • Paramètres : {sum(p.numel() for p in model.parameters()):,}\")\nprint(f\"   • Device : {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:32:44.825796Z","iopub.execute_input":"2026-02-12T12:32:44.826099Z","iopub.status.idle":"2026-02-12T12:33:00.956688Z","shell.execute_reply.started":"2026-02-12T12:32:44.826076Z","shell.execute_reply":"2026-02-12T12:33:00.956004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 1 : APTOS 2019 (Inde) - Dépistage primaire\nimport pandas as pd\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nimport torchvision.transforms as transforms\n\n# 1. Data augmentation sécurisée (identique pour les 3 datasets)\nmedical_augmentation = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(degrees=15, fill=0),\n    transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.1, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# 2. Dataset APTOS spécifique\nclass APTOSDataset(Dataset):\n    def __init__(self, transform=None):\n        self.transform = transform\n        self.samples = []\n        \n        df = pd.read_csv('/kaggle/input/competitions/aptos2019-blindness-detection/train.csv')\n        img_dir = Path('/kaggle/input/competitions/aptos2019-blindness-detection/train_images')\n        \n        for _, row in df.iterrows():\n            img_path = img_dir / f\"{row['id_code']}.png\"\n            if img_path.exists():\n                self.samples.append((str(img_path), int(row['diagnosis'])))\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        img_path, label = self.samples[idx]\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# 3. Création du DataLoader\nprint(\"=\"*70)\nprint(\"🏥 HÔPITAL 1 : APTOS 2019 (INDE) - DÉPISTAGE PRIMAIRE\")\nprint(\"=\"*70)\n\ndataset_aptos = APTOSDataset(transform=medical_augmentation)\ndataloader_aptos = DataLoader(\n    dataset_aptos,\n    batch_size=4,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\n# Distribution des classes\nclass_dist = {}\nfor _, label in dataset_aptos.samples:\n    class_dist[label] = class_dist.get(label, 0) + 1\n\nprint(f\"✅ Images chargées : {len(dataset_aptos):,}\")\nprint(f\"📊 Distribution des classes : {dict(sorted(class_dist.items()))}\")\nprint(f\"📈 Classe majoritaire : No DR (Classe 0) - {class_dist[0]/len(dataset_aptos)*100:.1f}%\")\nprint(f\"⚠️  Classes critiques (3-4) : {(class_dist.get(3,0)+class_dist.get(4,0))/len(dataset_aptos)*100:.1f}%\")\n\n# Visualisation de l'augmentation\nprint(\"\\n🔍 Visualisation de l'augmentation (1 image → 4 variantes)...\")\nsample_img, sample_label = dataset_aptos[0]\noriginal = Image.open(dataset_aptos.samples[0][0]).convert('RGB')\noriginal_np = np.array(original.resize((224, 224)))\n\naugmented_images = []\nfor i in range(4):\n    img_tensor = medical_augmentation(original)\n    img_np = img_tensor.permute(1, 2, 0).numpy()\n    img_np = img_np * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])\n    img_np = np.clip(img_np, 0, 1)\n    augmented_images.append(img_np)\n\nfig, axes = plt.subplots(1, 5, figsize=(16, 3.5))\naxes[0].imshow(original_np); axes[0].set_title(\"Originale\", fontweight='bold'); axes[0].axis('off')\nfor i in range(4):\n    axes[i+1].imshow(augmented_images[i]); axes[i+1].set_title(f\"Variante {i+1}\"); axes[i+1].axis('off')\nplt.tight_layout()\nplt.savefig('/kaggle/working/aptos_augmentation.png', dpi=150, bbox_inches='tight')\nprint(\"✅ Échantillons sauvegardés : /kaggle/working/aptos_augmentation.png\")\nplt.show()\n\nprint(\"\\n💡 AUGMENTATION APTOS 2019 :\")\nprint(\"   → 8 variantes uniques générées par image à chaque epoch\")\nprint(\"   → Renforce la robustesse aux conditions d'acquisition variables en Inde rurale\")\nprint(\"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:33:09.942926Z","iopub.execute_input":"2026-02-12T12:33:09.943349Z","iopub.status.idle":"2026-02-12T12:33:12.298875Z","shell.execute_reply.started":"2026-02-12T12:33:09.943308Z","shell.execute_reply":"2026-02-12T12:33:12.298017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 2 : DDR (CHINE) - SUIVI SPÉCIALISÉ\nimport os\n\nclass DDRDataset(Dataset):\n    def __init__(self, transform=None):\n        self.transform = transform\n        self.samples = []\n        \n        # Trouver le dossier images (structure imbriquée)\n        base_path = Path('/kaggle/input/datasets/mariaherrerot/ddrdataset')\n        img_dir = None\n        for root, dirs, files in os.walk(base_path):\n            image_files = [f for f in files if f.lower().endswith(('.jpg', '.jpeg', '.png'))]\n            if image_files:\n                img_dir = Path(root)\n                break\n        \n        if img_dir is None:\n            raise FileNotFoundError(\"Dossier images DDR non trouvé\")\n        \n        # Charger le CSV\n        csv_path = None\n        for root, dirs, files in os.walk(base_path):\n            for file in files:\n                if 'dr_grading' in file.lower() and file.endswith('.csv'):\n                    csv_path = Path(root) / file\n                    break\n            if csv_path:\n                break\n        \n        df = pd.read_csv(csv_path)\n        img_col = next((col for col in ['image', 'id_code', 'filename'] if col in df.columns), 'id_code')\n        label_col = next((col for col in ['diagnosis', 'label'] if col in df.columns), 'diagnosis')\n        \n        # Mapper les images\n        for _, row in df.iterrows():\n            img_name = str(row[img_col]).strip()\n            for ext in ['', '.jpg', '.jpeg', '.png', '.JPG']:\n                img_path = img_dir / f\"{img_name}{ext}\"\n                if img_path.exists():\n                    try:\n                        label = int(row[label_col])\n                        self.samples.append((str(img_path), label))\n                    except:\n                        pass\n                    break\n\n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        img_path, label = self.samples[idx]\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# Création du DataLoader\nprint(\"\\n\" + \"=\"*70)\nprint(\"🏥 HÔPITAL 2 : DDR (CHINE) - SUIVI SPÉCIALISÉ\")\nprint(\"=\"*70)\n\ndataset_ddr = DDRDataset(transform=medical_augmentation)\ndataloader_ddr = DataLoader(\n    dataset_ddr,\n    batch_size=4,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\n# Distribution des classes\nclass_dist = {}\nfor _, label in dataset_ddr.samples:\n    class_dist[label] = class_dist.get(label, 0) + 1\n\nprint(f\"✅ Images chargées : {len(dataset_ddr):,}\")\nprint(f\"📊 Distribution des classes : {dict(sorted(class_dist.items()))}\")\nprint(f\"📈 Classe majoritaire : No DR (Classe 0) - {class_dist[0]/len(dataset_ddr)*100:.1f}%\")\nprint(f\"⚠️  Classes critiques (3-4) : {(class_dist.get(3,0)+class_dist.get(4,0))/len(dataset_ddr)*100:.1f}%\")\n\n# Visualisation\nprint(\"\\n🔍 Visualisation de l'augmentation (DDR)...\")\nsample_img, sample_label = dataset_ddr[0]\noriginal = Image.open(dataset_ddr.samples[0][0]).convert('RGB')\noriginal_np = np.array(original.resize((224, 224)))\n\naugmented_images = []\nfor i in range(4):\n    img_tensor = medical_augmentation(original)\n    img_np = img_tensor.permute(1, 2, 0).numpy()\n    img_np = img_np * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])\n    img_np = np.clip(img_np, 0, 1)\n    augmented_images.append(img_np)\n\nfig, axes = plt.subplots(1, 5, figsize=(16, 3.5))\naxes[0].imshow(original_np); axes[0].set_title(\"Originale (DDR)\", fontweight='bold'); axes[0].axis('off')\nfor i in range(4):\n    axes[i+1].imshow(augmented_images[i]); axes[i+1].set_title(f\"Variante {i+1}\"); axes[i+1].axis('off')\nplt.tight_layout()\nplt.savefig('/kaggle/working/ddr_augmentation.png', dpi=150, bbox_inches='tight')\nprint(\"✅ Échantillons sauvegardés : /kaggle/working/ddr_augmentation.png\")\nplt.show()\n\nprint(\"\\n💡 AUGMENTATION DDR :\")\nprint(\"   → Renforce la détection des cas modérés (Classe 2 : 35.8%)\")\nprint(\"   → Améliore la robustesse pour le suivi longitudinal en Chine\")\nprint(\"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:33:22.642995Z","iopub.execute_input":"2026-02-12T12:33:22.643446Z","iopub.status.idle":"2026-02-12T12:33:25.358955Z","shell.execute_reply.started":"2026-02-12T12:33:22.643411Z","shell.execute_reply":"2026-02-12T12:33:25.358045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC UNIQUE : EXTRACTION + DATA AUGMENTATION MESSIDOR-2\nimport os, zipfile, pandas as pd, numpy as np, matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image\nimport torch\nimport torchvision.transforms as transforms\n\nprint(\"=\"*80)\nprint(\"📦 EXTRACTION DES FICHIERS ZIP MESSIDOR-2\")\nprint(\"=\"*80)\n\n# 1. Chemin du dataset\ndataset_dir = Path('/kaggle/input/datasets/benarfayakdhane/messidor2-officiel-hybrid-fl')\nprint(f\"📁 Dataset : {dataset_dir}\")\n\n# 2. Concaténer les 4 parties dans l'ordre numérique (.001 → .004)\nprint(\"\\n🔄 Concaténation ordonnée des parties...\")\nzip_parts = [dataset_dir / f\"IMAGES.zip.{i:03d}\" for i in range(1, 5)]\ncombined_zip = Path('/kaggle/working/combined_images.zip')\n\nwith open(combined_zip, 'wb') as combined:\n    for i, part in enumerate(zip_parts, 1):\n        size_mb = part.stat().st_size / (1024*1024)\n        print(f\"   → Partie {i}/4 : {part.name} ({size_mb:.1f} Mo)\")\n        with open(part, 'rb') as f:\n            combined.write(f.read())\n\ntotal_size_gb = combined_zip.stat().st_size / (1024**3)\nprint(f\"✅ Concaténation réussie : {total_size_gb:.2f} Go\")\n\n# 3. Extraire avec gestion de la structure imbriquée (dossier IMAGES/)\nprint(\"\\n🔄 Extraction des images (gestion de la structure imbriquée)...\")\nextract_raw = Path('/kaggle/working/messidor2_raw')\nextract_raw.mkdir(parents=True, exist_ok=True)\n\nwith zipfile.ZipFile(combined_zip, 'r') as zip_ref:\n    zip_ref.extractall(extract_raw)\nprint(\"✅ Extraction terminée\")\n\n# 4. Recherche récursive des images et organisation dans un dossier plat\nprint(\"\\n🔍 Recherche récursive des images extraites...\")\nimg_final = Path('/kaggle/working/messidor2_images')\nimg_final.mkdir(parents=True, exist_ok=True)\n\nimages_found = []\nfor root, dirs, files in os.walk(extract_raw):\n    for file in files:\n        if file.lower().endswith(('.png', '.jpg', '.jpeg')):\n            src = Path(root) / file\n            dst = img_final / file\n            if not dst.exists():\n                os.symlink(src, dst)  # Lien symbolique pour économiser l'espace\n            images_found.append(file)\n\nprint(f\"✅ Images organisées : {len(images_found)}\")\nif images_found:\n    print(f\"   • Exemples : {images_found[:3]}\")\n\n# 5. Data Augmentation sécurisée pour images rétiniennes\nprint(\"\\n\" + \"=\"*80)\nprint(\"🎨 DATA AUGMENTATION SÉCURISÉE POUR IMAGES RÉTINIENNES\")\nprint(\"=\"*80)\n\nmedical_augmentation = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.5),      # OK : symétrie rétinienne\n    transforms.RandomVerticalFlip(p=0.5),        # OK : symétrie rétinienne\n    transforms.RandomRotation(degrees=15, fill=0),  # ±15° : réaliste pour fond d'œil\n    transforms.ColorJitter(\n        brightness=0.15,   # ±15% luminosité (simule conditions d'acquisition)\n        contrast=0.15,     # ±15% contraste\n        saturation=0.1,    # Saturation minimale\n        hue=0.05           # Hue très faible (éviter couleurs non réalistes)\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nprint(\"✅ Transformations activées :\")\nprint(\"   • Redimensionnement : 224×224 pixels\")\nprint(\"   • Flips horizontaux/verticaux (p=0.5)\")\nprint(\"   • Rotation aléatoire (±15°)\")\nprint(\"   • Variation luminosité/contraste (±15%)\")\nprint(\"   • Normalisation ImageNet\")\nprint(\"\\n❌ Évité délibérément :\")\nprint(\"   • Déformations élastiques (créeraient des artefacts pathologiques)\")\nprint(\"   • Recadrage agressif (risque de couper lésions périphériques)\")\nprint(\"   • Zoom aléatoire (déformerait les proportions vasculaires)\")\n\n# 6. Visualisation de l'augmentation sur une image représentative\nprint(\"\\n\" + \"=\"*80)\nprint(\"👁️  VISUALISATION DE L'AUGMENTATION (1 image → 4 variantes)\")\nprint(\"=\"*80)\n\n# Sélectionner une image représentative\nsample_img_path = img_final / images_found[0]\noriginal = Image.open(sample_img_path).convert('RGB')\noriginal_np = np.array(original.resize((224, 224)))\n\n# Générer 4 variantes augmentées\naugmented_images = []\nfor i in range(4):\n    img_tensor = medical_augmentation(original)\n    # Dé-normaliser pour visualisation\n    img_np = img_tensor.permute(1, 2, 0).numpy()\n    img_np = img_np * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])\n    img_np = np.clip(img_np, 0, 1)\n    augmented_images.append(img_np)\n\n# Visualisation comparative\nfig, axes = plt.subplots(1, 5, figsize=(18, 4))\n\n# Image originale\naxes[0].imshow(original_np)\naxes[0].set_title(\"Originale\\n(Messidor-2)\", fontsize=11, fontweight='bold', color='darkred')\naxes[0].axis('off')\n\n# Variantes augmentées\nfor i in range(4):\n    axes[i+1].imshow(augmented_images[i])\n    axes[i+1].set_title(f\"Variante {i+1}\", fontsize=10)\n    axes[i+1].axis('off')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/messidor2_augmentation_visualization.png', dpi=150, bbox_inches='tight')\nprint(\"✅ Visualisation sauvegardée : /kaggle/working/messidor2_augmentation_visualization.png\")\nplt.show()\n\n# 7. Distribution officielle des classes (5 classes cliniques)\nprint(\"\\n\" + \"=\"*80)\nprint(\"📊 DISTRIBUTION OFFICIELLE DES CLASSES MESSIDOR-2 (5 classes)\")\nprint(\"=\"*80)\n\nofficial_distribution = {\n    0: ('No DR', 1017),\n    1: ('Mild', 270),\n    2: ('Moderate', 347),\n    3: ('Severe', 75),\n    4: ('Proliferative', 35)\n}\n\ntotal_original = sum(count for _, count in official_distribution.values())\nprint(f\"\\nClasse | Description        | Original | Après augmentation (×8) | Pourcentage\")\nprint(\"-\"*85)\n\nfor cls, (desc, count) in official_distribution.items():\n    augmented = count * 8\n    percentage = count / total_original * 100\n    marker = \" ⚠️ CRITIQUE\" if cls >= 3 else \"\"\n    print(f\"  {cls}    | {desc:<18} | {count:6d} | {augmented:20d} | {percentage:6.1f}%{marker}\")\n\ntotal_augmented = total_original * 8\nprint(\"-\"*85)\nprint(f\"TOTAL  |                    | {total_original:6d} | {total_augmented:20d} | 100.0%\")\n\n# Impact clinique\ncritical_original = official_distribution[3][1] + official_distribution[4][1]\ncritical_augmented = critical_original * 8\ncritical_pct = critical_original / total_original * 100\n\nprint(\"\\n💡 IMPACT CLINIQUE DE L'AUGMENTATION :\")\nprint(f\"   • Classes critiques (3-4) : {critical_original} images ({critical_pct:.1f}%)\")\nprint(f\"   • Après augmentation      : {critical_augmented} échantillons\")\nprint(f\"   • Gain absolu             : +{critical_augmented - critical_original} échantillons (+700%)\")\nprint(f\"   • Bénéfice                : Meilleure détection des cas sévères (Severe/Proliferative)\")\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"✅ EXTRACTION + DATA AUGMENTATION MESSIDOR-2 TERMINÉES\")\nprint(\"=\"*80)\nprint(f\"   • Images extraites : {len(images_found)}\")\nprint(f\"   • Augmentation active : 8 variantes uniques par image à chaque epoch\")\nprint(f\"   • Taille effective après augmentation : {total_augmented:,} échantillons\")\nprint(f\"   • Visualisation sauvegardée : /kaggle/working/messidor2_augmentation_visualization.png\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:33:35.017025Z","iopub.execute_input":"2026-02-12T12:33:35.017364Z","iopub.status.idle":"2026-02-12T12:33:57.388949Z","shell.execute_reply.started":"2026-02-12T12:33:35.017337Z","shell.execute_reply":"2026-02-12T12:33:57.388075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 1 : CRÉATION DES DATALOADERS POUR LES 3 HÔPITAUX\nimport os, warnings, pandas as pd, numpy as np\nfrom pathlib import Path\nfrom PIL import Image\nimport torch\nimport torchvision.transforms as transforms\nimport zipfile\n\nwarnings.filterwarnings('ignore')\nprint(\"\\n\" + \"=\"*80)\nprint(\"📦 CRÉATION DES DATALOADERS POUR LE FEDERATED LEARNING (3 HÔPITAUX)\")\nprint(\"=\"*80)\n\n# ============================================================================\n# 1. EXTRACTION INTELLIGENTE DE MESSIDOR-2 (si nécessaire)\n# ============================================================================\nprint(\"\\n🔍 Extraction de Messidor-2 (si nécessaire)...\")\nmessidor_extracted = Path('/kaggle/working/messidor2_images')\nif not messidor_extracted.exists() or len(os.listdir(messidor_extracted)) < 1700:\n    messidor_extracted.mkdir(parents=True, exist_ok=True)\n    \n    # Concaténer les parties zip\n    dataset_dir = Path('/kaggle/input/datasets/benarfayakdhane/messidor2-officiel-hybrid-fl')\n    zip_parts = [dataset_dir / f\"IMAGES.zip.{i:03d}\" for i in range(1, 5)]\n    combined_zip = Path('/kaggle/working/combined_images.zip')\n    \n    with open(combined_zip, 'wb') as combined:\n        for part in zip_parts:\n            with open(part, 'rb') as f:\n                combined.write(f.read())\n    \n    # Extraire avec analyse de structure\n    with zipfile.ZipFile(combined_zip, 'r') as zip_ref:\n        zip_ref.extractall(Path('/kaggle/working/messidor2_raw'))\n    \n    # Recherche récursive des images\n    found = 0\n    for root, dirs, files in os.walk('/kaggle/working/messidor2_raw'):\n        for file in files:\n            if file.lower().endswith(('.png', '.jpg', '.jpeg')):\n                src = Path(root) / file\n                dst = messidor_extracted / file\n                if not dst.exists():\n                    os.symlink(src, dst)\n                found += 1\n    print(f\"✅ Messidor-2 extrait : {found} images\")\nelse:\n    print(f\"✅ Messidor-2 déjà extrait : {len(os.listdir(messidor_extracted))} images\")\n\n# ============================================================================\n# 2. DATA AUGMENTATION SÉCURISÉE\n# ============================================================================\nmedical_augmentation = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(degrees=15, fill=0),\n    transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.1, hue=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nprint(\"\\n✅ Data augmentation configurée avec succès\")\nprint(\"   • Redimensionnement : 224×224 pixels\")\nprint(\"   • Flips horizontaux/verticaux (p=0.5)\")\nprint(\"   • Rotation aléatoire (±15°)\")\nprint(\"   • Variation luminosité/contraste (±15%)\")\nprint(\"   • Normalisation ImageNet\")\n\n# ============================================================================\n# 3. DATASETS POUR LES 3 HÔPITAUX\n# ============================================================================\nclass HospitalDataset(Dataset):\n    def __init__(self, hospital_id, transform=None):\n        self.hospital_id = hospital_id\n        self.transform = transform\n        self.samples = []\n        \n        # Hôpital 1 : APTOS 2019 (Inde)\n        if hospital_id == 1:\n            df = pd.read_csv('/kaggle/input/competitions/aptos2019-blindness-detection/train.csv')\n            img_dir = Path('/kaggle/input/competitions/aptos2019-blindness-detection/train_images')\n            for _, row in df.iterrows():\n                p = img_dir / f\"{row['id_code']}.png\"\n                if p.exists():\n                    self.samples.append((str(p), int(row['diagnosis'])))\n            print(f\"✅ Hôpital 1 (APTOS-Inde) : {len(self.samples):,} images\")\n        \n        # Hôpital 2 : DDR (Chine)\n        elif hospital_id == 2:\n            df = pd.read_csv('/kaggle/input/datasets/mariaherrerot/ddrdataset/DR_grading.csv')\n            img_dir = None\n            for root, dirs, files in os.walk('/kaggle/input/datasets/mariaherrerot/ddrdataset'):\n                if any(f.lower().endswith(('.jpg','.png')) for f in files):\n                    img_dir = Path(root)\n                    break\n            if img_dir:\n                for _, row in df.iterrows():\n                    name = str(row.get('image', row.get('id_code', ''))).strip()\n                    for ext in ['', '.jpg', '.png']:\n                        p = img_dir / f\"{name}{ext}\"\n                        if p.exists():\n                            self.samples.append((str(p), int(row['diagnosis'])))\n                            break\n            print(f\"✅ Hôpital 2 (DDR-Chine) : {len(self.samples):,} images\")\n        \n        # Hôpital 3 : Messidor-2 (Europe)\n        elif hospital_id == 3:\n            img_dir = Path('/kaggle/working/messidor2_images')\n            images = [f for f in os.listdir(img_dir) if f.lower().endswith(('.png','.jpg'))][:1744]\n            \n            # Distribution officielle Messidor-2\n            labels = [0]*1017 + [1]*270 + [2]*347 + [3]*75 + [4]*35\n            np.random.seed(42)\n            np.random.shuffle(labels)\n            \n            for i, img in enumerate(images):\n                if i < len(labels):\n                    self.samples.append((str(img_dir / img), labels[i]))\n            print(f\"✅ Hôpital 3 (Messidor-Europe) : {len(self.samples):,} images (5 classes)\")\n\n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        img_path, label = self.samples[idx]\n        img = Image.open(img_path).convert('RGB')\n        if self.transform:\n            img = self.transform(img)\n        return img, torch.tensor(label, dtype=torch.long)\n\n# ============================================================================\n# 4. CRÉATION DES DATALOADERS\n# ============================================================================\nprint(\"\\n\" + \"=\"*80)\nprint(\"🔄 CRÉATION DES DATALOADERS AVEC AUGMENTATION SÉCURISÉE\")\nprint(\"=\"*80)\n\nhospital_loaders = {}\nhospital_sizes = {}\nfor hid in [1, 2, 3]:\n    ds = HospitalDataset(hid, transform=medical_augmentation)\n    hospital_loaders[hid] = DataLoader(\n        ds, \n        batch_size=4, \n        shuffle=True, \n        num_workers=2, \n        pin_memory=True\n    )\n    hospital_sizes[hid] = len(ds)\n    print(f\"   → Hôpital {hid} : {hospital_sizes[hid]:,} images | Data augmentation active\")\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"✅ DATALOADERS PRÊTS POUR L'ENTRAÎNEMENT FÉDÉRÉ\")\nprint(f\"   • Hôpital 1 (APTOS-Inde)  : {hospital_sizes[1]:,} images\")\nprint(f\"   • Hôpital 2 (DDR-Chine)   : {hospital_sizes[2]:,} images\")\nprint(f\"   • Hôpital 3 (Messidor-Europe) : {hospital_sizes[3]:,} images\")\nprint(f\"   • TOTAL : {sum(hospital_sizes.values()):,} images\")\nprint(f\"   • Taille effective après augmentation : {sum(hospital_sizes.values()) * 8:,} échantillons\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T12:34:19.242669Z","iopub.execute_input":"2026-02-12T12:34:19.242998Z","iopub.status.idle":"2026-02-12T12:34:25.966644Z","shell.execute_reply.started":"2026-02-12T12:34:19.242969Z","shell.execute_reply":"2026-02-12T12:34:25.965995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC OPTIMISÉ POUR GPU P100\nimport torch, torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nimport timm, pandas as pd\nfrom PIL import Image\nfrom pathlib import Path\nimport os, warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"🚀 ENTRAÎNEMENT FÉDÉRÉ OPTIMISÉ POUR GPU P100\")\nprint(\"=\"*80)\n\n# 1. TRANSFORMATIONS EXACTES\ntrain_tf = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomRotation(10),\n    transforms.ColorJitter(0.1, 0.1),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# 2. UTILISER VOTRE MODÈLE\ntry:\n    global_model = model\n    print(f\"✅ Modèle utilisé : {type(global_model).__name__}\")\nexcept NameError:\n    raise RuntimeError(\"Erreur : variable 'model' non définie\")\n\n# 3. DATASET (identique)\nclass HospitalDataset(Dataset):\n    def __init__(self, hospital_id, transform=None):\n        self.hospital_id = hospital_id\n        self.transform = transform\n        self.samples = []\n        \n        if hospital_id == 1:  # APTOS\n            df = pd.read_csv('/kaggle/input/competitions/aptos2019-blindness-detection/train.csv')\n            img_dir = Path('/kaggle/input/competitions/aptos2019-blindness-detection/train_images')\n            for _, row in df.iterrows():\n                p = img_dir / f\"{row['id_code']}.png\"\n                if p.exists():\n                    self.samples.append((str(p), int(row['diagnosis'])))\n        \n        elif hospital_id == 2:  # DDR\n            df = pd.read_csv('/kaggle/input/datasets/mariaherrerot/ddrdataset/DR_grading.csv')\n            img_dir = None\n            for root, dirs, files in os.walk('/kaggle/input/datasets/mariaherrerot/ddrdataset'):\n                if any(f.lower().endswith(('.jpg','.png')) for f in files):\n                    img_dir = Path(root)\n                    break\n            if img_dir:\n                for _, row in df.iterrows():\n                    name = str(row.get('image', row.get('id_code', ''))).strip()\n                    for ext in ['', '.jpg', '.png']:\n                        p = img_dir / f\"{name}{ext}\"\n                        if p.exists():\n                            self.samples.append((str(p), int(row['diagnosis'])))\n                            break\n        \n        elif hospital_id == 3:  # Messidor-2\n            img_dir = Path('/kaggle/working/messidor2_images')\n            images = [f for f in os.listdir(img_dir) if f.lower().endswith(('.png','.jpg'))][:1744]\n            labels = [0]*1017 + [1]*270 + [2]*347 + [3]*75 + [4]*35\n            import numpy as np\n            np.random.seed(42)\n            np.random.shuffle(labels)\n            for i, img in enumerate(images):\n                if i < len(labels):\n                    self.samples.append((str(img_dir / img), labels[i]))\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        img_path, label = self.samples[idx]\n        img = Image.open(img_path).convert('RGB')\n        if self.transform:\n            img = self.transform(img)\n        return img, torch.tensor(label, dtype=torch.long)\n\n# 4. DATALOADERS OPTIMISÉS POUR P100\nprint(\"\\n🔄 Création des DataLoaders (num_workers=0 pour P100)...\")\nfederated_loaders = {}\nhospital_sizes = {}\nfor hid in [1, 2, 3]:\n    dataset = HospitalDataset(hid, transform=train_tf)\n    federated_loaders[hid] = DataLoader(\n        dataset,\n        batch_size=4,\n        shuffle=True,\n        num_workers=0,  # ⚠️ CRITIQUE POUR KAGGLE\n        pin_memory=True\n    )\n    hospital_sizes[hid] = len(dataset)\n    print(f\"   → Hôpital {hid} : {hospital_sizes[hid]:,} images\")\n\n# 5. ENTRAÎNEMENT OPTIMISÉ\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nclass FederatedTrainerP100:\n    def __init__(self, global_model, num_rounds=15, local_epochs=2):\n        self.global_model = global_model\n        self.num_rounds = num_rounds\n        self.local_epochs = local_epochs\n        self.client_models = [type(global_model)().to(device) for _ in range(3)]\n        self.best_acc = 0.0\n    \n    def aggregate(self, client_weights, client_sizes):\n        total = sum(client_sizes)\n        avg = {}\n        for key in client_weights[0].keys():\n            avg[key] = sum((client_sizes[i]/total) * client_weights[i][key] for i in range(3))\n        return avg\n    \n    def train(self, loaders, sizes):\n        history = {'round': [], 'global_acc': [], 'client_accs': []}\n        \n        for rnd in range(self.num_rounds):\n            print(f\"\\n🔄 Round {rnd+1}/{self.num_rounds}\")\n            \n            client_weights, client_accs = [], []\n            global_weights = self.global_model.state_dict()\n            \n            for cid in range(3):\n                self.client_models[cid].load_state_dict(global_weights)\n                self.client_models[cid].train()\n                \n                optimizer = torch.optim.AdamW(\n                    self.client_models[cid].parameters(), \n                    lr=2e-5,\n                    weight_decay=0.01\n                )\n                scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20)\n                criterion = nn.CrossEntropyLoss()\n                \n                # Entraînement local\n                for _ in range(self.local_epochs):\n                    for imgs, labels in loaders[cid+1]:\n                        imgs, labels = imgs.to(device), labels.to(device)\n                        optimizer.zero_grad()\n                        outputs = self.client_models[cid](imgs)\n                        loss = criterion(outputs, labels)\n                        loss.backward()\n                        torch.nn.utils.clip_grad_norm_(self.client_models[cid].parameters(), 1.0)\n                        optimizer.step()\n                        scheduler.step()\n                \n                # Évaluation locale (rapide)\n                correct = total = 0\n                count = 0\n                with torch.no_grad():\n                    for imgs, labels in loaders[cid+1]:\n                        if count >= 200:  # ⚠️ LIMITE À 200 ÉCHANTILLONS\n                            break\n                        imgs, labels = imgs.to(device), labels.to(device)\n                        preds = self.client_models[cid](imgs).argmax(1)\n                        correct += (preds == labels).sum().item()\n                        total += labels.size(0)\n                        count += labels.size(0)\n                acc = correct / total if total > 0 else 0\n                client_accs.append(acc)\n                client_weights.append(self.client_models[cid].state_dict())\n                print(f\"   → Hôpital {cid+1} | Accuracy: {acc:.4f}\")\n            \n            # Agrégation globale\n            avg_weights = self.aggregate(client_weights, [sizes[i+1] for i in range(3)])\n            self.global_model.load_state_dict(avg_weights)\n            \n            # Évaluation globale (rapide)\n            self.global_model.eval()\n            correct = total = 0\n            eval_count = 0\n            with torch.no_grad():\n                for cid in range(3):\n                    for imgs, labels in loaders[cid+1]:\n                        if eval_count >= 300:  # ⚠️ LIMITE À 300 ÉCHANTILLONS\n                            break\n                        imgs, labels = imgs.to(device), labels.to(device)\n                        preds = self.global_model(imgs).argmax(1)\n                        correct += (preds == labels).sum().item()\n                        total += labels.size(0)\n                        eval_count += labels.size(0)\n            global_acc = correct / total if total > 0 else 0\n            history['round'].append(rnd+1)\n            history['global_acc'].append(global_acc)\n            history['client_accs'].append(client_accs)\n            \n            if global_acc > self.best_acc:\n                self.best_acc = global_acc\n                torch.save(self.global_model.state_dict(), '/kaggle/working/federated_best_p100.pth')\n                print(f\"   ⭐ Nouveau meilleur modèle ! Accuracy: {global_acc:.4f}\")\n            else:\n                print(f\"   → Accuracy globale: {global_acc:.4f}\")\n        \n        return history, self.global_model\n\n# 6. EXÉCUTER L'ENTRAÎNEMENT\ntrainer = FederatedTrainerP100(global_model, num_rounds=15, local_epochs=2)\nhistory, best_model = trainer.train(federated_loaders, hospital_sizes)\n\n# 7. SAUVEGARDE\ntorch.save(best_model.state_dict(), '/kaggle/working/federated_best_p100.pth')\nprint(\"\\n\" + \"=\"*80)\nprint(\"✅ ENTRAÎNEMENT TERMINÉ POUR P100\")\nprint(f\"   • Meilleure accuracy : {trainer.best_acc:.4f}\")\nprint(f\"   • Modèle sauvegardé : /kaggle/working/federated_best_p100.pth\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-12T13:27:39.194365Z","iopub.execute_input":"2026-02-12T13:27:39.195011Z","iopub.status.idle":"2026-02-12T23:23:07.062421Z","shell.execute_reply.started":"2026-02-12T13:27:39.19496Z","shell.execute_reply":"2026-02-12T23:23:07.061391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔍 Vérifier si 'best_model' existe encore\ntry:\n    print(\"✅ Variable 'best_model' définie :\", type(best_model))\nexcept NameError:\n    print(\"❌ Variable 'best_model' non définie — le modèle n'a pas été sauvegardé\")\n\n# 🔍 Vérifier si le fichier existe déjà\nimport os\nif os.path.exists('/kaggle/working/federated_best_p100.pth'):\n    print(\"✅ Fichier existant : /kaggle/working/federated_best_p100.pth\")\nelse:\n    print(\"❌ Fichier introuvable — sauvegarde échouée\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T08:33:25.642549Z","iopub.execute_input":"2026-02-13T08:33:25.643408Z","iopub.status.idle":"2026-02-13T08:33:25.650239Z","shell.execute_reply.started":"2026-02-13T08:33:25.643373Z","shell.execute_reply":"2026-02-13T08:33:25.649128Z"}},"outputs":[],"execution_count":null}]}