{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"datasetVersion","sourceId":2542390,"datasetId":1541666,"databundleVersionId":2585349},{"sourceType":"datasetVersion","sourceId":15497858,"datasetId":9915331,"databundleVersionId":16423200},{"sourceType":"datasetVersion","sourceId":15803031,"datasetId":10129754,"databundleVersionId":16750348},{"sourceType":"datasetVersion","sourceId":15443056,"datasetId":9879442,"databundleVersionId":16363305},{"sourceType":"datasetVersion","sourceId":15782534,"datasetId":10115822,"databundleVersionId":16728281},{"sourceType":"datasetVersion","sourceId":15499316,"datasetId":9916286,"databundleVersionId":16424782},{"sourceType":"datasetVersion","sourceId":15160735,"datasetId":9708077,"databundleVersionId":16051896}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🧬 Multimodal Radiogenomics: MRI + Genomics for Brain Tumor Grading\n**Debasish Mondal — BT23CSE108**  \n","metadata":{}},{"cell_type":"code","source":"# ========================================\n# 1. SETUP & ENVIRONMENT (RUN AFTER RESTART)\n# ========================================\n\n!pip install nibabel SimpleITK grad-cam -q\n\nimport os\nimport tarfile\nimport random\nfrom glob import glob\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nimport nibabel as nib\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision.models as models\nimport torchvision.transforms as transforms\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score, f1_score, confusion_matrix\n\n# ----------------------------------------\n# Device configuration\n# ----------------------------------------\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"PyTorch:\", torch.__version__)\nprint(\"CUDA available:\", torch.cuda.is_available())\n\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\nelse:\n    print(\"GPU: None\")\n\n# ----------------------------------------\n# Reproducibility\n# ----------------------------------------\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n    import torch.backends.cudnn as cudnn\n\n    cudnn.deterministic = True\n    cudnn.benchmark = False\n\nseed_everything()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T05:31:23.899695Z","iopub.execute_input":"2026-04-28T05:31:23.899884Z","iopub.status.idle":"2026-04-28T05:31:42.640483Z","shell.execute_reply.started":"2026-04-28T05:31:23.899863Z","shell.execute_reply":"2026-04-28T05:31:42.639197Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔄 Section 2: Session Recovery\n> Run this every time after notebook restart. Auto re-extracts BraTS and rebuilds labels if needed.","metadata":{}},{"cell_type":"code","source":"#RUN THESE EVERYTHIME AFTER RESTART\nimport os\nimport tarfile\n\nBRATS_ROOT = \"/kaggle/working/brats2021\"\n\n# Re-extract BraTS if session reset\nif not os.path.exists(BRATS_ROOT) or len(os.listdir(BRATS_ROOT)) < 100:\n    print(\"BraTS data not found — re-extracting...\")\n    os.makedirs(BRATS_ROOT, exist_ok=True)\n    tar_path = \"/kaggle/input/datasets/dschettler8845/brats-2021-task1/BraTS2021_Training_Data.tar\"\n    with tarfile.open(tar_path, 'r') as tar:\n        tar.extractall(BRATS_ROOT)\n    print(\"Re-extraction done!\")\nelse:\n    print(f\"✅ BraTS data ready — {len(os.listdir(BRATS_ROOT))} patients found\")\n\n# Re-build labels if session reset\nlabels_path = \"/kaggle/working/labels_full.csv\"\nif not os.path.exists(labels_path):\n    print(\"Labels not found — rebuilding...\")\n    import pandas as pd\n    brats_patients = sorted([p for p in os.listdir(BRATS_ROOT) if p.startswith('BraTS')])\n    brats_ids = [int(p.split('_')[1]) for p in brats_patients]\n    df_full = pd.DataFrame({'BraTS21ID': brats_ids, 'patient_folder': brats_patients})\n    df_full['tumor_type'] = df_full['BraTS21ID'].apply(lambda x: 'GBM' if x < 1000 else 'LGG')\n    df_full['grade'] = df_full['tumor_type'].apply(lambda x: 2 if x == 'GBM' else 1)\n    df_full['idh'] = df_full['tumor_type'].apply(lambda x: 0 if x == 'GBM' else 1)\n    rsna = pd.read_csv(\"/kaggle/input/competitions/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv\")\n    df_full = df_full.merge(rsna, on='BraTS21ID', how='left')\n    df_full.rename(columns={'MGMT_value': 'mgmt'}, inplace=True)\n    df_full.to_csv(labels_path, index=False)\n    print(\"Labels rebuilt!\")\nelse:\n    df_full = pd.read_csv(labels_path)\n    print(f\"✅ Labels ready — {len(df_full)} patients\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T05:32:00.640890Z","iopub.execute_input":"2026-04-28T05:32:00.641159Z","iopub.status.idle":"2026-04-28T05:33:56.368245Z","shell.execute_reply.started":"2026-04-28T05:32:00.641132Z","shell.execute_reply":"2026-04-28T05:33:56.366572Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧠 Section 3: Preprocessing Pipeline\nSeg-mask guided slice selection — tumor center varies per patient (EDA confirmed).","metadata":{}},{"cell_type":"code","source":"# ── CELL 3: PREPROCESSING — load_patient_slice() ─────────────────────────────\n# RUN THESE EVERY TIME AFTER RESTART\nimport cv2\nimport numpy as np\n\ndef load_patient_slice(patient_dir, patient_id, slice_idx=None):\n    modalities = ['t1', 't1ce', 't2', 'flair']\n    channels = []\n\n    # SMARTER SLICE SELECTION — use seg mask to find tumor center\n    if slice_idx is None:\n        seg_path = f\"{patient_dir}/{patient_id}_seg.nii.gz\"\n        if os.path.exists(seg_path):\n            seg = nib.load(seg_path).get_fdata()\n            tumor_slices = np.where(seg.sum(axis=(0, 1)) > 0)[0]\n            if len(tumor_slices) > 0:\n                # Pick middle of actual tumor extent\n                slice_idx = int(tumor_slices[len(tumor_slices) // 2])\n            else:\n                slice_idx = 80  # EDA-confirmed fallback\n        else:\n            slice_idx = 80\n\n    for mod in modalities:\n        path = f\"{patient_dir}/{patient_id}_{mod}.nii.gz\"\n        vol = nib.load(path).get_fdata()\n        slc = vol[:, :, slice_idx].astype(np.float32)\n\n        # Z-score normalize using only non-zero (brain) voxels\n        mask = slc > 0\n        if mask.sum() > 0:\n            slc[mask] = (slc[mask] - slc[mask].mean()) / (slc[mask].std() + 1e-8)\n\n        # Resize to 224x224\n        slc = cv2.resize(slc, (224, 224))\n        channels.append(slc)\n\n    return np.stack(channels, axis=0)\n\n\n# ── VERIFY: show which slice was selected per patient ────────────────────────\ntest_patients = [\n    \"BraTS2021_00000\",  # GBM\n    \"BraTS2021_01000\",  # LGG\n    \"BraTS2021_01100\",  # LGG\n]\n\nprint(f\"{'Patient':<25} {'Type':<6} {'Seg Slice Selected':<20} {'Old (fixed 80)'}\")\nprint(\"-\" * 65)\nfor pid in test_patients:\n    pdir = f\"{BRATS_ROOT}/{pid}\"\n    seg = nib.load(f\"{pdir}/{pid}_seg.nii.gz\").get_fdata()\n    tumor_slices = np.where(seg.sum(axis=(0, 1)) > 0)[0]\n    if len(tumor_slices) > 0:\n        smart_idx = int(tumor_slices[len(tumor_slices) // 2])\n    else:\n        smart_idx = 80\n    ptype = \"GBM\" if int(pid.split(\"_\")[1]) < 1000 else \"LGG\"\n    print(f\"{pid:<25} {ptype:<6} {smart_idx:<20} 80\")\n\n# ── TEST OUTPUT SHAPE ────────────────────────────────────────────────────────\ntensor = load_patient_slice(f\"{BRATS_ROOT}/BraTS2021_00000\", \"BraTS2021_00000\")\nprint(\"\\nOutput shape:\", tensor.shape)\nprint(\"Channel 0 (T1)   — min:\", tensor[0].min().round(3), \"max:\", tensor[0].max().round(3))\nprint(\"Channel 1 (T1ce) — min:\", tensor[1].min().round(3), \"max:\", tensor[1].max().round(3))\nprint(\"Channel 2 (T2)   — min:\", tensor[2].min().round(3), \"max:\", tensor[2].max().round(3))\nprint(\"Channel 3 (FLAIR)— min:\", tensor[3].min().round(3), \"max:\", tensor[3].max().round(3))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T05:39:45.228370Z","iopub.execute_input":"2026-04-28T05:39:45.230249Z","iopub.status.idle":"2026-04-28T05:39:45.920789Z","shell.execute_reply.started":"2026-04-28T05:39:45.230163Z","shell.execute_reply":"2026-04-28T05:39:45.918889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Step 1: Build base dataframe from all BraTS 2021 patients\nbrats_patients = sorted([p for p in os.listdir(BRATS_ROOT) if p.startswith('BraTS')])\nbrats_ids = [int(p.split('_')[1]) for p in brats_patients]\n\ndf = pd.DataFrame({\n    'BraTS21ID': brats_ids,\n    'patient_folder': brats_patients\n})\n\n# Step 2: Assign tumor type based on ID range\ndf['tumor_type'] = df['BraTS21ID'].apply(lambda x: 'GBM' if x < 1000 else 'LGG')\n\n# Step 3: Assign Grade\n# GBM = Grade IV, LGG = Grade II or III (we'll refine later)\ndf['grade'] = df['tumor_type'].apply(lambda x: 2 if x == 'GBM' else 1)\n# 2 = GBM/Grade IV, 1 = LGG/Grade II-III\n\n# Step 4: Assign IDH\n# GBM = mostly IDH wildtype (0), LGG = mostly IDH mutant (1)\ndf['idh'] = df['tumor_type'].apply(lambda x: 0 if x == 'GBM' else 1)\n\n# Step 5: Merge MGMT from RSNA labels\nrsna = pd.read_csv(\"/kaggle/input/competitions/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv\")\ndf = df.merge(rsna, on='BraTS21ID', how='left')\ndf.rename(columns={'MGMT_value': 'mgmt'}, inplace=True)\n\nprint(df.shape)\nprint(df.head(10))\nprint(\"\\nTumor type distribution:\")\nprint(df['tumor_type'].value_counts())\nprint(\"\\nMGMT available:\", df['mgmt'].notna().sum())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T05:39:48.803264Z","iopub.execute_input":"2026-04-28T05:39:48.803543Z","iopub.status.idle":"2026-04-28T05:39:48.910126Z","shell.execute_reply.started":"2026-04-28T05:39:48.803522Z","shell.execute_reply":"2026-04-28T05:39:48.908767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Full dataset — all 1251 patients have grade and IDH\n# MGMT is NaN for patients without labels — that's intentional\n\ndf_full = df.copy()\n\n# For patients without MGMT, keep NaN — we'll handle this in the loss function\n# Grade and IDH are available for all 1251 patients\n\nprint(\"Total patients:\", len(df_full))\nprint(\"\\nGrade distribution (all 1251):\")\nprint(df_full['grade'].value_counts())\nprint(\"\\nIDH distribution (all 1251):\")\nprint(df_full['idh'].value_counts())\nprint(\"\\nMGMT available:\", df_full['mgmt'].notna().sum())\nprint(\"MGMT missing:\", df_full['mgmt'].isna().sum())\n\n# Save full label file\ndf_full.to_csv(\"/kaggle/working/labels_full.csv\", index=False)\nprint(\"\\nSaved to /kaggle/working/labels_full.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T05:39:51.992432Z","iopub.execute_input":"2026-04-28T05:39:51.992773Z","iopub.status.idle":"2026-04-28T05:39:52.009245Z","shell.execute_reply.started":"2026-04-28T05:39:51.992747Z","shell.execute_reply":"2026-04-28T05:39:52.008484Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🏷️ Section 4: Label Construction\nIDH + Grade from BraTS ID range split. MGMT from RSNA 2021 competition CSV. Partial label strategy for 677 missing MGMT patients.","metadata":{}},{"cell_type":"code","source":"# ========================================\n# DATASET (FINAL v5 — BALANCED AUGMENTATION)\n# ========================================\n\nfrom torch.utils.data import Dataset\nimport torch\nimport pandas as pd\nimport random\n\nclass BraTSDataset(Dataset):\n    \n    def __init__(self, df, brats_root, transform=False):\n        self.df = df.reset_index(drop=True)\n        self.brats_root = brats_root\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        \n        row = self.df.iloc[idx]\n        \n        # ======================\n        # LOAD MRI\n        # ======================\n        patient_id = row['patient_folder']\n        patient_dir = f\"{self.brats_root}/{patient_id}\"\n        \n        img = load_patient_slice(patient_dir, patient_id)  # (4, 224, 224)\n        img = torch.from_numpy(img).float()\n\n        # ======================\n        # AUGMENTATION (v5 FINAL)\n        # ======================\n        if self.transform:\n            \n            # 🔥 Horizontal flip\n            if random.random() > 0.5:\n                img = torch.flip(img, dims=[2])\n            \n            # 🔥 Rotation (balanced)\n            if random.random() > 0.3:\n                k = random.randint(0, 3)  # 0, 90, 180, 270\n                img = torch.rot90(img, k=k, dims=[1, 2])\n            \n            # 🔥 VERY LIGHT NOISE (KEY FIX)\n            if random.random() > 0.7:\n                img = img + torch.randn_like(img) * 0.005\n\n        # ======================\n        # LABELS\n        # ======================\n        \n        # IDH\n        idh = torch.tensor(row['idh'], dtype=torch.float32)\n\n        # Grade mapping\n        grade_val = row['grade']\n        if grade_val == 1:\n            grade = torch.tensor(0, dtype=torch.long)\n        elif grade_val == 2:\n            grade = torch.tensor(1, dtype=torch.long)\n        else:\n            raise ValueError(f\"Invalid grade value: {grade_val}\")\n\n        # MGMT (masked)\n        mgmt_val = row['mgmt']\n        if pd.isna(mgmt_val):\n            mgmt = torch.tensor(-1.0, dtype=torch.float32)\n        else:\n            mgmt = torch.tensor(mgmt_val, dtype=torch.float32)\n\n        labels = {\n            'idh': idh,\n            'mgmt': mgmt,\n            'grade': grade\n        }\n\n        return img, labels\n\n\n# ========================================\n# TEST\n# ========================================\n\nBRATS_ROOT = \"/kaggle/working/brats2021\"\n\ntrain_dataset = BraTSDataset(df, BRATS_ROOT, transform=True)\nval_dataset   = BraTSDataset(df, BRATS_ROOT, transform=False)\n\nimg, labels = train_dataset[0]\n\nprint(\"Image shape:\", img.shape)\nprint(\"Labels:\", labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:12:53.584292Z","iopub.execute_input":"2026-04-19T21:12:53.585509Z","iopub.status.idle":"2026-04-19T21:12:54.194521Z","shell.execute_reply.started":"2026-04-19T21:12:53.585459Z","shell.execute_reply":"2026-04-19T21:12:54.193728Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🏗️ Section 5: Dataset Class & DataLoaders","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nfrom sklearn.model_selection import train_test_split\n\n# ======================\n# SPLIT\n# ======================\ntrain_df, temp_df = train_test_split(\n    df_full,\n    test_size=0.3,\n    random_state=42,\n    stratify=df_full['grade']\n)\n\nval_df, test_df = train_test_split(\n    temp_df,\n    test_size=0.5,\n    random_state=42,\n    stratify=temp_df['grade']\n)\n\nprint(f\"Train: {len(train_df)} | Val: {len(val_df)} | Test: {len(test_df)}\")\n\n# ======================\n# DATASETS\n# ======================\n\ntrain_dataset = BraTSDataset(train_df, BRATS_ROOT, transform=True)\nval_dataset   = BraTSDataset(val_df,   BRATS_ROOT, transform=False)\ntest_dataset  = BraTSDataset(test_df,  BRATS_ROOT, transform=False)\n\n# ======================\n# DATALOADERS (FIXED)\n# ======================\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=16,\n    shuffle=True,\n    num_workers=0,          # 🔥 FIX: no multiprocessing issues\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=16,\n    shuffle=False,\n    num_workers=0,          # 🔥 FIX\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=16,\n    shuffle=False,\n    num_workers=0,          # 🔥 FIX\n    pin_memory=True\n)\n\nprint(f\"Train batches: {len(train_loader)}\")\nprint(f\"Val batches:   {len(val_loader)}\")\nprint(f\"Test batches:  {len(test_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:12:57.195272Z","iopub.execute_input":"2026-04-19T21:12:57.195590Z","iopub.status.idle":"2026-04-19T21:12:57.217755Z","shell.execute_reply.started":"2026-04-19T21:12:57.195564Z","shell.execute_reply":"2026-04-19T21:12:57.216660Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 📊 Section 6: Exploratory Data Analysis\n**Key findings:** Best slice = tumor center not fixed index. MGMT labels 97% GBM. Tumor volume similar across GBM/LGG.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nfig, axes = plt.subplots(1, 3, figsize=(15, 5))\nfig.suptitle('BraTS 2021 + TCGA Label Distribution', fontsize=14, fontweight='bold')\n\n# IDH\naxes[0].bar(['IDH Wildtype\\n(GBM)', 'IDH Mutant\\n(LGG)'], [585, 666], \n            color=['#E74C3C', '#2ECC71'], edgecolor='white', linewidth=1.5)\naxes[0].set_title('IDH Mutation Status', fontweight='bold')\naxes[0].set_ylabel('Number of Patients')\nfor i, v in enumerate([585, 666]):\n    axes[0].text(i, v+10, str(v), ha='center', fontweight='bold')\n\n# Grade\naxes[1].bar(['GBM\\n(Grade IV)', 'LGG\\n(Grade II/III)'], [585, 666],\n            color=['#E74C3C', '#3498DB'], edgecolor='white', linewidth=1.5)\naxes[1].set_title('Tumor Grade Distribution', fontweight='bold')\nfor i, v in enumerate([585, 666]):\n    axes[1].text(i, v+10, str(v), ha='center', fontweight='bold')\n\n# MGMT\naxes[2].bar(['Unmethylated', 'Methylated', 'Not Available'], [276, 301, 674],\n            color=['#E67E22', '#9B59B6', '#95A5A6'], edgecolor='white', linewidth=1.5)\naxes[2].set_title('MGMT Methylation Status', fontweight='bold')\nfor i, v in enumerate([276, 301, 674]):\n    axes[2].text(i, v+10, str(v), ha='center', fontweight='bold')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/label_distribution.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T13:55:44.308695Z","iopub.execute_input":"2026-03-30T13:55:44.309584Z","iopub.status.idle":"2026-03-30T13:55:45.352890Z","shell.execute_reply.started":"2026-03-30T13:55:44.309545Z","shell.execute_reply":"2026-03-30T13:55:45.351563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(3, 4, figsize=(16, 12))\nfig.suptitle('BraTS 2021 — 4 MRI Modalities Across 3 Patients', fontsize=14, fontweight='bold')\n\npatients = ['BraTS2021_00000', 'BraTS2021_00500', 'BraTS2021_01200']\nmodalities = ['t1', 't1ce', 't2', 'flair']\nmod_titles = ['T1', 'T1ce (Contrast)', 'T2', 'FLAIR']\n\nfor row, pid in enumerate(patients):\n    pdir = f\"{BRATS_ROOT}/{pid}\"\n    for col, (mod, title) in enumerate(zip(modalities, mod_titles)):\n        vol = nib.load(f\"{pdir}/{pid}_{mod}.nii.gz\").get_fdata()\n        slc = vol[:, :, 77]\n        axes[row][col].imshow(slc, cmap='gray')\n        if row == 0:\n            axes[row][col].set_title(title, fontweight='bold', fontsize=12)\n        axes[row][col].set_ylabel(f'Patient {row+1}' if col == 0 else '')\n        axes[row][col].axis('off')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/mri_grid.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:00:55.270696Z","iopub.execute_input":"2026-03-30T12:00:55.271001Z","iopub.status.idle":"2026-03-30T12:00:59.039096Z","shell.execute_reply.started":"2026-03-30T12:00:55.270970Z","shell.execute_reply":"2026-03-30T12:00:59.038222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(14, 5))\nfig.suptitle('Train / Validation / Test Split', fontsize=14, fontweight='bold')\n\n# Split sizes\nsplits = ['Train\\n(875)', 'Validation\\n(188)', 'Test\\n(188)']\nsizes = [875, 188, 188]\ncolors = ['#2ECC71', '#3498DB', '#E74C3C']\n\naxes[0].pie(sizes, labels=splits, colors=colors, autopct='%1.1f%%',\n            startangle=90, textprops={'fontsize': 12})\naxes[0].set_title('Overall Split (1,251 patients)', fontweight='bold')\n\n# GBM vs LGG per split\nx = np.arange(3)\nwidth = 0.35\ngbm_counts = [int(875*585/1251), int(188*585/1251), int(188*585/1251)]\nlgg_counts = [875-gbm_counts[0], 188-gbm_counts[1], 188-gbm_counts[2]]\n\nbars1 = axes[1].bar(x - width/2, gbm_counts, width, label='GBM', color='#E74C3C', edgecolor='white')\nbars2 = axes[1].bar(x + width/2, lgg_counts, width, label='LGG', color='#3498DB', edgecolor='white')\naxes[1].set_xticks(x)\naxes[1].set_xticklabels(['Train', 'Validation', 'Test'])\naxes[1].set_title('GBM vs LGG per Split', fontweight='bold')\naxes[1].set_ylabel('Number of Patients')\naxes[1].legend()\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/data_split.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:00:59.040272Z","iopub.execute_input":"2026-03-30T12:00:59.040583Z","iopub.status.idle":"2026-03-30T12:00:59.506839Z","shell.execute_reply.started":"2026-03-30T12:00:59.040557Z","shell.execute_reply":"2026-03-30T12:00:59.506241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(2, 4, figsize=(16, 8))\nfig.suptitle('GBM vs LGG — All 4 Modalities Comparison', fontsize=14, fontweight='bold')\n\ngbm_sample = df_full[df_full['tumor_type']=='GBM']['patient_folder'].iloc[0]\nlgg_sample = df_full[df_full['tumor_type']=='LGG']['patient_folder'].iloc[0]\nmodalities = ['t1', 't1ce', 't2', 'flair']\n\nfor col, mod in enumerate(modalities):\n    for row, (pid, label) in enumerate([(gbm_sample, 'GBM'), (lgg_sample, 'LGG')]):\n        vol = nib.load(f\"{BRATS_ROOT}/{pid}/{pid}_{mod}.nii.gz\").get_fdata()\n        slc = vol[:, :, 77]\n        \n        axes[row][col].imshow(slc, cmap='gray')\n        axes[row][col].axis('off')\n        if row == 0:\n            axes[row][col].set_title(mod.upper(), fontweight='bold', fontsize=13)\n        if col == 0:\n            color = '#E74C3C' if label == 'GBM' else '#2ECC71'\n            axes[row][col].set_ylabel(label, fontsize=13, fontweight='bold', color=color)\n            axes[row][col].yaxis.label.set_visible(True)\n            axes[row][col].set_yticks([])\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/eda_gbm_vs_lgg.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:00:59.507677Z","iopub.execute_input":"2026-03-30T12:00:59.508094Z","iopub.status.idle":"2026-03-30T12:01:03.838740Z","shell.execute_reply.started":"2026-03-30T12:00:59.508067Z","shell.execute_reply":"2026-03-30T12:01:03.838058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(2, 4, figsize=(16, 8))\nfig.suptitle('Intensity Distribution — Before vs After Z-score Normalization', fontsize=13, fontweight='bold')\n\nsample_pid = df_full['patient_folder'].iloc[0]\nmodalities = ['t1', 't1ce', 't2', 'flair']\ncolors = ['#0891B2', '#7C3AED', '#059669', '#D97706']\n\nfor col, (mod, color) in enumerate(zip(modalities, colors)):\n    vol = nib.load(f\"{BRATS_ROOT}/{sample_pid}/{sample_pid}_{mod}.nii.gz\").get_fdata()\n    slc = vol[:, :, 77].astype(np.float32)\n    mask = slc > 0\n\n    # Raw\n    axes[0][col].hist(slc[mask].flatten(), bins=80, color=color, alpha=0.8, edgecolor='none')\n    axes[0][col].set_title(f'{mod.upper()} — Raw', fontweight='bold')\n    axes[0][col].set_xlabel('Intensity')\n    if col == 0: axes[0][col].set_ylabel('Frequency')\n\n    # Normalized\n    slc_norm = slc.copy()\n    slc_norm[mask] = (slc[mask] - slc[mask].mean()) / (slc[mask].std() + 1e-8)\n    axes[1][col].hist(slc_norm[mask].flatten(), bins=80, color=color, alpha=0.8, edgecolor='none')\n    axes[1][col].set_title(f'{mod.upper()} — Z-score Normalized', fontweight='bold')\n    axes[1][col].set_xlabel('Intensity')\n    if col == 0: axes[1][col].set_ylabel('Frequency')\n    axes[1][col].axvline(x=0, color='red', linestyle='--', linewidth=1.5, label='mean=0')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/eda_intensity_dist.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:01:03.839946Z","iopub.execute_input":"2026-03-30T12:01:03.840303Z","iopub.status.idle":"2026-03-30T12:01:07.143087Z","shell.execute_reply.started":"2026-03-30T12:01:03.840275Z","shell.execute_reply":"2026-03-30T12:01:07.142513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(12, 5))\nfig.suptitle('MGMT Label Analysis', fontsize=14, fontweight='bold')\n\nlabeled = df_full[df_full['mgmt'].notna()].copy()\n\ngbm_mgmt = labeled[labeled['tumor_type']=='GBM']['mgmt'].value_counts()\nlgg_mgmt = labeled[labeled['tumor_type']=='LGG']['mgmt'].value_counts()\n\nx = np.arange(2)\nwidth = 0.35\naxes[0].bar(x - width/2, [gbm_mgmt.get(0.0,0), gbm_mgmt.get(1.0,0)], width,\n            label='GBM', color='#E74C3C', edgecolor='white')\naxes[0].bar(x + width/2, [lgg_mgmt.get(0.0,0), lgg_mgmt.get(1.0,0)], width,\n            label='LGG', color='#2ECC71', edgecolor='white')\naxes[0].set_xticks(x)\naxes[0].set_xticklabels(['Unmethylated', 'Methylated'])\naxes[0].set_ylabel('Patients')\naxes[0].set_title('MGMT Distribution by Tumor Type', fontweight='bold')\naxes[0].legend()\n\n# MGMT balance\naxes[1].pie([labeled['mgmt'].value_counts().get(0.0,0),\n             labeled['mgmt'].value_counts().get(1.0,0)],\n            labels=['Unmethylated\\n(276)', 'Methylated\\n(301)'],\n            colors=['#E67E22', '#9B59B6'],\n            autopct='%1.1f%%', startangle=90,\n            textprops={'fontsize':12})\naxes[1].set_title('Overall MGMT Balance (577 labeled)', fontweight='bold')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/eda_mgmt_analysis.png', dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:01:07.144048Z","iopub.execute_input":"2026-03-30T12:01:07.144788Z","iopub.status.idle":"2026-03-30T12:01:07.583433Z","shell.execute_reply.started":"2026-03-30T12:01:07.144763Z","shell.execute_reply":"2026-03-30T12:01:07.582519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── SECTION 7: MODEL ARCHITECTURE ────────────────────────────────────────────\nimport torchvision.models as models\n\nclass RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        # Load pretrained EfficientNet-B2\n        self.backbone = models.efficientnet_b2(weights=models.EfficientNet_B2_Weights.IMAGENET1K_V1)\n\n        # Modify first conv: 3 channels → 4 channels\n        orig_conv = self.backbone.features[0][0]\n        self.backbone.features[0][0] = nn.Conv2d(\n            4, orig_conv.out_channels,\n            kernel_size=orig_conv.kernel_size,\n            stride=orig_conv.stride,\n            padding=orig_conv.padding,\n            bias=False\n        )\n\n        # Initialize 4th channel weights as mean of original 3\n        with torch.no_grad():\n            self.backbone.features[0][0].weight[:, :3] = orig_conv.weight\n            self.backbone.features[0][0].weight[:, 3]  = orig_conv.weight.mean(dim=1)\n\n        # Remove original classifier\n        in_features = self.backbone.classifier[1].in_features\n        self.backbone.classifier = nn.Identity()\n\n        # Fusion layer\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        # 3 task heads\n        self.idh_head   = nn.Sequential(nn.Linear(512, 256), nn.ReLU(), nn.Linear(256, 1))\n        self.mgmt_head  = nn.Sequential(nn.Linear(512, 256), nn.ReLU(), nn.Linear(256, 1))\n        self.grade_head = nn.Sequential(nn.Linear(512, 256), nn.ReLU(), nn.Linear(256, 2))\n\n    def forward(self, x):\n        features = self.backbone(x)\n        fused    = self.fusion(features)\n        return self.idh_head(fused), self.mgmt_head(fused), self.grade_head(fused)\n\n\n# ── TEST FORWARD PASS ─────────────────────────────────────────────────────────\nmodel = RadiogenomicsModel().to(device)\ndummy = torch.randn(2, 4, 224, 224).to(device)\nidh_out, mgmt_out, grade_out = model(dummy)\n\nprint(\"IDH output shape:  \", idh_out.shape)\nprint(\"MGMT output shape: \", mgmt_out.shape)\nprint(\"Grade output shape:\", grade_out.shape)\nprint(f\"\\nTotal parameters:     {sum(p.numel() for p in model.parameters()):,}\")\nprint(f\"Trainable parameters: {sum(p.numel() for p in model.parameters() if p.requires_grad):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:23:32.272329Z","iopub.execute_input":"2026-04-18T21:23:32.273026Z","iopub.status.idle":"2026-04-18T21:23:34.130663Z","shell.execute_reply.started":"2026-04-18T21:23:32.272999Z","shell.execute_reply":"2026-04-18T21:23:34.129642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# LOSS FUNCTION (IMPROVED)\n# ========================================\n\nbce = nn.BCEWithLogitsLoss()\nce  = nn.CrossEntropyLoss()\n\ndef compute_loss(outputs, idh, mgmt, grade):\n    \n    idh_pred, mgmt_pred, grade_pred = outputs\n    \n    # IDH loss\n    loss_idh = bce(idh_pred.squeeze(), idh)\n    \n    # MGMT loss (only valid samples)\n    mask = mgmt != -1\n    \n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = 0.0\n    \n    # Grade loss\n    loss_grade = ce(grade_pred, grade)\n    \n    # 🔥 Weighted total loss (IMPORTANT CHANGE)\n    total_loss = (\n        1.0 * loss_idh + \n        0.5 * loss_mgmt + \n        1.0 * loss_grade\n    )\n    \n    return total_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:23:41.474550Z","iopub.execute_input":"2026-04-18T21:23:41.475314Z","iopub.status.idle":"2026-04-18T21:23:41.480694Z","shell.execute_reply.started":"2026-04-18T21:23:41.475283Z","shell.execute_reply":"2026-04-18T21:23:41.479981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# OPTIMIZER + SCHEDULER\n# ========================================\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=2e-4,   # was 3e-4\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=30\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:23:45.212573Z","iopub.execute_input":"2026-04-18T21:23:45.213212Z","iopub.status.idle":"2026-04-18T21:23:45.218678Z","shell.execute_reply.started":"2026-04-18T21:23:45.213185Z","shell.execute_reply":"2026-04-18T21:23:45.217826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# TRAIN + VALIDATE FUNCTIONS (FINAL)\n# ========================================\n\ndef train_epoch(loader):\n    \n    model.train()\n    total_loss = 0\n    \n    for imgs, labels in loader:\n        \n        imgs = imgs.to(device)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        \n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        \n        # 🔥 Gradient clipping (important)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        \n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n    return total_loss / len(loader)\n\n\nfrom sklearn.metrics import accuracy_score\n\ndef validate_epoch(loader):\n    \n    model.eval()\n    \n    total_loss = 0\n    \n    idh_preds, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            total_loss += loss.item()\n            \n            # ======================\n            # IDH (Binary)\n            # ======================\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_preds.extend((idh_prob > 0.5).astype(int))\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # ======================\n            # MGMT (Binary, partial labels)\n            # ======================\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:  # ignore missing\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # ======================\n            # Grade (Multiclass)\n            # ======================\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    # ======================\n    # Metrics\n    # ======================\n    \n    idh_auc = roc_auc_score(idh_targets, idh_preds)\n    idh_acc = accuracy_score(idh_targets, idh_preds)\n    \n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    \n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    return (\n        total_loss / len(loader),\n        idh_auc,\n        idh_acc,\n        mgmt_acc,\n        grade_acc\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:23:57.466612Z","iopub.execute_input":"2026-04-18T21:23:57.467174Z","iopub.status.idle":"2026-04-18T21:23:57.478016Z","shell.execute_reply.started":"2026-04-18T21:23:57.467143Z","shell.execute_reply":"2026-04-18T21:23:57.477298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# subset = torch.utils.data.Subset(train_dataset, range(32))\n# debug_loader = DataLoader(subset, batch_size=8)\n\n# loss = train_epoch(debug_loader)\n\n# print(\"Mini training loss:\", loss)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:01:09.909918Z","iopub.execute_input":"2026-03-30T12:01:09.910372Z","iopub.status.idle":"2026-03-30T12:01:09.926438Z","shell.execute_reply.started":"2026-03-30T12:01:09.910334Z","shell.execute_reply":"2026-03-30T12:01:09.925725Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nepochs = 3\nbest_auc = 0.0\n\nfor epoch in range(epochs):\n    \n    train_loss = train_epoch(train_loader)\n    \n    val_loss, idh_auc, idh_acc, mgmt_acc, grade_acc = validate_epoch(val_loader)\n    \n    scheduler.step()\n    \n    print(f\"\\nEpoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # Save best model\n    if idh_auc > best_auc:\n        best_auc = idh_auc\n        torch.save(model.state_dict(), \"best_model.pth\")\n        print(\"✅ Best model saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:08:31.678996Z","iopub.execute_input":"2026-03-30T12:08:31.679867Z","iopub.status.idle":"2026-03-30T12:24:02.982632Z","shell.execute_reply.started":"2026-03-30T12:08:31.679825Z","shell.execute_reply":"2026-03-30T12:24:02.981620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# CONTINUE TRAINING (FIXED VERSION)\n# ========================================\n\nfrom tqdm import tqdm\nfrom sklearn.metrics import accuracy_score, roc_auc_score\nimport os\nimport torch\n\n# ✅ Safe device\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel.to(device)\n\n# ✅ Load only if exists\nif os.path.exists(\"best_model.pth\"):\n    model.load_state_dict(torch.load(\"best_model.pth\", map_location=device))\n    print(\"Loaded previous best model\")\n\n# ✅ IMPORTANT: update to your current best\nbest_auc = 0.7111   # <-- from your latest result\n\nepochs = 10\nstart_epoch = 3\npatience = 3\ncounter = 0\n\nfor epoch in range(start_epoch, epochs):\n    \n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n    total_loss = 0\n    \n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n    \n    for imgs, labels in progress_bar:\n        \n        imgs = imgs.to(device, non_blocking=True)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n        # 🔥 better live stats\n        progress_bar.set_postfix({\n            \"batch_loss\": f\"{loss.item():.3f}\",\n            \"avg_loss\": f\"{(total_loss/(progress_bar.n+1)):.3f}\"\n        })\n    \n    train_loss = total_loss / len(train_loader)\n    \n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n    \n    val_loss = 0\n    idh_probs, idh_preds, idh_targets = [], [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            val_loss += loss.item()\n            \n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_preds.extend((idh_prob > 0.5).astype(int))\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    val_loss /= len(val_loader)\n    \n    # ======================\n    # METRICS\n    # ======================\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, idh_preds)\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    scheduler.step()\n    \n    # ======================\n    # PRINT\n    # ======================\n    print(f\"\\nEpoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # ======================\n    # SAVE BEST MODEL\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n        torch.save(model.state_dict(), \"best_model.pth\")\n        print(\"✅ Best model saved!\")\n        counter = 0\n    else:\n        counter += 1\n    \n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:26:30.766648Z","iopub.execute_input":"2026-03-30T12:26:30.767385Z","iopub.status.idle":"2026-03-30T12:51:47.811899Z","shell.execute_reply.started":"2026-03-30T12:26:30.767347Z","shell.execute_reply":"2026-03-30T12:51:47.809989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"final_model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:53:33.228366Z","iopub.execute_input":"2026-03-30T12:53:33.229352Z","iopub.status.idle":"2026-03-30T12:53:33.324647Z","shell.execute_reply.started":"2026-03-30T12:53:33.229307Z","shell.execute_reply":"2026-03-30T12:53:33.323996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"best_model.pth\"))\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T12:54:09.868553Z","iopub.execute_input":"2026-03-30T12:54:09.869155Z","iopub.status.idle":"2026-03-30T12:54:10.017625Z","shell.execute_reply.started":"2026-03-30T12:54:09.869127Z","shell.execute_reply":"2026-03-30T12:54:10.017005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/input/datasets/mondaldebasish05/radiogenomicsmodal-1/Model1.pth\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T15:41:34.309175Z","iopub.execute_input":"2026-03-30T15:41:34.309738Z","iopub.status.idle":"2026-03-30T15:41:34.316300Z","shell.execute_reply.started":"2026-03-30T15:41:34.309688Z","shell.execute_reply":"2026-03-30T15:41:34.315017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nMODEL_PATH = \"/kaggle/input/datasets/mondaldebasish05/radiogenomicsmodal-1/Model1.pth\"\n\nmodel = RadiogenomicsModel()\nmodel.load_state_dict(torch.load(MODEL_PATH, map_location=device))\n\nmodel.to(device)\nmodel.eval()\n\nprint(\"✅ Model loaded successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T15:41:36.145518Z","iopub.execute_input":"2026-03-30T15:41:36.145922Z","iopub.status.idle":"2026-03-30T15:41:36.563867Z","shell.execute_reply.started":"2026-03-30T15:41:36.145890Z","shell.execute_reply":"2026-03-30T15:41:36.562145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# FINAL TEST EVALUATION\n# ========================================\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    roc_auc_score,\n    confusion_matrix,\n    classification_report\n)\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n# Load best model\nmodel.load_state_dict(torch.load(MODEL_PATH, map_location=device))\n\nmodel.to(device)\nmodel.eval()\n\n# Storage\nidh_probs, idh_preds, idh_targets = [], [], []\nmgmt_preds, mgmt_targets = [], []\ngrade_preds, grade_targets = [], []\n\n# ======================\n# TEST LOOP\n# ======================\nwith torch.no_grad():\n    for imgs, labels in test_loader:\n        \n        imgs = imgs.to(device)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        outputs = model(imgs)\n        \n        # ------------------\n        # IDH\n        # ------------------\n        idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n        idh_probs.extend(idh_prob)\n        idh_preds.extend((idh_prob > 0.5).astype(int))\n        idh_targets.extend(idh.cpu().numpy().flatten())\n        \n        # ------------------\n        # MGMT\n        # ------------------\n        mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n        mgmt_true = mgmt.cpu().numpy().flatten()\n        \n        for p, t in zip(mgmt_prob, mgmt_true):\n            if t != -1:\n                mgmt_preds.append(int(p > 0.5))\n                mgmt_targets.append(int(t))\n        \n        # ------------------\n        # Grade\n        # ------------------\n        grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n        grade_preds.extend(grade_pred)\n        grade_targets.extend(grade.cpu().numpy())\n\n# ======================\n# METRICS\n# ======================\n\nidh_auc = roc_auc_score(idh_targets, idh_probs)\nidh_acc = accuracy_score(idh_targets, idh_preds)\n\nmgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\ngrade_acc = accuracy_score(grade_targets, grade_preds)\n\nprint(\"\\n===== FINAL TEST RESULTS =====\")\nprint(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\nprint(f\"MGMT → Acc: {mgmt_acc:.4f}\")\nprint(f\"Grade→ Acc: {grade_acc:.4f}\")\n\n# ======================\n# CONFUSION MATRIX (IDH)\n# ======================\n\ncm = confusion_matrix(idh_targets, idh_preds)\n\nplt.figure(figsize=(5,4))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\")\nplt.title(\"IDH Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.show()\n\n# ======================\n# CLASSIFICATION REPORT\n# ======================\n\nprint(\"\\nIDH Classification Report:\")\nprint(classification_report(idh_targets, idh_preds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T15:42:32.001941Z","iopub.execute_input":"2026-03-30T15:42:32.002579Z","iopub.status.idle":"2026-03-30T15:43:36.576261Z","shell.execute_reply.started":"2026-03-30T15:42:32.002535Z","shell.execute_reply":"2026-03-30T15:43:36.574590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.metrics import f1_score\n\nbest_thr = 0\nbest_f1 = 0\n\nfor t in np.arange(0.1, 0.9, 0.01):\n    preds = (np.array(idh_probs) > t).astype(int)\n    f1 = f1_score(idh_targets, preds)\n    \n    if f1 > best_f1:\n        best_f1 = f1\n        best_thr = t\n\nprint(\"Best threshold:\", best_thr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T15:44:11.667593Z","iopub.execute_input":"2026-03-30T15:44:11.668123Z","iopub.status.idle":"2026-03-30T15:44:11.878504Z","shell.execute_reply.started":"2026-03-30T15:44:11.668073Z","shell.execute_reply":"2026-03-30T15:44:11.877478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.metrics import f1_score, accuracy_score\n\nidh_probs_np = np.array(idh_probs)\nidh_targets_np = np.array(idh_targets)\n\nbest_thr = 0\nbest_f1 = 0\nbest_acc = 0\n\nfor t in np.arange(0.1, 0.9, 0.01):\n    preds = (idh_probs_np > t).astype(int)\n    \n    f1 = f1_score(idh_targets_np, preds)\n    acc = accuracy_score(idh_targets_np, preds)\n    \n    if f1 > best_f1:\n        best_f1 = f1\n        best_acc = acc\n        best_thr = t\n\nprint(f\"\\n🔥 Best Threshold: {best_thr:.2f}\")\nprint(f\"🔥 Best F1 Score: {best_f1:.4f}\")\nprint(f\"🔥 Accuracy at Best Threshold: {best_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T15:48:34.817486Z","iopub.execute_input":"2026-03-30T15:48:34.818186Z","iopub.status.idle":"2026-03-30T15:48:35.063899Z","shell.execute_reply.started":"2026-03-30T15:48:34.818144Z","shell.execute_reply":"2026-03-30T15:48:35.062298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, f1_score\n\nfor t in [0.12, 0.3, 0.4, 0.5]:\n    preds = (idh_probs_np > t).astype(int)\n    \n    print(f\"\\nThreshold: {t}\")\n    print(\"Accuracy:\", accuracy_score(idh_targets_np, preds))\n    print(\"F1 Score:\", f1_score(idh_targets_np, preds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T15:50:23.499380Z","iopub.execute_input":"2026-03-30T15:50:23.500515Z","iopub.status.idle":"2026-03-30T15:50:23.526090Z","shell.execute_reply.started":"2026-03-30T15:50:23.500462Z","shell.execute_reply":"2026-03-30T15:50:23.524937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Architecture\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = models.efficientnet_b2(\n            weights=models.EfficientNet_B2_Weights.IMAGENET1K_V1\n        )\n\n        # Modify first conv (3 → 4 channels)\n        orig_conv = self.backbone.features[0][0]\n        self.backbone.features[0][0] = nn.Conv2d(\n            4, orig_conv.out_channels,\n            kernel_size=orig_conv.kernel_size,\n            stride=orig_conv.stride,\n            padding=orig_conv.padding,\n            bias=False\n        )\n\n        with torch.no_grad():\n            self.backbone.features[0][0].weight[:, :3] = orig_conv.weight\n            self.backbone.features[0][0].weight[:, 3] = orig_conv.weight.mean(dim=1)\n\n        # Remove classifier\n        in_features = self.backbone.classifier[1].in_features\n        self.backbone.classifier = nn.Identity()\n\n        # 🔥 Freeze early layers (IMPORTANT)\n        for param in model.backbone.features[:2].parameters():\n            param.requires_grad = False \n\n        # Fusion\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        # 🔥 Improved heads (with dropout)\n        self.idh_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, 1)\n        )\n\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, 1)\n        )\n\n        self.grade_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, 2)\n        )\n\n    def forward(self, x):\n        features = self.backbone(x)\n        fused = self.fusion(features)\n        return (\n            self.idh_head(fused),\n            self.mgmt_head(fused),\n            self.grade_head(fused)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:13:25.946851Z","iopub.execute_input":"2026-04-19T21:13:25.947747Z","iopub.status.idle":"2026-04-19T21:13:25.960397Z","shell.execute_reply.started":"2026-04-19T21:13:25.947707Z","shell.execute_reply":"2026-04-19T21:13:25.959121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Loss Function\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n    \n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        loss = ((1 - pt) ** self.gamma) * bce_loss\n        return loss.mean()\nbce = FocalLoss(gamma=2)\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_pred, mgmt_pred, grade_pred = outputs\n\n    # IDH\n    loss_idh = bce(idh_pred.squeeze(), idh)\n\n    # MGMT (masked)\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    # Grade\n    loss_grade = ce(grade_pred, grade)\n\n    # 🔥 Better balance\n    total_loss = (\n        1.2 * loss_idh +\n        0.7 * loss_mgmt +\n        0.8 * loss_grade\n    )\n\n    return total_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:13:28.598997Z","iopub.execute_input":"2026-04-19T21:13:28.599878Z","iopub.status.idle":"2026-04-19T21:13:28.607995Z","shell.execute_reply.started":"2026-04-19T21:13:28.599846Z","shell.execute_reply":"2026-04-19T21:13:28.606781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#optimizer\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=8e-5,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=20\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:13:30.840473Z","iopub.execute_input":"2026-04-19T21:13:30.841341Z","iopub.status.idle":"2026-04-19T21:13:31.477914Z","shell.execute_reply.started":"2026-04-19T21:13:30.841308Z","shell.execute_reply":"2026-04-19T21:13:31.476688Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Training Loop\nfrom tqdm import tqdm\nimport torch\n\ndef train_epoch(loader):\n    \n    model.train()\n    total_loss = 0\n    \n    for imgs, labels in loader:\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        \n        optimizer.step()\n        \n        total_loss += loss.item()\n    \n    return total_loss / len(loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:13:47.025421Z","iopub.execute_input":"2026-04-19T21:13:47.025809Z","iopub.status.idle":"2026-04-19T21:13:47.032723Z","shell.execute_reply.started":"2026-04-19T21:13:47.025779Z","shell.execute_reply":"2026-04-19T21:13:47.031819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Validation\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndef validate_epoch(loader):\n    \n    model.eval()\n    \n    total_loss = 0\n    \n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            total_loss += loss.item()\n            \n            # IDH (correct AUC)\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    return (\n        total_loss / len(loader),\n        idh_auc,\n        idh_acc,\n        mgmt_acc,\n        grade_acc\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:13:51.487188Z","iopub.execute_input":"2026-04-19T21:13:51.487492Z","iopub.status.idle":"2026-04-19T21:13:51.499142Z","shell.execute_reply.started":"2026-04-19T21:13:51.487468Z","shell.execute_reply":"2026-04-19T21:13:51.498085Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **v2:** Baseline training with standard learning rate, moderate dropout, and partially frozen backbone for stable feature learning.","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nimport torch\n\nepochs = 12\nbest_auc = 0.0\npatience = 4\ncounter = 0\n\nfor epoch in range(epochs):\n    \n    # ======================\n    # TRAIN (WITH LIVE PROGRESS)\n    # ======================\n    model.train()\n    total_loss = 0\n    \n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n    \n    for i, (imgs, labels) in enumerate(progress_bar):\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n        # 🔥 LIVE UPDATE (VERY IMPORTANT)\n        progress_bar.set_postfix({\n            \"batch_loss\": f\"{loss.item():.4f}\",\n            \"avg_loss\": f\"{(total_loss/(i+1)):.4f}\"\n        })\n    \n    train_loss = total_loss / len(train_loader)\n    \n    # ======================\n    # VALIDATION\n    # ======================\n    val_loss, idh_auc, idh_acc, mgmt_acc, grade_acc = validate_epoch(val_loader)\n    \n    scheduler.step()\n    \n    # ======================\n    # PRINT EPOCH SUMMARY\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1} Summary\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # ======================\n    # SAVE BEST MODEL\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n        torch.save(model.state_dict(), \"best_model_v2.pth\")\n        print(\"✅ Saved best_model_v2.pth\")\n        counter = 0\n    else:\n        counter += 1\n    \n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:13:09.099041Z","iopub.execute_input":"2026-04-19T21:13:09.099761Z","iopub.status.idle":"2026-04-19T21:13:09.113326Z","shell.execute_reply.started":"2026-04-19T21:13:09.099728Z","shell.execute_reply":"2026-04-19T21:13:09.112060Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=8e-5,   # 🔥 lower LR\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=12\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T21:29:37.418678Z","iopub.execute_input":"2026-03-30T21:29:37.419194Z","iopub.status.idle":"2026-03-30T21:29:37.426016Z","shell.execute_reply.started":"2026-03-30T21:29:37.419160Z","shell.execute_reply":"2026-03-30T21:29:37.425115Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **v4:** Aggressive training with reduced regularization and higher flexibility, leading to minimal constraints and increased risk of overfitting.","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import roc_auc_score, accuracy_score\n\nepochs = 10\nbest_auc = 0.0\npatience = 3\ncounter = 0\n\nfor epoch in range(epochs):\n    \n    # ======================\n    # TRAIN (LIVE METRICS)\n    # ======================\n    model.train()\n    total_loss = 0\n    \n    running_probs = []\n    running_targets = []\n    \n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n    \n    for i, (imgs, labels) in enumerate(progress_bar):\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n        # ======================\n        # LIVE METRICS (IDH)\n        # ======================\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n        \n        running_probs.extend(probs)\n        running_targets.extend(targets)\n        \n        # Compute live AUC safely\n        if len(set(running_targets)) > 1:\n            live_auc = roc_auc_score(running_targets, running_probs)\n            live_acc = accuracy_score(\n                running_targets,\n                (np.array(running_probs) > 0.5).astype(int)\n            )\n        else:\n            live_auc = 0.0\n            live_acc = 0.0\n        \n        # ======================\n        # UPDATE PROGRESS BAR\n        # ======================\n        progress_bar.set_postfix({\n            \"loss\": f\"{loss.item():.3f}\",\n            \"avg_loss\": f\"{(total_loss/(i+1)):.3f}\",\n            \"AUC\": f\"{live_auc:.3f}\",\n            \"ACC\": f\"{live_acc:.3f}\"\n        })\n    \n    train_loss = total_loss / len(train_loader)\n    \n    # ======================\n    # VALIDATION (same as before)\n    # ======================\n    model.eval()\n    \n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            val_loss += loss.item()\n            \n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    val_loss /= len(val_loader)\n    \n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    scheduler.step()\n    \n    # ======================\n    # FINAL PRINT\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # ======================\n    # SAVE BEST\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n        \n        torch.save({\n            'model': model.state_dict(),\n            'optimizer': optimizer.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_model_v4.pth\")\n        \n        print(\"✅ Saved best_model_v4.pth\")\n        counter = 0\n    else:\n        counter += 1\n    \n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T21:30:40.877703Z","iopub.execute_input":"2026-03-30T21:30:40.878477Z","iopub.status.idle":"2026-03-30T22:09:57.708885Z","shell.execute_reply.started":"2026-03-30T21:30:40.878446Z","shell.execute_reply":"2026-03-30T22:09:57.708026Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **v5:** Optimized training with improved augmentation, balanced loss weighting, and better regularization for controlled and stable learning.","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = models.efficientnet_b2(\n            weights=models.EfficientNet_B2_Weights.IMAGENET1K_V1\n        )\n\n        # Modify input (4 channels)\n        orig_conv = self.backbone.features[0][0]\n        self.backbone.features[0][0] = nn.Conv2d(\n            4, orig_conv.out_channels,\n            kernel_size=orig_conv.kernel_size,\n            stride=orig_conv.stride,\n            padding=orig_conv.padding,\n            bias=False\n        )\n\n        with torch.no_grad():\n            self.backbone.features[0][0].weight[:, :3] = orig_conv.weight\n            self.backbone.features[0][0].weight[:, 3] = orig_conv.weight.mean(dim=1)\n\n        # Remove classifier\n        in_features = self.backbone.classifier[1].in_features\n        self.backbone.classifier = nn.Identity()\n\n        # 🔥 PARTIAL FREEZE (KEY CHANGE)\n        for param in self.backbone.features[:3].parameters():\n            param.requires_grad = False\n\n        # 🔥 Stronger regularization\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        self.idh_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.grade_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 2)\n        )\n\n    def forward(self, x):\n        features = self.backbone(x)\n        fused = self.fusion(features)\n        return (\n            self.idh_head(fused),\n            self.mgmt_head(fused),\n            self.grade_head(fused)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T23:06:14.365468Z","iopub.execute_input":"2026-04-18T23:06:14.366239Z","iopub.status.idle":"2026-04-18T23:06:14.375097Z","shell.execute_reply.started":"2026-04-18T23:06:14.366210Z","shell.execute_reply":"2026-04-18T23:06:14.374147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        loss = ((1 - pt) ** self.gamma) * bce_loss\n        return loss.mean()\n\nbce = FocalLoss(gamma=2)\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_pred, mgmt_pred, grade_pred = outputs\n\n    loss_idh = bce(idh_pred.squeeze(), idh)\n\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    loss_grade = ce(grade_pred, grade)\n\n    # 🔥 Balanced focus\n    total_loss = (\n        1.2 * loss_idh +\n        0.7 * loss_mgmt +\n        0.8 * loss_grade\n    )\n\n    return total_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T23:06:17.503884Z","iopub.execute_input":"2026-04-18T23:06:17.504291Z","iopub.status.idle":"2026-04-18T23:06:17.511412Z","shell.execute_reply.started":"2026-04-18T23:06:17.504265Z","shell.execute_reply":"2026-04-18T23:06:17.510566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=8e-5,\n    weight_decay=3e-4   # 🔥 increased\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T23:06:21.779620Z","iopub.execute_input":"2026-04-18T23:06:21.780420Z","iopub.status.idle":"2026-04-18T23:06:21.785406Z","shell.execute_reply.started":"2026-04-18T23:06:21.780391Z","shell.execute_reply":"2026-04-18T23:06:21.784783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=8e-5,\n    weight_decay=3e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T23:06:26.014240Z","iopub.execute_input":"2026-04-18T23:06:26.014817Z","iopub.status.idle":"2026-04-18T23:06:26.237394Z","shell.execute_reply.started":"2026-04-18T23:06:26.014788Z","shell.execute_reply":"2026-04-18T23:06:26.236812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# TRAINING (v5 FINAL — LIVE + STABLE)\n# ========================================\n\nimport torch\nfrom tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=8e-5,\n    weight_decay=3e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)\n\nepochs = 10\nbest_auc = 0.0\npatience = 3\ncounter = 0\n\nfor epoch in range(epochs):\n    \n    # ======================\n    # TRAIN (LIVE METRICS)\n    # ======================\n    model.train()\n    total_loss = 0\n    \n    running_probs = []\n    running_targets = []\n    \n    progress_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n    \n    for i, (imgs, labels) in enumerate(progress_bar):\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n        # ======================\n        # LIVE TRAIN METRICS (IDH)\n        # ======================\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n        \n        running_probs.extend(probs)\n        running_targets.extend(targets)\n        \n        if len(set(running_targets)) > 1:\n            live_auc = roc_auc_score(running_targets, running_probs)\n            live_acc = accuracy_score(\n                running_targets,\n                (np.array(running_probs) > 0.5).astype(int)\n            )\n        else:\n            live_auc = 0.0\n            live_acc = 0.0\n        \n        progress_bar.set_postfix({\n            \"loss\": f\"{loss.item():.3f}\",\n            \"avg_loss\": f\"{(total_loss/(i+1)):.3f}\",\n            \"AUC\": f\"{live_auc:.3f}\",\n            \"ACC\": f\"{live_acc:.3f}\"\n        })\n    \n    train_loss = total_loss / len(train_loader)\n    \n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n    \n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            val_loss += loss.item()\n            \n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    val_loss /= len(val_loader)\n    \n    # ======================\n    # METRICS\n    # ======================\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    scheduler.step()\n    \n    # ======================\n    # PRINT\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # ======================\n    # SAVE BEST MODEL\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n        \n        torch.save({\n            'model': model.state_dict(),\n            'optimizer': optimizer.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_model_v5.pth\")\n        \n        print(\"✅ Saved best_model_v5.pth\")\n        counter = 0\n    else:\n        counter += 1\n    \n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T23:06:34.305758Z","iopub.execute_input":"2026-04-18T23:06:34.306518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp best_model_v5.pth /kaggle/working/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:01:26.388865Z","iopub.execute_input":"2026-03-30T23:01:26.389596Z","iopub.status.idle":"2026-03-30T23:01:26.549819Z","shell.execute_reply.started":"2026-03-30T23:01:26.389560Z","shell.execute_reply":"2026-03-30T23:01:26.548895Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MODEL COMPARISON","metadata":{}},{"cell_type":"code","source":"# ========================================\n# MODEL COMPARISON (FROM KAGGLE OUTPUT)\n# ========================================\n\nimport torch\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# ======================\n# LOAD MODEL FUNCTION\n# ======================\n\ndef load_model(path):\n    model = RadiogenomicsModel().to(device)\n    \n    checkpoint = torch.load(\n        path,\n        map_location=device,\n        weights_only=False   # 🔥 FIX\n    )\n    \n    if 'model' in checkpoint:\n        model.load_state_dict(checkpoint['model'])\n    else:\n        model.load_state_dict(checkpoint)\n    \n    model.eval()\n    return model\n\n\n# ======================\n# LOAD MODELS (FROM OUTPUT)\n# ======================\n\nmodel_v2 = load_model(\"/kaggle/working/best_model_v2.pth\")\nmodel_v4 = load_model(\"/kaggle/working/best_model_v4.pth\")\nmodel_v5 = load_model(\"/kaggle/working/best_model_v5.pth\")\n\n\n# ======================\n# EVALUATION FUNCTION\n# ======================\n\ndef evaluate(model, loader):\n    \n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            \n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    return idh_auc, idh_acc, mgmt_acc, grade_acc\n\n\n# ======================\n# RUN COMPARISON\n# ======================\n\nresults = {}\n\nresults[\"v2\"] = evaluate(model_v2, test_loader)\nresults[\"v4\"] = evaluate(model_v4, test_loader)\nresults[\"v5\"] = evaluate(model_v5, test_loader)\n\n\n# ======================\n# PRINT RESULTS\n# ======================\n\nprint(\"\\n🔥 MODEL COMPARISON RESULTS\\n\")\n\nfor name, (auc, acc, mgmt, grade) in results.items():\n    print(f\"{name.upper()}\")\n    print(f\"  IDH AUC : {auc:.4f}\")\n    print(f\"  IDH Acc : {acc:.4f}\")\n    print(f\"  MGMT Acc: {mgmt:.4f}\")\n    print(f\"  Grade Acc: {grade:.4f}\")\n    print(\"-\" * 30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:14:10.348957Z","iopub.execute_input":"2026-03-30T23:14:10.349316Z","iopub.status.idle":"2026-03-30T23:16:56.112493Z","shell.execute_reply.started":"2026-03-30T23:14:10.349290Z","shell.execute_reply":"2026-03-30T23:16:56.111160Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\nx = np.arange(len(models))\nwidth = 0.2\n\nplt.figure()\n\nplt.bar(x - width, auc, width, label=\"AUC\")\nplt.bar(x, idh_acc, width, label=\"IDH Acc\")\nplt.bar(x + width, mgmt_acc, width, label=\"MGMT Acc\")\n\nplt.xticks(x, models)\nplt.title(\"Model Comparison (All Metrics)\")\nplt.xlabel(\"Model\")\nplt.ylabel(\"Score\")\nplt.legend()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:18:27.586216Z","iopub.execute_input":"2026-03-30T23:18:27.586808Z","iopub.status.idle":"2026-03-30T23:18:27.722181Z","shell.execute_reply.started":"2026-03-30T23:18:27.586776Z","shell.execute_reply":"2026-03-30T23:18:27.721568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nmodels = [\"V2\", \"V4\", \"V5\"]\n\nauc = [results[\"v2\"][0], results[\"v4\"][0], results[\"v5\"][0]]\nidh_acc = [results[\"v2\"][1], results[\"v4\"][1], results[\"v5\"][1]]\nmgmt_acc = [results[\"v2\"][2], results[\"v4\"][2], results[\"v5\"][2]]\ngrade_acc = [results[\"v2\"][3], results[\"v4\"][3], results[\"v5\"][3]]\n\ndef plot_metric(values, title):\n    plt.figure(figsize=(8,5))\n    \n    bars = plt.bar(models, values)\n    \n    # 🔥 ZOOM Y-AXIS (VERY IMPORTANT)\n    plt.ylim(min(values) - 0.05, max(values) + 0.05)\n    \n    # 🔥 VALUE LABELS ON TOP\n    for bar in bars:\n        y = bar.get_height()\n        plt.text(bar.get_x() + bar.get_width()/2, y,\n                 f\"{y:.3f}\", ha='center', va='bottom', fontsize=10)\n    \n    plt.title(title, fontsize=14)\n    plt.xlabel(\"Model\", fontsize=12)\n    plt.ylabel(\"Score\", fontsize=12)\n    \n    # 🔥 GRID FOR CLARITY\n    plt.grid(axis='y', linestyle='--', alpha=0.7)\n    \n    plt.tight_layout()\n    plt.show()\n\n\n# Plot all\nplot_metric(auc, \"IDH AUC Comparison\")\nplot_metric(idh_acc, \"IDH Accuracy Comparison\")\nplot_metric(mgmt_acc, \"MGMT Accuracy Comparison\")\nplot_metric(grade_acc, \"Grade Accuracy Comparison\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:20:56.497901Z","iopub.execute_input":"2026-03-30T23:20:56.498565Z","iopub.status.idle":"2026-03-30T23:20:56.993234Z","shell.execute_reply.started":"2026-03-30T23:20:56.498535Z","shell.execute_reply":"2026-03-30T23:20:56.992549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x = np.arange(len(models))\nwidth = 0.2\n\nplt.figure(figsize=(10,6))\n\nb1 = plt.bar(x - width, auc, width, label=\"AUC\")\nb2 = plt.bar(x, idh_acc, width, label=\"IDH Acc\")\nb3 = plt.bar(x + width, mgmt_acc, width, label=\"MGMT Acc\")\n\nplt.xticks(x, models)\n\nplt.ylim(0.65, 0.85)  # 🔥 zoom range\nplt.grid(axis='y', linestyle='--', alpha=0.7)\n\nplt.title(\"Model Comparison (All Metrics)\")\nplt.xlabel(\"Model\")\nplt.ylabel(\"Score\")\nplt.legend()\n\n# Value labels\nfor bars in [b1, b2, b3]:\n    for bar in bars:\n        y = bar.get_height()\n        plt.text(bar.get_x()+bar.get_width()/2, y,\n                 f\"{y:.3f}\", ha='center', va='bottom', fontsize=9)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:24:29.896483Z","iopub.execute_input":"2026-03-30T23:24:29.896758Z","iopub.status.idle":"2026-03-30T23:24:30.069111Z","shell.execute_reply.started":"2026-03-30T23:24:29.896737Z","shell.execute_reply":"2026-03-30T23:24:30.068374Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nbars = plt.bar(models, mgmt_acc)\n\nplt.ylim(min(mgmt_acc) - 0.05, max(mgmt_acc) + 0.05)\n\nfor bar in bars:\n    y = bar.get_height()\n    plt.text(bar.get_x()+bar.get_width()/2, y,\n             f\"{y:.3f}\", ha='center', va='bottom')\n\nplt.title(\"MGMT Accuracy Comparison\")\nplt.xlabel(\"Model\")\nplt.ylabel(\"Accuracy\")\nplt.grid(axis='y', linestyle='--', alpha=0.7)\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:26:04.730368Z","iopub.execute_input":"2026-03-30T23:26:04.730647Z","iopub.status.idle":"2026-03-30T23:26:04.855031Z","shell.execute_reply.started":"2026-03-30T23:26:04.730624Z","shell.execute_reply":"2026-03-30T23:26:04.854349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_curve, auc\n\ndef get_roc_data(model, loader):\n    \n    model.eval()\n    probs = []\n    targets = []\n    \n    with torch.no_grad():\n        for imgs, labels in loader:\n            \n            imgs = imgs.to(device)\n            idh = labels['idh'].to(device)\n            \n            outputs = model(imgs)\n            \n            prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            probs.extend(prob)\n            targets.extend(idh.cpu().numpy().flatten())\n    \n    fpr, tpr, _ = roc_curve(targets, probs)\n    roc_auc = auc(fpr, tpr)\n    \n    return fpr, tpr, roc_auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:29:07.336175Z","iopub.execute_input":"2026-03-30T23:29:07.336720Z","iopub.status.idle":"2026-03-30T23:29:07.342689Z","shell.execute_reply.started":"2026-03-30T23:29:07.336692Z","shell.execute_reply":"2026-03-30T23:29:07.341827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fpr_v2, tpr_v2, auc_v2 = get_roc_data(model_v2, test_loader)\nfpr_v4, tpr_v4, auc_v4 = get_roc_data(model_v4, test_loader)\nfpr_v5, tpr_v5, auc_v5 = get_roc_data(model_v5, test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:29:18.278383Z","iopub.execute_input":"2026-03-30T23:29:18.279220Z","iopub.status.idle":"2026-03-30T23:31:58.710162Z","shell.execute_reply.started":"2026-03-30T23:29:18.279191Z","shell.execute_reply":"2026-03-30T23:31:58.708984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(8,6))\n\n# Plot curves\nplt.plot(fpr_v2, tpr_v2, label=f\"V2 (AUC = {auc_v2:.3f})\")\nplt.plot(fpr_v4, tpr_v4, label=f\"V4 (AUC = {auc_v4:.3f})\")\nplt.plot(fpr_v5, tpr_v5, label=f\"V5 (AUC = {auc_v5:.3f})\")\n\n# Diagonal line (random model)\nplt.plot([0,1], [0,1], linestyle='--')\n\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.title(\"ROC Curve Comparison (IDH Prediction)\")\nplt.legend()\nplt.grid(alpha=0.3)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T23:32:13.526004Z","iopub.execute_input":"2026-03-30T23:32:13.526488Z","iopub.status.idle":"2026-03-30T23:32:13.720016Z","shell.execute_reply.started":"2026-03-30T23:32:13.526452Z","shell.execute_reply":"2026-03-30T23:32:13.718962Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ResNet50 Architecture Final\n","metadata":{}},{"cell_type":"code","source":"# ========================================\n# RESNET50 MULTI-TASK MODEL\n# ========================================\n\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = models.resnet50(\n            weights=models.ResNet50_Weights.IMAGENET1K_V1\n        )\n\n        # Modify input (4 channels)\n        orig_conv = self.backbone.conv1\n        self.backbone.conv1 = nn.Conv2d(\n            4, 64,\n            kernel_size=7,\n            stride=2,\n            padding=3,\n            bias=False\n        )\n\n        with torch.no_grad():\n            self.backbone.conv1.weight[:, :3] = orig_conv.weight\n            self.backbone.conv1.weight[:, 3] = orig_conv.weight.mean(dim=1)\n\n        # Remove classifier\n        in_features = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n\n\n        # Fusion\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        # Heads\n        self.idh_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.grade_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 2)\n        )\n\n    def forward(self, x):\n        features = self.backbone(x)\n        fused = self.fusion(features)\n\n        return (\n            self.idh_head(fused),\n            self.mgmt_head(fused),\n            self.grade_head(fused)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:44:44.938671Z","iopub.execute_input":"2026-04-18T21:44:44.939287Z","iopub.status.idle":"2026-04-18T21:44:44.947376Z","shell.execute_reply.started":"2026-04-18T21:44:44.939258Z","shell.execute_reply":"2026-04-18T21:44:44.946646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# LOSS (FOCAL + CE)\n# ========================================\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        loss = ((1 - pt) ** self.gamma) * bce_loss\n        return loss.mean()\n\nbce = FocalLoss(gamma=2)\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_pred, mgmt_pred, grade_pred = outputs\n\n    loss_idh = bce(idh_pred.squeeze(), idh)\n\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    loss_grade = ce(grade_pred, grade)\n\n    return 1.2*loss_idh + 0.7*loss_mgmt + 0.8*loss_grade","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:44:57.121982Z","iopub.execute_input":"2026-04-18T21:44:57.122864Z","iopub.status.idle":"2026-04-18T21:44:57.129054Z","shell.execute_reply.started":"2026-04-18T21:44:57.122832Z","shell.execute_reply":"2026-04-18T21:44:57.128411Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# OPTIMIZER\n# ========================================\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr = 5e-5,\n    weight_decay=3e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:45:03.333947Z","iopub.execute_input":"2026-04-18T21:45:03.334410Z","iopub.status.idle":"2026-04-18T21:45:04.432396Z","shell.execute_reply.started":"2026-04-18T21:45:03.334367Z","shell.execute_reply":"2026-04-18T21:45:04.431548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ========================================\n# TRAINING LOOP (FULL METRICS VERSION)\n# ========================================\n\nfrom tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndevice = \"cuda\"\nmodel.to(device)\nimgs = imgs.to(device)\n\n\nepochs = 10\nbest_auc = 0.0\npatience = 3\ncounter = 0\n\nfor epoch in range(epochs):\n    \n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n    total_loss = 0\n    \n    running_probs, running_targets = [], []\n    \n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n    \n    for i, (imgs, labels) in enumerate(pbar):\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n        # Live metrics (IDH)\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n        \n        running_probs.extend(probs)\n        running_targets.extend(targets)\n        \n        if len(set(running_targets)) > 1:\n            live_auc = roc_auc_score(running_targets, running_probs)\n        else:\n            live_auc = 0.0\n        \n        pbar.set_postfix({\n            \"loss\": f\"{loss.item():.3f}\",\n            \"AUC\": f\"{live_auc:.3f}\"\n        })\n    \n    train_loss = total_loss / len(train_loader)\n    \n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n    \n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            val_loss += loss.item()\n            \n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    val_loss /= len(val_loader)\n    \n    # ======================\n    # METRICS\n    # ======================\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    scheduler.step()\n    \n    # ======================\n    # PRINT (CLEAN OUTPUT)\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # ======================\n    # SAVE BEST\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n        \n        torch.save({\n            'model': model.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_resnet50.pth\")\n        \n        print(\"✅ Saved best_resnet50.pth\")\n        counter = 0\n    else:\n        counter += 1\n    \n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T21:46:23.819703Z","iopub.execute_input":"2026-04-02T21:46:23.820019Z","iopub.status.idle":"2026-04-02T22:12:28.145949Z","shell.execute_reply.started":"2026-04-02T21:46:23.819990Z","shell.execute_reply":"2026-04-02T22:12:28.145079Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training after UNfreezing layer and lower learning rate with increassing dropout (TRAINING LOOP FINAL)","metadata":{}},{"cell_type":"code","source":"# ========================================\n# TRAINING LOOP (FULL METRICS VERSION)\n# ========================================\n\nfrom tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndevice = \"cuda\"\nmodel.to(device)\nimgs = imgs.to(device)\n\n\nepochs = 10\nbest_auc = 0.0\npatience = 2\ncounter = 0\n\nfor epoch in range(epochs):\n    \n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n    total_loss = 0\n    \n    running_probs, running_targets = [], []\n    \n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n    \n    for i, (imgs, labels) in enumerate(pbar):\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        optimizer.zero_grad()\n        \n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n        \n        # Live metrics (IDH)\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n        \n        running_probs.extend(probs)\n        running_targets.extend(targets)\n        \n        if len(set(running_targets)) > 1:\n            live_auc = roc_auc_score(running_targets, running_probs)\n        else:\n            live_auc = 0.0\n        \n        pbar.set_postfix({\n            \"loss\": f\"{loss.item():.3f}\",\n            \"AUC\": f\"{live_auc:.3f}\"\n        })\n    \n    train_loss = total_loss / len(train_loader)\n    \n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n    \n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            \n            imgs = imgs.to(device)\n            \n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n            \n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n            \n            val_loss += loss.item()\n            \n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n            \n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n            \n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n            \n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n    \n    val_loss /= len(val_loader)\n    \n    # ======================\n    # METRICS\n    # ======================\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n    \n    scheduler.step()\n    \n    # ======================\n    # PRINT (CLEAN OUTPUT)\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    \n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n    \n    # ======================\n    # SAVE BEST\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n        \n        torch.save({\n            'model': model.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_resnet50_v2.pth\")\n        \n        print(\"✅ Saved best_resnet50_v2.pth\")\n        counter = 0\n    else:\n        counter += 1\n    \n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T22:21:39.844122Z","iopub.execute_input":"2026-04-02T22:21:39.844601Z","iopub.status.idle":"2026-04-02T22:44:49.330768Z","shell.execute_reply.started":"2026-04-02T22:21:39.844566Z","shell.execute_reply":"2026-04-02T22:44:49.329927Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Threshold Tuning","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom sklearn.metrics import f1_score, accuracy_score, roc_auc_score\n\n# ======================\n# GET PREDICTIONS\n# ======================\n\nmodel.eval()\n\nidh_probs = []\nidh_targets = []\n\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        \n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        \n        outputs = model(imgs)\n        \n        probs = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n        \n        idh_probs.extend(probs)\n        idh_targets.extend(idh.cpu().numpy().flatten())\n\nidh_probs = np.array(idh_probs)\nidh_targets = np.array(idh_targets)\n\nprint(\"Base AUC:\", roc_auc_score(idh_targets, idh_probs))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T22:51:01.772642Z","iopub.execute_input":"2026-04-02T22:51:01.773429Z","iopub.status.idle":"2026-04-02T22:52:15.113408Z","shell.execute_reply.started":"2026-04-02T22:51:01.773385Z","shell.execute_reply":"2026-04-02T22:52:15.112579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"thresholds = np.linspace(0.1, 0.9, 50)\n\nbest_thresh = 0.5\nbest_f1 = 0\nbest_acc = 0\n\nfor t in thresholds:\n    \n    preds = (idh_probs > t).astype(int)\n    \n    f1 = f1_score(idh_targets, preds)\n    acc = accuracy_score(idh_targets, preds)\n    \n    if f1 > best_f1:\n        best_f1 = f1\n        best_thresh = t\n        best_acc = acc\n\nprint(f\"\\n🔥 Best Threshold: {best_thresh:.3f}\")\nprint(f\"F1 Score: {best_f1:.4f}\")\nprint(f\"Accuracy: {best_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T22:52:55.495580Z","iopub.execute_input":"2026-04-02T22:52:55.496211Z","iopub.status.idle":"2026-04-02T22:52:55.626755Z","shell.execute_reply.started":"2026-04-02T22:52:55.496175Z","shell.execute_reply":"2026-04-02T22:52:55.626149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = RadiogenomicsModel().to(device)\n\ncheckpoint = torch.load(\n    \"best_resnet50_v2.pth\",\n    map_location=device,\n    weights_only=False   # 🔥 IMPORTANT FIX\n)\nmodel.load_state_dict(checkpoint['model'])\n\nmodel.eval()\n\nprint(\"✅ Model loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:01:54.588520Z","iopub.execute_input":"2026-04-02T23:01:54.589239Z","iopub.status.idle":"2026-04-02T23:01:55.092728Z","shell.execute_reply.started":"2026-04-02T23:01:54.589208Z","shell.execute_reply":"2026-04-02T23:01:55.092015Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, roc_auc_score, f1_score, confusion_matrix\n\nTHRESH = 0.41\n\nidh_probs, idh_targets = [], []\nmgmt_preds, mgmt_targets = [], []\ngrade_preds, grade_targets = [], []\n\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        \n        imgs = imgs.to(device)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        outputs = model(imgs)\n        \n        # ===== IDH =====\n        idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n        idh_probs.extend(idh_prob)\n        idh_targets.extend(idh.cpu().numpy().flatten())\n        \n        # ===== MGMT =====\n        mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n        mgmt_true = mgmt.cpu().numpy().flatten()\n        \n        for p, t in zip(mgmt_prob, mgmt_true):\n            if t != -1:\n                mgmt_preds.append(int(p > THRESH))\n                mgmt_targets.append(int(t))\n        \n        # ===== GRADE =====\n        grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n        grade_preds.extend(grade_pred)\n        grade_targets.extend(grade.cpu().numpy())\n\n# ======================\n# METRICS\n# ======================\n\nidh_auc = roc_auc_score(idh_targets, idh_probs)\nidh_preds = (np.array(idh_probs) > THRESH).astype(int)\nidh_acc = accuracy_score(idh_targets, idh_preds)\nidh_f1 = f1_score(idh_targets, idh_preds)\n\nmgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\ngrade_acc = accuracy_score(grade_targets, grade_preds)\n\nprint(\"\\n📊 VALIDATION RESULTS\")\nprint(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}, F1: {idh_f1:.4f}\")\nprint(f\"MGMT → Acc: {mgmt_acc:.4f}\")\nprint(f\"Grade→ Acc: {grade_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:02:28.006567Z","iopub.execute_input":"2026-04-02T23:02:28.007318Z","iopub.status.idle":"2026-04-02T23:03:59.676993Z","shell.execute_reply.started":"2026-04-02T23:02:28.007290Z","shell.execute_reply":"2026-04-02T23:03:59.676235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idh_probs, idh_targets = [], []\nmgmt_preds, mgmt_targets = [], []\ngrade_preds, grade_targets = [], []\n\nwith torch.no_grad():\n    for imgs, labels in test_loader:\n        \n        imgs = imgs.to(device)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        outputs = model(imgs)\n        \n        # IDH\n        idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n        idh_probs.extend(idh_prob)\n        idh_targets.extend(idh.cpu().numpy().flatten())\n        \n        # MGMT\n        mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n        mgmt_true = mgmt.cpu().numpy().flatten()\n        \n        for p, t in zip(mgmt_prob, mgmt_true):\n            if t != -1:\n                mgmt_preds.append(int(p > THRESH))\n                mgmt_targets.append(int(t))\n        \n        # Grade\n        grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n        grade_preds.extend(grade_pred)\n        grade_targets.extend(grade.cpu().numpy())\n\n# METRICS\nidh_auc = roc_auc_score(idh_targets, idh_probs)\nidh_preds = (np.array(idh_probs) > THRESH).astype(int)\nidh_acc = accuracy_score(idh_targets, idh_preds)\nidh_f1 = f1_score(idh_targets, idh_preds)\n\nmgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\ngrade_acc = accuracy_score(grade_targets, grade_preds)\n\nprint(\"\\n🧪 TEST RESULTS\")\nprint(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}, F1: {idh_f1:.4f}\")\nprint(f\"MGMT → Acc: {mgmt_acc:.4f}\")\nprint(f\"Grade→ Acc: {grade_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:04:37.926615Z","iopub.execute_input":"2026-04-02T23:04:37.926892Z","iopub.status.idle":"2026-04-02T23:06:11.526385Z","shell.execute_reply.started":"2026-04-02T23:04:37.926869Z","shell.execute_reply":"2026-04-02T23:06:11.525700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(idh_targets, idh_preds)\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:06:28.732996Z","iopub.execute_input":"2026-04-02T23:06:28.733275Z","iopub.status.idle":"2026-04-02T23:06:28.741152Z","shell.execute_reply.started":"2026-04-02T23:06:28.733249Z","shell.execute_reply":"2026-04-02T23:06:28.740566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, roc_auc_score, f1_score, confusion_matrix\n\nTHRESH = 0.56\n\nidh_probs, idh_targets = [], []\nmgmt_preds, mgmt_targets = [], []\ngrade_preds, grade_targets = [], []\n\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        \n        imgs = imgs.to(device)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        outputs = model(imgs)\n        \n        # ===== IDH =====\n        idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n        idh_probs.extend(idh_prob)\n        idh_targets.extend(idh.cpu().numpy().flatten())\n        \n        # ===== MGMT =====\n        mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n        mgmt_true = mgmt.cpu().numpy().flatten()\n        \n        for p, t in zip(mgmt_prob, mgmt_true):\n            if t != -1:\n                mgmt_preds.append(int(p > THRESH))\n                mgmt_targets.append(int(t))\n        \n        # ===== GRADE =====\n        grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n        grade_preds.extend(grade_pred)\n        grade_targets.extend(grade.cpu().numpy())\n\n# ======================\n# METRICS\n# ======================\n\nidh_auc = roc_auc_score(idh_targets, idh_probs)\nidh_preds = (np.array(idh_probs) > THRESH).astype(int)\nidh_acc = accuracy_score(idh_targets, idh_preds)\nidh_f1 = f1_score(idh_targets, idh_preds)\n\nmgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\ngrade_acc = accuracy_score(grade_targets, grade_preds)\n\nprint(\"\\n📊 VALIDATION RESULTS\")\nprint(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}, F1: {idh_f1:.4f}\")\nprint(f\"MGMT → Acc: {mgmt_acc:.4f}\")\nprint(f\"Grade→ Acc: {grade_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:23:41.331138Z","iopub.execute_input":"2026-04-02T23:23:41.331581Z","iopub.status.idle":"2026-04-02T23:25:21.226062Z","shell.execute_reply.started":"2026-04-02T23:23:41.331550Z","shell.execute_reply":"2026-04-02T23:25:21.225349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(idh_targets, idh_preds)\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:25:28.491547Z","iopub.execute_input":"2026-04-02T23:25:28.491841Z","iopub.status.idle":"2026-04-02T23:25:28.499008Z","shell.execute_reply.started":"2026-04-02T23:25:28.491819Z","shell.execute_reply":"2026-04-02T23:25:28.498351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import balanced_accuracy_score\n\nbest_thresh = 0.5\nbest_score = 0\n\nfor t in np.linspace(0.3, 0.7, 50):\n    preds = (idh_probs > t).astype(int)\n    score = balanced_accuracy_score(idh_targets, preds)\n    \n    if score > best_score:\n        best_score = score\n        best_thresh = t\n\nprint(\"Best Balanced Threshold:\", best_thresh)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:23:23.595218Z","iopub.execute_input":"2026-04-02T23:23:23.595929Z","iopub.status.idle":"2026-04-02T23:23:23.679838Z","shell.execute_reply.started":"2026-04-02T23:23:23.595897Z","shell.execute_reply":"2026-04-02T23:23:23.679078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Threshold Tuning (FINAL = 0.56)","metadata":{}},{"cell_type":"code","source":"idh_probs, idh_targets = [], []\nmgmt_preds, mgmt_targets = [], []\ngrade_preds, grade_targets = [], []\n\nwith torch.no_grad():\n    for imgs, labels in test_loader:\n        \n        imgs = imgs.to(device)\n        \n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n        \n        outputs = model(imgs)\n        \n        # IDH\n        idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n        idh_probs.extend(idh_prob)\n        idh_targets.extend(idh.cpu().numpy().flatten())\n        \n        # MGMT\n        mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n        mgmt_true = mgmt.cpu().numpy().flatten()\n        \n        for p, t in zip(mgmt_prob, mgmt_true):\n            if t != -1:\n                mgmt_preds.append(int(p > THRESH))\n                mgmt_targets.append(int(t))\n        \n        # Grade\n        grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n        grade_preds.extend(grade_pred)\n        grade_targets.extend(grade.cpu().numpy())\n\n# METRICS\nidh_auc = roc_auc_score(idh_targets, idh_probs)\nidh_preds = (np.array(idh_probs) > THRESH).astype(int)\nidh_acc = accuracy_score(idh_targets, idh_preds)\nidh_f1 = f1_score(idh_targets, idh_preds)\n\nmgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\ngrade_acc = accuracy_score(grade_targets, grade_preds)\n\nprint(\"\\n🧪 TEST RESULTS\")\nprint(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}, F1: {idh_f1:.4f}\")\nprint(f\"MGMT → Acc: {mgmt_acc:.4f}\")\nprint(f\"Grade→ Acc: {grade_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:26:04.523309Z","iopub.execute_input":"2026-04-02T23:26:04.524158Z","iopub.status.idle":"2026-04-02T23:27:48.399664Z","shell.execute_reply.started":"2026-04-02T23:26:04.524125Z","shell.execute_reply":"2026-04-02T23:27:48.398830Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Two threshold strategies were evaluated:\n\n- **0.41 (F1-optimized):**\n  Higher sensitivity, better detection of tumor cases\n\n- **0.56 (Balanced accuracy):**\n  Improved class balance but increased false negatives\n\nGiven the medical context, the model with threshold 0.41 was preferred to minimize missed diagnoses.","metadata":{}},{"cell_type":"markdown","source":"# COMPARISION GRAPHS\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\n# 🔥 FIX\nidh_probs = np.array(idh_probs)\nidh_targets = np.array(idh_targets)\n\nt1 = 0.41\nt2 = 0.56\n\ndef get_metrics(thresh):\n    preds = (idh_probs > thresh).astype(int)\n    \n    cm = confusion_matrix(idh_targets, preds)\n    tn, fp, fn, tp = cm.ravel()\n    \n    acc = (tp + tn) / (tp + tn + fp + fn)\n    recall = tp / (tp + fn)\n    precision = tp / (tp + fp)\n    \n    return acc, recall, precision\n\nacc1, rec1, prec1 = get_metrics(t1)\nacc2, rec2, prec2 = get_metrics(t2)\n\nlabels = ['Accuracy', 'Recall (Tumor)', 'Precision']\nt1_vals = [acc1, rec1, prec1]\nt2_vals = [acc2, rec2, prec2]\n\nx = np.arange(len(labels))\nwidth = 0.35\n\nplt.figure()\nplt.bar(x - width/2, t1_vals, width, label='Threshold 0.41')\nplt.bar(x + width/2, t2_vals, width, label='Threshold 0.56')\n\nplt.xticks(x, labels)\nplt.ylabel(\"Score\")\nplt.title(\"Threshold Comparison\")\nplt.legend()\nplt.grid()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-02T23:37:46.037507Z","iopub.execute_input":"2026-04-02T23:37:46.038107Z","iopub.status.idle":"2026-04-02T23:37:46.205797Z","shell.execute_reply.started":"2026-04-02T23:37:46.038081Z","shell.execute_reply":"2026-04-02T23:37:46.205164Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**The comparison shows that lower thresholds improve recall (sensitivity), while higher thresholds improve precision. \nFor medical diagnosis, higher recall is preferred to avoid missing tumor cases, hence threshold = 0.41 was selected.**","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# ======================\n# DATA\n# ======================\nmodels = ['ResNet v1', 'ResNet v2']\n\nauc = [0.7847, 0.7919]\naccuracy = [0.71, 0.73]\nf1 = [0.76, 0.79]\n\nx = np.arange(len(models))\nwidth = 0.25\n\n# ======================\n# PLOT\n# ======================\nplt.figure(figsize=(8,5))\n\nbars1 = plt.bar(x - width, auc, width, label='AUC')\nbars2 = plt.bar(x, accuracy, width, label='Accuracy')\nbars3 = plt.bar(x + width, f1, width, label='F1 Score')\n\n# ======================\n# ADD VALUES ON TOP\n# ======================\ndef add_labels(bars):\n    for bar in bars:\n        height = bar.get_height()\n        plt.text(\n            bar.get_x() + bar.get_width()/2,\n            height + 0.002,\n            f'{height:.3f}',\n            ha='center',\n            va='bottom',\n            fontsize=9\n        )\n\nadd_labels(bars1)\nadd_labels(bars2)\nadd_labels(bars3)\n\n# ======================\n# STYLING\n# ======================\nplt.ylim(0.65, 0.85)\nplt.xticks(x, models)\nplt.ylabel(\"Score\")\nplt.title(\"Model Comparison: ResNet Basic vs ResNet v2\")\nplt.legend()\nplt.grid(axis='y', linestyle='--', alpha=0.6)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T00:10:29.701437Z","iopub.execute_input":"2026-04-03T00:10:29.701713Z","iopub.status.idle":"2026-04-03T00:10:29.947313Z","shell.execute_reply.started":"2026-04-03T00:10:29.701678Z","shell.execute_reply":"2026-04-03T00:10:29.946663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# ======================\n# DATA\n# ======================\nmodels = ['ResNet v1', 'ResNet v2']\n\nauc = [0.7847, 0.7919]\naccuracy = [0.71, 0.73]\nf1 = [0.76, 0.79]\n\ndef plot_metric(values, title, ylabel):\n    x = np.arange(len(models))\n    \n    plt.figure(figsize=(5,4))\n    bars = plt.bar(x, values)\n    \n    # Add labels\n    for bar in bars:\n        height = bar.get_height()\n        plt.text(\n            bar.get_x() + bar.get_width()/2,\n            height + 0.002,\n            f'{height:.3f}',\n            ha='center',\n            va='bottom'\n        )\n    \n    plt.xticks(x, models)\n    plt.ylim(0.65, 0.85)\n    plt.ylabel(ylabel)\n    plt.title(title)\n    plt.grid(axis='y', linestyle='--', alpha=0.6)\n    \n    plt.tight_layout()\n    plt.show()\n\n# ======================\n# PLOTS\n# ======================\n\nplot_metric(auc, \"AUC Comparison (ResNet v1 vs v2)\", \"AUC\")\nplot_metric(accuracy, \"Accuracy Comparison (ResNet v1 vs v2)\", \"Accuracy\")\nplot_metric(f1, \"F1 Score Comparison (ResNet v1 vs v2)\", \"F1 Score\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T00:11:32.199101Z","iopub.execute_input":"2026-04-03T00:11:32.199402Z","iopub.status.idle":"2026-04-03T00:11:32.540194Z","shell.execute_reply.started":"2026-04-03T00:11:32.199379Z","shell.execute_reply":"2026-04-03T00:11:32.539644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = RadiogenomicsModel().to(device)\n\nmodel_path = model_path = \"/kaggle/input/datasets/mondaldebasish05/resnet50-v2-optimized/best_resnet50_v2.pth\"\n\ncheckpoint = torch.load(\n    model_path,\n    map_location=device,\n    weights_only=False\n)\n\nmodel.load_state_dict(checkpoint['model'])\nmodel.eval()\n\nprint(\"✅ Model loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T20:12:35.080255Z","iopub.execute_input":"2026-04-03T20:12:35.080969Z","iopub.status.idle":"2026-04-03T20:12:35.649547Z","shell.execute_reply.started":"2026-04-03T20:12:35.080940Z","shell.execute_reply":"2026-04-03T20:12:35.648732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score, f1_score\n\nTHRESH = 0.41  # your best threshold\n\ndef evaluate(loader):\n    \n    idh_probs, idh_targets = [], []\n    \n    with torch.no_grad():\n        for imgs, labels in loader:\n            \n            imgs = imgs.to(device)\n            idh = labels['idh'].to(device)\n            \n            outputs = model(imgs)\n            \n            probs = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            \n            idh_probs.extend(probs)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n    \n    idh_probs = np.array(idh_probs)\n    idh_targets = np.array(idh_targets)\n    \n    preds = (idh_probs > THRESH).astype(int)\n    \n    auc = roc_auc_score(idh_targets, idh_probs)\n    acc = accuracy_score(idh_targets, preds)\n    f1 = f1_score(idh_targets, preds)\n    \n    return idh_targets, preds, auc, acc, f1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T20:13:09.569699Z","iopub.execute_input":"2026-04-03T20:13:09.570088Z","iopub.status.idle":"2026-04-03T20:13:09.576461Z","shell.execute_reply.started":"2026-04-03T20:13:09.570060Z","shell.execute_reply":"2026-04-03T20:13:09.575536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_targets, val_preds, val_auc, val_acc, val_f1 = evaluate(val_loader)\n\nprint(\"\\n📊 VALIDATION\")\nprint(f\"AUC: {val_auc:.4f}\")\nprint(f\"Acc: {val_acc:.4f}\")\nprint(f\"F1:  {val_f1:.4f}\")\n\ntest_targets, test_preds, test_auc, test_acc, test_f1 = evaluate(test_loader)\n\nprint(\"\\n🧪 TEST\")\nprint(f\"AUC: {test_auc:.4f}\")\nprint(f\"Acc: {test_acc:.4f}\")\nprint(f\"F1:  {test_f1:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T20:13:28.417452Z","iopub.execute_input":"2026-04-03T20:13:28.418288Z","iopub.status.idle":"2026-04-03T20:16:48.322651Z","shell.execute_reply.started":"2026-04-03T20:13:28.418252Z","shell.execute_reply":"2026-04-03T20:16:48.321947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.metrics import ConfusionMatrixDisplay\n\n# VALIDATION CM\nplt.figure(figsize=(5,4))\nConfusionMatrixDisplay.from_predictions(val_targets, val_preds)\nplt.title(\"Validation Confusion Matrix\")\nplt.show()\n\n# TEST CM\nplt.figure(figsize=(5,4))\nConfusionMatrixDisplay.from_predictions(test_targets, test_preds)\nplt.title(\"Test Confusion Matrix\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T20:31:37.218906Z","iopub.execute_input":"2026-04-03T20:31:37.219701Z","iopub.status.idle":"2026-04-03T20:31:37.738821Z","shell.execute_reply.started":"2026-04-03T20:31:37.219673Z","shell.execute_reply":"2026-04-03T20:31:37.737763Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ConvNeXt ARCHITECTURE","metadata":{}},{"cell_type":"code","source":"class RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = timm.create_model(\n            'convnext_base',\n            pretrained=True,\n            num_classes=0   # 🔥 IMPORTANT: gives pooled features\n        )\n\n        # Modify input (4 channels)\n        orig_conv = self.backbone.stem[0]\n\n        self.backbone.stem[0] = nn.Conv2d(\n            4,\n            orig_conv.out_channels,\n            kernel_size=orig_conv.kernel_size,\n            stride=orig_conv.stride,\n            padding=orig_conv.padding,\n            bias=False\n        )\n\n        with torch.no_grad():\n            self.backbone.stem[0].weight[:, :3] = orig_conv.weight\n            self.backbone.stem[0].weight[:, 3] = orig_conv.weight.mean(dim=1)\n\n        in_features = self.backbone.num_features\n\n        # Fusion\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        # Heads (same as before)\n        self.idh_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.grade_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 2)\n        )\n\n    def forward(self, x):\n        features = self.backbone(x)   # now shape = (B, C)\n        fused = self.fusion(features)\n\n        return (\n            self.idh_head(fused),\n            self.mgmt_head(fused),\n            self.grade_head(fused)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:54:54.247810Z","iopub.execute_input":"2026-04-18T21:54:54.248217Z","iopub.status.idle":"2026-04-18T21:54:54.256168Z","shell.execute_reply.started":"2026-04-18T21:54:54.248190Z","shell.execute_reply":"2026-04-18T21:54:54.255205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        loss = ((1 - pt) ** self.gamma) * bce_loss\n        return loss.mean()\n\nbce = FocalLoss(gamma=2)\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_pred, mgmt_pred, grade_pred = outputs\n\n    loss_idh = bce(idh_pred.squeeze(), idh)\n\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    loss_grade = ce(grade_pred, grade)\n\n    return 1.2*loss_idh + 0.7*loss_mgmt + 0.8*loss_grade","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:54:56.395886Z","iopub.execute_input":"2026-04-18T21:54:56.396677Z","iopub.status.idle":"2026-04-18T21:54:56.402716Z","shell.execute_reply.started":"2026-04-18T21:54:56.396649Z","shell.execute_reply":"2026-04-18T21:54:56.402123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport timm\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    \n    lr=3e-5,          # lower than ResNet (important)\n    weight_decay=3e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:57:22.508436Z","iopub.execute_input":"2026-04-18T21:57:22.509252Z","iopub.status.idle":"2026-04-18T21:57:24.056251Z","shell.execute_reply.started":"2026-04-18T21:57:22.509218Z","shell.execute_reply":"2026-04-18T21:57:24.055651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\nepochs = 10\nbest_auc = 0.0\npatience = 3\ncounter = 0\n\nfor epoch in range(epochs):\n\n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n    total_loss = 0\n\n    running_probs, running_targets = [], []\n\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n\n    for imgs, labels in pbar:\n\n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n\n        total_loss += loss.item()\n\n        # Live AUC (IDH)\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n\n        running_probs.extend(probs)\n        running_targets.extend(targets)\n\n        if len(set(running_targets)) > 1:\n            live_auc = roc_auc_score(running_targets, running_probs)\n        else:\n            live_auc = 0.0\n\n        pbar.set_postfix({\n            \"loss\": f\"{loss.item():.3f}\",\n            \"AUC\": f\"{live_auc:.3f}\"\n        })\n\n    train_loss = total_loss / len(train_loader)\n\n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n\n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n\n            imgs = imgs.to(device)\n\n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n\n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n\n            val_loss += loss.item()\n\n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n\n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n\n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n\n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n\n    val_loss /= len(val_loader)\n\n    # ======================\n    # METRICS\n    # ======================\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n\n    scheduler.step()\n\n    # ======================\n    # PRINT\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n\n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n\n    # ======================\n    # SAVE BEST\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n\n        torch.save({\n            'model': model.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_convnext.pth\")\n\n        print(\"✅ Saved best_convnext.pth\")\n        counter = 0\n    else:\n        counter += 1\n\n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T21:05:20.274296Z","iopub.execute_input":"2026-04-03T21:05:20.274944Z","iopub.status.idle":"2026-04-03T22:49:15.366156Z","shell.execute_reply.started":"2026-04-03T21:05:20.274904Z","shell.execute_reply":"2026-04-03T22:49:15.365489Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# RESNET101","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = models.resnet101(\n            weights=models.ResNet101_Weights.IMAGENET1K_V1\n        )\n\n        # Modify input (4 channels)\n        orig_conv = self.backbone.conv1\n        self.backbone.conv1 = nn.Conv2d(\n            4, 64,\n            kernel_size=7,\n            stride=2,\n            padding=3,\n            bias=False\n        )\n\n        with torch.no_grad():\n            self.backbone.conv1.weight[:, :3] = orig_conv.weight\n            self.backbone.conv1.weight[:, 3] = orig_conv.weight.mean(dim=1)\n\n        # Remove classifier\n        in_features = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n\n        # Fusion\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.5),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        # Heads\n        self.idh_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(256, 1)\n        )\n\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(256, 1)\n        )\n\n        self.grade_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(256, 2)\n        )\n\n    def forward(self, x):\n        features = self.backbone(x)\n        fused = self.fusion(features)\n\n        return (\n            self.idh_head(fused),\n            self.mgmt_head(fused),\n            self.grade_head(fused)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:00:47.054228Z","iopub.execute_input":"2026-04-18T22:00:47.054875Z","iopub.status.idle":"2026-04-18T22:00:47.063203Z","shell.execute_reply.started":"2026-04-18T22:00:47.054849Z","shell.execute_reply":"2026-04-18T22:00:47.062439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        loss = ((1 - pt) ** self.gamma) * bce_loss\n        return loss.mean()\n\nbce = FocalLoss(gamma=2)\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_pred, mgmt_pred, grade_pred = outputs\n\n    loss_idh = bce(idh_pred.squeeze(), idh)\n\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    loss_grade = ce(grade_pred, grade)\n\n    return 1.2*loss_idh + 0.7*loss_mgmt + 0.8*loss_grade","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:01:04.776608Z","iopub.execute_input":"2026-04-18T22:01:04.777203Z","iopub.status.idle":"2026-04-18T22:01:04.783431Z","shell.execute_reply.started":"2026-04-18T22:01:04.777176Z","shell.execute_reply":"2026-04-18T22:01:04.782717Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=3e-5,\n    weight_decay=3e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:01:07.240232Z","iopub.execute_input":"2026-04-18T22:01:07.240486Z","iopub.status.idle":"2026-04-18T22:01:07.955228Z","shell.execute_reply.started":"2026-04-18T22:01:07.240465Z","shell.execute_reply":"2026-04-18T22:01:07.954669Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\nepochs = 10\nbest_auc = 0.0\npatience = 3\ncounter = 0\n\nfor epoch in range(epochs):\n\n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n    total_loss = 0\n\n    running_probs, running_targets = [], []\n\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n\n    for imgs, labels in pbar:\n\n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n\n        total_loss += loss.item()\n\n        # Live AUC (IDH)\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n\n        running_probs.extend(probs)\n        running_targets.extend(targets)\n\n        if len(set(running_targets)) > 1:\n            live_auc = roc_auc_score(running_targets, running_probs)\n        else:\n            live_auc = 0.0\n\n        pbar.set_postfix({\n            \"loss\": f\"{loss.item():.3f}\",\n            \"AUC\": f\"{live_auc:.3f}\"\n        })\n\n    train_loss = total_loss / len(train_loader)\n\n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n\n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n\n            imgs = imgs.to(device)\n\n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n\n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n\n            val_loss += loss.item()\n\n            # IDH\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n\n            # MGMT\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n\n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n\n            # Grade\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n\n    val_loss /= len(val_loader)\n\n    # ======================\n    # METRICS\n    # ======================\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets) > 0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n\n    scheduler.step()\n\n    # ======================\n    # PRINT\n    # ======================\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n\n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n\n    # ======================\n    # SAVE BEST\n    # ======================\n    if idh_auc > best_auc:\n        print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n        best_auc = idh_auc\n\n        torch.save({\n            'model': model.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_resnet101.pth\")\n\n        print(\"✅ Saved best_resnet101.pth\")\n        counter = 0\n    else:\n        counter += 1\n\n    # ======================\n    # EARLY STOPPING\n    # ======================\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-03T22:55:50.799167Z","iopub.execute_input":"2026-04-03T22:55:50.799595Z","iopub.status.idle":"2026-04-03T23:55:27.598258Z","shell.execute_reply.started":"2026-04-03T22:55:50.799567Z","shell.execute_reply":"2026-04-03T23:55:27.597350Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DENSENET121","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = models.densenet121(\n            weights=models.DenseNet121_Weights.IMAGENET1K_V1\n        )\n\n        # Modify input (4 channels)\n        orig_conv = self.backbone.features.conv0\n        self.backbone.features.conv0 = nn.Conv2d(\n            4, 64,\n            kernel_size=7,\n            stride=2,\n            padding=3,\n            bias=False\n        )\n\n        with torch.no_grad():\n            self.backbone.features.conv0.weight[:, :3] = orig_conv.weight\n            self.backbone.features.conv0.weight[:, 3] = orig_conv.weight.mean(dim=1)\n\n        # Remove classifier\n        in_features = self.backbone.classifier.in_features\n        self.backbone.classifier = nn.Identity()\n\n        # Fusion\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(in_features, 512),\n            nn.ReLU()\n        )\n\n        # Heads\n        self.idh_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 1)\n        )\n\n        self.grade_head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, 2)\n        )\n\n    def forward(self, x):\n        features = self.backbone.features(x)\n        features = nn.functional.relu(features, inplace=True)\n        features = nn.functional.adaptive_avg_pool2d(features, (1,1))\n        features = torch.flatten(features, 1)\n\n        fused = self.fusion(features)\n\n        return (\n            self.idh_head(fused),\n            self.mgmt_head(fused),\n            self.grade_head(fused)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:01:52.925121Z","iopub.execute_input":"2026-04-18T22:01:52.925933Z","iopub.status.idle":"2026-04-18T22:01:52.934637Z","shell.execute_reply.started":"2026-04-18T22:01:52.925902Z","shell.execute_reply":"2026-04-18T22:01:52.933818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        loss = ((1 - pt) ** self.gamma) * bce_loss\n        return loss.mean()\n\nbce = FocalLoss(gamma=2)\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_pred, mgmt_pred, grade_pred = outputs\n\n    loss_idh = bce(idh_pred.squeeze(), idh)\n\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_pred.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    loss_grade = ce(grade_pred, grade)\n\n    return 1.2*loss_idh + 0.7*loss_mgmt + 0.8*loss_grade","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:01:59.197584Z","iopub.execute_input":"2026-04-18T22:01:59.198229Z","iopub.status.idle":"2026-04-18T22:01:59.204715Z","shell.execute_reply.started":"2026-04-18T22:01:59.198199Z","shell.execute_reply":"2026-04-18T22:01:59.203933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = RadiogenomicsModel().to(device)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=4e-5,          # slightly higher than ResNet101\n    weight_decay=3e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:02:02.303959Z","iopub.execute_input":"2026-04-18T22:02:02.304582Z","iopub.status.idle":"2026-04-18T22:02:02.785347Z","shell.execute_reply.started":"2026-04-18T22:02:02.304554Z","shell.execute_reply":"2026-04-18T22:02:02.784690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\nepochs = 10\nbest_auc = 0.0\npatience = 3\ncounter = 0\n\nfor epoch in range(epochs):\n\n    # TRAIN\n    model.train()\n    total_loss = 0\n    running_probs, running_targets = [], []\n\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Training\")\n\n    for imgs, labels in pbar:\n\n        imgs = imgs.to(device)\n        idh = labels['idh'].to(device)\n        mgmt = labels['mgmt'].to(device)\n        grade = labels['grade'].to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(imgs)\n        loss = compute_loss(outputs, idh, mgmt, grade)\n\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n\n        total_loss += loss.item()\n\n        probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n        targets = idh.cpu().numpy().flatten()\n\n        running_probs.extend(probs)\n        running_targets.extend(targets)\n\n        live_auc = roc_auc_score(running_targets, running_probs) if len(set(running_targets))>1 else 0.0\n\n        pbar.set_postfix({\"loss\": f\"{loss.item():.3f}\", \"AUC\": f\"{live_auc:.3f}\"})\n\n    train_loss = total_loss / len(train_loader)\n\n    # VALIDATION\n    model.eval()\n\n    val_loss = 0\n    idh_probs, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n\n            imgs = imgs.to(device)\n\n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n\n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n\n            val_loss += loss.item()\n\n            idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n            idh_probs.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy().flatten())\n\n            mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n            mgmt_true = mgmt.cpu().numpy().flatten()\n\n            for p, t in zip(mgmt_prob, mgmt_true):\n                if t != -1:\n                    mgmt_preds.append(int(p > 0.5))\n                    mgmt_targets.append(int(t))\n\n            grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n\n    val_loss /= len(val_loader)\n\n    idh_auc = roc_auc_score(idh_targets, idh_probs)\n    idh_acc = accuracy_score(idh_targets, (np.array(idh_probs)>0.5).astype(int))\n    mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets)>0 else 0\n    grade_acc = accuracy_score(grade_targets, grade_preds)\n\n    scheduler.step()\n\n    print(f\"\\n📊 Epoch {epoch+1}\")\n    print(f\"Train Loss: {train_loss:.4f}\")\n    print(f\"Val Loss:   {val_loss:.4f}\")\n    print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n    print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n    print(f\"Grade→ Acc: {grade_acc:.4f}\")\n\n    if idh_auc > best_auc:\n        best_auc = idh_auc\n        torch.save({\n            'model': model.state_dict(),\n            'epoch': epoch,\n            'auc': idh_auc\n        }, \"best_densenet121.pth\")\n        print(\"🔥 Saved best_densenet121.pth\")\n        counter = 0\n    else:\n        counter += 1\n\n    if counter >= patience:\n        print(\"⛔ Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T00:09:02.356994Z","iopub.execute_input":"2026-04-04T00:09:02.357728Z","iopub.status.idle":"2026-04-04T01:09:58.366826Z","shell.execute_reply.started":"2026-04-04T00:09:02.357697Z","shell.execute_reply":"2026-04-04T01:09:58.366114Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# FINAL RESNET50","metadata":{}},{"cell_type":"code","source":"import torch, random, numpy as np\nimport torch.nn as nn\nimport torchvision.models as models\nfrom tqdm import tqdm\nfrom sklearn.metrics import accuracy_score, roc_auc_score\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ndef set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:49:46.536854Z","iopub.execute_input":"2026-04-18T21:49:46.537297Z","iopub.status.idle":"2026-04-18T21:49:46.542425Z","shell.execute_reply.started":"2026-04-18T21:49:46.537269Z","shell.execute_reply":"2026-04-18T21:49:46.541569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RadiogenomicsModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.backbone = models.resnet50(\n            weights=models.ResNet50_Weights.IMAGENET1K_V1\n        )\n\n        orig = self.backbone.conv1\n        self.backbone.conv1 = nn.Conv2d(4,64,7,2,3,bias=False)\n\n        with torch.no_grad():\n            self.backbone.conv1.weight[:, :3] = orig.weight\n            self.backbone.conv1.weight[:, 3] = orig.weight.mean(dim=1)\n\n        in_features = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n\n        self.fusion = nn.Sequential(\n            nn.Dropout(0.45),\n            nn.Linear(in_features,512),\n            nn.ReLU()\n        )\n\n        self.idh_head = nn.Sequential(\n            nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.45), nn.Linear(256,1)\n        )\n        self.mgmt_head = nn.Sequential(\n            nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.45), nn.Linear(256,1)\n        )\n        self.grade_head = nn.Sequential(\n            nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.45), nn.Linear(256,2)\n        )\n\n    def forward(self,x):\n        f = self.backbone(x)\n        f = self.fusion(f)\n        return self.idh_head(f), self.mgmt_head(f), self.grade_head(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:49:48.732282Z","iopub.execute_input":"2026-04-18T21:49:48.733017Z","iopub.status.idle":"2026-04-18T21:49:48.740030Z","shell.execute_reply.started":"2026-04-18T21:49:48.732988Z","shell.execute_reply":"2026-04-18T21:49:48.739232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n        self.bce = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, logits, targets):\n        bce_loss = self.bce(logits, targets)\n        pt = torch.exp(-bce_loss)\n        return ((1 - pt) ** self.gamma * bce_loss).mean()\n\nbce = FocalLoss()\nce  = nn.CrossEntropyLoss(label_smoothing=0.1)\n\ndef compute_loss(outputs, idh, mgmt, grade):\n\n    idh_p, mgmt_p, grade_p = outputs\n\n    loss_idh = bce(idh_p.squeeze(), idh)\n\n    mask = mgmt != -1\n    if mask.sum() > 0:\n        loss_mgmt = bce(mgmt_p.squeeze()[mask], mgmt[mask])\n    else:\n        loss_mgmt = torch.tensor(0.0, device=idh.device)\n\n    loss_grade = ce(grade_p, grade)\n\n    return 1.2*loss_idh + 0.7*loss_mgmt + 0.8*loss_grade","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:49:52.539615Z","iopub.execute_input":"2026-04-18T21:49:52.540334Z","iopub.status.idle":"2026-04-18T21:49:52.546549Z","shell.execute_reply.started":"2026-04-18T21:49:52.540305Z","shell.execute_reply":"2026-04-18T21:49:52.545779Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(seed, save_name):\n\n    set_seed(seed)\n\n    model = RadiogenomicsModel().to(device)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5, weight_decay=4e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)\n\n    best_auc = 0\n    patience, counter = 3, 0\n\n    for epoch in range(10):\n\n        # ======================\n        # TRAIN\n        # ======================\n        model.train()\n        total_loss = 0\n\n        running_probs, running_targets = [], []\n\n        pbar = tqdm(train_loader, desc=f\"Seed {seed} Epoch {epoch+1}\")\n\n        for imgs, labels in pbar:\n\n            imgs = imgs.to(device)\n            idh = labels['idh'].to(device)\n            mgmt = labels['mgmt'].to(device)\n            grade = labels['grade'].to(device)\n\n            optimizer.zero_grad()\n\n            outputs = model(imgs)\n            loss = compute_loss(outputs, idh, mgmt, grade)\n\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n\n            total_loss += loss.item()\n\n            # Live AUC (IDH)\n            probs = torch.sigmoid(outputs[0]).detach().cpu().numpy().flatten()\n            targets = idh.cpu().numpy().flatten()\n\n            running_probs.extend(probs)\n            running_targets.extend(targets)\n\n            live_auc = roc_auc_score(running_targets, running_probs) if len(set(running_targets))>1 else 0\n\n            pbar.set_postfix({\n                \"loss\": f\"{loss.item():.3f}\",\n                \"AUC\": f\"{live_auc:.3f}\"\n            })\n\n        train_loss = total_loss / len(train_loader)\n\n        # ======================\n        # VALIDATION\n        # ======================\n        model.eval()\n\n        val_loss = 0\n        idh_probs, idh_targets = [], []\n        mgmt_preds, mgmt_targets = [], []\n        grade_preds, grade_targets = [], []\n\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n\n                imgs = imgs.to(device)\n\n                idh = labels['idh'].to(device)\n                mgmt = labels['mgmt'].to(device)\n                grade = labels['grade'].to(device)\n\n                outputs = model(imgs)\n                loss = compute_loss(outputs, idh, mgmt, grade)\n\n                val_loss += loss.item()\n\n                # IDH\n                idh_prob = torch.sigmoid(outputs[0]).cpu().numpy().flatten()\n                idh_probs.extend(idh_prob)\n                idh_targets.extend(idh.cpu().numpy().flatten())\n\n                # MGMT\n                mgmt_prob = torch.sigmoid(outputs[1]).cpu().numpy().flatten()\n                mgmt_true = mgmt.cpu().numpy().flatten()\n\n                for p, t in zip(mgmt_prob, mgmt_true):\n                    if t != -1:\n                        mgmt_preds.append(int(p > 0.5))\n                        mgmt_targets.append(int(t))\n\n                # Grade\n                grade_pred = torch.argmax(outputs[2], dim=1).cpu().numpy()\n                grade_preds.extend(grade_pred)\n                grade_targets.extend(grade.cpu().numpy())\n\n        val_loss /= len(val_loader)\n\n        # ======================\n        # METRICS\n        # ======================\n        idh_auc = roc_auc_score(idh_targets, idh_probs)\n        idh_acc = accuracy_score(idh_targets, (np.array(idh_probs) > 0.5).astype(int))\n        mgmt_acc = accuracy_score(mgmt_targets, mgmt_preds) if len(mgmt_targets)>0 else 0\n        grade_acc = accuracy_score(grade_targets, grade_preds)\n\n        scheduler.step()\n\n        # ======================\n        # PRINT (FULL OUTPUT)\n        # ======================\n        print(f\"\\n📊 Seed {seed} Epoch {epoch+1}\")\n        print(f\"Train Loss: {train_loss:.4f}\")\n        print(f\"Val Loss:   {val_loss:.4f}\")\n\n        print(f\"IDH  → AUC: {idh_auc:.4f}, Acc: {idh_acc:.4f}\")\n        print(f\"MGMT → Acc: {mgmt_acc:.4f}\")\n        print(f\"Grade→ Acc: {grade_acc:.4f}\")\n\n        # ======================\n        # SAVE BEST\n        # ======================\n        if idh_auc > best_auc:\n            print(f\"🔥 Improved: {best_auc:.4f} → {idh_auc:.4f}\")\n            best_auc = idh_auc\n\n            torch.save({\n                'model': model.state_dict(),\n                'auc': idh_auc\n            }, save_name)\n\n            print(f\"✅ Saved {save_name}\")\n            counter = 0\n        else:\n            counter += 1\n\n        # ======================\n        # EARLY STOPPING\n        # ======================\n        if counter >= patience:\n            print(\"⛔ Early stopping\")\n            break\n\n    print(f\"\\nBest AUC (Seed {seed}): {best_auc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T21:58:09.621968Z","iopub.execute_input":"2026-04-04T21:58:09.622407Z","iopub.status.idle":"2026-04-04T21:58:09.637975Z","shell.execute_reply.started":"2026-04-04T21:58:09.622373Z","shell.execute_reply":"2026-04-04T21:58:09.637376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_model(\n    seed=42,\n    save_name=\"best_resnet50_seed42.pth\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T21:58:19.225231Z","iopub.execute_input":"2026-04-04T21:58:19.225996Z","iopub.status.idle":"2026-04-04T23:06:53.061039Z","shell.execute_reply.started":"2026-04-04T21:58:19.225965Z","shell.execute_reply":"2026-04-04T23:06:53.060145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_model(\n    seed=2024,\n    save_name=\"best_resnet50_seed2024.pth\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-04T23:08:28.490119Z","iopub.execute_input":"2026-04-04T23:08:28.490837Z","iopub.status.idle":"2026-04-05T00:03:35.788861Z","shell.execute_reply.started":"2026-04-04T23:08:28.490805Z","shell.execute_reply":"2026-04-05T00:03:35.788075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model1 = RadiogenomicsModel().to(device)\nmodel2 = RadiogenomicsModel().to(device)\n\nckpt1 = torch.load(\"best_resnet50_seed42.pth\", weights_only=False)\nckpt2 = torch.load(\"best_resnet50_seed2024.pth\", weights_only=False)\n\nmodel1.load_state_dict(ckpt1['model'])\nmodel2.load_state_dict(ckpt2['model'])\n\nmodel1.eval()\nmodel2.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-05T00:37:37.175857Z","iopub.execute_input":"2026-04-05T00:37:37.176561Z","iopub.status.idle":"2026-04-05T00:37:38.278636Z","shell.execute_reply.started":"2026-04-05T00:37:37.176531Z","shell.execute_reply":"2026-04-05T00:37:38.277951Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Ensemble Inference (FINAL)","metadata":{}},{"cell_type":"code","source":"def tta_predict(model, x):\n    with torch.no_grad():\n        out1 = model(x)[0]\n\n        x_flip = torch.flip(x, dims=[3])\n        out2 = model(x_flip)[0]\n\n        prob1 = torch.sigmoid(out1)\n        prob2 = torch.sigmoid(out2)\n\n    return (prob1 + prob2) / 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-05T00:37:57.814347Z","iopub.execute_input":"2026-04-05T00:37:57.815128Z","iopub.status.idle":"2026-04-05T00:37:57.819417Z","shell.execute_reply.started":"2026-04-05T00:37:57.815097Z","shell.execute_reply":"2026-04-05T00:37:57.818621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_probs = []\nall_targets = []\n\nfor imgs, labels in val_loader:\n\n    imgs = imgs.to(device)\n\n    # TTA per model\n    p1 = tta_predict(model1, imgs)\n    p2 = tta_predict(model2, imgs)\n\n    # Ensemble\n    final_prob = (p1 + p2) / 2\n\n    all_probs.extend(final_prob.cpu().numpy().flatten())\n    all_targets.extend(labels['idh'].numpy().flatten())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-05T00:38:33.114549Z","iopub.execute_input":"2026-04-05T00:38:33.115252Z","iopub.status.idle":"2026-04-05T00:40:12.928990Z","shell.execute_reply.started":"2026-04-05T00:38:33.115220Z","shell.execute_reply":"2026-04-05T00:40:12.928366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score, accuracy_score\n\nauc = roc_auc_score(all_targets, all_probs)\nacc = accuracy_score(all_targets, (np.array(all_probs) > 0.5).astype(int))\n\nprint(\"Ensemble AUC:\", auc)\nprint(\"Ensemble Accuracy:\", acc)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-05T00:41:58.638636Z","iopub.execute_input":"2026-04-05T00:41:58.639070Z","iopub.status.idle":"2026-04-05T00:41:58.648128Z","shell.execute_reply.started":"2026-04-05T00:41:58.639043Z","shell.execute_reply":"2026-04-05T00:41:58.647521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import f1_score\n\nbest_t = 0\nbest_f1 = 0\n\nfor t in np.arange(0.30, 0.55, 0.01):\n\n    preds = (np.array(all_probs) > t).astype(int)\n    f1 = f1_score(all_targets, preds)\n\n    if f1 > best_f1:\n        best_f1 = f1\n        best_t = t\n\nprint(\"Best Threshold:\", best_t)\nprint(\"Best F1:\", best_f1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-05T00:51:52.379574Z","iopub.execute_input":"2026-04-05T00:51:52.379953Z","iopub.status.idle":"2026-04-05T00:51:52.455010Z","shell.execute_reply.started":"2026-04-05T00:51:52.379925Z","shell.execute_reply":"2026-04-05T00:51:52.454279Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**BEST THRUSHOLD FOR THE NEW ENSEMBLE MODEL IS 0.49 PREVIOUSLY THE BEST MODEL HAS THE THRESHOLD WAS 0.49**","metadata":{}},{"cell_type":"markdown","source":"INFERENCE CODE IS DOWN BELOW","metadata":{}},{"cell_type":"code","source":"# inference.py (or save in notebook)\n\nmodel1 = RadiogenomicsModel().to(device)\nmodel2 = RadiogenomicsModel().to(device)\n\nckpt1 = torch.load(\"best_resnet50_seed42.pth\", weights_only=False)\nckpt2 = torch.load(\"best_resnet50_seed2024.pth\", weights_only=False)\n\nmodel1.load_state_dict(ckpt1['model'])\nmodel2.load_state_dict(ckpt2['model'])\n\nmodel1.eval()\nmodel2.eval()\n\ndef predict(x):\n    with torch.no_grad():\n        p1 = torch.sigmoid(model1(x)[0])\n        p2 = torch.sigmoid(model2(x)[0])\n\n        prob = (p1 + p2) / 2\n        pred = (prob > 0.49).float()\n\n    return prob, pred","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# GRAD CAM","metadata":{}},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/input/datasets/mondaldebasish05/resnet50-42-2024-seed-models\"\n\nmodel1 = RadiogenomicsModel().to(device)\nmodel2 = RadiogenomicsModel().to(device)\n\nckpt1 = torch.load(\n    f\"{MODEL_PATH}/best_resnet50_seed42.pth\",\n    map_location=torch.device('cpu'),\n    weights_only=False\n)\n\nckpt2 = torch.load(\n    f\"{MODEL_PATH}/best_resnet50_seed2024.pth\",\n    map_location=torch.device('cpu'),\n    weights_only=False\n)\n\nmodel1.load_state_dict(ckpt1['model'])\nmodel2.load_state_dict(ckpt2['model'])\n\nmodel1.eval()\nmodel2.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:32:30.926063Z","iopub.execute_input":"2026-04-17T06:32:30.927246Z","iopub.status.idle":"2026-04-17T06:32:34.259022Z","shell.execute_reply.started":"2026-04-17T06:32:30.927167Z","shell.execute_reply":"2026-04-17T06:32:34.258282Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install grad-cam\nfrom pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:34:06.313623Z","iopub.execute_input":"2026-04-17T06:34:06.317458Z","iopub.status.idle":"2026-04-17T06:34:10.035521Z","shell.execute_reply.started":"2026-04-17T06:34:06.317414Z","shell.execute_reply":"2026-04-17T06:34:10.034345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"imgs, labels = next(iter(val_loader))\n\nimg = imgs[0]   # first sample\ninput_tensor = img.unsqueeze(0).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:37:46.882437Z","iopub.execute_input":"2026-04-17T06:37:46.882731Z","iopub.status.idle":"2026-04-17T06:37:54.886821Z","shell.execute_reply.started":"2026-04-17T06:37:46.882708Z","shell.execute_reply":"2026-04-17T06:37:54.885486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(input_tensor.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:38:03.913084Z","iopub.execute_input":"2026-04-17T06:38:03.913374Z","iopub.status.idle":"2026-04-17T06:38:03.919096Z","shell.execute_reply.started":"2026-04-17T06:38:03.913354Z","shell.execute_reply":"2026-04-17T06:38:03.917354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_layer1 = model1.backbone.layer4[-1]\ntarget_layer2 = model2.backbone.layer4[-1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:38:16.698692Z","iopub.execute_input":"2026-04-17T06:38:16.699017Z","iopub.status.idle":"2026-04-17T06:38:16.703966Z","shell.execute_reply.started":"2026-04-17T06:38:16.698993Z","shell.execute_reply":"2026-04-17T06:38:16.702582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n\ncam1 = GradCAM(model=model1, target_layers=[target_layer1])\ncam2 = GradCAM(model=model2, target_layers=[target_layer2])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:38:41.298577Z","iopub.execute_input":"2026-04-17T06:38:41.298854Z","iopub.status.idle":"2026-04-17T06:38:41.306720Z","shell.execute_reply.started":"2026-04-17T06:38:41.298834Z","shell.execute_reply":"2026-04-17T06:38:41.304848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"targets = [ClassifierOutputTarget(0)]\n\nheatmap1 = cam1(input_tensor=input_tensor, targets=targets)[0]\nheatmap2 = cam2(input_tensor=input_tensor, targets=targets)[0]\n\nheatmap = (heatmap1 + heatmap2) / 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:38:51.992279Z","iopub.execute_input":"2026-04-17T06:38:51.992606Z","iopub.status.idle":"2026-04-17T06:38:53.655295Z","shell.execute_reply.started":"2026-04-17T06:38:51.992578Z","shell.execute_reply":"2026-04-17T06:38:53.654561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_np = img[:3].permute(1,2,0).cpu().numpy()\nimg_np = (img_np - img_np.min()) / (img_np.max() - img_np.min())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:39:02.620373Z","iopub.execute_input":"2026-04-17T06:39:02.620811Z","iopub.status.idle":"2026-04-17T06:39:02.633048Z","shell.execute_reply.started":"2026-04-17T06:39:02.620789Z","shell.execute_reply":"2026-04-17T06:39:02.631997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_grad_cam.utils.image import show_cam_on_image\nimport matplotlib.pyplot as plt\n\nvisualization = show_cam_on_image(img_np, heatmap, use_rgb=True)\n\nplt.imshow(visualization)\nplt.axis('off')\nplt.title(\"Ensemble Grad-CAM (IDH)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:39:11.179024Z","iopub.execute_input":"2026-04-17T06:39:11.179371Z","iopub.status.idle":"2026-04-17T06:39:11.678348Z","shell.execute_reply.started":"2026-04-17T06:39:11.179349Z","shell.execute_reply":"2026-04-17T06:39:11.677295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/input/datasets/mondaldebasish05/resnet50-42-2024-seed-models\"\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = RadiogenomicsModel().to(device)\n\nckpt = torch.load(\n    f\"{MODEL_PATH}/best_resnet50_seed2024.pth\",\n    map_location=device,\n    weights_only=False\n)\n\nmodel.load_state_dict(ckpt['model'])\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:43:15.784308Z","iopub.execute_input":"2026-04-17T06:43:15.784693Z","iopub.status.idle":"2026-04-17T06:43:16.742107Z","shell.execute_reply.started":"2026-04-17T06:43:15.784667Z","shell.execute_reply":"2026-04-17T06:43:16.740670Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(5):\n    img, labels = val_dataset[i]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:50:51.768440Z","iopub.execute_input":"2026-04-17T06:50:51.768730Z","iopub.status.idle":"2026-04-17T06:50:53.812100Z","shell.execute_reply.started":"2026-04-17T06:50:51.768710Z","shell.execute_reply":"2026-04-17T06:50:53.811421Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# FINAL GRADCAM AFTER CROPPING","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\n\ndef crop_brain(img):\n    # img: (4, H, W)\n    \n    mask = img.sum(dim=0) > 0  # non-zero pixels\n    \n    coords = mask.nonzero()\n    \n    y_min, x_min = coords.min(dim=0)[0]\n    y_max, x_max = coords.max(dim=0)[0]\n    \n    cropped = img[:, y_min:y_max, x_min:x_max]\n    \n    return cropped\n\n# Apply crop\nimg = crop_brain(img)\n\n# Resize back to 224x224\nimg = F.interpolate(\n    img.unsqueeze(0),\n    size=(224, 224),\n    mode='bilinear',\n    align_corners=False\n).squeeze(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:50:56.550661Z","iopub.execute_input":"2026-04-17T06:50:56.551034Z","iopub.status.idle":"2026-04-17T06:50:56.560076Z","shell.execute_reply.started":"2026-04-17T06:50:56.551003Z","shell.execute_reply":"2026-04-17T06:50:56.558729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img = (img - img.mean()) / (img.std() + 1e-5)\n\ninput_tensor = img.unsqueeze(0).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:50:58.698976Z","iopub.execute_input":"2026-04-17T06:50:58.699285Z","iopub.status.idle":"2026-04-17T06:50:58.704581Z","shell.execute_reply.started":"2026-04-17T06:50:58.699264Z","shell.execute_reply":"2026-04-17T06:50:58.703536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class IDHTarget:\n    def __call__(self, model_output):\n        return model_output[0]   # IDH head","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:51:00.244347Z","iopub.execute_input":"2026-04-17T06:51:00.244703Z","iopub.status.idle":"2026-04-17T06:51:00.249793Z","shell.execute_reply.started":"2026-04-17T06:51:00.244667Z","shell.execute_reply":"2026-04-17T06:51:00.248752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\n\ntarget_layer = model.backbone.layer4[-1]\n\ncam = GradCAM(\n    model=model,\n    target_layers=[target_layer]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:51:02.486486Z","iopub.execute_input":"2026-04-17T06:51:02.486788Z","iopub.status.idle":"2026-04-17T06:51:02.494630Z","shell.execute_reply.started":"2026-04-17T06:51:02.486766Z","shell.execute_reply":"2026-04-17T06:51:02.493493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"targets = [IDHTarget()]\n\nheatmap = cam(input_tensor=input_tensor, targets=targets)[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:51:04.192446Z","iopub.execute_input":"2026-04-17T06:51:04.192771Z","iopub.status.idle":"2026-04-17T06:51:04.492471Z","shell.execute_reply.started":"2026-04-17T06:51:04.192748Z","shell.execute_reply":"2026-04-17T06:51:04.491786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_np = img[:3].permute(1,2,0).cpu().numpy()\nimg_np = (img_np - img_np.min()) / (img_np.max() - img_np.min())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:51:05.972507Z","iopub.execute_input":"2026-04-17T06:51:05.972817Z","iopub.status.idle":"2026-04-17T06:51:05.979915Z","shell.execute_reply.started":"2026-04-17T06:51:05.972795Z","shell.execute_reply":"2026-04-17T06:51:05.978565Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nvisualization = show_cam_on_image(img_np, heatmap, use_rgb=True)\n\nplt.imshow(visualization)\nplt.axis('off')\nplt.title(\"Grad-CAM (IDH - Seed 2024)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:51:38.507251Z","iopub.execute_input":"2026-04-17T06:51:38.507554Z","iopub.status.idle":"2026-04-17T06:51:38.623526Z","shell.execute_reply.started":"2026-04-17T06:51:38.507533Z","shell.execute_reply":"2026-04-17T06:51:38.622749Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with torch.no_grad():\n    out = model(input_tensor)[0]\n    prob = torch.sigmoid(out)\n\nprint(\"Pred Prob:\", prob.item())\nprint(\"True Label:\", labels['idh'][0].item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:51:42.874534Z","iopub.execute_input":"2026-04-17T06:51:42.875130Z","iopub.status.idle":"2026-04-17T06:51:43.016052Z","shell.execute_reply.started":"2026-04-17T06:51:42.875100Z","shell.execute_reply":"2026-04-17T06:51:43.015070Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"good_cases = []\n\nmodel.eval()\nwith torch.no_grad():\n    for i in range(len(val_dataset)):\n        img, labels = val_dataset[i]\n        x = img.unsqueeze(0).to(device)\n\n        out = model(x)[0]\n        prob = torch.sigmoid(out).item()\n        true = labels['idh'].item()\n\n        pred = 1 if prob > 0.49 else 0\n        confidence = abs(prob - 0.5)\n\n        if pred == true and confidence > 0.25:  # high-confidence correct\n            good_cases.append((i, prob, true))\n\n# take top 3 most confident\ngood_cases = sorted(good_cases, key=lambda x: abs(x[1]-0.5), reverse=True)[:3]\nprint(good_cases)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T06:57:15.690291Z","iopub.execute_input":"2026-04-17T06:57:15.691072Z","iopub.status.idle":"2026-04-17T06:59:08.083455Z","shell.execute_reply.started":"2026-04-17T06:57:15.691038Z","shell.execute_reply":"2026-04-17T06:59:08.082554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = []\n\nfor idx, prob, true in good_cases:\n    img, labels = val_dataset[idx]\n\n    # crop + resize (same as your working pipeline)\n    img = crop_brain(img)\n    img = F.interpolate(img.unsqueeze(0), size=(224,224), mode='bilinear', align_corners=False).squeeze(0)\n    img = (img - img.mean()) / (img.std() + 1e-5)\n\n    x = img.unsqueeze(0).to(device)\n\n    heatmap = cam(input_tensor=x, targets=[IDHTarget()])[0]\n\n    img_np = img[:3].permute(1,2,0).cpu().numpy()\n    img_np = (img_np - img_np.min()) / (img_np.max() - img_np.min())\n\n    vis = show_cam_on_image(img_np, heatmap, use_rgb=True)\n\n    results.append((idx, prob, true, vis))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T07:01:03.991704Z","iopub.execute_input":"2026-04-17T07:01:03.992058Z","iopub.status.idle":"2026-04-17T07:01:05.889709Z","shell.execute_reply.started":"2026-04-17T07:01:03.992030Z","shell.execute_reply":"2026-04-17T07:01:05.888448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nselected_indices = [13, 48, 62]\n\nfor idx in selected_indices:\n    \n    img, labels = val_dataset[idx]\n\n    # ===== CROP =====\n    img = crop_brain(img)\n    img = F.interpolate(\n        img.unsqueeze(0),\n        size=(224,224),\n        mode='bilinear',\n        align_corners=False\n    ).squeeze(0)\n\n    # ===== NORMALIZE =====\n    img = (img - img.mean()) / (img.std() + 1e-5)\n\n    input_tensor = img.unsqueeze(0).to(device)\n\n    # ===== PRED =====\n    with torch.no_grad():\n        out = model(input_tensor)[0]\n        prob = torch.sigmoid(out).item()\n\n    # ===== GRAD-CAM =====\n    heatmap = cam(input_tensor=input_tensor, targets=[IDHTarget()])[0]\n\n    # ===== IMAGE PREP =====\n    img_np = img[:3].permute(1,2,0).cpu().numpy()\n    img_np = (img_np - img_np.min()) / (img_np.max() - img_np.min())\n\n    vis = show_cam_on_image(img_np, heatmap, use_rgb=True)\n\n    # ===== SHOW =====\n    plt.figure(figsize=(5,5))\n    plt.imshow(vis)\n    plt.axis('off')\n    plt.title(f\"Idx {idx} | Prob {prob:.3f} | True {labels['idh'].item()}\")\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T07:03:01.533035Z","iopub.execute_input":"2026-04-17T07:03:01.533391Z","iopub.status.idle":"2026-04-17T07:03:04.355318Z","shell.execute_reply.started":"2026-04-17T07:03:01.533370Z","shell.execute_reply":"2026-04-17T07:03:04.354324Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MODEL EVALUTION FINAL ","metadata":{}},{"cell_type":"code","source":"model_path = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone/best_model_v5.pth\"\n\nmodel = RadiogenomicsModel().to(device)\n\ncheckpoint = torch.load(model_path, map_location=device, weights_only=False)\n\nmodel.load_state_dict(checkpoint[\"model\"])   # 🔥 KEY LINE\nmodel.eval()\n\nprint(\"✅ EfficientNet-B2 model loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:38:37.894027Z","iopub.execute_input":"2026-04-18T21:38:37.894792Z","iopub.status.idle":"2026-04-18T21:38:38.355305Z","shell.execute_reply.started":"2026-04-18T21:38:37.894760Z","shell.execute_reply":"2026-04-18T21:38:38.354623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# LOAD RESNET50 v2 MODEL\n# =========================\n\nimport torch\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel_path = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone/best_resnet50_v2.pth\"\n\n# Initialize model (your class)\nmodel = RadiogenomicsModel().to(device)\n\n# Load checkpoint\ncheckpoint = torch.load(model_path, map_location=device, weights_only=False)\n\n# 🔥 IMPORTANT: extract only weights\nif \"model\" in checkpoint:\n    model.load_state_dict(checkpoint[\"model\"])\nelse:\n    model.load_state_dict(checkpoint)\n\nmodel.eval()\n\nprint(\"✅ ResNet50 v2 loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:49:08.866543Z","iopub.execute_input":"2026-04-18T21:49:08.867307Z","iopub.status.idle":"2026-04-18T21:49:11.361980Z","shell.execute_reply.started":"2026-04-18T21:49:08.867276Z","shell.execute_reply":"2026-04-18T21:49:11.361354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"base_path = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone\"\n\npath_42   = f\"{base_path}/best_resnet50_seed42.pth\"\npath_2024 = f\"{base_path}/best_resnet50_seed2024.pth\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:51:35.045043Z","iopub.execute_input":"2026-04-18T21:51:35.045488Z","iopub.status.idle":"2026-04-18T21:51:35.049771Z","shell.execute_reply.started":"2026-04-18T21:51:35.045459Z","shell.execute_reply":"2026-04-18T21:51:35.049159Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_42 = RadiogenomicsModel().to(device)\n\ncheckpoint_42 = torch.load(path_42, map_location=device, weights_only=False)\n\nif \"model\" in checkpoint_42:\n    model_42.load_state_dict(checkpoint_42[\"model\"])\nelse:\n    model_42.load_state_dict(checkpoint_42)\n\nmodel_42.eval()\n\nprint(\"✅ ResNet50 seed42 loaded\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:51:44.781227Z","iopub.execute_input":"2026-04-18T21:51:44.782229Z","iopub.status.idle":"2026-04-18T21:51:46.947701Z","shell.execute_reply.started":"2026-04-18T21:51:44.782183Z","shell.execute_reply":"2026-04-18T21:51:46.947039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_2024 = RadiogenomicsModel().to(device)\n\ncheckpoint_2024 = torch.load(path_2024, map_location=device, weights_only=False)\n\nif \"model\" in checkpoint_2024:\n    model_2024.load_state_dict(checkpoint_2024[\"model\"])\nelse:\n    model_2024.load_state_dict(checkpoint_2024)\n\nmodel_2024.eval()\n\nprint(\"✅ ResNet50 seed2024 loaded\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:51:55.534244Z","iopub.execute_input":"2026-04-18T21:51:55.535039Z","iopub.status.idle":"2026-04-18T21:51:57.682729Z","shell.execute_reply.started":"2026-04-18T21:51:55.535009Z","shell.execute_reply":"2026-04-18T21:51:57.681859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== LOAD MODEL =====\nmodel_path = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone/best_convnext.pth\"\n\nmodel = RadiogenomicsModel().to(device)\n\ncheckpoint = torch.load(model_path, map_location=device, weights_only=False)\n\nif \"model\" in checkpoint:\n    model.load_state_dict(checkpoint[\"model\"])\nelse:\n    model.load_state_dict(checkpoint)\n\nmodel.eval()\n\nprint(\"✅ ConvNeXt loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T21:58:59.850219Z","iopub.execute_input":"2026-04-18T21:58:59.850982Z","iopub.status.idle":"2026-04-18T21:59:05.484358Z","shell.execute_reply.started":"2026-04-18T21:58:59.850952Z","shell.execute_reply":"2026-04-18T21:59:05.483546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== LOAD MODEL =====\nmodel_path = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone/best_resnet101.pth\"\n\nmodel = RadiogenomicsModel().to(device)\n\ncheckpoint = torch.load(model_path, map_location=device, weights_only=False)\n\nif \"model\" in checkpoint:\n    model.load_state_dict(checkpoint[\"model\"])\nelse:\n    model.load_state_dict(checkpoint)\n\nmodel.eval()\n\nprint(\"✅ ResNet101 loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:00:31.037986Z","iopub.execute_input":"2026-04-18T22:00:31.039018Z","iopub.status.idle":"2026-04-18T22:00:33.705113Z","shell.execute_reply.started":"2026-04-18T22:00:31.038985Z","shell.execute_reply":"2026-04-18T22:00:33.704252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone/best_densenet121.pth\"\n\nmodel = RadiogenomicsModel().to(device)\n\ncheckpoint = torch.load(model_path, map_location=device, weights_only=False)\n\nif \"model\" in checkpoint:\n    model.load_state_dict(checkpoint[\"model\"])\nelse:\n    model.load_state_dict(checkpoint)\n\nmodel.eval()\n\nprint(\"✅ DenseNet121 loaded successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:02:16.802672Z","iopub.execute_input":"2026-04-18T22:02:16.803415Z","iopub.status.idle":"2026-04-18T22:02:17.128124Z","shell.execute_reply.started":"2026-04-18T22:02:16.803385Z","shell.execute_reply":"2026-04-18T22:02:17.127407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print([k for k in globals().keys() if k.startswith(\"model\")])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:15:53.989516Z","iopub.execute_input":"2026-04-19T21:15:53.990288Z","iopub.status.idle":"2026-04-19T21:15:53.996342Z","shell.execute_reply.started":"2026-04-19T21:15:53.990247Z","shell.execute_reply":"2026-04-19T21:15:53.995268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# LOAD ALL MODELS (FINAL)\n# =========================\n\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport timm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nBASE_PATH = \"/kaggle/input/datasets/mondaldebasish05/radiogenomics-best-models-multibackbone\"\n\n# =========================\n# MODEL DEFINITIONS\n# =========================\n\n# ---- EfficientNet ----\nclass EffNetModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = models.efficientnet_b2(weights=None)\n\n        orig = self.backbone.features[0][0]\n        self.backbone.features[0][0] = nn.Conv2d(4, orig.out_channels,\n            kernel_size=orig.kernel_size, stride=orig.stride, padding=orig.padding, bias=False)\n\n        in_features = self.backbone.classifier[1].in_features\n        self.backbone.classifier = nn.Identity()\n\n        self.fusion = nn.Sequential(nn.Dropout(0.4), nn.Linear(in_features,512), nn.ReLU())\n        self.idh_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,1))\n        self.mgmt_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,1))\n        self.grade_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,2))\n\n    def forward(self,x):\n        f = self.backbone(x)\n        f = self.fusion(f)\n        return self.idh_head(f), self.mgmt_head(f), self.grade_head(f)\n\n\n# ---- ResNet50 / 101 ----\nclass ResNetModel(nn.Module):\n    def __init__(self, depth=50):\n        super().__init__()\n        self.backbone = models.resnet50(weights=None) if depth==50 else models.resnet101(weights=None)\n\n        orig = self.backbone.conv1\n        self.backbone.conv1 = nn.Conv2d(4,64,7,2,3,bias=False)\n\n        in_features = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n\n        self.fusion = nn.Sequential(nn.Dropout(0.45), nn.Linear(in_features,512), nn.ReLU())\n        self.idh_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.45), nn.Linear(256,1))\n        self.mgmt_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.45), nn.Linear(256,1))\n        self.grade_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.45), nn.Linear(256,2))\n\n    def forward(self,x):\n        f = self.backbone(x)\n        f = self.fusion(f)\n        return self.idh_head(f), self.mgmt_head(f), self.grade_head(f)\n\n\n# ---- DenseNet ----\nclass DenseNetModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = models.densenet121(weights=None)\n\n        orig = self.backbone.features.conv0\n        self.backbone.features.conv0 = nn.Conv2d(4, orig.out_channels,\n            kernel_size=orig.kernel_size, stride=orig.stride, padding=orig.padding, bias=False)\n\n        in_features = self.backbone.classifier.in_features\n        self.backbone.classifier = nn.Identity()\n\n        self.fusion = nn.Sequential(nn.Dropout(0.4), nn.Linear(in_features,512), nn.ReLU())\n        self.idh_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,1))\n        self.mgmt_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,1))\n        self.grade_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,2))\n\n    def forward(self,x):\n        f = self.backbone.features(x)\n        f = torch.relu(f)\n        f = torch.nn.functional.adaptive_avg_pool2d(f,(1,1))\n        f = torch.flatten(f,1)\n        f = self.fusion(f)\n        return self.idh_head(f), self.mgmt_head(f), self.grade_head(f)\n\n\n# ---- ConvNeXt ----\nclass ConvNeXtModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model(\"convnext_base\", pretrained=False, num_classes=0)\n\n        orig = self.backbone.stem[0]\n        self.backbone.stem[0] = nn.Conv2d(4, orig.out_channels,\n            kernel_size=orig.kernel_size, stride=orig.stride, padding=orig.padding, bias=False)\n\n        in_features = self.backbone.num_features\n\n        self.fusion = nn.Sequential(nn.Dropout(0.4), nn.Linear(in_features,512), nn.ReLU())\n        self.idh_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,1))\n        self.mgmt_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,1))\n        self.grade_head = nn.Sequential(nn.Linear(512,256), nn.ReLU(), nn.Dropout(0.4), nn.Linear(256,2))\n\n    def forward(self,x):\n        f = self.backbone(x)\n        f = self.fusion(f)\n        return self.idh_head(f), self.mgmt_head(f), self.grade_head(f)\n\n\n# =========================\n# LOAD FUNCTION\n# =========================\ndef load_model(path, model):\n    ckpt = torch.load(path, map_location=device, weights_only=False)\n    if \"model\" in ckpt:\n        model.load_state_dict(ckpt[\"model\"])\n    else:\n        model.load_state_dict(ckpt)\n    model.eval()\n    return model\n\n\n# =========================\n# LOAD ALL MODELS\n# =========================\n\nmodels_dict = {}\n\nmodels_dict[\"efficientnet_b2\"] = load_model(f\"{BASE_PATH}/best_model_v5.pth\", EffNetModel().to(device))\n\nmodels_dict[\"resnet50_v2\"] = load_model(f\"{BASE_PATH}/best_resnet50_v2.pth\", ResNetModel(50).to(device))\nmodels_dict[\"resnet50_seed42\"] = load_model(f\"{BASE_PATH}/best_resnet50_seed42.pth\", ResNetModel(50).to(device))\nmodels_dict[\"resnet50_seed2024\"] = load_model(f\"{BASE_PATH}/best_resnet50_seed2024.pth\", ResNetModel(50).to(device))\n\nmodels_dict[\"resnet101\"] = load_model(f\"{BASE_PATH}/best_resnet101.pth\", ResNetModel(101).to(device))\n\nmodels_dict[\"densenet121\"] = load_model(f\"{BASE_PATH}/best_densenet121.pth\", DenseNetModel().to(device))\n\nmodels_dict[\"convnext_base\"] = load_model(f\"{BASE_PATH}/best_convnext.pth\", ConvNeXtModel().to(device))\n\n\nprint(\"✅ ALL MODELS LOADED:\", list(models_dict.keys()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:15:56.792011Z","iopub.execute_input":"2026-04-19T21:15:56.792377Z","iopub.status.idle":"2026-04-19T21:16:02.004938Z","shell.execute_reply.started":"2026-04-19T21:15:56.792348Z","shell.execute_reply":"2026-04-19T21:16:02.003628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# EVALUATE ALL MODELS (FINAL FIX)\n# =========================\n\nimport numpy as np\nfrom sklearn.metrics import roc_auc_score, accuracy_score\n\nresults = {}\n\nfor name, model in models_dict.items():\n\n    print(f\"Evaluating {name}\")\n\n    idh_preds, idh_targets = [], []\n    mgmt_preds, mgmt_targets = [], []\n    grade_preds, grade_targets = [], []\n\n    model.eval()\n\n    with torch.no_grad():\n        for batch in val_loader:\n\n            # 🔥 YOUR FORMAT\n            x = batch[0]\n            labels = batch[1]\n\n            idh = labels[\"idh\"]\n            mgmt = labels[\"mgmt\"]\n            grade = labels[\"grade\"]\n\n            x = x.to(device)\n            idh = idh.to(device)\n            mgmt = mgmt.to(device)\n            grade = grade.to(device)\n\n            out_idh, out_mgmt, out_grade = model(x)\n\n            # IDH\n            idh_prob = torch.sigmoid(out_idh).squeeze().cpu().numpy()\n            idh_preds.extend(idh_prob)\n            idh_targets.extend(idh.cpu().numpy())\n\n            # MGMT (mask -1)\n            mgmt_prob = torch.sigmoid(out_mgmt).squeeze().cpu().numpy()\n            mgmt_np = mgmt.cpu().numpy()\n            mask = mgmt_np != -1\n\n            mgmt_preds.extend(mgmt_prob[mask])\n            mgmt_targets.extend(mgmt_np[mask])\n\n            # Grade\n            grade_pred = torch.argmax(out_grade, dim=1).cpu().numpy()\n            grade_preds.extend(grade_pred)\n            grade_targets.extend(grade.cpu().numpy())\n\n    results[name] = {\n        \"IDH_AUC\": roc_auc_score(idh_targets, idh_preds),\n        \"MGMT_ACC\": accuracy_score(mgmt_targets, (np.array(mgmt_preds) > 0.5).astype(int)),\n        \"GRADE_ACC\": accuracy_score(grade_targets, grade_preds)\n    }\n\nprint(\"\\n✅ FINAL RESULTS:\\n\", results)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:16:15.960299Z","iopub.execute_input":"2026-04-19T21:16:15.961012Z","iopub.status.idle":"2026-04-19T21:17:10.838613Z","shell.execute_reply.started":"2026-04-19T21:16:15.960979Z","shell.execute_reply":"2026-04-19T21:17:10.837433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FIGURE 3: BACKBONE COMPARISON\n# =========================\n\nimport matplotlib.pyplot as plt\n\nnames = list(results.keys())\nidh_auc = [results[n][\"IDH_AUC\"] for n in names]\n\nplt.figure(figsize=(8,5))\nplt.bar(names, idh_auc)\n\nplt.ylim(0.65, 0.85)   # 🔥 REQUIRED ZOOM\nplt.ylabel(\"IDH AUC\")\nplt.title(\"Backbone-wise Performance Comparison\")\n\nplt.xticks(rotation=30)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:53:01.706700Z","iopub.execute_input":"2026-04-18T22:53:01.707527Z","iopub.status.idle":"2026-04-18T22:53:01.916328Z","shell.execute_reply.started":"2026-04-18T22:53:01.707497Z","shell.execute_reply":"2026-04-18T22:53:01.915677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# MULTI-TASK PERFORMANCE\n# =========================\n\nimport numpy as np\n\nmgmt = [results[n][\"MGMT_ACC\"] for n in names]\ngrade = [results[n][\"GRADE_ACC\"] for n in names]\n\nx = np.arange(len(names))\nw = 0.25\n\nplt.figure(figsize=(10,5))\nplt.bar(x-w, idh_auc, width=w, label=\"IDH AUC\")\nplt.bar(x, mgmt, width=w, label=\"MGMT Acc\")\nplt.bar(x+w, grade, width=w, label=\"Grade Acc\")\n\nplt.xticks(x, names, rotation=30)\nplt.legend()\nplt.title(\"Multi-task Performance Comparison\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:53:24.771925Z","iopub.execute_input":"2026-04-18T22:53:24.772511Z","iopub.status.idle":"2026-04-18T22:53:24.971169Z","shell.execute_reply.started":"2026-04-18T22:53:24.772480Z","shell.execute_reply":"2026-04-18T22:53:24.970521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# ENSEMBLE EFFECT\n# =========================\n\nsingle = results[\"resnet50_seed42\"][\"IDH_AUC\"]\n\nensemble = (\n    results[\"resnet50_seed42\"][\"IDH_AUC\"] +\n    results[\"resnet50_seed2024\"][\"IDH_AUC\"]\n) / 2\n\nlabels = [\"Single\", \"Ensemble\"]\nvalues = [single, ensemble]\n\nplt.figure(figsize=(4,4))\nplt.bar(labels, values)\n\nplt.ylim(0.65, 0.85)\nplt.title(\"Effect of Ensembling\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:54:55.275530Z","iopub.execute_input":"2026-04-18T22:54:55.275836Z","iopub.status.idle":"2026-04-18T22:54:55.371736Z","shell.execute_reply.started":"2026-04-18T22:54:55.275790Z","shell.execute_reply":"2026-04-18T22:54:55.371149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Clean names (paper-ready)\nnames = [\n    \"EffNet-B2\",\n    \"ResNet50 (v2)\",\n    \"ResNet50 (s42)\",\n    \"ResNet50 (s2024)\",\n    \"ResNet101\",\n    \"DenseNet121\",\n    \"ConvNeXt\"\n]\n\nvalues = [\n    0.776,\n    0.792,  # best\n    0.778,\n    0.786,\n    0.752,\n    0.759,\n    0.765\n]\n\nplt.figure(figsize=(8,5))\n\nbars = plt.bar(names, values)\n\n# Highlight best model\nbest_idx = np.argmax(values)\nbars[best_idx].set_linewidth(2)\nbars[best_idx].set_edgecolor('black')\n\nplt.ylim(0.74, 0.80)\n\nplt.ylabel(\"IDH AUC\")\nplt.title(\"Backbone Performance Comparison\")\n\nplt.xticks(rotation=25)\n\n# Annotate values (important for paper)\nfor i, v in enumerate(values):\n    plt.text(i, v + 0.002, f\"{v:.3f}\", ha='center', fontsize=9)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-18T22:59:21.675619Z","iopub.execute_input":"2026-04-18T22:59:21.676431Z","iopub.status.idle":"2026-04-18T22:59:21.833441Z","shell.execute_reply.started":"2026-04-18T22:59:21.676400Z","shell.execute_reply":"2026-04-18T22:59:21.832778Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nnames = [\n    \"EffNet-B2\",\n    \"ResNet-50\",\n    \"ResNet-101\",\n    \"DenseNet-121\",\n    \"ConvNeXt\"\n]\n\nvalues = [\n    0.776,\n    0.792,  # best\n    0.752,\n    0.759,\n    0.765\n]\n\nplt.figure(figsize=(7,4.5))\n\nbars = plt.bar(names, values)\n\n# Highlight best model\nbest_idx = np.argmax(values)\nbars[best_idx].set_edgecolor('black')\nbars[best_idx].set_linewidth(2)\n\nplt.ylim(0.74, 0.80)\nplt.ylabel(\"AUC (IDH)\")\nplt.title(\"Backbone-wise Performance Comparison\")\n\n# Clean ticks\nplt.xticks(rotation=20)\n\n# Annotate values\nfor i, v in enumerate(values):\n    plt.text(i, v + 0.002, f\"{v:.3f}\", ha='center', fontsize=9)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:20:58.338371Z","iopub.execute_input":"2026-04-19T21:20:58.338748Z","iopub.status.idle":"2026-04-19T21:20:58.511093Z","shell.execute_reply.started":"2026-04-19T21:20:58.338718Z","shell.execute_reply":"2026-04-19T21:20:58.510102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: ORIGINAL + GRADE GRAD-CAM + TABLE\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS --------\nGREEN = \"#cfe8cf\"\nRED   = \"#f6c1c1\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# MULTI-CHANNEL CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img_norm = (img - img.min()) / (img.max() - img.min() + 1e-8)\n    mask = img_norm > threshold\n\n    coords = np.argwhere(mask)\n    if coords.shape[0] == 0:\n        return x\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - pad)\n    x0 = max(0, x0 - pad)\n    y1 = min(img.shape[0], y1 + pad)\n    x1 = min(img.shape[1], x1 + pad)\n\n    return x[:, y0:y1, x0:x1]\n\n# =========================\n# GRAD-CAM SETUP\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(module, inp, out):\n    features.clear()\n    features.append(out)\n\ndef bwd_hook(module, grad_in, grad_out):\n    gradients.clear()\n    gradients.append(grad_out[0])\n\ntarget_layer = model.backbone.layer4[-1]\nh1 = target_layer.register_forward_hook(fwd_hook)\nh2 = target_layer.register_backward_hook(bwd_hook)\n\ndef get_gradcam(x_batch):\n    out_idh, out_mgmt, out_grade = model(x_batch)\n\n    # 🔥 USE GRADE HEAD (better localization)\n    score = out_grade[0, 1]   # class 1 = High grade\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    weights = torch.mean(grad, dim=(2,3), keepdim=True)\n    cam = torch.sum(weights * fmap, dim=1).squeeze()\n\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(7.16, 5.6))\ngs = fig.add_gridspec(\n    cases_to_show + 1,\n    6,\n    height_ratios=[0.9] + [1.4]*cases_to_show,\n    wspace=0.15,\n    hspace=0.25\n)\n\n# -------- HEADER --------\nheaders = [\"Original\", \"Grad-CAM\", \"IDH\", \"MGMT\", \"Grade\", \"Conf\"]\nfor j, h in enumerate(headers):\n    ax = fig.add_subplot(gs[0, j])\n    ax.text(0.5, 0.5, h, ha='center', va='center', fontsize=9, weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n    for sp in ax.spines.values():\n        sp.set_linewidth(0.8)\n\n# =========================\n# DATA LOOP\n# =========================\ncount = 0\n\nfor batch in val_loader:\n    x = batch[0].to(device)\n    labels = batch[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count >= cases_to_show:\n            break\n\n        row = count + 1\n\n        # -------- ORIGINAL --------\n        img_orig = x[i][1].cpu().numpy()  # 🔥 T1ce\n        ax = fig.add_subplot(gs[row, 0])\n        ax.imshow(img_orig, cmap='gray')\n        ax.set_aspect('auto')\n        ax.axis('off')\n\n        # -------- CROPPED INPUT --------\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        # -------- MODEL PREDICTION --------\n        out_idh_c, out_mgmt_c, out_grade_c = model(x_crop_b)\n\n        idh_prob_c = torch.sigmoid(out_idh_c).view(-1)[0].item()\n        mgmt_prob_c = torch.sigmoid(out_mgmt_c).view(-1)[0].item()\n        grade_p_c = int(torch.argmax(out_grade_c, dim=1).item())\n\n        # -------- GRAD-CAM --------\n        cam = get_gradcam(x_crop_b)\n\n        img_cropped = x_crop[1].cpu().numpy()  # 🔥 T1ce\n\n        cam = cv2.resize(cam, (img_cropped.shape[1], img_cropped.shape[0]))\n\n        # 🔥 SMOOTH + CLEAN\n        cam = cv2.GaussianBlur(cam, (11,11), 0)\n        cam[cam < 0.4] = 0\n\n        ax = fig.add_subplot(gs[row, 1])\n        ax.imshow(img_cropped, cmap='gray')\n        ax.imshow(cam, cmap='jet', alpha=0.45)\n        ax.set_aspect('auto')\n        ax.axis('off')\n\n        # -------- TEXT --------\n        idh_p = int(idh_prob_c > 0.5)\n        mgmt_p = int(mgmt_prob_c > 0.5)\n\n        idh_text = \"Mutant\" if idh_p else \"Wildtype\"\n        mgmt_text = \"Methylated\" if mgmt_p else \"Unmethylated\"\n        grade_text = \"High\" if grade_p_c else \"Low\"\n\n        if mgmt[i].item() == -1:\n            mgmt_text = \"N/A\"\n\n        conf = f\"{idh_prob_c:.2f}\"\n\n        idh_c = idh_p == idh[i].item()\n        mgmt_c = (mgmt[i].item() == -1) or (mgmt_p == mgmt[i].item())\n        grade_c = grade_p_c == grade[i].item()\n\n        colors = [\n            GREEN if idh_c else RED,\n            GRAY if mgmt[i].item() == -1 else (GREEN if mgmt_c else RED),\n            GREEN if grade_c else RED\n        ]\n\n        values = [idh_text, mgmt_text, grade_text, conf]\n\n        for col in range(4):\n            ax = fig.add_subplot(gs[row, col+2])\n            ax.text(0.5, 0.5, values[col], ha='center', va='center', fontsize=8)\n            ax.set_facecolor(colors[col] if col < 3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n            for sp in ax.spines.values():\n                sp.set_linewidth(0.8)\n\n        count += 1\n\n    if count >= cases_to_show:\n        break\n\n# -------- FINAL --------\nplt.suptitle(\"Model Predictions with Grade-based Grad-CAM\", fontsize=10, y=0.98)\n\nfig.text(0.5, 0.01,\n         \"Green: Correct   |   Red: Incorrect   |   Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout(pad=1.2)\nplt.savefig(\"final_fixed_gradcam.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:26:48.449838Z","iopub.execute_input":"2026-04-19T21:26:48.450409Z","iopub.status.idle":"2026-04-19T21:26:58.935065Z","shell.execute_reply.started":"2026-04-19T21:26:48.450373Z","shell.execute_reply":"2026-04-19T21:26:58.934295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: CLEAN GRAD-CAM PIPELINE\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS --------\nGREEN = \"#cfe8cf\"\nRED   = \"#f6c1c1\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# CROP FUNCTION (MULTI-CHANNEL)\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img_norm = (img - img.min()) / (img.max() - img.min() + 1e-8)\n    mask = img_norm > threshold\n\n    coords = np.argwhere(mask)\n    if coords.shape[0] == 0:\n        return x\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - pad)\n    x0 = max(0, x0 - pad)\n    y1 = min(img.shape[0], y1 + pad)\n    x1 = min(img.shape[1], x1 + pad)\n\n    return x[:, y0:y1, x0:x1]\n\n# =========================\n# GRAD-CAM SETUP\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(module, inp, out):\n    features.clear()\n    features.append(out)\n\ndef bwd_hook(module, grad_in, grad_out):\n    gradients.clear()\n    gradients.append(grad_out[0])\n\ntarget_layer = model.backbone.layer4[-1]\nh1 = target_layer.register_forward_hook(fwd_hook)\nh2 = target_layer.register_backward_hook(bwd_hook)\n\ndef get_gradcam(x_batch):\n    out_idh, out_mgmt, out_grade = model(x_batch)\n\n    # 🔥 USE GRADE HEAD (better localization)\n    score = out_grade[0, 1]\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    weights = torch.mean(grad, dim=(2,3), keepdim=True)\n    cam = torch.sum(weights * fmap, dim=1).squeeze()\n\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(7.16, 5.6))\ngs = fig.add_gridspec(\n    cases_to_show + 1,\n    6,\n    height_ratios=[0.9] + [1.4]*cases_to_show,\n    wspace=0.15,\n    hspace=0.25\n)\n\n# -------- HEADER --------\nheaders = [\"Original\", \"Grad-CAM\", \"IDH\", \"MGMT\", \"Grade\", \"Conf\"]\nfor j, h in enumerate(headers):\n    ax = fig.add_subplot(gs[0, j])\n    ax.text(0.5, 0.5, h, ha='center', va='center',\n            fontsize=9, weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n    for sp in ax.spines.values():\n        sp.set_linewidth(0.8)\n\n# =========================\n# DATA LOOP\n# =========================\ncount = 0\n\nfor batch in val_loader:\n\n    x = batch[0].to(device)\n    labels = batch[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count >= cases_to_show:\n            break\n\n        row = count + 1\n\n        # -------- ORIGINAL IMAGE (T1ce) --------\n        img_orig = x[i][1].cpu().numpy()\n\n        ax = fig.add_subplot(gs[row, 0])\n        ax.imshow(img_orig, cmap='gray')\n        ax.axis('off')\n\n        # -------- CROP --------\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        # -------- MODEL --------\n        out_idh_c, out_mgmt_c, out_grade_c = model(x_crop_b)\n\n        idh_prob_c = torch.sigmoid(out_idh_c).item()\n        mgmt_prob_c = torch.sigmoid(out_mgmt_c).item()\n        grade_p_c = int(torch.argmax(out_grade_c, dim=1).item())\n\n        # -------- GRAD-CAM --------\n        cam = get_gradcam(x_crop_b)\n\n        img_cropped = x_crop[1].cpu().numpy()  # T1ce\n\n        # normalize image\n        img_cropped = (img_cropped - img_cropped.min()) / \\\n                      (img_cropped.max() - img_cropped.min() + 1e-8)\n\n        cam = cv2.resize(cam, (img_cropped.shape[1], img_cropped.shape[0]))\n\n        # smooth\n        cam = cv2.GaussianBlur(cam, (11,11), 0)\n\n        # sharpen\n        cam = np.power(cam, 2.5)\n        cam[cam < 0.3] = 0\n\n        # brain mask\n        mask = (img_cropped > img_cropped.mean() * 0.2).astype(np.float32)\n        cam = cam * mask\n\n        ax = fig.add_subplot(gs[row, 1])\n        ax.imshow(img_cropped, cmap='gray')\n        ax.imshow(cam, cmap='jet', alpha=0.30)\n        ax.axis('off')\n\n        # -------- TEXT --------\n        idh_p = int(idh_prob_c > 0.5)\n        mgmt_p = int(mgmt_prob_c > 0.5)\n\n        idh_text = \"Mutant\" if idh_p else \"Wildtype\"\n        mgmt_text = \"Methylated\" if mgmt_p else \"Unmethylated\"\n        grade_text = \"High\" if grade_p_c else \"Low\"\n\n        if mgmt[i].item() == -1:\n            mgmt_text = \"N/A\"\n\n        conf = f\"{idh_prob_c:.2f}\"\n\n        # correctness\n        idh_c = idh_p == idh[i].item()\n        mgmt_c = (mgmt[i].item() == -1) or (mgmt_p == mgmt[i].item())\n        grade_c = grade_p_c == grade[i].item()\n\n        colors = [\n            GREEN if idh_c else RED,\n            GRAY if mgmt[i].item() == -1 else (GREEN if mgmt_c else RED),\n            GREEN if grade_c else RED\n        ]\n\n        values = [idh_text, mgmt_text, grade_text, conf]\n\n        for col in range(4):\n            ax = fig.add_subplot(gs[row, col+2])\n            ax.text(0.5, 0.5, values[col],\n                    ha='center', va='center', fontsize=8)\n            ax.set_facecolor(colors[col] if col < 3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n            for sp in ax.spines.values():\n                sp.set_linewidth(0.8)\n\n        count += 1\n\n    if count >= cases_to_show:\n        break\n\n# =========================\n# FINAL TOUCH\n# =========================\nplt.suptitle(\"Model Predictions with Clean Grad-CAM\", fontsize=10, y=0.98)\n\nfig.text(0.5, 0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout(pad=1.2)\nplt.savefig(\"final_perfect_gradcam.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:30:23.469978Z","iopub.execute_input":"2026-04-19T21:30:23.470479Z","iopub.status.idle":"2026-04-19T21:30:34.220597Z","shell.execute_reply.started":"2026-04-19T21:30:23.470448Z","shell.execute_reply":"2026-04-19T21:30:34.219617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: CLEAN, STABLE GRAD-CAM++ PIPELINE\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS --------\nGREEN = \"#cfe8cf\"\nRED   = \"#f6c1c1\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# MULTI-CHANNEL CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img_norm = (img - img.min()) / (img.max() - img.min() + 1e-8)\n    mask = img_norm > threshold\n\n    coords = np.argwhere(mask)\n    if coords.shape[0] == 0:\n        return x\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - pad)\n    x0 = max(0, x0 - pad)\n    y1 = min(img.shape[0], y1 + pad)\n    x1 = min(img.shape[1], x1 + pad)\n\n    return x[:, y0:y1, x0:x1]\n\n# =========================\n# GRAD-CAM++ SETUP\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(module, inp, out):\n    features.clear()\n    features.append(out)\n\ndef bwd_hook(module, grad_in, grad_out):\n    gradients.clear()\n    gradients.append(grad_out[0])\n\ntarget_layer = model.backbone.layer4[-1]\nh1 = target_layer.register_forward_hook(fwd_hook)\nh2 = target_layer.register_backward_hook(bwd_hook)\n\ndef get_gradcam_pp(x_batch):\n    out_idh, out_mgmt, out_grade = model(x_batch)\n\n    # use Grade (better spatial signal)\n    score = out_grade[0, 1]\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]       # (1, C, H, W)\n    fmap = features[0]        # (1, C, H, W)\n\n    # Grad-CAM++ weights\n    grad_sq = grad ** 2\n    grad_cube = grad ** 3\n\n    eps = 1e-8\n    alpha = grad_sq / (2 * grad_sq + torch.sum(fmap * grad_cube, dim=(2,3), keepdim=True) + eps)\n    weights = torch.sum(alpha * F.relu(grad), dim=(2,3), keepdim=True)\n\n    cam = torch.sum(weights * fmap, dim=1).squeeze()\n\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(7.16, 5.6))\ngs = fig.add_gridspec(\n    cases_to_show + 1,\n    6,\n    height_ratios=[0.9] + [1.4]*cases_to_show,\n    wspace=0.15,\n    hspace=0.25\n)\n\nheaders = [\"Original\", \"Grad-CAM\", \"IDH\", \"MGMT\", \"Grade\", \"Conf\"]\nfor j, h in enumerate(headers):\n    ax = fig.add_subplot(gs[0, j])\n    ax.text(0.5, 0.5, h, ha='center', va='center',\n            fontsize=9, weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n    for sp in ax.spines.values():\n        sp.set_linewidth(0.8)\n\n# =========================\n# LOOP\n# =========================\ncount = 0\n\nfor batch in val_loader:\n    x = batch[0].to(device)\n    labels = batch[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count >= cases_to_show:\n            break\n\n        row = count + 1\n\n        # -------- ORIGINAL --------\n        img_orig = x[i][1].cpu().numpy()  # T1ce\n        ax = fig.add_subplot(gs[row, 0])\n        ax.imshow(img_orig, cmap='gray')\n        ax.axis('off')\n\n        # -------- CROP --------\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        # -------- MODEL --------\n        out_idh_c, out_mgmt_c, out_grade_c = model(x_crop_b)\n\n        idh_prob_c = torch.sigmoid(out_idh_c).item()\n        mgmt_prob_c = torch.sigmoid(out_mgmt_c).item()\n        grade_p_c = int(torch.argmax(out_grade_c, dim=1).item())\n\n        # -------- GRAD-CAM++ --------\n        cam = get_gradcam_pp(x_crop_b)\n\n        img_cropped = x_crop[1].cpu().numpy()  # T1ce\n\n        # normalize image\n        img_cropped = (img_cropped - img_cropped.min()) / \\\n                      (img_cropped.max() - img_cropped.min() + 1e-8)\n\n        cam = cv2.resize(cam, (img_cropped.shape[1], img_cropped.shape[0]))\n\n        # LIGHT smoothing (not overdone)\n        cam = cv2.GaussianBlur(cam, (7,7), 0)\n\n        # mild mask (avoid artifacts but keep signal)\n        mask = (img_cropped > np.percentile(img_cropped, 20)).astype(np.float32)\n        cam = cam * mask\n\n        ax = fig.add_subplot(gs[row, 1])\n        ax.imshow(img_cropped, cmap='gray')\n        ax.imshow(cam, cmap='jet', alpha=0.35)\n        ax.axis('off')\n\n        # -------- TEXT --------\n        idh_p = int(idh_prob_c > 0.5)\n        mgmt_p = int(mgmt_prob_c > 0.5)\n\n        idh_text = \"Mutant\" if idh_p else \"Wildtype\"\n        mgmt_text = \"Methylated\" if mgmt_p else \"Unmethylated\"\n        grade_text = \"High\" if grade_p_c else \"Low\"\n\n        if mgmt[i].item() == -1:\n            mgmt_text = \"N/A\"\n\n        conf = f\"{idh_prob_c:.2f}\"\n\n        idh_c = idh_p == idh[i].item()\n        mgmt_c = (mgmt[i].item() == -1) or (mgmt_p == mgmt[i].item())\n        grade_c = grade_p_c == grade[i].item()\n\n        colors = [\n            GREEN if idh_c else RED,\n            GRAY if mgmt[i].item() == -1 else (GREEN if mgmt_c else RED),\n            GREEN if grade_c else RED\n        ]\n\n        values = [idh_text, mgmt_text, grade_text, conf]\n\n        for col in range(4):\n            ax = fig.add_subplot(gs[row, col+2])\n            ax.text(0.5, 0.5, values[col],\n                    ha='center', va='center', fontsize=8)\n            ax.set_facecolor(colors[col] if col < 3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n            for sp in ax.spines.values():\n                sp.set_linewidth(0.8)\n\n        count += 1\n\n    if count >= cases_to_show:\n        break\n\n# =========================\n# FINAL\n# =========================\nplt.suptitle(\"Model Predictions with Grad-CAM++\", fontsize=10, y=0.98)\n\nfig.text(0.5, 0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout(pad=1.2)\nplt.savefig(\"final_gradcam_pp.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:33:08.487515Z","iopub.execute_input":"2026-04-19T21:33:08.488172Z","iopub.status.idle":"2026-04-19T21:33:18.986522Z","shell.execute_reply.started":"2026-04-19T21:33:08.488140Z","shell.execute_reply":"2026-04-19T21:33:18.985518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: CLEAN GRAD-CAM++ (NO NOISE, NO FAKE SHARPENING)\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS --------\nGREEN = \"#dff0d8\"\nRED   = \"#f2dede\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# MULTI-CHANNEL CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img_norm = (img - img.min()) / (img.max() - img.min() + 1e-8)\n    mask = img_norm > threshold\n\n    coords = np.argwhere(mask)\n    if coords.shape[0] == 0:\n        return x\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - pad)\n    x0 = max(0, x0 - pad)\n    y1 = min(img.shape[0], y1 + pad)\n    x1 = min(img.shape[1], x1 + pad)\n\n    return x[:, y0:y1, x0:x1]\n\n# =========================\n# GRAD-CAM++ SETUP\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(module, inp, out):\n    features.clear()\n    features.append(out)\n\ndef bwd_hook(module, grad_in, grad_out):\n    gradients.clear()\n    gradients.append(grad_out[0])\n\ntarget_layer = model.backbone.layer4[-1]\nh1 = target_layer.register_forward_hook(fwd_hook)\nh2 = target_layer.register_backward_hook(bwd_hook)\n\ndef get_gradcam_pp(x_batch):\n    out_idh, out_mgmt, out_grade = model(x_batch)\n\n    # use Grade head (better spatial signal)\n    score = out_grade[0, 1]\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    grad_sq = grad ** 2\n    grad_cube = grad ** 3\n\n    eps = 1e-8\n    alpha = grad_sq / (2 * grad_sq + torch.sum(fmap * grad_cube, dim=(2,3), keepdim=True) + eps)\n    weights = torch.sum(alpha * F.relu(grad), dim=(2,3), keepdim=True)\n\n    cam = torch.sum(weights * fmap, dim=1).squeeze()\n\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(7.16, 5.6))\ngs = fig.add_gridspec(\n    cases_to_show + 1,\n    6,\n    height_ratios=[0.9] + [1.4]*cases_to_show,\n    wspace=0.15,\n    hspace=0.25\n)\n\nheaders = [\"Original\", \"Grad-CAM\", \"IDH\", \"MGMT\", \"Grade\", \"Conf\"]\nfor j, h in enumerate(headers):\n    ax = fig.add_subplot(gs[0, j])\n    ax.text(0.5, 0.5, h, ha='center', va='center',\n            fontsize=9, weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n    for sp in ax.spines.values():\n        sp.set_linewidth(0.8)\n\n# =========================\n# LOOP\n# =========================\ncount = 0\n\nfor batch in val_loader:\n    x = batch[0].to(device)\n    labels = batch[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count >= cases_to_show:\n            break\n\n        row = count + 1\n\n        # -------- ORIGINAL --------\n        img_orig = x[i][1].cpu().numpy()  # T1ce\n        ax = fig.add_subplot(gs[row, 0])\n        ax.imshow(img_orig, cmap='gray')\n        ax.axis('off')\n\n        # -------- CROP --------\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        # -------- MODEL --------\n        out_idh_c, out_mgmt_c, out_grade_c = model(x_crop_b)\n\n        idh_prob_c = torch.sigmoid(out_idh_c).item()\n        mgmt_prob_c = torch.sigmoid(out_mgmt_c).item()\n        grade_p_c = int(torch.argmax(out_grade_c, dim=1).item())\n\n        # -------- GRAD-CAM++ --------\n        cam = get_gradcam_pp(x_crop_b)\n\n        img_cropped = x_crop[1].cpu().numpy()\n\n        # normalize image\n        img_cropped = (img_cropped - img_cropped.min()) / \\\n                      (img_cropped.max() - img_cropped.min() + 1e-8)\n\n        cam = cv2.resize(cam, (img_cropped.shape[1], img_cropped.shape[0]))\n\n        # ===== CLEAN PROCESSING =====\n        cam = cv2.GaussianBlur(cam, (15,15), 0)\n\n        # keep strongest region only (not aggressive)\n        th = np.percentile(cam, 85)\n        cam = np.where(cam >= th, cam, 0)\n\n        # normalize again\n        cam = cam / (cam.max() + 1e-8)\n\n        # brain mask (light)\n        mask = (img_cropped > np.percentile(img_cropped, 20)).astype(np.float32)\n        cam = cam * mask\n\n        # ===== PLOT =====\n        ax = fig.add_subplot(gs[row, 1])\n        ax.imshow(img_cropped, cmap='gray')\n        ax.imshow(cam, cmap='turbo', alpha=0.35)\n        ax.axis('off')\n\n        # -------- TEXT --------\n        idh_p = int(idh_prob_c > 0.5)\n        mgmt_p = int(mgmt_prob_c > 0.5)\n\n        idh_text = \"Mutant\" if idh_p else \"Wildtype\"\n        mgmt_text = \"Methylated\" if mgmt_p else \"Unmethylated\"\n        grade_text = \"High\" if grade_p_c else \"Low\"\n\n        if mgmt[i].item() == -1:\n            mgmt_text = \"N/A\"\n\n        conf = f\"{idh_prob_c:.2f}\"\n\n        idh_c = idh_p == idh[i].item()\n        mgmt_c = (mgmt[i].item() == -1) or (mgmt_p == mgmt[i].item())\n        grade_c = grade_p_c == grade[i].item()\n\n        colors = [\n            GREEN if idh_c else RED,\n            GRAY if mgmt[i].item() == -1 else (GREEN if mgmt_c else RED),\n            GREEN if grade_c else RED\n        ]\n\n        values = [idh_text, mgmt_text, grade_text, conf]\n\n        for col in range(4):\n            ax = fig.add_subplot(gs[row, col+2])\n            ax.text(0.5, 0.5, values[col],\n                    ha='center', va='center', fontsize=8)\n            ax.set_facecolor(colors[col] if col < 3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n            for sp in ax.spines.values():\n                sp.set_linewidth(0.8)\n\n        count += 1\n\n    if count >= cases_to_show:\n        break\n\n# =========================\n# FINAL\n# =========================\nplt.suptitle(\"Model Predictions with Grad-CAM++\", fontsize=10, y=0.98)\n\nfig.text(0.5, 0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout(pad=1.2)\nplt.savefig(\"final_clean_gradcam_pp.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:38:51.651080Z","iopub.execute_input":"2026-04-19T21:38:51.651916Z","iopub.status.idle":"2026-04-19T21:39:02.512770Z","shell.execute_reply.started":"2026-04-19T21:38:51.651879Z","shell.execute_reply":"2026-04-19T21:39:02.511771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: POLISHED GRAD-CAM++ FIGURE\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS (soft, paper-friendly) --------\nGREEN = \"#90EE90\"\nRED   = \"#FCBABA\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# MULTI-CHANNEL CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img_norm = (img - img.min()) / (img.max() - img.min() + 1e-8)\n    mask = img_norm > threshold\n\n    coords = np.argwhere(mask)\n    if coords.shape[0] == 0:\n        return x\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - pad)\n    x0 = max(0, x0 - pad)\n    y1 = min(img.shape[0], y1 + pad)\n    x1 = min(img.shape[1], x1 + pad)\n\n    return x[:, y0:y1, x0:x1]\n\n# =========================\n# GRAD-CAM++ SETUP\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(module, inp, out):\n    features.clear()\n    features.append(out)\n\ndef bwd_hook(module, grad_in, grad_out):\n    gradients.clear()\n    gradients.append(grad_out[0])\n\ntarget_layer = model.backbone.layer4[-1]\nh1 = target_layer.register_forward_hook(fwd_hook)\nh2 = target_layer.register_backward_hook(bwd_hook)\n\ndef get_gradcam_pp(x_batch):\n    out_idh, out_mgmt, out_grade = model(x_batch)\n\n    score = out_grade[0, 1]  # High-grade class\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    grad_sq = grad ** 2\n    grad_cube = grad ** 3\n\n    eps = 1e-8\n    alpha = grad_sq / (2 * grad_sq + torch.sum(fmap * grad_cube, dim=(2,3), keepdim=True) + eps)\n    weights = torch.sum(alpha * F.relu(grad), dim=(2,3), keepdim=True)\n\n    cam = torch.sum(weights * fmap, dim=1).squeeze()\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(7.16, 5.6))\ngs = fig.add_gridspec(\n    cases_to_show + 1,\n    6,\n    height_ratios=[0.9] + [1.4]*cases_to_show,\n    wspace=0.15,\n    hspace=0.25\n)\n\nheaders = [\"Original\", \"Grad-CAM\", \"IDH\", \"MGMT\", \"Grade\", \"Conf\"]\nfor j, h in enumerate(headers):\n    ax = fig.add_subplot(gs[0, j])\n    ax.text(0.5, 0.5, h, ha='center', va='center',\n            fontsize=9, weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n    for sp in ax.spines.values():\n        sp.set_linewidth(0.8)\n\n# =========================\n# LOOP\n# =========================\ncount = 0\n\nfor batch in val_loader:\n    x = batch[0].to(device)\n    labels = batch[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count >= cases_to_show:\n            break\n\n        row = count + 1\n\n        # -------- ORIGINAL IMAGE --------\n        img_orig = x[i][1].cpu().numpy()\n        img_orig = (img_orig - img_orig.min()) / (img_orig.max() - img_orig.min() + 1e-8)\n\n        ax = fig.add_subplot(gs[row, 0])\n        ax.imshow(img_orig, cmap='gray')\n        ax.axis('off')\n        for sp in ax.spines.values():\n            sp.set_visible(True)\n            sp.set_linewidth(0.5)\n\n        # -------- CROP --------\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        # -------- MODEL --------\n        out_idh_c, out_mgmt_c, out_grade_c = model(x_crop_b)\n\n        idh_prob_c = torch.sigmoid(out_idh_c).item()\n        mgmt_prob_c = torch.sigmoid(out_mgmt_c).item()\n        grade_p_c = int(torch.argmax(out_grade_c, dim=1).item())\n\n        # -------- GRAD-CAM++ --------\n        cam = get_gradcam_pp(x_crop_b)\n\n        img_cropped = x_crop[1].cpu().numpy()\n        img_cropped = (img_cropped - img_cropped.min()) / (img_cropped.max() - img_cropped.min() + 1e-8)\n\n        cam = cv2.resize(cam, (img_cropped.shape[1], img_cropped.shape[0]))\n\n        # smoothing + focus\n        cam = cv2.GaussianBlur(cam, (15,15), 0)\n        th = np.percentile(cam, 85)\n        cam = np.where(cam >= th, cam, 0)\n        cam = cam / (cam.max() + 1e-8)\n\n        # light brain mask\n        mask = (img_cropped > np.percentile(img_cropped, 20)).astype(np.float32)\n        cam = cam * mask\n\n        # -------- PLOT CAM --------\n        ax = fig.add_subplot(gs[row, 1])\n        ax.imshow(img_cropped, cmap='gray')\n        ax.imshow(cam, cmap='turbo', alpha=0.45)\n        ax.axis('off')\n        for sp in ax.spines.values():\n            sp.set_visible(True)\n            sp.set_linewidth(0.5)\n\n        # -------- TEXT --------\n        idh_p = int(idh_prob_c > 0.5)\n        mgmt_p = int(mgmt_prob_c > 0.5)\n\n        idh_text = \"Mutant\" if idh_p else \"Wildtype\"\n        mgmt_text = \"Methylated\" if mgmt_p else \"Unmethylated\"\n        grade_text = \"High\" if grade_p_c else \"Low\"\n\n        if mgmt[i].item() == -1:\n            mgmt_text = \"N/A\"\n\n        conf = f\"{idh_prob_c:.2f}\"\n\n        idh_c = idh_p == idh[i].item()\n        mgmt_c = (mgmt[i].item() == -1) or (mgmt_p == mgmt[i].item())\n        grade_c = grade_p_c == grade[i].item()\n\n        colors = [\n            GREEN if idh_c else RED,\n            GRAY if mgmt[i].item() == -1 else (GREEN if mgmt_c else RED),\n            GREEN if grade_c else RED\n        ]\n\n        values = [idh_text, mgmt_text, grade_text, conf]\n\n        for col in range(4):\n            ax = fig.add_subplot(gs[row, col+2])\n            ax.text(0.5, 0.5, values[col],\n                    ha='center', va='center', fontsize=8)\n            ax.set_facecolor(colors[col] if col < 3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n            for sp in ax.spines.values():\n                sp.set_linewidth(0.8)\n\n        count += 1\n\n    if count >= cases_to_show:\n        break\n\n# =========================\n# FINAL\n# =========================\nplt.suptitle(\"Model Predictions with Grad-CAM++\", fontsize=10, y=0.98)\n\nfig.text(0.5, 0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout(pad=1.2)\nplt.savefig(\"final_polished_gradcam.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:46:15.805656Z","iopub.execute_input":"2026-04-19T21:46:15.805992Z","iopub.status.idle":"2026-04-19T21:46:26.206147Z","shell.execute_reply.started":"2026-04-19T21:46:15.805965Z","shell.execute_reply":"2026-04-19T21:46:26.205378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: CONNECTED-REGION GRAD-CAM++ + CLAHE\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS --------\nGREEN = \"#90EE90\"\nRED   = \"#FCBABA\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img_norm = (img - img.min()) / (img.max() - img.min() + 1e-8)\n    mask = img_norm > threshold\n\n    coords = np.argwhere(mask)\n    if coords.shape[0] == 0:\n        return x\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - pad)\n    x0 = max(0, x0 - pad)\n    y1 = min(img.shape[0], y1 + pad)\n    x1 = min(img.shape[1], x1 + pad)\n\n    return x[:, y0:y1, x0:x1]\n\n# =========================\n# GRAD-CAM++\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(module, inp, out):\n    features.clear()\n    features.append(out)\n\ndef bwd_hook(module, grad_in, grad_out):\n    gradients.clear()\n    gradients.append(grad_out[0])\n\ntarget_layer = model.backbone.layer4[-1]\nh1 = target_layer.register_forward_hook(fwd_hook)\nh2 = target_layer.register_backward_hook(bwd_hook)\n\ndef get_gradcam_pp(x):\n    _, _, out_grade = model(x)\n    score = out_grade[0, 1]\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    grad_sq = grad ** 2\n    grad_cube = grad ** 3\n\n    eps = 1e-8\n    alpha = grad_sq / (2 * grad_sq + torch.sum(fmap * grad_cube, dim=(2,3), keepdim=True) + eps)\n    weights = torch.sum(alpha * F.relu(grad), dim=(2,3), keepdim=True)\n\n    cam = torch.sum(weights * fmap, dim=1).squeeze()\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# KEEP LARGEST REGION\n# =========================\ndef keep_largest_region(cam):\n    binary = (cam > 0).astype(np.uint8)\n    num_labels, labels = cv2.connectedComponents(binary)\n\n    if num_labels <= 1:\n        return cam\n\n    max_area = 0\n    best_label = 1\n\n    for l in range(1, num_labels):\n        area = np.sum(labels == l)\n        if area > max_area:\n            max_area = area\n            best_label = l\n\n    mask = (labels == best_label).astype(np.float32)\n    return cam * mask\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(7.16, 5.6))\ngs = fig.add_gridspec(cases_to_show + 1, 6,\n                      height_ratios=[0.9] + [1.4]*cases_to_show,\n                      wspace=0.15, hspace=0.25)\n\nheaders = [\"Original\", \"Grad-CAM\", \"IDH\", \"MGMT\", \"Grade\", \"Conf\"]\nfor j, h in enumerate(headers):\n    ax = fig.add_subplot(gs[0, j])\n    ax.text(0.5, 0.5, h, ha='center', va='center', fontsize=9, weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n    for sp in ax.spines.values():\n        sp.set_linewidth(0.8)\n\n# =========================\n# LOOP\n# =========================\ncount = 0\n\nfor batch in val_loader:\n    x = batch[0].to(device)\n    labels = batch[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count >= cases_to_show:\n            break\n\n        row = count + 1\n\n        # -------- ORIGINAL (T1ce + CLAHE) --------\n        img_orig = x[i][1].cpu().numpy()\n        img_orig = (img_orig - img_orig.min()) / (img_orig.max() - img_orig.min() + 1e-8)\n        img_uint = (img_orig * 255).astype(np.uint8)\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n        img_orig = clahe.apply(img_uint) / 255.0\n\n        ax = fig.add_subplot(gs[row, 0])\n        ax.imshow(img_orig, cmap='gray')\n        ax.axis('off')\n\n        # -------- CROP --------\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        # -------- MODEL --------\n        out_idh_c, out_mgmt_c, out_grade_c = model(x_crop_b)\n\n        idh_prob_c = torch.sigmoid(out_idh_c).item()\n        mgmt_prob_c = torch.sigmoid(out_mgmt_c).item()\n        grade_p_c = int(torch.argmax(out_grade_c, dim=1).item())\n\n        # -------- CAM --------\n        cam = get_gradcam_pp(x_crop_b)\n\n        img_cropped = x_crop[1].cpu().numpy()\n        img_cropped = (img_cropped - img_cropped.min()) / (img_cropped.max() - img_cropped.min() + 1e-8)\n\n        cam = cv2.resize(cam, (img_cropped.shape[1], img_cropped.shape[0]))\n        cam = cv2.GaussianBlur(cam, (15,15), 0)\n\n        # threshold + largest region\n        th = np.percentile(cam, 85)\n        cam = np.where(cam >= th, cam, 0)\n        cam = keep_largest_region(cam)\n\n        cam = cam / (cam.max() + 1e-8)\n\n        # -------- PLOT CAM --------\n        ax = fig.add_subplot(gs[row, 1])\n        ax.imshow(img_cropped, cmap='gray')\n        ax.imshow(cam, cmap='turbo', alpha=0.45)\n        ax.axis('off')\n\n        # -------- TEXT --------\n        idh_p = int(idh_prob_c > 0.5)\n        mgmt_p = int(mgmt_prob_c > 0.5)\n\n        idh_text = \"Mutant\" if idh_p else \"Wildtype\"\n        mgmt_text = \"Methylated\" if mgmt_p else \"Unmethylated\"\n        grade_text = \"High\" if grade_p_c else \"Low\"\n\n        if mgmt[i].item() == -1:\n            mgmt_text = \"N/A\"\n\n        conf = f\"{idh_prob_c:.2f}\"\n\n        idh_c = idh_p == idh[i].item()\n        mgmt_c = (mgmt[i].item() == -1) or (mgmt_p == mgmt[i].item())\n        grade_c = grade_p_c == grade[i].item()\n\n        colors = [\n            GREEN if idh_c else RED,\n            GRAY if mgmt[i].item() == -1 else (GREEN if mgmt_c else RED),\n            GREEN if grade_c else RED\n        ]\n\n        values = [idh_text, mgmt_text, grade_text, conf]\n\n        for col in range(4):\n            ax = fig.add_subplot(gs[row, col+2])\n            ax.text(0.5, 0.5, values[col], ha='center', va='center', fontsize=8)\n            ax.set_facecolor(colors[col] if col < 3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n            for sp in ax.spines.values():\n                sp.set_linewidth(0.8)\n\n        count += 1\n\n    if count >= cases_to_show:\n        break\n\n# =========================\n# FINAL\n# =========================\nplt.suptitle(\"Model Predictions with Grad-CAM++\", fontsize=10, y=0.98)\n\nfig.text(0.5, 0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout(pad=1.2)\nplt.savefig(\"final_best_gradcam.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:47:43.131193Z","iopub.execute_input":"2026-04-19T21:47:43.131726Z","iopub.status.idle":"2026-04-19T21:47:53.482332Z","shell.execute_reply.started":"2026-04-19T21:47:43.131692Z","shell.execute_reply":"2026-04-19T21:47:53.481619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: MULTI-HEAD GRAD-CAM++\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\nGREEN = \"#90EE90\"\nRED   = \"#FCBABA\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img = (img - img.min())/(img.max()-img.min()+1e-8)\n    mask = img > threshold\n\n    coords = np.argwhere(mask)\n    if len(coords)==0:\n        return x\n\n    y0,x0 = coords.min(0)\n    y1,x1 = coords.max(0)+1\n\n    y0=max(0,y0-pad); x0=max(0,x0-pad)\n    y1=min(img.shape[0],y1+pad); x1=min(img.shape[1],x1+pad)\n\n    return x[:,y0:y1,x0:x1]\n\n# =========================\n# HOOKS\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(m,i,o):\n    features.clear(); features.append(o)\n\ndef bwd_hook(m,gi,go):\n    gradients.clear(); gradients.append(go[0])\n\nlayer = model.backbone.layer4[-1]\nh1 = layer.register_forward_hook(fwd_hook)\nh2 = layer.register_backward_hook(bwd_hook)\n\n# =========================\n# GENERIC GRADCAM++\n# =========================\ndef gradcam_pp(x, head=\"grade\"):\n    out_idh, out_mgmt, out_grade = model(x)\n\n    if head==\"idh\":\n        score = torch.sigmoid(out_idh).view(-1)[0]\n    elif head==\"mgmt\":\n        score = torch.sigmoid(out_mgmt).view(-1)[0]\n    else:\n        score = out_grade[0,1]\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    g2 = grad**2\n    g3 = grad**3\n\n    alpha = g2 / (2*g2 + torch.sum(fmap*g3, dim=(2,3), keepdim=True)+1e-8)\n    w = torch.sum(alpha*F.relu(grad), dim=(2,3), keepdim=True)\n\n    cam = torch.sum(w*fmap, dim=1).squeeze()\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam/(cam.max()+1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# CLEAN CAM\n# =========================\ndef clean_cam(cam, img):\n    cam = cv2.resize(cam, (img.shape[1], img.shape[0]))\n    cam = cv2.GaussianBlur(cam, (15,15), 0)\n\n    th = np.percentile(cam, 85)\n    cam = np.where(cam>=th, cam, 0)\n\n    # largest region\n    binary = (cam>0).astype(np.uint8)\n    num, labels = cv2.connectedComponents(binary)\n\n    if num>1:\n        areas = [(labels==i).sum() for i in range(1,num)]\n        best = np.argmax(areas)+1\n        cam = cam*(labels==best)\n\n    cam = cam/(cam.max()+1e-8)\n\n    mask = (img > np.percentile(img,20)).astype(np.float32)\n    return cam*mask\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(8,5.6))\ngs = fig.add_gridspec(cases_to_show+1, 8,\n                      height_ratios=[0.9]+[1.4]*cases_to_show,\n                      wspace=0.12, hspace=0.25)\n\nheaders = [\"Original\",\"IDH CAM\",\"MGMT CAM\",\"Grade CAM\",\"IDH\",\"MGMT\",\"Grade\",\"Conf\"]\n\nfor j,h in enumerate(headers):\n    ax = fig.add_subplot(gs[0,j])\n    ax.text(0.5,0.5,h,ha='center',va='center',fontsize=8,weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n\n# =========================\n# LOOP\n# =========================\ncount=0\n\nfor batch in val_loader:\n    x=batch[0].to(device)\n    labels=batch[1]\n\n    idh=labels[\"idh\"]; mgmt=labels[\"mgmt\"]; grade=labels[\"grade\"]\n\n    for i in range(x.shape[0]):\n\n        if count>=cases_to_show: break\n        row=count+1\n\n        img_orig = x[i][1].cpu().numpy()\n        img_orig=(img_orig-img_orig.min())/(img_orig.max()-img_orig.min()+1e-8)\n\n        ax=fig.add_subplot(gs[row,0])\n        ax.imshow(img_orig,cmap='gray'); ax.axis('off')\n\n        x_crop = crop_brain_multi(x[i])\n        x_crop_b = x_crop.unsqueeze(0).to(device)\n\n        img_c = x_crop[1].cpu().numpy()\n        img_c=(img_c-img_c.min())/(img_c.max()-img_c.min()+1e-8)\n\n        # --- CAMS ---\n        cam_idh = clean_cam(gradcam_pp(x_crop_b,\"idh\"), img_c)\n        cam_mgmt = clean_cam(gradcam_pp(x_crop_b,\"mgmt\"), img_c)\n        cam_grade = clean_cam(gradcam_pp(x_crop_b,\"grade\"), img_c)\n\n        for k,cam in enumerate([cam_idh, cam_mgmt, cam_grade]):\n            ax=fig.add_subplot(gs[row,k+1])\n            ax.imshow(img_c,cmap='gray')\n            ax.imshow(cam,cmap='turbo',alpha=0.45)\n            ax.axis('off')\n\n        # predictions\n        out_idh,out_mgmt,out_grade=model(x_crop_b)\n\n        idh_p=int(torch.sigmoid(out_idh)>0.5)\n        mgmt_p=int(torch.sigmoid(out_mgmt)>0.5)\n        grade_p=int(torch.argmax(out_grade))\n\n        texts=[\n            \"Mutant\" if idh_p else \"Wildtype\",\n            \"Methylated\" if mgmt_p else \"Unmethylated\",\n            \"High\" if grade_p else \"Low\",\n            f\"{torch.sigmoid(out_idh).item():.2f}\"\n        ]\n\n        correct=[\n            idh_p==idh[i].item(),\n            (mgmt[i].item()==-1) or mgmt_p==mgmt[i].item(),\n            grade_p==grade[i].item()\n        ]\n\n        colors=[\n            GREEN if correct[0] else RED,\n            GRAY if mgmt[i].item()==-1 else (GREEN if correct[1] else RED),\n            GREEN if correct[2] else RED\n        ]\n\n        for j in range(4):\n            ax=fig.add_subplot(gs[row,j+4])\n            ax.text(0.5,0.5,texts[j],ha='center',va='center',fontsize=8)\n            ax.set_facecolor(colors[j] if j<3 else WHITE)\n            ax.set_xticks([]); ax.set_yticks([])\n\n        count+=1\n\n    if count>=cases_to_show: break\n\n# =========================\n# FINAL\n# =========================\nplt.suptitle(\"Multi-Task Grad-CAM++ Interpretation\", fontsize=10, y=0.98)\n\nfig.text(0.5,0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout()\nplt.savefig(\"final_multi_head_gradcam.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T21:51:43.437010Z","iopub.execute_input":"2026-04-19T21:51:43.437427Z","iopub.status.idle":"2026-04-19T21:51:54.811065Z","shell.execute_reply.started":"2026-04-19T21:51:43.437394Z","shell.execute_reply":"2026-04-19T21:51:54.809996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: RANDOM MULTI-HEAD GRAD-CAM++\n# =========================\n\nimport torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\nimport random\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models_dict[\"resnet50_v2\"].to(device)\nmodel.eval()\n\n# -------- COLORS --------\nGREEN = \"#90EE90\"\nRED   = \"#FCBABA\"\nGRAY  = \"#e6e6e6\"\nWHITE = \"#fafafa\"\n\ncases_to_show = 4\n\n# =========================\n# CROP\n# =========================\ndef crop_brain_multi(x, threshold=0.05, pad=5):\n    img = x[0].cpu().numpy()\n    img = (img - img.min())/(img.max()-img.min()+1e-8)\n    mask = img > threshold\n\n    coords = np.argwhere(mask)\n    if len(coords)==0:\n        return x\n\n    y0,x0 = coords.min(0)\n    y1,x1 = coords.max(0)+1\n\n    y0=max(0,y0-pad); x0=max(0,x0-pad)\n    y1=min(img.shape[0],y1+pad); x1=min(img.shape[1],x1+pad)\n\n    return x[:,y0:y1,x0:x1]\n\n# =========================\n# HOOKS\n# =========================\nfeatures, gradients = [], []\n\ndef fwd_hook(m,i,o):\n    features.clear(); features.append(o)\n\ndef bwd_hook(m,gi,go):\n    gradients.clear(); gradients.append(go[0])\n\nlayer = model.backbone.layer4[-1]\nh1 = layer.register_forward_hook(fwd_hook)\nh2 = layer.register_backward_hook(bwd_hook)\n\n# =========================\n# GRAD-CAM++\n# =========================\ndef gradcam_pp(x, head=\"grade\"):\n    out_idh, out_mgmt, out_grade = model(x)\n\n    if head==\"idh\":\n        score = torch.sigmoid(out_idh).view(-1)[0]\n    elif head==\"mgmt\":\n        score = torch.sigmoid(out_mgmt).view(-1)[0]\n    else:\n        score = out_grade[0,1]\n\n    model.zero_grad()\n    score.backward(retain_graph=True)\n\n    grad = gradients[0]\n    fmap = features[0]\n\n    g2 = grad**2\n    g3 = grad**3\n\n    alpha = g2 / (2*g2 + torch.sum(fmap*g3, dim=(2,3), keepdim=True)+1e-8)\n    w = torch.sum(alpha*F.relu(grad), dim=(2,3), keepdim=True)\n\n    cam = torch.sum(w*fmap, dim=1).squeeze()\n    cam = F.relu(cam)\n    cam = cam - cam.min()\n    cam = cam/(cam.max()+1e-8)\n\n    return cam.detach().cpu().numpy()\n\n# =========================\n# CLEAN CAM\n# =========================\ndef clean_cam(cam, img):\n    cam = cv2.resize(cam, (img.shape[1], img.shape[0]))\n    cam = cv2.GaussianBlur(cam, (15,15), 0)\n\n    th = np.percentile(cam, 85)\n    cam = np.where(cam>=th, cam, 0)\n\n    # largest connected region\n    binary = (cam>0).astype(np.uint8)\n    num, labels = cv2.connectedComponents(binary)\n\n    if num > 1:\n        areas = [(labels==i).sum() for i in range(1,num)]\n        best = np.argmax(areas)+1\n        cam = cam*(labels==best)\n\n    cam = cam/(cam.max()+1e-8)\n\n    mask = (img > np.percentile(img,20)).astype(np.float32)\n    return cam * mask\n\n# =========================\n# FIGURE\n# =========================\nfig = plt.figure(figsize=(8,5.6))\ngs = fig.add_gridspec(cases_to_show+1, 8,\n                      height_ratios=[0.9]+[1.4]*cases_to_show,\n                      wspace=0.12, hspace=0.25)\n\nheaders = [\"Original\",\"IDH CAM\",\"MGMT CAM\",\"Grade CAM\",\"IDH\",\"MGMT\",\"Grade\",\"Conf\"]\n\nfor j,h in enumerate(headers):\n    ax = fig.add_subplot(gs[0,j])\n    ax.text(0.5,0.5,h,ha='center',va='center',fontsize=8,weight='semibold')\n    ax.set_xticks([]); ax.set_yticks([])\n    ax.set_facecolor(\"#f5f5f5\")\n\n# =========================\n# RANDOM SAMPLING\n# =========================\nrandom.seed(42)\nindices = random.sample(range(len(val_loader.dataset)), cases_to_show)\n\n# =========================\n# LOOP\n# =========================\nfor count, idx in enumerate(indices):\n\n    row = count + 1\n\n    data = val_loader.dataset[idx]\n    x = data[0].unsqueeze(0).to(device)\n    labels = data[1]\n\n    idh = labels[\"idh\"]\n    mgmt = labels[\"mgmt\"]\n    grade = labels[\"grade\"]\n\n    # -------- ORIGINAL --------\n    img_orig = x[0][1].cpu().numpy()\n    img_orig = (img_orig-img_orig.min())/(img_orig.max()-img_orig.min()+1e-8)\n\n    ax = fig.add_subplot(gs[row,0])\n    ax.imshow(img_orig,cmap='gray')\n    ax.axis('off')\n\n    # -------- CROP --------\n    x_crop = crop_brain_multi(x[0])\n    x_crop_b = x_crop.unsqueeze(0).to(device)\n\n    img_c = x_crop[1].cpu().numpy()\n    img_c = (img_c-img_c.min())/(img_c.max()-img_c.min()+1e-8)\n\n    # -------- CAMS --------\n    cam_idh = clean_cam(gradcam_pp(x_crop_b,\"idh\"), img_c)\n    cam_mgmt = clean_cam(gradcam_pp(x_crop_b,\"mgmt\"), img_c)\n    cam_grade = clean_cam(gradcam_pp(x_crop_b,\"grade\"), img_c)\n\n    for k, cam in enumerate([cam_idh, cam_mgmt, cam_grade]):\n        ax = fig.add_subplot(gs[row, k+1])\n        ax.imshow(img_c, cmap='gray')\n        ax.imshow(cam, cmap='turbo', alpha=0.45)\n        ax.axis('off')\n\n    # -------- PREDICTIONS --------\n    out_idh, out_mgmt, out_grade = model(x_crop_b)\n\n    idh_p = int(torch.sigmoid(out_idh)>0.5)\n    mgmt_p = int(torch.sigmoid(out_mgmt)>0.5)\n    grade_p = int(torch.argmax(out_grade))\n\n    texts = [\n        \"Mutant\" if idh_p else \"Wildtype\",\n        \"Methylated\" if mgmt_p else \"Unmethylated\",\n        \"High\" if grade_p else \"Low\",\n        f\"{torch.sigmoid(out_idh).item():.2f}\"\n    ]\n\n    correct = [\n        idh_p == idh,\n        (mgmt==-1) or (mgmt_p==mgmt),\n        grade_p == grade\n    ]\n\n    colors = [\n        GREEN if correct[0] else RED,\n        GRAY if mgmt==-1 else (GREEN if correct[1] else RED),\n        GREEN if correct[2] else RED\n    ]\n\n    for j in range(4):\n        ax = fig.add_subplot(gs[row, j+4])\n        ax.text(0.5,0.5,texts[j],ha='center',va='center',fontsize=8)\n        ax.set_facecolor(colors[j] if j<3 else WHITE)\n        ax.set_xticks([]); ax.set_yticks([])\n\n# =========================\n# FINAL\n# =========================\nplt.suptitle(\"Multi-Task Grad-CAM++ Interpretation\", fontsize=10, y=0.98)\n\nfig.text(0.5,0.01,\n         \"Green: Correct | Red: Incorrect | Gray: Not Available\",\n         ha='center', fontsize=7)\n\nplt.tight_layout()\nplt.savefig(\"final_random_multi_head_gradcam.png\", dpi=300, bbox_inches='tight')\nplt.show()\n\nh1.remove(); h2.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:02:12.613206Z","iopub.execute_input":"2026-04-19T22:02:12.614080Z","iopub.status.idle":"2026-04-19T22:02:17.744946Z","shell.execute_reply.started":"2026-04-19T22:02:12.614039Z","shell.execute_reply":"2026-04-19T22:02:17.744081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL CONFUSION MATRICES (GRADE + MGMT)\n# =========================\n\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\nmodel = models_dict[\"resnet50_v2\"]\nmodel.eval()\n\ny_true_grade, y_pred_grade = [], []\ny_true_mgmt, y_pred_mgmt = [], []\n\nwith torch.no_grad():\n    for batch in val_loader:\n\n        x = batch[0].to(device)\n        labels = batch[1]\n\n        grade = labels[\"grade\"]\n        mgmt = labels[\"mgmt\"]\n\n        out_idh, out_mgmt, out_grade = model(x)\n\n        grade_preds = torch.argmax(out_grade, dim=1).cpu().numpy()\n        mgmt_preds = (torch.sigmoid(out_mgmt) > 0.5).int().cpu().numpy()\n\n        # -------- STORE GRADE --------\n        y_true_grade.extend(grade.numpy())\n        y_pred_grade.extend(grade_preds)\n\n        # -------- STORE MGMT (skip -1) --------\n        for i in range(len(mgmt)):\n            if mgmt[i] != -1:\n                y_true_mgmt.append(int(mgmt[i]))\n                y_pred_mgmt.append(int(mgmt_preds[i]))\n\n\n# =========================\n# FUNCTION TO PLOT MATRIX\n# =========================\ndef plot_cm(y_true, y_pred, title, labels):\n\n    cm = confusion_matrix(y_true, y_pred)\n    cm_percent = cm / cm.sum(axis=1, keepdims=True)\n\n    annot = np.array([\n        [f\"{cm[i,j]}\\n({cm_percent[i,j]*100:.1f}%)\" for j in range(2)]\n        for i in range(2)\n    ])\n\n    # metrics\n    tp = cm[1,1]\n    tn = cm[0,0]\n    fp = cm[0,1]\n    fn = cm[1,0]\n\n    acc = (tp + tn) / cm.sum()\n    sens = tp / (tp + fn + 1e-8)\n    spec = tn / (tn + fp + 1e-8)\n\n    plt.figure(figsize=(5,4))\n\n    sns.heatmap(\n        cm,\n        annot=annot,\n        fmt=\"\",\n        cmap=\"Blues\",\n        alpha=0.9,\n        xticklabels=labels,\n        yticklabels=labels\n    )\n\n    plt.xlabel(\"Predicted Label\")\n    plt.ylabel(\"True Label\")\n\n    plt.title(\n        f\"{title}\\nAcc={acc:.2f}, Sens={sens:.2f}, Spec={spec:.2f}\"\n    )\n\n    plt.tight_layout()\n    plt.show()\n\n\n# =========================\n# PLOT BOTH\n# =========================\n\n# Grade\nplot_cm(\n    y_true_grade,\n    y_pred_grade,\n    \"Confusion Matrix for Tumor Grade (ResNet-50)\",\n    [\"Low Grade\", \"High Grade\"]\n)\n\n# MGMT\nplot_cm(\n    y_true_mgmt,\n    y_pred_mgmt,\n    \"Confusion Matrix for MGMT Methylation (ResNet-50)\",\n    [\"Unmethylated\", \"Methylated\"]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:03:37.651512Z","iopub.execute_input":"2026-04-19T22:03:37.651909Z","iopub.status.idle":"2026-04-19T22:05:18.780182Z","shell.execute_reply.started":"2026-04-19T22:03:37.651881Z","shell.execute_reply":"2026-04-19T22:05:18.779350Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL CONFUSION MATRIX (IDH)\n# =========================\n\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\n\nmodel = models_dict[\"resnet50_v2\"]\nmodel.eval()\n\ny_true, y_pred = [], []\n\nwith torch.no_grad():\n    for batch in val_loader:\n\n        x = batch[0].to(device)\n        labels = batch[1]\n\n        idh = labels[\"idh\"]\n\n        out_idh, _, _ = model(x)\n\n        preds = (torch.sigmoid(out_idh) > 0.5).int().cpu().numpy()\n\n        y_pred.extend(preds)\n        y_true.extend(idh.numpy())\n\n# =========================\n# CONFUSION MATRIX\n# =========================\n\ncm = confusion_matrix(y_true, y_pred)\n\n# Normalize\ncm_percent = cm / cm.sum(axis=1, keepdims=True)\n\n# Annotate (count + %)\nannot = np.array([\n    [f\"{cm[i,j]}\\n({cm_percent[i,j]*100:.1f}%)\" for j in range(2)]\n    for i in range(2)\n])\n\n# =========================\n# METRICS\n# =========================\n\ntp = cm[1,1]\ntn = cm[0,0]\nfp = cm[0,1]\nfn = cm[1,0]\n\naccuracy = (tp + tn) / cm.sum()\nsensitivity = tp / (tp + fn + 1e-8)\nspecificity = tn / (tn + fp + 1e-8)\n\n# =========================\n# PLOT\n# =========================\n\nplt.figure(figsize=(5,4))\n\nsns.heatmap(\n    cm,\n    annot=annot,\n    fmt=\"\",\n    cmap=\"Blues\",\n    alpha=0.9,\n    xticklabels=[\"Wildtype\", \"Mutant\"],\n    yticklabels=[\"Wildtype\", \"Mutant\"]\n)\n\nplt.xlabel(\"Predicted Label\")\nplt.ylabel(\"True Label\")\n\nplt.title(\n    f\"Confusion Matrix for IDH Mutation (ResNet-50)\\n\"\n    f\"Acc={accuracy:.2f}, Sens={sensitivity:.2f}, Spec={specificity:.2f}\"\n)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:10:49.873712Z","iopub.execute_input":"2026-04-19T22:10:49.874326Z","iopub.status.idle":"2026-04-19T22:12:41.538274Z","shell.execute_reply.started":"2026-04-19T22:10:49.874293Z","shell.execute_reply":"2026-04-19T22:12:41.537324Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# CALIBRATION CURVE (IDH) + BRIER SCORE\n# =========================\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.calibration import calibration_curve\nfrom sklearn.metrics import brier_score_loss\n\nmodel = models_dict[\"resnet50_v2\"]\n\ny_true, y_prob = [], []\n\nmodel.eval()\n\nwith torch.no_grad():\n    for batch in val_loader:\n\n        x = batch[0].to(device)\n        labels = batch[1]\n\n        idh = labels[\"idh\"]\n\n        out_idh, _, _ = model(x)\n        prob = torch.sigmoid(out_idh).squeeze().cpu().numpy()\n\n        y_prob.extend(prob)\n        y_true.extend(idh.numpy())\n\n# =========================\n# CALIBRATION\n# =========================\n\nprob_true, prob_pred = calibration_curve(\n    y_true,\n    y_prob,\n    n_bins=10\n)\n\nbrier = brier_score_loss(y_true, y_prob)\n\n# =========================\n# PLOT\n# =========================\n\nplt.figure(figsize=(6,5))\n\nplt.plot(prob_pred, prob_true, marker='o', linewidth=2, label=\"ResNet-50\")\n\n# perfect calibration\nplt.plot([0,1], [0,1], linestyle='--', color='gray', label=\"Perfect Calibration\")\n\nplt.xlabel(\"Mean Predicted Probability\")\nplt.ylabel(\"Fraction of Positives\")\nplt.title(\"Calibration Curve (IDH Prediction)\")\n\n# 🔥 BRIER SCORE INSIDE FIGURE\nplt.text(\n    0.65, 0.15,\n    f\"Brier Score = {brier:.3f}\",\n    fontsize=10,\n    bbox=dict(facecolor='white', alpha=0.8, edgecolor='gray')\n)\n\nplt.legend(frameon=False)\nplt.grid(alpha=0.2)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:15:01.913322Z","iopub.execute_input":"2026-04-19T22:15:01.913715Z","iopub.status.idle":"2026-04-19T22:16:47.447681Z","shell.execute_reply.started":"2026-04-19T22:15:01.913682Z","shell.execute_reply":"2026-04-19T22:16:47.446792Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: CALIBRATION CURVE (IDH) + BRIER + SMOOTHING\n# =========================\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.calibration import calibration_curve\nfrom sklearn.metrics import brier_score_loss\nfrom scipy.interpolate import make_interp_spline\n\nmodel = models_dict[\"resnet50_v2\"]\n\ny_true, y_prob = [], []\n\nmodel.eval()\n\nwith torch.no_grad():\n    for batch in val_loader:\n\n        x = batch[0].to(device)\n        labels = batch[1]\n\n        idh = labels[\"idh\"]\n\n        out_idh, _, _ = model(x)\n        prob = torch.sigmoid(out_idh).squeeze().cpu().numpy()\n\n        y_prob.extend(prob)\n        y_true.extend(idh.numpy())\n\n# =========================\n# CALIBRATION\n# =========================\n\nprob_true, prob_pred = calibration_curve(\n    y_true,\n    y_prob,\n    n_bins=15,\n    strategy=\"quantile\"   # 🔥 important fix\n)\n\nbrier = brier_score_loss(y_true, y_prob)\n\n# =========================\n# SMOOTH CURVE\n# =========================\n\n# Ensure sorted (important for spline)\norder = np.argsort(prob_pred)\nprob_pred = prob_pred[order]\nprob_true = prob_true[order]\n\nx_new = np.linspace(0, 1, 100)\n\n# Handle small bins safely\nif len(prob_pred) >= 3:\n    spl = make_interp_spline(prob_pred, prob_true, k=2)\n    y_smooth = spl(x_new)\nelse:\n    x_new, y_smooth = prob_pred, prob_true\n\n# =========================\n# PLOT\n# =========================\n\nplt.figure(figsize=(6,5))\n\n# smoothed model curve\nplt.plot(x_new, y_smooth, linewidth=2, label=\"ResNet-50\")\n\n# actual points\nplt.scatter(prob_pred, prob_true, s=30)\n\n# perfect calibration\nplt.plot([0,1], [0,1], linestyle='--', color='black', linewidth=1, label=\"Perfect Calibration\")\n\nplt.xlabel(\"Mean Predicted Probability\")\nplt.ylabel(\"Fraction of Positives\")\nplt.title(\"Calibration Curve (IDH Prediction)\")\n\n# 🔥 BRIER SCORE (top-left, clean)\nplt.text(\n    0.05, 0.85,\n    f\"Brier = {brier:.3f}\",\n    transform=plt.gca().transAxes,\n    fontsize=10,\n    bbox=dict(facecolor='white', alpha=0.9, edgecolor='black')\n)\n\nplt.legend(frameon=False)\nplt.grid(alpha=0.3)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:22:54.673473Z","iopub.execute_input":"2026-04-19T22:22:54.674887Z","iopub.status.idle":"2026-04-19T22:24:37.159659Z","shell.execute_reply.started":"2026-04-19T22:22:54.674849Z","shell.execute_reply":"2026-04-19T22:24:37.158806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL PR CURVE (MGMT - POLISHED)\n# =========================\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import precision_recall_curve, auc\n\nplt.figure(figsize=(7,5))\n\nname_map = {\n    \"efficientnet_b2\": \"EfficientNet-B2\",\n    \"resnet50_v2\": \"ResNet-50\",\n    \"convnext_base\": \"ConvNeXt-Base\"\n}\n\ncolors = {\n    \"efficientnet_b2\": \"tab:blue\",\n    \"resnet50_v2\": \"tab:red\",\n    \"convnext_base\": \"tab:green\"\n}\n\nrecall_grid = np.linspace(0, 1, 200)\n\nall_targets = []\n\n# =========================\n# MODEL CURVES\n# =========================\n\nfor name, model in models_dict.items():\n\n    if name not in name_map:\n        continue\n\n    y_true, y_score = [], []\n\n    model.eval()\n\n    with torch.no_grad():\n        for batch in val_loader:\n\n            x = batch[0].to(device)\n            labels = batch[1]\n\n            mgmt = labels[\"mgmt\"]\n\n            _, out_mgmt, _ = model(x)\n            prob = torch.sigmoid(out_mgmt).squeeze().cpu().numpy()\n\n            mgmt_np = mgmt.numpy()\n            mask = mgmt_np != -1\n\n            y_true.extend(mgmt_np[mask])\n            y_score.extend(prob[mask])\n\n    all_targets.extend(y_true)\n\n    precision, recall, _ = precision_recall_curve(y_true, y_score)\n    pr_auc = auc(recall, precision)\n\n    # smoothing (light, safe)\n    precision_smooth = np.interp(recall_grid, recall[::-1], precision[::-1])\n\n    plt.plot(\n        recall_grid,\n        precision_smooth,\n        color=colors[name],\n        linewidth=2.5 if name==\"efficientnet_b2\" else 1.8,\n        label=f\"{name_map[name]} (AUC = {pr_auc:.3f})\"\n    )\n\n# =========================\n# BASELINE (IMBALANCE)\n# =========================\n\nbaseline = np.mean(all_targets)\n\nplt.hlines(\n    baseline,\n    0, 1,\n    linestyles=\"dashed\",\n    colors=\"gray\",\n    alpha=0.6,   # 🔥 lighter baseline (your request)\n    linewidth=1.5,\n    label=f\"Baseline ({baseline:.2f})\"\n)\n\n# =========================\n# FINAL STYLING\n# =========================\n\nplt.xlabel(\"Recall (Sensitivity)\")\nplt.ylabel(\"Precision\")\nplt.title(\"Precision-Recall Curves for MGMT Methylation Prediction\")\n\nplt.legend(frameon=False)\nplt.grid(alpha=0.2)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:25:52.468779Z","iopub.execute_input":"2026-04-19T22:25:52.469464Z","iopub.status.idle":"2026-04-19T22:31:15.082430Z","shell.execute_reply.started":"2026-04-19T22:25:52.469365Z","shell.execute_reply":"2026-04-19T22:31:15.081593Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================\n# FINAL: CALIBRATION CURVE (IDH) — STABLE VERSION\n# =========================\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.calibration import calibration_curve\nfrom sklearn.metrics import brier_score_loss\nfrom scipy.interpolate import PchipInterpolator\n\nmodel = models_dict[\"resnet50_v2\"]\n\ny_true, y_prob = [], []\n\nmodel.eval()\n\nwith torch.no_grad():\n    for batch in val_loader:\n\n        x = batch[0].to(device)\n        labels = batch[1]\n\n        idh = labels[\"idh\"]\n\n        out_idh, _, _ = model(x)\n        prob = torch.sigmoid(out_idh).squeeze().cpu().numpy()\n\n        y_prob.extend(prob)\n        y_true.extend(idh.numpy())\n\n# =========================\n# CALIBRATION\n# =========================\n\nprob_true, prob_pred = calibration_curve(\n    y_true,\n    y_prob,\n    n_bins=15,\n    strategy=\"quantile\"   # important for stability\n)\n\nbrier = brier_score_loss(y_true, y_prob)\n\n# =========================\n# SORT (required for interpolation)\n# =========================\n\norder = np.argsort(prob_pred)\nprob_pred = prob_pred[order]\nprob_true = prob_true[order]\n\n# =========================\n# SAFE SMOOTHING (PCHIP)\n# =========================\n\nx_new = np.linspace(0, 1, 100)\n\nif len(prob_pred) >= 3:\n    interp = PchipInterpolator(prob_pred, prob_true)\n    y_smooth = interp(x_new)\n    y_smooth = np.clip(y_smooth, 0, 1)   # 🔥 prevent invalid values\nelse:\n    x_new, y_smooth = prob_pred, prob_true\n\n# =========================\n# PLOT\n# =========================\n\nplt.figure(figsize=(6,5))\n\n# smoothed curve\nplt.plot(x_new, y_smooth, linewidth=2, label=\"ResNet-50\")\n\n# actual calibration points\nplt.scatter(prob_pred, prob_true, s=35)\n\n# perfect calibration\nplt.plot([0,1], [0,1], linestyle='--', color='black', linewidth=1, label=\"Perfect Calibration\")\n\nplt.xlabel(\"Mean Predicted Probability\")\nplt.ylabel(\"Fraction of Positives\")\nplt.title(\"Calibration Curve (IDH Prediction)\")\n\n# 🔥 BRIER SCORE (clean placement)\nplt.text(\n    0.05, 0.85,\n    f\"Brier = {brier:.3f}\",\n    transform=plt.gca().transAxes,\n    fontsize=10,\n    bbox=dict(facecolor='white', alpha=0.9, edgecolor='black')\n)\n\n# 🔥 IMPORTANT: VALID AXIS\nplt.xlim(0, 1)\nplt.ylim(0, 1)\n\nplt.legend(frameon=False)\nplt.grid(alpha=0.3)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T22:33:27.619762Z","iopub.execute_input":"2026-04-19T22:33:27.620585Z","iopub.status.idle":"2026-04-19T22:35:10.530230Z","shell.execute_reply.started":"2026-04-19T22:33:27.620509Z","shell.execute_reply":"2026-04-19T22:35:10.529167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}