{"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":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================================================\n# CELL 1: Installation and Imports\n# =============================================================================\n\n!pip install --no-index --find-links=\"/kaggle/input/surface-detect-package-scraper\" --no-deps -q pytorch_lightning monai albumentations imagecodecs\n!pip uninstall -q -y tensorflow  # preventing AttributeError\n!pip install imagecodecs\n!pip install monai\n!pip install albumentations\n!pip uninstall -y numpy scipy torchmetrics pytorch-lightning\n!pip install numpy==1.26.4 scipy==1.15.3\n!pip install torchmetrics pytorch-lightning\nimport os\nimport imagecodecs\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\nimport tifffile\nimport albumentations as A\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom typing import Tuple, List, Literal, Optional, Dict\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.auto import tqdm\nfrom scipy import ndimage\nfrom scipy.stats import entropy\nfrom skimage.morphology import remove_small_objects, remove_small_holes, ball\nfrom skimage.segmentation import mark_boundaries, watershed\nfrom skimage.feature import peak_local_max\nfrom functools import lru_cache, partial\nimport multiprocessing\nimport warnings\nimport re\nimport zipfile\n\nwarnings.simplefilter(action='ignore', category=FutureWarning)\n\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"CUDA device: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:14:53.067024Z","iopub.execute_input":"2026-01-08T15:14:53.067196Z","iopub.status.idle":"2026-01-08T15:17:10.841051Z","shell.execute_reply.started":"2026-01-08T15:14:53.067172Z","shell.execute_reply":"2026-01-08T15:17:10.840318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 2: Configuration\n# =============================================================================\n\nCACHE_DIR = Path(\"/kaggle/tmp/dataset_cache\")  # Cache for preprocessed slices\n# Data paths\nDATA_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nTRAIN_IMAGES_DIR = DATA_DIR / \"train_images\"\nTRAIN_LABELS_DIR = DATA_DIR / \"train_labels\"\nTEST_IMAGES_DIR = DATA_DIR / \"test_images\"\nOUTPUT_DIR = Path(\".\")\nCHECKPOINT_DIR = \"/kaggle/input/vesuvius-surface-detection-2d-checkpoints/pytorch/default/5\"\n\n# Dataset configuration\nSLICE_AXIS: Literal[\"x\", \"y\", \"z\"] = \"z\"\nIMAGE_SIZE = (320, 320)\nUSE_CACHE = True\n\n# Training configuration\nBATCH_SIZE = 28\nNUM_EPOCHS = 3\nLEARNING_RATE = 1e-3\nVAL_SPLIT = 0.1\nSEED = 42\n\n# Loss weights (optimized for competition metrics)\nLAMBDA_DICE_CE = 1.0\nLAMBDA_SURFACE = 0.4\nLAMBDA_BOUNDARY = 0.2\nLAMBDA_CONNECTIVITY = 0.3\n\n# Post-processing configuration\nMIN_COMPONENT_SIZE = 25 * 25 * 25\nMAX_HOLE_SIZE = 500\nCLOSING_RADIUS = 3\nSMOOTHING_SIGMA = 1.0\nTHIN_BRIDGE_THRESHOLD = 5\n\n# Device\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Set seed\npl.seed_everything(SEED)\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"Configuration loaded successfully!\")\nprint(f\"Device: {DEVICE}\")\nprint(f\"Batch size: {BATCH_SIZE}\")\nprint(f\"Learning rate: {LEARNING_RATE}\")\nprint(f\"Number of epochs: {NUM_EPOCHS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:17:10.842602Z","iopub.execute_input":"2026-01-08T15:17:10.843008Z","iopub.status.idle":"2026-01-08T15:17:10.855175Z","shell.execute_reply.started":"2026-01-08T15:17:10.842983Z","shell.execute_reply":"2026-01-08T15:17:10.854479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 3: Dataset Class\n# =============================================================================\n\nclass VesuviusSliceDataset(Dataset):\n    \"\"\"Dataset for slicing 3D volumes into 2D segmentation samples.\"\"\"\n\n    def __init__(\n        self,\n        images_dir: Path,\n        labels_dir: Optional[Path],\n        volume_files: List[str],\n        slice_axis: Literal[\"x\", \"y\", \"z\"] = \"z\",\n        cache_dir: Optional[Path] = None,\n        use_cache: bool = True,\n        transform: Optional[A.Compose] = None,\n    ):\n        super().__init__()\n        self.images_dir = images_dir\n        self.labels_dir = labels_dir\n        self.volume_files = volume_files\n        self.slice_axis = slice_axis\n        self.cache_dir = cache_dir\n        self.use_cache = use_cache\n        self.transform = transform\n\n        if self.cache_dir is not None:\n            self.cache_dir.mkdir(parents=True, exist_ok=True)\n\n        print(\"Building slice index...\")\n        self.slice_index = self._build_slice_index(\n            self.images_dir,\n            tuple(self.volume_files),\n            self.slice_axis,\n            self.cache_dir\n        )\n        print(f\"Dataset initialized with {len(self.slice_index)} slices.\")\n\n    @staticmethod\n    def _map_axis_name_to_index(slice_axis: str) -> int:\n        axis_map = {\"z\": 0, \"y\": 1, \"x\": 2}\n        return axis_map[slice_axis]\n\n    @staticmethod\n    @lru_cache(maxsize=9)\n    def _build_slice_index(\n        images_dir: Path,\n        volume_files: Tuple[str, ...],\n        slice_axis: Literal[\"x\", \"y\", \"z\"],\n        cache_dir: Optional[Path] = None,\n    ) -> List[Tuple[str, str, int]]:\n        \"\"\"Build slice index for all volumes.\"\"\"\n        slice_index = []\n\n        for volume_file in tqdm(list(volume_files), desc=\"Building slice index\"):\n            num_slices = 0\n\n            if cache_dir is not None:\n                volume_stem = Path(volume_file).stem\n                marker_dir = cache_dir / volume_stem\n                marker_path = marker_dir / f\"axis_{slice_axis}.done\"\n                try:\n                    num_slices = int(marker_path.read_text().strip())\n                except Exception:\n                    pass\n                \n            if not num_slices:\n                volume_tiff = tifffile.imread(str(images_dir / volume_file))\n                if len(volume_tiff.shape) == 2:\n                    num_slices = 1\n                else:\n                    axis_idx_tiff = VesuviusSliceDataset._map_axis_name_to_index(slice_axis)\n                    num_slices = volume_tiff.shape[axis_idx_tiff]\n\n            slice_index += [(volume_file, slice_axis, i) for i in range(num_slices)]\n\n        return slice_index\n\n    @staticmethod\n    def _extract_slice_from_volume(\n        volume: np.ndarray,\n        slice_idx: int,\n        axis: Literal[\"x\", \"y\", \"z\"]\n    ) -> np.ndarray:\n        \"\"\"Extract slice from 3D volume.\"\"\"\n        if len(volume.shape) == 2:\n            return volume[np.newaxis, ...]\n\n        if axis == \"z\":\n            slice_2d = volume[slice_idx, :, :]\n        elif axis == \"y\":\n            slice_2d = volume[:, slice_idx, :]\n        else:\n            slice_2d = volume[:, :, slice_idx]\n\n        return slice_2d[np.newaxis, ...]\n\n    @staticmethod\n    def _get_cache_path(\n        cache_dir: Path, volume_file: str, slice_axis: str, slice_idx: int\n    ) -> Path:\n        \"\"\"Get path to cached slice file.\"\"\"\n        volume_stem = Path(volume_file).stem\n        filename = f\"{slice_axis}_slice-{slice_idx:04d}.npz\"\n        return cache_dir / volume_stem / filename\n\n    @staticmethod\n    def _get_image_and_mask_from_volume(\n        image_volume: np.ndarray,\n        label_volume: Optional[np.ndarray],\n        slice_idx: int,\n        slice_axis: str,\n    ) -> Tuple[np.ndarray, np.ndarray]:\n        \"\"\"Extract and return image and mask slices.\"\"\"\n        image_slice = VesuviusSliceDataset._extract_slice_from_volume(\n            image_volume, slice_idx, slice_axis)\n\n        if label_volume is not None:\n            label_slice = VesuviusSliceDataset._extract_slice_from_volume(\n                label_volume, slice_idx, slice_axis)\n        else:\n            label_slice = np.zeros(image_slice.shape, dtype=np.uint8)\n\n        image_slice = image_slice.astype(np.uint8)\n        label_slice = label_slice.astype(np.uint8)\n        \n        return image_slice, label_slice\n\n    def _load_from_raw(\n        self, volume_file: str, slice_axis: str, slice_idx: int,\n    ) -> Tuple[np.ndarray, np.ndarray]:\n        \"\"\"Load image and mask from raw TIFF files.\"\"\"\n        image_path = self.images_dir / volume_file\n        image_volume = tifffile.imread(str(image_path))\n\n        label_volume = None\n        if self.labels_dir:\n            label_path = self.labels_dir / volume_file\n            if label_path.exists():\n                label_volume = tifffile.imread(str(label_path))\n\n        image, label_slice = VesuviusSliceDataset._get_image_and_mask_from_volume(\n            image_volume, label_volume, slice_idx, slice_axis)\n        return image, label_slice\n\n    @staticmethod\n    def _save_cache(\n        cache_dir: Path,\n        volume_file: str,\n        slice_axis: str,\n        slice_idx: int,\n        image_data: np.ndarray,\n        mask_data: np.ndarray\n    ):\n        \"\"\"Save image and mask data to cache.\"\"\"\n        cache_path = VesuviusSliceDataset._get_cache_path(\n            cache_dir, volume_file, slice_axis, slice_idx)\n        cache_path.parent.mkdir(parents=True, exist_ok=True)\n        np.savez_compressed(str(cache_path), image=image_data.astype(np.uint8), \n                           mask=mask_data.astype(np.uint8))\n\n    @staticmethod\n    def _load_cache(\n        cache_dir: Path, volume_file: str, slice_axis: str, slice_idx: int,\n    ) -> Tuple[np.ndarray, np.ndarray]:\n        \"\"\"Load image and mask data from cache.\"\"\"\n        cache_path = VesuviusSliceDataset._get_cache_path(\n            cache_dir, volume_file, slice_axis, slice_idx)\n        data = np.load(str(cache_path), allow_pickle=True)\n        return data['image'].astype(np.uint8), data['mask'].astype(np.uint8)\n\n    def __len__(self) -> int:\n        return len(self.slice_index)\n\n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        volume_file, slice_axis, slice_idx = self.slice_index[idx]\n\n        image = None\n        label_slice = None\n\n        if self.use_cache and self.cache_dir is not None:\n            cache_path = VesuviusSliceDataset._get_cache_path(\n                self.cache_dir, volume_file, slice_axis, slice_idx)\n\n            if cache_path.exists():\n                image, label_slice = VesuviusSliceDataset._load_cache(\n                    self.cache_dir, volume_file, slice_axis, slice_idx)\n            else:\n                image, label_slice = self._load_from_raw(\n                    volume_file, slice_axis, slice_idx)\n                VesuviusSliceDataset._save_cache(\n                    self.cache_dir, volume_file, slice_axis, slice_idx, image, label_slice)\n        else:\n            image, label_slice = self._load_from_raw(\n                volume_file, slice_axis, slice_idx)\n\n        # Normalize to 0-1\n        image = image.astype(np.float32) / 255.0\n        image = np.transpose(image, (1, 2, 0))\n        label_slice = label_slice.squeeze(0)\n\n        if self.transform is not None:\n            transformed = self.transform(image=image, mask=label_slice)\n            image = transformed[\"image\"]\n            label_slice = transformed[\"mask\"]\n        else:\n            image = torch.from_numpy(image).permute(2, 0, 1)\n            label_slice = torch.from_numpy(label_slice).unsqueeze(0)\n\n        return image.float(), label_slice.float()\n\n    @staticmethod\n    def _process_volume_for_cache(\n        volume_file: str,\n        images_dir: Path,\n        labels_dir: Optional[Path],\n        slice_axis: str,\n        cache_dir: Path\n    ) -> List[Tuple[str, str, int]]:\n        \"\"\"Process a single volume for cache creation.\"\"\"\n        volume_stem = Path(volume_file).stem\n        marker_dir = cache_dir / volume_stem\n        marker_path = marker_dir / f\"axis_{slice_axis}.done\"\n\n        if marker_path.exists():\n            return []\n\n        image_path = images_dir / volume_file\n        image_volume = tifffile.imread(str(image_path))\n\n        label_volume = None\n        if labels_dir is not None:\n            label_path = labels_dir / volume_file\n            if label_path.exists():\n                label_volume = tifffile.imread(str(label_path))\n\n        if len(image_volume.shape) == 2:\n            num_slices = 1\n        else:\n            axis_idx = VesuviusSliceDataset._map_axis_name_to_index(slice_axis)\n            num_slices = image_volume.shape[axis_idx]\n\n        slice_entries = []\n        for slice_idx in range(num_slices):\n            image_slice, label_slice = VesuviusSliceDataset._get_image_and_mask_from_volume(\n                image_volume, label_volume, slice_idx, slice_axis)\n\n            VesuviusSliceDataset._save_cache(\n                cache_dir, volume_file, slice_axis, slice_idx, image_slice, label_slice)\n            slice_entries.append((volume_file, slice_axis, slice_idx))\n\n        marker_dir.mkdir(parents=True, exist_ok=True)\n        marker_path.write_text(str(num_slices))\n\n        return slice_entries\n\n    @staticmethod\n    def prepopulate_cache(\n        images_dir: Path,\n        labels_dir: Optional[Path],\n        slice_axis: Literal[\"x\", \"y\", \"z\"],\n        cache_dir: Path,\n        volume_files: Optional[List[str]] = None,\n        num_processes: int = -1,\n    ) -> None:\n        \"\"\"Prepopulate the entire slice cache using multiprocessing.\"\"\"\n        if not cache_dir:\n            raise ValueError(\"cache_dir must be provided.\")\n        if not images_dir.exists():\n            raise ValueError(f\"Images directory {images_dir} does not exist.\")\n\n        cache_dir.mkdir(parents=True, exist_ok=True)\n\n        if num_processes < 0:\n            num_processes = multiprocessing.cpu_count()\n        if volume_files is None:\n            volume_files = sorted([f.name for f in images_dir.glob(\"*.tif\")])\n        if not volume_files:\n            raise ValueError(f\"No TIFF files found in {images_dir}\")\n\n        process_func = partial(\n            VesuviusSliceDataset._process_volume_for_cache,\n            images_dir=images_dir,\n            labels_dir=labels_dir,\n            slice_axis=slice_axis,\n            cache_dir=cache_dir\n        )\n\n        print(f\"Prepopulating cache for {len(volume_files)} volumes...\")\n        all_slice_index = []\n        with multiprocessing.Pool(processes=num_processes) as pool:\n            for volume_slices in tqdm(\n                pool.imap_unordered(process_func, volume_files),\n                total=len(volume_files),\n                desc=\"Prepopulating cache\",\n            ):\n                all_slice_index.extend(volume_slices)\n\n        print(f\"Cache prepopulated with {len(all_slice_index)} slices.\")\n\nprint(\"Dataset class defined successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:17:10.855951Z","iopub.execute_input":"2026-01-08T15:17:10.856244Z","iopub.status.idle":"2026-01-08T15:17:10.884941Z","shell.execute_reply.started":"2026-01-08T15:17:10.856225Z","shell.execute_reply":"2026-01-08T15:17:10.884343Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 4: Topology-Aware Loss Functions (OPTIMIZED)\n# =============================================================================\n\nclass BoundaryLoss(nn.Module):\n    \"\"\"\n    Explicit boundary detection loss using Sobel filters.\n    Runs entirely on GPU.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        # Sobel filters for edge detection\n        sobel_x = torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], \n                               dtype=torch.float32).view(1, 1, 3, 3)\n        sobel_y = torch.tensor([[-1, -2, -1], [0, 0, 0], [1, 2, 1]], \n                               dtype=torch.float32).view(1, 1, 3, 3)\n        self.register_buffer('sobel_x', sobel_x)\n        self.register_buffer('sobel_y', sobel_y)\n    \n    def get_boundaries(self, mask: torch.Tensor) -> torch.Tensor:\n        \"\"\"Extract boundaries using Sobel operator on GPU.\"\"\"\n        if mask.dim() == 3:\n            mask = mask.unsqueeze(1)\n        \n        # Apply Sobel filters\n        edges_x = F.conv2d(mask, self.sobel_x, padding=1)\n        edges_y = F.conv2d(mask, self.sobel_y, padding=1)\n        \n        # Combine\n        edges = torch.sqrt(edges_x**2 + edges_y**2 + 1e-8)\n        return edges\n    \n    def forward(self, pred_proba: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        pred_fg = pred_proba[:, 1:2, :, :]\n        target_binary = (target == 1).float()\n        \n        pred_boundaries = self.get_boundaries(pred_fg)\n        target_boundaries = self.get_boundaries(target_binary)\n        \n        # MSE between predicted boundaries and target boundaries\n        return F.mse_loss(pred_boundaries, target_boundaries)\n\nclass ConnectivityLoss(nn.Module):\n    \"\"\"\n    Penalizes isolated pixels to encourage connectivity.\n    Runs entirely on GPU via convolution.\n    \"\"\"\n    def __init__(self, kernel_size: int = 3):\n        super().__init__()\n        self.kernel_size = kernel_size\n        kernel = torch.ones(1, 1, kernel_size, kernel_size) / (kernel_size ** 2)\n        self.register_buffer('avg_kernel', kernel)\n    \n    def forward(self, pred_proba: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        pred_fg = pred_proba[:, 1:2, :, :]\n        padding = self.kernel_size // 2\n        local_avg = F.conv2d(pred_fg, self.avg_kernel, padding=padding)\n        # Penalize pixels that are high but surrounded by low pixels\n        isolation_score = pred_fg * (1 - local_avg)\n        return isolation_score.mean()\n\nclass TopologyAwareLoss(nn.Module):\n    \"\"\"\n    Optimized Loss: Removed CPU-bound SurfaceLoss.\n    Increased weight of BoundaryLoss to compensate.\n    \"\"\"\n    def __init__(\n        self,\n        lambda_dice_ce: float = 1.0,\n        lambda_surface: float = 0.0, # Deprecated (too slow)\n        lambda_boundary: float = 0.5, # Increased from 0.2\n        lambda_connectivity: float = 0.2,\n        ignore_index: int = 2,\n    ):\n        super().__init__()\n        self.lambda_dice_ce = lambda_dice_ce\n        self.lambda_boundary = lambda_boundary\n        self.lambda_connectivity = lambda_connectivity\n        self.ignore_index = ignore_index\n        \n        from monai.losses import DiceCELoss\n        self.dice_ce_loss = DiceCELoss(\n            include_background=True,\n            to_onehot_y=False,\n            softmax=False,\n            lambda_dice=0.5,\n            lambda_ce=0.5,\n        )\n        self.boundary_loss = BoundaryLoss()\n        self.connectivity_loss = ConnectivityLoss()\n    \n    def forward(self, logits: torch.Tensor, target: torch.Tensor) -> Tuple[torch.Tensor, dict]:\n        pred_proba = torch.softmax(logits, dim=1)\n        valid_mask = (target != self.ignore_index).float()\n        \n        # Prepare targets\n        target_for_onehot = target.clone()\n        target_for_onehot[target_for_onehot == self.ignore_index] = 0\n        target_one_hot = F.one_hot(\n            target_for_onehot.squeeze(1).long(), \n            num_classes=logits.shape[1]\n        ).permute(0, 3, 1, 2).float()\n        \n        # Masking\n        target_one_hot = target_one_hot * valid_mask\n        pred_proba_masked = pred_proba * valid_mask\n        \n        # Compute Losses (All GPU)\n        loss_dice_ce = self.dice_ce_loss(pred_proba_masked, target_one_hot)\n        loss_boundary = self.boundary_loss(pred_proba, target)\n        loss_connectivity = self.connectivity_loss(pred_proba, target)\n        \n        total_loss = (\n            self.lambda_dice_ce * loss_dice_ce +\n            self.lambda_boundary * loss_boundary +\n            self.lambda_connectivity * loss_connectivity\n        )\n        \n        loss_dict = {\n            'dice_ce': loss_dice_ce.item(),\n            'boundary': loss_boundary.item(),\n            'connectivity': loss_connectivity.item(),\n            'total': total_loss.item(),\n        }\n        \n        return total_loss, loss_dict\n\nprint(\"Optimized Loss functions defined successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:44:14.372904Z","iopub.execute_input":"2026-01-08T15:44:14.373242Z","iopub.status.idle":"2026-01-08T15:44:14.388058Z","shell.execute_reply.started":"2026-01-08T15:44:14.373216Z","shell.execute_reply":"2026-01-08T15:44:14.387391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 5: Topology-Aware Metrics\n# =============================================================================\n\nclass TopologyMetrics:\n    \"\"\"Compute metrics that align with competition scoring.\"\"\"\n    \n    @staticmethod\n    def compute_surface_dice(\n        pred: np.ndarray, \n        target: np.ndarray, \n        tau: float = 2.0,\n        spacing: Tuple[float, ...] = (1.0, 1.0)\n    ) -> float:\n        \"\"\"Compute Surface Dice score with tolerance τ.\"\"\"\n        if pred.sum() == 0 and target.sum() == 0:\n            return 1.0\n        if pred.sum() == 0 or target.sum() == 0:\n            return 0.0\n        \n        pred_surface = TopologyMetrics._get_surface(pred)\n        target_surface = TopologyMetrics._get_surface(target)\n        \n        if pred_surface.sum() == 0 or target_surface.sum() == 0:\n            return 0.0\n        \n        pred_dist = ndimage.distance_transform_edt(~target, sampling=spacing)\n        target_dist = ndimage.distance_transform_edt(~pred, sampling=spacing)\n        \n        pred_matches = (pred_dist[pred_surface] <= tau).sum()\n        target_matches = (target_dist[target_surface] <= tau).sum()\n        \n        surface_dice = (pred_matches + target_matches) / (\n            pred_surface.sum() + target_surface.sum()\n        )\n        \n        return float(surface_dice)\n    \n    @staticmethod\n    def _get_surface(mask: np.ndarray) -> np.ndarray:\n        \"\"\"Extract surface voxels using morphological erosion.\"\"\"\n        if mask.sum() == 0:\n            return mask\n        \n        struct = ndimage.generate_binary_structure(mask.ndim, 1)\n        eroded = ndimage.binary_erosion(mask, structure=struct)\n        surface = mask & ~eroded\n        return surface\n    \n    @staticmethod\n    def compute_voi_score(\n        pred: np.ndarray, \n        target: np.ndarray,\n        alpha: float = 0.3\n    ) -> Tuple[float, float, float]:\n        \"\"\"Compute VOI-based score.\"\"\"\n        pred_labels, n_pred = ndimage.label(pred)\n        target_labels, n_target = ndimage.label(target)\n        \n        if n_pred == 0 and n_target == 0:\n            return 1.0, 0.0, 0.0\n        if n_pred == 0 or n_target == 0:\n            return 0.0, float('inf'), float('inf')\n        \n        max_label = max(pred_labels.max(), target_labels.max()) + 1\n        contingency = np.zeros((max_label, max_label), dtype=np.float64)\n        \n        for p, t in zip(pred_labels.ravel(), target_labels.ravel()):\n            contingency[p, t] += 1\n        \n        total = contingency.sum()\n        if total == 0:\n            return 1.0, 0.0, 0.0\n        \n        p_joint = contingency / total\n        p_pred = p_joint.sum(axis=1)\n        p_target = p_joint.sum(axis=0)\n        \n        voi_split = 0.0\n        for i in range(max_label):\n            if p_pred[i] > 0:\n                cond_prob = p_joint[i, :] / p_pred[i]\n                cond_prob = cond_prob[cond_prob > 0]\n                voi_split += p_pred[i] * entropy(cond_prob, base=2)\n        \n        voi_merge = 0.0\n        for j in range(max_label):\n            if p_target[j] > 0:\n                cond_prob = p_joint[:, j] / p_target[j]\n                cond_prob = cond_prob[cond_prob > 0]\n                voi_merge += p_target[j] * entropy(cond_prob, base=2)\n        \n        voi_total = voi_split + voi_merge\n        voi_score = 1.0 / (1.0 + alpha * voi_total)\n        \n        return voi_score, voi_split, voi_merge\n    \n    @staticmethod\n    def compute_betti_numbers(mask: np.ndarray) -> Tuple[int, int, int]:\n        \"\"\"Compute Betti numbers for topology scoring.\"\"\"\n        _, b0 = ndimage.label(mask)\n        \n        if mask.ndim == 2:\n            inverse = ~mask\n            labeled_bg, n_bg = ndimage.label(inverse)\n            b1 = max(0, n_bg - 1)\n            b2 = 0\n        else:\n            inverse = ~mask\n            labeled_bg, n_bg = ndimage.label(inverse)\n            b1 = max(0, n_bg - 1)\n            b2 = 0\n        \n        return b0, b1, b2\n    \n    @staticmethod\n    def compute_topo_score(\n        pred: np.ndarray,\n        target: np.ndarray,\n        weights: Tuple[float, float, float] = (0.34, 0.33, 0.33)\n    ) -> float:\n        \"\"\"Compute topology score based on Betti number matching.\"\"\"\n        pred_betti = TopologyMetrics.compute_betti_numbers(pred)\n        target_betti = TopologyMetrics.compute_betti_numbers(target)\n        \n        scores = []\n        active_dims = []\n        \n        for k, (pred_b, target_b, w) in enumerate(zip(pred_betti, target_betti, weights)):\n            if pred_b == 0 and target_b == 0:\n                continue\n            \n            active_dims.append(k)\n            \n            if pred_b == 0 or target_b == 0:\n                scores.append(0.0)\n            else:\n                match = min(pred_b, target_b)\n                precision = match / pred_b\n                recall = match / target_b\n                f1 = 2 * precision * recall / (precision + recall + 1e-8)\n                scores.append(f1 * w)\n        \n        if not scores:\n            return 1.0\n        \n        active_weights = [weights[k] for k in active_dims]\n        weight_sum = sum(active_weights)\n        \n        return sum(scores) / weight_sum if weight_sum > 0 else 0.0\n    \n    @staticmethod\n    def compute_competition_score(\n        pred: np.ndarray,\n        target: np.ndarray,\n        tau: float = 2.0,\n        spacing: Tuple[float, ...] = (1.0, 1.0)\n    ) -> dict:\n        \"\"\"\n        Compute full competition score.\n        Score = 0.30 × TopoScore + 0.35 × SurfaceDice@τ + 0.35 × VOI_score\n        \"\"\"\n        pred_binary = pred > 0\n        target_binary = target > 0\n        \n        surface_dice = TopologyMetrics.compute_surface_dice(\n            pred_binary, target_binary, tau, spacing\n        )\n        voi_score, voi_split, voi_merge = TopologyMetrics.compute_voi_score(\n            pred_binary, target_binary\n        )\n        topo_score = TopologyMetrics.compute_topo_score(\n            pred_binary, target_binary\n        )\n        \n        total_score = 0.30 * topo_score + 0.35 * surface_dice + 0.35 * voi_score\n        \n        return {\n            'total_score': total_score,\n            'surface_dice': surface_dice,\n            'voi_score': voi_score,\n            'voi_split': voi_split,\n            'voi_merge': voi_merge,\n            'topo_score': topo_score,\n        }\n\nprint(\"Topology metrics defined successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:44:20.238046Z","iopub.execute_input":"2026-01-08T15:44:20.238331Z","iopub.status.idle":"2026-01-08T15:44:20.256943Z","shell.execute_reply.started":"2026-01-08T15:44:20.23831Z","shell.execute_reply":"2026-01-08T15:44:20.256196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 6: Optimized Model\n# =============================================================================\n\nclass OptimizedVesuviusModel(pl.LightningModule):\n    \"\"\"\n    Optimized segmentation model with topology-aware loss and metrics.\n    \"\"\"\n    \n    def __init__(\n        self,\n        net: nn.Module,\n        out_channels: int = 2,\n        learning_rate: float = 1e-3,\n        lambda_dice_ce: float = 1.0,\n        lambda_surface: float = 0.3,\n        lambda_boundary: float = 0.2,\n        lambda_connectivity: float = 0.2,\n        compute_full_metrics_every_n_epochs: int = 1,\n    ):\n        super().__init__()\n        self.save_hyperparameters(ignore=[\"net\"])\n        \n        self.net = net\n        self.out_channels = out_channels\n        self.learning_rate = learning_rate\n        self.ignore_index = 2\n        self.compute_full_metrics_every_n_epochs = compute_full_metrics_every_n_epochs\n        \n        self.criterion = TopologyAwareLoss(\n            lambda_dice_ce=lambda_dice_ce,\n            lambda_surface=lambda_surface,\n            lambda_boundary=lambda_boundary,\n            lambda_connectivity=lambda_connectivity,\n            ignore_index=self.ignore_index,\n        )\n        \n        self.val_predictions = []\n        self.val_targets = []\n    \n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return self.net(x)\n    \n    def _compute_fast_metrics(\n        self, \n        logits: torch.Tensor, \n        targets: torch.Tensor\n    ) -> Dict[str, float]:\n        \"\"\"Fast metrics for every step (Dice, IoU).\"\"\"\n        pred_proba = torch.softmax(logits, dim=1)\n        pred_hard = torch.argmax(pred_proba, dim=1, keepdim=True)\n        \n        valid_mask = (targets != self.ignore_index).float()\n        \n        pred_fg = (pred_hard == 1).float() * valid_mask\n        target_fg = (targets == 1).float() * valid_mask\n        \n        intersection = (pred_fg * target_fg).sum()\n        union = pred_fg.sum() + target_fg.sum()\n        \n        dice = (2 * intersection + 1e-8) / (union + 1e-8)\n        iou = (intersection + 1e-8) / (union - intersection + 1e-8)\n        \n        return {'dice': dice.item(), 'iou': iou.item()}\n    \n    def training_step(self, batch: Tuple[torch.Tensor, torch.Tensor], batch_idx: int):\n        images, masks = batch\n        if masks.ndim == 3:\n            masks = masks.unsqueeze(1)\n        \n        logits = self(images)\n        loss, loss_dict = self.criterion(logits, masks)\n        \n        metrics = self._compute_fast_metrics(logits, masks)\n        \n        self.log('train_loss', loss, prog_bar=True)\n        self.log('train_dice', metrics['dice'], prog_bar=True)\n        self.log('train_iou', metrics['iou'])\n        \n        for k, v in loss_dict.items():\n            if k != 'total':\n                self.log(f'train_loss_{k}', v)\n        \n        return loss\n    \n    def validation_step(self, batch: Tuple[torch.Tensor, torch.Tensor], batch_idx: int):\n        images, masks = batch\n        if masks.ndim == 3:\n            masks = masks.unsqueeze(1)\n        \n        logits = self(images)\n        loss, loss_dict = self.criterion(logits, masks)\n        \n        metrics = self._compute_fast_metrics(logits, masks)\n        \n        self.log('val_loss', loss, prog_bar=True, sync_dist=True)\n        self.log('val_dice', metrics['dice'], prog_bar=True, sync_dist=True)\n        self.log('val_iou', metrics['iou'], sync_dist=True)\n        \n        if self.current_epoch % self.compute_full_metrics_every_n_epochs == 0:\n            pred_proba = torch.softmax(logits, dim=1)\n            pred_hard = torch.argmax(pred_proba, dim=1).cpu().numpy()\n            target_np = masks.squeeze(1).cpu().numpy()\n            \n            if len(self.val_predictions) < 50:\n                for i in range(min(4, pred_hard.shape[0])):\n                    self.val_predictions.append(pred_hard[i])\n                    self.val_targets.append(target_np[i])\n        \n        return loss\n    \n    def on_validation_epoch_end(self):\n        \"\"\"Compute topology-aware metrics at epoch end.\"\"\"\n        if self.current_epoch % self.compute_full_metrics_every_n_epochs != 0:\n            return\n        \n        if not self.val_predictions:\n            return\n        \n        surface_dices = []\n        voi_scores = []\n        topo_scores = []\n        \n        for pred, target in zip(self.val_predictions, self.val_targets):\n            valid = target != self.ignore_index\n            pred_valid = (pred == 1) & valid\n            target_valid = (target == 1) & valid\n            \n            try:\n                metrics = TopologyMetrics.compute_competition_score(\n                    pred_valid.astype(np.uint8),\n                    target_valid.astype(np.uint8),\n                    tau=2.0\n                )\n                surface_dices.append(metrics['surface_dice'])\n                voi_scores.append(metrics['voi_score'])\n                topo_scores.append(metrics['topo_score'])\n            except Exception as e:\n                continue\n        \n        if surface_dices:\n            avg_surface_dice = np.mean(surface_dices)\n            avg_voi = np.mean(voi_scores)\n            avg_topo = np.mean(topo_scores)\n            \n            competition_score = 0.30 * avg_topo + 0.35 * avg_surface_dice + 0.35 * avg_voi\n            \n            self.log('val_surface_dice', avg_surface_dice, prog_bar=True)\n            self.log('val_voi_score', avg_voi)\n            self.log('val_topo_score', avg_topo)\n            self.log('val_competition_score', competition_score, prog_bar=True)\n        \n        self.val_predictions = []\n        self.val_targets = []\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(\n            self.parameters(), \n            lr=self.learning_rate, \n            weight_decay=1e-2,\n            betas=(0.9, 0.999)\n        )\n        \n        # Calculate total steps\n        if self.trainer and self.trainer.estimated_stepping_batches:\n            total_steps = self.trainer.estimated_stepping_batches\n        else:\n            total_steps = 10000  # Default fallback\n        \n        from torch.optim.lr_scheduler import OneCycleLR\n        scheduler = OneCycleLR(\n            optimizer,\n            max_lr=self.learning_rate,\n            total_steps=total_steps,\n            pct_start=0.1,\n            anneal_strategy='cos',\n            div_factor=25,\n            final_div_factor=1000,\n        )\n        \n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\n                \"scheduler\": scheduler,\n                \"interval\": \"step\",\n            }\n        }\n\nprint(\"Optimized model defined successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:44:26.663287Z","iopub.execute_input":"2026-01-08T15:44:26.663577Z","iopub.status.idle":"2026-01-08T15:44:26.681043Z","shell.execute_reply.started":"2026-01-08T15:44:26.663555Z","shell.execute_reply":"2026-01-08T15:44:26.680439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 7: Post-Processing Functions\n# =============================================================================\n\ndef gpu_morphological_close(\n    volume: np.ndarray, \n    radius: int, \n    device: str = \"cuda\"\n) -> np.ndarray:\n    \"\"\"GPU-accelerated morphological closing using spherical element.\"\"\"\n    kernel_np = ball(radius)\n    kernel = torch.from_numpy(kernel_np.astype(np.float32)).unsqueeze(0).unsqueeze(0)\n    kernel = kernel.to(device)\n    \n    input_tensor = torch.from_numpy(volume.astype(np.float32))\n    input_tensor = input_tensor.unsqueeze(0).unsqueeze(0).to(device)\n    \n    padding = radius\n    \n    dilated = (F.conv3d(input_tensor, kernel, padding=padding) > 0).float()\n    eroded = 1.0 - (F.conv3d(1.0 - dilated, kernel, padding=padding) > 0).float()\n    \n    return eroded.squeeze().cpu().numpy().astype(bool)\n\n\ndef remove_thin_bridges(\n    volume: np.ndarray, \n    threshold: int = 5\n) -> np.ndarray:\n    \"\"\"Remove thin bridges that connect separate components.\"\"\"\n    if volume.sum() == 0:\n        return volume\n    \n    dist = ndimage.distance_transform_edt(volume)\n    thin_mask = dist < threshold\n    \n    labeled_orig, n_orig = ndimage.label(volume)\n    tentative = volume & ~thin_mask\n    labeled_new, n_new = ndimage.label(tentative)\n    \n    if n_new > n_orig:\n        for label_id in range(1, n_new + 1):\n            component = labeled_new == label_id\n            if component.sum() < 100:\n                tentative = tentative | (volume & component)\n        return tentative\n    else:\n        return volume\n\n\ndef detect_and_separate_wraps(\n    volume: np.ndarray,\n    min_separation: int = 10,\n    n_iterations: int = 3,\n) -> np.ndarray:\n    \"\"\"Detect and separate adjacent papyrus wraps that are incorrectly merged.\"\"\"\n    binary_vol = volume > 0\n    \n    for _ in range(n_iterations):\n        dist = ndimage.distance_transform_edt(binary_vol)\n        \n        local_max = peak_local_max(\n            dist, \n            min_distance=min_separation,\n            threshold_abs=min_separation / 2,\n            exclude_border=False\n        )\n        \n        if len(local_max) <= 1:\n            break\n        \n        markers = np.zeros(binary_vol.shape, dtype=np.int32)\n        for i, coord in enumerate(local_max):\n            markers[tuple(coord)] = i + 1\n        \n        watershed_result = watershed(-dist, markers, mask=binary_vol)\n        \n        struct = ndimage.generate_binary_structure(3, 1)\n        dilated_regions = []\n        for label_id in range(1, len(local_max) + 1):\n            region = watershed_result == label_id\n            dilated = ndimage.binary_dilation(region, struct)\n            dilated_regions.append(dilated)\n        \n        if len(dilated_regions) >= 2:\n            boundaries = np.zeros_like(binary_vol)\n            for i in range(len(dilated_regions)):\n                for j in range(i + 1, len(dilated_regions)):\n                    boundaries |= dilated_regions[i] & dilated_regions[j]\n            \n            thin_boundary = boundaries & (dist < min_separation / 2)\n            binary_vol = binary_vol & ~thin_boundary\n    \n    return binary_vol.astype(np.uint8)\n\n\ndef adaptive_threshold_refinement(\n    prediction_proba: np.ndarray,\n    base_threshold: float = 0.5,\n    edge_threshold: float = 0.7,\n    edge_width: int = 20,\n) -> np.ndarray:\n    \"\"\"Apply adaptive thresholding with higher threshold at edges.\"\"\"\n    h, w = prediction_proba.shape[-2:]\n    edge_weight = np.ones_like(prediction_proba)\n    \n    if prediction_proba.ndim == 2:\n        edge_weight[:edge_width, :] = 0\n        edge_weight[-edge_width:, :] = 0\n        edge_weight[:, :edge_width] = 0\n        edge_weight[:, -edge_width:] = 0\n        \n        for i in range(edge_width):\n            factor = i / edge_width\n            edge_weight[i, :] = max(edge_weight[i, :].min(), factor)\n            edge_weight[-i-1, :] = max(edge_weight[-i-1, :].min(), factor)\n            edge_weight[:, i] = max(edge_weight[:, i].min(), factor)\n            edge_weight[:, -i-1] = max(edge_weight[:, -i-1].min(), factor)\n    \n    threshold_map = base_threshold + (edge_threshold - base_threshold) * (1 - edge_weight)\n    \n    return (prediction_proba > threshold_map).astype(np.uint8)\n\n\ndef topology_aware_postprocess(\n    volume: np.ndarray,\n    min_component_size: int = 1000,\n    max_hole_size: int = 500,\n    closing_radius: int = 3,\n    smoothing_sigma: float = 1.0,\n    thin_bridge_threshold: int = 5,\n    device: str = \"cuda\",\n) -> np.ndarray:\n    \"\"\"Topology-aware post-processing optimized for competition metrics.\"\"\"\n    binary_vol = volume > 0\n    \n    if smoothing_sigma > 0:\n        smoothed = ndimage.gaussian_filter(binary_vol.astype(np.float32), sigma=smoothing_sigma)\n        binary_vol = smoothed > 0.5\n    \n    if closing_radius > 0:\n        try:\n            binary_vol = gpu_morphological_close(binary_vol, closing_radius, device)\n        except Exception as e:\n            print(f\"GPU morphological close failed, using CPU: {e}\")\n            from skimage.morphology import binary_closing\n            binary_vol = binary_closing(binary_vol, ball(closing_radius))\n    \n    binary_vol = remove_small_holes(binary_vol, area_threshold=max_hole_size)\n    \n    if thin_bridge_threshold > 0:\n        binary_vol = remove_thin_bridges(binary_vol, threshold=thin_bridge_threshold)\n    \n    binary_vol = remove_small_objects(binary_vol, min_size=min_component_size)\n    \n    return binary_vol.astype(np.uint8)\n\n\ndef remove_edge_artifacts(mask: np.ndarray, margin: int = 20) -> np.ndarray:\n    \"\"\"Zeroes out margins of a 2D mask to remove edge artifacts.\"\"\"\n    h, w = mask.shape\n    if h > 2 * margin:\n        mask[:margin, :] = 0\n        mask[-margin:, :] = 0\n    if w > 2 * margin:\n        mask[:, :margin] = 0\n        mask[:, -margin:] = 0\n    return mask\n\n\ndef resize_nearest_neighbor(\n    image_array: np.ndarray, target_height: int, target_width: int\n) -> np.ndarray:\n    \"\"\"Resizes a 2D numpy array using nearest neighbor interpolation.\"\"\"\n    original_h, original_w = image_array.shape\n    y_indices = np.linspace(0, original_h - 1, target_height, dtype=np.int32)\n    x_indices = np.linspace(0, original_w - 1, target_width, dtype=np.int32)\n    resized_image = image_array[np.ix_(y_indices, x_indices)]\n    return resized_image\n\nprint(\"Post-processing functions defined successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:44:35.902214Z","iopub.execute_input":"2026-01-08T15:44:35.903202Z","iopub.status.idle":"2026-01-08T15:44:35.921963Z","shell.execute_reply.started":"2026-01-08T15:44:35.903162Z","shell.execute_reply":"2026-01-08T15:44:35.921136Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 8: Data Module with Custom MONAI-style Augmentations (Albumentations)\n# =============================================================================\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport numpy as np\n\n\nclass ShiftIntensity(A.ImageOnlyTransform):\n    \"\"\"Shift intensity values by a random offset (MONAI-style).\"\"\"\n    \n    def __init__(self, offset_range: float = 0.1, always_apply=False, p=0.5):\n        super().__init__(always_apply, p)\n        self.offset_range = offset_range\n    \n    def apply(self, img, offset=0, **params):\n        return np.clip(img + offset, 0, 1).astype(img.dtype)\n    \n    def get_params(self):\n        return {\"offset\": np.random.uniform(-self.offset_range, self.offset_range)}\n    \n    def get_transform_init_args_names(self):\n        return (\"offset_range\",)\n\n\nclass ScaleIntensityRange(A.ImageOnlyTransform):\n    \"\"\"Scale intensity from [a_min, a_max] to [b_min, b_max] (MONAI-style).\"\"\"\n    \n    def __init__(\n        self, \n        a_min: float = 0, \n        a_max: float = 65535,  # 16-bit\n        b_min: float = 0.0, \n        b_max: float = 1.0,\n        clip: bool = True,\n        always_apply=True, \n        p=1.0\n    ):\n        super().__init__(always_apply, p)\n        self.a_min = a_min\n        self.a_max = a_max\n        self.b_min = b_min\n        self.b_max = b_max\n        self.clip = clip\n    \n    def apply(self, img, **params):\n        img = img.astype(np.float32)\n        img = (img - self.a_min) / (self.a_max - self.a_min + 1e-8)\n        img = img * (self.b_max - self.b_min) + self.b_min\n        if self.clip:\n            img = np.clip(img, self.b_min, self.b_max)\n        return img\n    \n    def get_transform_init_args_names(self):\n        return (\"a_min\", \"a_max\", \"b_min\", \"b_max\", \"clip\")\n\n\nclass RandSpatialCrop(A.DualTransform):\n    \"\"\"Random spatial crop (MONAI-style) - applies same crop to image and mask.\"\"\"\n    \n    def __init__(\n        self, \n        crop_height: int, \n        crop_width: int,\n        always_apply=False, \n        p=1.0\n    ):\n        super().__init__(always_apply, p)\n        self.crop_height = crop_height\n        self.crop_width = crop_width\n    \n    def apply(self, img, x_start=0, y_start=0, **params):\n        return img[y_start:y_start + self.crop_height, x_start:x_start + self.crop_width]\n    \n    def apply_to_mask(self, mask, x_start=0, y_start=0, **params):\n        return mask[y_start:y_start + self.crop_height, x_start:x_start + self.crop_width]\n    \n    def get_params_dependent_on_targets(self, params):\n        img = params[\"image\"]\n        h, w = img.shape[:2]\n        \n        if h < self.crop_height or w < self.crop_width:\n            raise ValueError(\n                f\"Image size ({h}, {w}) is smaller than crop size \"\n                f\"({self.crop_height}, {self.crop_width})\"\n            )\n        \n        x_start = np.random.randint(0, w - self.crop_width + 1)\n        y_start = np.random.randint(0, h - self.crop_height + 1)\n        \n        return {\"x_start\": x_start, \"y_start\": y_start}\n    \n    @property\n    def targets_as_params(self):\n        return [\"image\"]\n    \n    def get_transform_init_args_names(self):\n        return (\"crop_height\", \"crop_width\")\n\nimport os\nclass VesuviusDataModule(pl.LightningDataModule):\n    \"\"\"DataModule with MONAI-style augmentations using Albumentations.\"\"\"\n\n    def __init__(\n        self,\n        images_dir: Path,\n        labels_dir: Path,\n        slice_axis: Literal[\"x\", \"y\", \"z\"],\n        image_size: Tuple[int, int],\n        cache_dir: Optional[Path],\n        use_cache: bool,\n        batch_size: int,\n        val_split: float,\n        num_workers: 4,\n        seed: int = 42,\n        # NEW: Augmentation parameters\n        crop_size: Optional[Tuple[int, int]] = None,\n        intensity_shift_range: float = 0.1,\n        a_min: float = 0,\n        a_max: float = 65535,\n    ):\n        super().__init__()\n        self.images_dir = images_dir\n        self.labels_dir = labels_dir\n        self.slice_axis = slice_axis\n        self.image_size = image_size\n        self.crop_size = crop_size or (int(image_size[0] * 0.8), int(image_size[1] * 0.8))\n        self.cache_dir = cache_dir\n        self.use_cache = use_cache\n        self.batch_size = batch_size\n        if num_workers <= 0:\n            num_workers = min(8, os.cpu_count())\n        self.num_workers = num_workers\n        self.val_split = val_split\n        self.seed = seed\n        self.intensity_shift_range = intensity_shift_range\n        self.a_min = a_min\n        self.a_max = a_max\n\n    def setup(self, stage: Optional[str] = None):\n        \"\"\"Setup datasets.\"\"\"\n        volume_files = sorted([f.name for f in self.images_dir.glob(\"*.tif\")])\n\n        if not volume_files:\n            raise ValueError(f\"No TIFF files found in {self.images_dir}\")\n\n        print(f\"Found {len(volume_files)} volume files\")\n\n        train_files, val_files = train_test_split(\n            volume_files, test_size=self.val_split, random_state=self.seed\n        )\n\n        print(f\"Train: {len(train_files)} volumes, Val: {len(val_files)} volumes\")\n\n        # =====================================================================\n        # Training Transforms with MONAI-style augmentations\n        # =====================================================================\n        self.train_transform = A.Compose([\n            # 1. ScaleIntensityRange: Normalize to [0, 1]\n            ScaleIntensityRange(\n                a_min=self.a_min,\n                a_max=self.a_max,\n                b_min=0.0,\n                b_max=1.0,\n                clip=True,\n            ),\n            \n            # 2. Resize to larger size for cropping\n            A.Resize(\n                height=int(self.crop_size[0] * 1.25),\n                width=int(self.crop_size[1] * 1.25),\n            ),\n            \n            # 3. RandSpatialCrop: Random crop\n            RandSpatialCrop(\n                crop_height=self.crop_size[0],\n                crop_width=self.crop_size[1],\n                p=1.0,\n            ),\n            \n            # 4. Final resize to target size\n            A.Resize(\n                height=self.image_size[0],\n                width=self.image_size[1],\n            ),\n            \n            # 5. ShiftIntensity: Random intensity shift\n            ShiftIntensity(\n                offset_range=self.intensity_shift_range,\n                p=0.5,\n            ),\n                 \n            # 7. Additional augmentations\n            A.GaussNoise(var_limit=(0.001, 0.01), p=0.3),\n            A.GaussianBlur(blur_limit=(3, 5), p=0.2),\n            \n            # 8. Final normalization for model input\n            A.Normalize(mean=[0.5], std=[0.5]),\n            ToTensorV2(),\n        ])\n        \n        # =====================================================================\n        # Validation Transforms (no augmentation, just preprocessing)\n        # =====================================================================\n        self.val_transform = A.Compose([\n            ScaleIntensityRange(\n                a_min=self.a_min,\n                a_max=self.a_max,\n                b_min=0.0,\n                b_max=1.0,\n                clip=True,\n            ),\n            A.Resize(\n                height=self.image_size[0],\n                width=self.image_size[1],\n            ),\n            A.Normalize(mean=[0.5], std=[0.5]),\n            ToTensorV2(),\n        ])\n\n        self.train_dataset = VesuviusSliceDataset(\n            images_dir=self.images_dir,\n            labels_dir=self.labels_dir,\n            volume_files=train_files,\n            slice_axis=self.slice_axis,\n            cache_dir=self.cache_dir,\n            use_cache=self.use_cache,\n            transform=self.train_transform,\n        )\n        \n        self.val_dataset = VesuviusSliceDataset(\n            images_dir=self.images_dir,\n            labels_dir=self.labels_dir,\n            volume_files=val_files,\n            slice_axis=self.slice_axis,\n            cache_dir=self.cache_dir,\n            use_cache=self.use_cache,\n            transform=self.val_transform,\n        )\n\n        print(f\"Train dataset: {len(self.train_dataset)} slices\")\n        print(f\"Val dataset: {len(self.val_dataset)} slices\")\n\n    def train_dataloader(self) -> DataLoader:\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            shuffle=True,\n            num_workers=self.num_workers,\n            pin_memory=True,\n        )\n\n    def val_dataloader(self) -> DataLoader:\n        return DataLoader(\n            self.val_dataset,\n            batch_size=self.batch_size,\n            shuffle=False,\n            num_workers=self.num_workers,\n            pin_memory=True,\n        )\n\n\nprint(\"Data module with MONAI-style augmentations defined successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:45:30.407117Z","iopub.execute_input":"2026-01-08T15:45:30.407868Z","iopub.status.idle":"2026-01-08T15:45:30.429824Z","shell.execute_reply.started":"2026-01-08T15:45:30.407841Z","shell.execute_reply":"2026-01-08T15:45:30.42902Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 9: Prepopulate Cache (if needed)\n# =============================================================================\n\n# Only run if cache doesn't exist\nif not CACHE_DIR.exists() or len(list(CACHE_DIR.glob(\"*/*.npz\"))) == 0:\n    print(\"Cache not found, prepopulating...\")\n    VesuviusSliceDataset.prepopulate_cache(\n        images_dir=TRAIN_IMAGES_DIR,\n        labels_dir=TRAIN_LABELS_DIR,\n        slice_axis=SLICE_AXIS,\n        cache_dir=CACHE_DIR,\n    )\nelse:\n    print(f\"Using existing cache at {CACHE_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:44:50.58547Z","iopub.execute_input":"2026-01-08T15:44:50.586226Z","iopub.status.idle":"2026-01-08T15:44:52.016177Z","shell.execute_reply.started":"2026-01-08T15:44:50.58619Z","shell.execute_reply":"2026-01-08T15:44:52.015521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 10: Create Data Module\n# =============================================================================\nCROP_SIZE = (280,280)\ndata_module = VesuviusDataModule(\n    images_dir=TRAIN_IMAGES_DIR,\n    labels_dir=TRAIN_LABELS_DIR,\n    slice_axis=SLICE_AXIS,\n    image_size=IMAGE_SIZE,\n    crop_size=CROP_SIZE,\n    cache_dir=CACHE_DIR,\n    use_cache=USE_CACHE,\n    batch_size=BATCH_SIZE,\n    val_split=VAL_SPLIT,\n    num_workers=4,\n    seed=SEED,\n)\ndata_module.setup()\n\nprint(\"\\nData module created successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:45:35.121865Z","iopub.execute_input":"2026-01-08T15:45:35.122633Z","iopub.status.idle":"2026-01-08T15:45:35.146925Z","shell.execute_reply.started":"2026-01-08T15:45:35.122574Z","shell.execute_reply":"2026-01-08T15:45:35.145988Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 11: Visualize Sample\n# =============================================================================\n\ndef visualize_sample(\n    dataset,\n    idx: int,\n    figsize: Tuple[int, int] = (18, 5)\n):\n    \"\"\"Visualize a sample from the dataset.\"\"\"\n    image, mask = dataset[idx]\n    image = image.cpu().numpy()\n    mask = mask.cpu().numpy()\n\n    fig, axes = plt.subplots(1, 3, figsize=figsize)\n\n    mid_channel = image.shape[0] // 2\n    axes[0].imshow(image[mid_channel], cmap='gray')\n    axes[0].set_title(f'Image (channel {mid_channel}/{image.shape[0]})')\n    axes[0].axis('off')\n\n    img4contour = image[mid_channel].copy()\n    if img4contour.max() > 1.0:\n        img4contour = img4contour / img4contour.max()\n    \n    # Normalize for display\n    img4contour = (img4contour - img4contour.min()) / (img4contour.max() - img4contour.min() + 1e-8)\n    \n    binary_mask = (mask > 0).astype(int)\n    if binary_mask.ndim > 2:\n        binary_mask = binary_mask.squeeze()\n    \n    contoured_image = mark_boundaries(img4contour, binary_mask, outline_color=(1, 0, 0))\n    axes[1].imshow(contoured_image)\n    axes[1].set_title('Image with Mask Contours')\n    axes[1].axis('off')\n\n    mask_display = mask.squeeze() if mask.ndim > 2 else mask\n    im = axes[2].imshow(mask_display, cmap='hot', vmin=0, vmax=2)\n    axes[2].set_title('Ground Truth Mask (0: Bg, 1: Fg, 2: Ignored)')\n    axes[2].axis('off')\n\n    cbar = fig.colorbar(im, ax=axes[2], fraction=0.046, pad=0.04)\n    cbar.set_ticks([0, 1, 2])\n    cbar.set_ticklabels(['0 (Background)', '1 (Foreground)', '2 (Ignored)'])\n\n    plt.tight_layout()\n    plt.show()\n\n# Visualize a few samples\nfor i in [0, len(data_module.train_dataset) // 2, len(data_module.train_dataset) - 1]:\n    print(f\"Sample {i}:\")\n    visualize_sample(data_module.train_dataset, i)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:29:15.354975Z","iopub.execute_input":"2026-01-08T15:29:15.355244Z","iopub.status.idle":"2026-01-08T15:29:17.173815Z","shell.execute_reply.started":"2026-01-08T15:29:15.355225Z","shell.execute_reply":"2026-01-08T15:29:17.172966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 12: Create Model\n# =============================================================================\n\nfrom monai.networks.nets import SegResNet\n\n# Create network\nnet = SegResNet(\n    spatial_dims=2,\n    in_channels=1,\n    out_channels=2,\n    init_filters=32,  # Increased capacity\n    dropout_prob=0.2,\n)\n\nnet_name = net.__class__.__name__\nprint(f\"Network: {net_name}\")\n\n# Count parameters\ntotal_params = sum(p.numel() for p in net.parameters())\ntrainable_params = sum(p.numel() for p in net.parameters() if p.requires_grad)\nprint(f\"Total parameters: {total_params:,}\")\nprint(f\"Trainable parameters: {trainable_params:,}\")\n\n# Create optimized model\nmodel = OptimizedVesuviusModel(\n    net=net,\n    learning_rate=LEARNING_RATE,\n    lambda_dice_ce=LAMBDA_DICE_CE,\n    lambda_surface=LAMBDA_SURFACE,\n    lambda_boundary=LAMBDA_BOUNDARY,\n    lambda_connectivity=LAMBDA_CONNECTIVITY,\n    compute_full_metrics_every_n_epochs=1,\n)\n\nprint(\"\\nModel created successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:45:39.306293Z","iopub.execute_input":"2026-01-08T15:45:39.306943Z","iopub.status.idle":"2026-01-08T15:45:39.381502Z","shell.execute_reply.started":"2026-01-08T15:45:39.306908Z","shell.execute_reply":"2026-01-08T15:45:39.380728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 13: Find Best Checkpoint\n# =============================================================================\n\ndef get_best_checkpoint(\n    checkpoint_dirs: List[Path],\n    name: str = \"\",\n) -> Tuple[Optional[str], Optional[float]]:\n    \"\"\"Finds the checkpoint with the highest val_competition_score or val_dice.\"\"\"\n    if not isinstance(checkpoint_dirs, list):\n        checkpoint_dirs = [checkpoint_dirs]\n    checkpoint_dirs = [d for d in checkpoint_dirs if Path(d).exists()]\n    \n    if not checkpoint_dirs:\n        print(\"No valid checkpoint directories provided.\")\n        return None, None\n    \n    checkpoints = []\n    \n    # Try competition score first, then dice\n    patterns = [\n        (re.compile(r\"val_competition_score=?([0-9]+\\.[0-9]+)\"), \"competition_score\"),\n        (re.compile(r\"val_dice=?([0-9]+\\.[0-9]+)\"), \"dice\"),\n    ]\n    \n    for path in [f for d in checkpoint_dirs for f in Path(d).glob(f\"{name}*.ckpt\")]:\n        for pattern, metric_name in patterns:\n            match = pattern.search(path.name)\n            if match:\n                checkpoints.append((float(match.group(1)), str(path), metric_name))\n                break\n\n    if not checkpoints:\n        print(\"No valid checkpoints found.\")\n        return None, None\n\n    checkpoints.sort(key=lambda x: x[0], reverse=True)\n    best_score, best_path, metric = checkpoints[0]\n    print(f\"Found {len(checkpoints)} checkpoints.\")\n    print(f\"Best ({metric}={best_score:.4f}): {Path(best_path).name}\")\n    return best_path, best_score\n\nckpt_path, ckpt_score = get_best_checkpoint(\n    [OUTPUT_DIR, Path(CHECKPOINT_DIR)], \n    name=net_name\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:45:40.757834Z","iopub.execute_input":"2026-01-08T15:45:40.758331Z","iopub.status.idle":"2026-01-08T15:45:40.767188Z","shell.execute_reply.started":"2026-01-08T15:45:40.758302Z","shell.execute_reply":"2026-01-08T15:45:40.766434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 14: Setup Training\n# =============================================================================\n\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor\nfrom pytorch_lightning.loggers import CSVLogger\n\n# Callbacks\ncheckpoint_callback = ModelCheckpoint(\n    dirpath=OUTPUT_DIR,\n    filename=f\"{net_name}-{{epoch:02d}}-{{val_competition_score:.4f}}\",\n    monitor=\"val_competition_score\",\n    mode=\"max\",\n    save_top_k=3,\n    save_last=True,\n)\n\n# Fallback checkpoint if competition score not available\ncheckpoint_callback_fallback = ModelCheckpoint(\n    dirpath=OUTPUT_DIR,\n    filename=f\"{net_name}-{{epoch:02d}}-{{val_dice:.4f}}\",\n    monitor=\"val_dice\",\n    mode=\"max\",\n    save_top_k=2,\n)\n\nearly_stop_callback = EarlyStopping(\n    monitor=\"val_dice\",\n    patience=15,\n    mode=\"max\",\n    verbose=True,\n)\n\nlr_monitor = LearningRateMonitor(logging_interval=\"step\")\n\ncsv_logger = CSVLogger(save_dir=OUTPUT_DIR)\n\n# Trainer\ntrainer = pl.Trainer(\n    max_epochs=NUM_EPOCHS,\n    accelerator=\"auto\",\n    callbacks=[\n        checkpoint_callback, \n        checkpoint_callback_fallback,\n        lr_monitor,\n        early_stop_callback,\n    ],\n    logger=csv_logger,\n    log_every_n_steps=1,       # Reduce logging overhead\n    accumulate_grad_batches=1,  # Speed up loop (batch 64 is big enough, no need to accumulate)\n    val_check_interval=1.0,     # Only validate once at the END of the epoch\n    check_val_every_n_epoch=1,\n    precision='16-mixed',\n    gradient_clip_val=1.0,\n    benchmark=True              # Optimizes CUDNN\n)\n\nprint(\"Trainer configured successfully!\")\nprint(f\"Max epochs: {NUM_EPOCHS}\")\nprint(f\"Accumulate grad batches: 4\")\nprint(f\"Precision: 16-mixed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:45:42.424362Z","iopub.execute_input":"2026-01-08T15:45:42.425242Z","iopub.status.idle":"2026-01-08T15:45:42.486793Z","shell.execute_reply.started":"2026-01-08T15:45:42.425207Z","shell.execute_reply":"2026-01-08T15:45:42.486037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 15: Train Model\n# =============================================================================\n\nfrom pytorch_lightning.utilities.exceptions import MisconfigurationException\n\nprint(\"=\"*60)\nprint(\"Starting Training\")\nprint(\"=\"*60)\n\ntry:\n    trainer.fit(model, datamodule=data_module, ckpt_path=ckpt_path)\n    print(\"\\nTraining completed successfully!\")\nexcept MisconfigurationException as ex:\n    print(f\"Configuration error: {ex}\")\nexcept Exception as ex:\n    print(f\"Training error: {ex}\")\n    raise","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:45:44.285292Z","iopub.execute_input":"2026-01-08T15:45:44.2859Z","iopub.status.idle":"2026-01-08T15:46:14.065762Z","shell.execute_reply.started":"2026-01-08T15:45:44.285875Z","shell.execute_reply":"2026-01-08T15:46:14.064476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 16: Plot Training Metrics\n# =============================================================================\n\nimport pandas as pd\nimport seaborn as sns\nfrom IPython.display import display\n\nsns.set_style(\"whitegrid\")\n\nlog_base_dir = Path(trainer.logger.save_dir) / 'lightning_logs'\nmetrics_path = log_base_dir / f\"version_{trainer.logger._version}\" / 'metrics.csv'\n\nprint(f\"Loading metrics from: {metrics_path}\")\n\nif metrics_path.exists():\n    metrics = pd.read_csv(metrics_path)\n    display(metrics.dropna(axis=1, how=\"all\").tail(10))\n    \n    metrics.ffill(inplace=True)\n    \n    # Define metric groups\n    metric_groups = {\n        'Total Loss': [c for c in metrics.columns if c in ['train_loss', 'val_loss']],\n        'Component Losses': [c for c in metrics.columns if 'loss_' in c],\n        'Dice Score': [c for c in metrics.columns if 'dice' in c.lower()],\n        'Competition Metrics': [c for c in metrics.columns if c in \n                                ['val_surface_dice', 'val_voi_score', 'val_topo_score', 'val_competition_score']],\n    }\n    \n    for title, metric_list in metric_groups.items():\n        if not metric_list or not any(c in metrics.columns for c in metric_list):\n            continue\n            \n        available_metrics = [c for c in metric_list if c in metrics.columns]\n        if not available_metrics:\n            continue\n        \n        fig, ax = plt.subplots(figsize=(12, 5))\n        \n        for metric_name in available_metrics:\n            valid_data = metrics[['epoch', metric_name]].dropna()\n            if len(valid_data) > 0:\n                ax.plot(valid_data['epoch'], valid_data[metric_name], \n                       label=metric_name, linewidth=2)\n        \n        ax.set_title(f'{title} over Epochs', fontsize=14, fontweight='bold')\n        ax.set_xlabel('Epoch', fontsize=12)\n        ax.set_ylabel(title, fontsize=12)\n        ax.legend(loc='best')\n        ax.grid(True, alpha=0.3)\n        \n        if 'Loss' in title:\n            ax.set_yscale('log')\n        \n        plt.tight_layout()\n        plt.show()\nelse:\n    print(\"Metrics file not found!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.936004Z","iopub.status.idle":"2026-01-08T15:30:28.936319Z","shell.execute_reply.started":"2026-01-08T15:30:28.936189Z","shell.execute_reply":"2026-01-08T15:30:28.936204Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 17: Load Best Model for Inference\n# =============================================================================\n\n# Find best checkpoint\nbest_checkpoint_path, best_score = get_best_checkpoint(\n    [OUTPUT_DIR, Path(CHECKPOINT_DIR)], \n    name=net_name\n)\n\nif best_checkpoint_path:\n    print(f\"\\nLoading best checkpoint: {best_checkpoint_path}\")\n    model = OptimizedVesuviusModel.load_from_checkpoint(\n        best_checkpoint_path, \n        net=net,\n        strict=False,\n    )\nelse:\n    print(\"No checkpoint found, using current model state.\")\n\nmodel.eval()\nmodel.to(DEVICE)\n\nprint(f\"Model loaded and moved to {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.937632Z","iopub.status.idle":"2026-01-08T15:30:28.937899Z","shell.execute_reply.started":"2026-01-08T15:30:28.93777Z","shell.execute_reply":"2026-01-08T15:30:28.937784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 18: Inference Functions\n# =============================================================================\n\ndef run_inference(\n    model: nn.Module,\n    volume: np.ndarray,\n    slice_axis: str = \"z\",\n    transform=None,\n    device: str = \"cuda\",\n    use_adaptive_threshold: bool = True,\n) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"\n    Run inference on a 3D volume.\n    \n    Returns:\n        predictions_volume: Binary predictions\n        predictions_proba: Probability predictions\n    \"\"\"\n    model.eval()\n    predictions_list = []\n    proba_list = []\n    \n    axis_idx = VesuviusSliceDataset._map_axis_name_to_index(slice_axis)\n    num_slices = volume.shape[axis_idx]\n    \n    for i in tqdm(range(num_slices), desc=\"Inference\"):\n        image_slice = VesuviusSliceDataset._extract_slice_from_volume(\n            volume, i, slice_axis\n        ).squeeze(0)\n        \n        original_h, original_w = image_slice.shape\n        \n        # Normalize\n        image_normalized = image_slice.astype(np.float32) / 255.0\n        \n        if transform:\n            transformed = transform(image=image_normalized)\n            image_tensor = transformed[\"image\"].unsqueeze(0).to(device)\n        else:\n            image_tensor = torch.from_numpy(image_normalized).unsqueeze(0).unsqueeze(0).to(device)\n        \n        with torch.no_grad():\n            logits = model(image_tensor)\n            pred_proba = torch.softmax(logits, dim=1)[:, 1]  # Foreground probability\n            pred_proba_np = pred_proba.cpu().numpy()[0]\n        \n        # Resize to original size\n        pred_resized = resize_nearest_neighbor(pred_proba_np, original_h, original_w)\n        \n        # Apply thresholding\n        if use_adaptive_threshold:\n            pred_binary = adaptive_threshold_refinement(\n                pred_resized,\n                base_threshold=0.5,\n                edge_threshold=0.7,\n                edge_width=20\n            )\n        else:\n            pred_binary = (pred_resized > 0.5).astype(np.uint8)\n        \n        predictions_list.append(pred_binary)\n        proba_list.append(pred_resized)\n    \n    predictions_volume = np.stack(predictions_list, axis=0)\n    proba_volume = np.stack(proba_list, axis=0)\n    \n    return predictions_volume, proba_volume\n\n\ndef visualize_prediction(\n    image: np.ndarray,\n    pred: np.ndarray,\n    target: np.ndarray = None,\n    figsize: Tuple[int, int] = (18, 6)\n):\n    \"\"\"Visualize prediction vs ground truth.\"\"\"\n    n_plots = 3 if target is not None else 2\n    fig, axes = plt.subplots(1, n_plots, figsize=figsize)\n    \n    if len(image.shape) == 3:\n        mid_channel = image.shape[0] // 2\n        axes[0].imshow(image[mid_channel], cmap='gray')\n    else:\n        axes[0].imshow(image, cmap='gray')\n    axes[0].set_title('Input Image')\n    axes[0].axis('off')\n    \n    axes[1].imshow(pred, cmap='hot', vmin=0, vmax=1)\n    axes[1].set_title('Prediction')\n    axes[1].axis('off')\n    \n    if target is not None:\n        axes[2].imshow(target, cmap='hot', vmin=0, vmax=1)\n        axes[2].set_title('Ground Truth')\n        axes[2].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\nprint(\"Inference functions defined!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.939214Z","iopub.status.idle":"2026-01-08T15:30:28.939662Z","shell.execute_reply.started":"2026-01-08T15:30:28.939512Z","shell.execute_reply":"2026-01-08T15:30:28.939523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 19: Run Inference on Test Data\n# =============================================================================\n\ntest_volume_files = sorted([f.name for f in TEST_IMAGES_DIR.glob(\"*.tif\")])\n\nif not test_volume_files:\n    print(f\"No TIFF files found in {TEST_IMAGES_DIR}\")\nelse:\n    print(f\"Found {len(test_volume_files)} test volumes\")\n\ntest_filenames = []\n\nfor test_volume_file in tqdm(test_volume_files, desc=\"Processing test volumes\"):\n    volume_path = TEST_IMAGES_DIR / test_volume_file\n    print(f\"\\nProcessing: {volume_path.name}\")\n    \n    # Load volume\n    volume = tifffile.imread(str(volume_path))\n    print(f\"  Volume shape: {volume.shape}\")\n    \n    # Run inference\n    predictions_volume, proba_volume = run_inference(\n        model=model,\n        volume=volume,\n        slice_axis=SLICE_AXIS,\n        transform=data_module.val_transform,\n        device=DEVICE,\n        use_adaptive_threshold=True,\n    )\n    \n    print(f\"  Raw predictions shape: {predictions_volume.shape}\")\n    print(f\"  Unique values before post-processing: {np.unique(predictions_volume)}\")\n    \n    # Apply topology-aware post-processing\n    print(\"  Applying topology-aware post-processing...\")\n    predictions_volume = topology_aware_postprocess(\n        predictions_volume,\n        min_component_size=MIN_COMPONENT_SIZE,\n        max_hole_size=MAX_HOLE_SIZE,\n        closing_radius=CLOSING_RADIUS,\n        smoothing_sigma=SMOOTHING_SIGMA,\n        thin_bridge_threshold=THIN_BRIDGE_THRESHOLD,\n        device=DEVICE,\n    )\n    \n    # Separate merged wraps\n    print(\"  Separating merged wraps...\")\n    try:\n        predictions_volume = detect_and_separate_wraps(\n            predictions_volume,\n            min_separation=10,\n            n_iterations=2\n        )\n    except Exception as e:\n        print(f\"  Warning: Wrap separation failed: {e}\")\n    \n    print(f\"  Final predictions shape: {predictions_volume.shape}\")\n    print(f\"  Unique values after post-processing: {np.unique(predictions_volume)}\")\n    \n    # Save predictions\n    test_filename = f\"{Path(volume_path).stem}.tif\"\n    tifffile.imwrite(test_filename, predictions_volume.astype(np.uint8))\n    test_filenames.append(test_filename)\n    print(f\"  Saved: {test_filename}\")\n\nprint(f\"\\nProcessed {len(test_filenames)} test volumes\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.940755Z","iopub.status.idle":"2026-01-08T15:30:28.940999Z","shell.execute_reply.started":"2026-01-08T15:30:28.940888Z","shell.execute_reply":"2026-01-08T15:30:28.940899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 20: Visualize Sample Predictions\n# =============================================================================\n\nif test_filenames and len(test_volume_files) > 0:\n    # Load first test volume and its predictions for visualization\n    sample_volume_path = TEST_IMAGES_DIR / test_volume_files[0]\n    sample_volume = tifffile.imread(str(sample_volume_path))\n    sample_predictions = tifffile.imread(test_filenames[0])\n    \n    # Visualize a few slices\n    num_slices = sample_volume.shape[0]\n    slice_indices = [0, num_slices // 4, num_slices // 2, 3 * num_slices // 4, num_slices - 1]\n    \n    for idx in slice_indices:\n        if idx < num_slices:\n            original_slice = VesuviusSliceDataset._extract_slice_from_volume(\n                sample_volume, idx, SLICE_AXIS\n            ).squeeze(0)\n            pred_slice = sample_predictions[idx]\n            \n            print(f\"Slice {idx}/{num_slices}:\")\n            visualize_prediction(original_slice, pred_slice)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.942337Z","iopub.status.idle":"2026-01-08T15:30:28.942759Z","shell.execute_reply.started":"2026-01-08T15:30:28.942557Z","shell.execute_reply":"2026-01-08T15:30:28.942573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 21: Create Submission\n# =============================================================================\n\nsubmission_filename = 'submission.zip'\n\nprint(f\"Creating submission file: {submission_filename}\")\n\nwith zipfile.ZipFile(submission_filename, 'w', zipfile.ZIP_DEFLATED) as zipf:\n    for filename in tqdm(test_filenames, desc=\"Zipping files\"):\n        if not os.path.exists(filename):\n            print(f\"Warning: Missing file {filename}\")\n            continue\n        zipf.write(filename)\n        print(f\"  Added: {filename}\")\n\n# Clean up individual files\nfor filename in test_filenames:\n    if os.path.exists(filename):\n        os.remove(filename)\n\nprint(f\"\\nSubmission file created: {submission_filename}\")\nprint(f\"File size: {os.path.getsize(submission_filename) / 1024 / 1024:.2f} MB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.944015Z","iopub.status.idle":"2026-01-08T15:30:28.944332Z","shell.execute_reply.started":"2026-01-08T15:30:28.944186Z","shell.execute_reply":"2026-01-08T15:30:28.944202Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 22: Summary and Final Notes\n# =============================================================================\n\nprint(\"=\"*60)\nprint(\"TRAINING SUMMARY\")\nprint(\"=\"*60)\n\nprint(\"\\n📊 Competition Metrics Optimization:\")\nprint(\"  - SurfaceDice@τ: Surface Loss + Boundary Loss\")\nprint(\"  - VOI_score: Connectivity Loss + Wrap Separation\")\nprint(\"  - TopoScore: Hole Filling + Bridge Removal\")\n\nprint(\"\\n🔧 Loss Function Weights:\")\nprint(f\"  - DiceCE Loss: {LAMBDA_DICE_CE}\")\nprint(f\"  - Surface Loss: {LAMBDA_SURFACE}\")\nprint(f\"  - Boundary Loss: {LAMBDA_BOUNDARY}\")\nprint(f\"  - Connectivity Loss: {LAMBDA_CONNECTIVITY}\")\n\nprint(\"\\n📈 Post-Processing Steps:\")\nprint(f\"  - Gaussian smoothing (σ={SMOOTHING_SIGMA})\")\nprint(f\"  - Morphological closing (radius={CLOSING_RADIUS})\")\nprint(f\"  - Remove small holes (max_size={MAX_HOLE_SIZE})\")\nprint(f\"  - Remove thin bridges (threshold={THIN_BRIDGE_THRESHOLD})\")\nprint(f\"  - Remove small components (min_size={MIN_COMPONENT_SIZE})\")\nprint(f\"  - Separate merged wraps\")\n\nprint(\"\\n💡 Tips for Further Improvement:\")\nprint(\"  1. Increase init_filters in SegResNet (32 → 48 or 64)\")\nprint(\"  2. Train for more epochs with early stopping\")\nprint(\"  3. Use test-time augmentation (TTA)\")\nprint(\"  4. Ensemble multiple models\")\nprint(\"  5. Fine-tune loss weights based on validation metrics\")\nprint(\"  6. Add more aggressive wrap separation\")\n\nprint(\"\\n✅ Done!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-08T15:30:28.945506Z","iopub.status.idle":"2026-01-08T15:30:28.945821Z","shell.execute_reply.started":"2026-01-08T15:30:28.945698Z","shell.execute_reply":"2026-01-08T15:30:28.945714Z"}},"outputs":[],"execution_count":null}]}