{"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":117682,"databundleVersionId":15062069,"sourceType":"competition"}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"THIS is draft ..... And a test ...","metadata":{}},{"cell_type":"code","source":"!pip install imagecodecs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T21:01:12.943697Z","iopub.execute_input":"2026-02-07T21:01:12.944523Z","iopub.status.idle":"2026-02-07T21:01:19.750912Z","shell.execute_reply.started":"2026-02-07T21:01:12.944488Z","shell.execute_reply":"2026-02-07T21:01:19.749969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import imagecodecs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T21:01:19.752561Z","iopub.execute_input":"2026-02-07T21:01:19.752858Z","iopub.status.idle":"2026-02-07T21:01:19.762630Z","shell.execute_reply.started":"2026-02-07T21:01:19.752825Z","shell.execute_reply":"2026-02-07T21:01:19.762086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tifffile import TiffFile","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T21:01:19.763619Z","iopub.execute_input":"2026-02-07T21:01:19.763919Z","iopub.status.idle":"2026-02-07T21:01:19.997192Z","shell.execute_reply.started":"2026-02-07T21:01:19.763889Z","shell.execute_reply":"2026-02-07T21:01:19.996545Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The \"Sandwich\" Concept\n\nIf you are trying to find ink on slice 100 and Z_CONTEXT = 2, the model doesn't just look at slice 100. It looks at a 5-layer \"sandwich\":\n\nSlice 98\n\nSlice 99\n\nSlice 100 (Target)\n\nSlice 101\n\nSlice 102\n\n2. Why is it used? (2.5D vs. 2D)\n\n2D Approach (Z_CONTEXT = 0): The model sees only 1 slice. It is hard to distinguish ink from charred papyrus because they look very similar on a single plane.\n\n2.5D Approach (Z_CONTEXT > 0): By providing layers above and below, the model can see the 3D structure of the ink. Real ink has a specific \"height\" and \"depth\" profile within the papyrus fibers that a single slice can't capture.","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom pathlib import Path\nfrom PIL import Image, ImageSequence\nfrom tqdm import tqdm\n\nDATA_PATH = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nOUT_DIR = Path(\"/kaggle/temp/vesuvius_npy\")\nOUT_DIR.mkdir(exist_ok=True, parents=True)\n\nIMG_SRC = DATA_PATH / \"train_images\"\nLBL_SRC = DATA_PATH / \"train_labels\"\n\ndef tiff_to_npy(tif_path: Path, out_path: Path):\n    if out_path.exists():\n        return\n    with Image.open(str(tif_path)) as img:\n        frames = [np.array(frame) for frame in ImageSequence.Iterator(img)]\n        vol = np.stack(frames, axis=0)\n    np.save(out_path, vol)\n\n# images\nfor tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\")), desc=\"Images\"):\n    out = OUT_DIR / f\"{tif.stem}_img.npy\"\n    tiff_to_npy(tif, out)\n\n# labels\nfor tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\")), desc=\"Labels\"):\n    out = OUT_DIR / f\"{tif.stem}_lbl.npy\"\n    tiff_to_npy(tif, out)\n\nprint(\"Preprocessing done. NPY volumes stored in:\", OUT_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T21:01:19.998537Z","iopub.execute_input":"2026-02-07T21:01:19.998802Z","iopub.status.idle":"2026-02-07T21:35:59.591729Z","shell.execute_reply.started":"2026-02-07T21:01:19.998779Z","shell.execute_reply":"2026-02-07T21:35:59.590909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport time\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport matplotlib.pyplot as plt\n\nfrom pathlib import Path\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Dataset\n\nfrom skimage.measure import label, regionprops\n\n# ==========================================\n# 1. CONFIGURATION\n# ==========================================\nMAX_BATCHES_PER_EPOCH = 1200\nMAX_TRAIN_TIME_HOURS = 6\nMAX_TRAIN_SECONDS = MAX_TRAIN_TIME_HOURS * 3600\n\nNPY_DIR = Path(\"/kaggle/temp/vesuvius_npy\")  # Preprocessed volumes\nCHECKPOINT_DIR = Path(\"checkpoints\")\nCHECKPOINT_DIR.mkdir(exist_ok=True)\n# For savinfg best model\nbest_g_loss = float('inf')\n\nPATCH_SIZE = 256\nZ_CONTEXT = 3\n# BATCH_SIZE = 64\nBATCH_SIZE = 160\nLEARNING_RATE_G = 2e-4\n# LEARNING_RATE_D = 1e-4\nLEARNING_RATE_D = 5e-5\nSAMPLES_PER_ID = 24\n\n\n# ==========================================\n# 2. PREFETCHER\n# ==========================================\nclass FastPrefetcher:\n    def __init__(self, loader, device):\n        self.loader = iter(loader)\n        self.device = device\n        self.stream = torch.cuda.Stream()\n        self.next_input = None\n        self.next_target = None\n        self.preload()\n\n    def preload(self):\n        try:\n            self.next_input, self.next_target = next(self.loader)\n        except StopIteration:\n            self.next_input = None\n            self.next_target = None\n            return\n\n        with torch.cuda.stream(self.stream):\n            self.next_input = self.next_input.to(self.device, non_blocking=True)\n            self.next_target = self.next_target.to(self.device, non_blocking=True)\n\n    def next(self):\n        torch.cuda.current_stream().wait_stream(self.stream)\n        inp = self.next_input\n        tgt = self.next_target\n        if inp is not None:\n            self.preload()\n        return inp, tgt\n\n\n# ==========================================\n# 3. DATASET\n# ==========================================\ndef normalize_patch(patch: np.ndarray) -> np.ndarray:\n    patch = patch.astype(np.float32)\n    return (patch - patch.mean()) / (patch.std() + 1e-6)\n\n\nclass VesuviusSurfaceDataset(Dataset):\n    def __init__(self, ids, npy_dir, samples_per_id):\n        self.patch_size = PATCH_SIZE\n        self.z_context = Z_CONTEXT\n        self.rng = np.random.default_rng(42)\n        self.coords = []\n        self.volumes = {}\n        self.labels = {}\n        self.npy_dir = npy_dir\n\n        for vid in ids:\n            img_npy = npy_dir / f\"{vid}_img.npy\"\n            lbl_npy = npy_dir / f\"{vid}_lbl.npy\"\n            if not img_npy.exists() or not lbl_npy.exists():\n                continue\n\n            img_vol = np.load(img_npy, mmap_mode=\"r\")\n            z_max, h, w = img_vol.shape\n\n            z_min = max(self.z_context, 5)\n            z_max_r = min(z_max - self.z_context - 1, z_max - 5)\n            if z_max_r <= z_min:\n                continue\n\n            zs = self.rng.integers(z_min, z_max_r, size=samples_per_id)\n            ys = self.rng.integers(0, max(1, h - self.patch_size), size=samples_per_id)\n            xs = self.rng.integers(0, max(1, w - self.patch_size), size=samples_per_id)\n\n            for z, y, x in zip(zs, ys, xs):\n                self.coords.append(\n                    {\n                        \"vid\": vid,\n                        \"z\": int(z),\n                        \"y\": int(y),\n                        \"x\": int(x),\n                        \"img_npy\": img_npy,\n                        \"lbl_npy\": lbl_npy,\n                    }\n                )\n\n    def __len__(self):\n        return len(self.coords)\n\n    def _load_volume_pair(self, item):\n        vid = item[\"vid\"]\n        if vid in self.volumes:\n            return self.volumes[vid], self.labels[vid]\n\n        img_vol = np.load(item[\"img_npy\"], mmap_mode=\"r\")\n        lbl_vol = np.load(item[\"lbl_npy\"], mmap_mode=\"r\")\n\n        self.volumes[vid] = img_vol\n        self.labels[vid] = lbl_vol\n\n        if len(self.volumes) > 8:\n            self.volumes.clear()\n            self.labels.clear()\n            self.volumes[vid] = img_vol\n            self.labels[vid] = lbl_vol\n\n        return img_vol, lbl_vol\n\n    def __getitem__(self, idx):\n        item = self.coords[idx]\n        img_vol, lbl_vol = self._load_volume_pair(item)\n\n        z, y, x = item[\"z\"], item[\"y\"], item[\"x\"]\n        img_patch = img_vol[\n            z - self.z_context : z + self.z_context + 1,\n            y : y + self.patch_size,\n            x : x + self.patch_size,\n        ]\n        img_patch = normalize_patch(img_patch)\n\n        z_lbl = z if lbl_vol.shape[0] > 1 else 0\n        target_patch = lbl_vol[z_lbl, y : y + self.patch_size, x : x + self.patch_size]\n\n        img_tensor = torch.from_numpy(img_patch.copy()).float()\n        target_tensor = torch.from_numpy(target_patch.copy()).float().unsqueeze(0)\n        return img_tensor, target_tensor\n\n\n# ==========================================\n# 4. MODELS\n# ==========================================\nclass Pix2PixGenerator(nn.Module):\n    def __init__(self, in_ch=5, out_ch=1):\n        super().__init__()\n\n        def block(in_c, out_c, down=True):\n            if down:\n                return nn.Sequential(\n                    nn.Conv2d(in_c, out_c, 4, 2, 1),\n                    nn.BatchNorm2d(out_c),\n                    nn.LeakyReLU(0.2, inplace=True),\n                )\n            else:\n                return nn.Sequential(\n                    nn.ConvTranspose2d(in_c, out_c, 4, 2, 1),\n                    nn.BatchNorm2d(out_c),\n                    nn.ReLU(inplace=True),\n                )\n\n        self.e1 = nn.Sequential(\n            nn.Conv2d(in_ch, 64, 4, 2, 1),\n            nn.LeakyReLU(0.2, inplace=True),\n        )\n        self.e2 = block(64, 128, down=True)\n        self.e3 = block(128, 256, down=True)\n\n        self.d1 = block(256, 128, down=False)\n        self.d2 = block(256, 64, down=False)\n        self.final = nn.ConvTranspose2d(128, out_ch, 4, 2, 1)\n\n    def forward(self, x):\n        s1 = self.e1(x)\n        s2 = self.e2(s1)\n        s3 = self.e3(s2)\n\n        u1 = self.d1(s3)\n        u2 = self.d2(torch.cat([u1, s2], dim=1))\n        out = self.final(torch.cat([u2, s1], dim=1))\n        return out\n\n\nclass TopologyDiscriminator(nn.Module):\n    def __init__(self, in_ch=6):\n        super().__init__()\n\n        def d_block(in_c, out_c):\n            return nn.Sequential(\n                nn.Conv2d(in_c, out_c, 4, 2, 1, bias=False),\n                nn.BatchNorm2d(out_c),\n                nn.LeakyReLU(0.2, inplace=True),\n            )\n\n        self.model = nn.Sequential(\n            d_block(in_ch, 64),\n            d_block(64, 128),\n            d_block(128, 256),\n            nn.Conv2d(256, 1, 4, 1, 1),\n        )\n\n    def forward(self, x, label):\n        return self.model(torch.cat([x, label], dim=1))\n\n\n# ==========================================\n# 5. LOSS & VISUALIZATION\n# ==========================================\ndef hybrid_loss(logits, target):\n    valid_mask = (target != 2).float()\n    target_binary = (target == 1).float()\n\n    bce = nn.functional.binary_cross_entropy_with_logits(\n        logits, target_binary, reduction=\"none\"\n    )\n    masked_bce = (bce * valid_mask).sum() / (valid_mask.sum() + 1e-6)\n\n    pred = torch.sigmoid(logits)\n    dice = 1 - (\n        2.0 * (pred * target_binary * valid_mask).sum() + 1e-6\n    ) / (\n        (pred * valid_mask).sum()\n        + (target_binary * valid_mask).sum()\n        + 1e-6\n    )\n\n    return 0.5 * masked_bce + 0.5 * dice\n\n\ndef visualize_cc_impact(model, dataset, device, epoch, num_samples=5, min_cc_size=150):\n    model.eval()\n    rows = num_samples\n    cols = 7  # NEW: added overlay-on-label column\n    fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4))\n\n    for i in range(rows):\n        idx = random.randint(0, len(dataset) - 1)\n        img, target = dataset[idx]\n\n        with torch.no_grad():\n            pred = torch.sigmoid(model(img.unsqueeze(0).to(device))).cpu().squeeze().numpy()\n\n        # --- Input slice ---\n        input_slice = img[Z_CONTEXT]\n\n        # --- Target (label) ---\n        target_np = target.squeeze().numpy()\n        target_viz = np.zeros((*target_np.shape, 3), dtype=np.float32)\n        target_viz[target_np == 1] = [1, 1, 1]       # white\n        target_viz[target_np == 2] = [0.5, 0.5, 0.5] # gray\n\n        # --- Raw prediction ---\n        raw_pred = pred\n\n        # --- Threshold ---\n        thr = raw_pred > 0.5\n\n        # --- Clean CC ---\n        lbl = label(thr)\n        cleaned = np.zeros_like(thr)\n        for region in regionprops(lbl):\n            if region.area >= min_cc_size:\n                cleaned[lbl == region.label] = 1\n\n        # --- Overlay on input ---\n        overlay_input = np.zeros((*cleaned.shape, 3), dtype=np.float32)\n        overlay_input[..., 2] = 0.3\n        overlay_input[..., 0] = cleaned * 1.0\n\n        # --- Overlay on label (NEW) ---\n        overlay_label = target_viz.copy()\n        overlay_label[..., 0] = np.maximum(overlay_label[..., 0], cleaned * 1.0)\n\n        # --- Plot ---\n        axes[i, 0].imshow(input_slice, cmap=\"gray\")\n        axes[i, 0].set_title(\"Input Slice\")\n\n        axes[i, 1].imshow(target_viz)\n        axes[i, 1].set_title(\"Target (Label)\")\n\n        axes[i, 2].imshow(raw_pred, cmap=\"magma\", vmin=0, vmax=1)\n        axes[i, 2].set_title(\"Raw Prediction\")\n\n        axes[i, 3].imshow(thr, cmap=\"gray\")\n        axes[i, 3].set_title(\"Thresholded\")\n\n        axes[i, 4].imshow(cleaned, cmap=\"gray\")\n        axes[i, 4].set_title(\"Cleaned CC\")\n\n        axes[i, 5].imshow(overlay_input)\n        axes[i, 5].set_title(\"Overlay on Input\")\n\n        axes[i, 6].imshow(overlay_label)\n        axes[i, 6].set_title(\"Overlay on Label (NEW)\")\n\n        for j in range(cols):\n            axes[i, j].axis(\"off\")\n\n    plt.tight_layout()\n    plt.savefig(f\"epoch_{epoch}_visual.png\")\n    plt.show()\n    plt.close()\n# ==========================================\n# 6. TRAINING LOOP\n# ==========================================\nif __name__ == \"__main__\":\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    train_ids = sorted({p.stem.replace(\"_img\", \"\") for p in NPY_DIR.glob(\"*_img.npy\")})\n\n    ds = VesuviusSurfaceDataset(train_ids, NPY_DIR, SAMPLES_PER_ID)\n\n    loader = DataLoader(\n        ds,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        drop_last=True,\n        num_workers=4,          # Kaggle sweet spot\n        pin_memory=True,\n        persistent_workers=True,\n    )\n\n    netG = Pix2PixGenerator(in_ch=2 * Z_CONTEXT + 1).to(device)\n    netD = TopologyDiscriminator(in_ch=2 * Z_CONTEXT + 2).to(device)\n\n    optG = optim.Adam(netG.parameters(), lr=LEARNING_RATE_G, betas=(0.5, 0.999))\n    optD = optim.Adam(netD.parameters(), lr=LEARNING_RATE_D, betas=(0.5, 0.999))\n    scaler = torch.amp.GradScaler(\"cuda\")\n\n    start_time = time.time()\n    epoch = 1\n\n    while (time.time() - start_time) < MAX_TRAIN_SECONDS:\n        netG.train()\n        netD.train()\n\n        total_batches = min(len(loader), MAX_BATCHES_PER_EPOCH)\n        prefetcher = FastPrefetcher(loader, device)\n\n        pbar = tqdm(range(total_batches), desc=f\"Epoch {epoch}\")\n        for _ in pbar:\n            imgs, targets = prefetcher.next()\n            if imgs is None:\n                break\n\n            # --- Train Discriminator ---\n            # optD.zero_grad(set_to_none=True)\n            # with torch.amp.autocast(\"cuda\"):\n            #     fake_targets = netG(imgs).detach()\n            #     d_real = netD(imgs, (targets == 1).float())\n            #     d_fake = netD(imgs, torch.sigmoid(fake_targets))\n\n            #     loss_D_real = nn.functional.mse_loss(d_real, torch.ones_like(d_real))\n            #     loss_D_fake = nn.functional.mse_loss(d_fake, torch.zeros_like(d_fake))\n            #     loss_D = loss_D_real + loss_D_fake\n\n            # scaler.scale(loss_D).backward()\n            # scaler.step(optD)\n\n            # # --- Train Generator ---\n            # optG.zero_grad(set_to_none=True)\n            # with torch.amp.autocast(\"cuda\"):\n            #     gen_logits = netG(imgs)\n            #     g_task_loss = hybrid_loss(gen_logits, targets)\n            #     g_adv_loss = nn.functional.mse_loss(\n            #         netD(imgs, torch.sigmoid(gen_logits)),\n            #         torch.ones_like(d_real),\n            #     )\n            #     loss_G = g_task_loss + 0.1 * g_adv_loss\n            if _ % 2 == 0: \n                optD.zero_grad(set_to_none=True)\n                with torch.amp.autocast(\"cuda\"):\n                    fake_targets = netG(imgs).detach()\n                    d_real = netD(imgs, (targets == 1).float())\n                    d_fake = netD(imgs, torch.sigmoid(fake_targets))\n\n                    loss_D_real = nn.functional.mse_loss(d_real, torch.ones_like(d_real))\n                    loss_D_fake = nn.functional.mse_loss(d_fake, torch.zeros_like(d_fake))\n                    loss_D = loss_D_real + loss_D_fake\n\n                scaler.scale(loss_D).backward()\n                scaler.step(optD)\n\n            # --- Train Generator (Updated every batch) ---\n            optG.zero_grad(set_to_none=True)\n            with torch.amp.autocast(\"cuda\"):\n                gen_logits = netG(imgs)\n                g_task_loss = hybrid_loss(gen_logits, targets)\n                # g_adv_loss = nn.functional.mse_loss(\n                #     netD(imgs, torch.sigmoid(gen_logits)),\n                #     torch.ones_like(d_real),\n                # )\n                \n                d_output = netD(imgs, torch.sigmoid(gen_logits))\n                g_adv_loss = nn.functional.mse_loss(\n                    d_output,\n                    torch.ones_like(d_output), # <--- Always matches current batch size\n                )\n                loss_G = g_task_loss + 0.1 * g_adv_loss\n            scaler.scale(loss_G).backward()\n            scaler.step(optG)\n            scaler.update()\n\n            pbar.set_postfix(\n                G_loss=f\"{loss_G.item():.4f}\",\n                D_loss=f\"{loss_D.item():.4f}\",\n            )\n        if epoch % 2 == 0:\n            visualize_cc_impact(netG, ds, device, epoch)\n        \n        torch.save(netG.state_dict(), CHECKPOINT_DIR / \"latest_model.pth\")\n\n        # 2. Save the \"Best\" checkpoint if loss improved\n        current_g_loss = loss_G.item()\n        if current_g_loss < best_g_loss:\n            best_g_loss = current_g_loss\n            # This saves the weights exactly like your old method\n            torch.save(netG.state_dict(), CHECKPOINT_DIR / \"best_model.pth\")\n            torch.save(netD.state_dict(), CHECKPOINT_DIR / \"netD_best_model.pth\")\n            print(f\"⭐ New best model saved (G_Loss: {best_g_loss:.4f})\")\n        # 3. Save every 7 epochs (Creates a new unique file)\n        if epoch % 10 == 0:\n            checkpoint_name = f\"model_epoch_{epoch}.pth\"\n            torch.save(netG.state_dict(), CHECKPOINT_DIR / checkpoint_name)\n            print(f\"💾 Periodic checkpoint saved: {checkpoint_name}\")\n            torch.save(netD.state_dict(), CHECKPOINT_DIR / f\"netD_epoch_{epoch}.pth\")\n        epoch += 1\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T21:55:59.241532Z","iopub.execute_input":"2026-02-07T21:55:59.241888Z","iopub.status.idle":"2026-02-07T22:10:06.835043Z","shell.execute_reply.started":"2026-02-07T21:55:59.241853Z","shell.execute_reply":"2026-02-07T22:10:06.820916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport zipfile\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom scipy import ndimage  # Added for Connected Components\n\n# ==========================================\n# 6. INFERENCE WITH CONNECTED COMPONENTS\n# ==========================================\n\ndef read_volume(path: Path):\n    try:\n        with Image.open(str(path)) as img:\n            frames = [np.array(frame) for frame in ImageSequence.Iterator(img)]\n            vol = np.stack(frames, axis=0)\n    except Exception: return np.zeros((1, PATCH_SIZE, PATCH_SIZE), dtype=np.uint8)\n    if vol.ndim == 2: vol = vol[None, ...]\n    return vol\ndef apply_cc_filter(mask: np.ndarray, min_size: int = 100):\n    \"\"\"\n    Removes small disconnected components from the binary mask.\n    min_size: Minimum pixel area to keep a component.\n    \"\"\"\n    # Label distinct islands of pixels\n    labeled, num_features = ndimage.label(mask)\n    if num_features == 0:\n        return mask\n    \n    # Count pixels in each component\n    component_sizes = np.bincount(labeled.ravel())\n    \n    # Create a mask of components that meet the size requirement\n    too_small = component_sizes < min_size\n    remove_mask = too_small[labeled]\n    \n    mask[remove_mask] = 0\n    return mask\n\ndef run_inference(model, device, min_cc_size=150):\n    \"\"\"\n    Inference loop with 2.5D stacking and CC filtering. [cite: 28]\n    \"\"\"\n    model.eval()\n    test_dir = DATA_PATH / 'test_images'\n    if not test_dir.exists():\n        print(\"Test directory not found. Skipping inference.\") [cite: 29]\n        return\n\n    test_ids = [f.stem for f in test_dir.glob('*.tif')]\n    submission_files = []\n\n    print(f\"🏁 Starting Inference with CC Filter (min_size={min_cc_size})...\")\n\n    for vid in test_ids:\n        vol = read_volume(test_dir / f\"{vid}.tif\")\n        z_max, h, w = vol.shape\n        volume_pages = []\n        \n        for z in tqdm(range(z_max), desc=f\"Infer {vid}\", leave=False):\n            z0, z1 = z - Z_CONTEXT, z + Z_CONTEXT + 1\n            \n            # Handling edges with blank pages [cite: 31]\n            if z0 < 0 or z1 > z_max:\n                volume_pages.append(Image.fromarray(np.zeros((h, w), dtype=np.uint8)))\n                continue\n                \n            page_pred = np.zeros((h, w), dtype=np.float32)\n            \n            # Tile processing to save VRAM [cite: 32]\n            for y in range(0, h, PATCH_SIZE):\n                for x in range(0, w, PATCH_SIZE):\n                    y1, x1 = min(y + PATCH_SIZE, h), min(x + PATCH_SIZE, w)\n                    y0, x0 = y1 - PATCH_SIZE, x1 - PATCH_SIZE\n                    \n                    patch = normalize_patch(vol[z0:z1, y0:y1, x0:x1])\n                    patch_t = torch.from_numpy(patch).float().unsqueeze(0).to(device)\n                    \n                    with torch.no_grad():\n                        # Use the generator (netG) to get logits\n                        logits = model(patch_t)\n                        pred = torch.sigmoid(logits).cpu().squeeze().numpy()\n                    \n                    # Accumulate predictions (useful if patches overlap)\n                    page_pred[y0:y1, x0:x1] = pred\n\n            # 1. Convert to Binary [cite: 35]\n            binary_mask = (page_pred > 0.5).astype(np.uint8)\n            \n            # 2. Apply Connected Components Filter\n            filtered_mask = apply_cc_filter(binary_mask, min_size=min_cc_size)\n            \n            volume_pages.append(Image.fromarray(filtered_mask))\n\n        # Save multi-page TIFF [cite: 36]\n        out_name = f\"{vid}.tif\"\n        if volume_pages:\n            volume_pages[0].save(out_name, save_all=True, append_images=volume_pages[1:], compression=\"tiff_deflate\")\n            submission_files.append(out_name)\n\n    # Final Zip [cite: 36]\n    with zipfile.ZipFile('submission.zip', 'w') as zipf:\n        for f in submission_files:\n            zipf.write(f)\n            os.remove(f) \n    print(\"📦 Submission.zip created with filtered results.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T22:10:15.375736Z","iopub.execute_input":"2026-02-07T22:10:15.376037Z","iopub.status.idle":"2026-02-07T22:10:15.418499Z","shell.execute_reply.started":"2026-02-07T22:10:15.376008Z","shell.execute_reply":"2026-02-07T22:10:15.417871Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # After training loop is complete:\n# if (CHECKPOINT_DIR / \"best_model.pth\").exists():\n#     checkpoint = torch.load(CHECKPOINT_DIR / \"best_model.pth\")\n#     # Wrap netG in DataParallel if using 2x T4 GPUs\n#     netG.load_state_dict(checkpoint['state_dict'])\n    \nrun_inference(netG, device, min_cc_size=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T22:11:30.668014Z","iopub.execute_input":"2026-02-07T22:11:30.668372Z","iopub.status.idle":"2026-02-07T22:11:38.788413Z","shell.execute_reply.started":"2026-02-07T22:11:30.668342Z","shell.execute_reply":"2026-02-07T22:11:38.787561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tifffile\nimport matplotlib.pyplot as plt\nimport zipfile\nimport os\nimport numpy as np\n\ndef visualize_submission_samples(zip_path='submission.zip', num_samples=3):\n    # 1. Extract the submission files\n    extract_path = 'temp_viz'\n    with zipfile.ZipFile(zip_path, 'r') as zip_ref:\n        zip_ref.extractall(extract_path)\n    \n    tif_files = [f for f in os.listdir(extract_path) if f.endswith('.tif')]\n    \n    for tif_file in tif_files[:num_samples]:\n        file_path = os.path.join(extract_path, tif_file)\n        \n        # Load the generated volume (Binary masks after CC filter)\n        pred_volume = tifffile.imread(file_path)\n        \n        # Select a middle slice for visualization\n        mid_idx = pred_volume.shape[0] // 2\n        pred_slice = pred_volume[mid_idx]\n        \n        # Load original image for context (if available in test_images)\n        test_img_path = DATA_PATH / 'test_images' / tif_file\n        if test_img_path.exists():\n            orig_volume = read_volume(test_img_path)\n            orig_slice = orig_volume[mid_idx]\n        else:\n            orig_slice = np.zeros_like(pred_slice)\n\n        # Plotting\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n        \n        axes[0].imshow(orig_slice, cmap='gray')\n        axes[0].set_title(f\"Original Slice (z={mid_idx})\")\n        \n        axes[1].imshow(pred_slice, cmap='magma')\n        axes[1].set_title(\"Prediction (After CC Filter)\")\n        \n        # Overlay to see alignment\n        axes[2].imshow(orig_slice, cmap='gray')\n        axes[2].imshow(pred_slice, cmap='jet', alpha=0.4)\n        axes[2].set_title(\"Overlay (Surface Alignment)\")\n        \n        for ax in axes: ax.axis('off')\n        plt.suptitle(f\"Sample: {tif_file}\")\n        plt.show()\n\n# Run the visualization\nif os.path.exists('submission.zip'):\n    visualize_submission_samples()\nelse:\n    print(\"submission.zip not found. Run inference first!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T22:11:44.890672Z","iopub.execute_input":"2026-02-07T22:11:44.891311Z","iopub.status.idle":"2026-02-07T22:11:46.965371Z","shell.execute_reply.started":"2026-02-07T22:11:44.891265Z","shell.execute_reply":"2026-02-07T22:11:46.964498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nimport random\nfrom scipy import ndimage\n\ndef visualize_cc_impact(generator, dataset, device, num_samples=5, min_cc_size=150):\n    generator.eval()\n    \n    # Select random indices from the dataset\n    indices = random.sample(range(len(dataset)), min(num_samples, len(dataset)))\n    \n    for idx in indices:\n        img, target = dataset[idx]\n        \n        with torch.no_grad():\n            # Get model output (logits) using netG \n            logits = generator(img.unsqueeze(0).to(device))\n            # Convert to probability map \n            prob_map = torch.sigmoid(logits).cpu().squeeze().numpy()\n            \n        # 1. Prediction Before CC Filter (Thresholded at 0.5) \n        pre_filter_mask = (prob_map > 0.5).astype(np.uint8)\n        \n        # 2. Manual CC Filter application for visualization \n        labeled, num_features = ndimage.label(pre_filter_mask)\n        component_sizes = np.bincount(labeled.ravel())\n        too_small = component_sizes < min_cc_size\n        remove_mask = too_small[labeled]\n        \n        post_filter_mask = pre_filter_mask.copy()\n        post_filter_mask[remove_mask] = 0\n        \n        # Plotting the 4-stage pipeline\n        fig, axes = plt.subplots(1, 4, figsize=(24, 6))\n        \n        # Show mid-slice of the 2.5D context \n        img_mid = img[Z_CONTEXT].numpy() \n        axes[0].imshow(img_mid, cmap='gray')\n        axes[0].set_title(f\"Input Mid-Slice (Idx: {idx})\")\n        \n        # Show raw confidence/probability map\n        axes[1].imshow(prob_map, cmap='magma')\n        axes[1].set_title(\"Raw Probabilities (Before Filter)\")\n        \n        # Show binary mask with noise\n        axes[2].imshow(pre_filter_mask, cmap='gray')\n        axes[2].set_title(\"Thresholded (Pre-CC Filter)\")\n        \n        # Show cleaned result\n        axes[3].imshow(post_filter_mask, cmap='jet')\n        axes[3].set_title(f\"Cleaned (Post-CC > {min_cc_size}px)\")\n        \n        for ax in axes: ax.axis('off')\n        plt.tight_layout()\n        plt.show()\n\n# Run using 'netG' as defined in your training script \nvisualize_cc_impact(netG, ds, device, num_samples=5, min_cc_size=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-07T22:12:19.244258Z","iopub.execute_input":"2026-02-07T22:12:19.244579Z","iopub.status.idle":"2026-02-07T22:12:23.017876Z","shell.execute_reply.started":"2026-02-07T22:12:19.244551Z","shell.execute_reply":"2026-02-07T22:12:23.016620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-07T20:51:46.850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}