{"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":99552,"databundleVersionId":13851420,"sourceType":"competition"},{"sourceId":13848468,"sourceType":"datasetVersion","datasetId":8821011}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# import pydicom\n# import json\n\n# d = pydicom.dcmread(\"/kaggle/input/rsna-intracranial-aneurysm-detection/series/1.2.826.0.1.3680043.8.498.10004044428023505108375152878107656647/1.2.826.0.1.3680043.8.498.10124807242473374136099471315028464450.dcm\")\n# print(d)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import pandas as pd\n\n# # Replace 'your_file.csv' with the actual path to your CSV file\n# file_path = '/kaggle/input/rsna-intracranial-aneurysm-detection/train_localizers.csv' \n# df = pd.read_csv(file_path)\n\n# # Display the first few rows of the DataFrame\n# print(\"train_localizers.CSV Data (first 5 rows):\")\n# print(df.head())\n# print(\"train_localizers.CSV Data describe()\")\n# print(df.describe())\n# print(\"train_localizers.CSV Data Info()\")\n# print(df.info())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 0: Install dependencies (safe versions). Kaggle already has many packages;\n# we force-install monai and ensure protobuf and tqdm are present.\n# This is idempotent and quiet to avoid noisy logs.\n\nimport sys\nimport os\n\n# Only run pip installs if packages missing to avoid slow installs on Kaggle\ndef pip_install_if_missing(pkg, import_name=None, version=None):\n    try:\n        __import__(import_name or pkg)\n    except Exception:\n        pkg_str = pkg + ((\"==\" + version) if version else \"\")\n        print(f\"Installing {pkg_str} ...\")\n        os.system(f\"{sys.executable} -m pip install --no-warn-script-location --quiet {pkg_str}\")\n\n# Required libs\npip_install_if_missing(\"tqdm\", \"tqdm\")\npip_install_if_missing(\"monai\", \"monai\")   # adjust version if needed\npip_install_if_missing(\"timm\", \"timm\")\npip_install_if_missing(\"nibabel\", \"nibabel\")\npip_install_if_missing(\"pydicom\", \"pydicom\")\npip_install_if_missing(\"scikit-learn\", \"sklearn\")\npip_install_if_missing(\"opencv-python-headless\", \"cv2\")\npip_install_if_missing(\"protobuf\", \"google.protobuf\")\n\n!pip install -q \"protobuf\"\nprint(\"All required packages should be present.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:46:11.073686Z","iopub.execute_input":"2025-11-29T13:46:11.073842Z","iopub.status.idle":"2025-11-29T13:47:36.165866Z","shell.execute_reply.started":"2025-11-29T13:46:11.073827Z","shell.execute_reply":"2025-11-29T13:47:36.164841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 1 — Imports + AMP + Global Seed\n# ============================================================\n\nimport os, sys, math, time, gc, random, json\nfrom glob import glob\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport nibabel as nib\nimport pydicom\nimport cv2\n\nfrom scipy.ndimage import zoom as ndi_zoom\nfrom scipy.ndimage import label as cc_label\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\n# New recommended AMP API\nfrom torch import amp\n\nimport timm\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.metrics import roc_auc_score\n\n# reproducibility\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.enabled = True\ntorch.backends.cudnn.benchmark = True\n\nprint(\"PyTorch Version:\", torch.__version__)\nprint(\"CUDA Available:\", torch.cuda.is_available())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:51:08.509906Z","iopub.execute_input":"2025-11-29T13:51:08.510461Z","iopub.status.idle":"2025-11-29T13:51:08.932897Z","shell.execute_reply.started":"2025-11-29T13:51:08.510429Z","shell.execute_reply":"2025-11-29T13:51:08.932221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 2 — Global Config (Matches Your Hardware Precisely)\n# ============================================================\n\nclass CFG:\n    # paths\n    DATA_ROOT = \"/kaggle/input/rsna-intracranial-aneurysm-detection\"\n    FILTERED_ROOT = \"/kaggle/input/rsna-filtered-set\"\n    SERIES_DIR = \"series\"\n    SEG_DIR = \"segmentations\"\n\n    OUT_DIR = \"/kaggle/working/RSNA_FULL_PIPELINE\"\n    os.makedirs(OUT_DIR, exist_ok=True)\n\n    # hardware\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    device_type = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    num_workers = 0               # Kaggle safe\n    mixed_precision = True        # use AMP\n    batch_size = 1                # MUST for 3D\n    grad_accum = 8\n    max_epochs_a = 8             # Segmentation stage\n    max_epochs_b = 5             # Classifier stage\n    save_every = 3\n\n    # preprocessing (Stage A)\n    target_spacing = 0.8\n    a_patch_shape = (80, 96, 96)\n    a_clip = (-300, 1000)\n    base_ch = 24\n    lr_a = 2e-4\n\n    # candidate extraction\n    cand_prob_thr = 0.20\n    cand_min_dist_mm = 6.0\n    max_cands_per_series = 150\n\n    # Stage B (2.5D classifier)\n    b_num_slices = 9\n    b_patch_size = 160\n    b_lr = 1e-4\n    b_emb_dim = 512\n\ncfg = CFG()\n\nprint(cfg.__dict__)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:51:58.295093Z","iopub.execute_input":"2025-11-29T13:51:58.295891Z","iopub.status.idle":"2025-11-29T13:51:58.302262Z","shell.execute_reply.started":"2025-11-29T13:51:58.295865Z","shell.execute_reply":"2025-11-29T13:51:58.301594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 3 — Load CSV Metadata\n# ============================================================\n\ntrain_csv = os.path.join(cfg.DATA_ROOT, \"train.csv\")\ntrain_masked_csv = os.path.join(cfg.FILTERED_ROOT, \"train_masked.csv\")\ntrain_localizers_csv = os.path.join(cfg.DATA_ROOT, \"train_localizers.csv\")\n\ntrain_df = pd.read_csv(train_csv)\ntrain_masked_df = pd.read_csv(train_masked_csv)\nlocalizers_df = pd.read_csv(train_localizers_csv)\n\nprint(\"train_df:\", train_df.shape)\nprint(\"train_masked_df:\", train_masked_df.shape)\nprint(\"localizers_df:\", localizers_df.shape)\n\n# Only keep UIDs that have cowseg available\ndef cowseg_exists(uid):\n    p = Path(cfg.DATA_ROOT) / cfg.SEG_DIR / f\"{uid}_cowseg.nii\"\n    p2 = Path(cfg.FILTERED_ROOT) / \"segmentations\" / f\"{uid}_cowseg.nii\"\n    return p.exists() or p2.exists()\n\ntrain_masked_df[\"has_cowseg\"] = train_masked_df[\"SeriesInstanceUID\"].apply(cowseg_exists)\ndf_a = train_masked_df[train_masked_df[\"has_cowseg\"]].reset_index(drop=True)\ndf_a = df_a[0:60] #comment this for full dataset\nprint(\"Stage A usable:\", df_a.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:52:02.075468Z","iopub.execute_input":"2025-11-29T13:52:02.076059Z","iopub.status.idle":"2025-11-29T13:52:02.22252Z","shell.execute_reply.started":"2025-11-29T13:52:02.076014Z","shell.execute_reply":"2025-11-29T13:52:02.221916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 4 — DICOM + NIfTI Loader (No warnings, stable)\n# ============================================================\n\ndef load_dicom_series(series_uid):\n    \"\"\"Load DICOM series → (volume[z,y,x], spacing[z,y,x]).\"\"\"\n    series_path = Path(cfg.DATA_ROOT) / cfg.SERIES_DIR / series_uid\n    files = sorted(list(series_path.glob(\"*.dcm\")))\n    if len(files) == 0:\n        raise FileNotFoundError(f\"No DICOM found for {series_uid}\")\n\n    # sort by InstanceNumber robustly\n    inst_list = []\n    for f in files:\n        try:\n            d = pydicom.dcmread(str(f), stop_before_pixels=True)\n            inst = int(getattr(d, \"InstanceNumber\", 0))\n        except Exception:\n            inst = 0\n        inst_list.append((inst, f))\n\n    inst_list = sorted(inst_list, key=lambda x: x[0])\n    slices = [pydicom.dcmread(str(f)) for _, f in inst_list]\n\n    vol = np.stack([s.pixel_array.astype(np.float32) for s in slices], axis=0)\n\n    intercept = float(getattr(slices[0], \"RescaleIntercept\", 0))\n    slope = float(getattr(slices[0], \"RescaleSlope\", 1))\n    vol = vol * slope + intercept\n\n    px = getattr(slices[0], \"PixelSpacing\", [1.0, 1.0])\n    px = np.array(px, dtype=np.float32)\n\n    # spacing in z\n    try:\n        z0 = slices[0].ImagePositionPatient[2]\n        z1 = slices[1].ImagePositionPatient[2]\n        sz = abs(z1 - z0)\n    except Exception:\n        sz = float(getattr(slices[0], \"SliceThickness\", 1.0))\n\n    spacing = np.array([sz, px[0], px[1]], dtype=np.float32)\n\n    # (z,y,x), spacing[z,y,x]\n    return vol, spacing\n\n\ndef load_nifti(path):\n    \"\"\"Load NIfTI → array(z,y,x), spacing[z,y,x].\"\"\"\n    nii = nib.load(path)\n    arr = nii.get_fdata().astype(np.float32)\n    arr = np.transpose(arr, (2,1,0))  # to (z,y,x)\n    pixdim = nii.header.get_zooms()\n    spacing = np.array([pixdim[2], pixdim[1], pixdim[0]], dtype=np.float32)\n    return arr, spacing\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:52:08.209173Z","iopub.execute_input":"2025-11-29T13:52:08.209682Z","iopub.status.idle":"2025-11-29T13:52:08.217841Z","shell.execute_reply.started":"2025-11-29T13:52:08.209657Z","shell.execute_reply":"2025-11-29T13:52:08.217084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 5 — Preprocessing (Window, Resample, Crop/PAD)\n# ============================================================\n\ndef hu_window(vol, min_hu, max_hu):\n    vol = np.clip(vol, min_hu, max_hu)\n    return ((vol - min_hu) / (max_hu - min_hu)).astype(np.float32)\n\n\ndef resample_to_spacing(vol, spacing, target_spacing, order=1):\n    factors = spacing / target_spacing\n    factors = np.clip(factors, 0.1, 10.0)\n    res = ndi_zoom(vol, zoom=factors, order=order, mode=\"nearest\")\n    return res, np.array([target_spacing]*3, dtype=np.float32)\n\n\ndef center_crop_or_pad(vol, shape):\n    \"\"\"Produce exact shape (Z,Y,X).\"\"\"\n    Z, Y, X = shape\n    out = np.zeros(shape, dtype=vol.dtype)\n\n    z0 = max((Z - vol.shape[0])//2, 0)\n    y0 = max((Y - vol.shape[1])//2, 0)\n    x0 = max((X - vol.shape[2])//2, 0)\n\n    z1 = min(z0 + vol.shape[0], Z)\n    y1 = min(y0 + vol.shape[1], Y)\n    x1 = min(x0 + vol.shape[2], X)\n\n    vz0 = max(0, -(Z - vol.shape[0])//2)\n    vy0 = max(0, -(Y - vol.shape[1])//2)\n    vx0 = max(0, -(X - vol.shape[2])//2)\n\n    out[z0:z1, y0:y1, x0:x1] = vol[\n        vz0:(vz0 + (z1-z0)),\n        vy0:(vy0 + (y1-y0)),\n        vx0:(vx0 + (x1-x0))\n    ]\n\n    return out\n\n\ndef preprocess_volume_for_stage_a(vol, spacing):\n    vol = hu_window(vol, cfg.a_clip[0], cfg.a_clip[1])\n    res, _ = resample_to_spacing(vol, spacing, cfg.target_spacing, order=1)\n    res = center_crop_or_pad(res, cfg.a_patch_shape)\n    return res, cfg.target_spacing\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:52:13.08742Z","iopub.execute_input":"2025-11-29T13:52:13.087954Z","iopub.status.idle":"2025-11-29T13:52:13.095382Z","shell.execute_reply.started":"2025-11-29T13:52:13.087932Z","shell.execute_reply":"2025-11-29T13:52:13.09476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 6 — cowseg Mask Loader + Multi-channel GT creator\n# ============================================================\n\ndef find_cowseg(uid):\n    p1 = Path(cfg.DATA_ROOT) / cfg.SEG_DIR / f\"{uid}_cowseg.nii\"\n    p2 = Path(cfg.FILTERED_ROOT) / \"segmentations\" / f\"{uid}_cowseg.nii\"\n    if p1.exists(): return str(p1)\n    if p2.exists(): return str(p2)\n    return None\n\n\ndef load_clean_mask(uid):\n    \"\"\"Load _cowseg.nii and sanitize to 0..13.\"\"\"\n    path = find_cowseg(uid)\n    if path is None:\n        return None, None\n    arr, sp = load_nifti(path)\n    arr = np.nan_to_num(arr).astype(np.int32)\n    arr[arr < 0] = 0\n    arr[arr > 13] = 0\n    return arr, sp\n\n\ndef build_multichannel_mask(mask, spacing):\n    \"\"\"Convert mask → (14,Z,Y,X) one-hot.\"\"\"\n    if mask is None:\n        return np.zeros((14,) + cfg.a_patch_shape, dtype=np.uint8)\n\n    res, _ = resample_to_spacing(mask.astype(np.float32),\n                                 spacing,\n                                 cfg.target_spacing,\n                                 order=0)\n    res = center_crop_or_pad(res.astype(np.int32), cfg.a_patch_shape)\n\n    out = np.zeros((14,) + cfg.a_patch_shape, dtype=np.uint8)\n    for lbl in range(1,14):\n        out[lbl-1] = (res == lbl).astype(np.uint8)\n    return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:52:17.246267Z","iopub.execute_input":"2025-11-29T13:52:17.246769Z","iopub.status.idle":"2025-11-29T13:52:17.253359Z","shell.execute_reply.started":"2025-11-29T13:52:17.246748Z","shell.execute_reply":"2025-11-29T13:52:17.252577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 7 — Stage A Dataset + Loader\n# ============================================================\n\nclass StageADataset(Dataset):\n    \"\"\"Loads DICOM + cowseg mask → normalized 3D tensor + 14-channel GT.\"\"\"\n    def __init__(self, df):\n        self.df = df.reset_index(drop=True)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        uid = row[\"SeriesInstanceUID\"]\n\n        try:\n            vol, sp = load_dicom_series(uid)\n        except Exception:\n            vol = np.zeros(cfg.a_patch_shape, dtype=np.float32)\n            sp = np.array([cfg.target_spacing]*3)\n\n        vol_prep, _ = preprocess_volume_for_stage_a(vol, sp)\n\n        mask, msp = load_clean_mask(uid)\n        gt = build_multichannel_mask(mask, msp if msp is not None else sp)\n\n        x = torch.from_numpy(vol_prep[None]).float()     # (1,Z,Y,X)\n        y = torch.from_numpy(gt).float()                 # (14,Z,Y,X)\n\n        return x, y, uid\n\n\n# 5-fold split (like earlier)\ngkf = GroupKFold(n_splits=5)\nfold = 0\nindices = list(gkf.split(df_a, groups=df_a[\"SeriesInstanceUID\"]))\ntrain_idx, val_idx = indices[fold]\n\ntrain_df_a = df_a.iloc[train_idx].reset_index(drop=True)\nval_df_a   = df_a.iloc[val_idx].reset_index(drop=True)\n\ntrain_ds_a = StageADataset(train_df_a)\nval_ds_a   = StageADataset(val_df_a)\n\ntrain_loader_a = DataLoader(\n    train_ds_a,\n    batch_size=cfg.batch_size,\n    shuffle=True,\n    num_workers=cfg.num_workers,\n    pin_memory=True,\n)\n\nval_loader_a = DataLoader(\n    val_ds_a,\n    batch_size=1,\n    shuffle=False,\n    num_workers=cfg.num_workers,\n    pin_memory=True,\n)\n\nprint(\"Stage A train:\", len(train_ds_a), \"val:\", len(val_ds_a))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T13:52:22.219138Z","iopub.execute_input":"2025-11-29T13:52:22.219938Z","iopub.status.idle":"2025-11-29T13:52:22.235375Z","shell.execute_reply.started":"2025-11-29T13:52:22.219911Z","shell.execute_reply":"2025-11-29T13:52:22.234678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ============================================================\n# # Cell 8 — LightUNet3D (VRAM-safe, stable for T4/P100)\n# # ============================================================\n\n# class DoubleConv3D(nn.Module):\n#     def __init__(self, in_ch, out_ch):\n#         super().__init__()\n#         self.block = nn.Sequential(\n#             nn.Conv3d(in_ch, out_ch, kernel_size=3, padding=1, bias=False),\n#             nn.InstanceNorm3d(out_ch),\n#             nn.LeakyReLU(0.1, inplace=True),\n\n#             nn.Conv3d(out_ch, out_ch, kernel_size=3, padding=1, bias=False),\n#             nn.InstanceNorm3d(out_ch),\n#             nn.LeakyReLU(0.1, inplace=True),\n#         )\n\n#     def forward(self, x):\n#         return self.block(x)\n\n\n# class LightUNet3D(nn.Module):\n#     \"\"\"Lightweight 3D UNet suitable for 13GB VRAM.\"\"\"\n#     def __init__(self, in_ch=1, out_ch=14, base_ch=16):\n#         super().__init__()\n#         b = base_ch\n\n#         # Encoder\n#         self.enc1 = DoubleConv3D(in_ch, b)\n#         self.pool1 = nn.MaxPool3d(2)\n\n#         self.enc2 = DoubleConv3D(b, b*2)\n#         self.pool2 = nn.MaxPool3d(2)\n\n#         self.enc3 = DoubleConv3D(b*2, b*4)\n#         self.pool3 = nn.MaxPool3d(2)\n\n#         self.enc4 = DoubleConv3D(b*4, b*8)\n#         self.pool4 = nn.MaxPool3d(2)\n\n#         # Bottleneck\n#         self.bottleneck = DoubleConv3D(b*8, b*16)\n\n#         # Decoder\n#         self.up4 = nn.ConvTranspose3d(b*16, b*8, kernel_size=2, stride=2)\n#         self.dec4 = DoubleConv3D(b*16, b*8)\n\n#         self.up3 = nn.ConvTranspose3d(b*8, b*4, kernel_size=2, stride=2)\n#         self.dec3 = DoubleConv3D(b*8, b*4)\n\n#         self.up2 = nn.ConvTranspose3d(b*4, b*2, kernel_size=2, stride=2)\n#         self.dec2 = DoubleConv3D(b*4, b*2)\n\n#         self.up1 = nn.ConvTranspose3d(b*2, b, kernel_size=2, stride=2)\n#         self.dec1 = DoubleConv3D(b*2, b)\n\n#         self.out_conv = nn.Conv3d(b, out_ch, kernel_size=1)\n\n#     def forward(self, x):\n#         c1 = self.enc1(x)\n#         p1 = self.pool1(c1)\n\n#         c2 = self.enc2(p1)\n#         p2 = self.pool2(c2)\n\n#         c3 = self.enc3(p2)\n#         p3 = self.pool3(c3)\n\n#         c4 = self.enc4(p3)\n#         p4 = self.pool4(c4)\n\n#         bn = self.bottleneck(p4)\n\n#         u4 = self.up4(bn)\n#         d4 = self.dec4(torch.cat([u4, c4], dim=1))\n\n#         u3 = self.up3(d4)\n#         d3 = self.dec3(torch.cat([u3, c3], dim=1))\n\n#         u2 = self.up2(d3)\n#         d2 = self.dec2(torch.cat([u2, c2], dim=1))\n\n#         u1 = self.up1(d2)\n#         d1 = self.dec1(torch.cat([u1, c1], dim=1))\n\n#         return self.out_conv(d1)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# ---------------------------\n# DoubleConv3D (keeps same style)\n# ---------------------------\nclass DoubleConv3D(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(out_ch),\n            nn.LeakyReLU(0.1, inplace=True),\n\n            nn.Conv3d(out_ch, out_ch, kernel_size=3, padding=1, bias=False),\n            nn.InstanceNorm3d(out_ch),\n            nn.LeakyReLU(0.1, inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n# ---------------------------\n# Attention Gate (3D)\n# ---------------------------\nclass AttentionGate3D(nn.Module):\n    \"\"\"\n    Attention gate that takes decoder gating signal g and encoder features x,\n    produces attention coefficients to scale x.\n    \"\"\"\n    def __init__(self, F_g, F_l, F_int):\n        super().__init__()\n        # 1x1 conv to reduce channels\n        self.W_g = nn.Sequential(\n            nn.Conv3d(F_g, F_int, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.InstanceNorm3d(F_int)\n        )\n        self.W_x = nn.Sequential(\n            nn.Conv3d(F_l, F_int, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.InstanceNorm3d(F_int)\n        )\n        self.psi = nn.Sequential(\n            nn.Conv3d(F_int, 1, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.InstanceNorm3d(1),\n            nn.Sigmoid()\n        )\n        self.relu = nn.LeakyReLU(0.1, inplace=True)\n\n    def forward(self, g, x):\n        \"\"\"\n        g: gating signal (from decoder), shape [B, F_g, ...]\n        x: encoder feature map to be modulated, shape [B, F_l, ...]\n        returns: x * attention_map (same shape as x)\n        \"\"\"\n        # Reduce and add\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)  # [B,1,D,H,W] with values in (0,1)\n        return x * psi  # broadcast multiply\n\n# ---------------------------\n# LightUNet3D (Attention U-Net implementation, same API)\n# ---------------------------\nclass LightUNet3D(nn.Module):\n    \"\"\"Attention U-Net in the same footprint as previous LightUNet3D.\n       Signature kept identical: LightUNet3D(in_ch=1, out_ch=14, base_ch=16)\n    \"\"\"\n    def __init__(self, in_ch=1, out_ch=14, base_ch=16):\n        super().__init__()\n        b = base_ch\n\n        # Encoder (same layout as your previous LightUNet3D)\n        self.enc1 = DoubleConv3D(in_ch, b)\n        self.pool1 = nn.MaxPool3d(2)\n\n        self.enc2 = DoubleConv3D(b, b*2)\n        self.pool2 = nn.MaxPool3d(2)\n\n        self.enc3 = DoubleConv3D(b*2, b*4)\n        self.pool3 = nn.MaxPool3d(2)\n\n        self.enc4 = DoubleConv3D(b*4, b*8)\n        self.pool4 = nn.MaxPool3d(2)\n\n        # Bottleneck\n        self.bottleneck = DoubleConv3D(b*8, b*16)\n\n        # Attention gates (match channels of encoder + decoder)\n        # F_g = gating channels (decoder), F_l = encoder feature channels\n        self.att4 = AttentionGate3D(F_g=b*8, F_l=b*8, F_int=b*4)\n        self.att3 = AttentionGate3D(F_g=b*4, F_l=b*4, F_int=b*2)\n        self.att2 = AttentionGate3D(F_g=b*2, F_l=b*2, F_int=b)\n        self.att1 = AttentionGate3D(F_g=b,   F_l=b,   F_int=max(b//2, 1))\n\n        # Decoder (ConvTranspose3d + DoubleConv)\n        self.up4 = nn.ConvTranspose3d(b*16, b*8, kernel_size=2, stride=2)\n        self.dec4 = DoubleConv3D(b*16, b*8)\n\n        self.up3 = nn.ConvTranspose3d(b*8, b*4, kernel_size=2, stride=2)\n        self.dec3 = DoubleConv3D(b*8, b*4)\n\n        self.up2 = nn.ConvTranspose3d(b*4, b*2, kernel_size=2, stride=2)\n        self.dec2 = DoubleConv3D(b*4, b*2)\n\n        self.up1 = nn.ConvTranspose3d(b*2, b, kernel_size=2, stride=2)\n        self.dec1 = DoubleConv3D(b*2, b)\n\n        # Output conv (same out_ch as before)\n        self.out_conv = nn.Conv3d(b, out_ch, kernel_size=1)\n\n    def forward(self, x):\n        # Encoder\n        c1 = self.enc1(x)     # [B, b, ...]\n        p1 = self.pool1(c1)\n\n        c2 = self.enc2(p1)    # [B, 2b, ...]\n        p2 = self.pool2(c2)\n\n        c3 = self.enc3(p2)    # [B, 4b, ...]\n        p3 = self.pool3(c3)\n\n        c4 = self.enc4(p3)    # [B, 8b, ...]\n        p4 = self.pool4(c4)\n\n        bn = self.bottleneck(p4)  # [B, 16b, ...]\n\n        # Decoder + attention-modulated skips\n        u4 = self.up4(bn)             # -> [B, 8b, ...]\n        a4 = self.att4(g=u4, x=c4)    # apply attention on encoder features\n        d4 = self.dec4(torch.cat([u4, a4], dim=1))\n\n        u3 = self.up3(d4)\n        a3 = self.att3(g=u3, x=c3)\n        d3 = self.dec3(torch.cat([u3, a3], dim=1))\n\n        u2 = self.up2(d3)\n        a2 = self.att2(g=u2, x=c2)\n        d2 = self.dec2(torch.cat([u2, a2], dim=1))\n\n        u1 = self.up1(d2)\n        a1 = self.att1(g=u1, x=c1)\n        d1 = self.dec1(torch.cat([u1, a1], dim=1))\n\n        return self.out_conv(d1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:03:49.524143Z","iopub.execute_input":"2025-11-29T14:03:49.52474Z","iopub.status.idle":"2025-11-29T14:03:49.541187Z","shell.execute_reply.started":"2025-11-29T14:03:49.524719Z","shell.execute_reply":"2025-11-29T14:03:49.540432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 9 — DiceLoss + FocalBCE + CombinedLoss\n# ============================================================\n\n# class DiceLoss(nn.Module):\n#     def __init__(self, smooth=1e-6):\n#         super().__init__()\n#         self.smooth = smooth\n\n#     def forward(self, logits, targets):\n#         probs = torch.sigmoid(logits)\n#         probs = probs.reshape(probs.shape[0], probs.shape[1], -1)\n#         targets = targets.reshape(targets.shape[0], targets.shape[1], -1)\n\n#         intersection = (probs * targets).sum(-1)\n#         denom = probs.sum(-1) + targets.sum(-1)\n\n#         dice = (2*intersection + self.smooth) / (denom + self.smooth)\n#         return 1 - dice.mean()\n\n\n# class FocalBCE(nn.Module):\n#     \"\"\"Good for heavy label imbalance.\"\"\"\n#     def __init__(self, gamma=2.0, alpha=0.25):\n#         super().__init__()\n#         self.gamma = gamma\n#         self.alpha = alpha\n\n#     def forward(self, logits, targets):\n#         bce = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')\n#         pt = torch.exp(-bce)\n#         focal = self.alpha * ((1-pt)**self.gamma) * bce\n#         return focal.mean()\n\n\n# class HeatmapLoss(nn.Module):\n#     \"\"\"Dice + Focal = stable training for micro-aneurysm-like lesions.\"\"\"\n#     def __init__(self):\n#         super().__init__()\n#         self.dice = DiceLoss()\n#         self.focal = FocalBCE()\n\n#     def forward(self, logits, targets):\n#         return 0.7 * self.dice(logits, targets) + 0.3 * self.focal(logits, targets)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 9 — DiceLoss + FocalBCE + HeatmapLoss (Focal-Tversky + FocalBCE)\n# ============================================================\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass DiceLoss(nn.Module):\n    \"\"\"Classic soft-Dice (kept for compatibility).\"\"\"\n    def __init__(self, smooth=1e-6):\n        super().__init__()\n        self.smooth = smooth\n\n    def forward(self, logits, targets):\n        \"\"\"\n        logits: tensor [B, C, ...] (raw logits)\n        targets: tensor [B, C, ...] (binary or one-hot for each channel)\n        \"\"\"\n        probs = torch.sigmoid(logits)\n        # flatten per sample per channel\n        probs = probs.reshape(probs.shape[0], probs.shape[1], -1)\n        targets = targets.reshape(targets.shape[0], targets.shape[1], -1)\n\n        intersection = (probs * targets).sum(-1)\n        denom = probs.sum(-1) + targets.sum(-1)\n\n        dice = (2 * intersection + self.smooth) / (denom + self.smooth)\n        # dice shape: [B, C] -> mean over channels then over batch\n        return 1.0 - dice.mean()\n\n\nclass FocalBCE(nn.Module):\n    \"\"\"Focal variant of BCE; same API as before.\"\"\"\n    def __init__(self, gamma=2.0, alpha=0.25):\n        super().__init__()\n        self.gamma = gamma\n        self.alpha = alpha\n\n    def forward(self, logits, targets):\n        \"\"\"\n        logits: [B, C, ...]\n        targets: [B, C, ...]\n        \"\"\"\n        bce = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')  # same shape as logits\n        pt = torch.exp(-bce)  # pt = sigmoid(logit) when target=1, (1-sigmoid) when target=0 analog\n        focal = self.alpha * ((1 - pt) ** self.gamma) * bce\n        return focal.mean()\n\n\n# --- Utility: Tversky / Focal-Tversky (multi-channel aware) ---\ndef tversky_index_from_logits(logits, targets, alpha=0.7, beta=0.3, eps=1e-6):\n    \"\"\"\n    Compute Tversky index per-sample per-channel from logits.\n    logits: [B, C, ...]\n    targets: [B, C, ...] (binary)\n    returns: tensor [B, C] of Tversky indices\n    \"\"\"\n    prob = torch.sigmoid(logits)\n    # flatten spatial dims\n    prob = prob.reshape(prob.shape[0], prob.shape[1], -1)\n    targ = targets.reshape(targets.shape[0], targets.shape[1], -1)\n\n    TP = (prob * targ).sum(-1)\n    FP = (prob * (1 - targ)).sum(-1)\n    FN = ((1 - prob) * targ).sum(-1)\n\n    tversky = (TP + eps) / (TP + alpha * FP + beta * FN + eps)\n    return tversky  # [B, C]\n\n\ndef focal_tversky_loss(logits, targets, alpha=0.7, beta=0.3, gamma=0.75):\n    \"\"\"\n    Focal-Tversky loss averaged over channels and batch.\n    \"\"\"\n    tversky = tversky_index_from_logits(logits, targets, alpha=alpha, beta=beta)\n    # tversky in [0,1]; loss = (1 - T)^gamma\n    loss = (1 - tversky) ** gamma\n    return loss.mean()\n\n\nclass HeatmapLoss(nn.Module):\n    \"\"\"\n    Combined loss for heatmap training.\n    Uses Focal-Tversky (for tiny lesion sensitivity) + Focal BCE (for stability).\n    We keep the same outward behavior as previous HeatmapLoss (single forward(logits, targets)).\n    \"\"\"\n    def __init__(self, wt_tversky=0.7, wt_focalbce=0.3, \n                 tversky_alpha=0.7, tversky_beta=0.3, tversky_gamma=0.75,\n                 focal_gamma=2.0, focal_alpha=0.25):\n        super().__init__()\n        self.wt_tversky = wt_tversky\n        self.wt_focalbce = wt_focalbce\n        self.tversky_alpha = tversky_alpha\n        self.tversky_beta = tversky_beta\n        self.tversky_gamma = tversky_gamma\n        self.focal = FocalBCE(gamma=focal_gamma, alpha=focal_alpha)\n\n    def forward(self, logits, targets):\n        \"\"\"\n        logits: [B, C, ...], targets: [B, C, ...] (binary)\n        returns: scalar loss\n        \"\"\"\n        # ensure float tensors\n        logits = logits.float()\n        targets = targets.float()\n\n        ft = focal_tversky_loss(logits, targets, alpha=self.tversky_alpha, beta=self.tversky_beta, gamma=self.tversky_gamma)\n        fb = self.focal(logits, targets)\n        loss = self.wt_tversky * ft + self.wt_focalbce * fb\n        return loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:03:57.712513Z","iopub.execute_input":"2025-11-29T14:03:57.71279Z","iopub.status.idle":"2025-11-29T14:03:57.725142Z","shell.execute_reply.started":"2025-11-29T14:03:57.712769Z","shell.execute_reply":"2025-11-29T14:03:57.724404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 10 — Stage A Training Loop (AMP-safe, tqdm, stable)\n# ============================================================\n\ndef train_stage_a():\n    device = cfg.device\n    device_type = cfg.device_type\n\n    model = LightUNet3D(\n        in_ch=1,\n        out_ch=14,\n        base_ch=cfg.base_ch\n    ).to(device)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr_a)\n    criterion = HeatmapLoss()\n\n    scaler = amp.GradScaler(enabled=cfg.mixed_precision)\n    best_val = 1e9\n\n    history = {\"train\": [], \"val\": []}\n\n    for epoch in range(cfg.max_epochs_a):\n        # ---- Training ----\n        model.train()\n        train_loss = 0.0\n\n        pbar = tqdm(train_loader_a, desc=f\"Epoch {epoch+1}/{cfg.max_epochs_a} [Train]\", leave=False)\n\n        optimizer.zero_grad()\n\n        for step, (x, y, uid) in enumerate(pbar):\n            x = x.to(device)\n            y = y.to(device)\n\n            with amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n                logits = model(x)\n                loss = criterion(logits, y) / cfg.grad_accum\n\n            scaler.scale(loss).backward()\n\n            if (step+1) % cfg.grad_accum == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n            train_loss += loss.item() * cfg.grad_accum\n            pbar.set_postfix(loss=train_loss / (step+1))\n\n        train_loss /= len(train_loader_a)\n        history[\"train\"].append(train_loss)\n\n        # ---- Validation ----\n        model.eval()\n        val_loss = 0.0\n\n        pbar = tqdm(val_loader_a, desc=f\"Epoch {epoch+1}/{cfg.max_epochs_a} [Val]\", leave=False)\n\n        with torch.no_grad():\n            for x, y, uid in pbar:\n                x = x.to(device)\n                y = y.to(device)\n\n                with amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n                    logits = model(x)\n                    loss = criterion(logits, y)\n\n                val_loss += loss.item()\n                pbar.set_postfix(loss=val_loss / (len(val_loader_a)))\n\n        val_loss /= len(val_loader_a)\n        history[\"val\"].append(val_loss)\n\n        print(f\"[Epoch {epoch+1}] Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n\n        # Save best checkpoint\n        best_path = os.path.join(cfg.OUT_DIR, \"stageA_best.pth\")\n        if val_loss < best_val:\n            best_val = val_loss\n            torch.save(model.state_dict(), best_path)\n            print(\"   → Saved BEST model\")\n\n        # Periodic checkpoint\n        if (epoch+1) % cfg.save_every == 0:\n            ckpt_path = os.path.join(cfg.OUT_DIR, f\"stageA_epoch{epoch+1}.pth\")\n            torch.save(model.state_dict(), ckpt_path)\n            print(\"   → Saved checkpoint\", ckpt_path)\n\n    return model, history\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:04.082358Z","iopub.execute_input":"2025-11-29T14:04:04.082632Z","iopub.status.idle":"2025-11-29T14:04:04.092471Z","shell.execute_reply.started":"2025-11-29T14:04:04.082611Z","shell.execute_reply":"2025-11-29T14:04:04.091768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 11 — Stage A Inference + Candidate Extraction\n# ============================================================\n\ndef extract_candidates_from_heatmap(hm, spacing, prob_thr=cfg.cand_prob_thr,\n                                    min_dist_mm=cfg.cand_min_dist_mm, max_cands=cfg.max_cands_per_series):\n    \"\"\"\n    hm: 3D numpy array (Z,Y,X) with values 0..1\n    spacing: scalar or (z,y,x) mm spacing (here we use cfg.target_spacing)\n    \"\"\"\n    mask = hm > prob_thr\n    if mask.sum() == 0:\n        return []\n    labeled, num = cc_label(mask)\n    candidates = []\n    for lab in range(1, num+1):\n        comp = (labeled == lab)\n        if comp.sum() == 0:\n            continue\n        # centroid by max probability\n        idx = np.unravel_index(np.argmax(hm * comp), hm.shape)\n        prob = float(hm[idx])\n        candidates.append((int(idx[0]), int(idx[1]), int(idx[2]), prob))\n    # sort by prob desc\n    candidates = sorted(candidates, key=lambda x: x[3], reverse=True)\n    # NMS by physical distance\n    selected = []\n    for cz, cy, cx, p in candidates:\n        keep = True\n        for s in selected:\n            dz = (cz - s[0]) * cfg.target_spacing\n            dy = (cy - s[1]) * cfg.target_spacing\n            dx = (cx - s[2]) * cfg.target_spacing\n            if math.sqrt(dx*dx + dy*dy + dz*dz) < min_dist_mm:\n                keep = False; break\n        if keep:\n            selected.append((cz, cy, cx, p))\n        if len(selected) >= max_cands:\n            break\n    return selected\n\n\ndef stage_a_inference_and_save(model_path, meta_df=df_a, out_candidates=os.path.join(cfg.OUT_DIR, \"stageA_candidates.csv\")):\n    device = cfg.device\n    device_type = cfg.device_type\n    model = LightUNet3D(in_ch=1, out_ch=14, base_ch=cfg.base_ch).to(device)\n    model.load_state_dict(torch.load(model_path, map_location=device))\n    model.eval()\n\n    cand_rows = []\n    for idx, row in tqdm(meta_df.iterrows(), total=len(meta_df), desc=\"Stage A inference\"):\n        uid = row[\"SeriesInstanceUID\"]\n        try:\n            vol, spacing = load_dicom_series(uid)\n        except Exception:\n            continue\n        vol_prep, _ = preprocess_volume_for_stage_a(vol, spacing)\n        img = torch.from_numpy(vol_prep[None,None].astype(np.float32)).to(device)\n        with torch.no_grad(), amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n            logits = model(img)  # (1,14,Z,Y,X)\n            probs = torch.sigmoid(logits)[0].cpu().numpy()  # (14,Z,Y,X)\n            heatmap = probs.max(0)  # (Z,Y,X)\n            # normalize heatmap for safety\n            if heatmap.max() > 0:\n                heatmap = heatmap / heatmap.max()\n            else:\n                heatmap = heatmap\n            cands = extract_candidates_from_heatmap(heatmap, cfg.target_spacing)\n            for cz, cy, cx, p in cands:\n                cand_rows.append({\"SeriesInstanceUID\": uid, \"z\": int(cz), \"y\": int(cy), \"x\": int(cx), \"prob_a\": float(p)})\n    cand_df = pd.DataFrame(cand_rows)\n    cand_df.to_csv(out_candidates, index=False)\n    print(\"Saved candidates:\", out_candidates, \"count:\", len(cand_df))\n    return cand_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:07.56312Z","iopub.execute_input":"2025-11-29T14:04:07.563683Z","iopub.status.idle":"2025-11-29T14:04:07.574934Z","shell.execute_reply.started":"2025-11-29T14:04:07.563659Z","shell.execute_reply":"2025-11-29T14:04:07.574203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 12 — Build Stage B Labels (GT overlap)\n# ============================================================\n\ndef build_stageb_labels(cand_df, pos_radius_mm=4.0):\n    \"\"\"\n    For each candidate (in Stage A preprocessed coordinates), label positive if within pos_radius_mm\n    of any non-zero voxel in the cowseg mask (after same preprocess/resample/crop).\n    \"\"\"\n    labeled = []\n    for idx, row in tqdm(cand_df.iterrows(), total=len(cand_df), desc=\"Labeling candidates\"):\n        uid = row[\"SeriesInstanceUID\"]\n        cz, cy, cx = int(row[\"z\"]), int(row[\"y\"]), int(row[\"x\"])\n        mask, msp = load_clean_mask(uid)\n        if mask is None:\n            row['label'] = 0\n            labeled.append(row)\n            continue\n        # resample mask and crop to stage A shape\n        mask_res, _ = resample_to_spacing(mask.astype(np.float32), msp, cfg.target_spacing, order=0)\n        mask_res = center_crop_or_pad(mask_res.astype(np.int32), cfg.a_patch_shape)\n        # get GT coords\n        coords = np.argwhere(mask_res > 0)\n        if coords.shape[0] == 0:\n            row['label'] = 0\n            labeled.append(row)\n            continue\n        dz = (coords[:,0] - cz) * cfg.target_spacing\n        dy = (coords[:,1] - cy) * cfg.target_spacing\n        dx = (coords[:,2] - cx) * cfg.target_spacing\n        dist = np.sqrt(dx*dx + dy*dy + dz*dz)\n        row['label'] = int((dist <= pos_radius_mm).any())\n        labeled.append(row)\n    labeled_df = pd.DataFrame(labeled)\n    pos = labeled_df['label'].sum()\n    print(\"Labeled candidates:\", len(labeled_df), \"positives:\", int(pos))\n    return labeled_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:11.143209Z","iopub.execute_input":"2025-11-29T14:04:11.143772Z","iopub.status.idle":"2025-11-29T14:04:11.150789Z","shell.execute_reply.started":"2025-11-29T14:04:11.143747Z","shell.execute_reply":"2025-11-29T14:04:11.149961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 13 — Stage B Dataset (2.5D stacks) + DataLoader\n# ============================================================\n\nclass StageBDataset(Dataset):\n    def __init__(self, cand_df):\n        self.df = cand_df.reset_index(drop=True)\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        uid = row[\"SeriesInstanceUID\"]\n        cz, cy, cx = int(row[\"z\"]), int(row[\"y\"]), int(row[\"x\"])\n        label = float(row.get(\"label\", 0.0))\n\n        # load preprocessed volume (Stage A preprocessing)\n        try:\n            vol, sp = load_dicom_series(uid)\n            vol_prep, _ = preprocess_volume_for_stage_a(vol, sp)\n        except Exception:\n            vol_prep = np.zeros(cfg.a_patch_shape, dtype=np.float32)\n\n        K = cfg.b_num_slices\n        half = K // 2\n        Z,Y,X = vol_prep.shape\n\n        # axial slices\n        z_idxs = np.clip(np.arange(cz-half, cz+half+1), 0, Z-1)\n        axial = vol_prep[z_idxs]  # (K,Y,X)\n\n        # coronal: fix y axis\n        y_idxs = np.clip(np.arange(cy-half, cy+half+1), 0, Y-1)\n        coronal = vol_prep[:, y_idxs, :]  # (Z,K,X) -> transpose to (K,Z,X)\n        coronal = np.transpose(coronal, (1,0,2))\n\n        # sagittal: fix x axis\n        x_idxs = np.clip(np.arange(cx-half, cx+half+1), 0, X-1)\n        sagittal = vol_prep[:, :, x_idxs]  # (Z,Y,K) -> transpose to (K,Z,Y)\n        sagittal = np.transpose(sagittal, (2,0,1))\n\n        # helper crop+resize function\n        def crop_resize(stack):\n            K, H, W = stack.shape\n            min_hw = min(H, W)\n            y0 = (H - min_hw)//2\n            x0 = (W - min_hw)//2\n            crop = stack[:, y0:y0+min_hw, x0:x0+min_hw]\n            out = np.zeros((K, cfg.b_patch_size, cfg.b_patch_size), dtype=np.float32)\n            for i in range(K):\n                out[i] = cv2.resize(crop[i], (cfg.b_patch_size, cfg.b_patch_size), interpolation=cv2.INTER_LINEAR)\n            return out\n\n        axial = crop_resize(axial)\n        coronal = crop_resize(coronal)\n        sagittal = crop_resize(sagittal)\n\n        inp = np.concatenate([axial, coronal, sagittal], axis=0)  # (3K, H, W)\n        inp = torch.from_numpy(inp).float()\n\n        return inp, torch.tensor(label).float(), uid\n\ndef get_stageb_loaders(labeled_df, batch_size=8):\n    # group split by series uid to avoid leakage\n    gkf = GroupKFold(n_splits=5)\n    fold = 0\n    tr_idx, va_idx = list(gkf.split(labeled_df, groups=labeled_df[\"SeriesInstanceUID\"]))[fold]\n    tr_df = labeled_df.iloc[tr_idx].reset_index(drop=True)\n    va_df = labeled_df.iloc[va_idx].reset_index(drop=True)\n\n    tr_ds = StageBDataset(tr_df)\n    va_ds = StageBDataset(va_df)\n\n    tr_loader = DataLoader(tr_ds, batch_size=batch_size, shuffle=True, num_workers=cfg.num_workers, pin_memory=True)\n    va_loader = DataLoader(va_ds, batch_size=batch_size, shuffle=False, num_workers=cfg.num_workers, pin_memory=True)\n    return tr_loader, va_loader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:14.45426Z","iopub.execute_input":"2025-11-29T14:04:14.454543Z","iopub.status.idle":"2025-11-29T14:04:14.466307Z","shell.execute_reply.started":"2025-11-29T14:04:14.454522Z","shell.execute_reply":"2025-11-29T14:04:14.465501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 14 — Stage B Model (EfficientNet-B0) and helper\n# ============================================================\n\nclass StageBModel(nn.Module):\n    def __init__(self, in_ch, backbone_name=\"tf_efficientnet_b0\", emb_dim=cfg.b_emb_dim):\n        super().__init__()\n        self.preconv = nn.Conv2d(in_ch, 3, kernel_size=1)  # map to 3 channels\n        self.backbone = timm.create_model(backbone_name, pretrained=True, features_only=True)\n        feats = self.backbone.feature_info.channels()[-1]\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(feats, emb_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(emb_dim, 1)\n        )\n\n    def forward(self, x):\n        # x: (B, C, H, W)\n        x = self.preconv(x)\n        feats = self.backbone(x)[-1]\n        pooled = self.pool(feats).flatten(1)\n        logit = self.fc(pooled).squeeze(1)\n        return logit\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:19.601266Z","iopub.execute_input":"2025-11-29T14:04:19.601979Z","iopub.status.idle":"2025-11-29T14:04:19.607364Z","shell.execute_reply.started":"2025-11-29T14:04:19.601954Z","shell.execute_reply":"2025-11-29T14:04:19.606518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 15 — Train Stage B (AMP-safe, tqdm, AUC)\n# ============================================================\n\ndef train_stage_b(tr_loader, va_loader, save_path=os.path.join(cfg.OUT_DIR, \"stageB_best.pth\")):\n    device = cfg.device\n    device_type = cfg.device_type\n    in_ch = 3 * cfg.b_num_slices\n\n    model = StageBModel(in_ch=in_ch).to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.b_lr)\n    criterion = nn.BCEWithLogitsLoss()\n    scaler = amp.GradScaler(enabled=cfg.mixed_precision)\n\n    best_auc = 0.0\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_auc\": []}\n\n    for epoch in range(cfg.max_epochs_b):\n        model.train()\n        tr_loss = 0.0\n        pbar = tqdm(tr_loader, desc=f\"Stage B Train Epoch {epoch+1}/{cfg.max_epochs_b}\", leave=False)\n        for imgs, labels, uids in pbar:\n            imgs = imgs.to(device)\n            labels = labels.to(device)\n\n            with amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n                logits = model(imgs)\n                loss = criterion(logits, labels)\n\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n            tr_loss += loss.item() * imgs.size(0)\n            pbar.set_postfix(train_loss=tr_loss / ((pbar.n + 1) * (imgs.size(0))))\n\n        tr_loss /= len(tr_loader.dataset)\n        history[\"train_loss\"].append(tr_loss)\n\n        # Validation\n        model.eval()\n        val_loss = 0.0\n        all_preds, all_gts = [], []\n        with torch.no_grad():\n            pbar = tqdm(va_loader, desc=f\"Stage B Val Epoch {epoch+1}/{cfg.max_epochs_b}\", leave=False)\n            for imgs, labels, uids in pbar:\n                imgs = imgs.to(device)\n                labels = labels.to(device)\n                with amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n                    logits = model(imgs)\n                    loss = criterion(logits, labels)\n                val_loss += loss.item() * imgs.size(0)\n                preds = torch.sigmoid(logits).detach().cpu().numpy()\n                all_preds.append(preds)\n                all_gts.append(labels.detach().cpu().numpy())\n\n        val_loss /= len(va_loader.dataset)\n        all_preds = np.concatenate(all_preds)\n        all_gts = np.concatenate(all_gts)\n        try:\n            auc = roc_auc_score(all_gts, all_preds)\n        except:\n            auc = 0.0\n\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_auc\"].append(auc)\n\n        print(f\"[Stage B] Epoch {epoch+1} | train_loss={tr_loss:.4f} | val_loss={val_loss:.4f} | val_auc={auc:.4f}\")\n\n        if auc > best_auc:\n            best_auc = auc\n            torch.save(model.state_dict(), save_path)\n            print(\"   → Saved best Stage B model:\", save_path)\n\n    return model, history\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:23.601644Z","iopub.execute_input":"2025-11-29T14:04:23.602511Z","iopub.status.idle":"2025-11-29T14:04:23.612698Z","shell.execute_reply.started":"2025-11-29T14:04:23.602476Z","shell.execute_reply":"2025-11-29T14:04:23.611889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 16 — Build Final Series-level CSV (Stage A -> per-artery)\n# ============================================================\n\ndef build_final_csv_from_stage_a(model_a_path, meta_df=df_a, out_csv=os.path.join(cfg.OUT_DIR, \"final_output_stageA.csv\"), artery_thr=0.25):\n    device = cfg.device\n    device_type = cfg.device_type\n    model = LightUNet3D(in_ch=1, out_ch=14, base_ch=cfg.base_ch).to(device)\n    model.load_state_dict(torch.load(model_a_path, map_location=device))\n    model.eval()\n\n    rows = []\n    for idx, row in tqdm(meta_df.iterrows(), total=len(meta_df), desc=\"Stage A series -> final CSV\"):\n        uid = row[\"SeriesInstanceUID\"]\n        try:\n            vol, sp = load_dicom_series(uid)\n        except Exception:\n            continue\n        vol_prep, _ = preprocess_volume_for_stage_a(vol, sp)\n        img = torch.from_numpy(vol_prep[None,None].astype(np.float32)).to(device)\n        with torch.no_grad(), amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n            logits = model(img)\n            probs = torch.sigmoid(logits)[0].cpu().numpy()  # (14,Z,Y,X)\n            # per-artery max in channels 0..12 (13 arteries)\n            per_artery_max = probs[:13].reshape(13,-1).max(axis=1)\n            # flat = probs[:13].reshape(13, -1)\n            # per_artery_max = np.percentile(flat, 99, axis=1) #chat says this two lines start from flat= is for some changes\n            bin_preds = per_artery_max\n            # bin_preds = (per_artery_max > artery_thr).astype(int) #this gives round off value so we comment it\n            aneurysm_present = int(bin_preds.sum() > 0)\n            # map to required column order (from your doc)\n            col_order = [\n                \"Left Infraclinoid ICA\",\"Right Infraclinoid ICA\",\n                \"Left Supraclinoid ICA\",\"Right Supraclinoid ICA\",\n                \"Left MCA\",\"Right MCA\",\"AComm\",\n                \"Left ACA\",\"Right ACA\",\n                \"Left PComm\",\"Right PComm\",\n                \"Basilar Tip\",\"Other Posterior Circulation\"\n            ]\n            # mapping assumption: our channel index 0..12 corresponds to doc order? If not, remap accordingly.\n            # Here we assign channels 0..12 directly to columns in that order.\n            rec = {\"SeriesInstanceUID\": uid}\n            for i, col in enumerate(col_order):\n                rec[col] = float(bin_preds[i])\n            rec[\"Aneurysm Present\"] = int(aneurysm_present)\n            rows.append(rec)\n    out_df = pd.DataFrame(rows)\n    out_df.to_csv(out_csv, index=False)\n    print(\"Saved final CSV:\", out_csv, \"rows:\", len(out_df))\n    return out_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:27.638097Z","iopub.execute_input":"2025-11-29T14:04:27.638452Z","iopub.status.idle":"2025-11-29T14:04:27.649173Z","shell.execute_reply.started":"2025-11-29T14:04:27.638421Z","shell.execute_reply":"2025-11-29T14:04:27.648171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 17 — Plot and Save Training Histories\n# ============================================================\n\ndef plot_history_stagea(history_a, save_path=os.path.join(cfg.OUT_DIR, \"stageA_loss.png\")):\n    plt.figure(figsize=(6,4))\n    plt.plot(history_a[\"train\"], label=\"train_loss\")\n    plt.plot(history_a[\"val\"], label=\"val_loss\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Loss\")\n    plt.title(\"Stage A: Train/Val Loss\")\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=200)\n    print(\"Saved Stage A loss plot:\", save_path)\n    plt.show()\n\ndef plot_history_stageb(history_b, save_path=os.path.join(cfg.OUT_DIR, \"stageB_metrics.png\")):\n    plt.figure(figsize=(6,4))\n    plt.plot(history_b[\"train_loss\"], label=\"train_loss\")\n    plt.plot(history_b[\"val_loss\"], label=\"val_loss\")\n    plt.plot(history_b[\"val_auc\"], label=\"val_auc\")\n    plt.xlabel(\"Epoch\")\n    plt.title(\"Stage B: Loss & AUC\")\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=200)\n    print(\"Saved Stage B metrics plot:\", save_path)\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:31.372063Z","iopub.execute_input":"2025-11-29T14:04:31.372303Z","iopub.status.idle":"2025-11-29T14:04:31.378432Z","shell.execute_reply.started":"2025-11-29T14:04:31.372286Z","shell.execute_reply":"2025-11-29T14:04:31.377686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 18 — Sample segmentation visualizer (saves images)\n# ============================================================\n\ndef save_sample_segmentations(model_path, meta_df=df_a.sample(min(6, len(df_a))), out_dir=os.path.join(cfg.OUT_DIR, \"samples\")):\n    os.makedirs(out_dir, exist_ok=True)\n    device = cfg.device\n    device_type = cfg.device_type\n    model = LightUNet3D(in_ch=1, out_ch=14, base_ch=cfg.base_ch).to(device)\n    model.load_state_dict(torch.load(model_path, map_location=device))\n    model.eval()\n\n    for idx, row in meta_df.iterrows():\n        uid = row[\"SeriesInstanceUID\"]\n        try:\n            vol, sp = load_dicom_series(uid)\n        except Exception:\n            continue\n        vol_prep, _ = preprocess_volume_for_stage_a(vol, sp)\n        img_t = torch.from_numpy(vol_prep[None,None].astype(np.float32)).to(device)\n        with torch.no_grad(), amp.autocast(device_type=device_type, enabled=cfg.mixed_precision):\n            logits = model(img_t)\n            probs = torch.sigmoid(logits)[0].cpu().numpy()  # (14,Z,Y,X)\n            heatmap = probs.max(0)\n        # pick central axial slice\n        zc = heatmap.shape[0] // 2\n        axial_slice = vol_prep[zc]\n        hm_slice = heatmap[zc]\n        # normalize images\n        axial_img = (axial_slice - axial_slice.min()) / (axial_slice.max() - axial_slice.min() + 1e-8)\n        hm_img = (hm_slice - hm_slice.min()) / (hm_slice.max() - hm_slice.min() + 1e-8)\n        fig, ax = plt.subplots(1,2, figsize=(8,4))\n        ax[0].imshow(axial_img, cmap='gray')\n        ax[0].set_title(f\"{uid} axial (z={zc})\")\n        ax[1].imshow(axial_img, cmap='gray')\n        ax[1].imshow(hm_img, cmap='jet', alpha=0.5)\n        ax[1].set_title(\"Heatmap overlay\")\n        plt.tight_layout()\n        save_path = os.path.join(out_dir, f\"{uid}_sample_z{zc}.png\")\n        plt.savefig(save_path, dpi=200)\n        plt.close(fig)\n        print(\"Saved sample:\", save_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:35.002856Z","iopub.execute_input":"2025-11-29T14:04:35.003225Z","iopub.status.idle":"2025-11-29T14:04:35.01329Z","shell.execute_reply.started":"2025-11-29T14:04:35.003202Z","shell.execute_reply":"2025-11-29T14:04:35.012657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DRIVER (run as a single cell after Cells 1..18)\n\n# 1) Train Stage A\nmodel_a, history_a = train_stage_a()\n\n# 2) Save Stage A plots\nplot_history_stagea(history_a)\n\n# 3) Stage A inference -> candidates\ncand_df = stage_a_inference_and_save(os.path.join(cfg.OUT_DIR, \"stageA_best.pth\"))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:04:37.909924Z","iopub.execute_input":"2025-11-29T14:04:37.910695Z","iopub.status.idle":"2025-11-29T14:54:19.357689Z","shell.execute_reply.started":"2025-11-29T14:04:37.910667Z","shell.execute_reply":"2025-11-29T14:54:19.357088Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4) Label candidates (Stage B labels)\nlabeled_df = build_stageb_labels(cand_df)\nlabeled_df.to_csv(os.path.join(cfg.OUT_DIR, \"stageB_labeled_candidates.csv\"), index=False)\n\n# 5) Build Stage B loaders\ntr_loader_b, va_loader_b = get_stageb_loaders(labeled_df, batch_size=8)\n\n# 6) Train Stage B\nmodel_b, history_b = train_stage_b(tr_loader_b, va_loader_b)\nplot_history_stageb(history_b)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T14:57:50.375815Z","iopub.execute_input":"2025-11-29T14:57:50.376239Z","iopub.status.idle":"2025-11-29T15:19:03.562882Z","shell.execute_reply.started":"2025-11-29T14:57:50.376212Z","shell.execute_reply":"2025-11-29T15:19:03.562135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 7) Build final per-series CSV via Stage A\nfinal_df = build_final_csv_from_stage_a(os.path.join(cfg.OUT_DIR, \"stageA_best.pth\"))\nfinal_df.head()\n\n# 8) Save sample segmentations (optional)\nsave_sample_segmentations(os.path.join(cfg.OUT_DIR, \"stageA_best.pth\"))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# final_df = build_final_csv_from_stage_a(os.path.join(cfg.OUT_DIR, \"stageA_best.pth\"))\n# final_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# uid = train_df_a.iloc[0][\"SeriesInstanceUID\"]\n# mask, msp = load_clean_mask(uid)\n# print(\"Original mask unique:\", np.unique(mask)[:20])\n\n# mask_res, _ = resample_to_spacing(mask.astype(np.float32), msp, cfg.target_spacing, order=0)\n# mask_res = center_crop_or_pad(mask_res.astype(np.int32), cfg.a_patch_shape)\n# print(\"Resampled mask unique:\", np.unique(mask_res)[:20])\n\n# print(\"Foreground voxels:\", (mask_res > 0).sum())\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# uid = df_a.iloc[0][\"SeriesInstanceUID\"]\n\n# vol, sp = load_dicom_series(uid)\n# vol_prep, _ = preprocess_volume_for_stage_a(vol, sp)\n# img = torch.from_numpy(vol_prep[None,None].astype(np.float32)).to(cfg.device)\n\n# model = LightUNet3D(in_ch=1, out_ch=14, base_ch=cfg.base_ch).to(cfg.device)\n# model.load_state_dict(torch.load(os.path.join(cfg.OUT_DIR, \"stageA_best.pth\"), map_location=cfg.device))\n# model.eval()\n\n# with torch.no_grad(), amp.autocast(device_type=cfg.device_type, enabled=cfg.mixed_precision):\n#     logits = model(img)\n#     probs = torch.sigmoid(logits)[0].cpu().numpy()\n\n# print(\"Global min prob:\", probs.min())\n# print(\"Global max prob:\", probs.max())\n# print(\"Mean prob:\", probs.mean())\n# print(\"Median prob:\", np.median(probs))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}