{"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":99552,"databundleVersionId":13851420,"sourceType":"competition"},{"sourceId":13866621,"sourceType":"datasetVersion","datasetId":8834739},{"sourceId":13913205,"sourceType":"datasetVersion","datasetId":8834618},{"sourceId":13927856,"sourceType":"datasetVersion","datasetId":8858929}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q monai","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:18:15.272904Z","iopub.execute_input":"2025-11-30T15:18:15.273502Z","iopub.status.idle":"2025-11-30T15:19:37.883187Z","shell.execute_reply.started":"2025-11-30T15:18:15.273476Z","shell.execute_reply":"2025-11-30T15:19:37.882523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Libraries from Demo\nimport os\nimport shutil\nfrom collections import defaultdict\n\nimport pandas as pd\nimport polars as pl\nimport pydicom as dicom\n\n\n#Libraries from attempt\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom monai.losses import DiceLoss\nimport random\nimport scipy.ndimage as ndi\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\n\nfrom collections import Counter\nfrom scipy import ndimage\n\nfrom scipy.ndimage import zoom as ndi_zoom\nfrom sklearn.model_selection import train_test_split, StratifiedShuffleSplit\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom typing import Tuple, List\n\n\nfrom sklearn.preprocessing import StandardScaler","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:19:37.884428Z","iopub.execute_input":"2025-11-30T15:19:37.884795Z","iopub.status.idle":"2025-11-30T15:20:15.391027Z","shell.execute_reply.started":"2025-11-30T15:19:37.884765Z","shell.execute_reply":"2025-11-30T15:20:15.390468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    \"\"\"\n    Set random seeds for reproducibility in deep learning projects.\n    \n    Args:\n        seed (int): Random seed value (default: 42)\n    \"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)  # if using multi-GPU\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nseed_everything()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.391791Z","iopub.execute_input":"2025-11-30T15:20:15.392499Z","iopub.status.idle":"2025-11-30T15:20:15.40117Z","shell.execute_reply.started":"2025-11-30T15:20:15.392477Z","shell.execute_reply":"2025-11-30T15:20:15.400434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ID_COL = 'SeriesInstanceUID'\nLABEL_COLS = [\n    'Left Infraclinoid Internal Carotid Artery',\n    'Right Infraclinoid Internal Carotid Artery',\n    'Left Supraclinoid Internal Carotid Artery',\n    'Right Supraclinoid Internal Carotid Artery',\n    'Left Middle Cerebral Artery',\n    'Right Middle Cerebral Artery',\n    'Anterior Communicating Artery',\n    'Left Anterior Cerebral Artery',\n    'Right Anterior Cerebral Artery',\n    'Left Posterior Communicating Artery',\n    'Right Posterior Communicating Artery',\n    'Basilar Tip',\n    'Other Posterior Circulation',\n    'Aneurysm Present',\n]\n\nNUM_LABELS = len(LABEL_COLS) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.402858Z","iopub.execute_input":"2025-11-30T15:20:15.403074Z","iopub.status.idle":"2025-11-30T15:20:15.41563Z","shell.execute_reply.started":"2025-11-30T15:20:15.40306Z","shell.execute_reply":"2025-11-30T15:20:15.414917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nTRAIN_CSV = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\"\ntest_frac = 0.2\nval_frac = 0.1\nval_frac_within_trainval = val_frac / (1 - test_frac)\nseed = 42  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.4164Z","iopub.execute_input":"2025-11-30T15:20:15.416639Z","iopub.status.idle":"2025-11-30T15:20:15.430484Z","shell.execute_reply.started":"2025-11-30T15:20:15.416618Z","shell.execute_reply":"2025-11-30T15:20:15.429789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(TRAIN_CSV)\nprint(f\"On the original Dataset, the percentage of aneurysms is: {100 * sum(train['Aneurysm Present'])/len(train)}%\")\nprint(f\"The original dataset has {len(train)} samples.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.431208Z","iopub.execute_input":"2025-11-30T15:20:15.431429Z","iopub.status.idle":"2025-11-30T15:20:15.490201Z","shell.execute_reply.started":"2025-11-30T15:20:15.431409Z","shell.execute_reply":"2025-11-30T15:20:15.489635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"excluded = np.load('/kaggle/input/succesful/usefull.npz')[\"lst\"]\nprint(len(excluded))\ntrain = train[train['SeriesInstanceUID'].isin(excluded)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.490797Z","iopub.execute_input":"2025-11-30T15:20:15.491008Z","iopub.status.idle":"2025-11-30T15:20:15.510284Z","shell.execute_reply.started":"2025-11-30T15:20:15.490989Z","shell.execute_reply":"2025-11-30T15:20:15.509702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y = train[\"Aneurysm Present\"].astype(int).values\nsss = StratifiedShuffleSplit(n_splits=1, test_size=test_frac, random_state=seed)\n(trainval_idx, test_idx), = sss.split(np.zeros(len(y)), y)\ny_trainval = y[trainval_idx]\n\nsss_2 = StratifiedShuffleSplit(n_splits=1, test_size=val_frac_within_trainval, random_state=seed)\n(train_rel_idx, val_rel_idx), = sss_2.split(np.zeros(len(y_trainval)), y_trainval)\ntrain_idx = trainval_idx[train_rel_idx]\nval_idx   = trainval_idx[val_rel_idx]\n\ntrain_ds = Subset(train, train_idx.tolist())\nval_ds   = Subset(train, val_idx.tolist())\ntest_ds  = Subset(train, test_idx.tolist())\n\nprint(f\"Train/Val/Test sizes: {len(train_ds)} / {len(val_ds)} / {len(test_ds)}\")\nprint(len(train.iloc[train_idx]), len(train.iloc[val_idx]), len(train.iloc[test_idx]))\nprint(\"Train positive rate:\", train.iloc[train_idx][\"Aneurysm Present\"].mean())\nprint(\"Val   positive rate:\", train.iloc[val_idx][\"Aneurysm Present\"].mean())\nprint(\"Test   positive rate:\", train.iloc[test_idx][\"Aneurysm Present\"].mean())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.51093Z","iopub.execute_input":"2025-11-30T15:20:15.511143Z","iopub.status.idle":"2025-11-30T15:20:15.527038Z","shell.execute_reply.started":"2025-11-30T15:20:15.511127Z","shell.execute_reply":"2025-11-30T15:20:15.526287Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Datasets","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom pathlib import Path\nimport os, numpy as np, torch\n\nfrom torch.utils.data import Dataset\nimport os, numpy as np, torch\n\nclass CachedVolumeDataset(Dataset):\n    def __init__(self, df, vols_dirs, mask_dirs, id_col, transform=None):\n        self.df = df.reset_index(drop=True).copy()\n        self.vols_dirs = list(vols_dirs)\n        self.mask_dirs = list(mask_dirs)\n        self.id_col = id_col\n        self.transform = transform\n\n        assert len(self.vols_dirs) == len(self.mask_dirs), \\\n            \"Volume and mask directory lists must be the same length\"\n\n    def __len__(self):\n        return len(self.df)\n\n    def _find_files(self, sid):\n        fname = f\"{sid}.npz\"\n\n        # Iterate over corresponding (vol_dir, mask_dir) pairs\n        for vol_dir, mask_dir in zip(self.vols_dirs, self.mask_dirs):\n\n            vol_path  = os.path.join(vol_dir,  fname)\n            mask_path = os.path.join(mask_dir, fname)\n\n            if os.path.exists(vol_path) and os.path.exists(mask_path):\n                return vol_path, mask_path\n\n        return None\n\n    def __getitem__(self, idx):\n        sid = str(self.df[self.id_col].iloc[idx])\n\n        paths = self._find_files(sid)\n        if paths is None:\n            raise FileNotFoundError(\n                f\"Missing volume/mask for ID {sid}\\n\"\n                f\"Searched in:\\n\"\n                + \"\\n\".join(self.vols_dirs + self.mask_dirs)\n            )\n\n        vol_path, mask_path = paths\n\n        # Load files\n        vol  = np.load(vol_path)[\"vol\"].astype(np.float32)      # [D, H, W]\n        mask = np.load(mask_path)[\"vol\"].astype(np.float32)    # [D, H, W]\n\n        # Convert to PyTorch tensors with channel dimension\n        x = torch.from_numpy(vol).unsqueeze(0)    # [1, D, H, W]\n        y = torch.from_numpy(mask).unsqueeze(0)   # [1, D, H, W]\n\n        if self.transform:\n            x, y = self.transform(x, y)\n\n        return x, y\n\n    def verify(self):\n        missing = []\n        for sid in self.df[self.id_col].astype(str):\n            if self._find_files(sid) is None:\n                missing.append(sid)\n        return missing\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.527811Z","iopub.execute_input":"2025-11-30T15:20:15.528053Z","iopub.status.idle":"2025-11-30T15:20:15.537054Z","shell.execute_reply.started":"2025-11-30T15:20:15.528037Z","shell.execute_reply":"2025-11-30T15:20:15.536438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PROCESSED_DATA_DIRS_MASKS = [\n    \"/kaggle/input/binary-masks-dataset/masks_quart1_2_5\",\n    \"/kaggle/input/binary-masks-dataset/masks_quart2_2_5\",\n    \"/kaggle/input/binary-masks-dataset/masks_quart3_2_5\",\n    \"/kaggle/input/binary-masks-dataset/masks_quart4_2_5\",\n]\n\nPROCESSED_DATA_DIRS_VOLS = [\n    \"/kaggle/input/vol-dataset-1-25-iso/vols_quart1 (1)\",\n    \"/kaggle/input/vol-dataset-1-25-iso/vols_quart2 (1)\",\n    \"/kaggle/input/vol-dataset-1-25-iso/vols_quart3 (1)\",\n    \"/kaggle/input/vol-dataset-1-25-iso/vols_quart4 (1)\",\n]\n\ntrain_ds = CachedVolumeDataset(\n    df=train.iloc[train_idx],\n    vols_dirs=PROCESSED_DATA_DIRS_VOLS,\n    mask_dirs=PROCESSED_DATA_DIRS_MASKS,\n    id_col=ID_COL\n)\n\nvalid_ds = CachedVolumeDataset(\n    df=train.iloc[val_idx],\n    vols_dirs=PROCESSED_DATA_DIRS_VOLS,\n    mask_dirs=PROCESSED_DATA_DIRS_MASKS,\n    id_col=ID_COL\n)\n\ntest_ds = CachedVolumeDataset(\n    df=train.iloc[test_idx],\n    vols_dirs=PROCESSED_DATA_DIRS_VOLS,\n    mask_dirs=PROCESSED_DATA_DIRS_MASKS,\n    id_col=ID_COL\n)\nprint(len(train_ds), len(valid_ds), len(test_ds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.538877Z","iopub.execute_input":"2025-11-30T15:20:15.539301Z","iopub.status.idle":"2025-11-30T15:20:15.558786Z","shell.execute_reply.started":"2025-11-30T15:20:15.539285Z","shell.execute_reply":"2025-11-30T15:20:15.558199Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Reviewing X=volume, Y=mask","metadata":{}},{"cell_type":"code","source":"x, y = train_ds[95]\nprint(\"Volume:\", x.shape)\nprint(\"Mask:\", y.shape)\n\nplt.figure(figsize=(10,4))\nplt.subplot(1,3,1)\nplt.imshow(x[0, 57].numpy())\nplt.title(\"Volume slice\")\n\nplt.subplot(1,3,2)\nplt.imshow(y[0, 57].numpy())\nplt.title(\"Mask slice\")\n\nplt.subplot(1,3,3)\nplt.imshow(x[0, 57].numpy() *  y[0, 57].numpy())\nplt.title(\"ROI slice\")\n\n# Find all slices with mask pixels\n# y shape = [C, D, H, W]\n# Step 1: Max over H and W\nmask_2d_max = y.max(dim=3)[0].max(dim=2)[0]   # shape: [1, D]\nmask_1d = mask_2d_max.squeeze(0)  # shape: [D]\n\n# Step 3: Get indices of slices containing mask\nmask_slices = (mask_1d == 1).nonzero(as_tuple=True)[0]\n\nprint(\"Slices containing mask pixels:\", mask_slices)\nprint(\"Number of slices with mask:\", len(mask_slices))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:15.559387Z","iopub.execute_input":"2025-11-30T15:20:15.559603Z","iopub.status.idle":"2025-11-30T15:20:16.266339Z","shell.execute_reply.started":"2025-11-30T15:20:15.559589Z","shell.execute_reply":"2025-11-30T15:20:16.265691Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class ConvBlock3D(nn.Module):\n    \"\"\"\n    3D convolutional downsampling block:\n    Conv3d -> ReLU -> Conv3d -> ReLU (+ optional dropout) (+ optional MaxPool3d)\n\n    Returns:\n      next_layer: tensor after optional max pooling\n      skip: tensor before pooling (for skip connection)\n    \"\"\"\n    def __init__(self, in_channels, out_channels, dropout_prob=0.0, use_maxpool=True):\n        super().__init__()\n        self.use_maxpool = use_maxpool\n\n        self.conv1 = nn.Conv3d(in_channels, out_channels, kernel_size=3, padding=1)\n        self.conv2 = nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm3d(out_channels)\n        self.bn2 = nn.BatchNorm3d(out_channels)\n\n        self.dropout = nn.Dropout3d(dropout_prob) if dropout_prob > 0 else None\n        self.pool = nn.MaxPool3d(kernel_size=2, stride=2) if use_maxpool else None\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = F.relu(x, inplace=True)\n\n        x = self.conv2(x)\n        x = self.bn2(x)\n        x = F.relu(x, inplace=True)\n\n        if self.dropout is not None:\n            x = self.dropout(x)\n\n        skip = x  # skip connection\n\n        if self.use_maxpool:\n            next_layer = self.pool(x)\n        else:\n            next_layer = x\n\n        return next_layer, skip\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:16.267049Z","iopub.execute_input":"2025-11-30T15:20:16.26731Z","iopub.status.idle":"2025-11-30T15:20:16.274254Z","shell.execute_reply.started":"2025-11-30T15:20:16.267293Z","shell.execute_reply":"2025-11-30T15:20:16.273543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UpBlock3D(nn.Module):\n    \"\"\"\n    Convolutional upsampling block\n    \n    Arguments:\n        expansive_input -- Input tensor from previous layer\n        contractive_input -- Input tensor from previous skip layer\n    Returns: \n        conv -- Tensor output\n    \"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n\n        # upsample by factor 2 in D, H, W\n        self.upconv = nn.ConvTranspose3d(\n            in_channels, out_channels,\n            kernel_size=2, stride=2\n        )\n\n        # after concatenation: out_channels (upsampled) + out_channels (skip)\n        self.conv1 = nn.Conv3d(out_channels * 2, out_channels, kernel_size=3, padding=1)\n        self.conv2 = nn.Conv3d(out_channels, out_channels, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm3d(out_channels)\n        self.bn2 = nn.BatchNorm3d(out_channels)\n\n    def forward(self, expansive_input, contractive_input):\n        # upsample\n        x = self.upconv(expansive_input)\n\n        # handle possible size mismatches in D, H, W (odd sizes)\n        if x.shape[-3:] != contractive_input.shape[-3:]:\n            diff_d = contractive_input.size(-3) - x.size(-3)\n            diff_h = contractive_input.size(-2) - x.size(-2)\n            diff_w = contractive_input.size(-1) - x.size(-1)\n\n            contractive_input = contractive_input[\n                :,\n                :,\n                diff_d // 2 : contractive_input.size(-3) - diff_d // 2,\n                diff_h // 2 : contractive_input.size(-2) - diff_h // 2,\n                diff_w // 2 : contractive_input.size(-1) - diff_w // 2,\n            ]\n\n        # concat along channels\n        x = torch.cat([contractive_input, x], dim=1)\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = F.relu(x, inplace=True)\n\n        x = self.conv2(x)\n        x = self.bn2(x)\n        x = F.relu(x, inplace=True)\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:16.275162Z","iopub.execute_input":"2025-11-30T15:20:16.27541Z","iopub.status.idle":"2025-11-30T15:20:16.307619Z","shell.execute_reply.started":"2025-11-30T15:20:16.275388Z","shell.execute_reply":"2025-11-30T15:20:16.306911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet3D(nn.Module):\n    def __init__(self, in_channels=1, n_filters=16, n_classes=1):\n        \"\"\"\n        in_channels: channels of input volume (e.g., 1 for CT/MRI)\n        n_filters: base number of filters (you may want 16 or 32 depending on GPU)\n        n_classes: output channels (1 for binary mask)\n        \"\"\"\n        super().__init__()\n\n        # Contracting path\n        self.enc1 = ConvBlock3D(in_channels,      n_filters,      dropout_prob=0.0, use_maxpool=True)\n        self.enc2 = ConvBlock3D(n_filters,        n_filters * 2,  dropout_prob=0.0, use_maxpool=True)\n        self.enc3 = ConvBlock3D(n_filters * 2,    n_filters * 4,  dropout_prob=0.0, use_maxpool=True)\n        self.enc4 = ConvBlock3D(n_filters * 4,    n_filters * 8,  dropout_prob=0.0, use_maxpool=True)\n\n        # Bottleneck (no maxpool)\n        self.bottleneck = ConvBlock3D(n_filters * 8, n_filters * 16,\n                                      dropout_prob=0.0, use_maxpool=False)\n\n        # Expansive path\n        self.up4 = UpBlock3D(n_filters * 16, n_filters * 8)   # connects to enc4 skip\n        self.up3 = UpBlock3D(n_filters * 8,  n_filters * 4)   # connects to enc3 skip\n        self.up2 = UpBlock3D(n_filters * 4,  n_filters * 2)   # connects to enc2 skip\n        self.up1 = UpBlock3D(n_filters * 2,  n_filters * 1)   # connects to enc1 skip\n\n        # Final 1x1x1 conv (logits)\n        self.final_conv = nn.Conv3d(n_filters, n_classes, kernel_size=1)\n\n    def forward(self, x):\n        # x: (B, C, D, H, W)\n\n        # Encoding path\n        x1, skip1 = self.enc1(x)\n        x2, skip2 = self.enc2(x1)\n        x3, skip3 = self.enc3(x2)\n        x4, skip4 = self.enc4(x3)\n\n        # Bottleneck\n        bottleneck, _ = self.bottleneck(x4)\n\n        # Decoding path\n        d4 = self.up4(bottleneck, skip4)\n        d3 = self.up3(d4,        skip3)\n        d2 = self.up2(d3,        skip2)\n        d1 = self.up1(d2,        skip1)\n\n        logits = self.final_conv(d1)  # (B, n_classes, D, H, W)\n        return logits\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:16.308489Z","iopub.execute_input":"2025-11-30T15:20:16.308673Z","iopub.status.idle":"2025-11-30T15:20:16.329464Z","shell.execute_reply.started":"2025-11-30T15:20:16.308659Z","shell.execute_reply":"2025-11-30T15:20:16.328789Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# BCE+DICE Loss","metadata":{}},{"cell_type":"code","source":"class BCEDiceMonai(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.dice = DiceLoss(\n            include_background=True,\n            sigmoid=True,\n            reduction=\"mean\",\n            smooth_nr=1e-6,\n            smooth_dr=1e-6,\n        )\n        self.bce  = nn.BCEWithLogitsLoss()\n\n    def forward(self, logits, targets):\n        return self.dice(logits, targets) + self.bce(logits, targets)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:16.330147Z","iopub.execute_input":"2025-11-30T15:20:16.330379Z","iopub.status.idle":"2025-11-30T15:20:16.346621Z","shell.execute_reply.started":"2025-11-30T15:20:16.330365Z","shell.execute_reply":"2025-11-30T15:20:16.345945Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def train_model_volumes(\n    train_ds,\n    val_ds,\n    test_ds,\n    num_epochs: int = 10,\n    batch_size: int = 1,       # small by default for 3D volumes\n    val_batch_size: int = 1,\n    test_batch_size: int = 1,\n    threshold: float = 0.5,\n    save_path: str = \"unet3d_seg_weights.pth\",\n    base_filters: int = 8,\n    lr: float = 1e-3,\n    weight_decay: float = 1e-4,\n    device: torch.device = DEVICE,\n):\n\n    # --- Model ---\n    model = UNet3D(\n        in_channels=1,      # CachedVolumeDataset gives [1, D, H, W]\n        n_filters=base_filters,\n        n_classes=1         # binary mask\n    ).to(device).float()\n\n    # --- Dataloaders ---\n    train_loader = DataLoader(\n        train_ds,\n        batch_size=batch_size,\n        shuffle=True,\n        num_workers=0,\n        pin_memory=True,\n    )\n\n    val_loader = DataLoader(\n        val_ds,\n        batch_size=val_batch_size,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=True,\n    )\n\n    test_loader = DataLoader(\n            test_ds,\n            batch_size=test_batch_size,\n            shuffle=False,\n            num_workers=0,\n            pin_memory=True,\n        )\n\n    # --- Loss / Optimizer / Scheduler ---\n    #criterion = nn.BCEWithLogitsLoss()\n    criterion = BCEDiceMonai()\n    optimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode=\"min\", patience=2, factor=0.5\n    )\n\n    best_val = float(\"inf\")\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_dice\": []}\n\n    print(\"Starting 3D U-Net training on volumes\")\n    for epoch in range(num_epochs):\n        # ===================== TRAIN =====================\n        model.train()\n        running_train_loss = 0.0\n\n        for batch_idx, (volumes, masks) in enumerate(train_loader):\n            # volumes, masks: [B, 1, D, H, W]\n            volumes = volumes.to(device, non_blocking=True).float()\n            masks   = masks.to(device, non_blocking=True).float()\n\n            optimizer.zero_grad(set_to_none=True)\n            logits = model(volumes)              # [B, 1, D, H, W]\n            loss   = criterion(logits, masks)\n\n            loss.backward()\n            optimizer.step()\n\n            running_train_loss += loss.item() * volumes.size(0)\n\n            if batch_idx % 50 == 0:\n                print(\n                    f\"Epoch {epoch+1}/{num_epochs} | \"\n                    f\"Batch {batch_idx}/{len(train_loader)} | \"\n                    f\"Loss {loss.item():.4f}\"\n                )\n\n        epoch_train_loss = running_train_loss / max(1, len(train_loader.dataset))\n        history[\"train_loss\"].append(epoch_train_loss)\n\n        # ===================== VALIDATION =====================\n        model.eval()\n        running_val_loss = 0.0\n        dice_sum = 0.0\n        n_samples = 0\n\n        with torch.no_grad():\n            for volumes, masks in val_loader:\n                volumes = volumes.to(device, non_blocking=True).float()\n                masks   = masks.to(device, non_blocking=True).float()\n\n                logits = model(volumes)\n                loss   = criterion(logits, masks)\n\n                running_val_loss += loss.item() * volumes.size(0)\n\n                # ---- Dice score over whole volumes ----\n                probs = torch.sigmoid(logits)\n                preds = (probs >= threshold).float()\n\n                # flatten per sample: [B, 1, D, H, W] -> sum over dims\n                smooth = 1e-6\n                intersection = (preds * masks).sum(dim=(1, 2, 3, 4))\n                union        = preds.sum(dim=(1, 2, 3, 4)) + masks.sum(dim=(1, 2, 3, 4))\n                dice         = (2.0 * intersection + smooth) / (union + smooth)\n                # mean over batch\n                dice_sum += dice.sum().item()\n                n_samples += volumes.size(0)\n\n        epoch_val_loss = running_val_loss / max(1, len(val_loader.dataset))\n        mean_dice      = dice_sum / max(1, n_samples)\n\n        history[\"val_loss\"].append(epoch_val_loss)\n        history[\"val_dice\"].append(mean_dice)\n\n        # step scheduler on validation loss\n        scheduler.step(epoch_val_loss)\n\n        print(\n            f\"Epoch {epoch+1}/{num_epochs} done | \"\n            f\"train_loss={epoch_train_loss:.4f} | \"\n            f\"val_loss={epoch_val_loss:.4f} | \"\n            f\"val_dice={mean_dice:.4f}\"\n        )\n\n        # Save best model by val loss\n        if epoch_val_loss < best_val:\n            best_val = epoch_val_loss\n            torch.save(model.state_dict(), f'epoch{epoch}_'+save_path)\n            print(f\"  ↳ New best model saved to {save_path}\")\n\n    #####TEST LOOP#######\n    print(\"\\nEvaluating on TEST set...\")\n    model.eval()\n    running_test_loss = 0.0\n    dice_sum = 0.0\n    n_samples = 0\n    example_pred_mask = None  \n    example_vol = None\n    with torch.no_grad():\n        for batch_idx, (volumes, masks) in enumerate(test_loader):\n            volumes = volumes.to(device, non_blocking=True).float()\n            masks   = masks.to(device, non_blocking=True).float()\n            logits = model(volumes)\n            loss   = criterion(logits, masks)\n            running_test_loss += loss.item() * volumes.size(0)\n            probs = torch.sigmoid(logits)\n            preds = (probs >= threshold).float()\n\n            smooth = 1e-6\n            intersection = (preds * masks).sum(dim=(1, 2, 3, 4))\n            union = preds.sum(dim=(1, 2, 3, 4)) + masks.sum(dim=(1, 2, 3, 4))\n            dice = (2.0 * intersection + smooth) / (union + smooth)\n            dice_sum += dice.sum().item()\n            n_samples += volumes.size(0)\n            \n            vol_example = volumes[0, 0].detach().cpu().numpy()\n            mask_example = masks[0, 0].detach().cpu().numpy()\n            pred_mask_example = preds[0, 0].detach().cpu().numpy()\n    \n    test_loss = running_test_loss / max(1, len(test_loader.dataset))\n    test_dice = dice_sum / max(1, n_samples)\n    print(f\"TEST results -> loss={test_loss:.4f} | dice={test_dice:.4f}\")\n    return model, history, vol_example, mask_example, pred_mask_example\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:16.347265Z","iopub.execute_input":"2025-11-30T15:20:16.347433Z","iopub.status.idle":"2025-11-30T15:20:16.365635Z","shell.execute_reply.started":"2025-11-30T15:20:16.34742Z","shell.execute_reply":"2025-11-30T15:20:16.365029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model, history, vol_example, mask_example, pred_mask_example = train_model_volumes(\n    train_ds=train_ds,\n    val_ds=valid_ds,\n    test_ds=test_ds,\n    num_epochs=8,\n    batch_size=6,\n    val_batch_size=3,\n    test_batch_size=1,\n    save_path=\"unet3d_seg.pth\"\n)\ntorch.save(history, \"unet3d_history.pth\")\ntorch.save(vol_example, \"unet3d_vol_example.pth\")\ntorch.save(mask_example, \"unet3d_mask_example.pth\")\ntorch.save(pred_mask_example, \"unet3d_pred_mask_example.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-30T15:20:16.366362Z","iopub.execute_input":"2025-11-30T15:20:16.366591Z","iopub.status.idle":"2025-11-30T15:21:23.74326Z","shell.execute_reply.started":"2025-11-30T15:20:16.366576Z","shell.execute_reply":"2025-11-30T15:21:23.742288Z"}},"outputs":[],"execution_count":null}]}