{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":13773433,"sourceType":"datasetVersion","datasetId":8766236},{"sourceId":14044945,"sourceType":"datasetVersion","datasetId":8941352}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#IMPORTS\nimport sys\nimport subprocess\nimport os\nimport shutil\nimport re\nimport importlib","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# --- CONFIGURATION ---\nDATASET_PATH = \"/kaggle/input/vesuvius-segresmamba-libs-2025\" \nWRITABLE_PATH = \"/kaggle/working/fixed_model\"\n\nprint(\" Starting Final Setup...\")\n\n# --- PART 1: INSTALL LIBRARIES ---\nwheels_path = DATASET_PATH \nfor root, dirs, files in os.walk(DATASET_PATH):\n    if any(f.startswith(\"mamba_ssm\") and f.endswith(\".whl\") for f in files):\n        wheels_path = root\n        break\n\ntry:\n    import mamba_ssm\n    print(\" Mamba already installed.\")\nexcept ImportError:\n    print(f\"...Installing libraries from {wheels_path}...\")\n    try:\n        subprocess.check_call([\n            sys.executable, \"-m\", \"pip\", \"install\", \n            \"--no-index\", \n            \"--no-deps\", \n            f\"--find-links={wheels_path}\", \n            \"causal_conv1d\", \"mamba_ssm\", \"monai\"\n        ])\n        print(\"✅ Libraries installed successfully.\")\n    except Exception as e:\n        print(f\" Installation failed: {e}\")\n\n# --- PART 2: FIND & COPY MODEL CODE ---\nprint(\"...Locating Model Code...\")\nsource_code_path = None\nfor root, dirs, files in os.walk(DATASET_PATH):\n    if \"segresmamba.py\" in files:\n        source_code_path = os.path.dirname(root) \n        break\n\nif not source_code_path:\n    print(\" CRITICAL: Could not find segresmamba.py\")\nelse:\n    # Clear \"Zombie\" memory\n    modules_to_kill = [m for m in sys.modules if \"segresmamba\" in m or \"segmamba\" in m]\n    for m in modules_to_kill:\n        del sys.modules[m]\n\n    # Copy code\n    if os.path.exists(WRITABLE_PATH):\n        shutil.rmtree(WRITABLE_PATH)\n    try:\n        shutil.copytree(source_code_path, WRITABLE_PATH)\n        print(f\" Copied code to: {WRITABLE_PATH}\")\n    except Exception as e:\n        print(f\" Copy warning: {e}. Trying fallback copy...\")\n        os.makedirs(WRITABLE_PATH, exist_ok=True)\n        subprocess.call([\"cp\", \"-r\", f\"{source_code_path}/.\", WRITABLE_PATH])\n\n    # --- PART 3: APPLY THE FIXES (Patch 'bimamba_type' AND 'nslices') ---\n    print(\"...Applying multiple fixes to code...\")\n    patched_count = 0\n    for root, dirs, files in os.walk(WRITABLE_PATH):\n        for filename in files:\n            if filename.endswith(\".py\"):\n                file_path = os.path.join(root, filename)\n                with open(file_path, 'r') as f:\n                    content = f.read()\n                \n                original_content = content\n                \n                # FIX 1: Remove 'bimamba_type'\n                pattern_bimamba = r',?\\s*bimamba_type\\s*=\\s*[^,)]+'\n                if re.search(pattern_bimamba, content):\n                    content = re.sub(pattern_bimamba, '', content)\n                    \n                # FIX 2: Remove 'nslices' (The new error)\n                pattern_nslices = r',?\\s*nslices\\s*=\\s*[^,)]+'\n                if re.search(pattern_nslices, content):\n                    content = re.sub(pattern_nslices, '', content)\n\n                # FIX 3: Remove 'use_fast_path' (Common error in v2, removing preemptively)\n                pattern_fast = r',?\\s*use_fast_path\\s*=\\s*[^,)]+'\n                if re.search(pattern_fast, content):\n                    content = re.sub(pattern_fast, '', content)\n\n                # If we changed anything, save the file\n                if content != original_content:\n                    with open(file_path, 'w') as f:\n                        f.write(content)\n                    patched_count += 1\n                    print(f\"   🔧 Fixed file: {filename}\")\n\n    print(f\" Patched {patched_count} files.\")\n\n    # --- PART 4: IMPORT & TEST ---\n    if WRITABLE_PATH not in sys.path:\n        sys.path.insert(0, WRITABLE_PATH)\n\n    import torch\n    try:\n        from model.segresmamba import SegResMamba\n        print(\" SUCCESS! Model imported.\")\n        \n        if torch.cuda.is_available():\n            # Initialize with CORRECT arguments\n            model = SegResMamba(spatial_dims=3, in_chans=1, out_chans=1).cuda()\n            print(\" TEST PASSED: Model initialized on GPU!\")\n            print(\"   (Code is now compatible with Mamba v2)\")\n        else:\n            print(\" Warning: CUDA not available.\")\n            \n    except Exception as e:\n        print(f\" Final Error: {e}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#IMPORTS\nimport inspect\n\n# Ensure the model is in the path (if you haven't run the Setup Block yet)\n# If you HAVE run the Setup Block, you can skip these 3 lines.\nDATASET_PATH = \"/kaggle/input/vesuvius-segresmamba-libs-2025\"\nfor root, dirs, files in os.walk(DATASET_PATH):\n    if \"segresmamba.py\" in files:\n        sys.path.append(os.path.dirname(root))\n        break\n\n# --- INSPECTION CODE ---\nfrom model.segresmamba import SegResMamba\n\nprint(\" Inspecting SegResMamba Arguments...\")\nprint(\"-\" * 30)\n\n# Method 1: The Signature (Best for seeing variable names)\nsig = inspect.signature(SegResMamba.__init__)\nprint(str(sig))\n\nprint(\"-\" * 30)\n\n# Method 2: The Docstring (Best for reading descriptions)\nprint(SegResMamba.__doc__)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#IMPORTS\nimport 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\nfrom ipywidgets import interact, IntSlider\nfrom torch.utils.data import Dataset, DataLoader\nimport glob\nfrom sklearn.model_selection import train_test_split","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/input/vesuvius-npy/\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/vesuvius-npy/train_images\" \nMASK_DIR  = \"/kaggle/input/vesuvius-npy/train_labels\"\n\ntrain_files = sorted(glob.glob(os.path.join(DATA_DIR, \"*.npy\")))\nlabel_files = sorted(glob.glob(os.path.join(MASK_DIR, \"*.npy\")))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_train= np.load(train_files[0])\nvolume_train = data_train\n# Load the very first file from your list\nprint(f\"1. Shape:     {volume_train.shape}\")\nprint(f\"2. Data Type: {volume_train.dtype}\")\nprint(f\"3. Max Value: {volume_train.max()}\")\nprint(f\"4. Min Value: {volume_train.min()}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_label= np.load(label_files[0])\nvolume_label = data_label\n# Load the very first file from your list\nprint(f\"1. Shape:     {volume_label.shape}\")\nprint(f\"2. Data Type: {volume_label.dtype}\")\nprint(f\"3. Max Value: {volume_label.max()}\")\nprint(f\"4. Min Value: {volume_label.min()}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.device_count()","metadata":{"trusted":true},"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_label[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_label.shape[0] - 1));","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass VesuviusCubeDataset(Dataset):\n    def __init__(self, img_dir, mask_dir, ids, mode=\"train\", patch_size=128, enable_occlusion=False):\n        \"\"\"\n        Args:\n            img_dir: Path to train_images\n            mask_dir: Path to train_labels\n            ids: List of file indices (e.g. [0, 1, 2...]) to use\n            mode: 'train' or 'val'\n            patch_size: Size of the 3D crop for training (default 224)\n        \"\"\"\n        self.img_dir = img_dir\n        self.mask_dir = mask_dir\n        self.mode = mode\n        self.patch_size = patch_size\n        self.enable_occlusion = enable_occlusion\n\n        #get list of all .npy files\n        self.all_files = sorted([f for f in os.listdir(img_dir) if f.endswith('.npy')])\n        self.files = [self.all_files[i] for i in ids if i < len(self.all_files)]\n\n    def __len__(self):\n        return len(self.files)\n\n    def __getitem__(self, idx):\n        filename = self.files[idx]\n        \n        # 1. Load Data\n        vol_path = os.path.join(self.img_dir, filename)\n        mask_path = os.path.join(self.mask_dir, filename)\n        \n        volume = np.load(vol_path)   # Shape: (320, 320, 320)\n        mask = np.load(mask_path)    # Shape: (320, 320, 320)\n        \n        # 2. Normalize (uint8 -> 0.0 to 1.0)\n        volume = volume.astype(np.float32) / 255.0\n        \n        # 3. Create Targets\n        target = (mask == 1).astype(np.float32)\n        valid_mask = (mask != 2).astype(np.float32)\n        \n        # TRAIN MODE: Patching + Augmentation\n        if self.mode == 'train':\n            d, h, w = volume.shape\n            ps = self.patch_size\n            \n            # Biased sampling\n            surface_indices =  np.argwhere(mask==1)\n            \n            # 60% chance to force crop on the sheet (Foreground)\n            if len(surface_indices)>0 and np.random.rand()<0.6:\n                center = surface_indices[np.random.randint(len(surface_indices))]\n                z = max(0, min(d - ps, center[0] - ps // 2))       \n                y = max(0, min(h - ps, center[1] - ps // 2))\n                x = max(0, min(w - ps, center[2] - ps // 2))\n            #40% chance of background\n            else:\n                z = np.random.randint(0, max(1, d - ps))\n                y = np.random.randint(0, max(1, h - ps))\n                x = np.random.randint(0, max(1, w - ps))\n                \n            #crop\n            volume = volume[z:z+ps, y:y+ps, x:x+ps]\n            target = target[z:z+ps, y:y+ps, x:x+ps]\n            valid_mask = valid_mask[z:z+ps, y:y+ps, x:x+ps]\n\n            # 1.Augmentation: Random Flips\n            if np.random.rand() > 0.5: # Vertical\n                volume = np.flip(volume, axis=1)\n                target = np.flip(target, axis=1)\n                valid_mask = np.flip(valid_mask, axis=1)\n            \n            if np.random.rand() > 0.5: # Horizontal\n                volume = np.flip(volume, axis=2)\n                target = np.flip(target, axis=2)\n                valid_mask = np.flip(valid_mask, axis=2)\n                \n            if np.random.rand() > 0.5: # Depth Flip (Invert Scroll)\n                volume = np.flip(volume, axis=0)\n                target = np.flip(target, axis=0)\n                valid_mask = np.flip(valid_mask, axis=0)\n                \n            # 2.Augmentation: Random 90-Degree Rotations\n            k = np.random.randint(0, 4) # 0, 1, 2, 3 rotations\n            if k > 0:\n                # Rotate axes 1 and 2 (Height and Width)\n                volume = np.rot90(volume, k=k, axes=(1, 2))\n                target = np.rot90(target, k=k, axes=(1, 2))\n                valid_mask = np.rot90(valid_mask, k=k, axes=(1, 2))\n\n            # 3.Augmentation: Occlusion\n            # Adding random black boxes to INPUT, but NOT to TARGET.\n            if np.random.rand() < 0.5 and self.enable_occlusion: # 50% chance\n                num_holes = np.random.randint(1, 4) # 1 to 3 holes\n                for _ in range(num_holes):\n                    # Random size (10 to 40 pixels)\n                    h_size = np.random.randint(2, 8)\n                    \n                    # Random position\n                    oz = np.random.randint(0, max(1, ps - h_size))\n                    oy = np.random.randint(0, max(1, ps - h_size))\n                    ox = np.random.randint(0, max(1, ps - h_size))\n                    \n                    # Set ONLY the input volume to 0 (Black)\n                    # We leave the target ALONE so the model learns to \"fill in\" the missing data\n                    volume[oz:oz+h_size, oy:oy+h_size, ox:ox+h_size] = 0.0\n\n        # VAL MODE: Return Full Volume (No Cropping)\n        # 4. Final Formatting (Add Channel Dimension)\n        # Must use .copy() to fix negative strides from flipping\n        return (\n            torch.from_numpy(volume.copy()).unsqueeze(0).float(),      # (1, D, H, W)\n            torch.from_numpy(target.copy()).unsqueeze(0).float(),      # (1, D, H, W)\n            torch.from_numpy(valid_mask.copy()).unsqueeze(0).float()   # (1, D, H, W)\n        )\n\n            ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# 2. GET ALL FILE INDICES\n# We list all files to count them, then create a list of numbers [0, 1, 2, ... N]\nall_files = sorted([f for f in os.listdir(DATA_DIR) if f.endswith('.npy')])\nall_ids = list(range(len(all_files)))\n\n# 3. CALCULATE SPLIT POINT (90% / 10%)\nsplit_pct = 0.90\nsplit_index = int(len(all_ids) * split_pct)\n\n# 4. PERFORM CONTIGUOUS SPLIT\n# Train gets the first 90% (e.g., top of the scroll)\n# Val gets the last 10% (e.g., bottom of the scroll)\ntrain_ids = all_ids[:split_index]\nval_ids = all_ids[split_index:]\n\n# Anchor Step\n# Picking the first 4 volumes of the validation set to chek every epoch\n# We'll check the 81 volumes of the validation set only after training fininshes(after 20 epochs)\nanchor_val_ids = val_ids[:4]\n\nprint(f\"Total Volumes: {len(all_ids)}\")\nprint(f\"Training on:   {len(train_ids)} volumes (Indices: {train_ids[0]} to {train_ids[-1]})\")\nprint(f\"Validating on: {len(val_ids)} volumes (Indices: {val_ids[0]} to {val_ids[-1]})\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nPATCH_SIZE = 128     \nBATCH_SIZE = 8       \nNUM_WORKERS = 4      # 2 workers per GPU\nPIN_MEMORY = True    \nPERSISTENT = True # Keeps workers alive (Speed Boost)\nPREFETCH_FACTOR = 2\n\ntrain_ds = VesuviusCubeDataset(\n    img_dir=DATA_DIR, \n    mask_dir=MASK_DIR, \n    ids=train_ids, \n    mode=\"train\", \n    patch_size=PATCH_SIZE\n)\n\nfull_val_ds = VesuviusCubeDataset(\n    img_dir=DATA_DIR, \n    mask_dir=MASK_DIR, \n    ids=val_ids, \n    mode=\"val\")\n\nanchor_val_ds = VesuviusCubeDataset(\n    img_dir=DATA_DIR, \n    mask_dir=MASK_DIR, \n    ids=anchor_val_ids, \n    mode=\"val\")\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=BATCH_SIZE,\n    shuffle=True,             \n    num_workers=NUM_WORKERS,\n    pin_memory=PIN_MEMORY,\n    drop_last=True,           # Good for BatchNorm stability\n    persistent_workers=True,  \n    prefetch_factor=PREFETCH_FACTOR\n)\n\nfull_val_loader = DataLoader(\n    full_val_ds,\n    batch_size=1,             # Must be 1 for Sliding Window Validation\n    shuffle=False,            \n    num_workers=NUM_WORKERS,\n    pin_memory=PIN_MEMORY,\n    persistent_workers=True, \n    prefetch_factor=PREFETCH_FACTOR\n)\n\nanchor_val_loader = DataLoader(\n    anchor_val_ds,\n    batch_size=1,             # Must be 1 for Sliding Window Validation\n    shuffle=False,            \n    num_workers=NUM_WORKERS,\n    pin_memory=PIN_MEMORY,\n    persistent_workers=True, \n    prefetch_factor=PREFETCH_FACTOR\n)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#IMPORTS\nimport torch.optim as optim\nfrom torch.cuda.amp import autocast, GradScaler\nfrom monai.inferers import sliding_window_inference\nfrom monai.metrics import compute_surface_dice\nfrom scipy.ndimage import label as nd_label\nfrom tqdm import tqdm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# CONFIGURATION:\nCONFIG = {\n    \"DEVICE\": \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    \"DATA_DIR\": \"/kaggle/input/vesuvius-npy/train_images\",\n    \"MASK_DIR\": \"/kaggle/input/vesuvius-npy/train_labels\",\n    \n    # Kaggle T4 Optimization\n    \"PATCH_SIZE\": 128,      # 128^3 fits comfortably in memory\n    \"BATCH_SIZE\": 8,        # 4 samples per GPU (Dual T4)\n    \"LR\": 1e-4,             \n    \"EPOCHS\": 120,\n    \"VAL_INTERVAL\": 10,\n    \"NUM_WORKERS\": 4,\n    \"PIN_MEMORY\": True,\n    \"PERSISTENT_WORKERS\": True\n}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # METRICS & UTILS\n# def calculate_voi_score(pred_mask, true_mask, alpha = 0.3,threshold = 0.5):\n#     \"\"\"\n#     Calculates Variation of Information (VOI).\n#     Measures if the SHEET is fragmented (Split) or accidentally merged.\n#     Input: (D, H, W) numpy arrays\n#     \"\"\"\n#     if pred_mask.max() > 1.0: # Check if logits\n#         pred_mask = torch.sigmoid(pred_mask)\n        \n#     p = (pred_mask > threshold).detach().cpu().numpy().astype(np.int32)\n#     t = (true_mask > threshold).detach().cpu().numpy().astype(np.int32)\n\n#     # Connected Componenets\n#     p_labeled, p_n = nd_label(p, structure=np.ones((3,3,3)))\n#     t_labeled, t_n = nd_label(t, structure=np.ones((3,3,3)))\n\n#     # Joint Histogram\n#     joint_hist = np.histogram2d(\n#         p_labeled.flatten(), t_labeled.flatten(), \n#         bins=[p_n+1, t_n+1], range=[[0, p_n+1], [0, t_n+1]]\n#     )[0]\n\n#     joint_prob = joint_hist[1:, 1:] \n#     total = joint_prob.sum()\n#     if total == 0: return 0.0\n#     joint_prob /= total\n    \n#     p_prob = joint_prob.sum(axis=1)\n#     t_prob = joint_prob.sum(axis=0)\n\n#     # Entropy Calculation\n#     h_p = -np.sum(p_prob[p_prob > 0] * np.log2(p_prob[p_prob > 0]))\n#     h_t = -np.sum(t_prob[t_prob > 0] * np.log2(t_prob[t_prob > 0]))\n#     h_pt = -np.sum(joint_prob[joint_prob > 0] * np.log2(joint_prob[joint_prob > 0]))\n\n#     voi_split = h_pt - h_p\n#     voi_merge = h_pt - h_t\n#     return 1.0 / (1.0 + alpha * (voi_split + voi_merge))\n\n# class SoftDiceCLDice(nn.Module):\n#     def __init__(self, iter_=3, alpha=0.5, smooth=1e-6):\n#         \"\"\"\n#         Combined Dice + clDice Loss\n#         \"\"\" \n#         super(SoftDiceCLDice, self).__init__()\n#         self.iter = iter_\n#         self.alpha = alpha\n#         self.smooth = smooth\n\n#     def forward(self, pred, target, valid_mask):\n#         \"\"\"\n#         Args:\n#             pred (Tensor): Raw Logits from model [B, C, D, H, W]\n#             target (Tensor): Ground Truth [B, C, D, H, W] (0 or 1)\n#             valid_mask (Tensor): Mask [B, C, D, H, W] (1=Valid, 0=Ignore)\n#         \"\"\"\n#         probs = torch.sigmoid(pred)\n\n#         # Masking\n#         probs = probs * valid_mask\n#         target = target * valid_mask\n\n#         # Dice loss\n#         dice_loss = self.compute_dice_loss(probs, target)\n\n#         # CLDice loss\n#         cldice_loss = self.compute_cldice_loss(probs, target)\n\n#         # Weighted sum\n#         return (1.0 - self.alpha) * dice_loss + self.alpha * cldice_loss\n\n#     def compute_dice_loss(self, probs, target):\n#         # Flatten\n#         p_flat = probs.reshape(probs.size(0), -1)\n#         t_flat = target.reshape(target.size(0), -1)\n\n#         # Intersection and union\n#         intersection = (p_flat * t_flat).sum(dim=1)\n#         union  = p_flat.sum(dim=1) + t_flat.sum(dim=1)\n\n#         dice = (2.0 * intersection + self.smooth) / (union + self.smooth)\n\n#         return 1.0 - dice.mean()\n\n#     def compute_cldice_loss(self, probs, target):\n#         # Skeletonize \n#         skel_pred = self.soft_skeletonize(probs)\n#         skel_true = self.soft_skeletonize(target)\n\n#         # Topology precision\n#         t_prec = (skel_pred * target).sum() / (skel_pred.sum() + self.smooth)\n\n#         # Topology Sensitivity\n#         t_sens = (skel_true * probs).sum() / (skel_true.sum() + self.smooth)\n\n#         # CLDice (Harmonic mean)\n#         cl_dice = 2.0 * (t_prec * t_sens) / (t_prec + t_sens + self.smooth)\n\n#         return 1.0 - cl_dice\n\n#     def soft_skeletonize(self, x):\n#         \"\"\"\n#         Iterative Skeletonization via Morphological Opening (Top-Hat).\n#         Skeleton = Image - Open(Image)\n#         \"\"\"\n#         p1 = self.soft_erode(x)\n#         p2 = self.soft_dilate(p1) # This sequence (Erode -> Dilate) is \"Opening\"\n\n#         return F.relu(x - p2)\n\n#     def soft_erode(self, x):\n#         \"\"\"\n#         Differentiable Erosion using Min-Pooling.\n#         \"\"\"\n#         for i in range(self.iter):\n#             x = -F.max_pool3d(-x, kernel_size=3, stride=1, padding=1)   \n#         return x\n\n#     def soft_dilate(self, x):\n#         \"\"\"\n#         Differentiable Dilation using Max-Pooling.\n#         \"\"\"\n#         for i in range(self.iter):\n#             x = F.max_pool3d(x, kernel_size=3, stride=1, padding=1)\n#         return x\n        \n         ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1e-6):\n        super(DiceLoss, self).__init__()\n        self.smooth = smooth\n\n    def forward(self, pred, target, valid_mask=None):\n        \"\"\"\n        Args:\n            pred (Tensor): Raw Logits from model [B, C, D, H, W]\n            target (Tensor): Ground Truth [B, C, D, H, W] (0 or 1)\n            valid_mask (Tensor, optional): Mask [B, C, D, H, W] (1=Valid, 0=Ignore)\n        \"\"\"\n        # Apply Sigmoid to logits\n        probs = torch.sigmoid(pred)\n\n        # Apply mask if provided\n        if valid_mask is not None:\n            probs = probs * valid_mask\n            target = target * valid_mask\n\n        # Flatten: [B, C, D, H, W] -> [B, N]\n        p_flat = probs.reshape(probs.size(0), -1)\n        t_flat = target.reshape(target.size(0), -1)\n\n        # Dice Calculation\n        intersection = (p_flat * t_flat).sum(dim=1)\n        union = p_flat.sum(dim=1) + t_flat.sum(dim=1)\n\n        dice = (2.0 * intersection + self.smooth) / (union + self.smooth)\n\n        return 1.0 - dice.mean()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MAIN EXECUTION\n\n# Training and Validation\nif True:\n    \n    # Model Setup\n    try:\n        model = SegResMamba(\n            spatial_dims=3, \n            in_chans=1,     # Input: Grayscale Volume\n            out_chans=1,    # Output: Binary Sheet Mask\n        ).to(CONFIG[\"DEVICE\"])\n        \n        if torch.cuda.device_count() > 1:\n            print(f\"Dual GPU Detected: Using {torch.cuda.device_count()} GPUs\")\n            model = nn.DataParallel(model)\n            \n    except NameError:\n        print(\" Error: SegResMamba class not found. Ensure you ran the import cell.\")\n        sys.exit(1)\n\n    # Optimizer & Scaler\n    optimizer = optim.AdamW(model.parameters(), lr=CONFIG[\"LR\"], weight_decay=1e-5, eps=1e-5)\n    loss_function = DiceLoss().cuda() \n    scaler = torch.amp.GradScaler('cuda') # Update for newer PyTorch versions if needed\n    best_metric = -1\n\n    # Start Training\n    print(\"\\n Starting Training...\")\n\n    for epoch in range(CONFIG[\"EPOCHS\"]):\n        print(f\"\\n Epoch {epoch + 1}/{CONFIG['EPOCHS']}\")\n\n        # --- CURRICULUM SWITCH ---\n        # Fixed: Matched comment to code. Assuming you want Epoch 25.\n        if epoch == 24: # Epoch 0-indexed, so 24 is the 25th epoch\n            print(f\"\\n>>> CURRICULUM UPDATE (Epoch {epoch+1}): Enabling Occlusion Augmentation! <<<\")\n            \n            # Re-create Dataset/Loader logic...\n            train_ds = VesuviusCubeDataset(\n                img_dir=DATA_DIR, \n                mask_dir=MASK_DIR, \n                ids=train_ids, \n                mode=\"train\", \n                patch_size=PATCH_SIZE,\n                enable_occlusion=True \n            )\n            train_loader = DataLoader(\n                train_ds,\n                batch_size=BATCH_SIZE,\n                shuffle=True,             \n                num_workers=NUM_WORKERS,\n                pin_memory=PIN_MEMORY,\n                drop_last=True,\n                persistent_workers=True,\n                prefetch_factor=PREFETCH_FACTOR\n            )\n\n        # --- TRAINING PHASE ---\n        model.train()\n        epoch_loss = 0\n        step = 0\n        pbar = tqdm(train_loader, desc=\"Training\")\n        \n        for batch_data in pbar:\n            inputs = batch_data[0].to(CONFIG[\"DEVICE\"])\n            labels = batch_data[1].to(CONFIG[\"DEVICE\"])\n            valid_mask = batch_data[2].to(CONFIG[\"DEVICE\"])\n            \n            optimizer.zero_grad()\n            \n            # Forward (Mixed Precision)\n            with torch.amp.autocast('cuda'):\n                outputs = model(inputs)\n                loss = loss_function(outputs, labels, valid_mask)\n            \n            if torch.isnan(loss):\n                print(\" NaN Loss detected! Skipping batch.\")\n                optimizer.zero_grad()\n                continue\n            \n            # Backward\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=2.0)\n            scaler.step(optimizer)\n            scaler.update()\n            \n            epoch_loss += loss.item()\n            step += 1\n            pbar.set_postfix({\"Loss\": f\"{loss.item():.4f}\"})\n            \n            \n        print(f\" Avg Train Loss: {epoch_loss/max(step,1):.4f}\")\n\n        # SAVE CHECKPOINTS\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': epoch_loss,\n        }, \"last_checkpoint.pth\")\n        \n        if (epoch + 1) % 5 == 0:\n            ckpt_name = f\"checkpoint_epoch_{epoch+1}.pth\"\n            torch.save(model.state_dict(), ckpt_name)\n            print(f\" [Checkpoint] Saved {ckpt_name}\")\n\n        # --- VALIDATION PHASE ---\n        if (epoch + 1) % CONFIG[\"VAL_INTERVAL\"] == 0:\n            model.eval()\n            val_dice_scores = []\n            \n            with torch.no_grad():\n                for val_data in tqdm(anchor_val_loader, desc=\"Quick Validation\"):\n                    val_inputs = val_data[0].to(CONFIG[\"DEVICE\"])\n                    val_labels = val_data[1].to(CONFIG[\"DEVICE\"])\n                    val_valid  = val_data[2].to(CONFIG[\"DEVICE\"])\n                    \n                    val_outputs = sliding_window_inference(\n                        inputs=val_inputs, roi_size=(128, 128, 128), \n                        sw_batch_size=4, predictor=model, overlap=0.5\n                    )\n                    \n                    # Post-Processing\n                    probs = torch.sigmoid(val_outputs)\n                    # No need to threshold for DiceLoss, it handles probs directly if designed correctly\n                    # But if you want exact \"Dice Score\" of binary output:\n                    preds = (probs > 0.5).float() \n                    \n                    # --- FIX: Calculate Score, not Loss ---\n                    # We reuse loss_function (DiceLoss)\n                    # DiceLoss = 1 - Dice => Dice = 1 - DiceLoss\n                    \n                    # Note: Passing 'preds' (0 or 1) to DiceLoss is fine if it expects logits \n                    # BUT your DiceLoss has torch.sigmoid inside it!\n                    # If you pass binary preds (0/1) to a function that applies sigmoid, \n                    # 0 becomes 0.5, 1 becomes ~0.73. THAT IS WRONG.\n                    \n                    d_loss = loss_function(val_outputs, val_labels, val_valid)\n                    dice_score = 1.0 - d_loss.item()\n                    \n                    if not torch.isnan(d_loss): val_dice_scores.append(dice_score)\n                    else: val_dice_scores.append(0.0)\n            \n            avg_dice = np.mean(val_dice_scores)\n            print(f\"Val Dice Score : {avg_dice: .4f}\")\n            \n            if avg_dice > best_metric:\n                best_metric = avg_dice\n                torch.save(model.state_dict(), \"best_surface_model.pth\")\n                print(f\" New Best Model Saved! ({best_metric:.4f})\")\n                \n    print(\"\\nTraining Complete. Running FINAL Full Evaluation on 81 Volumes...\")\n\n    # --- FINAL FULL CHECK ---\n    model.load_state_dict(torch.load(\"best_surface_model.pth\"))\n    model.eval()\n    \n    final_dice_scores = []\n    \n    with torch.no_grad():\n        for val_data in tqdm(full_val_loader, desc=\"Final Full Evaluation\"):\n            val_inputs = val_data[0].to(CONFIG[\"DEVICE\"])\n            val_labels = val_data[1].to(CONFIG[\"DEVICE\"])\n            val_valid  = val_data[2].to(CONFIG[\"DEVICE\"])\n            \n            val_outputs = sliding_window_inference(\n                inputs=val_inputs, roi_size=(128, 128, 128), \n                sw_batch_size=4, predictor=model, overlap=0.5\n            )\n            \n            # Same Fix: Calculate Score from Loss\n            d_loss = loss_function(val_outputs, val_labels, val_valid)\n            dice_score = 1.0 - d_loss.item()\n\n            if not torch.isnan(d_loss): final_dice_scores.append(dice_score)\n            else: final_dice_scores.append(0.0)\n\nfinal_avg_dice = np.mean(final_dice_scores)\n\nprint(\"-\" * 30)\nprint(f\"FINAL RESULT ON {len(full_val_loader)} VOLUMES\")\nprint(f\"Surface Dice: {final_avg_dice:.4f}\")\nprint(\"-\" * 30)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}