{"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":"gpu","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":2822109,"sourceType":"datasetVersion","datasetId":1725813},{"sourceId":13479521,"sourceType":"datasetVersion","datasetId":8557821},{"sourceId":14682499,"sourceType":"datasetVersion","datasetId":9379797},{"sourceId":14877635,"sourceType":"datasetVersion","datasetId":9518172},{"sourceId":2812287,"sourceType":"datasetVersion","datasetId":1719146}],"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":"# 🔹 BLOC 1 : IMPORTATIONS ET CONFIGURATION\nimport os, warnings, pandas as pd, numpy as np\nfrom pathlib import Path\nfrom PIL import Image\nimport torch, torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nimport timm\nimport glob\n\n# Reproductibilité\ntorch.manual_seed(42)\nnp.random.seed(42)\n\nwarnings.filterwarnings('ignore')\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"🚀 REPRISE DE L'ENTRAÎNEMENT FÉDÉRÉ (Rounds 12-15)\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:43:31.171174Z","iopub.execute_input":"2026-02-18T12:43:31.171443Z","iopub.status.idle":"2026-02-18T12:43:43.271454Z","shell.execute_reply.started":"2026-02-18T12:43:31.171421Z","shell.execute_reply":"2026-02-18T12:43:43.270701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 2 : HARMONISATION RAPIDE\ndef harmonize_image(img):\n    if img.mode != 'RGB':\n        img = img.convert('RGB')\n    return img.resize((224, 224), Image.BILINEAR)\n\ntrain_tf = transforms.Compose([\n    transforms.Lambda(harmonize_image),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.05, hue=0.02),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:44:51.354923Z","iopub.execute_input":"2026-02-18T12:44:51.355556Z","iopub.status.idle":"2026-02-18T12:44:51.360628Z","shell.execute_reply.started":"2026-02-18T12:44:51.355527Z","shell.execute_reply":"2026-02-18T12:44:51.360000Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 3 : ÉQUILIBRAGE SÉCURISÉ\ndef balance_classes(samples):\n    if not samples:\n        return []\n    from collections import Counter\n    class_counts = Counter([s[1] for s in samples])\n    max_count = max(class_counts.values())\n    balanced_samples = []\n    \n    for img_path, label in samples:\n        factor = 12 if label in [3, 4] else 8  # Boost critique\n        factor = min(factor, max(1, int(max_count / class_counts[label])))\n        for _ in range(factor):\n            balanced_samples.append((img_path, label))\n    return balanced_samples","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:44:54.942683Z","iopub.execute_input":"2026-02-18T12:44:54.943478Z","iopub.status.idle":"2026-02-18T12:44:54.948203Z","shell.execute_reply.started":"2026-02-18T12:44:54.943449Z","shell.execute_reply":"2026-02-18T12:44:54.947549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 4 : MODÈLE HYBRIDE AVEC TÊTES LOCALES + CHARGEMENT ROBUSTE\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\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.relu = nn.ReLU()\n        self.dropout = nn.Dropout(0.3)\n        self.local_heads = nn.ModuleDict({\n            '1': nn.Linear(512, num_classes),\n            '2': nn.Linear(512, num_classes),\n            '3': nn.Linear(512, num_classes)\n        })\n    \n    def forward(self, x, hospital_id=None):\n        res_feat = self.resnet(x)\n        deit_feat = self.deit(x)\n        fused = torch.cat([res_feat, deit_feat], dim=1)\n        features = self.dropout(self.relu(self.fusion(fused)))\n        \n        if hospital_id and hospital_id in self.local_heads:\n            return self.local_heads[hospital_id](features)\n        return self.local_heads['1'](features)  # Par défaut\n\n# Initialisation du modèle\nglobal_model = HybridResNet50_DeitBasic().to(device)\nprint(\"✅ Modèle initialisé.\")\n\n# 🛑 ÉTAPE CRITIQUE : RECHERCHE AUTOMATIQUE DU CHECKPOINT ROUND 11\nimport glob\ncheckpoint_found = False\ncheckpoint_path = None\n\n# Cherche récursivement dans TOUS les dossiers de /kaggle/input/\nmatches = glob.glob('/kaggle/input/**/model_round_11.pth', recursive=True)\n\nif matches:\n    checkpoint_path = matches[0] # Prend le premier trouvé\n    print(f\"\\n🔄 CHECKPOINT TROUVÉ : {checkpoint_path}\")\n    try:\n        global_model.load_state_dict(torch.load(checkpoint_path))\n        print(\"✅ Modèle du Round 11 chargé avec succès !\")\n        checkpoint_found = True\n    except Exception as e:\n        print(f\"❌ Erreur lors du chargement : {e}\")\nelse:\n    print(\"\\n⚠️ ATTENTION : Aucun fichier model_round_11.pth trouvé dans les inputs !\")\n    print(\"Vérifiez que le dataset 'Output 1' est bien ajouté dans le panneau Input à droite.\")\n    print(\"L'entraînement repartira de zéro si on continue.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:48:13.072423Z","iopub.execute_input":"2026-02-18T12:48:13.072792Z","iopub.status.idle":"2026-02-18T12:48:55.750167Z","shell.execute_reply.started":"2026-02-18T12:48:13.072762Z","shell.execute_reply":"2026-02-18T12:48:55.749511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 5 : DATASET AVEC IDRID (CORRIGÉ)\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:  # IDRiD\n            img_dir = Path('/kaggle/input/datasets/mariaherrerot/idrid-dataset/Imagenes/Imagenes')\n            if not img_dir.exists():\n                img_dir = Path('/kaggle/input/datasets/mariaherrerot/idrid-dataset/Imagenes')\n            df = pd.read_csv('/kaggle/input/datasets/mariaherrerot/idrid-dataset/idrid_labels.csv')\n            for _, row in df.iterrows():\n                img_name = str(row['id_code']).strip()\n                for ext in ['.jpg', '.png']:\n                    p = img_dir / f\"{img_name}{ext}\"\n                    if p.exists():\n                        self.samples.append((str(p), int(row['diagnosis'])))\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        try:\n            img = Image.open(img_path).convert('RGB')\n        except:\n            img = Image.new('RGB', (224, 224), (0, 0, 0))\n        if self.transform:\n            img = self.transform(img)\n        return img, torch.tensor(label, dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:55:07.805281Z","iopub.execute_input":"2026-02-18T12:55:07.806118Z","iopub.status.idle":"2026-02-18T12:55:07.816068Z","shell.execute_reply.started":"2026-02-18T12:55:07.806084Z","shell.execute_reply":"2026-02-18T12:55:07.815312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 6 : CRÉATION DES DATALOADERS\nprint(\"\\n🔄 Création des DataLoaders...\")\nfederated_loaders = {}\nhospital_sizes = {}\n\nfor hid in [1, 2, 3]:\n    dataset = HospitalDataset(hid, transform=train_tf)\n    if len(dataset.samples) == 0:\n        print(f\"⚠️ Attention: Hôpital {hid} n'a aucun échantillon. Vérifiez les chemins.\")\n        continue\n        \n    dataset.samples = balance_classes(dataset.samples)\n    federated_loaders[hid] = DataLoader(\n        dataset,\n        batch_size=8,          # Augmenté pour GPU P100\n        shuffle=True,\n        num_workers=2,         # Activé pour chargement parallèle\n        pin_memory=True\n    )\n    hospital_sizes[hid] = len(dataset.samples)\n    print(f\"   → Hôpital {hid} : {hospital_sizes[hid]:,} échantillons\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:55:12.336472Z","iopub.execute_input":"2026-02-18T12:55:12.337317Z","iopub.status.idle":"2026-02-18T12:55:19.868274Z","shell.execute_reply.started":"2026-02-18T12:55:12.337284Z","shell.execute_reply":"2026-02-18T12:55:19.867512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 7 : ENTRAÎNEUR OPTIMISÉ POUR REPRISE (Méthode de calcul identique à l'original)\nclass FederatedTrainerResume:\n    def __init__(self, global_model, start_round, total_rounds=15, local_epochs=2):\n        self.global_model = global_model\n        self.start_round = start_round \n        self.total_rounds = total_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_accuracies):\n        # Pondération inverse de l'erreur (comme dans l'original optimisé)\n        weights = [1.0 / (1.0 - acc + 1e-5) for acc in client_accuracies]\n        total = sum(weights)\n        weights = [w / total for w in weights]\n        avg = {}\n        for key in client_weights[0].keys():\n            if 'local_heads' in key: continue\n            avg[key] = sum(w * client_weights[i][key] for i, w in enumerate(weights))\n        return avg\n\n    def train(self, loaders, sizes):\n        history = {'round': [], 'global_acc': [], 'client_accs': []}\n        \n        # Boucle pour les rounds restants (ex: 12, 13, 14, 15)\n        for rnd in range(self.start_round, self.total_rounds):\n            print(f\"\\n🔄 Round {rnd+1}/{self.total_rounds} (Reprise depuis {self.start_round})\")\n            \n            client_weights, client_accs = [], []\n            global_weights = self.global_model.state_dict()\n            \n            # --- Entraînement Local ---\n            for cid in range(3):\n                if (cid+1) not in loaders: continue\n                self.client_models[cid].load_state_dict(global_weights)\n                self.client_models[cid].train()\n                \n                optimizer = torch.optim.AdamW(self.client_models[cid].parameters(), lr=2e-5, weight_decay=0.01)\n                scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20)\n                criterion = nn.CrossEntropyLoss()\n                \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, hospital_id=str(cid+1))\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 (limitée à 100 batches pour gagner du temps)\n                correct = total = 0\n                eval_batches = 0\n                with torch.no_grad():\n                    for imgs, labels in loaders[cid+1]:\n                        if eval_batches >= 100: break\n                        imgs, labels = imgs.to(device), labels.to(device)\n                        preds = self.client_models[cid](imgs, hospital_id=str(cid+1)).argmax(1)\n                        correct += (preds == labels).sum().item()\n                        total += labels.size(0)\n                        eval_batches += 1\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 Fédérée ---\n            avg_weights = self.aggregate(client_weights, client_accs)\n            global_state = self.global_model.state_dict()\n            global_state.update(avg_weights)\n            self.global_model.load_state_dict(global_state)\n            \n            # Sauvegarde du modèle du round\n            torch.save(self.global_model.state_dict(), f'/kaggle/working/model_round_{rnd+1}.pth')\n            print(f\"💾 Modèle sauvegardé : model_round_{rnd+1}.pth\")\n            \n            # --- ÉVALUATION GLOBALE (MÉTHODE IDENTIQUE À L'ORIGINAL) ---\n            # Cette section reproduit exactement la logique de federated-learning-finale-0 (6).ipynb\n            self.global_model.eval()\n            all_preds, all_labels = [], []\n            \n            with torch.no_grad():\n                for cid in range(3):\n                    if (cid+1) not in loaders: continue\n                    # On parcourt TOUTES les données pour le calcul global (pas de limite ici)\n                    for imgs, labels in loaders[cid+1]:\n                        imgs, labels = imgs.to(device), labels.to(device)\n                        # Utilisation explicite de la tête locale correspondante\n                        preds = self.global_model(imgs, hospital_id=str(cid+1)).argmax(1)\n                        all_preds.append(preds.cpu())\n                        all_labels.append(labels.cpu())\n            \n            # Calcul de la moyenne globale pondérée par la taille des datasets\n            if all_preds:\n                all_preds = torch.cat(all_preds)\n                all_labels = torch.cat(all_labels)\n                # C'est ici que se fait le calcul exact : (Vrais Positifs + Vrais Négatifs) / Total\n                global_acc = (all_preds == all_labels).float().mean().item()\n                \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_final_v2.pth')\n                    print(f\"   ⭐ Nouveau meilleur modèle ! Accuracy: {global_acc:.4f}\")\n                else:\n                    print(f\"   → Accuracy globale: {global_acc:.4f}\")\n            else:\n                print(\"   ⚠️ Aucune donnée pour évaluer la accuracy globale.\")\n        \n        return history, self.global_model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 🔹 BLOC 8 : EXÉCUTION FINALE\nif checkpoint_found:\n    # On lance seulement les 4 rounds restants (de 11 à 14, donc rounds 12, 13, 14, 15)\n    trainer = FederatedTrainerResume(global_model, start_round=11, total_rounds=15, local_epochs=2)\n    history, best_model = trainer.train(federated_loaders, hospital_sizes)\n    \n    torch.save(best_model.state_dict(), '/kaggle/working/federated_best_final_v2.pth')\n    print(\"\\n\" + \"=\"*80)\n    print(f\"✅ ENTRAÎNEMENT TERMINÉ (Rounds 12 à 15 effectués)\")\n    print(f\"• Meilleure accuracy finale : {trainer.best_acc:.4f}\")\n    print(\"=\"*80)\nelse:\n    print(\"\\n❌ Impossible de lancer la reprise car le checkpoint n'a pas été trouvé.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-18T12:55:38.464086Z","iopub.execute_input":"2026-02-18T12:55:38.464390Z","execution_failed":"2026-02-18T14:25:04.914Z"}},"outputs":[],"execution_count":null}]}