{"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":22990,"databundleVersionId":2048213,"sourceType":"competition"},{"sourceId":14579787,"sourceType":"datasetVersion","datasetId":9313333},{"sourceId":14637593,"sourceType":"datasetVersion","datasetId":9350586},{"sourceId":732872,"sourceType":"modelInstanceVersion","modelInstanceId":558501,"modelId":571068},{"sourceId":733024,"sourceType":"modelInstanceVersion","modelInstanceId":558628,"modelId":571190},{"sourceId":734082,"sourceType":"modelInstanceVersion","modelInstanceId":559488,"modelId":572062}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install monai\n!pip install -q zarr imagecodecs","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-28T16:24:53.745796Z","iopub.execute_input":"2026-01-28T16:24:53.746466Z","iopub.status.idle":"2026-01-28T16:25:00.123582Z","shell.execute_reply.started":"2026-01-28T16:24:53.746429Z","shell.execute_reply":"2026-01-28T16:25:00.122329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model.py\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\n\n#https://github.com/clemkoa/u-net/blob/master/unet/unet.py\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        x = self.conv(x)\n        return x\n\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(Up, self).__init__()\n        self.up_scale = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)\n\n    def forward(self, x1, x2):\n        x2 = self.up_scale(x2)\n\n        diffY = x1.size()[2] - x2.size()[2]\n        diffX = x1.size()[3] - x2.size()[3]\n\n        x2 = F.pad(x2, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2])\n        x = torch.cat([x2, x1], dim=1)\n        return x\n\n\nclass DownLayer(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(DownLayer, self).__init__()\n        self.pool = nn.MaxPool2d(2, stride=2, padding=0)\n        self.conv = DoubleConv(in_ch, out_ch)\n\n    def forward(self, x):\n        x = self.conv(self.pool(x))\n        return x\n\n\nclass UpLayer(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(UpLayer, self).__init__()\n        self.up = Up(in_ch, out_ch)\n        self.conv = DoubleConv(in_ch, out_ch)\n\n    def forward(self, x1, x2):\n        a = self.up(x1, x2)\n        x = self.conv(a)\n        return x\n\n\nclass UNet(nn.Module):\n    def __init__(self, dimensions=1):\n        super(UNet, self).__init__()\n        self.conv1 = DoubleConv(3, 64)\n        self.down1 = DownLayer(64, 128)\n        self.down2 = DownLayer(128, 256)\n        self.down3 = DownLayer(256, 512)\n        self.down4 = DownLayer(512, 1024)\n        self.up1 = UpLayer(1024, 512)\n        self.up2 = UpLayer(512, 256)\n        self.up3 = UpLayer(256, 128)\n        self.up4 = UpLayer(128, 64)\n        self.last_conv = nn.Conv2d(64, dimensions, 1)\n\n    def forward(self, x):\n        x1 = self.conv1(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x1_up = self.up1(x4, x5)\n        x2_up = self.up2(x3, x1_up)\n        x3_up = self.up3(x2, x2_up)\n        x4_up = self.up4(x1, x3_up)\n        output = self.last_conv(x4_up)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-28T16:25:00.125515Z","iopub.execute_input":"2026-01-28T16:25:00.125863Z","iopub.status.idle":"2026-01-28T16:25:00.140363Z","shell.execute_reply.started":"2026-01-28T16:25:00.125827Z","shell.execute_reply":"2026-01-28T16:25:00.139156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom torch import nn\nimport gc\n\n# =========================================================\n# 1. CONFIGURATION DES CHEMINS\n# =========================================================\nTILES_IMG_DIR = \"/kaggle/input/hubmap-tiles-lite/hubmap-tiles-filtered/tiles/tiles/images\"\nTILES_MSK_DIR = \"/kaggle/input/hubmap-tiles-lite/hubmap-tiles-filtered/tiles/tiles/masks\"\nCKPT_PATH = \"/kaggle/input/bestmodel2/pytorch/default/1/best_model (1).pth\"\n\nOUTPUT_DIR = \"/kaggle/working/test_tiles_results\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nTHRESHOLD = 0.3 # Ton seuil optimisé\nNUM_TESTS = 50  # Nombre de tiles à tester\nNUM_PLOTS = 10  # Nombre de graphiques à afficher et enregistrer\n\n# =========================================================\n# 2. MÉTRIQUES ET UTILS\n# =========================================================\ndef calculate_dice(pred, target):\n    smooth = 1e-7\n    p = pred.flatten()\n    t = target.flatten()\n    intersection = np.sum(p * t)\n    return (2. * intersection + smooth) / (np.sum(p) + np.sum(t) + smooth)\n\ndef load_model():\n    print(f\"📦 Chargement du modèle...\")\n    model = UNet(dimensions=1).to(DEVICE)\n    ckpt = torch.load(CKPT_PATH, map_location=DEVICE)\n    state_dict = ckpt[\"model_state_dict\"] if \"model_state_dict\" in ckpt else ckpt\n    model.load_state_dict(state_dict)\n    model.eval()\n    return model\n\n# =========================================================\n# 3. EXÉCUTION DU TEST SUR TILES\n# =========================================================\ndef run_tile_comparison():\n    model = load_model()\n    \n    # Sélectionner 50 images au hasard\n    all_images = sorted([f for f in os.listdir(TILES_IMG_DIR) if f.endswith('.png')])\n    selected_files = np.random.choice(all_images, min(NUM_TESTS, len(all_images)), replace=False)\n    \n    dice_results = []\n    \n    print(f\"🚀 Analyse de {len(selected_files)} tiles...\")\n\n    for i, img_name in enumerate(selected_files):\n        # Chemins\n        msk_name = img_name.replace(\"_tile_\", \"_mask_\")\n        img_path = os.path.join(TILES_IMG_DIR, img_name)\n        msk_path = os.path.join(TILES_MSK_DIR, msk_name)\n        \n        # 1. Chargement (convert(\"RGB\") pour éviter le vert)\n        img_pil = Image.open(img_path).convert(\"RGB\")\n        img_np = np.array(img_pil)\n        \n        gt_pil = Image.open(msk_path).convert(\"L\")\n        gt_np = (np.array(gt_pil) > 0).astype(np.uint8)\n        \n        # 2. Inférence\n        img_t = torch.from_numpy(img_np.transpose(2,0,1)).float().unsqueeze(0).to(DEVICE) / 255.0\n        with torch.no_grad():\n            output = model(img_t)\n            logits = output[0] if isinstance(output, tuple) else output\n            pred_probs = torch.sigmoid(logits).cpu().numpy().squeeze()\n            pred_bin = (pred_probs > THRESHOLD).astype(np.uint8)\n        \n        # 3. Calcul du Dice\n        score = calculate_dice(pred_bin, gt_np)\n        dice_results.append(score)\n        \n        # 4. Visualisation (seulement pour les NUM_PLOTS premières)\n        if i < NUM_PLOTS:\n            fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n            \n            # Image Originale\n            axes[0].imshow(img_np)\n            axes[0].set_title(f\"Tile: {img_name}\")\n            \n            # Ground Truth (Masque réel)\n            axes[1].imshow(gt_np, cmap='gray')\n            axes[1].set_title(\"Vérité Terrain (GT)\")\n            \n            # Prédiction IA\n            axes[2].imshow(pred_bin, cmap='magma')\n            axes[2].set_title(f\"Prédiction IA\\nDice: {score:.4f}\")\n            \n            # Overlay (Superposition)\n            overlay = img_np.copy()\n            # On colore la prédiction en rouge sur l'image originale\n            overlay[pred_bin > 0] = [255, 0, 0] \n            axes[3].imshow(overlay)\n            axes[3].set_title(\"Overlay (IA en Rouge)\")\n            \n            for ax in axes: ax.axis('off')\n            \n            plt.savefig(os.path.join(OUTPUT_DIR, f\"tile_test_{i}.png\"), bbox_inches='tight')\n            plt.show()\n            \n        if (i+1) % 10 == 0:\n            print(f\"Progression : {i+1}/{NUM_TESTS} tiles traitées.\")\n\n    # 5. BILAN STATISTIQUE\n    dice_arr = np.array(dice_results)\n    print(\"\\n\" + \"=\"*40)\n    print(\"📊 RÉSULTATS FINAUX SUR LES TILES\")\n    print(f\"DICE MOYEN      : {np.mean(dice_arr):.4f}\")\n    print(f\"Écart-type      : {np.std(dice_arr):.4f}\")\n    print(f\"Meilleur Dice   : {np.max(dice_arr):.4f}\")\n    print(f\"Pire Dice       : {np.min(dice_arr):.4f}\")\n    print(\"=\"*40)\n    print(f\"📂 Les images de test sont dans : {OUTPUT_DIR}\")\n\nif __name__ == \"__main__\":\n    run_tile_comparison()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-28T16:25:00.141661Z","iopub.execute_input":"2026-01-28T16:25:00.141888Z","iopub.status.idle":"2026-01-28T16:25:07.641164Z","shell.execute_reply.started":"2026-01-28T16:25:00.141864Z","shell.execute_reply":"2026-01-28T16:25:07.640483Z"}},"outputs":[],"execution_count":null}]}