{"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":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14286216,"datasetId":8787846,"databundleVersionId":15088516},{"sourceType":"modelInstanceVersion","sourceId":667609,"databundleVersionId":14706139,"modelInstanceId":499221},{"sourceType":"kernelVersion","sourceId":284794863}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🏛️ Vesuvius Challenge - Surface Detection Baseline\n\n### Simple 3D Segmentation Approach\n\nThis notebook demonstrates a **full 3D segmentation pipeline running entirely on GPU**, including data loading and augmentations. Key features:\n\n- **GPU-Accelerated Augmentations**: Leveraging MONAI transforms to perform data augmentations directly on the GPU, significantly speeding up the training process by minimizing data transfers.\n- **Faster I/O**: Utilizes pre-saved `.npy` and `.npz` volumes for both images and labels, which are much faster to load than `.tif` files, further enhancing data throughput.\n- **Optimized Data Module**: A custom `SurfaceDataset3D` and `SurfaceDataModule` handle variable-sized 3D volumes efficiently, resizing them on the fly to `MODEL_INPUT_SIZE` during GPU-accelerated augmentation.\n- **Robust Model Training**: Employs a MONAI UNet for 3D segmentation with PyTorch Lightning for a clean, reproducible training workflow, including metrics like Dice and IoU.\n- **Simple Baseline**: This is a straightforward baseline implementation, offering substantial room for improvement through more advanced architectures, data augmentation strategies, loss functions, and ensemble methods.\n\n### 📊 Data Structure\n\n```\nvesuvius-challenge-surface-detection/\n├── train_images/       # 3D TIFF volumes\n│   ├── 1004283650.tif\n│   └── ...\n├── train_labels/       # 3D mask annotations (same filenames)\n│   ├── 1004283650.tif\n│   └── ...\n└── test_images/        # Test volumes (no labels)\n    └── ...\n```","metadata":{}},{"cell_type":"code","source":"!pip install -q -U \"numpy<2\" \"pandas<2.2.0\" matplotlib pytorch_lightning monai albumentations imagecodecs scikit-learn scikit-image\n!pip uninstall -q -y tensorflow\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:07.051246Z","iopub.execute_input":"2025-12-30T14:13:07.051427Z","iopub.status.idle":"2025-12-30T14:13:32.282203Z","shell.execute_reply.started":"2025-12-30T14:13:07.051404Z","shell.execute_reply":"2025-12-30T14:13:32.281457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport warnings\nfrom pathlib import Path\nfrom typing import Tuple, Optional, Dict, List, Callable\n\nimport imagecodecs\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport pytorch_lightning as pl\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings(\"ignore\")\n\n# Paths\nDATA_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nCHECKPOINT_DIR = \"/kaggle/input/vesuvius-surface-baseline-v3-fixed\"\n# TRAIN_IMAGES_DIR = DATA_DIR / \"train_images\"\n# TRAIN_LABELS_DIR = DATA_DIR / \"train_labels\"\nTRAIN_IMAGES_DIR = Path(\"/kaggle/input/vesuvius-surface-npz/train_images\")\nTRAIN_LABELS_DIR = Path(\"/kaggle/input/vesuvius-surface-npz/train_labels\")\nTEST_IMAGES_DIR = DATA_DIR / \"test_images\"\nOUTPUT_DIR = Path(\".\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:32.283733Z","iopub.execute_input":"2025-12-30T14:13:32.284478Z","iopub.status.idle":"2025-12-30T14:13:44.499639Z","shell.execute_reply.started":"2025-12-30T14:13:32.284451Z","shell.execute_reply":"2025-12-30T14:13:44.499054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Device\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Model architecture\n# MODEL_INPUT_SIZE = (224, 224, 224)  # (depth, height, width) - resize volumes to this\nMODEL_INPUT_SIZE = (160, 160, 160)  # (depth, height, width) - resize volumes to this\nIN_CHANNELS = 1  # grayscale\nOUT_CHANNELS = 2  # background + papyrus (ignore class 2)\n\n# Training\nBATCH_SIZE = 1\nNUM_WORKERS = 2\nMAX_EPOCHS = 20\nLEARNING_RATE = 2e-3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:44.500341Z","iopub.execute_input":"2025-12-30T14:13:44.500823Z","iopub.status.idle":"2025-12-30T14:13:44.538505Z","shell.execute_reply.started":"2025-12-30T14:13:44.5008Z","shell.execute_reply":"2025-12-30T14:13:44.537967Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 📊 Dataset","metadata":{}},{"cell_type":"code","source":"import tifffile\nimport numpy as np\n\nclass SurfaceDataset3D(Dataset):\n    \"\"\"3D Surface Detection Dataset.\n\n    Updated to support volume-based loading.\n    Optimized for faster Torch conversion.\n    Supports .tif, .npy, .npz formats.\n    \"\"\"\n\n    def __init__(\n        self,\n        images_dir: Path,\n        labels_dir: Optional[Path],\n        volume_files: Optional[List[str]] = None,\n        volume_shape: Tuple[int, int, int] = (64, 64, 64), # Kept for compatibility, unused\n    ):\n        super().__init__()\n        self.images_dir = images_dir\n        # Ensure labels_dir is a Path to prevent errors in _load_from_raw\n        self.labels_dir = labels_dir if labels_dir is not None else Path(\"__no_label__\")\n        self.volume_files = volume_files\n        self.volume_shape = volume_shape\n\n        # Validate and populate volume_files\n        self._prepare_volume_files()\n\n    def _prepare_volume_files(self):\n        \"\"\"Validate and index provided volumes, updating self.volume_files.\"\"\"\n        # Determine source of files\n        if self.volume_files is None:\n             print(f\"No volume files specified. Scanning {self.images_dir}...\")\n             # Priority: npy > npz > tif\n             extensions = [\".npy\", \".npz\", \".tif\"]\n             self.volume_files = []\n             # TODO: Add check for duplicates if multiple formats exist for the same volume.\n             # Currently assuming each volume appears only once across these formats.\n             for ext in extensions:\n                 # glob returns full paths, we just want filenames\n                 files = sorted([p.name for p in self.images_dir.glob(f\"*{ext}\")])\n                 self.volume_files.extend(files)\n\n        print(f\"Indexing volumes...\")\n        valid_files = []\n        for filename in self.volume_files:\n            image_path = self.images_dir / filename\n            if not image_path.exists():\n                print(f\"Warning: {image_path} not found, skipping.\")\n                continue\n            valid_files.append(filename)\n\n        self.volume_files = valid_files\n        print(f\"Found {len(self.volume_files)} volumes.\")\n\n    def __len__(self) -> int:\n        return len(self.volume_files)\n\n    def __getitem__(self, idx: int):\n        filename = self.volume_files[idx]\n        # Load raw data -> (D, H, W)\n        image, mask = self._load_from_raw(filename)\n        # Optimization: Convert directly to Tensor to avoid intermediate numpy float64 copies\n        # 1. Convert raw uint8/uint16 -> Tensor\n        # 2. Cast to float16\n        # 3. Scale\n        # 4. Add channel dim\n        image_t = torch.from_numpy(image).half().div_(255.0).unsqueeze(0)\n\n        # Handle Mask\n        if mask is not None:\n             mask_t = torch.from_numpy(mask).long().unsqueeze(0)\n        else:\n             # Return dummy mask for test set (class 2 is ignored in loss)\n             mask_t = torch.full_like(image_t, 2, dtype=torch.long)\n\n        # Return fragment ID (filename without extension) for prediction grouping\n        frag_id = Path(filename).stem\n        return image_t, mask_t, frag_id\n\n    def _load_file(self, path: Path) -> np.ndarray:\n        \"\"\"Helper to load generic file formats.\"\"\"\n        if path.suffix == \".npy\":\n            return np.load(str(path))\n        if path.suffix == \".npz\":\n            data = np.load(str(path))\n            # Return the first array found in the archive\n            return data[list(data.files)[0]]\n        # Original dataset format\n        return tifffile.imread(str(path))\n\n    def _load_from_raw(\n        self,\n        volume_file: str,\n    ) -> Tuple[np.ndarray, Optional[np.ndarray]]:\n        \"\"\"Helper to load image and mask from raw TIFF files.\"\"\"\n        # Load from disk\n        image_path = self.images_dir / volume_file\n        image_volume = self._load_file(image_path)\n\n        label_volume = None\n        label_path = self.labels_dir / volume_file\n        if label_path.exists():\n            label_volume = self._load_file(label_path)\n\n        return image_volume, label_volume","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:44.539186Z","iopub.execute_input":"2025-12-30T14:13:44.539451Z","iopub.status.idle":"2025-12-30T14:13:44.709448Z","shell.execute_reply.started":"2025-12-30T14:13:44.539427Z","shell.execute_reply":"2025-12-30T14:13:44.708654Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 📦 DataModule","metadata":{}},{"cell_type":"code","source":"import random\nfrom torch.utils.data import DataLoader\nfrom monai import transforms as MT\n\ndef custom_collate(batch):\n    \"\"\"Custom collate to handle variable size 3D volumes.\n    Returns a list of items instead of stacking them, allowing GPU resizing later.\n    \"\"\"\n    return batch\n\nclass SurfaceDataModule(pl.LightningDataModule):\n    \"\"\"Lightning DataModule for Surface Detection.\n\n    Handles all data loading, splitting, and dataloader creation.\n    Updated to use MONAI 3D augmentations on GPU with dynamic resizing.\n    \"\"\"\n\n    def __init__(\n        self,\n        train_images_dir: Path,\n        train_labels_dir: Path,\n        volume_shape: Tuple[int, int, int] = (64, 64, 64),\n        val_split: float = 0.2,\n        batch_size: int = BATCH_SIZE,\n        num_workers: int = NUM_WORKERS\n    ):\n        super().__init__()\n        self.train_images_dir = train_images_dir\n        self.train_labels_dir = train_labels_dir\n        self.volume_shape = volume_shape\n        self.val_split = val_split\n        self.batch_size = batch_size\n        self.num_workers = num_workers\n\n        # Will be set in setup()\n        self.train_dataset = None\n        self.val_dataset = None\n        self.test_dataset = None\n\n        # Define GPU-based augmentations using MONAI\n        # 1. Resize to target shape (trilinear for image, nearest for label)\n        # 2. Apply Augmentations\n        self.gpu_augments = MT.Compose([\n            MT.Resized(keys=[\"image\", \"label\"], spatial_size=self.volume_shape, mode=[\"trilinear\", \"nearest\"]),\n            MT.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=0),\n            MT.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=1),\n            MT.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=2),\n            MT.RandRotated(keys=[\"image\", \"label\"], range_x=0.1, range_y=0.1, range_z=0.1, prob=0.3, keep_size=True, mode=[\"bilinear\", \"nearest\"]),\n            MT.RandShiftIntensityd(keys=[\"image\"], offsets=0.1, prob=0.5),\n            MT.RandGaussianNoised(keys=[\"image\"], prob=0.3, mean=0.0, std=0.01),\n        ])\n        # Validation transforms: Just Resize (for image AND label)\n        self.val_augments = MT.Compose([\n            MT.Resized(keys=[\"image\", \"label\"], spatial_size=self.volume_shape, mode=[\"trilinear\", \"nearest\"])\n        ])\n        # Validation transforms for Image ONLY (for test set where labels are None)\n        self.val_image_augments = MT.Compose([\n            MT.Resized(keys=[\"image\"], spatial_size=self.volume_shape, mode=[\"trilinear\"])\n        ])\n\n    def setup(self, stage: Optional[str] = None):\n        \"\"\"Setup datasets for different stages.\"\"\"\n        print(f\"\\nSetting up training data...\")\n\n        # Get all available training files (scanning for npy, npz, tif)\n        extensions = [\".npy\", \".npz\", \".tif\"]\n        all_files = []\n        for ext in extensions:\n             files = sorted([f.name for f in self.train_images_dir.glob(f\"*{ext}\")])\n             all_files.extend(files)\n\n        if not all_files:\n            raise RuntimeError(f\"No volume files found in {self.train_images_dir}. Check your data path.\")\n\n        # Shuffle for random split (deterministic with seed)\n        random.seed(42)\n        random.shuffle(all_files)\n        # Split into train/val\n        split_idx = int(len(all_files) * (1 - self.val_split))\n        train_files = all_files[:split_idx]\n        val_files = all_files[split_idx:]\n\n        print(f\"Total files: {len(all_files)}\")\n        print(f\"Train files: {len(train_files)}\")\n        print(f\"Val files: {len(val_files)}\")\n\n        # Create train dataset\n        self.train_dataset = SurfaceDataset3D(\n            images_dir=self.train_images_dir,\n            labels_dir=self.train_labels_dir,\n            volume_files=train_files,\n            volume_shape=self.volume_shape\n        )\n        # Create validation dataset\n        self.val_dataset = SurfaceDataset3D(\n            images_dir=self.train_images_dir,\n            labels_dir=self.train_labels_dir,\n            volume_files=val_files,\n            volume_shape=self.volume_shape\n        )\n\n    def train_dataloader(self) -> DataLoader:\n        \"\"\"Create train dataloader with custom collate for variable sizes.\"\"\"\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            persistent_workers=bool(self.num_workers > 0),\n            collate_fn=custom_collate\n        )\n\n    def val_dataloader(self) -> DataLoader:\n        \"\"\"Create validation dataloader with custom collate for variable sizes.\"\"\"\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            persistent_workers=bool(self.num_workers > 0),\n            collate_fn=custom_collate\n        )\n\n    def on_after_batch_transfer(self, batch, dataloader_idx):\n        \"\"\"Apply MONAI GPU-accelerated 3D augmentations to the batch.\"\"\"\n        # If custom_collate is used, batch is a list of tuples [(x, y, id), ...]\n        if not isinstance(batch, list):\n            return super().on_after_batch_transfer(batch, dataloader_idx)\n\n        x_list, y_list, frag_ids = [], [], []\n        # Determine device to ensure we process on GPU\n        # self.trainer.strategy.root_device is reliable in Lightning\n        device = self.trainer.strategy.root_device if self.trainer else torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        # Select transform\n        transforms = self.gpu_augments if self.trainer.training else self.val_augments\n\n        for item in batch:\n            x, y, frag_id = item\n            # IMPORTANT: Explicitly move to GPU now.\n            # This ensures the Resized and other transforms run on VRAM, avoiding CPU RAM spikes.\n            x = x.to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n\n            data = {\"image\": x, \"label\": y}\n            # Apply transforms (Resize + Augments)\n            data = transforms(data)\n\n            x_list.append(data[\"image\"])\n            y_list.append(data[\"label\"])\n            frag_ids.append(frag_id)\n\n        # Stack into tensors -> (B, C, D, H, W)\n        return torch.stack(x_list), torch.stack(y_list), frag_ids","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:44.711445Z","iopub.execute_input":"2025-12-30T14:13:44.711708Z","iopub.status.idle":"2025-12-30T14:13:54.520101Z","shell.execute_reply.started":"2025-12-30T14:13:44.711691Z","shell.execute_reply":"2025-12-30T14:13:54.519522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create DataModule\ndatamodule = SurfaceDataModule(\n    train_images_dir=TRAIN_IMAGES_DIR,\n    train_labels_dir=TRAIN_LABELS_DIR,\n    volume_shape=MODEL_INPUT_SIZE,\n)\ndatamodule.setup()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:54.520842Z","iopub.execute_input":"2025-12-30T14:13:54.521463Z","iopub.status.idle":"2025-12-30T14:13:56.356302Z","shell.execute_reply.started":"2025-12-30T14:13:54.52144Z","shell.execute_reply":"2025-12-30T14:13:56.355726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Get a batch\ntrain_loader = datamodule.train_dataloader()\nbatch = next(iter(train_loader))\n\n# With custom_collate, batch is a list: [(img, mask, id), ...]\n# We extract the first sample\nraw_img, raw_mask, frag_id = batch[0]\n\nprint(f\"Raw ID: {frag_id}\")\nprint(f\"Raw Image Shape: {raw_img.shape}\")\n\n# Manually apply the validation transform (Resize) to visualize the model input\n# We use val_augments which only contains Resized\n# Construct dictionary as expected by MONAI transforms\ndata = {\"image\": raw_img, \"label\": raw_mask}\ndata_resized = datamodule.val_augments(data)\n\nimages = data_resized[\"image\"]\nmasks = data_resized[\"label\"]\n\n# Select first channel\nimg = images[0].numpy()  # (D, H, W)\nmsk = masks[0].numpy()   # (D, H, W)\n\nprint(f\"Resized Image Shape: {img.shape}\")\nprint(f\"Mask Shape: {msk.shape}\")\nprint(f\"Image Range: {img.min()} - {img.max()}\")\n\n# Calculate middle indices\nd_mid = img.shape[0] // 2\nh_mid = img.shape[1] // 2\nw_mid = img.shape[2] // 2\n\n# Setup plot: 3 Rows (Axes), 2 Cols (Image, Mask)\nfig, axes = plt.subplots(3, 2, figsize=(10, 15))\n\n# Row 1: Z-axis (Depth/Axial)\naxes[0, 0].imshow(img[d_mid, :, :], cmap='gray')\naxes[0, 0].set_title(f'Axial (Depth={d_mid}) - Image')\naxes[0, 1].imshow(msk[d_mid, :, :], cmap='gray')\naxes[0, 1].set_title(f'Axial (Depth={d_mid}) - Mask')\n\n# Row 2: Y-axis (Height/Coronal)\naxes[1, 0].imshow(img[:, h_mid, :], cmap='gray')\naxes[1, 0].set_title(f'Coronal (Height={h_mid}) - Image')\naxes[1, 1].imshow(msk[:, h_mid, :], cmap='gray')\naxes[1, 1].set_title(f'Coronal (Height={h_mid}) - Mask')\n\n# Row 3: X-axis (Width/Sagittal)\naxes[2, 0].imshow(img[:, :, w_mid], cmap='gray')\naxes[2, 0].set_title(f'Sagittal (Width={w_mid}) - Image')\naxes[2, 1].imshow(msk[:, :, w_mid], cmap='gray')\naxes[2, 1].set_title(f'Sagittal (Width={w_mid}) - Mask')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-30T14:13:56.357116Z","iopub.execute_input":"2025-12-30T14:13:56.357424Z","iopub.status.idle":"2025-12-30T14:14:00.507238Z","shell.execute_reply.started":"2025-12-30T14:13:56.357398Z","shell.execute_reply":"2025-12-30T14:14:00.506021Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧠 Model","metadata":{}},{"cell_type":"code","source":"%%writefile surface_model.py\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport pytorch_lightning as pl\nfrom typing import Tuple, Dict\nfrom monai.losses import DiceCELoss, TverskyLoss\n\n\nclass SurfaceSegmentation3D(pl.LightningModule):\n    \"\"\"3D Surface Segmentation using a custom network.\n\n    Key Design Choices:\n    - **Loss**: Combined DiceCELoss + TverskyLoss to handle structural imbalance.\n    - **Metrics**: Manual computation of Dice and IoU ignoring class 2.\n    \"\"\"\n\n    def __init__(\n        self,\n        net: nn.Module,\n        out_channels: int = 2,\n        spatial_dims: int = 3,\n        learning_rate: float = 1e-3,\n        weight_decay: float = 1e-4,\n        ignore_index_val: int = 2\n    ):\n        super().__init__()\n        self.save_hyperparameters(ignore=[\"net\"])\n        self.net_module = net\n        self.learning_rate = learning_rate\n        self.weight_decay = weight_decay\n        self.ignore_index_val = ignore_index_val\n\n        # Loss function configuration\n        # TverskyLoss with alpha=0.7 emphasizes minimizing False Negatives (Recall)\n        self.criterion_tversky = TverskyLoss(\n            softmax=True,\n            to_onehot_y=False,\n            include_background=True,\n            alpha=0.7,\n            beta=0.3\n        )\n        # DiceCELoss combines Dice Loss and Cross Entropy Loss\n        self.criterion_dice_ce = DiceCELoss(\n            softmax=True,\n            to_onehot_y=False,\n            include_background=True,\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return self.net_module(x)\n\n    def _compute_loss(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n        \"\"\"Compute loss excluding class 2 (unlabeled). Optimized for GPU.\"\"\"\n        # targets shape: (B, 1, D, H, W)\n        mask = (targets != self.ignore_index_val)\n        # Prepare targets for One-Hot Encoding (replace ignore index with 0 temporary)\n        targets_sq = targets.squeeze(1)\n        targets_clean = torch.where(mask.squeeze(1), targets_sq, torch.tensor(0, device=targets.device))\n        # One-Hot Encode\n        targets_onehot = torch.nn.functional.one_hot(\n            targets_clean.long(),\n            num_classes=self.hparams.out_channels\n        ).float()\n        if self.hparams.spatial_dims == 3:\n            targets_onehot = targets_onehot.permute(0, 4, 1, 2, 3)\n        else:\n            targets_onehot = targets_onehot.permute(0, 3, 1, 2)\n        # Mask One-Hot Targets\n        targets_masked_ohe = targets_onehot * mask.half()\n\n        # Compute both losses and sum them\n        loss_tversky = self.criterion_tversky(logits, targets_masked_ohe)\n        loss_dice_ce = self.criterion_dice_ce(logits, targets_masked_ohe)\n\n        return loss_tversky + loss_dice_ce\n\n    def _compute_metrics(self, preds_logits: torch.Tensor, targets_class_indices: torch.Tensor) -> dict:\n        preds_proba = torch.softmax(preds_logits, dim=1)\n        preds_hard = torch.argmax(preds_proba, dim=1, keepdim=True)\n        valid_mask = (targets_class_indices != self.ignore_index_val).float() # (B, 1, D, H, W)\n        num_classes = preds_logits.shape[1] # This will be 2 (background, foreground)\n        dice_scores_per_class = []\n        iou_scores_per_class = []\n\n        for i in range(num_classes):\n            pred_class_i = (preds_hard == i).float() # (B, 1, D, H, W)\n            target_class_i = (targets_class_indices == i).float() # (B, 1, D, H, W)\n\n            pred_class_i_valid = pred_class_i * valid_mask\n            target_class_i_valid = target_class_i * valid_mask\n\n            intersection = (pred_class_i_valid * target_class_i_valid).sum()\n            union_sum_dice = pred_class_i_valid.sum() + target_class_i_valid.sum()\n            union_sum_iou = pred_class_i_valid.sum() + target_class_i_valid.sum() - intersection\n            dice = (2 * intersection + 1e-8) / (union_sum_dice + 1e-8)\n            iou = (intersection + 1e-8) / (union_sum_iou + 1e-8)\n            dice_scores_per_class.append(dice)\n            iou_scores_per_class.append(iou)\n\n        mean_dice = torch.mean(torch.stack(dice_scores_per_class))\n        mean_iou = torch.mean(torch.stack(iou_scores_per_class))\n        return {\"dice\": mean_dice, \"iou\": mean_iou}\n\n    def training_step(self, batch: Tuple, batch_idx: int) -> torch.Tensor:\n        inputs, targets, _ = batch\n        logits = self(inputs)\n        loss = self._compute_loss(logits, targets)\n\n        metrics = self._compute_metrics(logits, targets)\n\n        self.log(\"train_loss\", loss, on_step=True, on_epoch=True, prog_bar=True)\n        self.log(\"train_dice\", metrics[\"dice\"], on_step=True, on_epoch=True, prog_bar=True)\n        self.log(\"train_iou\", metrics[\"iou\"], on_step=True, on_epoch=True, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch: Tuple, batch_idx: int) -> torch.Tensor:\n        inputs, targets, _ = batch\n        logits = self(inputs)\n        loss = self._compute_loss(logits, targets)\n\n        metrics = self._compute_metrics(logits, targets)\n\n        self.log(\"val_loss\", loss, on_step=False, on_epoch=True, prog_bar=True)\n        self.log(\"val_dice\", metrics[\"dice\"], on_step=False, on_epoch=True, prog_bar=True)\n        self.log(\"val_iou\", metrics[\"iou\"], on_step=False, on_epoch=True, prog_bar=True)\n        return loss\n\n    def configure_optimizers(self):\n        optimizer = optim.AdamW(self.parameters(), lr=self.learning_rate, weight_decay=self.weight_decay)\n        # Cosine Annealing Scheduler\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer,\n            T_max=self.trainer.max_epochs if self.trainer else MAX_EPOCHS,\n            eta_min=1e-6\n        )\n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\n                \"scheduler\": scheduler,\n                \"interval\": \"epoch\"\n            }\n        }\n\n    def predict_step(self, batch: Tuple, batch_idx: int) -> Dict:\n        inputs, _, frag_id = batch\n        logits = self(inputs)\n        probs = torch.softmax(logits, dim=1)\n        pred_class = torch.argmax(probs, dim=1)\n        return {\"prediction\": pred_class, \"fragment_id\": frag_id}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:14:00.508637Z","iopub.execute_input":"2025-12-30T14:14:00.508939Z","iopub.status.idle":"2025-12-30T14:14:00.518183Z","shell.execute_reply.started":"2025-12-30T14:14:00.508903Z","shell.execute_reply":"2025-12-30T14:14:00.517481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from surface_model import SurfaceSegmentation3D\nfrom monai.networks.nets import SegResNet, SwinUNETR\n\n# Initialize model\n# Note: EfficientNet is typically 2D. For 3D volumes, SegResNet is the optimized, efficient standard.\nnet = SegResNet(\n    spatial_dims=3,\n    in_channels=IN_CHANNELS,\n    out_channels=OUT_CHANNELS,\n    init_filters=16,\n    dropout_prob=0.2\n)\n\n# Initialize SwinUNETR model\n# net = SwinUNETR(\n#     in_channels=1,\n#     out_channels=2,\n#     feature_size=48,\n#     use_v2=True,\n#     drop_rate=0.2,\n#     attn_drop_rate=0.2,\n#     dropout_path_rate=0.2,\n# )\n\nnet_name = net.__class__.__name__\nmodel = SurfaceSegmentation3D(net=net)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:14:00.519074Z","iopub.execute_input":"2025-12-30T14:14:00.519311Z","iopub.status.idle":"2025-12-30T14:14:00.587209Z","shell.execute_reply.started":"2025-12-30T14:14:00.519293Z","shell.execute_reply":"2025-12-30T14:14:00.586416Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🏋️ Training","metadata":{}},{"cell_type":"code","source":"import re\nfrom pathlib import Path\nfrom typing import List, Union, Tuple, Optional\n\ndef get_best_checkpoint(\n    checkpoint_dirs: Union[str, Path, List[Union[str, Path]]],\n    name: str = \"\",\n) -> Tuple[str, float]:\n    \"\"\"Finds the checkpoint with the highest val_dice score across multiple directories.\"\"\"\n    # Normalize input to a list of Paths\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    if not checkpoint_dirs:\n        print(\"No valid folder provided.\")\n        return None, None\n    \n    checkpoints = []\n    # Regex for val_dice\n    pattern = re.compile(r\"val_dice=?([0-9]+\\.[0-9]+)\")\n    # Iterate over all files in all valid directories\n    for path in [f for d in checkpoint_dirs for f in Path(d).glob(f\"{name}*.ckpt\")]:\n        match = pattern.search(path.name)\n        if not match:\n            continue\n        checkpoints.append((float(match.group(1)), str(path)))\n\n    if not checkpoints:\n        print(\"No valid checkpoints found.\")\n        return None, None\n\n    # Sort by score descending so the best is first\n    checkpoints.sort(key=lambda x: x[0], reverse=True)\n    best_score, best_path = checkpoints[0]\n    print(f\"Found {len(checkpoints)} checkpoints.\")\n    print(f\"Best  (Score={best_score}): {Path(best_path)}\")\n    return best_path, best_score\n\nckpt_path, ckpt_score = get_best_checkpoint(\n    [OUTPUT_DIR, CHECKPOINT_DIR], name=net_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:14:00.588104Z","iopub.execute_input":"2025-12-30T14:14:00.588316Z","iopub.status.idle":"2025-12-30T14:14:00.597103Z","shell.execute_reply.started":"2025-12-30T14:14:00.588299Z","shell.execute_reply":"2025-12-30T14:14:00.596427Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.utilities.exceptions import MisconfigurationException\n\n# Callbacks\ncheckpoint_callback = ModelCheckpoint(\n    dirpath=OUTPUT_DIR,\n    filename=net_name + \"-{epoch:02d}-{val_dice:.4f}\",\n    monitor=\"val_dice\",\n    mode=\"max\",\n    save_top_k=3,\n    verbose=True\n)\n\nearly_stop_callback = EarlyStopping(\n    monitor=\"val_dice\",\n    patience=10,\n    mode=\"max\",\n    verbose=True\n)\n\nlr_monitor = LearningRateMonitor(logging_interval=\"epoch\")\ncsv_logger = CSVLogger(save_dir=OUTPUT_DIR)\n\n# Trainer\ntrainer = pl.Trainer(\n    max_epochs=MAX_EPOCHS,\n    accelerator=\"auto\",\n    devices=\"auto\",\n    logger=csv_logger,\n    callbacks=[checkpoint_callback, early_stop_callback, lr_monitor],\n    precision=\"16-mixed\",\n    log_every_n_steps=1,\n    enable_progress_bar=True,\n    accumulate_grad_batches=18,\n    gradient_clip_val=1.0, # Clips gradient norm to 1.0 to prevent exploding gradients\n)\n\n# Train\ntry:\n    trainer.fit(model, datamodule=datamodule, ckpt_path=ckpt_path)\nexcept MisconfigurationException as ex:\n    print(ex)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T14:14:00.598667Z","iopub.execute_input":"2025-12-30T14:14:00.598878Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns\nfrom IPython.display import display # Ensure display is available\nsns.set()\n\n# Read the metrics.csv using the trainer's logger directory\n# We need to find the latest version directory within the logger's save_dir\nlog_base_dir = Path(trainer.logger.save_dir) / 'lightning_logs'\n\n# Get the latest version directory\nmetrics_path = log_base_dir / f\"version_{trainer.logger._version}\" / 'metrics.csv'\nprint(f\"Loading metrics from: {metrics_path}\")\n\nif metrics_path.exists():\n    metrics = pd.read_csv(metrics_path)\n    # Remove any columns that are entirely NaN (e.g., from different logging frequencies)\n    display(metrics.dropna(axis=1, how=\"all\").head())\n    # Fill any NaN values by propagating the last valid observation forward (useful for sparse logging)\n    metrics.ffill(inplace=True)\n    # Melt the DataFrame to long-form for plotting\n    # We assume 'epoch' is a reliable identifier for x-axis\n    metrics_melted = metrics.reset_index().melt(\n        id_vars='epoch', var_name='metric', value_name='value')\n    # Define metric groups based on available metrics from VesuviusSegmentationModel\n    metric_groups = {\n        'Loss': [c for c in metrics.columns if '_loss' in c],\n        'Dice Score': [c for c in metrics.columns if '_dice' in c],\n        'IoU Score': [c for c in metrics.columns if '_iou' in c],\n    }\n    # Plot metrics for each group in a separate chart\n    for title, metric_list in metric_groups.items():\n        # Filter melted DataFrame for the current group\n        group_metrics = metrics_melted[metrics_melted['metric'].isin(metric_list)]\n        plt.figure(figsize=(10, 5))\n        sns.lineplot(data=group_metrics, x='epoch', y='value', hue='metric')\n        plt.title(f'{title} over Epochs', fontsize=14, fontweight='bold')\n        plt.xlabel('Epoch', fontsize=12)\n        plt.ylabel(title, fontsize=12)\n        plt.grid(True, alpha=0.3)\n        # Apply log scale only for Loss, not for Dice/IoU which are typically 0-1\n        if title == 'Loss':\n            plt.yscale('log')\n        plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔮 Inference","metadata":{}},{"cell_type":"code","source":"# Get test dataset directly\ntest_dataset = datamodule.test_dataset\ntest_files = sorted([f.name for f in TEST_IMAGES_DIR.glob(\"*.tif\")])\nprint(f\"Test files: {test_files}\")\ntest_dataset = SurfaceDataset3D(\n    images_dir=TEST_IMAGES_DIR,\n    labels_dir=None,\n    volume_files=test_files,\n    volume_shape=datamodule.volume_shape\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best model\nbest_checkpoint_path, _ = get_best_checkpoint(\n    [OUTPUT_DIR, CHECKPOINT_DIR], name=net_name)\n\nassert best_checkpoint_path, \"No checkpoint found in trainer, using current model state.\"\nprint(f\"Loading best checkpoint: {best_checkpoint_path}\")\n# We must pass the 'net' argument because it was ignored in save_hyperparameters\nmodel = SurfaceSegmentation3D.load_from_checkpoint(best_checkpoint_path, net=net)\n\nmodel.eval()\nmodel.to(DEVICE)","metadata":{"trusted":true,"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from skimage.morphology import remove_small_objects, ball\nimport torch.nn.functional as F\n\ndef get_spherical_kernel(radius):\n    \"\"\"Generates a spherical kernel (structuring element) for 3D morphological operations.\"\"\"\n    # Generate boolean ball on CPU\n    kernel_np = ball(radius)\n    # Convert to float tensor: (1, 1, D, H, W)\n    kernel = torch.from_numpy(kernel_np.astype(np.float32)).unsqueeze(0).unsqueeze(0)\n    return kernel\n\ndef post_process_3d(\n    volume: np.ndarray,\n    min_size: int = 1000,\n    closing_radius: int = 5,\n    device: str = DEVICE,\n) -> np.ndarray:\n    \"\"\"Applies 3D morphological operations to clean up a segmentation volume using a spherical element.\n\n    Steps:\n    1. Performs morphological closing (Dilation -> Erosion) with a spherical kernel.\n    2. Removes small connected components.\n    \"\"\"\n    # Ensure input is boolean\n    binary_vol = volume > 0\n    clean_vol_np = binary_vol\n\n    # 1. Close gaps with Spherical Element (GPU accelerated)\n    if closing_radius > 0:\n        # print(f\"Closing gaps with spherical radius {closing_radius} (GPU accelerated)...\")\n        # Prepare Input: (1, 1, D, H, W)\n        input_tensor = torch.from_numpy(clean_vol_np.astype(np.float32)).unsqueeze(0).unsqueeze(0).to(device)\n        # Prepare Kernel\n        kernel = get_spherical_kernel(closing_radius).to(device)\n        # Dilation: (Input * Kernel) > 0\n        # We use padding=closing_radius to maintain the same spatial dimensions (same as 'same' padding)\n        dilated = (F.conv3d(input_tensor, kernel, padding=closing_radius) > 0).float()\n        # Erosion: 1 - ((1 - Dilated) * Kernel > 0)\n        # This relies on the duality: Erosion(A) = ~Dilation(~A)\n        eroded = 1.0 - (F.conv3d(1.0 - dilated, kernel, padding=closing_radius) > 0).float()\n        # Retrieve result\n        clean_vol_np = eroded.squeeze().cpu().numpy().astype(bool)\n\n    # 2. Remove small objects (CPU-bound)\n    # print(f\"Removing small objects < {min_size} voxels (CPU-bound)...\")\n    clean_vol_np = remove_small_objects(clean_vol_np, min_size=min_size)\n\n    return clean_vol_np.astype(np.uint8)","metadata":{"trusted":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions_tif = []\n\n# Iterate directly over the dataset\nfor image, _, frag_id in tqdm(test_dataset, desc=\"Processing and saving 3D predictions\"):\n    # image is (C, D, H, W)\n\n    # 1. Preprocess: Use val_image_augments for resizing\n    # Move image to device for GPU transform\n    image_on_device = image.to(model.device)\n\n    # Apply the image-only validation transform (which contains resizing)\n    processed_data = datamodule.val_image_augments({\"image\": image_on_device})\n    # Add batch dimension: (1, C, D, H, W)\n    inputs_resized = processed_data[\"image\"].unsqueeze(0)\n\n    # 2. Inference\n    with torch.no_grad():\n        # Construct batch tuple simulating a dataloader batch: (inputs, targets, [frag_id])\n        frag_id_list = [frag_id]\n        pred_dict = model.predict_step((inputs_resized, None, frag_id_list), 0)\n        # Get predicted class directly (already argmaxed in model)\n        # (D, H, W) relative to MODEL_INPUT_SIZE\n        pred_class = pred_dict[\"prediction\"][0]\n\n    # 3. Postprocess: Resize back to Original Shape\n    # Get original shape from file directly\n    image_path = test_dataset.images_dir / f\"{frag_id}.tif\"\n    with tifffile.TiffFile(str(image_path)) as tif:\n        original_shape = tif.series[0].shape # (D, H, W)\n    # Prediction is binary class index (0 or 1). Convert to float for interpolation.\n    pred_binary = pred_class.float()\n    # Resize to original 3D dimensions using Nearest Neighbor to preserve binary labels\n    # F.interpolate expects (B, C, D, H, W)\n    pred_input_tensor = pred_binary.unsqueeze(0).unsqueeze(0)\n    pred_restored = F.interpolate(\n        pred_input_tensor,\n        size=original_shape,\n        mode='nearest'\n    ).squeeze(0).squeeze(0) # Back to (D, H, W)\n    # Convert to uint8 (0, 1)\n    pred_saved = pred_restored.byte().cpu().numpy()\n    # pred_saved = post_process_3d(pred_saved, min_size=20*20*35)\n\n    # Save to TIFF\n    prediction_tif_name = f\"{frag_id}.tif\"\n    save_path = OUTPUT_DIR / prediction_tif_name\n    tifffile.imwrite(str(save_path), pred_saved)\n    predictions_tif.append(prediction_tif_name)\n    # print(f\"  {frag_id}: saved {save_path} with shape {pred_saved.shape}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_three_axis_cuts(image_vol_path, mask_vol_path):\n    \"\"\"Plots the middle slice of the XY, XZ, and YZ planes for both the image volume and the predicted mask.\"\"\"\n    print(f\"Visualizing cuts for: {os.path.basename(image_vol_path)}\")\n    # Load volumes\n    image_vol = tifffile.imread(image_vol_path)\n    mask_vol = tifffile.imread(mask_vol_path)\n    \n    # Get dimensions\n    d, h, w = image_vol.shape\n    z_mid, y_mid, x_mid = d // 2, h // 2, w // 2\n    \n    # Extract slices\n    slices = {\n        'XY Plane (Z-axis)': (image_vol[z_mid, :, :], mask_vol[z_mid, :, :]),\n        'XZ Plane (Y-axis)': (image_vol[:, y_mid, :], mask_vol[:, y_mid, :]),\n        'YZ Plane (X-axis)': (image_vol[:, :, x_mid], mask_vol[:, :, x_mid])\n    }\n    \n    fig, axes = plt.subplots(3, 2, figsize=(12, 15))\n    for i, (plane_name, (img_slice, mask_slice)) in enumerate(slices.items()):\n        # Image Volume\n        axes[i, 0].imshow(img_slice, cmap='gray')\n        axes[i, 0].set_title(f\"{plane_name} - Image Volume\")\n        axes[i, 0].axis('off')\n        \n        # Mask\n        axes[i, 1].imshow(mask_slice, cmap='gray')\n        axes[i, 1].set_title(f\"{plane_name} - Predicted Mask\")\n        axes[i, 1].axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n\n# Visualize the first processed volume\nif predictions_tif:\n    mask_path = predictions_tif[0]\n    name = os.path.basename(mask_path)\n    image_path = os.path.join(TEST_IMAGES_DIR, name)\n    plot_three_axis_cuts(image_path, mask_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import zipfile\n\nprint(f\"Zipping {len(predictions_tif)} files...\")\nwith zipfile.ZipFile('submission.zip', 'w', zipfile.ZIP_DEFLATED) as zipf:\n    for filename in tqdm(predictions_tif, desc=\"Zipping files\"):\n        if not os.path.exists(filename):\n            print(f\"Missing <> {filename}\")\n            continue\n        # Write to zip\n        zipf.write(filename)\n        # Remove original file to save space\n        os.remove(filename)\n\nprint(\"Submission.zip created successfully.\")","metadata":{"trusted":true,"_kg_hide-output":true},"outputs":[],"execution_count":null}]}