{"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":"none","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13441085}],"dockerImageVersionId":31090,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## References\n- https://www.kaggle.com/code/ahsuna123/voxel-by-voxel-3d-cnn-intracranial-aneurysms  \n\n## My Notebooks\n- Train: here\n- Inference: https://www.kaggle.com/code/ichigoe/voxel-by-voxel-3d-cnn-inference","metadata":{}},{"cell_type":"code","source":"# =========================\n# Imports & Global Settings\n# =========================\n\nimport os\nimport gc\nimport shutil\nfrom collections import OrderedDict\nfrom typing import Tuple, List, Dict\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\n\nimport pydicom\nfrom scipy import ndimage\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nfrom torch import amp\n\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import StratifiedShuffleSplit\n\nfrom tqdm.auto import tqdm\n\n# Reproducibility-ish (kept light for speed)\ndef seed_everything(seed: int = 42):\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-30T03:01:57.167407Z","iopub.execute_input":"2025-08-30T03:01:57.167795Z","iopub.status.idle":"2025-08-30T03:02:07.807447Z","shell.execute_reply.started":"2025-08-30T03:01:57.167744Z","shell.execute_reply":"2025-08-30T03:02:07.806494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================\n# Competition Constants\n# =====================\nID_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\n# Paths\nTRAIN_CSV_PATH = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\"\nSERIES_DIR = \"/kaggle/input/rsna-intracranial-aneurysm-detection/series\"\n\n# Processing / Model config\nTARGET_SIZE = (64, 64, 64)      # final (D,H,W)\nTARGET_SPACING_MM = 1.0         # isotropic resample\nCTA_WINDOW = (300.0, 700.0)     # (center, width) for CT (CTA)\nMRI_Z_CLIP = 3.0                # clip z-score to ±3σ\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nUSE_AMP = torch.cuda.is_available()  # enable mixed precision on GPU\n\n# Training knobs\nDO_TRAIN = True                 # This notebook is for training\nEPOCHS = 50\nBATCH_SIZE = 4\nLR = 1e-3\nWEIGHT_DECAY = 1e-5\nANEURYSM_PRESENT_BOOST = 1.0\nPATIENCE = 8  # early stopping\n\n# Runtime practicality\nTRAIN_MAX_SERIES = 512\nVAL_MAX_SERIES = 128\n\n# Avoid CUDA in DataLoader workers (to prevent fork-CUDA issue)\nNUM_WORKERS_TRAIN = 2 if torch.cuda.is_available() else 0\nNUM_WORKERS_VAL = 0\nPIN_MEMORY = False\nPERSISTENT_WORKERS = True if NUM_WORKERS_TRAIN > 0 else False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-30T03:02:07.810218Z","iopub.execute_input":"2025-08-30T03:02:07.810761Z","iopub.status.idle":"2025-08-30T03:02:07.823017Z","shell.execute_reply.started":"2025-08-30T03:02:07.810734Z","shell.execute_reply":"2025-08-30T03:02:07.821776Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================\n# DICOM Processing Utilities\n# ==========================\ndef _safe_zoom(volume: np.ndarray, zoom_factors: Tuple[float, ...], order: int = 1) -> np.ndarray:\n    \"\"\"Robust wrapper around ndimage.zoom to avoid rank mismatch and invalid factors.\"\"\"\n    volume = np.nan_to_num(volume, copy=False)\n    zf = tuple(float(max(1e-6, f)) for f in zoom_factors)  # avoid zeros/negatives\n    if len(zf) != volume.ndim:\n        if len(zf) > volume.ndim:\n            zf = zf[:volume.ndim]\n        else:\n            zf = (1.0,) * (volume.ndim - len(zf)) + zf\n    return ndimage.zoom(volume, zf, order=order)\n\ndef _resize_slice(arr: np.ndarray, out_h: int, out_w: int) -> np.ndarray:\n    \"\"\"Resize a 2D slice to (out_h, out_w) using safe zoom.\"\"\"\n    h, w = arr.shape\n    if h == out_h and w == out_w:\n        return arr.astype(np.float32, copy=False)\n    zy = out_h / max(h, 1)\n    zx = out_w / max(w, 1)\n    return _safe_zoom(arr, (zy, zx), order=1).astype(np.float32, copy=False)\n\n# ==========================\n# DICOM Series Processor\n# ==========================\nclass DICOMProcessor:\n    \"\"\"Process DICOM series into normalized 3D volumes.\"\"\"\n\n    def __init__(\n        self,\n        target_size: Tuple[int, int, int] = TARGET_SIZE,\n        target_spacing_mm: float = TARGET_SPACING_MM,\n        cta_window: Tuple[float, float] = CTA_WINDOW,\n        mri_z_clip: float = MRI_Z_CLIP,\n    ):\n        self.target_size = target_size\n        self.target_spacing_mm = target_spacing_mm\n        self.cta_window = cta_window\n        self.mri_z_clip = mri_z_clip\n        \n        # Adjustment counters\n        self.slope_adjustments = 0\n        self.intercept_adjustments = 0\n        self.adaptive_windowing_count = 0\n\n    def _validate_and_apply_rescale(self, sl: np.ndarray, ds) -> np.ndarray:\n        \"\"\"Validate slope/intercept values and apply rescaling.\"\"\"\n        slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n        intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n        \n        # Validate slope\n        if slope <= 0 or not np.isfinite(slope) or abs(slope) > 1000:\n            slope = 1.0\n            self.slope_adjustments += 1\n        \n        # Validate intercept with range clamping for extreme values\n        if not np.isfinite(intercept):\n            intercept = 0.0\n            self.intercept_adjustments += 1\n        elif abs(intercept) > 10000:\n            intercept = np.clip(intercept, -2000, 0)\n            self.intercept_adjustments += 1\n        \n        # Apply rescaling\n        rescaled = sl * slope + intercept\n        \n        # Validate result\n        if np.any(~np.isfinite(rescaled)):\n            rescaled = np.nan_to_num(rescaled, copy=False)\n        \n        # Post-rescale range check\n        min_val, max_val = rescaled.min(), rescaled.max()\n        if min_val < -5000 or max_val > 10000:\n            rescaled = np.clip(rescaled, -3000, 5000)\n        \n        return rescaled\n\n    def _log_adjustment_summary(self):\n        \"\"\"Log summary of adjustments made during processing.\"\"\"\n        print(f\"Processing adjustments - Slope: {self.slope_adjustments}, Intercept: {self.intercept_adjustments}, Adaptive windowing: {self.adaptive_windowing_count}\")\n\n    def load_dicom_series(self, series_path: str) -> np.ndarray:\n        \"\"\"Return (D,H,W) float32 volume in [0,1].\"\"\"\n        try:\n            # Collect DICOM datasets\n            dicoms = []\n            for root, _, files in os.walk(series_path):\n                for f in files:\n                    if f.endswith(\".dcm\"):\n                        try:\n                            ds = pydicom.dcmread(os.path.join(root, f), force=True)\n                            if hasattr(ds, \"PixelData\"):\n                                dicoms.append(ds)\n                        except Exception:\n                            continue\n            if not dicoms:\n                raise ValueError(f\"No valid DICOM files with pixel data in {series_path}\")\n\n            dicoms = self._sort_slices(dicoms)\n            has_multiframe = any(getattr(ds, \"NumberOfFrames\", 1) > 1 for ds in dicoms)\n            spacing = self._get_spacing(dicoms, has_multiframe=has_multiframe)\n\n            # Choose base HxW\n            base_h, base_w = self._choose_base_shape(dicoms)\n\n            modality_tag = (getattr(dicoms[0], \"Modality\", \"\") or \"\").upper()\n            vol_slices = []\n\n            for ds in dicoms:\n                arr = ds.pixel_array\n                # standardize to (N,H,W) where N=number of frames (1 if 2D)\n                if arr.ndim >= 3:\n                    h, w = arr.shape[-2], arr.shape[-1]\n                    n = int(np.prod(arr.shape[:-2]))\n                    arr = arr.reshape(n, h, w)\n                    frames = arr\n                else:\n                    frames = arr[np.newaxis, ...]  # shape (1,H,W)\n\n                for sl in frames:\n                    sl = sl.astype(np.float32)\n\n                    # Handle MONOCHROME1 inversion\n                    if getattr(ds, \"PhotometricInterpretation\", \"MONOCHROME2\") == \"MONOCHROME1\":\n                        sl = sl.max() - sl\n\n                    # Apply validated rescaling\n                    sl = self._validate_and_apply_rescale(sl, ds)\n\n                    sl = _resize_slice(sl, base_h, base_w)\n                    vol_slices.append(sl)\n\n            if len(vol_slices) == 0:\n                raise ValueError(\"No valid slices extracted.\")\n\n            volume = np.stack(vol_slices, axis=0).astype(np.float32)  # (D,H,W)\n\n            # Normalize by modality -> [0,1]\n            volume = self._normalize_by_modality(volume, modality_tag)\n\n            # Isotropic resample (mm-based)\n            if self.target_spacing_mm is not None:\n                dz, dy, dx = spacing\n                z, y, x = volume.shape\n                newD = max(1, int(round(z * dz / self.target_spacing_mm)))\n                newH = max(1, int(round(y * dy / self.target_spacing_mm)))\n                newW = max(1, int(round(x * dx / self.target_spacing_mm)))\n                volume = _safe_zoom(volume, (newD / z, newH / y, newW / x), order=1)\n\n            # Resize to target grid\n            tz, ty, tx = self.target_size\n            z, y, x = volume.shape\n            volume = _safe_zoom(volume, (tz / z, ty / y, tx / x), order=1).astype(np.float32)\n\n            return volume\n\n        except Exception:\n            return np.zeros(self.target_size, dtype=np.float32)\n\n    def _sort_slices(self, ds_list: List[pydicom.dataset.FileDataset]) -> List[pydicom.dataset.FileDataset]:\n        try:\n            orient = np.array(ds_list[0].ImageOrientationPatient, dtype=np.float32)\n            row = orient[:3]; col = orient[3:]\n            normal = np.cross(row, col)\n            def sort_key(ds):\n                ipp = np.array(getattr(ds, \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\n                return float(np.dot(ipp, normal))\n            return sorted(ds_list, key=sort_key)\n        except Exception:\n            return sorted(ds_list, key=lambda ds: getattr(ds, \"InstanceNumber\", 0))\n\n    def _get_spacing(self, ds_sorted: List[pydicom.dataset.FileDataset], has_multiframe: bool = False) -> Tuple[float, float, float]:\n        try:\n            dy, dx = map(float, ds_sorted[0].PixelSpacing)\n        except Exception:\n            ps = getattr(ds_sorted[0], \"PixelSpacing\", [1.0, 1.0])\n            dy, dx = float(ps[0]), float(ps[1])\n\n        if has_multiframe:\n            dz = float(getattr(ds_sorted[0], \"SpacingBetweenSlices\", getattr(ds_sorted[0], \"SliceThickness\", 1.0)))\n        else:\n            zs = []\n            for i in range(1, len(ds_sorted)):\n                p0 = np.array(getattr(ds_sorted[i-1], \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\n                p1 = np.array(getattr(ds_sorted[i], \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\n                d = np.linalg.norm(p1 - p0)\n                if d > 0:\n                    zs.append(d)\n            if zs:\n                dz = float(np.median(zs))\n            else:\n                dz = float(getattr(ds_sorted[0], \"SliceThickness\", 1.0))\n\n        dz = dz if (dz > 0 and np.isfinite(dz)) else 1.0\n        dy = dy if (dy > 0 and np.isfinite(dy)) else 1.0\n        dx = dx if (dx > 0 and np.isfinite(dx)) else 1.0\n        return (dz, dy, dx)\n\n    def _choose_base_shape(self, ds_list: List[pydicom.dataset.FileDataset]) -> Tuple[int, int]:\n        shapes = []\n        for ds in ds_list:\n            try:\n                h, w = int(ds.Rows), int(ds.Columns)\n            except Exception:\n                arr = ds.pixel_array\n                h, w = arr.shape[-2], arr.shape[-1]\n            shapes.append((h, w))\n        vals, counts = np.unique(shapes, return_counts=True, axis=0)\n        base = tuple(vals[counts.argmax()])\n        return int(base[0]), int(base[1])\n\n    def _normalize_by_modality(self, volume: np.ndarray, modality_tag: str) -> np.ndarray:\n        \"\"\"CT: adaptive windowing for extreme ranges; MR: z-score -> clip -> [0,1].\"\"\"\n        volume = np.nan_to_num(volume, copy=False)\n        \n        if modality_tag == \"CT\":\n            min_val, max_val = volume.min(), volume.max()\n            \n            # Check if values are in normal CT range\n            if min_val >= -2000 and max_val <= 4000:\n                # Normal range: use standard windowing\n                c, w = self.cta_window\n                lo, hi = c - w / 2.0, c + w / 2.0\n            else:\n                # Extreme range: use adaptive windowing\n                self.adaptive_windowing_count += 1\n                \n                # Percentile-based adaptive window\n                p1, p99 = np.percentile(volume, [1, 99])\n                margin = (p99 - p1) * 0.1\n                lo = p1 - margin\n                hi = p99 + margin\n                \n                # Ensure minimum window width\n                if hi - lo < 100:\n                    center = (hi + lo) / 2\n                    lo = center - 50\n                    hi = center + 50\n            \n            v = np.clip(volume, lo, hi)\n            v = (v - lo) / (hi - lo + 1e-6)\n            return v.astype(np.float32, copy=False)\n        else:\n            # MRI processing\n            mean = float(volume.mean())\n            std = float(volume.std() + 1e-6)\n            \n            # Validate statistics\n            if std < 1e-6 or not np.isfinite(mean) or not np.isfinite(std):\n                return np.full_like(volume, 0.5, dtype=np.float32)\n            \n            # Check dynamic range\n            min_val, max_val = volume.min(), volume.max()\n            if max_val - min_val < 1e-6:\n                return np.full_like(volume, 0.5, dtype=np.float32)\n            \n            v = (volume - mean) / std\n            zc = float(self.mri_z_clip)\n            v = np.clip(v, -zc, zc)\n            v = (v + zc) / (2.0 * zc)\n            return v.astype(np.float32, copy=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-30T03:02:07.82426Z","iopub.execute_input":"2025-08-30T03:02:07.824573Z","iopub.status.idle":"2025-08-30T03:02:07.867847Z","shell.execute_reply.started":"2025-08-30T03:02:07.824544Z","shell.execute_reply":"2025-08-30T03:02:07.866815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================\n# Preprocessed Dataset Container\n# ==============================\nclass PreprocessedDataset(Dataset):\n    \"\"\"Dataset for preprocessed volumes stored in memory.\"\"\"\n\n    def __init__(self, volumes: Dict[str, np.ndarray], data_df: pd.DataFrame):\n        self.volumes = volumes\n        self.data_df = data_df.reset_index(drop=True)\n\n    def __len__(self):\n        return len(self.data_df)\n\n    def __getitem__(self, idx):\n        row = self.data_df.iloc[idx]\n        series_id = row[ID_COL]\n        \n        volume = self.volumes[series_id]\n        if not volume.flags.writeable:\n            volume = volume.copy()\n\n        labels = row[LABEL_COLS].values.astype(np.float32)\n        volume_tensor = torch.from_numpy(volume).unsqueeze(0)  # (1,D,H,W)\n        labels_tensor = torch.from_numpy(labels)\n        return volume_tensor, labels_tensor\n\n# =======================\n# Model Class\n# =======================\nclass Simple3DCNN(nn.Module):\n    \"\"\"Lightweight 3D CNN for multi-label classification (returns logits).\"\"\"\n\n    def __init__(self, num_classes: int = len(LABEL_COLS)):\n        super(Simple3DCNN, self).__init__()\n        self.conv1 = nn.Conv3d(1, 16, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm3d(16)\n        self.pool1 = nn.MaxPool3d(2)\n\n        self.conv2 = nn.Conv3d(16, 32, kernel_size=3, padding=1)\n        self.bn2 = nn.BatchNorm3d(32)\n        self.pool2 = nn.MaxPool3d(2)\n\n        self.conv3 = nn.Conv3d(32, 64, kernel_size=3, padding=1)\n        self.bn3 = nn.BatchNorm3d(64)\n        self.pool3 = nn.MaxPool3d(2)\n\n        self.conv4 = nn.Conv3d(64, 128, kernel_size=3, padding=1)\n        self.bn4 = nn.BatchNorm3d(128)\n        self.pool4 = nn.MaxPool3d(2)\n\n        self.adaptive_pool = nn.AdaptiveAvgPool3d((2, 2, 2))\n        self.fc1 = nn.Linear(128 * 2 * 2 * 2, 256)\n        self.dropout1 = nn.Dropout(0.5)\n        self.fc2 = nn.Linear(256, 128)\n        self.dropout2 = nn.Dropout(0.3)\n        self.fc3 = nn.Linear(128, num_classes)\n\n    def forward(self, x):\n        # x: (B,1,D,H,W)\n        x = self.pool1(F.relu(self.bn1(self.conv1(x))))\n        x = self.pool2(F.relu(self.bn2(self.conv2(x))))\n        x = self.pool3(F.relu(self.bn3(self.conv3(x))))\n        x = self.pool4(F.relu(self.bn4(self.conv4(x))))\n        x = self.adaptive_pool(x)\n        x = x.view(x.size(0), -1)\n        x = F.relu(self.fc1(x)); x = self.dropout1(x)\n        x = F.relu(self.fc2(x)); x = self.dropout2(x)\n        x = self.fc3(x)  # logits\n        return x\n\n# ============================\n# Metric: Mean Weighted ColAUC\n# ============================\nAP_COL = 'Aneurysm Present'\nLOC_COLS = LABEL_COLS[:-1]\n\ndef mean_weighted_colwise_auc(y_true_df: pd.DataFrame, y_pred_df: pd.DataFrame):\n    \"\"\"Implements the competition metric.\"\"\"\n    aucs = {}\n    for c in LABEL_COLS:\n        y_t = y_true_df[c].values\n        y_p = y_pred_df[c].values\n        if len(np.unique(y_t)) < 2:\n            auc = 0.5\n        else:\n            auc = roc_auc_score(y_t, y_p)\n        aucs[c] = auc\n    ap = aucs[AP_COL]\n    others = float(np.mean([aucs[c] for c in LOC_COLS]))\n    final = 0.5 * (ap + others)\n    return final, aucs\n\n# =====================\n# Training & Evaluation\n# =====================\ndef compute_pos_weight(train_df: pd.DataFrame, label_cols: list, eps: float = 1.0) -> torch.Tensor:\n    \"\"\"pos_weight = (neg + eps) / (pos + eps) per column to counter class imbalance.\"\"\"\n    total = float(len(train_df))\n    pos = train_df[label_cols].sum(axis=0).astype(float)\n    neg = total - pos\n    w = (neg + eps) / (pos + eps)\n    return torch.tensor(w.values, dtype=torch.float32, device=DEVICE)\n\n@torch.no_grad()\ndef evaluate_model(model: nn.Module, val_dataset: PreprocessedDataset, batch_size: int = 1):\n    \"\"\"Validation using preprocessed data.\"\"\"\n    model.eval()\n    dl = DataLoader(\n        val_dataset, batch_size=batch_size, shuffle=False, num_workers=0,\n        pin_memory=False, persistent_workers=False\n    )\n\n    preds, trues = [], []\n    val_loss = 0.0\n    # Use pos_weight for consistency with training\n    pw = compute_pos_weight(val_dataset.data_df, LABEL_COLS, eps=1.0).clone()\n    criterion = nn.BCEWithLogitsLoss(pos_weight=pw)\n\n    for vols, labels in tqdm(dl, desc=\"Val\", leave=False):\n        vols = vols.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n        with amp.autocast(device_type='cuda', enabled=USE_AMP):\n            logits = model(vols)\n            loss = criterion(logits, labels)\n            probs = torch.sigmoid(logits)\n        val_loss += float(loss.item()) * vols.size(0)\n        preds.append(probs.cpu().numpy())\n        trues.append(labels.cpu().numpy())\n\n    y_pred = np.vstack(preds)\n    y_true = np.vstack(trues)\n    y_pred_df = pd.DataFrame(y_pred, columns=LABEL_COLS)\n    y_true_df = pd.DataFrame(y_true, columns=LABEL_COLS)\n\n    final, aucs = mean_weighted_colwise_auc(y_true_df, y_pred_df)\n    val_loss = val_loss / max(len(val_dataset), 1)\n    return final, aucs, val_loss\n\ndef preprocess_all_data(data_df: pd.DataFrame, series_dir: str, processor: DICOMProcessor) -> Dict[str, np.ndarray]:\n    \"\"\"Preprocess all DICOM series once at the beginning.\"\"\"\n    print(\"Preprocessing all DICOM series...\")\n    volumes = {}\n    \n    for idx, row in tqdm(data_df.iterrows(), total=len(data_df), desc=\"Processing\"):\n        series_id = row[ID_COL]\n        series_path = os.path.join(series_dir, series_id)\n        volume = processor.load_dicom_series(series_path)\n        volumes[series_id] = volume\n    \n    processor._log_adjustment_summary()\n    return volumes\n\ndef train_model(\n    train_df: pd.DataFrame,\n    series_dir: str,\n    processor: DICOMProcessor,\n    epochs: int = EPOCHS,\n    batch_size: int = BATCH_SIZE,\n    lr: float = LR,\n    weight_decay: float = WEIGHT_DECAY,\n    aneurysm_present_boost: float = ANEURYSM_PRESENT_BOOST,\n    patience: int = PATIENCE,\n    save_path: str = \"/kaggle/working/model_weights.pth\",\n    warm_start_path: str = None,\n    monitor: str = \"auc\",  # \"auc\" (maximize MW-ColAUC) or \"loss\" (minimize val_loss)\n) -> nn.Module:\n    \"\"\"Train 3D CNN with preprocessed data.\"\"\"\n    assert monitor in {\"auc\", \"loss\"}, \"monitor must be 'auc' or 'loss'\"\n\n    # Stratified split by AP (simple hold-out)\n    sss = StratifiedShuffleSplit(n_splits=1, test_size=0.2, random_state=42)\n    ap = train_df[AP_COL].values\n    train_idx, val_idx = next(sss.split(train_df, ap))\n\n    tr_df = train_df.iloc[train_idx].copy()\n    va_df = train_df.iloc[val_idx].copy()\n\n    # Optional subsampling for runtime practicality\n    if TRAIN_MAX_SERIES is not None and len(tr_df) > TRAIN_MAX_SERIES:\n        tr_df = tr_df.sample(TRAIN_MAX_SERIES, random_state=42)\n    if VAL_MAX_SERIES is not None and len(va_df) > VAL_MAX_SERIES:\n        va_df = va_df.sample(VAL_MAX_SERIES, random_state=42)\n\n    # Combine for preprocessing\n    combined_df = pd.concat([tr_df, va_df], ignore_index=True)\n    \n    # Preprocess all data once\n    all_volumes = preprocess_all_data(combined_df, series_dir, processor)\n    \n    # Create datasets\n    train_volumes = {sid: all_volumes[sid] for sid in tr_df[ID_COL]}\n    val_volumes = {sid: all_volumes[sid] for sid in va_df[ID_COL]}\n    \n    ds_tr = PreprocessedDataset(train_volumes, tr_df)\n    ds_va = PreprocessedDataset(val_volumes, va_df)\n    \n    dl_tr = DataLoader(\n        ds_tr, batch_size=batch_size, shuffle=True,\n        num_workers=NUM_WORKERS_TRAIN, pin_memory=PIN_MEMORY,\n        persistent_workers=PERSISTENT_WORKERS,\n        prefetch_factor=2 if NUM_WORKERS_TRAIN > 0 else None\n    )\n\n    # Create model\n    model = Simple3DCNN(num_classes=len(LABEL_COLS)).to(DEVICE)\n\n    # Warm start if provided\n    if warm_start_path and os.path.exists(warm_start_path):\n        try:\n            state = torch.load(warm_start_path, map_location='cpu')\n            model.load_state_dict(state, strict=False)\n            print(f\"Warm-started from {warm_start_path}\")\n        except Exception as e:\n            print(f\"[Warm start warn] {e}\")\n\n    # Loss with pos_weight\n    pos_weight = compute_pos_weight(tr_df, LABEL_COLS, eps=1.0).clone()\n    if aneurysm_present_boost != 1.0:\n        pos_weight[-1] = pos_weight[-1] * float(aneurysm_present_boost)\n    criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n\n    optimizer = optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        mode=(\"max\" if monitor == \"auc\" else \"min\"),\n        patience=1\n    )\n    scaler = amp.GradScaler(enabled=USE_AMP)\n\n    # Initialize best score\n    best_score = -float('inf') if monitor == \"auc\" else float('inf')\n    best_state = None\n    no_improve = 0\n\n    for epoch in range(1, epochs + 1):\n        model.train()\n        running = 0.0\n        n_samples = 0\n\n        pbar = tqdm(dl_tr, desc=f\"Train {epoch}/{epochs}\", leave=False)\n        for vols, labels in pbar:\n            vols = vols.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True)\n\n            optimizer.zero_grad(set_to_none=True)\n            with amp.autocast(device_type='cuda', enabled=USE_AMP):\n                logits = model(vols)\n                loss = criterion(logits, labels)\n\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            running += float(loss.item()) * vols.size(0)\n            n_samples += vols.size(0)\n            pbar.set_postfix(loss=f\"{running/max(n_samples,1):.4f}\")\n\n        train_loss = running / max(n_samples, 1)\n\n        # Validation\n        try:\n            final_auc, per_col, val_loss = evaluate_model(model, ds_va, batch_size=1)\n            print(f\"[Epoch {epoch}/{epochs}] train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | MW-ColAUC={final_auc:.4f}\")\n        except Exception as e:\n            print(f\"[Eval warning] {e}\")\n            final_auc, val_loss = -1.0, train_loss + 1.0\n\n        # Choose monitored score\n        score = final_auc if monitor == \"auc\" else val_loss\n\n        # Step scheduler with monitored score\n        scheduler.step(score)\n\n        # Early stopping with monitored score\n        is_better = (score > best_score) if monitor == \"auc\" else (score < best_score)\n        if is_better:\n            best_score = score\n            best_state = model.state_dict()\n            no_improve = 0\n        else:\n            no_improve += 1\n            if no_improve >= patience:\n                print(f\"Early stopping at epoch {epoch} (no improvement for {patience} epochs).\")\n                break\n\n    # Load best and save\n    if best_state is not None:\n        model.load_state_dict(best_state)\n    torch.save(model.state_dict(), save_path)\n    print(f\"Model saved to {save_path} (selected by monitor='{monitor}', best_score={best_score:.6f})\")\n\n    # Final validation summary\n    try:\n        final_auc, per_col, val_loss = evaluate_model(model, ds_va, batch_size=1)\n        print(f\"[Final Val] val_loss={val_loss:.4f} | MW-ColAUC={final_auc:.4f}\")\n    except Exception as e:\n        print(f\"[Eval warning] {e}\")\n\n    model.eval()\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-30T03:02:07.869055Z","iopub.execute_input":"2025-08-30T03:02:07.86935Z","iopub.status.idle":"2025-08-30T03:02:07.912047Z","shell.execute_reply.started":"2025-08-30T03:02:07.869316Z","shell.execute_reply":"2025-08-30T03:02:07.910747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =======================================\n# Global Initialization, Training & Save\n# =======================================\nprint(\"Initializing processor...\")\nprocessor = DICOMProcessor(\n    target_size=TARGET_SIZE,\n    target_spacing_mm=TARGET_SPACING_MM,\n    cta_window=CTA_WINDOW,\n    mri_z_clip=MRI_Z_CLIP,\n)\n\nmodel = None\n\nif DO_TRAIN:\n    try:\n        full_df = pd.read_csv(TRAIN_CSV_PATH)\n        print(f\"Training on up to {TRAIN_MAX_SERIES} train series, batch={BATCH_SIZE}, epochs={EPOCHS}, patience={PATIENCE} ...\")\n        model = train_model(\n            train_df=full_df,\n            series_dir=SERIES_DIR,\n            processor=processor,\n            epochs=EPOCHS,\n            batch_size=BATCH_SIZE,\n            lr=LR,\n            weight_decay=WEIGHT_DECAY,\n            aneurysm_present_boost=ANEURYSM_PRESENT_BOOST,\n            patience=PATIENCE,\n            save_path=\"/kaggle/working/model_weights.pth\",\n            warm_start_path=\"/kaggle/input/model_weights.pth\" if os.path.exists(\"/kaggle/input/model_weights.pth\") else None,\n            monitor=\"auc\",\n        )\n        \n    except Exception as e:\n        print(f\"[Train warning] {e}\")\n        model = Simple3DCNN(num_classes=len(LABEL_COLS)).to(DEVICE)\n        torch.save(model.state_dict(), \"/kaggle/working/model_weights.pth\")\n        print(\"Saved a randomly initialized model to /kaggle/working/model_weights.pth\")\nelse:\n    model = Simple3DCNN(num_classes=len(LABEL_COLS)).to(DEVICE)\n    torch.save(model.state_dict(), \"/kaggle/working/model_weights.pth\")\n    print(\"Training disabled. Saved untrained model to /kaggle/working/model_weights.pth\")\n\nprint(\"Training notebook completed. Best epoch weights are saved at /kaggle/working/model_weights.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-30T03:02:07.913257Z","iopub.execute_input":"2025-08-30T03:02:07.913631Z","execution_failed":"2025-08-30T03:02:26.91Z"}},"outputs":[],"execution_count":null}]}