{"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":283685877,"sourceType":"kernelVersion"},{"sourceId":283759136,"sourceType":"kernelVersion"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom pathlib import Path\nfrom collections import defaultdict\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.utils.data import Dataset, DataLoader, RandomSampler, SequentialSampler\nfrom torch.serialization import add_safe_globals\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import (\n    roc_auc_score,\n    confusion_matrix,\n    ConfusionMatrixDisplay,\n    classification_report,\n    roc_curve,\n    precision_recall_curve,\n    average_precision_score,\n)\n\nimport torchvision.transforms as VT\nimport torchvision.transforms.functional as VF","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:35:04.708742Z","iopub.execute_input":"2025-12-07T23:35:04.709021Z","iopub.status.idle":"2025-12-07T23:35:10.488246Z","shell.execute_reply.started":"2025-12-07T23:35:04.709003Z","shell.execute_reply":"2025-12-07T23:35:10.487475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CFG = {\n    # Preprocessed volumes\n    \"PREPROC_INDEX\": \"/kaggle/input/preprocessingv5/preproc_teacher/preproc_teacher_index.csv\",\n    \"NPZ_DIR\": \"/kaggle/input/preprocessingv5/preproc_teacher\",\n\n    # RSNA labels & localizers\n    \"TRAIN_LABELS\": \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\",\n    \"LOCALIZERS_CSV\": \"/kaggle/input/rsna-intracranial-aneurysm-detection/train_localizers.csv\",\n\n    # Training hyperparameters\n    \"epochs\": 12,\n    \"freeze_encoder_epochs\": 3,\n    \"batch_series\": 2,\n    \"num_workers\": 2,\n    \"lr\": 1e-4,\n    \"weight_decay\": 1e-4,\n    \"grad_clip\": 5.0,\n    \"aneurysm_weight\": 13.0,\n    \"label_smoothing\": 0.02,\n    \"seed\": 18,\n\n    # MIL / memory safety\n    \"max_s_train\": 96,\n    \"max_s_val\": 128,\n    \"slice_sample\": \"uniform\",   # \"uniform\" | \"center\"\n    \"encoder_chunk\": 32,\n    \"channels_last\": True,\n    \"cuda_alloc_conf\": \"expandable_segments:True,max_split_size_mb:128\",\n\n    # AMP & TTA\n    \"use_amp\": True,\n    \"tta_hflip\": True,\n    \"lambda_attn\": 0.08,\n\n    # CV / splits\n    \"folds\": 3,\n\n    # Device / fallback\n    \"fallback_to_cpu_if_cuda_unhealthy\": True,\n\n    # Checkpointing\n    \"ckpt_dir\": \"/kaggle/working\",\n    \"best_ckpt_name\": \"model_teacher_effs_mil_best.pt\",\n}\n\nPath(CFG[\"ckpt_dir\"]).mkdir(parents=True, exist_ok=True)\n\n# ============================================================\n# 1. DEVICE, SEEDING, HELPERS\n# ============================================================\ndef is_cuda_healthy() -> bool:\n    if not torch.cuda.is_available():\n        return False\n    try:\n        torch.cuda.synchronize()\n        _ = torch.empty(1, device=\"cuda\")\n        torch.cuda.synchronize()\n        return True\n    except Exception as e:\n        print(f\"[WARN] CUDA health check failed: {type(e).__name__}: {e}\")\n        return False\n\n\n# Env for CUDA allocation\nif \"PYTORCH_CUDA_ALLOC_CONF\" not in os.environ and CFG[\"cuda_alloc_conf\"]:\n    os.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = CFG[\"cuda_alloc_conf\"]\n\n_cuda_ok = is_cuda_healthy()\nif not _cuda_ok and CFG[\"fallback_to_cpu_if_cuda_unhealthy\"]:\n    print(\"[INFO] Falling back to CPU due to CUDA health check failure.\")\n    device = \"cpu\"\nelse:\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\ndef set_seed(seed: int = 42, allow_cuda_seed: bool = True):\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.random.default_generator.manual_seed(seed)\n    if allow_cuda_seed and torch.cuda.is_available():\n        try:\n            torch.cuda.manual_seed_all(seed)\n        except RuntimeError as e:\n            print(f\"[WARN] Skipping CUDA seeding ({type(e).__name__}: {e})\")\n\nset_seed(CFG[\"seed\"], allow_cuda_seed=(device == \"cuda\"))\n\nif CFG[\"channels_last\"] and device == \"cuda\":\n    try:\n        torch.set_float32_matmul_precision(\"medium\")\n    except Exception:\n        pass\n\ndef get_base_model(m):\n    return m.module if hasattr(m, \"module\") else m\n\ndef unwrap_state_dict(m):\n    return get_base_model(m).state_dict()\n\n# Safe torch.load for numpy scalars (PyTorch >= 2.6)\nadd_safe_globals([np.core.multiarray.scalar, np.dtype])\n\n# ============================================================\n# 2. LABELS & BASIC HELPERS\n# ============================================================\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]\nNUM_LABELS = len(LABEL_COLS)\nMODALITY_TO_ID = {\"CTA\": 0, \"MRA\": 1}\n\ndef load_labels_frame(labels_csv: str, uid_col: str = \"SeriesInstanceUID\") -> pd.DataFrame:\n    df = pd.read_csv(labels_csv)\n    needed = [uid_col] + LABEL_COLS\n    missing = [c for c in needed if c not in df.columns]\n    if missing:\n        raise ValueError(f\"Label CSV missing columns: {missing}\")\n    return df[needed].copy()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:35:12.174378Z","iopub.execute_input":"2025-12-07T23:35:12.175211Z","iopub.status.idle":"2025-12-07T23:35:12.437621Z","shell.execute_reply.started":"2025-12-07T23:35:12.175184Z","shell.execute_reply":"2025-12-07T23:35:12.436814Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 3. AUGMENTATIONS, LOCALIZER MAP, DATASET & COLLATE\n# ============================================================\nclass SliceAug:\n    \"\"\"Mild per-slice augmentations (3xHxW) on CPU.\"\"\"\n    def __init__(self, p=0.75):\n        self.p = p\n        self.jitter = VT.ColorJitter(brightness=0.15, contrast=0.15)\n    def __call__(self, x_np):  # (3,H,W) float32 in [0,1]\n        if np.random.rand() > self.p:\n            return x_np\n        x = torch.from_numpy(x_np)\n        if np.random.rand() < 0.4:\n            angle = float(np.random.uniform(-7.0, 7.0))\n            x = VF.rotate(x, angle, interpolation=VT.InterpolationMode.BILINEAR)\n        if np.random.rand() < 0.5:\n            x = self.jitter(x)\n        if np.random.rand() < 0.25:\n            x = VT.RandomErasing(\n                p=0.35, scale=(0.01, 0.05), ratio=(0.3, 3.0),\n                value=0.0, inplace=False\n            )(x)\n        return x.clamp_(0, 1).numpy()\n\nSLICE_AUG = SliceAug(p=0.75)\n\ndef build_series_to_roi_sops(localizers_csv: str) -> dict:\n    loc = pd.read_csv(localizers_csv)\n    if \"SOPInstanceUID\" not in loc.columns:\n        raise ValueError(\"train_localizers.csv missing SOPInstanceUID column\")\n    if \"any\" in loc.columns:\n        roi_loc = loc[loc[\"any\"] == 1]\n    else:\n        roi_loc = loc\n    series_to_roi = (\n        roi_loc.groupby(\"SeriesInstanceUID\")[\"SOPInstanceUID\"]\n        .apply(lambda x: set(map(str, x.values)))\n        .to_dict()\n    )\n    return series_to_roi\n\nclass RSNATeacherSeriesNPZ(Dataset):\n    \"\"\"\n    NPZ layout (from preproc_teacher):\n      - volume: [Z,3,H,W] uint8\n      - sops:   [Z] SOPInstanceUID strings\n      - spacing, modality, series_uid, bbox, z_bounds (some ignored here)\n    index CSV needs: SeriesInstanceUID, Modality, npz_path\n    \"\"\"\n    def __init__(\n        self,\n        index_csv: str,\n        labels_df: pd.DataFrame,\n        series_to_roi_sops: dict,\n        uid_col: str = \"SeriesInstanceUID\",\n        npz_path_col: str = \"npz_path\",\n        modality_col: str = \"Modality\",\n        load_into_mem: bool = False,\n        dtype: torch.dtype = torch.float32,\n        do_aug: bool = False,\n    ):\n        self.meta = pd.read_csv(index_csv)\n        need_cols = {\"SeriesInstanceUID\", npz_path_col, modality_col}\n        if not need_cols.issubset(self.meta.columns):\n            raise ValueError(f\"index CSV must include {need_cols}\")\n        self.uid_col = uid_col\n        self.npz_path_col = npz_path_col\n        self.modality_col = modality_col\n        self.labels_df = labels_df.set_index(uid_col)\n        self.series_to_roi_sops = series_to_roi_sops or {}\n        self.load_into_mem = load_into_mem\n        self.dtype = dtype\n        self.do_aug = do_aug\n\n        self.cache = {}\n        if load_into_mem:\n            for _, row in self.meta.iterrows():\n                uid = str(row[self.uid_col])\n                p = row[self.npz_path_col]\n                data = np.load(p, allow_pickle=True)\n                self.cache[uid] = {k: data[k] for k in data.files}\n\n    def __len__(self):\n        return len(self.meta)\n\n    def __getitem__(self, idx):\n        row = self.meta.iloc[idx]\n        uid = str(row[self.uid_col])\n        p = row[self.npz_path_col]\n\n        if self.load_into_mem and uid in self.cache:\n            data = self.cache[uid]\n        else:\n            data = np.load(p, allow_pickle=True)\n\n        if \"volume\" not in data.files:\n            raise ValueError(f\"{p} missing 'volume' array\")\n\n        vol = data[\"volume\"]  # [Z,3,H,W] uint8\n        if vol.ndim != 4:\n            raise ValueError(f\"{uid}: expected 4D volume, got {vol.shape}\")\n\n        slices = vol.astype(np.float32) / 255.0  # [Z,3,H,W]\n\n        sops = None\n        if \"sops\" in data.files:\n            sops_arr = data[\"sops\"]\n            sops = [str(s) for s in sops_arr.tolist()]\n            if len(sops) != slices.shape[0]:\n                if len(sops) == 1:\n                    sops = [sops[0] for _ in range(slices.shape[0])]\n                else:\n                    sops = sops[:slices.shape[0]]\n                    while len(sops) < slices.shape[0]:\n                        sops.append(sops[-1])\n\n        if self.do_aug:\n            for i in range(slices.shape[0]):\n                slices[i] = SLICE_AUG(slices[i])\n\n        modality_str = str(row[self.modality_col]).upper()\n        modality_id = MODALITY_TO_ID.get(modality_str, 0)\n\n        # labels\n        y = self.labels_df.loc[uid, LABEL_COLS].astype(float).values\n        has_aneurysm = (y[LABEL_COLS.index(\"Aneurysm Present\")] == 1.0)\n\n        S = slices.shape[0]\n        attn_target = np.zeros((S,), dtype=np.float32)\n        if has_aneurysm and sops is not None and uid in self.series_to_roi_sops:\n            roi_sops = self.series_to_roi_sops[uid]\n            for i, sop in enumerate(sops):\n                if sop in roi_sops:\n                    attn_target[i] = 1.0\n            # neighbor dilation\n            roi_idx = np.where(attn_target > 0.5)[0]\n            for i in roi_idx:\n                if i - 1 >= 0:\n                    attn_target[i-1] = max(attn_target[i-1], 0.5)\n                if i + 1 < S:\n                    attn_target[i+1] = max(attn_target[i+1], 0.5)\n\n        return {\n            \"uid\": uid,\n            \"slices\": torch.from_numpy(slices).to(self.dtype),  # [S,3,H,W]\n            \"modality_id\": torch.tensor(modality_id, dtype=torch.long),\n            \"labels\": torch.tensor(y, dtype=torch.float32),\n            \"attn_target\": torch.from_numpy(attn_target),\n            \"modality_str\": modality_str,\n        }\n\ndef _choose_indices(n: int, k: int, mode: str = \"uniform\", jitter: bool = False):\n    k = int(min(max(k, 1), n))\n    if k == n:\n        return np.arange(n, dtype=np.int64)\n    if mode == \"center\":\n        mid = n // 2; half = k // 2\n        s = max(0, mid - half); e = min(n, s + k); s = e - k\n        return np.arange(s, e, dtype=np.int64)\n    grid = np.linspace(0, n - 1, k + 2)[1:-1]\n    if jitter:\n        noise = np.random.uniform(-0.10, 0.10, size=grid.shape) * max(n-1, 1)\n        grid = np.clip(grid + noise, 0, n-1)\n    idx = np.unique(np.round(grid).astype(np.int64))\n    while len(idx) < k:\n        idx = np.unique(np.append(idx, np.random.randint(0, n)))\n    return idx[:k]\n\ndef collate_series_batch(batch):\n    if len(batch) == 0:\n        return None\n    max_s = max(x[\"slices\"].shape[0] for x in batch)\n    B = len(batch)\n    C, H, W = batch[0][\"slices\"].shape[1:]\n\n    batch_slices = torch.zeros((B, max_s, C, H, W), dtype=batch[0][\"slices\"].dtype)\n    batch_mask   = torch.zeros((B, max_s), dtype=torch.bool)\n    batch_attn_t = torch.zeros((B, max_s), dtype=torch.float32)\n    labels       = torch.zeros((B, NUM_LABELS), dtype=torch.float32)\n    modality_ids = torch.zeros((B,), dtype=torch.long)\n    uids         = []\n\n    for i, item in enumerate(batch):\n        s = item[\"slices\"]; n = s.shape[0]\n        batch_slices[i, :n] = s\n        batch_mask[i, :n]   = True\n        labels[i]           = item[\"labels\"]\n        modality_ids[i]     = item[\"modality_id\"]\n        uids.append(item[\"uid\"])\n        a = item[\"attn_target\"]\n        batch_attn_t[i, :min(n, a.shape[0])] = a[:n]\n\n    return {\n        \"uids\": uids,\n        \"slices\": batch_slices,\n        \"mask\": batch_mask,\n        \"labels\": labels,\n        \"modality_ids\": modality_ids,\n        \"attn_target\": batch_attn_t,\n    }\n\ndef cap_slices_on_cpu(batch_slices, batch_mask, max_s: int,\n                      mode: str = \"uniform\", stochastic: bool = False,\n                      attn_target=None):\n    \"\"\"\n    ROI-aware capping: always keep ROI slices (attn_target > 0.5),\n    then fill up to max_s from non-ROI.\n    \"\"\"\n    assert batch_slices.device.type == \"cpu\"\n    B, S, C, H, W = batch_slices.shape\n    n_valid = [int(batch_mask[i].sum().item()) for i in range(B)]\n    k_list  = [min(max_s, max(1, n)) for n in n_valid]\n    K = max(k_list) if B > 0 else 1\n\n    out_slices = torch.zeros((B, K, C, H, W), dtype=batch_slices.dtype)\n    out_masks  = torch.zeros((B, K), dtype=torch.bool)\n    out_attn_t = torch.zeros((B, K), dtype=torch.float32) if attn_target is not None else None\n\n    for i in range(B):\n        n_i = n_valid[i]; k_i = k_list[i]\n        if n_i <= 0:\n            out_slices[i, 0] = batch_slices[i, 0]\n            out_masks[i, 0]  = True\n            if out_attn_t is not None:\n                out_attn_t[i, 0] = 0.0\n            continue\n\n        valid_idx = np.arange(n_i)\n        roi_idx = np.array([], dtype=np.int64)\n        if attn_target is not None:\n            roi_idx = (attn_target[i, :n_i] > 0.5).nonzero(as_tuple=True)[0].cpu().numpy()\n        roi_idx = np.unique(roi_idx)\n\n        if k_i <= len(roi_idx):\n            idx = roi_idx[:k_i]\n        else:\n            non_roi = np.setdiff1d(valid_idx, roi_idx)\n            k_rest = k_i - len(roi_idx)\n            if len(non_roi) > 0:\n                base_idx = _choose_indices(len(non_roi), k_rest, mode=mode, jitter=stochastic)\n                idx_rest = non_roi[base_idx]\n                idx = np.concatenate([roi_idx, idx_rest])\n            else:\n                idx = roi_idx\n        while len(idx) < k_i:\n            idx = np.append(idx, idx[-1])\n        idx = idx[:k_i]\n\n        idx_t = torch.as_tensor(idx, dtype=torch.long)\n        sel = batch_slices[i, :n_i][idx_t]\n        out_slices[i, :k_i] = sel\n        out_masks[i, :k_i]  = True\n        if out_attn_t is not None:\n            out_attn_t[i, :k_i] = attn_target[i, :n_i][idx_t]\n\n    return out_slices, out_masks, out_attn_t","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:35:21.503974Z","iopub.execute_input":"2025-12-07T23:35:21.504272Z","iopub.status.idle":"2025-12-07T23:35:21.533725Z","shell.execute_reply.started":"2025-12-07T23:35:21.504248Z","shell.execute_reply":"2025-12-07T23:35:21.533106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 4. MODEL, LOSSES, METRICS\n# ============================================================\nclass EfficientNetV2S_Encoder(nn.Module):\n    def __init__(self, pretrained: bool = True, dropout: float = 0.0):\n        super().__init__()\n        from torchvision.models import efficientnet_v2_s, EfficientNet_V2_S_Weights\n        weights = EfficientNet_V2_S_Weights.IMAGENET1K_V1 if pretrained else None\n        model = efficientnet_v2_s(weights=weights)\n        self.features = model.features\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.drop = nn.Dropout(dropout) if dropout > 0 else nn.Identity()\n        with torch.no_grad():\n            x = torch.zeros(1, 3, 224, 224)\n            h = self.features(x); h = self.avgpool(h).flatten(1)\n            self.out_dim = h.shape[1]\n\n    def forward_slices(self, x: torch.Tensor) -> torch.Tensor:\n        h = self.features(x); h = self.avgpool(h).flatten(1)\n        return self.drop(h)\n\nclass TinyAttentionMIL(nn.Module):\n    def __init__(self, dim: int, hidden: int = 256, drop: float = 0.1):\n        super().__init__()\n        self.attn = nn.Sequential(\n            nn.Linear(dim, hidden),\n            nn.Tanh(),\n            nn.Dropout(drop),\n            nn.Linear(hidden, 1, bias=False),\n        )\n    def forward(self, H, mask):\n        A = self.attn(H).squeeze(-1)      # [B,S]\n        A = A.masked_fill(~mask, float(\"-inf\"))\n        A = torch.softmax(A, dim=1)\n        Z = torch.einsum(\"bs,bsd->bd\", A, H)\n        return Z, A\n\nclass AneurysmMILModel(nn.Module):\n    def __init__(\n        self,\n        mil_hidden: int = 256,\n        head_hidden: int = 768,\n        use_modality_embed: bool = True,\n        modality_embed_dim: int = 16,\n        encoder_pretrained: bool = True,\n        encoder_dropout: float = 0.2,\n    ):\n        super().__init__()\n        self.encoder = EfficientNetV2S_Encoder(pretrained=encoder_pretrained,\n                                               dropout=encoder_dropout)\n        enc_dim = self.encoder.out_dim\n        self.mil = TinyAttentionMIL(dim=enc_dim, hidden=mil_hidden, drop=0.1)\n        self.use_modality_embed = use_modality_embed\n        head_in = enc_dim + (modality_embed_dim if use_modality_embed else 0)\n        if use_modality_embed:\n            self.mod_embed = nn.Embedding(num_embeddings=len(MODALITY_TO_ID),\n                                          embedding_dim=modality_embed_dim)\n        self.head = nn.Sequential(\n            nn.LayerNorm(head_in),\n            nn.Linear(head_in, head_hidden),\n            nn.GELU(),\n            nn.Dropout(0.3),\n            nn.Linear(head_hidden, NUM_LABELS),\n        )\n\n    def forward(self, slices, mask, modality_ids=None):\n        B, S, C, H, W = slices.shape\n        chunk = max(1, int(CFG.get(\"encoder_chunk\", 64)))\n        feats_list = []\n        for start in range(0, S, chunk):\n            end = min(S, start + chunk)\n            x = slices[:, start:end].reshape(-1, C, H, W)\n            x = x.contiguous()\n            if CFG.get(\"channels_last\", False) and x.is_cuda and torch.cuda.device_count() <= 1:\n                x = x.contiguous(memory_format=torch.channels_last)\n            h = self.encoder.forward_slices(x)\n            feats_list.append(h.view(B, end - start, -1))\n        feats = torch.cat(feats_list, dim=1)   # [B,S,D]\n        Z, attn = self.mil(feats, mask)        # Z:[B,D], attn:[B,S]\n        if self.use_modality_embed and modality_ids is not None:\n            Z = torch.cat([Z, self.mod_embed(modality_ids)], dim=1)\n        logits = self.head(Z)\n        probs = torch.sigmoid(logits)\n        return {\"logits\": logits, \"probs\": probs, \"attn\": attn, \"bag\": Z}\n\nclass WeightedBCELossLS(nn.Module):\n    def __init__(self, aneurysm_weight: float = 13.0, smoothing: float = 0.02):\n        super().__init__()\n        w = torch.ones(NUM_LABELS, dtype=torch.float32)\n        w[LABEL_COLS.index(\"Aneurysm Present\")] = aneurysm_weight\n        self.register_buffer(\"w\", w)\n        self.smoothing = smoothing\n    def forward(self, logits, targets):\n        w = self.w if self.w.device == logits.device else self.w.to(logits.device)\n        eps = self.smoothing\n        targets = targets * (1 - eps) + 0.5 * eps\n        loss = F.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n        return (loss * w).mean()\n\ndef attention_supervision_loss(attn, attn_target, mask, alpha: float = 0.1):\n    attn = attn * mask\n    tgt  = attn_target * mask\n    sums = tgt.sum(dim=1, keepdim=True)\n    valid = (sums.squeeze(1) > 0)\n    if not valid.any():\n        return attn.sum() * 0.0\n    tgt_norm = torch.zeros_like(tgt)\n    tgt_norm[valid] = tgt[valid] / (sums[valid] + 1e-8)\n    loss = -(tgt_norm[valid] * (attn[valid] + 1e-8).log()).sum(dim=1).mean()\n    return alpha * loss\n\ndef columnwise_auc(y_true: np.ndarray, y_prob: np.ndarray, label_names):\n    aucs = {}\n    for i, name in enumerate(label_names):\n        yt, yp = y_true[:, i], y_prob[:, i]\n        if len(np.unique(yt)) < 2:\n            aucs[name] = np.nan\n        else:\n            aucs[name] = roc_auc_score(yt, yp)\n    return aucs\n\ndef weighted_mean_auc(aucs_dict: dict) -> float:\n    num = den = 0.0\n    for k, v in aucs_dict.items():\n        if np.isnan(v):\n            continue\n        w = 13.0 if k == \"Aneurysm Present\" else 1.0\n        num += w * v; den += w\n    return num / den if den > 0 else np.nan","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:35:35.466492Z","iopub.execute_input":"2025-12-07T23:35:35.467067Z","iopub.status.idle":"2025-12-07T23:35:35.484668Z","shell.execute_reply.started":"2025-12-07T23:35:35.467044Z","shell.execute_reply":"2025-12-07T23:35:35.483885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 5. LOAD LABELS, SPLIT DATA, BUILD DATALOADERS\n# ============================================================\nlabels_df = load_labels_frame(CFG[\"TRAIN_LABELS\"])\nindex_df  = pd.read_csv(CFG[\"PREPROC_INDEX\"])\n\n# fix npz paths if needed\nif \"npz_path\" in index_df.columns:\n    def fix_path(p):\n        fname = Path(p).name\n        return str(Path(CFG[\"NPZ_DIR\"]) / fname)\n    index_df[\"npz_path\"] = index_df[\"npz_path\"].apply(fix_path)\nelse:\n    raise ValueError(\"preproc_teacher_index.csv must have 'npz_path' column\")\n\n# keep only labeled series\nindex_labeled = index_df.merge(labels_df[[\"SeriesInstanceUID\"]],\n                               on=\"SeriesInstanceUID\", how=\"inner\")\ntmp = labels_df.set_index(\"SeriesInstanceUID\").loc[\n    index_labeled[\"SeriesInstanceUID\"], [\"Aneurysm Present\"]\n].reset_index()\nindex_labeled = index_labeled.merge(tmp, on=\"SeriesInstanceUID\", how=\"left\")\n\ny_strat = (\n    index_labeled[\"Modality\"].astype(str) + \"_\" +\n    tmp[\"Aneurysm Present\"].astype(int).astype(str)\n).values\n\nvc = pd.Series(y_strat).value_counts()\nsafe_folds = int(min(CFG[\"folds\"], max(2, vc.min()))) if len(vc) > 0 else 2\n\nif safe_folds < 2 or len(index_labeled) < 3:\n    rng = np.random.default_rng(CFG[\"seed\"])\n    perm = rng.permutation(len(index_labeled))\n    cut = max(1, int(0.8 * len(index_labeled)))\n    train_idx, val_idx = perm[:cut], perm[cut:]\nelse:\n    skf = StratifiedKFold(n_splits=safe_folds, shuffle=True,\n                          random_state=CFG[\"seed\"])\n    train_idx, val_idx = next(skf.split(index_labeled, y_strat))\n\ntrain_meta = index_labeled.iloc[train_idx].reset_index(drop=True)\nval_meta   = index_labeled.iloc[val_idx].reset_index(drop=True)\n\nif len(val_meta) == 0 and len(train_meta) > 0:\n    val_meta = train_meta.tail(1).copy()\n    train_meta = train_meta.iloc[:-1].reset_index(drop=True)\n\nTRAIN_SPLIT = \"/kaggle/working/train_split_teacher.csv\"\nVAL_SPLIT   = \"/kaggle/working/val_split_teacher.csv\"\ntrain_meta.to_csv(TRAIN_SPLIT, index=False)\nval_meta.to_csv(VAL_SPLIT, index=False)\nprint(f\"[SPLIT] train: {len(train_meta)} | val: {len(val_meta)} | folds used: {safe_folds}\")\n\nseries_to_roi_sops = build_series_to_roi_sops(CFG[\"LOCALIZERS_CSV\"])\n\ntrain_ds = RSNATeacherSeriesNPZ(\n    TRAIN_SPLIT,\n    labels_df=labels_df,\n    series_to_roi_sops=series_to_roi_sops,\n    load_into_mem=False,\n    do_aug=True,\n)\nval_ds = RSNATeacherSeriesNPZ(\n    VAL_SPLIT,\n    labels_df=labels_df,\n    series_to_roi_sops=series_to_roi_sops,\n    load_into_mem=False,\n    do_aug=False,\n)\n\ntrain_dl = DataLoader(\n    train_ds,\n    batch_size=CFG[\"batch_series\"],\n    sampler=RandomSampler(train_ds),\n    num_workers=CFG[\"num_workers\"],\n    pin_memory=True,\n    collate_fn=collate_series_batch,\n    persistent_workers=(CFG[\"num_workers\"] > 0),\n)\nval_dl = DataLoader(\n    val_ds,\n    batch_size=CFG[\"batch_series\"],\n    sampler=SequentialSampler(val_ds),\n    num_workers=CFG[\"num_workers\"],\n    pin_memory=True,\n    collate_fn=collate_series_batch,\n    persistent_workers=(CFG[\"num_workers\"] > 0),\n)\n\nanalysis_dl = DataLoader(\n    val_ds,\n    batch_size=1,\n    sampler=SequentialSampler(val_ds),\n    num_workers=0,\n    pin_memory=True,\n    collate_fn=collate_series_batch,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:35:41.023487Z","iopub.execute_input":"2025-12-07T23:35:41.023964Z","iopub.status.idle":"2025-12-07T23:35:41.240205Z","shell.execute_reply.started":"2025-12-07T23:35:41.023943Z","shell.execute_reply":"2025-12-07T23:35:41.239496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 6. MODEL, OPTIMIZER, EMA & TRAIN LOOP\n# ============================================================\nmodel = AneurysmMILModel(\n    mil_hidden=256,\n    head_hidden=768,\n    use_modality_embed=True,\n    modality_embed_dim=16,\n    encoder_pretrained=True,\n    encoder_dropout=0.2,\n).to(device)\n\nNUM_GPUS = torch.cuda.device_count() if device == \"cuda\" else 0\nif device == \"cuda\" and NUM_GPUS > 1:\n    print(f\"Using DataParallel on {NUM_GPUS} GPUs\")\n    model = nn.DataParallel(model)\n    CFG[\"channels_last\"] = False  # channels_last + DataParallel is messy\n\nif CFG[\"channels_last\"] and device == \"cuda\":\n    get_base_model(model).to(memory_format=torch.channels_last)\n\n# Freeze BN for stability\ndef freeze_bn(m):\n    for module in m.modules():\n        if isinstance(module, nn.BatchNorm2d):\n            module.eval()\n            for p in module.parameters():\n                p.requires_grad = False\n\nfreeze_bn(get_base_model(model))\n\ncriterion = WeightedBCELossLS(\n    aneurysm_weight=CFG[\"aneurysm_weight\"],\n    smoothing=CFG[\"label_smoothing\"],\n).to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=CFG[\"lr\"],\n    weight_decay=CFG[\"weight_decay\"],\n)\n\nsteps_per_epoch = max(1, len(train_dl))\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=3e-4,\n    div_factor=10.0,\n    total_steps=CFG[\"epochs\"] * steps_per_epoch,\n    pct_start=0.15,\n    final_div_factor=10.0,\n    anneal_strategy=\"cos\",\n)\n\nscaler = torch.amp.GradScaler(\n    'cuda',\n    enabled=(device == \"cuda\" and CFG[\"use_amp\"])\n)\n\n# Warm-up: freeze encoder\nfor p in get_base_model(model).encoder.parameters():\n    p.requires_grad = False\n\nclass EMA:\n    def __init__(self, model, decay=0.999):\n        self.decay = decay\n        self.shadow = {\n            k: v.clone().detach()\n            for k, v in get_base_model(model).state_dict().items()\n            if v.dtype.is_floating_point\n        }\n    @torch.no_grad()\n    def update(self, model):\n        sd = get_base_model(model).state_dict()\n        for k, v in self.shadow.items():\n            v.mul_(self.decay).add_(sd[k].detach(), alpha=1.0 - self.decay)\n    @torch.no_grad()\n    def copy_to(self, model):\n        sd = get_base_model(model).state_dict()\n        for k, v in self.shadow.items():\n            sd[k].copy_(v)\n\nema = EMA(model, decay=0.999)\n\nbest_wauc = -1.0\nbest_path = str(Path(CFG[\"ckpt_dir\"]) / CFG[\"best_ckpt_name\"])\nno_improve = 0\nEARLY_STOP = 6\n\ndef forward_pass(bag_slices, bag_mask, mods):\n    return model(bag_slices, bag_mask, mods)\n\ndef run_epoch(dataloader, train=True, stochastic_bag=False, tta=False, lambda_attn=0.0):\n    if len(dataloader) == 0:\n        return 0.0, float(\"nan\"), {}, 0.0, 0.0\n\n    model.train(train)\n    total_loss = total_loss_cls = total_loss_attn = 0.0\n    probs_cat, labels_cat = [], []\n\n    max_s_cap   = CFG[\"max_s_train\"] if train else CFG[\"max_s_val\"]\n    sample_mode = CFG[\"slice_sample\"]\n\n    for batch in dataloader:\n        if batch is None:\n            continue\n\n        slices_cpu = batch[\"slices\"]\n        mask_cpu   = batch[\"mask\"]\n        y_cpu      = batch[\"labels\"]\n        mods_cpu   = batch[\"modality_ids\"]\n        attn_t_cpu = batch[\"attn_target\"]\n\n        slices_cpu, mask_cpu, attn_t_cpu = cap_slices_on_cpu(\n            slices_cpu, mask_cpu, max_s=max_s_cap,\n            mode=sample_mode, stochastic=(train and stochastic_bag),\n            attn_target=attn_t_cpu,\n        )\n\n        slices = slices_cpu.to(device, non_blocking=True)\n        mask   = mask_cpu.to(device, non_blocking=True)\n        y      = y_cpu.to(device, non_blocking=True)\n        mods   = mods_cpu.to(device, non_blocking=True)\n        attn_t = attn_t_cpu.to(device, non_blocking=True)\n\n        bs = slices.size(0)\n\n        if train:\n            optimizer.zero_grad(set_to_none=True)\n            with torch.amp.autocast('cuda', enabled=(device == \"cuda\" and CFG[\"use_amp\"])):\n                out = forward_pass(slices, mask, mods)\n                loss_cls = criterion(out[\"logits\"], y)\n                if lambda_attn > 0:\n                    loss_attn = attention_supervision_loss(\n                        out[\"attn\"], attn_t, mask, alpha=lambda_attn\n                    )\n                else:\n                    loss_attn = torch.tensor(0.0, device=device)\n                loss = loss_cls + loss_attn\n            scaler.scale(loss).backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CFG[\"grad_clip\"])\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n            ema.update(model)\n        else:\n            with torch.no_grad():\n                with torch.amp.autocast('cuda', enabled=(device == \"cuda\" and CFG[\"use_amp\"])):\n                    if not tta:\n                        out = forward_pass(slices, mask, mods)\n                        loss_cls = criterion(out[\"logits\"], y)\n                    else:\n                        out1 = forward_pass(slices, mask, mods)\n                        if CFG[\"tta_hflip\"]:\n                            slices_flipped = slices.flip(-1)\n                            out2 = forward_pass(slices_flipped, mask, mods)\n                            logits = 0.5 * (out1[\"logits\"] + out2[\"logits\"])\n                            probs = torch.sigmoid(logits)\n                            out = {\"logits\": logits, \"probs\": probs}\n                        else:\n                            out = out1\n                        loss_cls = criterion(out[\"logits\"], y)\n                    loss_attn = torch.tensor(0.0, device=device)\n                    loss = loss_cls\n\n        total_loss      += float(loss.detach().cpu()) * bs\n        total_loss_cls  += float(loss_cls.detach().cpu()) * bs\n        total_loss_attn += float(loss_attn.detach().cpu()) * bs\n\n        probs_cat.append(out[\"probs\"].detach().cpu().numpy())\n        labels_cat.append(y.detach().cpu().numpy())\n\n    if len(labels_cat) == 0:\n        return 0.0, float(\"nan\"), {}, 0.0, 0.0\n\n    y_true_np = np.vstack(labels_cat)\n    y_prob_np = np.vstack(probs_cat)\n    aucs = columnwise_auc(y_true_np, y_prob_np, LABEL_COLS)\n    wauc = weighted_mean_auc(aucs)\n\n    N = len(dataloader.dataset)\n    avg_loss = total_loss / N\n    avg_cls  = total_loss_cls / N\n    avg_attn = total_loss_attn / N\n    return avg_loss, wauc, aucs, avg_cls, avg_attn\n\n@torch.no_grad()\ndef evaluate_with_ema(dataloader):\n    backup = {k: v.clone() for k, v in get_base_model(model).state_dict().items()}\n    ema.copy_to(model)\n    va_loss, va_wauc, va_aucs, _, _ = run_epoch(dataloader, train=False, tta=True)\n    get_base_model(model).load_state_dict(backup, strict=True)\n    return va_loss, va_wauc, va_aucs\n\n# history logging\nhistory = []\nhist_path = Path(CFG[\"ckpt_dir\"]) / \"training_history_teacher.csv\"\n\nfor epoch in range(1, CFG[\"epochs\"] + 1):\n    if epoch == CFG[\"freeze_encoder_epochs\"] + 1:\n        for p in get_base_model(model).encoder.parameters():\n            p.requires_grad = True\n\n    if   epoch <= 4:  lambda_attn = 0.0\n    elif epoch <= 8:  lambda_attn = CFG[\"lambda_attn\"] * 0.5\n    else:             lambda_attn = CFG[\"lambda_attn\"]\n\n    tr_loss, tr_wauc, _, tr_cls, tr_attn = run_epoch(\n        train_dl, train=True, stochastic_bag=True, lambda_attn=lambda_attn\n    )\n    va_loss, va_wauc, va_aucs = evaluate_with_ema(val_dl)\n\n    print(\n        f\"Epoch {epoch:02d} | train loss {tr_loss:.4f} \"\n        f\"(cls {tr_cls:.4f}, attn {tr_attn:.4f}) | \"\n        f\"val wAUC {va_wauc:.4f} | \"\n        f\"val Aneurysm AUC {va_aucs.get('Aneurysm Present', np.nan):.4f}\"\n    )\n\n    history.append({\n        \"epoch\": epoch,\n        \"train_loss\": float(tr_loss),\n        \"train_cls_loss\": float(tr_cls),\n        \"train_attn_loss\": float(tr_attn),\n        \"train_wAUC\": float(tr_wauc) if tr_wauc is not None else np.nan,\n        \"val_loss\": float(va_loss),\n        \"val_wAUC\": float(va_wauc),\n        \"val_auc_aneurysm\": float(va_aucs.get(\"Aneurysm Present\", np.nan)),\n    })\n    pd.DataFrame(history).to_csv(hist_path, index=False)\n\n    if not np.isnan(va_wauc) and va_wauc > best_wauc:\n        best_wauc = va_wauc\n        no_improve = 0\n        torch.save({\n            \"state_dict\": unwrap_state_dict(model),\n            \"best_wauc\": best_wauc,\n            \"label_order\": LABEL_COLS,\n        }, best_path)\n        print(f\"  ↳ Saved best to {best_path}\")\n    else:\n        no_improve += 1\n        if no_improve >= EARLY_STOP:\n            print(\"Early stopping (no val wAUC improvement).\")\n            last_path = str(Path(CFG[\"ckpt_dir\"]) / \"model_teacher_effs_mil_last.pt\")\n            torch.save({\n                \"state_dict\": unwrap_state_dict(model),\n                \"best_wauc\": float(best_wauc),\n                \"label_order\": list(LABEL_COLS),\n            }, last_path)\n            break\n\n    last_path = str(Path(CFG[\"ckpt_dir\"]) / \"model_teacher_effs_mil_last.pt\")\n    torch.save({\n        \"state_dict\": unwrap_state_dict(model),\n        \"epoch\": int(epoch),\n        \"label_order\": list(LABEL_COLS),\n    }, last_path)\n\nprint(\"Best weighted AUC:\", best_wauc)\n\n# Reload history for plots\nhist_df = pd.read_csv(hist_path)\nplt.figure(figsize=(6,4))\nplt.plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], marker=\"o\", label=\"Train loss\")\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_loss\"], marker=\"o\", label=\"Val loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend()\nplt.grid(True)\nplt.tight_layout()\nplt.show()\n\nplt.figure(figsize=(6,4))\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_wAUC\"], marker=\"o\", label=\"Val wAUC\")\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_auc_aneurysm\"], marker=\"o\", label=\"Val Aneurysm AUC\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"AUC\")\nplt.title(\"Validation AUCs over epochs\")\nplt.legend()\nplt.grid(True)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T18:31:16.218264Z","iopub.execute_input":"2025-12-07T18:31:16.21858Z","iopub.status.idle":"2025-12-07T22:30:03.828383Z","shell.execute_reply.started":"2025-12-07T18:31:16.218556Z","shell.execute_reply":"2025-12-07T22:30:03.827606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 7. ANALYSIS: LOAD BEST MODEL, COLLECT PREDICTIONS\n# ============================================================\n\nbest_path = \"/kaggle/input/efficientnetv2-mil-att-v3/model_teacher_effs_mil_best.pt\"\nstate = torch.load(best_path, map_location=device, weights_only=False)\nstate_dict = state[\"state_dict\"] if \"state_dict\" in state else state\n\nmodel = AneurysmMILModel(\n    mil_hidden=256,\n    head_hidden=768,\n    use_modality_embed=True,\n    modality_embed_dim=16,\n    encoder_pretrained=False,\n    encoder_dropout=0.2,\n).to(device)\nmodel.load_state_dict(state_dict, strict=True)\nmodel.eval()\nbase_model = get_base_model(model)\n\nANEUR_IDX = LABEL_COLS.index(\"Aneurysm Present\")\nprint(\"Model reloaded. Aneurysm label index:\", ANEUR_IDX)\n\ndef collect_series_predictions(dataloader, max_cases=None, thresh=0.5):\n    records = []\n    model.eval()\n    with torch.no_grad():\n        for i, batch in enumerate(dataloader):\n            if batch is None:\n                continue\n            if max_cases is not None and i >= max_cases:\n                break\n            uids   = batch[\"uids\"]\n            slices = batch[\"slices\"].to(device)\n            mask   = batch[\"mask\"].to(device)\n            labels = batch[\"labels\"].to(device)\n            mods   = batch[\"modality_ids\"].to(device)\n\n            out   = model(slices, mask, mods)\n            probs = out[\"probs\"].cpu().numpy()[0]\n            attn  = out[\"attn\"].cpu().numpy()[0]\n            mask_np = batch[\"mask\"].cpu().numpy()[0].astype(bool)\n\n            y_true_full = labels.cpu().numpy()[0]\n            uid   = uids[0]\n            mod_id = int(batch[\"modality_ids\"].cpu().numpy()[0])\n\n            gt   = int(y_true_full[ANEUR_IDX])\n            p    = float(probs[ANEUR_IDX])\n            pred = int(p >= thresh)\n\n            if   gt == 1 and pred == 1: cat = \"TP\"\n            elif gt == 0 and pred == 1: cat = \"FP\"\n            elif gt == 1 and pred == 0: cat = \"FN\"\n            else:                       cat = \"TN\"\n\n            records.append({\n                \"uid\": uid,\n                \"modality_id\": mod_id,\n                \"y_true\": gt,\n                \"prob\": p,\n                \"pred\": pred,\n                \"category\": cat,\n                \"attn\": attn[mask_np],\n            })\n    return records\n\nrecords = collect_series_predictions(analysis_dl, max_cases=None, thresh=0.5)\nprint(f\"Collected {len(records)} series for analysis.\")\n\n# Global metrics\ny_true = np.array([r[\"y_true\"] for r in records])\ny_prob = np.array([r[\"prob\"] for r in records])\ny_pred = np.array([r[\"pred\"] for r in records])\n\ncm = confusion_matrix(y_true, y_pred, labels=[0,1])\ntn, fp, fn, tp = cm.ravel()\nprint(\"Confusion matrix [[TN, FP], [FN, TP]]:\")\nprint(cm)\nprint(f\"\\nTN={tn}, FP={fp}, FN={fn}, TP={tp}\")\n\nsens = tp / (tp + fn + 1e-8)\nspec = tn / (tn + fp + 1e-8)\nprec = tp / (tp + fp + 1e-8)\nprint(f\"\\nSensitivity (Recall+): {sens:.3f}\")\nprint(f\"Specificity (Recall-): {spec:.3f}\")\nprint(f\"Precision:             {prec:.3f}\")\n\nprint(\"\\nClassification report:\")\nprint(classification_report(y_true, y_pred, target_names=[\"No aneurysm\", \"Aneurysm\"]))\n\nroc_auc = roc_auc_score(y_true, y_prob)\nprint(f\"ROC AUC (Aneurysm Present): {roc_auc:.3f}\")\n\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm,\n                              display_labels=[\"No aneurysm\", \"Aneurysm\"])\ndisp.plot(values_format=\"d\")\nplt.title(\"Confusion Matrix — Aneurysm Present\")\nplt.show()\n\nfpr, tpr, _ = roc_curve(y_true, y_prob)\nplt.figure()\nplt.plot(fpr, tpr, label=f\"AUC = {roc_auc:.3f}\")\nplt.plot([0, 1], [0, 1], linestyle=\"--\", color=\"grey\")\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.title(\"ROC Curve — Aneurysm Present\")\nplt.legend()\nplt.grid(True)\nplt.show()\n\nprec_curve, rec_curve, _ = precision_recall_curve(y_true, y_prob)\nap = average_precision_score(y_true, y_prob)\nplt.figure()\nplt.plot(rec_curve, prec_curve, label=f\"AP = {ap:.3f}\")\nplt.xlabel(\"Recall\")\nplt.ylabel(\"Precision\")\nplt.title(\"Precision–Recall Curve — Aneurysm Present\")\nplt.legend()\nplt.grid(True)\nplt.show()\n\n# Per-modality metrics\nmod_ids = np.array([r[\"modality_id\"] for r in records])\nfor mid, name in [(0, \"CTA\"), (1, \"MRA\")]:\n    mask = (mod_ids == mid)\n    if mask.sum() == 0:\n        continue\n    print(f\"\\n=== {name} only ===\")\n    y_true_m = y_true[mask]\n    y_prob_m = y_prob[mask]\n    y_pred_m = (y_prob_m >= 0.5).astype(int)\n    cm_m = confusion_matrix(y_true_m, y_pred_m, labels=[0,1])\n    print(\"Confusion Matrix:\\n\", cm_m)\n    print(\"ROC AUC:\", roc_auc_score(y_true_m, y_prob_m))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:49:47.644986Z","iopub.execute_input":"2025-12-07T23:49:47.645287Z","iopub.status.idle":"2025-12-07T23:57:24.699497Z","shell.execute_reply.started":"2025-12-07T23:49:47.645262Z","shell.execute_reply":"2025-12-07T23:57:24.69889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 8. GRAD-CAM + ATTENTION VISUALIZATION + DASHBOARD\n# ============================================================\n# Install once in the notebook\n!pip install grad-cam -q\n\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\n\nNPZ_DIR = Path(CFG[\"NPZ_DIR\"])\n\nclass MILGradCAMWrapper(nn.Module):\n    def __init__(self, mil_model, modality_id: int):\n        super().__init__()\n        self.mil_model = mil_model\n        self.modality_id = modality_id\n    def forward(self, x):\n        device_local = next(self.mil_model.parameters()).device\n        x = x.to(device_local)\n        S = x.shape[0]\n        B = 1\n        bag  = x.unsqueeze(0)\n        mask = torch.ones((B,S), dtype=torch.bool, device=device_local)\n        mods = torch.full((B,), self.modality_id, dtype=torch.long, device=device_local)\n        out  = self.mil_model(bag, mask, mods)\n        return out[\"logits\"]\n\n\n\ncheckpoint_path = \"/kaggle/input/efficientnetv2-mil-att-v3/model_teacher_effs_mil_best.pt\"\nprint(\"Loading best checkpoint from:\", checkpoint_path)\n\ncheckpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)\nstate_dict = checkpoint[\"state_dict\"] if \"state_dict\" in checkpoint else checkpoint\n\nbase_model = AneurysmMILModel(\n    mil_hidden=256,\n    head_hidden=768,\n    use_modality_embed=True,\n    modality_embed_dim=16,\n    encoder_pretrained=False,   \n    encoder_dropout=0.2,\n).to(device)\n\n# Load weights\nmissing, unexpected = base_model.load_state_dict(state_dict, strict=False)\nprint(\"Missing keys:\", missing)\nprint(\"Unexpected keys:\", unexpected)\n\nbase_model.eval()\nprint(\"Checkpoint loaded successfully.\")\n\n    \ntarget_layers = [base_model.encoder.features[-1]]\ncta_wrapper   = MILGradCAMWrapper(base_model, MODALITY_TO_ID[\"CTA\"]).to(device)\nmra_wrapper   = MILGradCAMWrapper(base_model, MODALITY_TO_ID[\"MRA\"]).to(device)\ncam = GradCAM(model=cta_wrapper, target_layers=target_layers)\n\ndef run_gradcam_for_uid_ax(uid: str,\n                           modality_str: str = \"CTA\",\n                           slice_idx: int = None,\n                           target_label_idx: int = ANEUR_IDX,\n                           ax=None):\n    npz_path = NPZ_DIR / f\"{uid}.npz\"\n    if not npz_path.exists():\n        print(f\"[WARN] NPZ not found: {npz_path}\")\n        return\n    data = np.load(npz_path, allow_pickle=True)\n    vol  = data[\"volume\"].astype(np.float32) / 255.0\n    Z, C, H, W = vol.shape\n    if slice_idx is None:\n        slice_idx = Z // 2\n    slice_idx = max(0, min(slice_idx, Z-1))\n    slice_chw = vol[slice_idx]\n    slice_hwc = np.transpose(slice_chw, (1,2,0))\n\n    mod_str_u = modality_str.upper()\n    if mod_str_u == \"CTA\":\n        cam.model = cta_wrapper\n    elif mod_str_u == \"MRA\":\n        cam.model = mra_wrapper\n    else:\n        print(f\"[WARN] Unknown modality {modality_str}, defaulting to CTA.\")\n        cam.model = cta_wrapper\n\n    input_tensor = torch.from_numpy(slice_chw).unsqueeze(0)\n    targets = [ClassifierOutputTarget(target_label_idx)]\n    grayscale_cam = cam(input_tensor=input_tensor, targets=targets)[0]\n    visualization = show_cam_on_image(slice_hwc, grayscale_cam, use_rgb=True)\n\n    if ax is None:\n        plt.figure(figsize=(6,6))\n        plt.imshow(visualization)\n        plt.axis(\"off\")\n        plt.title(f\"Grad-CAM: {modality_str}, UID={uid}, slice={slice_idx}\")\n        plt.show()\n    else:\n        ax.imshow(visualization)\n        ax.axis(\"off\")\n\ndef pick_examples(records, per_cat=3):\n    buckets = defaultdict(list)\n    for r in records:\n        if len(buckets[r[\"category\"]]) < per_cat:\n            buckets[r[\"category\"]].append(r)\n    return buckets\n\nexamples = pick_examples(records, per_cat=3)\n\ndef plot_attention_example(rec, label_str=\"Aneurysm Present\"):\n    uid  = rec[\"uid\"]\n    attn = rec[\"attn\"]\n    npz_path = NPZ_DIR / f\"{uid}.npz\"\n    if not npz_path.exists():\n        print(f\"[WARN] NPZ not found for {uid} at {npz_path}, skipping.\")\n        return\n    data = np.load(npz_path, allow_pickle=True)\n    vol  = data[\"volume\"]\n    Z, C, H, W = vol.shape\n    S     = len(attn)\n    z_idx = int(np.argmax(attn)) if S > 0 else 0\n    z_idx = min(z_idx, Z-1)\n    img = vol[z_idx,0].astype(np.float32) / 255.0\n\n    fig, axes = plt.subplots(1,2, figsize=(12,4))\n    axes[0].imshow(img, cmap=\"gray\")\n    axes[0].set_title(\n        f\"{uid}\\ncat={rec['category']}, prob={rec['prob']:.3f}, gt={rec['y_true']}\"\n    )\n    axes[0].axis(\"off\")\n    axes[1].plot(np.arange(S), attn, marker=\"o\")\n    axes[1].axvline(z_idx, linestyle=\"--\")\n    axes[1].set_xlabel(\"Slice index\")\n    axes[1].set_ylabel(\"Attention weight\")\n    axes[1].set_title(f\"Attention over slices ({label_str})\")\n    axes[1].grid(True)\n    plt.tight_layout()\n    plt.show()\n\ndef show_gradcam_examples_from_records(examples, category: str, k: int = 3):\n    recs = examples.get(category, [])\n    if not recs:\n        print(f\"No records found for category={category}\")\n        return\n    recs = recs[:k]\n    print(f\"\\n=== {category} Grad-CAM examples (up to {k}) ===\")\n    for i, rec in enumerate(recs):\n        uid   = rec[\"uid\"]\n        attn  = rec[\"attn\"]\n        y     = rec[\"y_true\"]\n        prob  = rec[\"prob\"]\n        slice_idx = int(np.argmax(attn)) if len(attn) > 0 else 0\n        modality_str = index_df.loc[\n            index_df[\"SeriesInstanceUID\"] == uid, \"Modality\"\n        ].iloc[0]\n        print(f\"[{category} #{i}] UID={uid}, Modality={modality_str}, y={y}, prob={prob:.3f}, slice_idx={slice_idx}\")\n        run_gradcam_for_uid_ax(uid, modality_str=modality_str,\n                               slice_idx=slice_idx, target_label_idx=ANEUR_IDX)\n        plot_attention_example(rec, label_str=\"Aneurysm Present\")\n\n\ndef plot_performance_dashboard(hist_df, records, index_df,\n                               figsize=(12,10), thresh=0.5):\n    y_true = np.array([r[\"y_true\"] for r in records])\n    y_prob = np.array([r[\"prob\"] for r in records])\n    y_pred = (y_prob >= thresh).astype(int)\n\n    cm  = confusion_matrix(y_true, y_pred, labels=[0,1])\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    roc_auc = roc_auc_score(y_true, y_prob)\n\n    # choose example: FN > FP > TP\n    cat_order = [\"FN\",\"FP\",\"TP\"]\n    chosen_rec = None\n    for cat in cat_order:\n        cand = [r for r in records if r[\"category\"] == cat]\n        if cand:\n            chosen_rec = cand[0]\n            break\n    if chosen_rec is None and records:\n        chosen_rec = records[0]\n\n    fig, axes = plt.subplots(2,2, figsize=figsize)\n\n    # (1,1) Train vs val loss (+ optional val_wAUC)\n    ax = axes[0,0]\n    ax.plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], marker=\"o\", label=\"Train loss\")\n    ax.plot(hist_df[\"epoch\"], hist_df[\"val_loss\"], marker=\"o\", label=\"Val loss\")\n    ax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"Loss\")\n    ax.set_title(\"Training vs Validation Loss\")\n    ax.legend(); ax.grid(True)\n    if \"val_wAUC\" in hist_df.columns:\n        ax2 = ax.twinx()\n        ax2.plot(hist_df[\"epoch\"], hist_df[\"val_wAUC\"],\n                 marker=\"s\", linestyle=\"--\", alpha=0.7, label=\"Val wAUC\")\n        ax2.set_ylabel(\"Val wAUC\")\n        ax2.legend(loc=\"lower right\")\n\n    # (1,2) ROC\n    ax = axes[0,1]\n    ax.plot(fpr, tpr, label=f\"ROC AUC = {roc_auc:.3f}\")\n    ax.plot([0,1],[0,1], linestyle=\"--\", color=\"grey\")\n    ax.set_xlabel(\"False Positive Rate\"); ax.set_ylabel(\"True Positive Rate\")\n    ax.set_title(\"ROC — Aneurysm Present\")\n    ax.legend(); ax.grid(True)\n\n    # (2,1) Confusion matrix\n    ax = axes[1,0]\n    disp = ConfusionMatrixDisplay(confusion_matrix=cm,\n                                  display_labels=[\"No aneurysm\", \"Aneurysm\"])\n    disp.plot(ax=ax, values_format=\"d\", colorbar=False)\n    ax.set_title(\"Confusion Matrix\")\n\n    # (2,2) Grad-CAM example\n    ax = axes[1,1]\n    if chosen_rec is not None:\n        uid   = chosen_rec[\"uid\"]\n        attn  = chosen_rec[\"attn\"]\n        prob  = chosen_rec[\"prob\"]\n        gt    = chosen_rec[\"y_true\"]\n        cat   = chosen_rec[\"category\"]\n        slice_idx = int(np.argmax(attn)) if len(attn) > 0 else 0\n        modality_str = index_df.loc[\n            index_df[\"SeriesInstanceUID\"] == uid, \"Modality\"\n        ].iloc[0]\n        run_gradcam_for_uid_ax(uid=uid, modality_str=modality_str,\n                               slice_idx=slice_idx, target_label_idx=ANEUR_IDX, ax=ax)\n        ax.set_title(\n            f\"Grad-CAM ({cat})\\nUID={uid}, Mod={modality_str}, \"\n            f\"gt={gt}, prob={prob:.3f}, slice={slice_idx}\"\n        )\n    else:\n        ax.text(0.5,0.5,\"No records available\", ha=\"center\", va=\"center\")\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:57:24.700846Z","iopub.execute_input":"2025-12-07T23:57:24.701078Z","iopub.status.idle":"2025-12-07T23:57:28.748989Z","shell.execute_reply.started":"2025-12-07T23:57:24.701046Z","shell.execute_reply":"2025-12-07T23:57:28.748274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for cat in [\"TP\", \"FP\", \"FN\"]:\n    show_gradcam_examples_from_records(examples, category=cat, k=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:57:28.749981Z","iopub.execute_input":"2025-12-07T23:57:28.750224Z","iopub.status.idle":"2025-12-07T23:57:35.847675Z","shell.execute_reply.started":"2025-12-07T23:57:28.750201Z","shell.execute_reply":"2025-12-07T23:57:35.846871Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_performance_dashboard(hist_df, records, index_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T22:40:07.619303Z","iopub.execute_input":"2025-12-07T22:40:07.619481Z","iopub.status.idle":"2025-12-07T22:40:08.571794Z","shell.execute_reply.started":"2025-12-07T22:40:07.619466Z","shell.execute_reply":"2025-12-07T22:40:08.571165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}