{"metadata":{"kernelspec":{"display_name":"Python 3","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"colab":{"name":"Vesuvius Challenge - Surface Detection ","provenance":[],"gpuType":"T4"},"accelerator":"GPU","widgets":{"application/vnd.jupyter.widget-state+json":{"262c29397f904d7796667567c29e3a1f":{"model_module":"@jupyter-widgets/controls","model_name":"VBoxModel","model_module_version":"1.5.0","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"VBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"VBoxView","box_style":"","children":["IPY_MODEL_bf2ca629e2c84cdfab5360cd2520e7d1","IPY_MODEL_757f782b966b4324abaf47fe6218a569","IPY_MODEL_82c070d1e24f4368a071e56c3bd16432","IPY_MODEL_a9ccb48cf219475b83bc8846d7ff6ab9","IPY_MODEL_27340da362de43d3995f5e3f7c267505"],"layout":"IPY_MODEL_b50eae212c824a91bb9fd12c291f9046"}},"bf2ca629e2c84cdfab5360cd2520e7d1":{"model_module":"@jupyter-widgets/controls","model_name":"HTMLModel","model_module_version":"1.5.0","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_7c6bb367c1904d14ab77e4a351ab21ab","placeholder":"​","style":"IPY_MODEL_9cb2adac060d42c4b0420a4e5810e73c","value":"<center> <img\nsrc=https://www.kaggle.com/static/images/site-logo.png\nalt='Kaggle'> <br> Create an API token from <a\nhref=\"https://www.kaggle.com/settings/account\" target=\"_blank\">your Kaggle\nsettings page</a> and paste it below along with your Kaggle username. <br> </center>"}},"757f782b966b4324abaf47fe6218a569":{"model_module":"@jupyter-widgets/controls","model_name":"TextModel","model_module_version":"1.5.0","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"TextModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"TextView","continuous_update":true,"description":"Username:","description_tooltip":null,"disabled":false,"layout":"IPY_MODEL_df004996c84f4189b0d2ec3cca2f659d","placeholder":"​","style":"IPY_MODEL_97520182dc034394b07b1677ce8cc49d","value":""}},"82c070d1e24f4368a071e56c3bd16432":{"model_module":"@jupyter-widgets/controls","model_name":"PasswordModel","model_module_version":"1.5.0","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"PasswordModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"PasswordView","continuous_update":true,"description":"Token:","description_tooltip":null,"disabled":false,"layout":"IPY_MODEL_3320a2c3c81543ccacb049b2ca45932c","placeholder":"​","style":"IPY_MODEL_4b3da9cb189e4fb5884b93398a77db0f","value":""}},"a9ccb48cf219475b83bc8846d7ff6ab9":{"model_module":"@jupyter-widgets/controls","model_name":"ButtonModel","model_module_version":"1.5.0","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ButtonModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ButtonView","button_style":"","description":"Login","disabled":false,"icon":"","layout":"IPY_MODEL_33d4b9524c3e4047ae707d684bd98d8f","style":"IPY_MODEL_6395612b56dd4bd8a86200a75254a4e0","tooltip":""}},"27340da362de43d3995f5e3f7c267505":{"model_module":"@jupyter-widgets/controls","model_name":"HTMLModel","model_module_version":"1.5.0","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_11dcb5251b3a4ef7b4cd2d15d9eaa261","placeholder":"​","style":"IPY_MODEL_d379d11c5cdd45198dd12aa69aef52fc","value":"\n<b>Thank You</b></center>"}},"b50eae212c824a91bb9fd12c291f9046":{"model_module":"@jupyter-widgets/base","model_name":"LayoutModel","model_module_version":"1.2.0","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":"center","align_self":null,"border":null,"bottom":null,"display":"flex","flex":null,"flex_flow":"column","grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":"50%"}},"7c6bb367c1904d14ab77e4a351ab21ab":{"model_module":"@jupyter-widgets/base","model_name":"LayoutModel","model_module_version":"1.2.0","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"9cb2adac060d42c4b0420a4e5810e73c":{"model_module":"@jupyter-widgets/controls","model_name":"DescriptionStyleModel","model_module_version":"1.5.0","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"df004996c84f4189b0d2ec3cca2f659d":{"model_module":"@jupyter-widgets/base","model_name":"LayoutModel","model_module_version":"1.2.0","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"97520182dc034394b07b1677ce8cc49d":{"model_module":"@jupyter-widgets/controls","model_name":"DescriptionStyleModel","model_module_version":"1.5.0","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"3320a2c3c81543ccacb049b2ca45932c":{"model_module":"@jupyter-widgets/base","model_name":"LayoutModel","model_module_version":"1.2.0","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"4b3da9cb189e4fb5884b93398a77db0f":{"model_module":"@jupyter-widgets/controls","model_name":"DescriptionStyleModel","model_module_version":"1.5.0","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"33d4b9524c3e4047ae707d684bd98d8f":{"model_module":"@jupyter-widgets/base","model_name":"LayoutModel","model_module_version":"1.2.0","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"6395612b56dd4bd8a86200a75254a4e0":{"model_module":"@jupyter-widgets/controls","model_name":"ButtonStyleModel","model_module_version":"1.5.0","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ButtonStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","button_color":null,"font_weight":""}},"11dcb5251b3a4ef7b4cd2d15d9eaa261":{"model_module":"@jupyter-widgets/base","model_name":"LayoutModel","model_module_version":"1.2.0","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"d379d11c5cdd45198dd12aa69aef52fc":{"model_module":"@jupyter-widgets/controls","model_name":"DescriptionStyleModel","model_module_version":"1.5.0","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}}}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"__author__=\"Kushvinth\"","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:29.276548Z","iopub.execute_input":"2025-11-15T08:37:29.277483Z","iopub.status.idle":"2025-11-15T08:37:29.28191Z","shell.execute_reply.started":"2025-11-15T08:37:29.277406Z","shell.execute_reply":"2025-11-15T08:37:29.281041Z"},"trusted":true,"id":"Ql1qD2ErfeiV"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Basic Module Imports","metadata":{"id":"3miZYByKfeiX"}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport random\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom PIL import Image\nimport cv2\nfrom scipy import ndimage\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\nfrom collections import defaultdict\nimport pickle\nimport gc\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import autocast, GradScaler\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n# import segmentation_models_pytorch as smp\n\n# Advanced augmentations\nfrom scipy.ndimage import gaussian_filter, rotate\nfrom skimage.transform import rescale, resize\nfrom skimage.filters import gaussian\nfrom skimage.exposure import equalize_adapthist\n\n# Memory optimization\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cudnn.deterministic = False\n","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:29.349332Z","iopub.execute_input":"2025-11-15T08:37:29.349557Z","iopub.status.idle":"2025-11-15T08:37:29.395654Z","shell.execute_reply.started":"2025-11-15T08:37:29.34954Z","shell.execute_reply":"2025-11-15T08:37:29.394906Z"},"trusted":true,"id":"iAFFQTLnfeia"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n# SAFE TIFF LOADER (Works in Kaggle)\n\n","metadata":{"id":"fkeYe-dWfeib"}},{"cell_type":"code","source":"def load_3d_tiff_optimized(path, z_range=None, normalize=True, cache=True):\n    \"\"\"Advanced TIFF loader with caching and preprocessing\"\"\"\n    if cache and hasattr(load_3d_tiff_optimized, '_cache'):\n        if str(path) in load_3d_tiff_optimized._cache:\n            return load_3d_tiff_optimized._cache[str(path)]\n\n    if not hasattr(load_3d_tiff_optimized, '_cache'):\n        load_3d_tiff_optimized._cache = {}\n\n    img = Image.open(path)\n    slices = []\n    try:\n        i = 0\n        while True:\n            img.seek(i)\n            if z_range is None or (z_range[0] <= i <= z_range[1]):\n                slice_data = np.array(img)\n                if normalize:\n                    # Assuming 16-bit images for this competition, normalize to [0, 1]\n                    slice_data = slice_data.astype(np.float32) / 65535.0\n                slices.append(slice_data)\n            i += 1\n            if z_range and i > z_range[1]:\n                break\n    except EOFError:\n        pass\n\n    volume = np.stack(slices, axis=0)\n\n    if cache and len(load_3d_tiff_optimized._cache) < 10:\n        load_3d_tiff_optimized._cache[str(path)] = volume\n\n    return volume\n\ndef adaptive_histogram_equalization_3d(volume, clip_limit=0.03):\n    \"\"\"Apply CLAHE to 3D volume slice by slice\"\"\"\n    enhanced = np.zeros_like(volume)\n    for i in range(volume.shape[0]):\n        enhanced[i] = equalize_adapthist(volume[i], clip_limit=clip_limit)\n    return enhanced\n\ndef multi_scale_pyramid(volume, scales=[1.0, 0.8, 1.2]):\n    \"\"\"Create multi-scale pyramid for robustness\"\"\"\n    pyramids = []\n    for scale in scales:\n        if scale != 1.0:\n            # Fixed: removed deprecated multichannel parameter\n            scaled = rescale(volume, scale, preserve_range=True, anti_aliasing=True, channel_axis=None)\n            # Crop or pad to original size\n            if scale > 1.0:  # Crop\n                start_z = (scaled.shape[0] - volume.shape[0]) // 2\n                start_y = (scaled.shape[1] - volume.shape[1]) // 2\n                start_x = (scaled.shape[2] - volume.shape[2]) // 2\n                scaled = scaled[start_z:start_z+volume.shape[0],\n                               start_y:start_y+volume.shape[1],\n                               start_x:start_x+volume.shape[2]]\n            else:  # Pad\n                pad_z = (volume.shape[0] - scaled.shape[0]) // 2\n                pad_y = (volume.shape[1] - scaled.shape[1]) // 2\n                pad_x = (volume.shape[2] - scaled.shape[2]) // 2\n                scaled = np.pad(scaled, ((pad_z, volume.shape[0]-scaled.shape[0]-pad_z),\n                                        (pad_y, volume.shape[1]-scaled.shape[1]-pad_y),\n                                        (pad_x, volume.shape[2]-scaled.shape[2]-pad_x)))\n        else:\n            scaled = volume.copy()\n        pyramids.append(scaled)\n    return pyramids","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:29.459838Z","iopub.execute_input":"2025-11-15T08:37:29.461072Z","iopub.status.idle":"2025-11-15T08:37:29.470371Z","shell.execute_reply.started":"2025-11-15T08:37:29.461053Z","shell.execute_reply":"2025-11-15T08:37:29.469826Z"},"trusted":true,"id":"vdpykL_pfeie"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{"id":"aQQyNUKZfeii"}},{"cell_type":"code","source":"class Advanced3DAugmentation:\n    def __init__(self, prob=0.5):\n        self.prob = prob\n\n    def __call__(self, volume, mask=None):\n        if random.random() > self.prob:\n            return volume, mask\n\n        # Random 3D rotation\n        if random.random() < 0.3:\n            angle = random.uniform(-15, 15)\n            axes = random.choice([(0,1), (0,2), (1,2)])\n            volume = rotate(volume, angle, axes=axes, reshape=False, mode='constant')\n            if mask is not None:\n                mask = rotate(mask, angle, axes=axes, reshape=False, mode='constant')\n\n        # Elastic deformation\n        if random.random() < 0.2:\n            sigma = random.uniform(2, 5)\n            volume = gaussian_filter(volume, sigma=sigma)\n\n        # Intensity variations\n        if random.random() < 0.4:\n            # Gamma correction\n            gamma = random.uniform(0.8, 1.2)\n            volume = np.power(volume / volume.max(), gamma) * volume.max()\n\n        # Noise injection\n        if random.random() < 0.3:\n            noise = np.random.normal(0, 0.02 * volume.std(), volume.shape)\n            volume = volume + noise\n\n        return volume, mask\n\nclass MultiScaleVesuviusDataset(Dataset):\n    def __init__(self, img_paths, mask_paths=None, patch_sizes=[(32,64,64), (64,128,128), (96,192,192)],\n                 transforms=None, mode='train', use_pyramid=True, fragment_weights=None):\n        self.img_paths = img_paths\n        self.mask_paths = mask_paths\n        self.patch_sizes = patch_sizes\n        self.transforms = transforms\n        self.mode = mode\n        self.use_pyramid = use_pyramid\n        self.fragment_weights = fragment_weights or [1.0] * len(img_paths)\n\n        # Advanced augmentation\n        self.aug_3d = Advanced3DAugmentation(prob=0.7 if mode=='train' else 0.0)\n\n        # Cache for frequently used fragments\n        self.cache = {}\n\n        # Pre-calculate fragment statistics for adaptive sampling\n        self.fragment_stats = self._calculate_fragment_stats()\n\n    def _calculate_fragment_stats(self):\n        stats = []\n        for i, (img_path, mask_path) in enumerate(zip(self.img_paths, self.mask_paths or [None]*len(self.img_paths))):\n            if mask_path and os.path.exists(mask_path):\n                mask = load_3d_tiff_optimized(mask_path, normalize=False)\n                ink_ratio = mask.sum() / mask.size\n                stats.append({'ink_ratio': ink_ratio, 'priority': 1.0 + ink_ratio})\n            else:\n                stats.append({'ink_ratio': 0.0, 'priority': 1.0})\n        return stats\n\n    def __len__(self):\n        return len(self.img_paths) * 3  # Multiple samples per fragment\n\n    def __getitem__(self, idx):\n        fragment_idx = idx % len(self.img_paths)\n        patch_size_idx = (idx // len(self.img_paths)) % len(self.patch_sizes)\n        patch_size = self.patch_sizes[patch_size_idx]\n\n        # Load or get from cache\n        cache_key = f\"{fragment_idx}_{patch_size}\"\n        if cache_key in self.cache:\n            img, mask = self.cache[cache_key]\n        else:\n            img = load_3d_tiff_optimized(self.img_paths[fragment_idx]).astype(np.float32)\n\n            # Advanced preprocessing\n            if self.use_pyramid:\n                pyramids = multi_scale_pyramid(img)\n                img = pyramids[random.randint(0, len(pyramids)-1)]\n\n            # Adaptive histogram equalization\n            if random.random() < 0.5:\n                img = adaptive_histogram_equalization_3d(img)\n\n            if self.mask_paths and self.mask_paths[fragment_idx]:\n                mask = load_3d_tiff_optimized(self.mask_paths[fragment_idx], normalize=False).astype(np.uint8)\n            else:\n                mask = np.zeros_like(img, dtype=np.uint8)\n\n            # Cache smaller volumes\n            if img.nbytes < 100 * 1024 * 1024:  # < 100MB\n                self.cache[cache_key] = (img.copy(), mask.copy())\n\n        Z, Y, X = img.shape\n        pz, py, px = patch_size\n\n        # Adaptive patch sampling based on ink density\n        if self.mode == 'train' and self.mask_paths and self.mask_paths[fragment_idx]:\n            ink_mask = mask > 0\n            if ink_mask.sum() > 100:  # If there's enough ink\n                # Sample around ink regions 70% of the time\n                if random.random() < 0.7:\n                    ink_coords = np.where(ink_mask)\n                    if len(ink_coords[0]) > 0:\n                        rand_idx = random.randint(0, len(ink_coords[0])-1)\n                        center_z = ink_coords[0][rand_idx]\n                        center_y = ink_coords[1][rand_idx]\n                        center_x = ink_coords[2][rand_idx]\n\n                        z0 = max(0, min(Z-pz, center_z - pz//2))\n                        y0 = max(0, min(Y-py, center_y - py//2))\n                        x0 = max(0, min(X-px, center_x - px//2))\n                    else:\n                        z0 = random.randint(0, max(0, Z-pz))\n                        y0 = random.randint(0, max(0, Y-py))\n                        x0 = random.randint(0, max(0, X-px))\n                else:\n                    z0 = random.randint(0, max(0, Z-pz))\n                    y0 = random.randint(0, max(0, Y-py))\n                    x0 = random.randint(0, max(0, X-px))\n            else:\n                z0 = random.randint(0, max(0, Z-pz))\n                y0 = random.randint(0, max(0, Y-py))\n                x0 = random.randint(0, max(0, X-px))\n        else:\n            z0 = random.randint(0, max(0, Z-pz))\n            y0 = random.randint(0, max(0, Y-py))\n            x0 = random.randint(0, max(0, X-px))\n\n        img_patch = img[z0:z0+pz, y0:y0+py, x0:x0+px].copy()\n        mask_patch = mask[z0:z0+pz, y0:y0+py, x0:x0+px].copy()\n\n        # Apply 3D augmentations\n        img_patch, mask_patch = self.aug_3d(img_patch, mask_patch)\n\n        # Apply 2D augmentations slice by slice\n        if self.transforms:\n            augmented_slices_img = []\n            augmented_slices_mask = []\n            for i in range(img_patch.shape[0]):\n                img_slice = np.clip(img_patch[i], 0, 255).astype(np.uint8)\n                mask_slice = mask_patch[i].astype(np.uint8)\n\n                aug = self.transforms(image=img_slice, mask=mask_slice)\n                augmented_slices_img.append(aug['image'])\n                augmented_slices_mask.append(aug['mask'])\n\n            img_patch = np.stack(augmented_slices_img)\n            mask_patch = np.stack(augmented_slices_mask)\n\n        # Normalize\n        img_patch = img_patch.astype(np.float32)\n        if img_patch.max() > 1.0:\n            img_patch = img_patch / 255.0\n\n        # Advanced normalization\n        img_patch = (img_patch - img_patch.mean()) / (img_patch.std() + 1e-8)\n        img_patch = np.clip(img_patch, -3, 3)  # Clip outliers\n\n        mask_patch = (mask_patch > 0).astype(np.float32)\n\n        return (\n            torch.tensor(img_patch[None], dtype=torch.float32),\n            torch.tensor(mask_patch[None], dtype=torch.float32)\n        )","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:29.6568Z","iopub.execute_input":"2025-11-15T08:37:29.657361Z","iopub.status.idle":"2025-11-15T08:37:29.677152Z","shell.execute_reply.started":"2025-11-15T08:37:29.657342Z","shell.execute_reply":"2025-11-15T08:37:29.676541Z"},"trusted":true,"id":"wYpONqt0feij"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# UNET 3D MODEL (Compact version)","metadata":{"id":"sLPd9zrifein"}},{"cell_type":"code","source":"class SpatialAttention3D(nn.Module):\n    def __init__(self, kernel_size=3):\n        super().__init__()\n        self.conv = nn.Conv3d(2, 1, kernel_size, padding=kernel_size//2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        attention = torch.cat([avg_out, max_out], dim=1)\n        attention = self.conv(attention)\n        return x * self.sigmoid(attention)\n\nclass ChannelAttention3D(nn.Module):\n    def __init__(self, in_channels, ratio=8):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool3d(1)\n        self.max_pool = nn.AdaptiveMaxPool3d(1)\n\n        self.fc = nn.Sequential(\n            nn.Conv3d(in_channels, in_channels // ratio, 1, bias=False),\n            nn.ReLU(),\n            nn.Conv3d(in_channels // ratio, in_channels, 1, bias=False)\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = self.fc(self.avg_pool(x))\n        max_out = self.fc(self.max_pool(x))\n        attention = avg_out + max_out\n        return x * self.sigmoid(attention)\n\nclass AttentionGate3D(nn.Module):\n    def __init__(self, F_g, F_l, F_int):\n        super().__init__()\n        self.W_g = nn.Sequential(\n            nn.Conv3d(F_g, F_int, kernel_size=1, bias=True),\n            nn.GroupNorm(max(1, F_int//8), F_int)\n        )\n        self.W_x = nn.Sequential(\n            nn.Conv3d(F_l, F_int, kernel_size=1, bias=True),\n            nn.GroupNorm(max(1, F_int//8), F_int)\n        )\n        self.psi = nn.Sequential(\n            nn.Conv3d(F_int, 1, kernel_size=1, bias=True),\n            nn.GroupNorm(1, 1),\n            nn.Sigmoid()\n        )\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, g, x):\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        return x * psi\n\nclass EnhancedResidualBlock3D(nn.Module):\n    def __init__(self, in_ch, out_ch, dropout=0.0, use_attention=True):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_ch, out_ch, 3, padding=1, bias=False)\n        self.gn1 = nn.GroupNorm(max(1, out_ch//8), out_ch)\n        self.act1 = nn.LeakyReLU(0.01, inplace=True)\n\n        self.conv2 = nn.Conv3d(out_ch, out_ch, 3, padding=1, bias=False)\n        self.gn2 = nn.GroupNorm(max(1, out_ch//8), out_ch)\n\n        # Squeeze-and-Excitation like mechanism\n        self.se = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Conv3d(out_ch, out_ch//4, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch//4, out_ch, 1),\n            nn.Sigmoid()\n        ) if use_attention else nn.Identity()\n\n        self.dropout = nn.Dropout3d(dropout) if dropout > 0 else nn.Identity()\n\n        if in_ch != out_ch:\n            self.res_conv = nn.Conv3d(in_ch, out_ch, 1, bias=False)\n        else:\n            self.res_conv = nn.Identity()\n\n        self.act2 = nn.LeakyReLU(0.01, inplace=True)\n\n    def forward(self, x):\n        res = self.res_conv(x)\n\n        x = self.conv1(x)\n        x = self.gn1(x)\n        x = self.act1(x)\n        x = self.dropout(x)\n\n        x = self.conv2(x)\n        x = self.gn2(x)\n\n        # Apply SE attention\n        se_weights = self.se(x)\n        x = x * se_weights\n\n        x = x + res\n        x = self.act2(x)\n        return x\n\nclass MultiScaleUpSample(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        # Multiple upsampling paths\n        self.path1 = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, 1),\n            nn.Upsample(scale_factor=2, mode='trilinear', align_corners=False)\n        )\n        self.path2 = nn.ConvTranspose3d(in_ch, out_ch, 4, stride=2, padding=1)\n        self.path3 = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch*2, 1),\n            nn.PixelShuffle(2) if hasattr(nn, 'PixelShuffle') else nn.Identity()\n        )\n        self.fusion = nn.Conv3d(out_ch*2, out_ch, 1)  # Adjust based on actual paths\n\n    def forward(self, x, target_shape=None):\n        out1 = self.path1(x)\n        out2 = self.path2(x)\n\n        if target_shape:\n            out1 = F.interpolate(out1, size=target_shape, mode='trilinear', align_corners=False)\n            out2 = F.interpolate(out2, size=target_shape, mode='trilinear', align_corners=False)\n\n        # Combine paths\n        combined = torch.cat([out1, out2], dim=1)\n        return self.fusion(combined)\n\nclass AdvancedResUNet3D(nn.Module):\n    def __init__(self, in_ch=1, out_ch=1, features=[32,64,128,256,512], dropout=0.1, deep_supervision=True):\n        super().__init__()\n        self.deep_supervision = deep_supervision\n\n        # Encoder\n        self.encs = nn.ModuleList()\n        self.pools = nn.ModuleList()\n        self.spatial_attns = nn.ModuleList()\n        self.channel_attns = nn.ModuleList()\n\n        ch = in_ch\n        for f in features:\n            self.encs.append(EnhancedResidualBlock3D(ch, f, dropout=dropout))\n            self.pools.append(nn.MaxPool3d(2))\n            self.spatial_attns.append(SpatialAttention3D())\n            self.channel_attns.append(ChannelAttention3D(f))\n            ch = f\n\n        # Bottleneck with self-attention\n        self.bottleneck = nn.Sequential(\n            EnhancedResidualBlock3D(features[-1], features[-1]*2, dropout=dropout),\n            SpatialAttention3D(),\n            EnhancedResidualBlock3D(features[-1]*2, features[-1]*2, dropout=dropout)\n        )\n\n        # Decoder with attention gates\n        self.upconvs = nn.ModuleList()\n        self.att_gates = nn.ModuleList()\n        self.decs = nn.ModuleList()\n\n        prev_ch = features[-1]*2\n        for i, f in enumerate(reversed(features)):\n            self.upconvs.append(MultiScaleUpSample(prev_ch, f))\n            self.att_gates.append(AttentionGate3D(f, f, f//2))\n            self.decs.append(EnhancedResidualBlock3D(f*2, f, dropout=dropout))\n            prev_ch = f\n\n        # Multi-scale outputs for deep supervision\n        self.deep_outputs = nn.ModuleList()\n        if deep_supervision:\n            for f in features:\n                self.deep_outputs.append(nn.Conv3d(f, out_ch, kernel_size=1))\n\n        self.final = nn.Conv3d(features[0], out_ch, kernel_size=1)\n\n    def forward(self, x):\n        # Encoder path with attention\n        skips = []\n        deep_outputs = []\n\n        for i, (enc, pool, spatial_attn, channel_attn) in enumerate(zip(\n            self.encs, self.pools, self.spatial_attns, self.channel_attns)):\n            x = enc(x)\n            x = spatial_attn(x)\n            x = channel_attn(x)\n            skips.append(x)\n\n            # Deep supervision outputs\n            if self.deep_supervision and self.training:\n                deep_out = self.deep_outputs[i](x)\n                deep_out = F.interpolate(deep_out, scale_factor=2**i, mode='trilinear', align_corners=False)\n                deep_outputs.append(deep_out)\n\n            x = pool(x)\n\n        # Bottleneck\n        x = self.bottleneck(x)\n\n        # Decoder path with attention gates\n        for up, att_gate, dec, skip in zip(self.upconvs, self.att_gates, self.decs, reversed(skips)):\n            x = up(x, target_shape=skip.shape[2:])\n            skip_att = att_gate(x, skip)  # Attention-gated skip connection\n            x = torch.cat([x, skip_att], dim=1)\n            x = dec(x)\n\n        final_out = self.final(x)\n\n        if self.deep_supervision and self.training:\n            return [final_out] + deep_outputs\n        else:\n            return final_out","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:29.837518Z","iopub.execute_input":"2025-11-15T08:37:29.83804Z","iopub.status.idle":"2025-11-15T08:37:29.859687Z","shell.execute_reply.started":"2025-11-15T08:37:29.838024Z","shell.execute_reply":"2025-11-15T08:37:29.858964Z"},"trusted":true,"id":"g2wN_JNcfein"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss","metadata":{"id":"Vy5WX7Sqfeiq"}},{"cell_type":"code","source":"class TverskyLoss(nn.Module):\n    def __init__(self, alpha=0.3, beta=0.7, smooth=1e-6):\n        super().__init__()\n        self.alpha = alpha  # False positive penalty\n        self.beta = beta    # False negative penalty\n        self.smooth = smooth\n\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        tp = (probs * targets).sum()\n        fp = (probs * (1 - targets)).sum()\n        fn = ((1 - probs) * targets).sum()\n\n        tversky = (tp + self.smooth) / (tp + self.alpha*fp + self.beta*fn + self.smooth)\n        return 1 - tversky\n\nclass FocalTverskyLoss(nn.Module):\n    def __init__(self, alpha=0.3, beta=0.7, gamma=2.0, smooth=1e-6):\n        super().__init__()\n        self.tversky = TverskyLoss(alpha, beta, smooth)\n        self.gamma = gamma\n\n    def forward(self, logits, targets):\n        tversky = self.tversky(logits, targets)\n        return torch.pow(tversky, self.gamma)\n\nclass ComboLoss(nn.Module):\n    def __init__(self, alpha=0.5, beta=0.5):\n        super().__init__()\n        self.alpha = alpha\n        self.beta = beta\n        self.bce = nn.BCEWithLogitsLoss()\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n\n        probs = torch.sigmoid(logits)\n        dice_num = 2 * (probs * targets).sum() + 1e-6\n        dice_den = probs.sum() + targets.sum() + 1e-6\n        dice_loss = 1 - dice_num / dice_den\n\n        return self.alpha * bce_loss + self.beta * dice_loss\n\nclass DeepSupervisionLoss(nn.Module):\n    def __init__(self, base_loss_fn, weights=None):\n        super().__init__()\n        self.base_loss_fn = base_loss_fn\n        self.weights = weights\n\n    def forward(self, outputs, targets):\n        if not isinstance(outputs, list):\n            return self.base_loss_fn(outputs, targets)\n\n        if self.weights is None:\n            weights = [1.0] + [0.5] * (len(outputs) - 1)\n        else:\n            weights = self.weights\n\n        total_loss = 0\n        for output, weight in zip(outputs, weights):\n            if output.shape != targets.shape:\n                # Resize target to match output\n                target_resized = F.interpolate(\n                    targets, size=output.shape[2:], mode='trilinear', align_corners=False\n                )\n            else:\n                target_resized = targets\n\n            loss = self.base_loss_fn(output, target_resized)\n            total_loss += weight * loss\n\n        return total_loss\n\nclass AdvancedLossCollection(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.tversky = FocalTverskyLoss()\n        self.combo = ComboLoss()\n        self.focal = FocalLoss(gamma=2.0, alpha=0.25)\n\n    def forward(self, logits, targets):\n        tversky_loss = self.tversky(logits, targets)\n        combo_loss = self.combo(logits, targets)\n        focal_loss = self.focal(logits, targets)\n\n        return 0.4 * tversky_loss + 0.4 * combo_loss + 0.2 * focal_loss\n\nclass DiceLoss(nn.Module):\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        num = 2 * (probs * targets).sum() + 1e-6\n        den = probs.sum() + targets.sum() + 1e-6\n        return 1 - num / den\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2.0, alpha=0.25):\n        super().__init__()\n        self.gamma, self.alpha = gamma, alpha\n    def forward(self, logits, targets):\n        p = torch.sigmoid(logits)\n        pt = p*targets + (1-p)*(1-targets)\n        w = self.alpha*targets + (1-self.alpha)*(1-targets)\n        return (-w*(1-pt)**self.gamma * torch.log(pt+1e-8)).mean()","metadata":{"trusted":true,"id":"UeZC2vi1feiq","execution":{"iopub.status.busy":"2025-11-15T08:37:29.916179Z","iopub.execute_input":"2025-11-15T08:37:29.916395Z","iopub.status.idle":"2025-11-15T08:37:29.929192Z","shell.execute_reply.started":"2025-11-15T08:37:29.916381Z","shell.execute_reply":"2025-11-15T08:37:29.928482Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Sliding Window Inference","metadata":{"id":"qT2fzmaRfeis"}},{"cell_type":"code","source":"def test_time_augmentation_3d(volume, model, num_augs=8, patch=(64,128,128), stride=(32,64,64)):\n    \"\"\"Enhanced sliding window with test-time augmentation\"\"\"\n    model.eval()\n    Z, Y, X = volume.shape\n    pz, py, px = patch\n    sz, sy, sx = stride\n\n    predictions = []\n\n    # Original prediction\n    pred_orig = sliding_window_base(volume, model, patch, stride)\n    predictions.append(pred_orig)\n\n    # Augmented predictions\n    for aug_idx in range(num_augs):\n        # Apply random augmentation\n        aug_volume = volume.copy()\n\n        # Random flip\n        if random.random() < 0.5:\n            axis = random.choice([0, 1, 2])\n            aug_volume = np.flip(aug_volume, axis=axis)\n\n        # Small rotation\n        if random.random() < 0.3:\n            angle = random.uniform(-5, 5)\n            axes = random.choice([(0,1), (0,2), (1,2)])\n            aug_volume = rotate(aug_volume, angle, axes=axes, reshape=False, mode='constant')\n\n        # Noise\n        if random.random() < 0.3:\n            noise = np.random.normal(0, 0.01 * aug_volume.std(), aug_volume.shape)\n            aug_volume = aug_volume + noise\n\n        pred_aug = sliding_window_base(aug_volume, model, patch, stride)\n        predictions.append(pred_aug)\n\n    # Ensemble predictions\n    final_pred = np.mean(predictions, axis=0)\n    return final_pred\n\ndef sliding_window_base(volume, model, patch=(64,128,128), stride=(32,64,64)):\n    \"\"\"Base sliding window implementation\"\"\"\n    model.eval()\n    Z, Y, X = volume.shape\n    pz, py, px = patch\n    sz, sy, sx = stride\n\n    out = np.zeros((Z, Y, X), np.float32)\n    cnt = np.zeros_like(out)\n\n    with torch.no_grad():\n        for z in range(0, max(1, Z-pz+1), sz):\n            for y in range(0, max(1, Y-py+1), sy):\n                for x in range(0, max(1, X-px+1), sx):\n                    # Ensure we don't go out of bounds\n                    z_end = min(z + pz, Z)\n                    y_end = min(y + py, Y)\n                    x_end = min(x + px, X)\n\n                    patch_data = volume[z:z_end, y:y_end, x:x_end]\n\n                    # Pad if necessary\n                    if patch_data.shape != patch:\n                        pad_z = pz - patch_data.shape[0]\n                        pad_y = py - patch_data.shape[1]\n                        pad_x = px - patch_data.shape[2]\n                        patch_data = np.pad(patch_data,\n                                          ((0, pad_z), (0, pad_y), (0, pad_x)),\n                                          mode='constant')\n\n                    # Normalize\n                    patch_norm = (patch_data - patch_data.mean()) / (patch_data.std() + 1e-8)\n                    patch_tensor = torch.tensor(patch_norm[None, None], dtype=torch.float32).cuda()\n\n                    pred = model(patch_tensor)\n                    if isinstance(pred, list):  # Deep supervision\n                        pred = pred[0]\n\n                    pred = torch.sigmoid(pred)[0, 0].cpu().numpy()\n\n                    # Only use the valid part of prediction\n                    pred_valid = pred[:z_end-z, :y_end-y, :x_end-x]\n                    out[z:z_end, y:y_end, x:x_end] += pred_valid\n                    cnt[z:z_end, y:y_end, x:x_end] += 1\n\n    return out / np.maximum(cnt, 1)\n\ndef multi_model_ensemble_inference(volume, models, weights=None):\n    \"\"\"Ensemble inference with multiple models\"\"\"\n    if weights is None:\n        weights = [1.0 / len(models)] * len(models)\n\n    predictions = []\n    for model in models:\n        pred = test_time_augmentation_3d(volume, model, num_augs=4)\n        predictions.append(pred)\n\n    # Weighted ensemble\n    final_pred = np.zeros_like(predictions[0])\n    for pred, weight in zip(predictions, weights):\n        final_pred += weight * pred\n\n    return final_pred","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:30.011343Z","iopub.execute_input":"2025-11-15T08:37:30.011729Z","iopub.status.idle":"2025-11-15T08:37:30.023551Z","shell.execute_reply.started":"2025-11-15T08:37:30.011712Z","shell.execute_reply":"2025-11-15T08:37:30.022811Z"},"trusted":true,"id":"C2-FqhCvfeit"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Setup: Paths / Splits","metadata":{"id":"P0faxSd-feiu"}},{"cell_type":"code","source":"ls -la","metadata":{"id":"Ij-l3Tn-icqb","outputId":"71717f1b-fbba-437e-9866-7c8a05f92c3e","trusted":true,"execution":{"iopub.status.busy":"2025-11-15T08:37:30.024621Z","iopub.execute_input":"2025-11-15T08:37:30.024829Z","iopub.status.idle":"2025-11-15T08:37:30.188085Z","shell.execute_reply.started":"2025-11-15T08:37:30.024808Z","shell.execute_reply":"2025-11-15T08:37:30.187364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ROOT = '/kaggle/input/vesuvius-challenge-surface-detection'\n\n# Simple hardcoded paths for Kaggle dataset\ntrain_imgs_dir = ROOT / \"train_images\"\ntrain_labels_dir = ROOT / \"train_labels\"\ntest_imgs_dir = ROOT / \"test_images\"\ntrain_csv = ROOT / \"train.csv\"\ntest_csv = ROOT / \"test.csv\"\n\n# Load training data\ntrain_imgs = sorted(train_imgs_dir.glob(\"*.tif\"))\ntrain_masks = sorted(train_labels_dir.glob(\"*.tif\"))\n\nprint(f\"Found {len(train_imgs)} training images\")\nprint(f\"Found {len(train_masks)} training masks\")\n\n# K-fold cross validation setup\nkfold = KFold(n_splits=5, shuffle=True, random_state=42)\nfold_splits = list(kfold.split(train_imgs))\n\n# Use first fold for this run\ntrain_idx, val_idx = fold_splits[0]\ntrain_list = [train_imgs[i] for i in train_idx]\nval_list = [train_imgs[i] for i in val_idx]\ntrain_mask_list = [train_masks[i] for i in train_idx]\nval_mask_list = [train_masks[i] for i in val_idx]\n\n# Simple fragment weights (higher for training data)\ntrain_weights = [1.5] * len(train_list)\n\nprint(f\"Train: {len(train_list)} images, Val: {len(val_list)} images\")","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:30.189765Z","iopub.execute_input":"2025-11-15T08:37:30.190059Z","iopub.status.idle":"2025-11-15T08:37:30.220454Z","shell.execute_reply.started":"2025-11-15T08:37:30.190037Z","shell.execute_reply":"2025-11-15T08:37:30.219712Z"},"trusted":true,"id":"DVe4Hbgvfeiu","outputId":"af099351-296d-4367-d2a8-bca8b5e2edeb"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DataLoaders","metadata":{"id":"MkBnWlXpfeiv"}},{"cell_type":"code","source":"# Advanced augmentation pipeline\ntransforms = A.Compose([\n    A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.7),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.3),\n    A.RandomRotate90(p=0.4),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),\n    A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=0.3),\n    A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.3),\n    A.OpticalDistortion(distort_limit=0.1, shift_limit=0.05, p=0.3),\n    A.GaussNoise(var_limit=(10.0, 50.0), p=0.4),\n    A.MotionBlur(blur_limit=3, p=0.2),\n    A.MedianBlur(blur_limit=3, p=0.2),\n    A.OneOf([\n        A.CLAHE(clip_limit=2.0, p=0.5),\n        A.Equalize(p=0.5),\n    ], p=0.5),\n])\n\n# Multi-scale datasets\npatch_sizes = [(32,64,64), (64,128,128), (96,192,192)]\n\ntrain_ds = MultiScaleVesuviusDataset(\n    train_list, train_mask_list,\n    patch_sizes=patch_sizes,\n    transforms=transforms,\n    mode='train',\n    fragment_weights=train_weights\n)\n\nval_ds = MultiScaleVesuviusDataset(\n    val_list, val_mask_list,\n    patch_sizes=[(64,128,128)],  # Fixed size for validation\n    transforms=None,\n    mode='val'\n)\n\n# Weighted sampling for imbalanced data\ntrain_sampler = WeightedRandomSampler(\n    weights=[train_weights[i % len(train_weights)] for i in range(len(train_ds))],\n    num_samples=len(train_ds),\n    replacement=True\n)\n\ntrain_dl = DataLoader(train_ds, batch_size=2, sampler=train_sampler, num_workers=4, pin_memory=True)\nval_dl = DataLoader(val_ds, batch_size=1, shuffle=False, num_workers=2, pin_memory=True)\n\nprint(f\"Training batches: {len(train_dl)}, Validation batches: {len(val_dl)}\")","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:30.221045Z","iopub.status.idle":"2025-11-15T08:37:30.221334Z","shell.execute_reply.started":"2025-11-15T08:37:30.221168Z","shell.execute_reply":"2025-11-15T08:37:30.221184Z"},"trusted":true,"id":"EMQ9mp2gfeiv","outputId":"ea053dfd-d16e-404a-c40b-090cf8f01c17"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.amp import autocast\nfrom torch.amp import GradScaler\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, ReduceLROnPlateau\n\n# Create multiple models for ensemble\nmodels = []\noptimizers = []\nschedulers = []\n\n# Model 1: Standard AdvancedResUNet3D\nmodel1 = AdvancedResUNet3D(features=[32,64,128,256], dropout=0.1).cuda()\nmodels.append(model1)\n\n# Model 2: Deeper variant\nmodel2 = AdvancedResUNet3D(features=[24,48,96,192,384], dropout=0.15).cuda()\nmodels.append(model2)\n\n# Model 3: Wider variant\nmodel3 = AdvancedResUNet3D(features=[48,96,192,384], dropout=0.08).cuda()\nmodels.append(model3)\n\n# Setup optimizers and schedulers for each model\nfor model in models:\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)\n    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)\n    optimizers.append(optimizer)\n    schedulers.append(scheduler)\n\n# Mixed precision training - Fixed deprecated usage\nscaler = GradScaler('cuda')\n\n# Advanced loss with deep supervision\nbase_criterion = AdvancedLossCollection()\ncriterion = DeepSupervisionLoss(base_criterion)\n\nEPOCHS = 15\n\n# Advanced metrics tracking\ndef calculate_metrics(preds, targets, threshold=0.5):\n    preds_bin = (preds > threshold).float()\n\n    tp = (preds_bin * targets).sum()\n    fp = (preds_bin * (1 - targets)).sum()\n    fn = ((1 - preds_bin) * targets).sum()\n    tn = ((1 - preds_bin) * (1 - targets)).sum()\n\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n    f1 = 2 * (precision * recall) / (precision + recall + 1e-8)\n\n    dice = (2 * tp + 1e-8) / (2 * tp + fp + fn + 1e-8)\n    iou = tp / (tp + fp + fn + 1e-8)\n\n    return {\n        'dice': dice.item(),\n        'iou': iou.item(),\n        'f1': f1.item(),\n        'precision': precision.item(),\n        'recall': recall.item()\n    }\n\ndef validate_ensemble(models, loader):\n    for model in models:\n        model.eval()\n\n    metrics_list = []\n\n    with torch.no_grad():\n        for x, y in tqdm(loader, desc=\"Validation\"):\n            x, y = x.cuda(), y.cuda()\n\n            # Ensemble prediction\n            ensemble_pred = 0\n            for model in models:\n                with autocast(\"cuda\"):\n                    pred = model(x)\n                    if isinstance(pred, list):\n                        pred = pred[0]\n                    pred = torch.sigmoid(pred)\n                ensemble_pred += pred\n\n            ensemble_pred = ensemble_pred / len(models)\n\n            # Calculate metrics\n            metrics = calculate_metrics(ensemble_pred, y)\n            metrics_list.append(metrics)\n\n    # Average metrics\n    avg_metrics = {}\n    for key in metrics_list[0].keys():\n        avg_metrics[key] = np.mean([m[key] for m in metrics_list])\n\n    return avg_metrics\n\n# Training loop with progressive learning\nbest_dice = 0\npatience = 0\nmax_patience = 5\n\n# Progressive patch sizes\npatch_schedules = [\n    [(32,64,64)] * 3,      # Small patches first\n    [(32,64,64), (64,128,128)] * 4,  # Mixed sizes\n    [(64,128,128), (96,192,192)] * 8  # Larger patches later\n]\n\nfor epoch in range(EPOCHS):\n    # Update patch sizes based on schedule\n    if epoch < 3:\n        current_patches = patch_schedules[0]\n    elif epoch < 7:\n        current_patches = patch_schedules[1]\n    else:\n        current_patches = patch_schedules[2]\n\n    # Update dataset patch sizes\n    train_ds.patch_sizes = current_patches[0] if len(current_patches) == 1 else current_patches\n\n    # Training phase\n    for model in models:\n        model.train()\n\n    total_losses = [0.0] * len(models)\n\n    pbar = tqdm(train_dl, desc=f\"Epoch {epoch}\")\n    for batch_idx, (x, y) in enumerate(pbar):\n        x, y = x.cuda(), y.cuda()\n\n        # Train each model\n        for model_idx, (model, optimizer, scheduler) in enumerate(zip(models, optimizers, schedulers)):\n            optimizer.zero_grad()\n\n            with autocast(\"cuda\"):\n                outputs = model(x)\n                loss = criterion(outputs, y)\n\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            total_losses[model_idx] += loss.item()\n\n        # Update progress\n        if batch_idx % 10 == 0:\n            avg_loss = np.mean([total_losses[i] / (batch_idx + 1) for i in range(len(models))])\n            pbar.set_postfix({'loss': f'{avg_loss:.4f}'})\n\n    # Update schedulers\n    for scheduler in schedulers:\n        scheduler.step()\n\n    # Validation\n    val_metrics = validate_ensemble(models, val_dl)\n\n    print(f\"Epoch {epoch}\")\n    print(f\"Train Loss: {[f'{total_losses[i]/len(train_dl):.4f}' for i in range(len(models))]}\")\n    print(f\"Val Metrics: Dice={val_metrics['dice']:.4f}, IoU={val_metrics['iou']:.4f}, F1={val_metrics['f1']:.4f}\")\n\n    # Save best models\n    if val_metrics['dice'] > best_dice:\n        best_dice = val_metrics['dice']\n        patience = 0\n\n        # Save all models\n        for i, model in enumerate(models):\n            torch.save(model.state_dict(), f\"/kaggle/working/best_model_{i}.pth\")\n\n        print(f\"Saved best ensemble! Dice: {best_dice:.4f}\")\n    else:\n        patience += 1\n\n    if patience >= max_patience:\n        print(f\"Early stopping at epoch {epoch}\")\n        break\n\n    # Memory cleanup\n    if epoch % 3 == 0:\n        torch.cuda.empty_cache()\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:30.343529Z","iopub.execute_input":"2025-11-15T08:37:30.343745Z","iopub.status.idle":"2025-11-15T08:37:41.1226Z","shell.execute_reply.started":"2025-11-15T08:37:30.343729Z","shell.execute_reply":"2025-11-15T08:37:41.119296Z"},"trusted":true,"id":"9g0GMQoLfeiw"},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference on Test Set","metadata":{"id":"ZxCQDI6xfeiw"}},{"cell_type":"code","source":"def advanced_rle_encode(mask):\n    pixels = mask.flatten(order=\"F\")\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    return \" \".join(str(x) for x in runs)\n\ndef post_process_prediction(pred, min_size=100, fill_holes=True):\n    \"\"\"Advanced post-processing of predictions\"\"\"\n    from scipy import ndimage\n    from skimage.morphology import remove_small_objects, binary_closing\n    from skimage.measure import label\n\n    # Convert to binary\n    binary_pred = pred > 0.5\n\n    # Fill small holes\n    if fill_holes:\n        binary_pred = ndimage.binary_fill_holes(binary_pred)\n\n    # Remove small objects\n    if min_size > 0:\n        # Process slice by slice for 3D\n        for i in range(binary_pred.shape[0]):\n            labeled_img = label(binary_pred[i])\n            binary_pred[i] = remove_small_objects(labeled_img, min_size=min_size) > 0\n\n    # Morphological closing to connect nearby regions\n    binary_pred = binary_closing(binary_pred)\n\n    return binary_pred.astype(np.uint8)\n\n# Load best models\nmodels = []\nfor i in range(3):\n    if i == 0:\n        model = AdvancedResUNet3D(features=[32,64,128,256], dropout=0.1).cuda()\n    elif i == 1:\n        model = AdvancedResUNet3D(features=[24,48,96,192,384], dropout=0.15).cuda()\n    else:\n        model = AdvancedResUNet3D(features=[48,96,192,384], dropout=0.08).cuda()\n\n    model.load_state_dict(torch.load(f\"/kaggle/working/best_model_{i}.pth\"))\n    model.eval()\n    models.append(model)\n\n# Test data processing\nimport os\nimport zipfile\nimport tifffile as tiff\n\nos.makedirs(\"/kaggle/working/pred_masks\", exist_ok=True)\n\n# Simple test file scanning\ntest_files = sorted(test_imgs_dir.glob(\"*.tif\"))\nprint(f\"Found {len(test_files)} test images\")\n\n# Load test CSV to get fragment IDs\ntest_df = pd.read_csv(test_csv)\nprint(f\"Test CSV has {len(test_df)} entries\")\n\n# Process each test image\nfor test_file in tqdm(test_files, desc=\"Processing test images\"):\n    fragment_id = test_file.stem  # Get filename without extension\n    print(f\"Processing fragment: {fragment_id}\")\n\n    # Load 3D volume\n    vol = load_3d_tiff_optimized(test_file).astype(np.float32)\n\n    # Preprocessing\n    vol = adaptive_histogram_equalization_3d(vol)\n    vol_norm = (vol - vol.mean()) / (vol.std() + 1e-8)\n\n    # Multi-model ensemble inference with TTA\n    prob = multi_model_ensemble_inference(\n        vol_norm,\n        models,\n        weights=[0.4, 0.35, 0.25]  # Weight models differently\n    )\n\n    # Advanced post-processing\n    mask = post_process_prediction(prob, min_size=50, fill_holes=True)\n\n    # Save with compression\n    out_path = f\"/kaggle/working/pred_masks/{fragment_id}.tif\"\n    tiff.imwrite(out_path, mask, dtype=\"uint8\", compression='lzw')\n\n    print(f\"Processed {fragment_id}: shape={mask.shape}, ink_pixels={mask.sum()}\")\n\nprint(\"All TIFF masks written successfully.\")\n\n# Create submission with validation\nzip_path = \"/kaggle/working/submission.zip\"\n\nwith zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as z:\n    for test_file in test_files:\n        fragment_id = test_file.stem\n        file_path = f\"/kaggle/working/pred_masks/{fragment_id}.tif\"\n        if os.path.exists(file_path):\n            z.write(file_path, arcname=f\"{fragment_id}.tif\")\n\n            # Verify the file\n            file_size = os.path.getsize(file_path)\n            print(f\"Added {fragment_id}.tif to submission (size: {file_size} bytes)\")\n\nprint(f\"Submission created: {zip_path}\")\nprint(f\"Submission size: {os.path.getsize(zip_path) / (1024*1024):.2f} MB\")\n\n# Memory cleanup\nfor model in models:\n    del model\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2025-11-15T08:37:41.124993Z","iopub.status.idle":"2025-11-15T08:37:41.125731Z","shell.execute_reply.started":"2025-11-15T08:37:41.125476Z","shell.execute_reply":"2025-11-15T08:37:41.125498Z"},"trusted":true,"id":"un6bNPQifeix"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"id":"qfAy9cQifeix"},"outputs":[],"execution_count":null}]}