{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":677849,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":514051,"modelId":528691}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport pydicom\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport timm\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom joblib import Parallel, delayed\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, confusion_matrix, roc_auc_score, f1_score\n\nclass Config:\n    BASE_PATH = \"/kaggle/input/rsna-intracranial-aneurysm-detection\"\n    CSV_PATH = os.path.join(BASE_PATH, \"train.csv\")\n    IMAGE_DIR = os.path.join(BASE_PATH, \"series\")\n    \n    # Saving as Binary Model\n    SAVE_PATH = \"best_binary_mra_model.pth\"\n    \n    # BINARY LABEL ONLY\n    LABEL_COLS = ['Aneurysm Present']\n    \n    MODEL_NAME = 'swin_tiny_patch4_window7_224' \n    IMG_SIZE = (224, 224)\n    BATCH_SIZE = 16\n    EPOCHS = 10 \n    LR = 5e-5\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nprint(f\"✅ Configuration Ready. Device: {Config.DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:04:22.344351Z","iopub.execute_input":"2025-12-09T23:04:22.344855Z","iopub.status.idle":"2025-12-09T23:04:34.195971Z","shell.execute_reply.started":"2025-12-09T23:04:22.344833Z","shell.execute_reply":"2025-12-09T23:04:34.195274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_mra_binary_series(row, image_dir, label_cols):\n    \"\"\"\n    Worker function: Checks if series is MRA, then extracts \n    paths for binary classification.\n    \"\"\"\n    sid = row['SeriesInstanceUID']\n    # Extract just the binary label (0 or 1)\n    # The CSV might have many columns, we only grabbed ['Aneurysm Present'] in Config\n    labels = row[label_cols].values.astype(np.float32)\n    \n    series_path = os.path.join(image_dir, str(sid))\n    if not os.path.exists(series_path): return []\n    \n    try:\n        dir_files = [f for f in os.listdir(series_path) if f.endswith('.dcm')]\n        if not dir_files: return []\n        \n        # 1. Check Modality (Must be MR)\n        first = pydicom.dcmread(os.path.join(series_path, dir_files[0]), stop_before_pixels=True)\n        if \"MR\" not in getattr(first, \"Modality\", \"\"): return []\n\n        # 2. Get all slices\n        files = []\n        for f in dir_files:\n            p = os.path.join(series_path, f)\n            d = pydicom.dcmread(p, stop_before_pixels=True)\n            files.append((int(d.InstanceNumber), p))\n            \n        if len(files) < 3: return []\n        files.sort(key=lambda x: x[0])\n        paths = [f[1] for f in files]\n        n_slices = len(paths)\n        \n        # 3. Smart Sampling\n        # Positive (1) -> 5 Slices\n        # Negative (0) -> 1 Slice (Center)\n        samples = []\n        if labels[0] == 1: \n            indices = np.linspace(0, n_slices - 1, num=5).astype(int)\n            indices = np.unique(indices)\n        else:\n            indices = [n_slices // 2]\n\n        for c_idx in indices:\n            samples.append({'paths': paths, 'center_idx': c_idx, 'labels': labels})\n        return samples\n    except: return []\n\nclass BinaryMRADataset(Dataset):\n    def __init__(self, df):\n        print(f\"⚡ Indexing {len(df)} series (Filtering MRA)...\")\n        results = Parallel(n_jobs=4, backend=\"threading\")(\n            delayed(process_mra_binary_series)(row, Config.IMAGE_DIR, Config.LABEL_COLS) \n            for _, row in tqdm(df.iterrows(), total=len(df))\n        )\n        self.samples = [item for sublist in results for item in sublist]\n        print(f\"✅ Found {len(self.samples)} valid samples.\")\n\n    def __len__(self): return len(self.samples)\n\n    def __getitem__(self, idx):\n        item = self.samples[idx]\n        imgs = []\n        for d in (-1, 0, 1):\n            k = np.clip(item['center_idx'] + d, 0, len(item['paths']) - 1)\n            try:\n                ds = pydicom.dcmread(item['paths'][k], stop_before_pixels=False)\n                img = ds.pixel_array.astype(np.float32)\n                c, w = 100, 200 # MRA Window\n                img = np.clip(img, c - w/2, c + w/2)\n                img = (img - (c - w/2)) / w\n                img = cv2.resize(img, Config.IMG_SIZE)\n            except: img = np.zeros(Config.IMG_SIZE, dtype=np.float32)\n            imgs.append(img)\n        return torch.tensor(np.stack(imgs), dtype=torch.float32), torch.tensor(item['labels'], dtype=torch.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:04:34.19719Z","iopub.execute_input":"2025-12-09T23:04:34.197617Z","iopub.status.idle":"2025-12-09T23:04:34.209065Z","shell.execute_reply.started":"2025-12-09T23:04:34.197597Z","shell.execute_reply":"2025-12-09T23:04:34.208414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNABinaryViT(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model(Config.MODEL_NAME, pretrained=True, in_chans=3, num_classes=0)\n        dim = self.backbone.num_features if hasattr(self.backbone, 'num_features') else self.backbone.embed_dim\n        \n        # Single output neuron for Binary Classification\n        self.head = nn.Sequential(\n            nn.LayerNorm(dim),\n            nn.Dropout(0.3),\n            nn.Linear(dim, 1) \n        )\n        \n    def forward(self, x): return self.head(self.backbone(x))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:04:34.209608Z","iopub.execute_input":"2025-12-09T23:04:34.209821Z","iopub.status.idle":"2025-12-09T23:04:34.23163Z","shell.execute_reply.started":"2025-12-09T23:04:34.209796Z","shell.execute_reply":"2025-12-09T23:04:34.231033Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_binary_3split_clean():\n    print(\"🔍 Loading & Analyzing MRA Data...\")\n    df = pd.read_csv(Config.CSV_PATH)\n    \n    # 1. Filter MRA\n    meta_path = os.path.join(Config.BASE_PATH, \"train_series_metadata.csv\")\n    if os.path.exists(meta_path):\n        meta = pd.read_csv(meta_path)\n        df = df.merge(meta[['SeriesInstanceUID', 'Modality']], on=\"SeriesInstanceUID\", how=\"left\")\n        df = df[df['Modality'] == 'MRA'].reset_index(drop=True)\n    \n    # 2. Index Dataset\n    full_ds = BinaryMRADataset(df)\n    if len(full_ds) == 0: return None, None\n\n    # 3. Stratified Split\n    all_labels = [int(x['labels'][0]) for x in full_ds.samples]\n    indices = np.arange(len(full_ds))\n    \n    # Split 80/10/10\n    train_idx, temp_idx, y_train, y_temp = train_test_split(indices, all_labels, stratify=all_labels, test_size=0.2, random_state=42)\n    val_idx, test_idx, y_val, y_test = train_test_split(temp_idx, y_temp, stratify=y_temp, test_size=0.5, random_state=42)\n    \n    # 4. PRINT DATA STATISTICS\n    print(\"\\n\" + \"=\"*50)\n    print(\"📊 DATASET DISTRIBUTION REPORT\")\n    print(\"=\"*50)\n    print(f\"{'SPLIT':<10} | {'TOTAL':<8} | {'HEALTHY (0)':<12} | {'SICK (1)':<12} | {'RATIO'}\")\n    print(\"-\" * 65)\n    \n    def print_stat(name, y_data):\n        c = np.bincount(y_data)\n        n0, n1 = c[0], c[1]\n        print(f\"{name:<10} | {len(y_data):<8} | {n0:<12} | {n1:<12} | 1:{n0/n1:.1f}\")\n        return n0, n1\n\n    n0_train, n1_train = print_stat(\"TRAIN\", y_train)\n    print_stat(\"VALIDATION\", y_val)\n    print_stat(\"TEST\", y_test)\n    print(\"-\" * 65)\n\n    # 5. Sampler (Balanced Training)\n    w0 = 1.0 / n0_train\n    w1 = 1.0 / n1_train\n    sample_weights = [w1 if t == 1 else w0 for t in y_train]\n    sampler = WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)\n    \n    # 6. Loaders\n    train_ds = torch.utils.data.Subset(full_ds, train_idx)\n    val_ds = torch.utils.data.Subset(full_ds, val_idx)\n    test_ds = torch.utils.data.Subset(full_ds, test_idx)\n    \n    train_loader = DataLoader(train_ds, batch_size=Config.BATCH_SIZE, sampler=sampler, num_workers=2)\n    val_loader = DataLoader(val_ds, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n    test_loader = DataLoader(test_ds, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n    \n    # 7. Model Setup\n    model = RSNABinaryViT().to(Config.DEVICE)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-2)\n    scheduler = CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    criterion = nn.BCEWithLogitsLoss()\n    \n    # 8. Clean Training Loop\n    best_auc = 0.0\n    print(f\"\\n🚀 Starting Training ({Config.EPOCHS} Epochs)...\")\n    \n    for epoch in range(Config.EPOCHS):\n        model.train()\n        \n        # TQDM Bar\n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{Config.EPOCHS}\", leave=True)\n        \n        for x, y in pbar:\n            x, y = x.to(Config.DEVICE), y.to(Config.DEVICE)\n            optimizer.zero_grad()\n            logits = model(x)\n            loss = criterion(logits, y)\n            loss.backward()\n            optimizer.step()\n            \n            pbar.set_postfix({'loss': f\"{loss.item():.4f}\"})\n        \n        scheduler.step()\n        \n        # Validation\n        model.eval()\n        y_true, y_probs = [], []\n        with torch.no_grad():\n            for x, y in val_loader:\n                x, y = x.to(Config.DEVICE), y.to(Config.DEVICE)\n                logits = model(x)\n                y_true.append(y.cpu().numpy())\n                y_probs.append(torch.sigmoid(logits).cpu().numpy())\n        \n        y_true = np.concatenate(y_true)\n        y_probs = np.concatenate(y_probs)\n        try:\n            val_auc = roc_auc_score(y_true, y_probs)\n        except: val_auc = 0.5\n        \n        print(f\"   └── Results: Val AUC: {val_auc:.4f} | Best: {max(best_auc, val_auc):.4f}\")\n        \n        if val_auc > best_auc:\n            best_auc = val_auc\n            torch.save({'state_dict': model.state_dict(), 'config': Config}, Config.SAVE_PATH)\n\n    # Load Best (FIXED LINE BELOW)\n    # weights_only=False fixes the UnpicklingError\n    ckpt = torch.load(Config.SAVE_PATH, weights_only=False) \n    model.load_state_dict(ckpt['state_dict'])\n    print(f\"\\n✅ Training Complete. Best Model AUC: {best_auc:.4f}\")\n    \n    return model, test_loader\n\n# RUN\nmodel, test_loader = train_binary_3split_clean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:04:34.233013Z","iopub.execute_input":"2025-12-09T23:04:34.233188Z","iopub.status.idle":"2025-12-09T23:37:31.373509Z","shell.execute_reply.started":"2025-12-09T23:04:34.233175Z","shell.execute_reply":"2025-12-09T23:37:31.372729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_final(model, loader):\n    model.eval()\n    y_true, y_probs = [], []\n    \n    print(\"🏥 Running Final Test Evaluation...\")\n    with torch.no_grad():\n        for x, y in tqdm(loader):\n            logits = model(x.to(Config.DEVICE))\n            y_true.append(y.cpu().numpy())\n            y_probs.append(torch.sigmoid(logits).cpu().numpy())\n            \n    y_true = np.concatenate(y_true)\n    y_probs = np.concatenate(y_probs)\n    y_preds = (y_probs > 0.5).astype(int) # Standard threshold\n    \n    # Metrics\n    auc_score = roc_auc_score(y_true, y_probs)\n    acc = accuracy_score(y_true, y_preds)\n    f1 = f1_score(y_true, y_preds)\n    tn, fp, fn, tp = confusion_matrix(y_true, y_preds).ravel()\n    sens = tp / (tp + fn)\n    spec = tn / (tn + fp)\n    \n    print(\"\\n\" + \"=\"*40)\n    print(\"FINAL TEST SET RESULTS\")\n    print(\"=\"*40)\n    print(f\"AUC Score:   {auc_score:.4f}\")\n    print(f\"Accuracy:    {acc:.4f}\")\n    print(f\"F1 Score:    {f1:.4f}\")\n    print(f\"Sensitivity: {sens:.1%} (Recall)\")\n    print(f\"Specificity: {spec:.1%}\")\n    print(\"-\" * 20)\n    print(f\"TP: {tp} | FN: {fn}\")\n    print(f\"FP: {fp} | TN: {tn}\")\n    \n    plt.figure(figsize=(5,4))\n    sns.heatmap([[tn, fp], [fn, tp]], annot=True, fmt='d', cmap='Blues', cbar=False)\n    plt.title(\"Test Set Confusion Matrix\")\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"Actual\")\n    plt.show()\n\nif test_loader:\n    evaluate_final(model, test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:37:31.37501Z","iopub.execute_input":"2025-12-09T23:37:31.375236Z","iopub.status.idle":"2025-12-09T23:37:47.320733Z","shell.execute_reply.started":"2025-12-09T23:37:31.375213Z","shell.execute_reply":"2025-12-09T23:37:47.320096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BinaryScreener:\n    def __init__(self, model_path, device='cuda'):\n        self.device = device\n        ckpt = torch.load(model_path, map_location=device)\n        self.cfg = ckpt['config']\n        \n        self.backbone = timm.create_model(self.cfg.MODEL_NAME, pretrained=False, in_chans=3, num_classes=0)\n        dim = self.backbone.num_features if hasattr(self.backbone, 'num_features') else self.backbone.embed_dim\n        self.head = nn.Sequential(nn.LayerNorm(dim), nn.Dropout(0.3), nn.Linear(dim, 1))\n        \n        full = nn.Sequential(self.backbone, self.head)\n        \n        # Clean keys\n        state = {k.replace('backbone.', '0.').replace('head.', '1.'): v for k, v in ckpt['state_dict'].items()}\n        full.load_state_dict(state)\n        self.model = full.to(device).eval()\n        print(\"✅ Screener Ready.\")\n\n    def predict(self, dcm_path):\n        try:\n            ds = pydicom.dcmread(dcm_path, stop_before_pixels=False)\n            img = ds.pixel_array.astype(np.float32)\n            c, w = 100, 200\n            img = np.clip(img, c - w/2, c + w/2)\n            img = (img - (c - w/2)) / w\n            img = cv2.resize(img, self.cfg.IMG_SIZE)\n            \n            stack = np.stack([img, img, img], axis=0)\n            t = torch.tensor(stack, dtype=torch.float32).unsqueeze(0).to(self.device)\n            \n            with torch.no_grad():\n                prob = torch.sigmoid(self.model(t)).item()\n            return f\"Probability: {prob:.4f} ({'SICK' if prob>0.5 else 'HEALTHY'})\"\n        except: return \"Error\"\n\n# screener = BinaryScreener(Config.SAVE_PATH)\n# print(screener.predict(\"test.dcm\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:37:47.321593Z","iopub.execute_input":"2025-12-09T23:37:47.321969Z","iopub.status.idle":"2025-12-09T23:37:47.330096Z","shell.execute_reply.started":"2025-12-09T23:37:47.321945Z","shell.execute_reply":"2025-12-09T23:37:47.32949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\nimport torch\nfrom sklearn.metrics import confusion_matrix, roc_curve, auc\nfrom tqdm import tqdm\n\ndef visualize_model_story(model, loader, device=Config.DEVICE, title=\"MRA Aneurysm Detection Model\"):\n    \"\"\"\n    Generates a 4-panel visualization storyboard explaining model performance.\n    \"\"\"\n    model.eval()\n    y_true, y_probs = [], []\n    \n    print(\"📊 Running inference for visualization...\")\n    with torch.no_grad():\n        for x, y in tqdm(loader, leave=False):\n            logits = model(x.to(device))\n            y_true.append(y.cpu().numpy())\n            y_probs.append(torch.sigmoid(logits).cpu().numpy())\n            \n    y_true = np.concatenate(y_true).ravel()\n    y_probs = np.concatenate(y_probs).ravel()\n    \n    # Set up the figure\n    fig, axes = plt.subplots(2, 2, figsize=(16, 14))\n    plt.suptitle(f\"The Story of Your Model: {title}\", fontsize=20, y=0.95)\n    \n    # ==================================================\n    # PANEL 1: Confusion Matrix (At Threshold 0.5)\n    # ==================================================\n    threshold = 0.5\n    y_preds = (y_probs > threshold).astype(int)\n    tn, fp, fn, tp = confusion_matrix(y_true, y_preds).ravel()\n    \n    labels = [f\"True Neg (Healthy)\\n{tn}\\n({tn/len(y_true):.1%})\", \n              f\"False Pos (False Alarm)\\n{fp}\\n({fp/len(y_true):.1%})\",\n              f\"False Neg (Missed)\\n{fn}\\n({fn/len(y_true):.1%})\", \n              f\"True Pos (Detected)\\n{tp}\\n({tp/len(y_true):.1%})\"]\n    labels = np.asarray(labels).reshape(2,2)\n    \n    sns.heatmap([[tn, fp], [fn, tp]], annot=labels, fmt='', cmap='Blues', ax=axes[0,0], cbar=False, annot_kws={\"size\": 12})\n    axes[0,0].set_title(f\"1. The Outcome (Confusion Matrix @ {threshold})\", fontsize=14)\n    axes[0,0].set_xlabel(\"AI Prediction\", fontsize=12)\n    axes[0,0].set_ylabel(\"Ground Truth\", fontsize=12)\n    axes[0,0].set_xticklabels(['Healthy', 'Sick'])\n    axes[0,0].set_yticklabels(['Healthy', 'Sick'])\n    \n    # ==================================================\n    # PANEL 2: ROC Curve (Ranking Ability)\n    # ==================================================\n    fpr, tpr, thresholds = roc_curve(y_true, y_probs)\n    roc_auc = auc(fpr, tpr)\n    \n    axes[0,1].plot(fpr, tpr, color='darkorange', lw=3, label=f'AUC = {roc_auc:.3f}')\n    axes[0,1].plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n    axes[0,1].set_xlim([-0.01, 1.0])\n    axes[0,1].set_ylim([0.0, 1.02])\n    axes[0,1].set_xlabel('False Positive Rate (1 - Specificity)', fontsize=12)\n    axes[0,1].set_ylabel('True Positive Rate (Sensitivity)', fontsize=12)\n    axes[0,1].set_title(\"2. Ranking Power (ROC Curve)\", fontsize=14)\n    axes[0,1].legend(loc=\"lower right\", fontsize=12)\n    axes[0,1].grid(alpha=0.3)\n    \n    # ==================================================\n    # PANEL 3: Confidence Histogram (Separation)\n    # ==================================================\n    sns.histplot(y_probs[y_true==0], color='green', label='Healthy Patients', kde=True, ax=axes[1,0], stat=\"density\", linewidth=0)\n    sns.histplot(y_probs[y_true==1], color='red', label='Sick Patients', kde=True, ax=axes[1,0], stat=\"density\", linewidth=0, alpha=0.6)\n    axes[1,0].axvline(threshold, color='black', linestyle='--', label=f'Threshold {threshold}')\n    axes[1,0].set_title(\"3. Model's Confidence (Separation)\", fontsize=14)\n    axes[1,0].set_xlabel('Predicted Probability of Aneurysm (0.0 to 1.0)', fontsize=12)\n    axes[1,0].set_ylabel('Density of Patients', fontsize=12)\n    axes[1,0].legend(fontsize=10)\n    axes[1,0].grid(alpha=0.3)\n    \n    # ==================================================\n    # PANEL 4: The Trade-off (Sens vs Spec Curve)\n    # ==================================================\n    sens, specs = [], []\n    thresh_range = np.linspace(0.01, 0.99, 100)\n    for t in thresh_range:\n        yp = (y_probs > t).astype(int)\n        tn_t, fp_t, fn_t, tp_t = confusion_matrix(y_true, yp).ravel()\n        sens.append(tp_t / (tp_t + fn_t))\n        specs.append(tn_t / (tn_t + fp_t))\n        \n    axes[1,1].plot(thresh_range, sens, color='red', lw=2, label='Sensitivity (Catching Sick)')\n    axes[1,1].plot(thresh_range, specs, color='green', lw=2, label='Specificity (Clearing Healthy)')\n    \n    # Find intersection (balance point)\n    idx = np.argwhere(np.diff(np.sign(np.array(sens) - np.array(specs)))).flatten()\n    if len(idx) > 0:\n        balance_thresh = thresh_range[idx[0]]\n        balance_val = sens[idx[0]]\n        axes[1,1].plot(balance_thresh, balance_val, 'ko', markersize=8, label=f'Balance Point (~{balance_thresh:.2f})')\n        axes[1,1].axvline(balance_thresh, color='gray', linestyle=':', alpha=0.5)\n        \n    axes[1,1].set_xlabel('Decision Threshold', fontsize=12)\n    axes[1,1].set_ylabel('Score (0.0 to 1.0)', fontsize=12)\n    axes[1,1].set_title(\"4. The Trade-off (Making a Decision)\", fontsize=14)\n    axes[1,1].set_ylim([0.0, 1.05])\n    axes[1,1].legend(fontsize=10)\n    axes[1,1].grid(alpha=0.3)\n    \n    plt.tight_layout()\n    plt.subplots_adjust(top=0.90) # Make room for suptitle\n    plt.show()\n    \n    print(\"\\n💡 THE STORY INTERPRETATION:\")\n    print(\"1. Confusion Matrix: Shows the raw results. Look at the top-right (False Positives) vs bottom-right (True Positives).\")\n    print(\"2. ROC Curve: AUC > 0.90 means the model is excellent at ranking sick patients higher than healthy ones.\")\n    print(\"3. Confidence Histogram: Ideally, you want a big green peak on the left (0.0) and a big red peak on the right (1.0). Overlap is where errors happen.\")\n    print(\"4. Trade-off Curve: This is for the doctor. If they want 95% Sensitivity, they can find the threshold on the red line and see what Specificity (green line) they will sacrifice.\")\n\n# Run the storyteller on your held-out TEST set\nif 'test_loader' in locals() and 'model' in locals():\n    visualize_model_story(model, test_loader)\nelse:\n    print(\"Please run training first to define 'model' and 'test_loader'.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T23:40:12.417324Z","iopub.execute_input":"2025-12-09T23:40:12.417627Z","iopub.status.idle":"2025-12-09T23:40:26.278503Z","shell.execute_reply.started":"2025-12-09T23:40:12.417606Z","shell.execute_reply":"2025-12-09T23:40:26.277605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, confusion_matrix, roc_auc_score, f1_score\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport pandas as pd\nimport os\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Subset\nfrom sklearn.model_selection import train_test_split\n\n# ==========================================\n# 1. RE-CREATE THE EXACT SPLITS\n# ==========================================\ndef get_all_loaders():\n    print(\"🔍 Re-loading Full Dataset...\")\n    df = pd.read_csv(Config.CSV_PATH)\n    \n    # Filter MRA\n    meta_path = os.path.join(Config.BASE_PATH, \"train_series_metadata.csv\")\n    if os.path.exists(meta_path):\n        meta = pd.read_csv(meta_path)\n        df = df.merge(meta[['SeriesInstanceUID', 'Modality']], on=\"SeriesInstanceUID\", how=\"left\")\n        df = df[df['Modality'] == 'MRA'].reset_index(drop=True)\n    \n    # Create Full Dataset\n    ds = BinaryMRADataset(df)\n    \n    # Re-create Stratified Split (Same Random State = Same Split)\n    all_labels = [int(x['labels'][0]) for x in ds.samples]\n    indices = np.arange(len(ds))\n    \n    train_idx, temp_idx, y_train, y_temp = train_test_split(indices, all_labels, stratify=all_labels, test_size=0.2, random_state=42)\n    val_idx, test_idx, y_val, y_test = train_test_split(temp_idx, y_temp, stratify=y_temp, test_size=0.5, random_state=42)\n    \n    print(f\"📊 Splits Re-created: Train={len(train_idx)}, Test={len(test_idx)}\")\n    \n    # Create Loaders\n    # Note: We do NOT use the Sampler here because we want to see true performance on the real imbalanced training set\n    train_ds = Subset(ds, train_idx)\n    test_ds = Subset(ds, test_idx)\n    \n    train_loader = DataLoader(train_ds, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n    test_loader = DataLoader(test_ds, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n    \n    return train_loader, test_loader\n\n# ==========================================\n# 2. EVALUATION FUNCTION\n# ==========================================\ndef get_metrics(model, loader, device=Config.DEVICE):\n    model.eval()\n    y_true, y_probs = [], []\n    \n    with torch.no_grad():\n        for x, y in tqdm(loader, leave=False):\n            logits = model(x.to(device))\n            y_true.append(y.cpu().numpy())\n            y_probs.append(torch.sigmoid(logits).cpu().numpy())\n            \n    y_true = np.concatenate(y_true).ravel()\n    y_probs = np.concatenate(y_probs).ravel()\n    y_preds = (y_probs > 0.5).astype(int)\n    \n    try: auc = roc_auc_score(y_true, y_probs)\n    except: auc = 0.5\n        \n    tn, fp, fn, tp = confusion_matrix(y_true, y_preds).ravel()\n    sens = tp / (tp + fn) if (tp + fn) > 0 else 0\n    spec = tn / (tn + fp) if (tn + fp) > 0 else 0\n    acc = accuracy_score(y_true, y_preds)\n    f1 = f1_score(y_true, y_preds)\n    \n    return {'auc': auc, 'acc': acc, 'f1': f1, 'sens': sens, 'spec': spec, 'cm': (tn, fp, fn, tp)}\n\n# ==========================================\n# 3. EXECUTE COMPARISON\n# ==========================================\ndef compare_train_vs_test(model):\n    # 1. Get Data\n    train_loader, test_loader = get_all_loaders()\n    \n    print(\"\\n🚀 Evaluating Training Set (Memory)...\")\n    t_m = get_metrics(model, train_loader)\n    \n    print(\"🚀 Evaluating Test Set (Generalization)...\")\n    v_m = get_metrics(model, test_loader)\n    \n    # 2. Print Table\n    print(\"\\n\" + \"=\"*70)\n    print(f\"{'METRIC':<15} | {'TRAIN (Learned)':<18} | {'TEST (Unseen)':<18} | {'GAP'}\")\n    print(\"=\"*70)\n    \n    metrics = ['AUC Score', 'Accuracy', 'F1 Score', 'Sensitivity', 'Specificity']\n    keys = ['auc', 'acc', 'f1', 'sens', 'spec']\n    \n    for name, key in zip(metrics, keys):\n        t_val = t_m[key]\n        v_val = v_m[key]\n        gap = t_val - v_val\n        status = \"✅\" if abs(gap) < 0.10 else \"⚠️\"\n        print(f\"{name:<15} | {t_val:.4f}             | {v_val:.4f}             | {gap:+.3f} {status}\")\n    print(\"-\" * 70)\n\n    # 3. Plot\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    \n    # Train Matrix\n    tn, fp, fn, tp = t_m['cm']\n    labels = [f\"TN\\n{tn}\", f\"FP\\n{fp}\", f\"FN\\n{fn}\", f\"TP\\n{tp}\"]\n    sns.heatmap([[tn, fp], [fn, tp]], annot=np.array(labels).reshape(2,2), fmt='', cmap='Blues', ax=axes[0], cbar=False)\n    axes[0].set_title(f\"TRAIN SET (AUC {t_m['auc']:.3f})\")\n    \n    # Test Matrix\n    tn, fp, fn, tp = v_m['cm']\n    labels = [f\"TN\\n{tn}\", f\"FP\\n{fp}\", f\"FN\\n{fn}\", f\"TP\\n{tp}\"]\n    sns.heatmap([[tn, fp], [fn, tp]], annot=np.array(labels).reshape(2,2), fmt='', cmap='Greens', ax=axes[1], cbar=False)\n    axes[1].set_title(f\"TEST SET (AUC {v_m['auc']:.3f})\")\n    \n    plt.show()\n\n# Run\nif 'model' in locals():\n    compare_train_vs_test(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T00:19:33.461639Z","iopub.execute_input":"2025-12-10T00:19:33.46233Z","iopub.status.idle":"2025-12-10T00:33:02.924267Z","shell.execute_reply.started":"2025-12-10T00:19:33.462303Z","shell.execute_reply":"2025-12-10T00:33:02.923432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FIX: calibration_curve is in sklearn.calibration\nfrom sklearn.metrics import precision_recall_curve, average_precision_score\nfrom sklearn.calibration import calibration_curve\nimport matplotlib.pyplot as plt\nimport torch\nimport numpy as np\nimport cv2\nimport random\n\ndef advanced_model_diagnostics(model, loader, device=Config.DEVICE):\n    model.eval()\n    y_true, y_probs = [], []\n    images = [] \n    \n    print(\"🔬 Running Advanced Diagnostics (Collecting all test images)...\")\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(device)\n            logits = model(x)\n            probs = torch.sigmoid(logits).cpu().numpy().flatten()\n            \n            # 1. Collect Data\n            y_true.extend(y.cpu().numpy().flatten())\n            y_probs.extend(probs)\n            \n            # 2. Collect Images for Visualization \n            # (We collect all to ensure true randomness)\n            imgs = x.cpu().numpy()\n            for i in range(len(imgs)):\n                # Take middle slice (channel 1) for visualization\n                img = imgs[i, 1, :, :] \n                images.append(img)\n    \n    y_true = np.array(y_true)\n    y_probs = np.array(y_probs)\n    images = np.array(images)\n    \n    # ==========================================\n    # 1. PRECISION-RECALL CURVE\n    # ==========================================\n    precision, recall, _ = precision_recall_curve(y_true, y_probs)\n    ap_score = average_precision_score(y_true, y_probs)\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"📊 METRIC EXPLANATIONS\")\n    print(\"=\"*60)\n    print(f\"1. Precision-Recall (PR) Curve (AP = {ap_score:.3f})\")\n    print(\"   • Why it matters: PR curves are better than ROC for imbalanced data.\")\n    print(\"   • Interpretation: High AP (>0.90) means your model is excellent at finding\")\n    print(\"     RARE positive cases without drowning in false alarms.\")\n    \n    # ==========================================\n    # 2. CALIBRATION CURVE\n    # ==========================================\n    prob_true, prob_pred = calibration_curve(y_true, y_probs, n_bins=10)\n    \n    print(\"\\n2. Calibration Curve\")\n    print(\"   • Why it matters: Tells us if the 'probability' (e.g. 80%) is real.\")\n    print(\"   • Interpretation:\")\n    print(\"     - Perfect: Dots hug the dashed gray line.\")\n    print(\"     - S-Shape / Below Line: The model is 'Overconfident' (Panic).\")\n    print(\"     - Above Line: The model is 'Underconfident' (Hesitant).\")\n    \n    # PLOT METRICS\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n    \n    # Plot PR Curve\n    axes[0].plot(recall, precision, marker='.', label=f'AP = {ap_score:.3f}')\n    axes[0].set_xlabel('Recall (Sensitivity)')\n    axes[0].set_ylabel('Precision (PPV)')\n    axes[0].set_title('Precision-Recall Curve')\n    axes[0].legend()\n    axes[0].grid(alpha=0.3)\n    \n    # Plot Calibration\n    axes[1].plot(prob_pred, prob_true, marker='o', label='Your Model')\n    axes[1].plot([0, 1], [0, 1], linestyle='--', color='gray', label='Perfectly Calibrated')\n    axes[1].set_xlabel('Mean Predicted Probability')\n    axes[1].set_ylabel('Fraction of Positives')\n    axes[1].set_title('Calibration Curve')\n    axes[1].legend()\n    axes[1].grid(alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # ==========================================\n    # 3. RANDOM SAMPLE GALLERY\n    # ==========================================\n    print(\"\\n\" + \"=\"*60)\n    print(\"📸 RANDOM SAMPLE PREDICTIONS\")\n    print(\"=\"*60)\n    print(\"• GREEN Title: The model was CORRECT.\")\n    print(\"• RED Title:   The model was WRONG.\")\n    \n    if len(images) > 0:\n        # Pick 4 random indices from the total set\n        random_indices = random.sample(range(len(images)), min(4, len(images)))\n        \n        fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n        \n        # If we have fewer than 4 images, ensure axes is iterable\n        if len(random_indices) == 1: axes = [axes]\n            \n        for i, idx in enumerate(random_indices):\n            ax = axes[i]\n            img = images[idx]\n            true_label = \"SICK\" if y_true[idx] == 1 else \"HEALTHY\"\n            pred_prob = y_probs[idx]\n            \n            is_correct = (round(pred_prob) == y_true[idx])\n            color = 'green' if is_correct else 'red'\n            \n            ax.imshow(img, cmap='gray')\n            ax.set_title(f\"True: {true_label}\\nPred: {pred_prob:.1%}\", color=color, fontweight='bold')\n            ax.axis('off')\n        plt.show()\n\n# Run it\nif 'test_loader' in locals():\n    advanced_model_diagnostics(model, test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-10T00:56:11.022008Z","iopub.execute_input":"2025-12-10T00:56:11.022843Z","iopub.status.idle":"2025-12-10T00:56:24.182461Z","shell.execute_reply.started":"2025-12-10T00:56:11.022811Z","shell.execute_reply":"2025-12-10T00:56:24.18155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}