{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13773433,"sourceType":"datasetVersion","datasetId":8766236},{"sourceId":674998,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":510291,"modelId":524964}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nimport os\n\n# 1. Clone the repository\n!git clone https://github.com/mirthAI/CSA-Net.git\n\n# 2. Install the specific library needed for their configs\n!pip install ml_collections\n\n# 3. Add the 'networks' folder to Python's path\n# The structure is CSA-Net -> CSANet -> networks\nsys.path.append('/kaggle/working/CSA-Net/CSANet')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:01.362354Z","iopub.execute_input":"2025-12-07T05:15:01.362579Z","iopub.status.idle":"2025-12-07T05:15:06.732765Z","shell.execute_reply.started":"2025-12-07T05:15:01.362561Z","shell.execute_reply":"2025-12-07T05:15:06.731867Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch \nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport tqdm as tqdm\nimport os\nfrom ipywidgets import interact, IntSlider\nfrom torch.utils.data import Dataset, DataLoader\nimport glob\nfrom sklearn.model_selection import train_test_split\nimport ml_collections\nimport torch.optim as optim\nfrom tqdm.notebook import tqdm\nimport torch.optim as optim\nfrom torch.amp import autocast, GradScaler\n\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:06.734869Z","iopub.execute_input":"2025-12-07T05:15:06.735094Z","iopub.status.idle":"2025-12-07T05:15:11.270691Z","shell.execute_reply.started":"2025-12-07T05:15:06.735073Z","shell.execute_reply":"2025-12-07T05:15:11.269914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Let's view the data once\n","metadata":{}},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/input/vesuvius-npy/\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.271524Z","iopub.execute_input":"2025-12-07T05:15:11.271900Z","iopub.status.idle":"2025-12-07T05:15:11.282743Z","shell.execute_reply.started":"2025-12-07T05:15:11.271877Z","shell.execute_reply":"2025-12-07T05:15:11.282063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/vesuvius-npy/train_images\" \nMASK_DIR  = \"/kaggle/input/vesuvius-npy/train_labels\"\ntrain_files = sorted(glob.glob(os.path.join(DATA_DIR, \"*.npy\")))\nlabel_files = sorted(glob.glob(os.path.join(MASK_DIR, \"*.npy\")))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.283537Z","iopub.execute_input":"2025-12-07T05:15:11.284114Z","iopub.status.idle":"2025-12-07T05:15:11.337755Z","shell.execute_reply.started":"2025-12-07T05:15:11.284092Z","shell.execute_reply":"2025-12-07T05:15:11.337233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"volume = np.load(label_files[0])\nvolume.shape\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.338333Z","iopub.execute_input":"2025-12-07T05:15:11.338546Z","iopub.status.idle":"2025-12-07T05:15:11.564023Z","shell.execute_reply.started":"2025-12-07T05:15:11.338530Z","shell.execute_reply":"2025-12-07T05:15:11.563214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.device_count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.564832Z","iopub.execute_input":"2025-12-07T05:15:11.565133Z","iopub.status.idle":"2025-12-07T05:15:11.597654Z","shell.execute_reply.started":"2025-12-07T05:15:11.565114Z","shell.execute_reply":"2025-12-07T05:15:11.597049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"volume.max()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.599857Z","iopub.execute_input":"2025-12-07T05:15:11.600277Z","iopub.status.idle":"2025-12-07T05:15:11.606997Z","shell.execute_reply.started":"2025-12-07T05:15:11.600254Z","shell.execute_reply":"2025-12-07T05:15:11.606341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  Define a function that plots a single slice\ndef explore_volume(layer_index):\n    plt.figure(figsize=(4, 4))\n    \n    #cmap='gray' is standard for CT scans\n    #removing the darkest and the lighest pixels in images\n    plt.imshow(volume[layer_index, :, :], cmap='gray', vmin=0, vmax=1)\n    plt.title(f\"Z-Axis Layer: {layer_index}\")\n    plt.axis('off')\n    plt.show()\n\n\n\n# This creates a slider from 0 to the max depth of the volume\ninteract(explore_volume, layer_index=(0, volume.shape[0] - 1));","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.607626Z","iopub.execute_input":"2025-12-07T05:15:11.607854Z","iopub.status.idle":"2025-12-07T05:15:11.791432Z","shell.execute_reply.started":"2025-12-07T05:15:11.607831Z","shell.execute_reply":"2025-12-07T05:15:11.790624Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Datasets and Dataloaders","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport random\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\n\n\nclass VesuviusPatchedDataset(Dataset):\n    \"\"\"\n    Patched dataset for Vesuvius volumes with per-worker memmap caching.\n\n    Returns:\n      - image:  (3, patch_size, patch_size), float32 in [0, 1]\n      - target: (1, patch_size, patch_size), float32 {0,1}\n      - valid_mask: (1, patch_size, patch_size), float32 {0,1}\n    \"\"\"\n\n    def __init__(\n        self,\n        img_dir,\n        mask_dir,\n        slice_map=None,\n        patch_size=224,\n        mode=\"train\",\n        transform=None,\n        cache_volumes=True,\n    ):\n        \"\"\"\n        Args:\n            img_dir: directory with volume .npy files\n            mask_dir: directory with mask .npy files\n            slice_map: optional list of (filename, z_idx). If None, we build it by scanning img_dir.\n            patch_size: output patch size (square)\n            mode: \"train\" or \"val\" / \"test\" (controls cropping)\n            transform: optional transform on the image tensor only (not target / valid_mask)\n            cache_volumes: if True, keep memmaps open per worker to avoid repeated np.load\n        \"\"\"\n        self.img_dir = img_dir\n        self.mask_dir = mask_dir\n        self.patch_size = patch_size\n        self.mode = mode\n        self.transform = transform\n        self.cache_volumes = cache_volumes\n\n        # Build or store slice_map\n        if slice_map is not None:\n            self.slice_map = slice_map\n        else:\n            self.slice_map = []\n            print(f\"[VesuviusPatchedDataset] Scanning {img_dir} for .npy volumes...\")\n            all_files = sorted(glob.glob(os.path.join(img_dir, \"*.npy\")))\n            for file_path in all_files:\n                fname = os.path.basename(file_path)\n                try:\n                    vol = np.load(file_path, mmap_mode=\"r\")\n                    depth = vol.shape[0]\n                    for z in range(depth):\n                        self.slice_map.append((fname, z))\n                except Exception as e:\n                    print(f\"  Error reading {file_path}: {e}\")\n            print(f\"[VesuviusPatchedDataset] Indexed {len(self.slice_map)} slices.\")\n\n        # Per-dataset (per-worker) caches\n        self._vol_cache = {}\n        self._mask_cache = {}\n\n    def __len__(self):\n        return len(self.slice_map)\n\n    # ----------------- caching helpers ----------------- #\n\n    def _get_arrays(self, filename):\n        \"\"\"\n        Get (vol_mmap, mask_mmap) for given filename, using a simple cache.\n        Cache lives inside each dataset instance (so per DataLoader worker).\n        \"\"\"\n        if not self.cache_volumes:\n            vol_path = os.path.join(self.img_dir, filename)\n            mask_path = os.path.join(self.mask_dir, filename)\n            vol = np.load(vol_path, mmap_mode=\"r\")\n            mask = np.load(mask_path, mmap_mode=\"r\")\n            return vol, mask\n\n        if filename not in self._vol_cache:\n            vol_path = os.path.join(self.img_dir, filename)\n            mask_path = os.path.join(self.mask_dir, filename)\n            self._vol_cache[filename] = np.load(vol_path, mmap_mode=\"r\")\n            self._mask_cache[filename] = np.load(mask_path, mmap_mode=\"r\")\n\n        return self._vol_cache[filename], self._mask_cache[filename]\n\n    # ----------------- cropping helpers ----------------- #\n\n    def _get_random_crop_coords(self, h, w):\n        \"\"\"Random top/left for training.\"\"\"\n        if h <= self.patch_size or w <= self.patch_size:\n            # will be handled in crop_or_pad\n            return 0, 0\n        top = random.randint(0, h - self.patch_size)\n        left = random.randint(0, w - self.patch_size)\n        return top, left\n\n    def _get_center_crop_coords(self, h, w):\n        \"\"\"Center crop for val/test.\"\"\"\n        top = max(0, (h - self.patch_size) // 2)\n        left = max(0, (w - self.patch_size) // 2)\n        return top, left\n\n    def _crop_or_pad(self, img2d, top=None, left=None):\n        \"\"\"\n        Crop a (H,W) image to (patch_size, patch_size).\n        If the image is smaller, pad with zeros, centered.\n        \"\"\"\n        h, w = img2d.shape\n        ps = self.patch_size\n\n        # Case 1: image is big enough -> just crop\n        if h >= ps and w >= ps:\n            if top is None or left is None:\n                top, left = self._get_center_crop_coords(h, w)\n            top = min(max(0, top), h - ps)\n            left = min(max(0, left), w - ps)\n            return img2d[top : top + ps, left : left + ps]\n\n        # Case 2: need padding\n        out = np.zeros((ps, ps), dtype=img2d.dtype)\n\n        # Clip to not exceed ps\n        crop_h = min(h, ps)\n        crop_w = min(w, ps)\n\n        # Place original into center of out\n        y0 = (ps - crop_h) // 2\n        x0 = (ps - crop_w) // 2\n\n        out[y0 : y0 + crop_h, x0 : x0 + crop_w] = img2d[:crop_h, :crop_w]\n        return out\n\n    # ----------------- main __getitem__ ----------------- #\n\n    def __getitem__(self, idx):\n        filename, z_idx = self.slice_map[idx]\n\n        # Get volume and mask arrays (possibly cached)\n        vol_mmap, mask_mmap = self._get_arrays(filename)\n\n        depth, h, w = vol_mmap.shape\n\n        # z indices: prev/current/next with clamping\n        z_prev = max(0, z_idx - 1)\n        z_next = min(depth - 1, z_idx + 1)\n\n        # Choose crop coords\n        if self.mode == \"train\":\n            top, left = self._get_random_crop_coords(h, w)\n        else:\n            top, left = self._get_center_crop_coords(h, w)\n\n        # Extract 3 slices\n        slice_prev = self._crop_or_pad(vol_mmap[z_prev], top, left)\n        slice_curr = self._crop_or_pad(vol_mmap[z_idx], top, left)\n        slice_next = self._crop_or_pad(vol_mmap[z_next], top, left)\n         # Mask handling:\n        #  - mask == 1 -> positive\n        #  - mask == 2 -> ignore region (valid_mask = 0)\n        #  - else -> background\n        mask_raw = self._crop_or_pad(mask_mmap[z_idx], top, left)\n\n        if random.random() > 0.5:\n                slice_prev = np.fliplr(slice_prev)\n                slice_curr = np.fliplr(slice_curr)\n                slice_next = np.fliplr(slice_next)\n                mask_raw   = np.fliplr(mask_raw)\n            \n            # 2. Random Flip Up-Down\n        if random.random() > 0.5:\n                slice_prev = np.flipud(slice_prev)\n                slice_curr = np.flipud(slice_curr)\n                slice_next = np.flipud(slice_next)\n                mask_raw   = np.flipud(mask_raw)\n            \n            # 3. Random Rotation (0, 90, 180, 270)\n        k = random.randint(0, 3)\n        if k > 0:\n                slice_prev = np.rot90(slice_prev, k)\n                slice_curr = np.rot90(slice_curr, k)\n                slice_next = np.rot90(slice_next, k)\n                mask_raw   = np.rot90(mask_raw, k)\n\n\n        # Stack channels, normalize to [0,1]\n        input_stack = np.stack(\n            [slice_prev, slice_curr, slice_next], axis=0\n        ).astype(np.float32)\n        input_stack = input_stack / 255.0\n\n\n        target = (mask_raw == 1).astype(np.float32)\n        valid_mask = (mask_raw != 2).astype(np.float32)\n\n        # Torch tensors\n        image = torch.from_numpy(input_stack).float()\n        target = torch.from_numpy(target).unsqueeze(0)       # (1, H, W)\n        valid_mask = torch.from_numpy(valid_mask).unsqueeze(0)\n\n        # Optional transform on image only\n        if self.transform is not None:\n            image = self.transform(image)\n\n        return image, target, valid_mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.792302Z","iopub.execute_input":"2025-12-07T05:15:11.792811Z","iopub.status.idle":"2025-12-07T05:15:11.811730Z","shell.execute_reply.started":"2025-12-07T05:15:11.792786Z","shell.execute_reply":"2025-12-07T05:15:11.811077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def vesuvius_collate(batch):\n    # batch is a list of (slices, mask, top, left)\n    # we just return the list as-is and handle stacking ourselves\n    return batch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:11.812323Z","iopub.execute_input":"2025-12-07T05:15:11.812595Z","iopub.status.idle":"2025-12-07T05:15:11.829963Z","shell.execute_reply.started":"2025-12-07T05:15:11.812567Z","shell.execute_reply":"2025-12-07T05:15:11.829325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# 1. GENERATE THE MAP ONCE\n# We create a dummy dataset just to run the scanning logic once.\nprint(\"Generating Master Index...\")\nmaster_dataset = VesuviusPatchedDataset(\n    img_dir=DATA_DIR, mask_dir=MASK_DIR, slice_map=None\n)\nmaster_map = master_dataset.slice_map # This is the list [(file, 0), (file, 1)...]\n\n# 2. SPLIT THE LIST (Instant)\ntrain_map, val_map = train_test_split(\n    master_map, \n    test_size=0.10, \n    random_state=42\n)\n\nprint(f\"Training Slices: {len(train_map)}\")\nprint(f\"Validation Slices: {len(val_map)}\")\n\n# 3. INITIALIZE DATASETS WITH PRE-MADE MAPS (Instant)\ntrain_dataset = VesuviusPatchedDataset(\n    img_dir=DATA_DIR, mask_dir=MASK_DIR, \n    slice_map=train_map,  # <--- PASS THE LIST\n    patch_size=224, \n    mode='train' # Random Crop\n)\n\nval_dataset = VesuviusPatchedDataset(\n    img_dir=DATA_DIR, mask_dir=MASK_DIR, \n    slice_map=val_map,    # <--- PASS THE LIST\n    patch_size=224, \n    mode='val'            # Center Crop\n)\n\n# Tune this based on CPU cores available\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:16:33.345193Z","iopub.execute_input":"2025-12-07T05:16:33.345851Z","iopub.status.idle":"2025-12-07T05:16:36.901940Z","shell.execute_reply.started":"2025-12-07T05:16:33.345816Z","shell.execute_reply":"2025-12-07T05:16:36.901159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_WORKERS = 4           # Try 6 first; if noisy HDD slow, reduce to 4\nPREFETCH_FACTOR = 2       # Default=2, can try 3 later\n\nworld_size = max(1, torch.cuda.device_count())\nBASE_BS = 18                      # what used to work on 1 GPU\nGLOBAL_BS = BASE_BS * world_size  # 32 if you have 2 GPUs\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=GLOBAL_BS,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    drop_last=True,\n    persistent_workers=True,\n    prefetch_factor=PREFETCH_FACTOR)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=GLOBAL_BS,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    persistent_workers=True,\n    prefetch_factor=PREFETCH_FACTOR)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:16:37.460317Z","iopub.execute_input":"2025-12-07T05:16:37.460997Z","iopub.status.idle":"2025-12-07T05:16:37.475480Z","shell.execute_reply.started":"2025-12-07T05:16:37.460966Z","shell.execute_reply":"2025-12-07T05:16:37.474559Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation Metric","metadata":{}},{"cell_type":"markdown","source":"I use a combination of BCE, Tverskty Loss (to penalize bridges, without high computation)","metadata":{}},{"cell_type":"code","source":"class VesuviusTverskyLoss(nn.Module):\n    def __init__(self, alpha=0.7, beta=0.3, gamma=1.0, smooth=1e-6, bce_pos_weight=1.0):\n        super().__init__()\n        self.alpha = alpha\n        self.beta = beta\n        self.gamma = gamma\n        self.smooth = smooth\n        \n        # 1. BCE with Positive Weight (Find the ink!)\n        self.register_buffer('pos_weight', torch.tensor([bce_pos_weight]))\n        self.bce = nn.BCEWithLogitsLoss(\n            reduction='none', \n            pos_weight=self.pos_weight\n        )\n\n    def forward(self, inputs, targets, mask=None):\n        # 2. Weighted BCE \n        bce_pixel_loss = self.bce(inputs, targets)\n        \n        if mask is not None:\n            bce_pixel_loss = bce_pixel_loss * mask\n            bce_loss = bce_pixel_loss.sum() / (mask.sum() + 1e-8)\n        else:\n            bce_loss = bce_pixel_loss.mean()\n\n        # 3. Tversky (Kill the bridges!)\n        probs = torch.sigmoid(inputs)\n        probs_flat = probs.view(-1)\n        targets_flat = targets.view(-1)\n        \n        if mask is not None:\n            mask_flat = mask.view(-1)\n            TP = (probs_flat * targets_flat * mask_flat).sum()\n            FP = (probs_flat * (1 - targets_flat) * mask_flat).sum()\n            FN = ((1 - probs_flat) * targets_flat * mask_flat).sum()\n        else:\n            TP = (probs_flat * targets_flat).sum()\n            FP = (probs_flat * (1 - targets_flat)).sum()\n            FN = ((1 - probs_flat) * targets_flat).sum()\n        \n        # High Alpha (0.7) means FP increases denominator fast -> Low Score -> High Loss\n        tversky_index = (TP + self.smooth) / (TP + self.alpha * FP + self.beta * FN + self.smooth)\n        tversky_loss = 1 - tversky_index\n\n        return bce_loss + (self.gamma * tversky_loss)\n\n# --- METRIC ---\ndef fast_surface_dice(y_pred, y_true, mask=None, tolerance=1):\n    \n    pred_mask = (y_pred > 0.5).float()\n    if mask is not None:\n        pred_mask = pred_mask * mask\n        y_true = y_true * mask\n    \n    kernel_size = 2 * tolerance + 1\n    dilated_pred = F.max_pool2d(pred_mask, kernel_size=kernel_size, stride=1, padding=tolerance)\n    dilated_true = F.max_pool2d(y_true, kernel_size=kernel_size, stride=1, padding=tolerance)\n    \n    intersection_1 = (pred_mask * dilated_true).sum()\n    intersection_2 = (y_true * dilated_pred).sum()\n    dice = (intersection_1 + intersection_2) / (pred_mask.sum() + y_true.sum() + 1e-6)\n    return dice.item()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:23.543683Z","iopub.execute_input":"2025-12-07T05:15:23.543981Z","iopub.status.idle":"2025-12-07T05:15:23.716073Z","shell.execute_reply.started":"2025-12-07T05:15:23.543960Z","shell.execute_reply":"2025-12-07T05:15:23.715382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# We initialize a model","metadata":{}},{"cell_type":"code","source":"def get_b16_config():\n    \"\"\"Returns the ViT-B/16 configuration.\"\"\"\n    config = ml_collections.ConfigDict()\n    config.patches = ml_collections.ConfigDict({'size': (16, 16)})\n    config.hidden_size = 768\n    config.transformer = ml_collections.ConfigDict()\n    config.transformer.mlp_dim = 3072\n    config.transformer.num_heads = 12\n    config.transformer.num_layers = 12\n    config.transformer.attention_dropout_rate = 0.0\n    config.transformer.dropout_rate = 0.1\n\n    config.classifier = 'seg'\n    config.representation_size = None\n    config.resnet_pretrained_path = None\n    config.pretrained_path = '../model/vit_checkpoint/imagenet21k/ViT-B_16.npz'\n    config.patch_size = 16\n\n    config.decoder_channels = (256, 128, 64, 16)\n    config.n_classes = 2\n    config.activation = 'softmax'\n    return config\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:23.716902Z","iopub.execute_input":"2025-12-07T05:15:23.717126Z","iopub.status.idle":"2025-12-07T05:15:23.731835Z","shell.execute_reply.started":"2025-12-07T05:15:23.717108Z","shell.execute_reply":"2025-12-07T05:15:23.731161Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_r50_b16_config():\n    \"\"\"Returns the Resnet50 + ViT-B/16 configuration.\"\"\"\n    config = get_b16_config()\n    config.patches.grid = (16, 16)\n    config.resnet = ml_collections.ConfigDict()\n    config.resnet.num_layers = (3, 4, 9)\n    config.resnet.width_factor = 1\n\n    config.classifier = 'seg'\n    config.pretrained_path = '../model/vit_checkpoint/imagenet21k/R50+ViT-B_16.npz'\n    config.decoder_channels = (256, 128, 64, 16)\n    config.skip_channels = [512, 256, 64, 16]\n    config.n_classes = 2\n    config.n_skip = 3\n    config.activation = 'softmax'\n\n    return config\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:23.732687Z","iopub.execute_input":"2025-12-07T05:15:23.732908Z","iopub.status.idle":"2025-12-07T05:15:23.745633Z","shell.execute_reply.started":"2025-12-07T05:15:23.732889Z","shell.execute_reply":"2025-12-07T05:15:23.744692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import the model class and configs\n# Note: We import from 'networks' because we added CSANet to the path above\nfrom networks.vit_seg_modeling import VisionTransformer as CSANet","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:23.746537Z","iopub.execute_input":"2025-12-07T05:15:23.746813Z","iopub.status.idle":"2025-12-07T05:15:23.768770Z","shell.execute_reply.started":"2025-12-07T05:15:23.746774Z","shell.execute_reply":"2025-12-07T05:15:23.767977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:23.769796Z","iopub.execute_input":"2025-12-07T05:15:23.770237Z","iopub.status.idle":"2025-12-07T05:15:23.773526Z","shell.execute_reply.started":"2025-12-07T05:15:23.770213Z","shell.execute_reply":"2025-12-07T05:15:23.772816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_SIZE = 224\n# 1. Get the base config for R50+ViT-B_16\nconfig = get_r50_b16_config()\n\n# 2. OVERRIDE: Point to your uploaded Kaggle weights\n# (Make sure you uploaded the file and check this path!)\nconfig.pretrained_path = '/kaggle/input/csa-net/pytorch/1-epoch-trained/5/model_epoch_2_fullDice_0.8170.pth'\n\n# 3. OVERRIDE: Set Grid Size for 384x384 images\n# The patch size is 16x16. \n# 384 / 16 = 24. So we need a 24x24 grid.\nconfig.patches.grid = (PATCH_SIZE//16, PATCH_SIZE//16)\n\n# 4. OVERRIDE: Binary Classification\n# We want 1 output channel (Ink probability) and no Softmax (we use Sigmoid later or in Loss)\nconfig.n_classes = 1\nconfig.activation = None \n\nimport torch\nimport os\n\n# ... (Your config setup code from above) ...\n\nmodel = CSANet(config, img_size=PATCH_SIZE, num_classes=1)\n\n# --- LOAD PRE-TRAINED WEIGHTS ---\nprint(f\"Loading weights from: {config.pretrained_path}\")\n\nif os.path.exists(config.pretrained_path):\n    # 1. Load the state dictionary\n    # map_location='cpu' ensures we don't run out of GPU memory during the load\n    state_dict = torch.load(config.pretrained_path, map_location='cpu')\n\n    # 2. Handle 'module.' prefix (in case it was saved from DataParallel)\n    # If your saved model has keys like 'module.backbone...', we need to remove 'module.'\n    new_state_dict = {}\n    for k, v in state_dict.items():\n        if k.startswith('module.'):\n            new_state_dict[k[7:]] = v\n        else:\n            new_state_dict[k] = v\n\n    # 3. Load into the model\n    # strict=True ensures every key matches exactly. \n    # Use strict=False if you encounter minor mismatches (e.g. head size differences), \n    # but for resuming training, True is better to catch errors.\n    missing_keys, unexpected_keys = model.load_state_dict(new_state_dict, strict=True)\n    \n    print(\"Weights loaded successfully!\")\n    if missing_keys:\n        print(\"Missing keys:\", missing_keys)\n    if unexpected_keys:\n        print(\"Unexpected keys:\", unexpected_keys)\n\nelse:\n    print(f\"ERROR: Checkpoint file not found at {config.pretrained_path}\")\n    print(\"Please check the path or your Kaggle dataset connection.\")\n\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:23.774307Z","iopub.execute_input":"2025-12-07T05:15:23.774637Z","iopub.status.idle":"2025-12-07T05:15:29.870668Z","shell.execute_reply.started":"2025-12-07T05:15:23.774615Z","shell.execute_reply":"2025-12-07T05:15:29.870090Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"cell_type":"code","source":"\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Number of trainable parameters: {trainable_params}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:15:29.871417Z","iopub.execute_input":"2025-12-07T05:15:29.871682Z","iopub.status.idle":"2025-12-07T05:15:29.877800Z","shell.execute_reply.started":"2025-12-07T05:15:29.871656Z","shell.execute_reply":"2025-12-07T05:15:29.877119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.notebook import tqdm\n\n\ntorch.backends.cudnn.benchmark = True\n# --- CONFIG ---\nACCUM_STEPS = 1 \nstart = 5\nEPOCHS = 6\nMAX_GRAD_NORM = 2.0  \nLEARNING_RATE = 8e-6\nSAVES_PER_EPOCH = 4  # Save 8 times per epoch\nMINI_VAL_STEPS = 50  # How many batches to check during the mid-epoch saves\n\n# --- SETUP ---\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS * len(train_loader))\n\n# Initialize Loss\ncriterion = VesuviusTverskyLoss(bce_pos_weight = 5.0).to(device)\n\nscaler = GradScaler(\"cuda\")\nbest_val_score = -1\nsave_step_interval = len(train_loader) // SAVES_PER_EPOCH\n\nprint(f\"Starting training: {len(train_loader)} batches | Saves every {save_step_interval} steps\")\n\n# --- HELPER: VALIDATION ---\ndef run_validation(loader, steps_to_run=None, desc=\"Validating\"):\n    model.eval()\n    val_dice_sum = 0.0\n    steps = 0\n    pbar = tqdm(loader, total=steps_to_run if steps_to_run else len(loader), desc=desc, leave=False)\n    \n    with torch.no_grad():\n        for images, targets, valid_mask in pbar:\n            images = images.to(device, non_blocking=True)\n            targets = targets.to(device, non_blocking=True)\n            valid_mask = valid_mask.to(device, non_blocking=True)\n            \n            with autocast(\"cuda\"):\n                outputs = model(images[:, 0:1, ...], images[:, 1:2, ...], images[:, 2:3, ...])\n                probs = torch.sigmoid(outputs)\n    \n            probs = probs * valid_mask\n            targets = targets * valid_mask\n\n            probs_cpu   = probs.detach().cpu()\n            targets_cpu = targets.detach().cpu()\n            mask_cpu    = valid_mask.detach().cpu()\n            \n            batch_dice = fast_surface_dice(probs_cpu, targets_cpu, mask=mask_cpu, tolerance=1)\n\n            val_dice_sum += batch_dice\n            steps += 1\n            if steps_to_run and steps >= steps_to_run: break\n    \n    return val_dice_sum / max(steps, 1)\n\n# --- MAIN LOOP ---\nfor epoch in range(start,EPOCHS):\n    model.train()\n    train_loss = 0.0\n    num_train_steps = 0\n    \n    loop = tqdm(train_loader, total=len(train_loader), desc=f\"Epoch {epoch+1}/{EPOCHS}\")\n    \n    for i, (images, targets, valid_mask) in enumerate(loop):\n        # 1. TRAIN STEP\n        images = images.to(device, non_blocking=True)\n        targets = targets.to(device, non_blocking=True)\n        valid_mask = valid_mask.to(device, non_blocking=True)\n        \n        with autocast(\"cuda\"):\n            outputs = model(images[:, 0:1, ...], images[:, 1:2, ...], images[:, 2:3, ...])\n            loss = criterion(outputs, targets, mask=valid_mask)\n            loss = loss / ACCUM_STEPS \n    \n        scaler.scale(loss).backward()\n        # if((i+1)%10==0):\n        #     break\n        if (i + 1) % ACCUM_STEPS == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n            scheduler.step()\n        \n        batch_loss = loss.item() * ACCUM_STEPS\n        train_loss += batch_loss\n        num_train_steps += 1\n        loop.set_postfix(loss=batch_loss, lr=optimizer.param_groups[0]['lr'])\n\n        # 2. MINI-SAVE (8 times per epoch)\n    \n        if (i + 1) % save_step_interval == 0:\n            print(f\"\\n[Step {i+1}] Running Mini-Validation...\")\n            mini_dice = run_validation(val_loader, steps_to_run=MINI_VAL_STEPS, desc=\"Mini-Val\")\n            \n            ckpt_name = f\"ckpt_ep{epoch+1}_step{i+1}_dice{mini_dice:.4f}.pth\"\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'best_score': mini_dice,\n            }, ckpt_name)\n            print(f\"[Step {i+1}] Mini-Val Dice: {mini_dice:.4f} -> Saved {ckpt_name}\")\n            model.train()\n        \n\n    # 3. FULL VALIDATION\n    avg_train_loss = train_loss / len(train_loader)\n    print(f\"\\nEpoch {epoch+1} Finished. Running FULL Validation...\")\n    full_val_dice = run_validation(val_loader, steps_to_run=None, desc=\"Full-Val\")\n    \n    print(f\"Epoch {epoch+1} Summary | Train Loss: {avg_train_loss:.4f} | Full Val Dice: {full_val_dice:.4f}\")\n\n    # Save Epoch Checkpoint\n    torch.save(model.state_dict(), f\"model_epoch_{epoch+1}_fullDice_{full_val_dice:.4f}.pth\")\n\n    if full_val_dice > best_val_score:\n        best_val_score = full_val_dice\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'best_score': best_val_score,\n        }, f\"checkpoint_full_ep{epoch+1}.pth\")\n        print(f\"New BEST model saved! Score: {best_val_score:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T05:16:44.002298Z","iopub.execute_input":"2025-12-07T05:16:44.003040Z","execution_failed":"2025-12-07T05:17:36.964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}