{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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":14443416,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Model Definition\n***Significance and Integration of the Model Definition***\n\nThe model_definition file, containing the `FullUNet` class, is the core of the entire pipeline—it represents the machine learning algorithm that learns to solve the problem.\n\n***The Learner:*** The model is the function ($f$) that takes the input X (the Vesuvius X-ray image patch) and predicts the output Y (the probability map of ink). Its parameters are what the training process (driven by the loss function and optimizer) modifies to improve accuracy.\n\n**Architecture Choice (U-Net):** The U-Net architecture is specifically designed for image segmentation tasks like ink detection. It features a symmetric encoder-decoder structure with skip connections. These skip connections are crucial because they allow fine-grained, low-level detail from the encoder to pass directly to the decoder, helping the model localize the precise boundaries of the ink, which is critical for good Dice and Topological scores.\n\n**Output Format:** Crucially, the model's final layer uses a Sigmoid activation. This ensures its output, $Y_{hat}$, is a floating-point tensor of the same size as the input label, with values ranging from $0$ to $1$. These values are interpreted as the probability of each pixel belonging to the ink class, which is the exact format required by your HybridTopoLossOrchestrator for differentiation.","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass SEBlock(nn.Module):\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.fc = nn.Sequential(\n            nn.Conv2d(channels, channels // reduction, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(channels // reduction, channels, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        w = self.fc(x)\n        return x * w\n\n\nclass Down(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_ch, out_ch)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, out_ch, bilinear=True):\n        super().__init__()\n\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode=\"bilinear\", align_corners=True)\n            self.conv = DoubleConv(in_ch, out_ch)\n        else:\n            self.up = nn.ConvTranspose2d(in_ch // 2, in_ch // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_ch, out_ch)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.conv = nn.Conv2d(in_ch, out_ch, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)\n\n\nclass FullUNet(nn.Module):\n    def __init__(self, n_channels=3, n_classes=3):\n        super().__init__()\n\n        self.inc = DoubleConv(n_channels, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        self.down4 = Down(512, 1024)\n\n        self.up1 = Up(1024 + 512, 512)\n        self.up2 = Up(512 + 256, 256)\n        self.up3 = Up(256 + 128, 128)\n        self.up4 = Up(128 + 64, 64)\n\n        self.outc = OutConv(64, n_classes)\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        logits = self.outc(x)\n        return logits\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T04:52:42.648167Z","iopub.execute_input":"2025-12-09T04:52:42.648330Z","iopub.status.idle":"2025-12-09T04:52:45.935186Z","shell.execute_reply.started":"2025-12-09T04:52:42.648315Z","shell.execute_reply":"2025-12-09T04:52:45.934625Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patching the train data\n**VesuviusPatchDataset, Analysis and Explanation**\n\nThis class inherits from PyTorch's Dataset and provides data samples to the DataLoader. It transforms the large Vesuvius TIFF files into small, optimized training patches.\n\n**The Core Role:** Three Outputs for Hybrid Loss\nThe primary distinction of this dataset is that its __getitem__ method is designed to return three essential tensors for every sample, preparing the data specifically for the HybridTopoLossOrchestrator:\n\n1. $X$ (Input Patch): The grayscale image slice.\n2. $Y_{true}$ (Binary Ground Truth): The standard ink/no-ink label used for $L_{BCE}$ and $L_{Dice}$.\n3. $Y_{dist}$ (Distance Map): The Euclidean Distance Transform (EDT) of $Y_{true}$, which provides topological shape information for $L_{Topo}$.\n\n**Patch Indexing Strategy (_prepare_patch_index):**\nThis function executes the complex, one-time task of deciding which $(y, x)$ coordinates to sample from the full images:\n    * **Sliding Window:** It iterates over the entire binary label mask using a sliding window defined by patch_size and stride.\n* **Foreground/Background Separation:** For every window location, it checks if patch.sum() > 0. If so, it's a Foreground (FG) patch and the coordinates are saved to coords_fg; otherwise, it's a Background (BG) patch.\n* **Controlled Sampling:** It uses the max_patches_per_image and oversample_fg to calculate how many FG and BG patches it needs. It then uses random.sample to select the required number of coordinates from both lists.\n* **Final Index:** The final self.patch_index list contains the coordinates for every single training sample in a tuple: (image_id, y_start, x_start).\n\nThis strategy is highly efficient because the slow coordinate calculation happens only once at startup.\n\n**Item Retrieval (__getitem__)**\nThis function is called by the DataLoader for every batch during training. It fetches the data for a single sample:\n1. **Loading and Cropping:** It loads the full TIFF files for the image (img) and the label (lbl) and extracts the precise patches defined by the precomputed $(y, x)$ coordinates.\n2. **Augmentation:** If enabled, it applies the random flips and rotations using albumentations.\n3. **CRITICAL STEP:** Distance Transform ($Y_{dist}$)dist_map_np = distance_transform_edt(lbl_patch). This function (from scipy.ndimage) is the source of the topological information. It calculates the shortest distance from every background pixel (0) to the nearest foreground pixel (1)Pixels far from the ink will have a large value; pixels right at the boundary will have a low value. This map provides a continuous, geometric representation of the ink's shape, which is what the topological loss uses.\n4. **Tensor Conversion:** The NumPy arrays are converted to PyTorch tensors.\n    * **Normalization:** The image patch ($X$) is normalized by $255.0$ to keep values between $[0, 1]$.\n    * **Channel Dimension:** unsqueeze(0) adds the necessary channel dimension (C=1) to all tensors, making them compatible with PyTorch's (N, C, H, W) format.\n\nOutput: Returns the final triplet: img_patch ($X$), lbl_patch ($Y_{true}$), and dist_map ($Y_{dist}$).","metadata":{}},{"cell_type":"code","source":"# %%writefile vesuvius_dataset.py\nimport os\nimport random\nfrom typing import List, Dict, Tuple, Optional, Any\nimport numpy as np\nimport tifffile\nfrom scipy.ndimage import distance_transform_edt\nimport torch\nfrom torch.utils.data import Dataset\nfrom tqdm import tqdm\nfrom PIL import Image\nfrom concurrent.futures import ProcessPoolExecutor\nimport multiprocessing\nimport math\nimport traceback\n\n# -----------------------\n# Helper I/O and sampling\n# -----------------------\ndef _safe_read_tiff(path: str) -> np.ndarray:\n    \"\"\"\n    Robustly read a TIFF that may be LZW-compressed and/or multi-page (Z, H, W).\n    - Try tifffile.memmap first (fast & uses installed imagecodecs).\n    - Fallback to PIL multi-page reading (handles LZW and page collection).\n    Returns a numpy ndarray with shape (Z, H, W) or raises RuntimeError.\n    \"\"\"\n    # 1. Try tifffile.memmap (will use imagecodecs if available)\n    try:\n        arr = tifffile.memmap(path)\n        arr = np.asarray(arr)\n        if arr.ndim == 3:\n            return arr\n    except Exception:\n        pass # Proceed to PIL fallback on failure or non-3D output\n\n    # 2. PIL multi-page fallback (robustly reads Z-stack)\n    try:\n        img = Image.open(path)\n        pages = []\n        i = 0\n        while True:\n            try:\n                img.seek(i)\n            except (EOFError, IndexError):\n                break\n            except Exception:\n                 break \n            \n            pages.append(np.array(img))\n            i += 1\n\n        if len(pages) == 0:\n            raise RuntimeError(\"PIL failed to extract any pages.\")\n            \n        pages = np.stack(pages, axis=0)\n        \n        # Handle case where PIL read a single page and result is 2D\n        if pages.ndim == 2:\n            pages = pages[np.newaxis, ...]\n            \n        if pages.ndim != 3:\n            raise RuntimeError(f\"PIL fallback resulted in non-3D array: {pages.shape}\")\n            \n        return pages\n    \n    except Exception as e:\n        # If both fail, raise a clear error message\n        raise RuntimeError(f\"Failed to read TIFF {path} after all attempts: {e}\")\n\n\ndef _safe_sample_list(lst: List[Any], n: int, rng: Optional[random.Random] = None) -> List[Any]:\n    \"\"\"\n    Safe sampling helper that is guaranteed not to raise 'Sample larger than population'.\n    It always clamps the requested size 'n' to the population size 'L'.\n    \"\"\"\n    if not lst or n <= 0:\n        return []\n    \n    L = len(lst)\n    n_clamped = min(n, L)\n    \n    if n_clamped >= L:\n        out = lst.copy()\n        (rng if rng else random).shuffle(out)\n        return out\n        \n    try:\n        sampler = rng.sample if rng else random.sample\n        return sampler(lst, n_clamped)\n    except Exception:\n        stride = max(1, L // n_clamped)\n        return [lst[i * stride] for i in range(n_clamped)]\n\n\n# -----------------------\n# Worker function for parallel indexing\n# -----------------------\ndef _index_single_image_worker(args: Dict) -> List[Dict]:\n    \"\"\"\n    Index patches for a single image id.\n    Returns list of dicts: {'id','z','y','x','is_fg'}\n    \"\"\"\n    try:\n        image_id = args[\"image_id\"]\n        images_dir = args[\"images_dir\"]\n        labels_dir = args[\"labels_dir\"]\n        patch_size = args[\"patch_size\"]\n        stride = args[\"stride\"]\n        num_slices = args[\"num_slices\"]\n        min_foreground_pixels = args[\"min_foreground_pixels\"]\n        oversample_fg = args[\"oversample_fg\"]\n        max_patches_per_image = args[\"max_patches_per_image\"]\n        seed = args.get(\"seed\", None)\n\n        rng = random.Random(f\"{seed}_{image_id}\") if seed is not None else random.Random(image_id)\n\n        img_fn = image_id + \".tif\"\n        img_path = os.path.join(images_dir, img_fn)\n        lbl_path = os.path.join(labels_dir, img_fn)\n\n        if not os.path.exists(img_path) and os.path.exists(os.path.join(images_dir, image_id + \".tiff\")):\n            img_path = os.path.join(images_dir, image_id + \".tiff\")\n        if not os.path.exists(lbl_path) and os.path.exists(os.path.join(labels_dir, image_id + \".tiff\")):\n            lbl_path = os.path.join(labels_dir, image_id + \".tiff\")\n\n        if not os.path.exists(img_path) or not os.path.exists(lbl_path):\n            return []\n\n        try:\n            # Only read label to determine foreground/background (Faster indexing)\n            lbl_vol = _safe_read_tiff(lbl_path)\n        except Exception:\n            return []\n            \n        lbl_vol = np.asarray(lbl_vol)\n        if lbl_vol.ndim != 3:\n            return []\n\n        lbl_vol = (lbl_vol > 0).astype(np.uint8)\n        Z, H, W = lbl_vol.shape\n\n        half = num_slices // 2\n        z_start = half\n        z_end = Z - (num_slices - 1 - half)\n        if z_end <= z_start:\n            return []\n\n        step_z = max(1, num_slices // 2)\n        candidate_z = list(range(z_start, z_end, step_z))\n\n        all_coords = []\n        for zc in candidate_z:\n            sl = lbl_vol[zc]\n            if H < patch_size or W < patch_size:\n                continue\n            for y in range(0, H - patch_size + 1, stride):\n                for x in range(0, W - patch_size + 1, stride):\n                    patch_mask = sl[y:y + patch_size, x:x + patch_size]\n                    is_fg = bool(patch_mask.sum() >= min_foreground_pixels)\n                    all_coords.append({'id': image_id, 'z': int(zc), 'y': int(y), 'x': int(x), 'is_fg': is_fg})\n\n        if not all_coords:\n            return []\n\n        fg = [c for c in all_coords if c['is_fg']]\n        bg = [c for c in all_coords if not c['is_fg']]\n\n        num_total = min(len(all_coords), int(max_patches_per_image))\n        num_total = max(0, num_total)\n\n        num_fg_desired = int(round(num_total * float(oversample_fg)))\n        num_fg = max(0, min(num_fg_desired, len(fg)))\n        num_bg = max(0, min(num_total - num_fg, len(bg)))\n\n        selected: List[Dict] = []\n        selected.extend(_safe_sample_list(fg, num_fg, rng=rng))\n        selected.extend(_safe_sample_list(bg, num_bg, rng=rng))\n\n        remaining_needed = num_total - len(selected)\n        if remaining_needed > 0:\n            pool = [c for c in all_coords if c not in selected]\n            selected.extend(_safe_sample_list(pool, remaining_needed, rng=rng))\n\n        if len(selected) > num_total:\n            selected = selected[:num_total]\n\n        return selected\n\n    except Exception:\n        # print(f\"Error in worker for {args.get('image_id')}: {traceback.format_exc()}\")\n        return []\n\n\n# -----------------------\n# Main Dataset\n# -----------------------\nclass VesuviusPatchDataset(Dataset):\n    \"\"\"\n    Multi-slice 2D patch dataset for Vesuvius surface detection.\n    CRITICALLY: It caches all preprocessed patches into memory at initialization\n    to eliminate slow I/O during the training loop's __getitem__.\n    \"\"\"\n    PatchCacheEntry = Tuple[np.ndarray, np.ndarray, np.ndarray]\n\n    def __init__(\n        self,\n        root_dir: str = \"/kaggle/input/vesuvius-challenge-surface-detection\",\n        split: str = \"train\",\n        patch_size: int = 384,\n        stride: int = 96,\n        num_slices: int = 10,\n        oversample_fg: float = 0.6,\n        max_patches_per_image: int = 100,\n        augment: bool = False,\n        min_foreground_pixels: int = 50,\n        seed: Optional[int] = 1337,\n        num_index_workers: Optional[int] = 4, # Defaulting to 4 for better indexing speed\n    ):\n        super().__init__()\n        self.root_dir = root_dir\n        self.split = split\n        self.patch_size = int(patch_size)\n        self.stride = int(stride)\n        self.num_slices = int(num_slices)\n        self.oversample_fg = float(oversample_fg)\n        self.max_patches_per_image = int(max_patches_per_image)\n        self.augment = bool(augment) \n        self.num_index_workers = int(num_index_workers)\n        self.min_foreground_pixels = int(min_foreground_pixels)\n        self.seed = seed\n\n        # Use the passed worker count\n        self.num_index_workers = max(1, int(num_index_workers) if num_index_workers is not None else 1)\n\n        self.images_dir = os.path.join(self.root_dir, f\"{self.split}_images\")\n        self.labels_dir = os.path.join(self.root_dir, f\"{self.split}_labels\")\n        self.patch_index: List[Dict] = []\n        self.ids: List[str] = []\n        self.cached_patches: List[VesuviusPatchDataset.PatchCacheEntry] = []\n\n        # --- 1. Discovery ---\n        if not os.path.exists(self.images_dir) or not os.path.exists(self.labels_dir):\n            raise RuntimeError(f\"Missing expected dirs: {self.images_dir} or {self.labels_dir}\")\n        image_files = sorted([f for f in os.listdir(self.images_dir) if f.lower().endswith((\".tif\", \".tiff\"))])\n        if len(image_files) == 0:\n            raise RuntimeError(f\"No TIFF images found in {self.images_dir}.\")\n\n        for fn in image_files:\n            base = os.path.splitext(fn)[0]\n            if any(os.path.exists(os.path.join(self.labels_dir, base + ext)) for ext in [\".tif\", \".tiff\"]):\n                self.ids.append(base)\n        if len(self.ids) == 0:\n            raise RuntimeError(\"No matching image/label pairs found.\")\n        print(f\"Found {len(self.ids)} image-label pairs.\")\n\n        # --- 2. Parallel Indexing (Finding Coordinates) ---\n        print(f\"Indexing patches using {self.num_index_workers} workers...\")\n        self._parallel_index_images()\n\n        # --- 3. Sequential Patch Loading (The new Caching step) ---\n        if len(self.patch_index) == 0:\n            raise RuntimeError(\"No patches indexed. Check volumes/parameters.\")\n            \n        print(f\"Indexed total coordinates: {len(self.patch_index)}. Now loading all patches into memory...\")\n        self._load_patches_into_cache()\n        \n        if len(self.cached_patches) == 0:\n            raise RuntimeError(\"Failed to load any patches into cache.\")\n            \n        print(f\"Successfully loaded {len(self.cached_patches)} patches into memory (ready for fast training).\")\n\n    def _parallel_index_images(self):\n        args_iterable = ({\n            \"image_id\": image_id,\n            \"images_dir\": self.images_dir,\n            \"labels_dir\": self.labels_dir,\n            \"patch_size\": self.patch_size,\n            \"stride\": self.stride,\n            \"num_slices\": self.num_slices,\n            \"min_foreground_pixels\": self.min_foreground_pixels,\n            \"oversample_fg\": self.oversample_fg,\n            \"max_patches_per_image\": self.max_patches_per_image,\n            \"seed\": self.seed,\n        } for image_id in self.ids)\n\n        with ProcessPoolExecutor(max_workers=self.num_index_workers) as ex:\n            results = ex.map(_index_single_image_worker, args_iterable, chunksize=1)\n            for selected in tqdm(results, total=len(self.ids), desc=\"Indexing (parallel)\"):\n                if selected:\n                    self.patch_index.extend(selected)\n\n    def _load_patches_into_cache(self):\n        \"\"\"Loads all indexed patches and their labels/distance maps into self.cached_patches.\"\"\"\n        \n        # Group coordinates by image ID to minimize repeated volume loading\n        images_to_load = {}\n        for entry in self.patch_index:\n            image_id = entry['id']\n            if image_id not in images_to_load:\n                images_to_load[image_id] = []\n            images_to_load[image_id].append(entry)\n\n        # Load each unique image volume once\n        vol_cache = {}\n        lbl_cache = {}\n        \n        # Iterate over unique images that provided patches\n        for image_id in tqdm(images_to_load.keys(), desc=\"Loading volumes\"):\n            # Path determination\n            img_path = os.path.join(self.images_dir, image_id + \".tif\")\n            lbl_path = os.path.join(self.labels_dir, image_id + \".tif\")\n            if not os.path.exists(img_path) and os.path.exists(os.path.join(self.images_dir, image_id + \".tiff\")):\n                img_path = os.path.join(self.images_dir, image_id + \".tiff\")\n            if not os.path.exists(lbl_path) and os.path.exists(os.path.join(self.labels_dir, image_id + \".tiff\")):\n                lbl_path = os.path.join(self.labels_dir, image_id + \".tiff\")\n\n            try:\n                # Load Image Volume\n                vol = _safe_read_tiff(img_path).astype(np.float32)\n                vol_cache[image_id] = vol\n                \n                # Load Label Volume\n                lbl_vol = _safe_read_tiff(lbl_path)\n                lbl_vol = (np.asarray(lbl_vol) > 0).astype(np.uint8)\n                lbl_cache[image_id] = lbl_vol\n                \n            except Exception as e:\n                print(f\"Warning: Failed to load volume {image_id} for caching: {e}\")\n                \n        \n        # Extract, Preprocess, and Cache Patches\n        cached_patches: List[VesuviusPatchDataset.PatchCacheEntry] = []\n        \n        for entry in tqdm(self.patch_index, desc=\"Extracting & Preprocessing Patches\"):\n            image_id = entry[\"id\"]\n            zc = int(entry[\"z\"])\n            y0 = int(entry[\"y\"])\n            x0 = int(entry[\"x\"])\n            \n            if image_id not in vol_cache or image_id not in lbl_cache:\n                continue\n\n            vol = vol_cache[image_id]\n            lbl_vol = lbl_cache[image_id]\n\n            Z, H, W = vol.shape\n            half = self.num_slices // 2\n            \n            # Z range calculation\n            z0 = max(0, zc - half)\n            z1 = z0 + self.num_slices\n            if z1 > Z:\n                z1 = Z\n                z0 = max(0, z1 - self.num_slices)\n\n            # Extract and normalize slices\n            slices = []\n            for z in range(z0, z1):\n                sl = vol[z].astype(np.float32)\n                mx = sl.max() if sl.max() > 0 else 1.0\n                sln = sl / mx\n                patch_sl = sln[y0:y0 + self.patch_size, x0:x0 + self.patch_size]\n                slices.append(patch_sl)\n                \n            # Pad if required\n            while len(slices) < self.num_slices:\n                slices.append(np.zeros((self.patch_size, self.patch_size), dtype=np.float32))\n\n            X = np.stack(slices, axis=0).astype(np.float32) # (C, H, W) where C=num_slices\n            \n            # Extract central label\n            central_z = min(max(0, zc), lbl_vol.shape[0] - 1)\n            Y_true_patch = (lbl_vol[central_z, y0:y0 + self.patch_size, x0:x0 + self.patch_size] > 0).astype(np.float32)\n\n            # Distance transform\n            # This is critical for the Hybrid Topo Loss's L_dist term\n            dt_in = distance_transform_edt(Y_true_patch)\n            dt_out = distance_transform_edt(1 - Y_true_patch)\n            Y_dist_patch = (dt_in - dt_out).astype(np.float32)\n            \n            # Store numpy arrays in cache (Tensor conversion happens in __getitem__)\n            # Store as (X, Y_true, Y_dist)\n            cached_patches.append((X, Y_true_patch, Y_dist_patch))\n\n        self.cached_patches = cached_patches\n        \n\n    def __len__(self) -> int:\n        # Length is now the size of the cache\n        return len(self.cached_patches)\n\n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n        # --- 1. Retrieve Preprocessed Patch from Cache (Instantaneous) ---\n        X_np, Y_true_np, Y_dist_np = self.cached_patches[idx]\n\n        # --- 2. Augmentation (Only rotational/flip needed now) ---\n        if self.augment:\n            # Copy here to ensure augmentation doesn't modify the cached array\n            X = X_np.copy()\n            Y_true_patch = Y_true_np.copy()\n            Y_dist_patch = Y_dist_np.copy()\n            \n            k = random.randint(0, 3)\n            # X is (C, H, W). Rotate H and W axes (1, 2).\n            X = np.rot90(X, k=k, axes=(1, 2)).copy()\n            Y_true_patch = np.rot90(Y_true_patch, k=k).copy()\n            Y_dist_patch = np.rot90(Y_dist_patch, k=k).copy()\n            \n            if random.random() < 0.5:\n                # Flip H and W axes for X (1, 2) and Y (0, 1)\n                X = np.flip(X, axis=2).copy() # flip along W\n                Y_true_patch = np.flip(Y_true_patch, axis=1).copy() # flip along W\n                Y_dist_patch = np.flip(Y_dist_patch, axis=1).copy() # flip along W\n        else:\n            X = X_np\n            Y_true_patch = Y_true_np\n            Y_dist_patch = Y_dist_np\n            \n        # --- 3. Convert to Tensors and Return ---\n        # X is (C, H, W). Ys are (H, W), unsqueeze to (1, H, W).\n        X_tensor = torch.from_numpy(X).float()\n        Y_true_tensor = torch.from_numpy(Y_true_patch).float().unsqueeze(0)\n        Y_dist_tensor = torch.from_numpy(Y_dist_patch).float().unsqueeze(0)\n\n        return X_tensor, Y_true_tensor, Y_dist_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T04:52:45.936473Z","iopub.execute_input":"2025-12-09T04:52:45.937306Z","iopub.status.idle":"2025-12-09T04:52:46.467571Z","shell.execute_reply.started":"2025-12-09T04:52:45.937285Z","shell.execute_reply":"2025-12-09T04:52:46.467004Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Topology-aware pipeline\n***The Optimized Training Pipeline: Data to Topological Loss***\n\nThe entire pipeline is designed to overcome the challenge of large, sparse data (Vesuvius Scrolls) and the computational difficulty of topology checking (Persistent Homology). This is achieved through a hybrid CPU/GPU workflow that localizes the computationally expensive topological checks.\n\n**I. Data Preparation Stage**\nGoal of VesuviusPatchDataset:\n1.  Load large 2D grayscale TIFF images (thousands × thousands of pixels).\n2.  Cut them into overlapping patches (e.g., 384×384) for training.\n3.  **Crucially**, it implements **foreground oversampling** to prioritize patches that contain positive mask pixels (ink) over blank, background-only regions.\n\nReturn batches suitable for PyTorch training:\n* `x`: `(1, P, P)` grayscale patch (normalized to `[0, 1]`)\n* `y`: `(1, P, P)` binary mask (label)\n\nWhat This Dataset Gives You:\n* Balanced sampling of rare positive pixels, preventing the model from learning to predict only background.\n* Large 2D sheets are handled efficiently via patch pre-indexing.\n* The data is perfectly compatible with: BCE + Dice, supervoxel loss, and the new **LocalBettiLoss**.\n\n\n\n**II. The Fast Topology Engine Core (CPU/NumPy)**\nThese helpers form the core of the topological detection engine. They run on the CPU (decoupled from the PyTorch graph) to identify regions where the predicted and ground truth Betti numbers diverge.\n\n#### ✅ **Helper Function: `_get_components_2d(mask)`**\n\nThis low-level function uses `skimage.measure.label` to find connected components. The choice of connectivity is vital for thin structures:\n\n* **Foreground (Ink):** Uses **8-connectivity** (diagonal connections allowed). This is especially important in scroll segmentation because thin fibers or noisy predictions often appear diagonal, and 8-connectivity prevents them from being treated as separate components.\n* **Background (Holes):** Uses **4-connectivity** (orthogonal connections allowed). This helps accurately define the boundaries of holes ($b_1$).\n\n#### ✅ **Helper Function: `compute_betti_numbers_2d(mask)`**\n\nThis function wraps the connectivity check to calculate the fundamental topological invariants:\n\n* $b_0$ = number of foreground components (islands/fragments)\n* $b_1$ = number of holes (loops/voids)\n\nThis is the exact comparison metric used to detect topological errors: if $b_0$ (pred) $\\neq$ $b_0$ (gt) or $b_1$ (pred) $\\neq$ $b_1$ (gt), the topology is wrong in that region.\n\n\n\n#### 🚀 **Optimization Helper: `extract_windows_2d(arr)`**\nThis is a critical performance component that maximizes the efficiency of the CPU check:\n\n* It uses the NumPy **`as_strided`** trick. This allows the function to generate millions of overlapping window views without moving or copying a single byte of memory.\n* It ensures the detection scan over the potentially huge patches is **instantaneous** and zero-copy.\n\n#### ✅ **Helper Function: `find_topology_mismatch_windows(pred_mask, gt_mask)`**\nThis is the orchestrator of the detection stage.\n\n**Goal:** On a (pred, gt) patch pair, slide a window (e.g., 64×64), compute $(b_0, b_1)$ in each window for both `pred` and `gt`, and report the coordinates where they differ.\n\n**Why this localization works:**\n1.  **Scan Cheaply:** The function quickly scans the entire input patch (which can be 384x384 or larger) using the fast, zero-copy window extractor.\n2.  **Find Suspicious Areas:** It identifies only the few windows where topology is suspicious (mismatched $b_0$ or $b_1$).\n3.  **Return Coordinates:** It returns a list of coordinates (`y0`, `x0`, `h`, `w`) that point to the problem areas on the original PyTorch tensor.\n\nThis reduces computational cost by orders of magnitude and is the **key step** in avoiding slow Persistent Homology libraries.\n\n### III. The Differentiable Integration (HybridTopoLossOrchestrator)\nThe `HybridTopoLossOrchestrator` module is the final training orchestrator. It manages a three-part loss strategy, using the CPU-based topological detection engine as a surgical mask to apply two distinct, localized correction terms on the GPU.\n\n#### **HybridTopoLossOrchestrator(yhat, y\\_true)**\n1.  **Base Loss Calculation:** Computes the standard **BCE** or **Dice Loss** over the *entire* input patch (`yhat` vs. `y_true`). This is the main semantic driver.\n    * *Location:* GPU/PyTorch (Differentiable).\n2.  **Topological Detection:** The input and ground truth are moved to the CPU, detached from the gradient graph (`.detach().cpu().numpy()`). The `find_topology_mismatch_windows` is called to get a list of mismatched window coordinates.\n    * *Location:* CPU/NumPy (Non-Differentiable).\n3.  **Topological Term Application:** For every set of coordinates\n4.  **Final Loss Combination:** The total loss is returned as the weighted average of all three components. This elegant separation of concerns allows the total gradient flow to be guided simultaneously by general segmentation quality, boundary precision, and strict topological fidelity, achieving a balance between semantic accuracy and structural correctness.","metadata":{}},{"cell_type":"code","source":"# %%writefile hybrid_topoloss.py\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom typing import List, Tuple, Dict, Any\nfrom numpy.lib.stride_tricks import as_strided\nfrom skimage.measure import label\nfrom concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor\nfrom math import log1p\n\n# =======================================================================\n# I. LOW-LEVEL NUMPY/CPU HELPERS (TOPOLOGY DETECTION CORE)\n# These run non-differentiably on the CPU to quickly locate topological errors.\n# =======================================================================\n\ndef _get_components_2d(mask: np.ndarray, connectivity: int = 2) -> Tuple[np.ndarray, int]:\n    \"\"\"Fast 2D connected components using skimage.\"\"\"\n    # Ensure mask is binary (0 or 1)\n    mask = (mask > 0).astype(np.uint8)\n    labels = label(mask, connectivity=connectivity)\n    num_components = labels.max()\n    return labels, num_components\n\ndef compute_betti_numbers_2d(mask: np.ndarray) -> Tuple[int, int, np.ndarray, np.ndarray]:\n    \"\"\"Compute Betti numbers b0 (components) and b1 (holes).\"\"\"\n    # Foreground components (b0) -> Connectivity=2 (corners count)\n    labels_fg, num_fg = _get_components_2d(mask, connectivity=2)\n    b0 = num_fg\n    \n    # Background components (b1) -> Connectivity=1 (corners don't count)\n    inverted = 1 - mask\n    labels_bg, num_bg = _get_components_2d(inverted, connectivity=1)\n    # The image border is always one component, so b1 = (bg components) - 1\n    b1 = max(num_bg - 1, 0)\n    return b0, b1, labels_fg, labels_bg\n\ndef extract_windows_2d(arr: np.ndarray, window_size: int = 64, stride: int = 32) -> Tuple[np.ndarray, List[Tuple[int, int]]]:\n    \"\"\"Extract sliding windows using the zero-copy NumPy as_strided trick.\"\"\"\n    H_orig, W_orig = arr.shape\n    ws, st = window_size, stride\n\n    # Calculate padding needed to ensure all parts of the image are covered\n    pad_h = (st - (H_orig - ws) % st) % st if H_orig >= ws else (ws - H_orig)\n    pad_w = (st - (W_orig - ws) % st) % st if W_orig >= ws else (ws - W_orig)\n\n    padded_arr = np.pad(arr, ((0, pad_h), (0, pad_w)), mode=\"constant\", constant_values=0)\n    H_pad, W_pad = padded_arr.shape\n\n    out_h = (H_pad - ws) // st + 1\n    out_w = (W_pad - ws) // st + 1\n\n    shape = (out_h, out_w, ws, ws)\n    strides = (padded_arr.strides[0] * st,\n               padded_arr.strides[1] * st,\n               padded_arr.strides[0],\n               padded_arr.strides[1])\n\n    windows_view = as_strided(padded_arr, shape=shape, strides=strides)\n    windows_4d = windows_view.reshape(-1, ws, ws)\n\n    # Generate coordinates for the top-left corner of each window\n    y_coords = np.arange(out_h) * st\n    x_coords = np.arange(out_w) * st\n    # Creates a flat list of (y, x) coordinates\n    coords = [ (y, x) for y in y_coords for x in x_coords ]\n\n    return windows_4d, coords\n\ndef compute_window_betti_and_centroids(window_mask: np.ndarray) -> Dict[str, Any]:\n    \"\"\"Computes b0 and b1 for a single window.\"\"\"\n    b0, b1, _, _ = compute_betti_numbers_2d(window_mask)\n    return {'b0': b0, 'b1': b1, 'fg': [], 'bg': []}\n\ndef match_betti_in_window(gt_win: np.ndarray, pred_win: np.ndarray) -> Dict[str, Any]:\n    \"\"\"Identifies the type of topological mismatch.\"\"\"\n    # Binarize prediction (common threshold for Betti check)\n    p_bin = (pred_win > 0.5).astype(np.uint8)\n    g_bin = (gt_win > 0).astype(np.uint8)\n\n    gt_info = compute_window_betti_and_centroids(g_bin)\n    pred_info = compute_window_betti_and_centroids(p_bin)\n\n    b0g, b1g = gt_info['b0'], gt_info['b1']\n    b0p, b1p = pred_info['b0'], pred_info['b1']\n\n    mismatches = []\n    # If predicted components > actual components, an object was split\n    if b0p > b0g: mismatches.append('split')\n    # If predicted components < actual components, objects were merged\n    if b0p < b0g: mismatches.append('merge')\n    # If predicted holes > actual holes, a hole was spuriously added\n    if b1p > b1g: mismatches.append('hole_added')\n    # If predicted holes < actual holes, a hole was spuriously filled\n    if b1p < b1g: mismatches.append('hole_removed')\n\n    return {'b0_gt': b0g, 'b1_gt': b1g, 'b0_pred': b0p, 'b1_pred': b1p, 'mismatches': mismatches}\n\ndef find_topology_mismatch_windows(\n    pred_mask: np.ndarray, gt_mask: np.ndarray,\n    window_size: int = 64, stride: int = 32, num_workers: int = 1\n) -> Tuple[float, List[Dict[str, Any]]]:\n    \"\"\"Main topology mismatch locator function (CPU/NumPy).\"\"\"\n    assert pred_mask.shape == gt_mask.shape\n\n    # Extract windows using the efficient as_strided view\n    windows_pred, coords = extract_windows_2d(pred_mask, window_size, stride)\n    windows_gt, _ = extract_windows_2d(gt_mask, window_size, stride)\n    total_windows = len(windows_gt)\n    \n    def worker(i: int) -> Dict[str, Any] | None:\n        p = windows_pred[i]\n        g = windows_gt[i]\n\n        # Ignore windows where GT is empty and Prediction is also effectively empty\n        if g.sum() == 0 and p.max() < 0.5:\n             return None\n\n        det = match_betti_in_window(g, p)\n\n        # Only return results if a mismatch was detected\n        if det['mismatches']:\n            y0, x0 = coords[i]\n            h, w = p.shape\n            return {'y0': int(y0), 'x0': int(x0), 'h': int(h), 'w': int(w), 'details': det}\n        \n        return None\n\n    results = []\n    # Use ThreadPoolExecutor for low latency, as the work is mostly I/O (NumPy/skimage)\n    if num_workers > 1:\n        # Using ThreadPoolExecutor is generally safer than ProcessPoolExecutor \n        # inside PyTorch loops due to multiprocessing start methods.\n        with ThreadPoolExecutor(max_workers=num_workers) as ex:\n            results = list(ex.map(worker, range(total_windows)))\n    else:\n        results = [worker(i) for i in range(total_windows)]\n\n    mismatches = [r for r in results if r is not None]\n    \n    # Calculate how many windows contained GT or a positive prediction\n    num_non_empty_windows = sum(1 for i in range(total_windows) if windows_gt[i].sum() > 0 or windows_pred[i].max() > 0.5)\n    # The score is the fraction of non-trivial windows containing a topological mismatch\n    mismatch_score = len(mismatches) / (num_non_empty_windows if num_non_empty_windows > 0 else total_windows)\n    \n    return mismatch_score, mismatches\n\n# =======================================================================\n# II. Differentiable Topological Term (Placeholder for PH/Betti Matching)\n# =======================================================================\n\ndef localized_betti_matching_loss(yhat_patch: torch.Tensor, y_true_patch: torch.Tensor, mismatch_type: List[str]) -> torch.Tensor:\n    \"\"\"\n    Simulated Localized Betti Matching Loss (BML).\n    \n    This uses a strongly scaled BCE proxy to penalize the local area causing \n    the Betti number mismatch, forcing the model to fix the topological error.\n    \"\"\"\n    \n    # 1. Base loss over the problematic window\n    # Ensure y_true_patch is float for BCE\n    base_loss = F.binary_cross_entropy(yhat_patch, y_true_patch.float(), reduction='mean')\n    \n    # 2. Scaling Factor (aggressively scales loss based on number of mismatch types)\n    severity_factor = log1p(len(mismatch_type)) \n    \n    # L_Betti_Match\n    return base_loss * severity_factor\n\n\n# =======================================================================\n# III. PYTORCH MODULE (HYBRID ORCHESTRATOR)\n# =======================================================================\n\nclass HybridTopoLossOrchestrator(torch.nn.Module):\n    \"\"\"\n    Combines Base Loss, Distance-Weighted Correction (L_dist), and \n    Localized Explicit Topological Correction (L_betti_match).\n    \n    L_Total = (1 - alpha - beta)*L_base + alpha*L_dist + beta*L_betti_match\n    \"\"\"\n    def __init__(self, window_size=64, stride=32, alpha=0.15, beta=0.2, base_loss='dice', num_workers=4, boundary_weight_factor=5.0):\n        super().__init__()\n        self.window_size = int(window_size)\n        self.stride = int(stride)\n        self.alpha = float(alpha) # Weight for Distance-Weighted Correction (L_dist)\n        self.beta = float(beta)   # Weight for Localized Betti Matching Term (L_betti_match)\n        \n        # Ensure total weight for corrections is less than 1.0\n        assert (self.alpha + self.beta) < 1.0, \"Alpha + Beta must be less than 1.0\"\n        \n        assert base_loss in ('bce','dice')\n        self.base_loss = base_loss\n        self.num_workers = num_workers\n        self.eps = 1e-7\n        self.boundary_weight_factor = boundary_weight_factor\n        \n    def _calculate_dice_loss(self, p: torch.Tensor, g: torch.Tensor) -> torch.Tensor:\n        \"\"\"Calculates 1 - Dice Score for PyTorch tensors.\"\"\"\n        p_flat, g_flat = p.reshape(-1), g.reshape(-1)\n        inter = (p_flat * g_flat).sum()\n        denom = p_flat.sum() + g_flat.sum()\n        return 1.0 - (2.0 * inter + self.eps) / (denom + self.eps)\n        \n    def _calculate_weighted_bce(self, p: torch.Tensor, g: torch.Tensor, d: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Calculates BCE loss weighted by the distance map (d).\n        This primarily corrects boundaries (Surface Dice).\n        \"\"\"\n        max_d = d.max()\n        if max_d == 0:\n              weights = torch.ones_like(p)\n        else:\n              normalized_d = d / max_d\n              # Exponential weighting for aggressive boundary focus\n              weights = 1.0 + self.boundary_weight_factor * torch.exp(-10 * normalized_d)\n        \n        loss = F.binary_cross_entropy(p, g, weight=weights, reduction='mean')\n        return loss\n        \n    def forward(self, yhat: torch.Tensor, y_true: torch.Tensor, y_dist: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        yhat: (B,1,H,W) probabilities (float) - Prediction\n        y_true: (B,1,H,W) binary - Ground Truth\n        y_dist: (B,1,H,W) distance transform map (float) - Distance Map\n        returns: scalar tensor\n        \"\"\"\n        assert yhat.shape == y_true.shape == y_dist.shape\n        B, C, H, W = yhat.shape\n        device = yhat.device\n\n        # --- 1. Base Loss (L_base) ---\n        if self.base_loss == 'bce':\n            L_base = F.binary_cross_entropy(yhat, y_true.float(), reduction='mean')\n        else:\n            p_flat = yhat.view(B, -1)\n            g_flat = y_true.float().view(B, -1)\n            # Calculate Dice for each sample in the batch and take the mean\n            dice_losses = [self._calculate_dice_loss(p_flat[i], g_flat[i]) for i in range(B)]\n            L_base = torch.stack(dice_losses).mean()\n\n        # --- 2. Topological Detection & Correction (Hybrid CPU/GPU) ---\n        L_distance_weighted_terms = []\n        L_betti_matching_terms = []\n        \n        for i in range(B):\n            # Move data from GPU/PyTorch to CPU/NumPy for fast, non-differentiable topology checks\n            p_np = yhat[i,0].detach().cpu().numpy()\n            g_np = y_true[i,0].detach().cpu().numpy().astype(np.uint8)\n\n            # Find all windows that contain a mismatch in Betti numbers\n            _, mismatches = find_topology_mismatch_windows(\n                p_np, g_np, window_size=self.window_size, stride=self.stride, num_workers=self.num_workers\n            )\n\n            if not mismatches:\n                # If no errors, append zero loss for this sample\n                L_distance_weighted_terms.append(torch.tensor(0.0, device=device))\n                L_betti_matching_terms.append(torch.tensor(0.0, device=device))\n                continue\n\n            # Compute differentiable losses for each mismatched window (back on GPU)\n            window_distance_losses = []\n            window_betti_losses = []\n\n            for w in mismatches:\n                y0, x0, h, w_ = int(w['y0']), int(w['x0']), int(w['h']), int(w['w'])\n                mismatch_types = w['details']['mismatches']\n                \n                # Slice the PyTorch tensors back on the GPU using the coordinates found on the CPU\n                p_patch = yhat[i,0, y0:y0+h, x0:x0+w_]\n                g_patch = y_true[i,0, y0:y0+h, x0:x0+w_].float()\n                d_patch = y_dist[i,0, y0:y0+h, x0:x0+w_].float() \n                \n                # L_distance_weighted (Addresses Surface Dice near the error)\n                dist_loss = self._calculate_weighted_bce(p_patch, g_patch, d_patch)\n                window_distance_losses.append(dist_loss)\n                \n                # L_betti_matching (Addresses TopoScore explicitly)\n                betti_loss = localized_betti_matching_loss(p_patch, g_patch, mismatch_types)\n                window_betti_losses.append(betti_loss)\n            \n            # The loss for the sample is the mean of all mismatched window losses\n            L_distance_weighted_terms.append(torch.stack(window_distance_losses).mean())\n            L_betti_matching_terms.append(torch.stack(window_betti_losses).mean())\n\n        # The final correction terms are the mean across the batch\n        L_distance_weighted_mean = torch.stack(L_distance_weighted_terms).mean()\n        L_betti_matching_mean = torch.stack(L_betti_matching_terms).mean()\n        \n        # --- 3. Combined Final Loss ---\n        L_total = (1.0 - self.alpha - self.beta) * L_base + \\\n                  self.alpha * L_distance_weighted_mean + \\\n                  self.beta * L_betti_matching_mean\n                  \n        return L_total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T04:52:46.468395Z","iopub.execute_input":"2025-12-09T04:52:46.468697Z","iopub.status.idle":"2025-12-09T04:52:46.511402Z","shell.execute_reply.started":"2025-12-09T04:52:46.468679Z","shell.execute_reply":"2025-12-09T04:52:46.510926Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main training\n***Main Training Orchestration***\n\n**Final Phase Analysis: Inputs and Outputs.**\nThis script is the final execution phase of the entire pipeline.\n\n**Core Components (Inputs)**\n1. The `train_main.py` script requires three previously defined custom modules to be available in the environment:`VesuviusPatchDataset` (Data): Provides the training samples.\n   * Input Data: Multichannel X-ray CT volume slices (X).\n   * Target Data: Binary ground truth ink mask (Y_true).\n   * Auxiliary Data: Euclidean Distance Transform map (Y_dist).\n     \n2. `FullUNet` (Model): The neural network architecture (in this case, a placeholder UNet).\n   * Input: Multi-layer X-ray patch (X).\n   * Output: Predicted ink probability map(Y_hat).\n\n3. `HybridTopoLossOrchestrator` (Loss Function): The combined loss mechanism.\n    * Inputs: Y_hat, Y_true, and Y_dist.\n    * Output: Scalar Tensor (the final weighted loss value, $L_{Total}$).\n\n\n**Outputs of the Training Loop**\nThe `main()` function's primary outputs are:\n1. **Model Checkpoint File (.pth):** The weights of the best performing model saved to the disk (best_vesuvius_topo_model.pth). This is essential for subsequent inference.\n2. **Performance Metrics:** Prints the average training loss per epoch and the total time taken.\n3. **Updated Model Weights:** The `model` object itself is updated in-place via `optimizer.step()`, carrying the improved weights ready for the next epoch.\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport time\nimport os\nimport sys\n\n# Notebook-Safe Path Setup \ncurrent_dir = os.getcwd()\nif current_dir not in sys.path:\n    sys.path.append(current_dir)\n\n# from vesuvius_dataset import VesuviusPatchDataset \n# from model_definition import FullUNet \n# from hybrid_topoloss import HybridTopoLossOrchestrator # Renamed import to match common conventions\n\n# Set the device to GPU if available, otherwise CPU\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {DEVICE}\")\n\n# Hyperparameters\nBATCH_SIZE = 4\nPATCH_SIZE = 384 # Standard patch size for Vesuvius Challenge\nNUM_EPOCHS = 10 \nLEARNING_RATE = 1e-4\nROOT_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection\" \nSAVE_PATH = \"best_vesuvius_topo_model.pth\"\n\n# --- Training Function ---\ndef train_one_epoch(\n    model: torch.nn.Module, \n    data_loader: DataLoader, \n    loss_fn: HybridTopoLossOrchestrator, \n    optimizer: optim.Adam, \n    epoch: int\n) -> float:\n    \"\"\"Trains the model for one epoch.\"\"\"\n    \n    model.train()\n    total_loss = 0.0\n    \n    # Use ThreadPoolExecutor within DataLoader for stable multiprocessing on most systems\n    progress_bar = tqdm(data_loader, desc=f\"Epoch {epoch}\", unit=\"batch\")\n    \n    for batch_idx, (X, Y_true, Y_dist) in enumerate(progress_bar):\n        # 1. Move data to the selected device (GPU/CPU)\n        X = X.to(DEVICE)        # Input Image (Multi-layer X-ray)\n        Y_true = Y_true.to(DEVICE) # Binary Ground Truth Ink Mask\n        Y_dist = Y_dist.to(DEVICE) # Euclidean Distance Transform (for L_dist term)\n\n        # 2. Zero the gradients before the forward pass\n        optimizer.zero_grad()\n\n        # 3. Forward Pass: Get the prediction (Y_hat, probabilities)\n        Y_hat_logits = model(X)\n        Y_hat = torch.sigmoid(Y_hat_logits)\n\n        # 4. Compute Loss: L_Total = (1-alpha-beta)*L_base + alpha*L_dist + beta*L_betti_match\n        loss = loss_fn(Y_hat, Y_true, Y_dist)\n\n        # 5. Backward Pass: Compute gradients\n        loss.backward()\n\n        # 6. Update Weights\n        optimizer.step()\n\n        total_loss += loss.item()\n        avg_loss = total_loss / (batch_idx + 1)\n        \n        progress_bar.set_postfix({\"Loss\": f\"{avg_loss:.4f}\"})\n\n    return total_loss / len(data_loader)\n\n# --- Main Orchestration Function ---\ndef main():\n    \"\"\"Sets up the pipeline and runs the training loop.\"\"\"\n    print(\"--- Pipeline Setup ---\")\n\n    # 1. Data Loading and Preparation\n    try:\n        train_dataset = VesuviusPatchDataset(\n            root_dir=ROOT_DIR, \n            patch_size=PATCH_SIZE, \n            max_patches_per_image=100, # Load 100 patches per image\n            augment=True\n        )\n        # Use half the available CPU cores for data loading for efficiency\n        num_workers = os.cpu_count() // 2 if os.cpu_count() else 2\n        train_loader = DataLoader(\n            train_dataset, \n            batch_size=BATCH_SIZE, \n            shuffle=True, \n            num_workers=num_workers,\n            pin_memory=True # Recommended for GPU training\n        )\n        print(f\"Total training patches indexed: {len(train_dataset)}. Batches: {len(train_loader)}\")\n    except Exception as e:\n        print(f\"Error during Data Loading. Check ROOT_DIR ({ROOT_DIR}) and VesuviusPatchDataset implementation. Error: {e}\")\n        return\n\n    # 2. Model Initialization\n    print(\"Initializing UNet model...\")\n    \n    NUM_SLICES = 10       # <--- set this to match your dataset\n    ADD_EDGE = False      # <--- unless your dataset adds Sobel edges\n    \n    in_channels = NUM_SLICES + (1 if ADD_EDGE else 0)\n    \n    model = FullUNet(\n        n_channels=in_channels,\n        n_classes=1\n    ).to(DEVICE)\n    \n    print(f\"Model initialized: {model.__class__.__name__} \"\n          f\"with {sum(p.numel() for p in model.parameters() if p.requires_grad):,} parameters.\")\n\n\n    # 3. Loss Function and Optimizer Initialization\n    # Setting alpha=0.35, beta=0.35 leaves 0.30 weight for L_base\n    loss_fn = HybridTopoLossOrchestrator(\n        alpha=0.35, beta=0.35, base_loss='dice'\n    ).to(DEVICE)\n    \n    optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\n    print(f\"Loss Function: Hybrid Topo Loss (alpha=0.35, beta=0.35)\")\n\n    print(\"--- Starting Training Loop ---\")\n    best_loss = float('inf')\n    start_time = time.time()\n\n    # 4. Training Loop\n    for epoch in range(1, NUM_EPOCHS + 1):\n        avg_epoch_loss = train_one_epoch(\n            model, train_loader, loss_fn, optimizer, epoch\n        )\n\n        print(f\"\\n[EPOCH {epoch}/{NUM_EPOCHS}] Finished. Avg Loss: {avg_epoch_loss:.4f}\")\n\n        # Save the best model state dictionary\n        if avg_epoch_loss < best_loss:\n            best_loss = avg_epoch_loss\n            torch.save(model.state_dict(), SAVE_PATH)\n            print(f\"Model checkpoint saved to {SAVE_PATH} (New best loss: {best_loss:.4f})\")\n            \n    end_time = time.time()\n    print(f\"\\nTraining complete in {(end_time - start_time) / 60:.2f} minutes. Final best loss: {best_loss:.4f}\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T04:52:46.512906Z","iopub.execute_input":"2025-12-09T04:52:46.513101Z","iopub.status.idle":"2025-12-09T04:58:42.103289Z","shell.execute_reply.started":"2025-12-09T04:52:46.513086Z","shell.execute_reply":"2025-12-09T04:58:42.102632Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission file\n***Inference and Submission Generation Script.***\n\nSince the training script is now complete, we need a separate file to handle:\n\n1. Loading the test data and the trained model weights. One must first initialize your model and load the best weights saved during training.\n    * **Load the Model:** Initialize FullNet and load the weights from best_vesuvius_topo_model.pth.\n    * **Load Test Data:** Use a test-specific version of the VesuviusPatchDataset to iterate through the test scroll(s) and extract patches.\n      \n2. Performing sliding window inference across the entire large test volume. The model needs to generate probability predictions ($\\hat{Y}$) for every pixel in the entire test scroll.\n\n   * **Sliding Window Inference:** Since the model was trained on small patches (384x384), you must iterate through the full test scroll using a sliding window approach (like a $384 \\times 384$ window with a small stride, e.g., 64).\n   * **Averaging Overlaps:** Due to the sliding window, most pixels will be predicted multiple times. The predictions for these overlapping regions must be averaged to create a smooth, final probability map for the entire volume.\n     \n3. Stitching the resulting predictions into a final probability volume. After all patches are predicted, they must be stitched back together into a single, high-resolution probability map.\n    * **Stitching:** The averaged patch predictions are recombined to form a single, 3D (H, W, D) probability volume representing the likelihood of ink at every pixel.\n\n4. Applying a threshold to generate the binary mask.\n    * **Thresholding (The Critical Step** The competition requires a binary mask (0s and 1s), not probabilities. You must apply a final threshold to the probability map.\n      \n5. Saving the mask as the required .tif file. Once the final binary mask ($M$) for a test scroll is available, package it for submission.\n\n   * **Format:** The mask must be saved as a single-channel, unsigned 8-bit integer (uint8) or similar data type, matching the dimensions of the original source image, in the TIF format.\n\n   * **Naming:** The file must be named exactly [image_id].tif (e.g., 2.tif).\n\n   * **Zipping:** All generated .tif files for the test set must be compressed into a single .zip file for upload to Kaggle.\n\nThis script bellow assumes the existence of the trained model file (best_vesuvius_topo_model.pth) and the required data classes.","metadata":{}},{"cell_type":"code","source":"# %%writefile make_submission_full_volume.py\nimport os\nimport sys\nimport time\nimport zipfile\nimport traceback\nfrom typing import List\n\nimport numpy as np\nimport torch\nfrom tqdm import tqdm\nimport pandas as pd\nfrom tifffile import imwrite\n\n# Replace with your safe tiff reader + model definition\n#from vesuvius_dataset import _safe_read_tiff   # must return ndarray (D,H,W)\n#from model_definition import FullNet        # model class used during training\n\n# -------------------------\n# Configuration (edit if needed)\n# -------------------------\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n\nROOT_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection\"\nTEST_SUBDIR = \"test_images\"\nOUTPUT_DIR = \"submission\"           # will contain <image_id>.tif (multi-page)\nZIPPATH = \"/kaggle/working/submission.zip\"\nMODEL_PATH = \"best_vesuvius_topo_model.pth\"   # path to .pth\nNUM_SLICES = 10                 # number of input slices the model expects (depth)\nPATCH_SIZE = 384\nSTRIDE = 64\nGLOBAL_THRESHOLD = 0.45         # final binarization threshold for probs -> 0/1\n\n# -------------------------\n# Helpers\n# -------------------------\ndef extract_center_slices(volume: np.ndarray, center_z: int, num_slices: int) -> np.ndarray:\n    \"\"\"\n    Return a stack of `num_slices` slices centered on center_z.\n    If near borders, pad with zeros slices to reach num_slices.\n    Result shape (C, H, W).\n    \"\"\"\n    D, H, W = volume.shape\n    half = num_slices // 2\n    z0 = center_z - half\n    z1 = z0 + num_slices\n\n    pad_before = 0\n    pad_after = 0\n    if z0 < 0:\n        pad_before = -z0\n        z0 = 0\n    if z1 > D:\n        pad_after = z1 - D\n        z1 = D\n\n    slices = []\n    if z0 < z1:\n        slices.append(volume[z0:z1])\n        slices = np.concatenate(slices, axis=0) if len(slices) > 1 else slices[0]\n    else:\n        slices = np.empty((0, H, W), dtype=volume.dtype)\n\n    # pad before / after as needed\n    if pad_before > 0:\n        pad_shape = (pad_before, H, W)\n        slices = np.concatenate([np.zeros(pad_shape, dtype=volume.dtype), slices], axis=0)\n    if pad_after > 0:\n        pad_shape = (pad_after, H, W)\n        slices = np.concatenate([slices, np.zeros(pad_shape, dtype=volume.dtype)], axis=0)\n\n    # if still shorter (shouldn't) pad\n    if slices.shape[0] < num_slices:\n        need = num_slices - slices.shape[0]\n        slices = np.concatenate([slices, np.zeros((need, H, W), dtype=volume.dtype)], axis=0)\n\n    # if longer, slice to exact\n    slices = slices[:num_slices]\n\n    # normalize each slice to [0,1] safe\n    out = []\n    for sl in slices:\n        mx = float(sl.max()) if sl.max() > 0 else 1.0\n        out.append((sl.astype(np.float32) / mx))\n    return np.stack(out, axis=0)  # (C, H, W)\n\n\ndef sliding_predict_on_stack(model, stack_c_hw: np.ndarray, patch_size: int, stride: int, device: torch.device):\n    \"\"\"\n    Performs sliding-window inference for a stack representing CxHxW but will\n    return a single HxW probability map for the *center* slice.\n    stack_c_hw: (C, H, W)\n    \"\"\"\n    C, H, W = stack_c_hw.shape\n\n    # pad if smaller than patch size\n    pad_y = max(0, patch_size - H)\n    pad_x = max(0, patch_size - W)\n    pad_yb = pad_y // 2\n    pad_ya = pad_y - pad_yb\n    pad_xb = pad_x // 2\n    pad_xa = pad_x - pad_xb\n\n    padded = np.pad(stack_c_hw, ((0,0),(pad_yb,pad_ya),(pad_xb,pad_xa)), mode=\"constant\", constant_values=0)\n    _, Hp, Wp = padded.shape\n\n    prediction_sum = np.zeros((Hp, Wp), dtype=np.float32)\n    prediction_count = np.zeros((Hp, Wp), dtype=np.int32)\n\n    max_y = Hp - patch_size\n    max_x = Wp - patch_size\n\n    # coords\n    y_coords = list(range(0, max_y + 1, stride)) if max_y >= 0 else [0]\n    x_coords = list(range(0, max_x + 1, stride)) if max_x >= 0 else [0]\n    if y_coords and y_coords[-1] < max_y:\n        y_coords.append(max_y)\n    if x_coords and x_coords[-1] < max_x:\n        x_coords.append(max_x)\n\n    # ensure there's at least one coordinate\n    if not y_coords:\n        y_coords = [0]\n    if not x_coords:\n        x_coords = [0]\n\n    # inference\n    model.eval()\n    with torch.no_grad():\n        for y0 in y_coords:\n            for x0 in x_coords:\n                patch = padded[:, y0:y0+patch_size, x0:x0+patch_size]\n                if patch.shape[1] != patch_size or patch.shape[2] != patch_size:\n                    # skip incomplete patches (shouldn't happen because we added final coords)\n                    continue\n                # to tensor: (1, C, P, P)\n                t = torch.from_numpy(patch).float().unsqueeze(0).to(device)\n                try:\n                    out = model(t)   # model should output (B,1,P,P) or (B,P,P)\n                except RuntimeError as re:\n                    # GPU OOM or similar: try CPU fallback\n                    print(\"[WARN] RuntimeError during forward, retrying on CPU:\", re)\n                    out = model.to(\"cpu\")(t.to(\"cpu\"))\n                # handle shapes\n                out_np = out.squeeze().detach().cpu().numpy()\n                # if out is (P,P) keep, if (1,P,P) squeeze to (P,P)\n                if out_np.ndim == 3:\n                    out_np = out_np.squeeze(0)\n                # aggregate\n                prediction_sum[y0:y0+patch_size, x0:x0+patch_size] += out_np\n                prediction_count[y0:y0+patch_size, x0:x0+patch_size] += 1\n\n    # avoid divide by zero\n    prediction_count[prediction_count == 0] = 1\n    final_padded = prediction_sum / prediction_count\n\n    # crop back to original H,W\n    result = final_padded[pad_yb:pad_yb+H, pad_xb:pad_xb+W]\n    return result  # H x W probability map\n\n\n# -------------------------\n# Top-level processing for one test volume\n# -------------------------\ndef predict_full_volume(model, full_vol: np.ndarray, num_slices: int, patch_size: int, stride: int, device: torch.device) -> np.ndarray:\n    \"\"\"\n    full_vol: (D, H, W)\n    returns pred_volume: (D, H, W) (float32 probabilities)\n    \"\"\"\n    D, H, W = full_vol.shape\n    pred_vol = np.zeros((D, H, W), dtype=np.float32)\n\n    for z in range(D):\n        try:\n            stack = extract_center_slices(full_vol, z, num_slices)  # (C,H,W)\n            prob_map = sliding_predict_on_stack(model, stack, patch_size, stride, device)  # (H,W)\n            pred_vol[z] = prob_map\n        except Exception as e:\n            print(f\"[ERROR] failed predicting slice z={z}: {e}\")\n            traceback.print_exc()\n            # fallback: zeros\n            pred_vol[z] = np.zeros((H, W), dtype=np.float32)\n\n    return pred_vol\n\n\n# -------------------------\n# Save & ZIP\n# -------------------------\ndef save_volume_tif(volume: np.ndarray, image_id: str, out_dir: str):\n    \"\"\"\n    Save volume as a multi-page tiff where pages are slices (z axis first).\n    volume shape: (D, H, W) with values 0/1 or 0..1 floats.\n    We'll save as uint8 (0/1) using zlib compression.\n    \"\"\"\n    os.makedirs(out_dir, exist_ok=True)\n    out_path = os.path.join(out_dir, f\"{image_id}.tif\")\n    try:\n        # Ensure bytes: cast to uint8 0/1\n        vol_u8 = (volume > 0.5).astype(np.uint8)\n        # tifffile can write the 3D array as a multi-page tiff\n        \n        imwrite(out_path, vol_u8, compression='zlib', compressionargs={'level': 6})\n        print(f\"[OK] Wrote {out_path} (shape {vol_u8.shape})\")\n    except Exception as e:\n        print(f\"[ERROR] Failed to write {out_path}: {e}\")\n        traceback.print_exc()\n\n\ndef make_submission_zip(output_dir: str, zip_path: str):\n    try:\n        with zipfile.ZipFile(zip_path, \"w\", compression=zipfile.ZIP_DEFLATED) as zf:\n            for root, _, files in os.walk(output_dir):\n                for fn in files:\n                    if fn.lower().endswith(\".tif\") or fn.lower().endswith(\".tiff\"):\n                        full = os.path.join(root, fn)\n                        arcname = os.path.relpath(full, output_dir)  # keep filenames inside zip\n                        zf.write(full, arcname=arcname)\n        print(f\"[OK] Created zip at {zip_path}\")\n    except Exception as e:\n        print(\"[ERROR] Failed to create zip:\", e)\n        traceback.print_exc()\n\n\n# -------------------------\n# Main\n# -------------------------\ndef main(test_ids: List[str] = None):\n    # 1) load model\n    try:\n        NUM_SLICES = 10       # <--- set this to match your dataset\n        ADD_EDGE = False      # <--- unless your dataset adds Sobel edges\n        \n        in_channels = NUM_SLICES + (1 if ADD_EDGE else 0)\n        \n        model = FullUNet(\n            n_channels=in_channels,\n            n_classes=1\n        ).to(DEVICE)\n        model.load_state_dict(torch.load(MODEL_PATH, map_location=DEVICE))\n        model.eval()\n        print(\"[OK] Model loaded:\", MODEL_PATH)\n    except Exception as e:\n        print(\"[ERROR] Could not load model:\", e)\n        traceback.print_exc()\n        return\n\n    # 2) discover test ids if not provided\n    test_dir = os.path.join(ROOT_DIR, TEST_SUBDIR)\n    if test_ids is None:\n        # test files are like <id>.tif\n        test_files = sorted([fn for fn in os.listdir(test_dir) if fn.lower().endswith((\".tif\", \".tiff\"))])\n        test_ids = [os.path.splitext(fn)[0] for fn in test_files]\n\n    if len(test_ids) == 0:\n        print(\"[ERROR] No test IDs found in\", test_dir)\n        return\n\n    os.makedirs(OUTPUT_DIR, exist_ok=True)\n\n    # 3) for each test volume: load, predict full D x H x W volume, save as multi-page tif\n    for image_id in test_ids:\n        print(f\"\\n--- Processing {image_id} ---\")\n        try:\n            vol_path = os.path.join(test_dir, f\"{image_id}.tif\")\n            vol = _safe_read_tiff(vol_path)\n            vol = np.asarray(vol).astype(np.float32)\n        except Exception as e:\n            print(f\"[ERROR] Failed to load test volume {image_id}: {e}\")\n            traceback.print_exc()\n            continue\n\n        if vol.ndim != 3:\n            print(f\"[ERROR] Unexpected volume dims for {image_id}: {vol.shape} (expected D,H,W). Skipping.\")\n            continue\n\n        try:\n            pred_vol = predict_full_volume(model, vol, NUM_SLICES, PATCH_SIZE, STRIDE, DEVICE)\n            # Binarize\n            bin_vol = (pred_vol > GLOBAL_THRESHOLD).astype(np.uint8)\n            save_volume_tif(bin_vol, image_id, OUTPUT_DIR)\n        except Exception as e:\n            print(f\"[ERROR] Prediction pipeline failed for {image_id}: {e}\")\n            traceback.print_exc()\n            continue\n\n    # 4) zip only tiffs.\n    make_submission_zip(OUTPUT_DIR, ZIPPATH)\n\n\nif __name__ == \"__main__\":\n    # Optionally pass specific test ids: main([\"1407735\"])\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T05:52:43.970009Z","iopub.execute_input":"2025-12-09T05:52:43.970846Z","iopub.status.idle":"2025-12-09T05:53:06.245653Z","shell.execute_reply.started":"2025-12-09T05:52:43.970822Z","shell.execute_reply":"2025-12-09T05:53:06.244945Z"}},"outputs":[],"execution_count":null}]}