{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom pathlib import Path\nfrom typing import Tuple\n \nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport torchvision.transforms as T\nimport timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:28:59.235214Z","iopub.execute_input":"2026-04-14T08:28:59.235390Z","iopub.status.idle":"2026-04-14T08:29:12.614195Z","shell.execute_reply.started":"2026-04-14T08:28:59.235369Z","shell.execute_reply":"2026-04-14T08:29:12.613605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib\nmatplotlib.use(\"Agg\")          # non-interactive backend — safe on servers & Colab\nimport matplotlib.pyplot as plt\nimport matplotlib.ticker as ticker\n \nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    roc_auc_score, f1_score, classification_report,\n    confusion_matrix, ConfusionMatrixDisplay,\n)\nfrom tqdm import tqdm\n \nwarnings.filterwarnings(\"ignore\", category=UserWarning)\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:30:30.480628Z","iopub.execute_input":"2026-04-14T08:30:30.481288Z","iopub.status.idle":"2026-04-14T08:30:31.214244Z","shell.execute_reply.started":"2026-04-14T08:30:30.481255Z","shell.execute_reply":"2026-04-14T08:30:31.213676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED         = 42\nDATA_DIR     = Path(\"/kaggle/input/competitions/rsna-pneumonia-detection-challenge\")        \nLABEL_CSV    = DATA_DIR / \"stage_2_train_labels.csv\"\nIMG_DIR      = DATA_DIR / \"stage_2_train_images\"\n \nIMG_SIZE     = 224        # ViT-B/16 canonical input size\nBATCH_SIZE   = 32\nNUM_EPOCHS   = 15\nLR           = 2e-5       # fine-tuning LR; ViT is pretrained\nWEIGHT_DECAY = 1e-4\nNUM_WORKERS  = 4\nDEVICE       = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n \n# Lung window (standard chest radiology)\nWC, WW       = -600, 1500  # window center, window width\n \n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:42.515039Z","iopub.execute_input":"2026-04-14T08:39:42.515698Z","iopub.status.idle":"2026-04-14T08:39:42.520955Z","shell.execute_reply.started":"2026-04-14T08:39:42.515663Z","shell.execute_reply":"2026-04-14T08:39:42.519971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"random.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:46.602406Z","iopub.execute_input":"2026-04-14T08:39:46.602812Z","iopub.status.idle":"2026-04-14T08:39:46.608602Z","shell.execute_reply.started":"2026-04-14T08:39:46.602783Z","shell.execute_reply":"2026-04-14T08:39:46.607798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_label_df(csv_path: Path) -> pd.DataFrame:\n    \"\"\"\n    The RSNA CSV has one row per bounding box.\n    A patient_id with ANY box → pneumonia (1), else normal (0).\n    Deduplicated to one row per patient.\n    \"\"\"\n    df = pd.read_csv(csv_path)\n    df = df[[\"patientId\", \"Target\"]].drop_duplicates(\"patientId\")\n    df = df.rename(columns={\"patientId\": \"patient_id\", \"Target\": \"label\"})\n    df[\"dcm_path\"] = df[\"patient_id\"].apply(\n        lambda pid: str(IMG_DIR / f\"{pid}.dcm\")\n    )\n    df = df[df[\"dcm_path\"].apply(os.path.exists)].reset_index(drop=True)\n    print(f\"Total samples : {len(df)}\")\n    print(f\"  Pneumonia   : {df['label'].sum()}\")\n    print(f\"  Normal      : {(df['label'] == 0).sum()}\")\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:48.100038Z","iopub.execute_input":"2026-04-14T08:39:48.100652Z","iopub.status.idle":"2026-04-14T08:39:48.106611Z","shell.execute_reply.started":"2026-04-14T08:39:48.100620Z","shell.execute_reply":"2026-04-14T08:39:48.105968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_splits(\n    df: pd.DataFrame,\n) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n    \"\"\"\n    1. Stratified hold-out of 20% as permanent test set.\n    2. Remaining 80% split 80/20 -> train / val.\n    Final proportions: ~64% train / ~16% val / 20% test.\n    \"\"\"\n    train_val_df, test_df = train_test_split(\n        df, test_size=0.20, stratify=df[\"label\"], random_state=SEED\n    )\n    train_df, val_df = train_test_split(\n        train_val_df, test_size=0.20,\n        stratify=train_val_df[\"label\"], random_state=SEED\n    )\n    print(f\"\\nSplit  ->  Train: {len(train_df)}\"\n          f\"  |  Val: {len(val_df)}\"\n          f\"  |  Test: {len(test_df)}\")\n    return (\n        train_df.reset_index(drop=True),\n        val_df.reset_index(drop=True),\n        test_df.reset_index(drop=True),\n    )\n \n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:50.330005Z","iopub.execute_input":"2026-04-14T08:39:50.330754Z","iopub.status.idle":"2026-04-14T08:39:50.335588Z","shell.execute_reply.started":"2026-04-14T08:39:50.330722Z","shell.execute_reply":"2026-04-14T08:39:50.334855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_dicom_tensor(\n    dcm_path: str,\n    wc: int = WC,\n    ww: int = WW,\n    size: int = IMG_SIZE,\n) -> torch.Tensor:\n    \"\"\"\n    Reads a DICOM file and returns a (3, size, size) float32 tensor in [0, 1].\n \n    Why not convert to PNG/JPG first?\n      PNG/JPG collapses 12-16 bit depth to 8-bit, discarding ~94% of the\n      diagnostic intensity range. Direct reading + lung windowing preserves\n      all clinically relevant HU values.\n \n    Steps:\n      1. Read raw pixel array (16-bit for chest X-ray DICOMs).\n      2. Apply RescaleSlope / RescaleIntercept -> Hounsfield Units.\n      3. Lung window: clip [wc - ww/2, wc + ww/2], rescale to [0, 1].\n      4. Resize to (size, size) via bilinear interpolation.\n      5. Replicate to 3 channels (ViT expects RGB input).\n    \"\"\"\n    ds      = pydicom.dcmread(dcm_path)\n    pixels  = ds.pixel_array.astype(np.float32)\n \n    slope     = float(getattr(ds, \"RescaleSlope\",     1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    pixels    = pixels * slope + intercept\n \n    lo      = wc - ww / 2\n    hi      = wc + ww / 2\n    pixels  = np.clip(pixels, lo, hi)\n    pixels  = (pixels - lo) / (hi - lo)      # [0, 1] float32\n \n    # (H, W) -> (1, 1, H, W) -> resize -> (1, size, size)\n    t = torch.from_numpy(pixels).unsqueeze(0).unsqueeze(0)\n    t = torch.nn.functional.interpolate(\n        t, size=(size, size), mode=\"bilinear\", align_corners=False\n    ).squeeze(0)                               # (1, size, size)\n \n    return t.repeat(3, 1, 1)                  # (3, size, size)\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:53.585825Z","iopub.execute_input":"2026-04-14T08:39:53.586612Z","iopub.status.idle":"2026-04-14T08:39:53.592067Z","shell.execute_reply.started":"2026-04-14T08:39:53.586578Z","shell.execute_reply":"2026-04-14T08:39:53.591514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        augment: bool = False,\n        img_size: int = IMG_SIZE,\n    ):\n        self.df      = df\n        self.augment = augment\n        self.size    = img_size\n \n        # Augmentation transforms (applied only during training run B)\n        self.aug = T.Compose([\n            T.RandomHorizontalFlip(p=0.5),\n            T.RandomRotation(degrees=10),\n            T.RandomAffine(degrees=0, translate=(0.05, 0.05)),\n            T.ColorJitter(brightness=0.15, contrast=0.15),\n        ]) if augment else None\n \n        # ImageNet normalisation — ViT was pretrained on ImageNet\n        self.normalize = T.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std =[0.229, 0.224, 0.225],\n        )\n \n    def __len__(self) -> int:\n        return len(self.df)\n \n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, int]:\n        row   = self.df.iloc[idx]\n        label = int(row[\"label\"])\n        img   = read_dicom_tensor(row[\"dcm_path\"], size=self.size)\n \n        if self.aug is not None:\n            img = self.aug(img)\n \n        return self.normalize(img), label\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:56.721864Z","iopub.execute_input":"2026-04-14T08:39:56.722127Z","iopub.status.idle":"2026-04-14T08:39:56.729215Z","shell.execute_reply.started":"2026-04-14T08:39:56.722105Z","shell.execute_reply":"2026-04-14T08:39:56.728418Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_class_weights(df: pd.DataFrame) -> torch.Tensor:\n    counts  = df[\"label\"].value_counts().sort_index()\n    weights = len(df) / (len(counts) * counts.values)\n    return torch.tensor(weights, dtype=torch.float32).to(DEVICE)\n \n \ndef get_weighted_sampler(df: pd.DataFrame) -> WeightedRandomSampler:\n    \"\"\"Over-samples the minority class so each training batch is balanced.\"\"\"\n    counts         = df[\"label\"].value_counts().sort_index()\n    sample_weights = df[\"label\"].map(1.0 / counts).values\n    return WeightedRandomSampler(\n        weights     = torch.DoubleTensor(sample_weights),\n        num_samples = len(df),\n        replacement = True,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:39:59.924008Z","iopub.execute_input":"2026-04-14T08:39:59.924748Z","iopub.status.idle":"2026-04-14T08:39:59.929883Z","shell.execute_reply.started":"2026-04-14T08:39:59.924714Z","shell.execute_reply":"2026-04-14T08:39:59.929144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model(num_classes: int = 2) -> nn.Module:\n    \"\"\"ViT-B/16 (ImageNet-21k pretrained) with a 2-class classification head.\"\"\"\n    model = timm.create_model(\n        \"vit_base_patch16_224\",\n        pretrained     = True,\n        num_classes    = num_classes,\n        drop_rate      = 0.1,\n        attn_drop_rate = 0.0,\n    )\n    return model.to(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:40:02.497619Z","iopub.execute_input":"2026-04-14T08:40:02.498234Z","iopub.status.idle":"2026-04-14T08:40:02.502384Z","shell.execute_reply.started":"2026-04-14T08:40:02.498202Z","shell.execute_reply":"2026-04-14T08:40:02.501658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_epoch(\n    model,\n    loader,\n    criterion,\n    optimizer=None,\n    phase: str = \"train\",\n) -> Tuple[float, float, float, float, list, list, list]:\n    \"\"\"\n    Runs one full pass over the loader.\n    Returns: avg_loss, accuracy, auc_roc, f1,\n             all_labels, all_probs, all_preds\n    \"\"\"\n    is_train = (phase == \"train\")\n    model.train() if is_train else model.eval()\n \n    total_loss                        = 0.0\n    all_preds, all_labels, all_probs  = [], [], []\n \n    with torch.set_grad_enabled(is_train):\n        for imgs, labels in tqdm(loader, desc=f\"  {phase:5s}\", leave=False):\n            imgs   = imgs.to(DEVICE)\n            labels = labels.to(DEVICE)\n \n            logits = model(imgs)\n            loss   = criterion(logits, labels)\n \n            if is_train:\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n \n            probs = torch.softmax(logits, dim=1)[:, 1].detach().cpu().numpy()\n            preds = logits.argmax(dim=1).detach().cpu().numpy()\n \n            total_loss  += loss.item() * len(labels)\n            all_probs.extend(probs)\n            all_preds.extend(preds)\n            all_labels.extend(labels.cpu().numpy())\n \n    n        = len(all_labels)\n    avg_loss = total_loss / n\n    auc      = roc_auc_score(all_labels, all_probs)\n    f1       = f1_score(all_labels, all_preds, average=\"binary\")\n    acc      = float(np.mean(np.array(all_preds) == np.array(all_labels)))\n \n    return avg_loss, acc, auc, f1, all_labels, all_probs, all_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:40:04.508828Z","iopub.execute_input":"2026-04-14T08:40:04.509244Z","iopub.status.idle":"2026-04-14T08:40:04.516815Z","shell.execute_reply.started":"2026-04-14T08:40:04.509209Z","shell.execute_reply":"2026-04-14T08:40:04.516110Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _plot_training_curves(history: dict, run_name: str) -> None:\n    \"\"\"\n    Saves a 2-panel figure:\n      Left  — Train vs Val Loss over all 15 epochs\n      Right — Train vs Val Accuracy over all 15 epochs\n    Output file: training_curves_{run_name}.png\n    \"\"\"\n    epochs = range(1, len(history[\"train_loss\"]) + 1)\n    fig, axes = plt.subplots(1, 2, figsize=(13, 5))\n    fig.suptitle(f\"Training curves  —  {run_name}\", fontsize=14, fontweight=\"bold\")\n \n    # ── Loss panel ──\n    ax = axes[0]\n    ax.plot(epochs, history[\"train_loss\"],\n            marker=\"o\", linewidth=2, markersize=5,\n            color=\"#4C72B0\", label=\"Train loss\")\n    ax.plot(epochs, history[\"val_loss\"],\n            marker=\"s\", linewidth=2, markersize=5,\n            color=\"#DD8452\", linestyle=\"--\", label=\"Val loss\")\n    ax.set_xlabel(\"Epoch\")\n    ax.set_ylabel(\"Loss (weighted CrossEntropy)\")\n    ax.set_title(\"Loss per epoch\")\n    ax.legend()\n    ax.xaxis.set_major_locator(ticker.MaxNLocator(integer=True))\n    ax.grid(True, linestyle=\":\", alpha=0.6)\n \n    # ── Accuracy panel ──\n    ax = axes[1]\n    tr_pct = [a * 100 for a in history[\"train_acc\"]]\n    va_pct = [a * 100 for a in history[\"val_acc\"]]\n    ax.plot(epochs, tr_pct,\n            marker=\"o\", linewidth=2, markersize=5,\n            color=\"#4C72B0\", label=\"Train acc\")\n    ax.plot(epochs, va_pct,\n            marker=\"s\", linewidth=2, markersize=5,\n            color=\"#DD8452\", linestyle=\"--\", label=\"Val acc\")\n    ax.set_xlabel(\"Epoch\")\n    ax.set_ylabel(\"Accuracy (%)\")\n    ax.set_title(\"Accuracy per epoch\")\n    ax.legend()\n    ax.xaxis.set_major_locator(ticker.MaxNLocator(integer=True))\n    ax.grid(True, linestyle=\":\", alpha=0.6)\n \n    plt.tight_layout()\n    out = f\"training_curves_{run_name}.png\"\n    plt.savefig(out, dpi=150, bbox_inches=\"tight\")\n    plt.close()\n    print(f\"  Saved training curves   ->  {out}\")\n \n \ndef _plot_confusion_matrix(labels: list, preds: list, run_name: str) -> None:\n    \"\"\"\n    Saves a side-by-side PNG:\n      Left  — raw counts\n      Right — row-normalised (each row sums to 1 = per-class recall)\n    Sensitivity and specificity are annotated below the figure.\n    Output file: confusion_matrix_{run_name}.png\n    \"\"\"\n    cm      = confusion_matrix(labels, preds)\n    cm_norm = confusion_matrix(labels, preds, normalize=\"true\")\n \n    fig, axes = plt.subplots(1, 2, figsize=(11, 4.5))\n    fig.suptitle(f\"Confusion matrix  —  {run_name}\", fontsize=14, fontweight=\"bold\")\n \n    class_names = [\"Normal\", \"Pneumonia\"]\n \n    ConfusionMatrixDisplay(cm, display_labels=class_names).plot(\n        ax=axes[0], colorbar=False, cmap=\"Blues\"\n    )\n    axes[0].set_title(\"Raw counts\")\n \n    ConfusionMatrixDisplay(cm_norm, display_labels=class_names).plot(\n        ax=axes[1], colorbar=False, cmap=\"Blues\", values_format=\".2%\"\n    )\n    axes[1].set_title(\"Normalised  (row = true class)\")\n \n    tn, fp, fn, tp = cm.ravel()\n    sensitivity = tp / (tp + fn) if (tp + fn) > 0 else 0.0\n    specificity = tn / (tn + fp) if (tn + fp) > 0 else 0.0\n    fig.text(\n        0.5, -0.05,\n        f\"Sensitivity (pneumonia recall): {sensitivity:.4f}    |    \"\n        f\"Specificity (normal recall): {specificity:.4f}\",\n        ha=\"center\", fontsize=11,\n    )\n \n    plt.tight_layout()\n    out = f\"confusion_matrix_{run_name}.png\"\n    plt.savefig(out, dpi=150, bbox_inches=\"tight\")\n    plt.close()\n    print(f\"  Saved confusion matrix  ->  {out}\")\n    print(f\"  Sensitivity : {sensitivity:.4f}   |   Specificity : {specificity:.4f}\")\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:40:07.858384Z","iopub.execute_input":"2026-04-14T08:40:07.859018Z","iopub.status.idle":"2026-04-14T08:40:07.872231Z","shell.execute_reply.started":"2026-04-14T08:40:07.858986Z","shell.execute_reply":"2026-04-14T08:40:07.871505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_and_evaluate(\n    train_df : pd.DataFrame,\n    val_df   : pd.DataFrame,\n    test_df  : pd.DataFrame,\n    augment  : bool = False,\n    run_name : str  = \"baseline\",\n) -> dict:\n \n    print(f\"\\n{'='*66}\")\n    print(f\"  Run  : {run_name}\")\n    print(f\"  Aug  : {augment}   |   Device : {DEVICE}   |   Epochs : {NUM_EPOCHS}\")\n    print(f\"{'='*66}\")\n \n    # ── Datasets & loaders ──\n    train_ds = RSNADataset(train_df, augment=augment)\n    val_ds   = RSNADataset(val_df,   augment=False)\n    test_ds  = RSNADataset(test_df,  augment=False)\n \n    train_loader = DataLoader(\n        train_ds, batch_size=BATCH_SIZE,\n        sampler=get_weighted_sampler(train_df),\n        num_workers=NUM_WORKERS, pin_memory=True,\n    )\n    val_loader = DataLoader(\n        val_ds, batch_size=BATCH_SIZE,\n        shuffle=False, num_workers=NUM_WORKERS, pin_memory=True,\n    )\n    test_loader = DataLoader(\n        test_ds, batch_size=BATCH_SIZE,\n        shuffle=False, num_workers=NUM_WORKERS, pin_memory=True,\n    )\n \n    # ── Model, loss, optimiser ──\n    model     = build_model(num_classes=2)\n    criterion = nn.CrossEntropyLoss(weight=get_class_weights(train_df))\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY\n    )\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=NUM_EPOCHS, eta_min=1e-7\n    )\n \n    best_val_auc   = 0.0\n    best_ckpt_path = f\"best_vit_{run_name}.pth\"\n    history = {\n        \"train_loss\": [], \"val_loss\": [],\n        \"train_acc\" : [], \"val_acc\" : [],\n        \"train_auc\" : [], \"val_auc\" : [],\n        \"train_f1\"  : [], \"val_f1\"  : [],\n    }\n \n    # ── Per-epoch header ──\n    SEP = \"─\" * 90\n    print(f\"\\n{SEP}\")\n    print(\n        f\"{'Epoch':>6}  \"\n        f\"{'Tr Loss':>9}  {'Tr Acc':>7}  {'Tr AUC':>7}  {'Tr F1':>7}    \"\n        f\"{'Va Loss':>9}  {'Va Acc':>7}  {'Va AUC':>7}  {'Va F1':>7}  {'':>4}\"\n    )\n    print(SEP)\n \n    for epoch in range(1, NUM_EPOCHS + 1):\n \n        tr_loss, tr_acc, tr_auc, tr_f1, _, _, _ = run_epoch(\n            model, train_loader, criterion, optimizer, phase=\"train\"\n        )\n        va_loss, va_acc, va_auc, va_f1, _, _, _ = run_epoch(\n            model, val_loader, criterion, phase=\"val\"\n        )\n        scheduler.step()\n \n        history[\"train_loss\"].append(tr_loss)\n        history[\"val_loss\"].append(va_loss)\n        history[\"train_acc\"].append(tr_acc)\n        history[\"val_acc\"].append(va_acc)\n        history[\"train_auc\"].append(tr_auc)\n        history[\"val_auc\"].append(va_auc)\n        history[\"train_f1\"].append(tr_f1)\n        history[\"val_f1\"].append(va_f1)\n \n        is_best = va_auc > best_val_auc\n        flag    = \"  * best\" if is_best else \"\"\n \n        print(\n            f\"{epoch:>6}  \"\n            f\"{tr_loss:>9.4f}  {tr_acc*100:>6.2f}%  {tr_auc:>7.4f}  {tr_f1:>7.4f}    \"\n            f\"{va_loss:>9.4f}  {va_acc*100:>6.2f}%  {va_auc:>7.4f}  {va_f1:>7.4f}\"\n            f\"{flag}\"\n        )\n \n        if is_best:\n            best_val_auc = va_auc\n            torch.save(model.state_dict(), best_ckpt_path)\n \n    print(SEP)\n    print(f\"  Best model saved  ->  {best_ckpt_path}  (val AUC = {best_val_auc:.4f})\")\n \n    # ────────────────────────────────────────\n    # Test evaluation on best checkpoint\n    # ────────────────────────────────────────\n    print(f\"\\nLoading best checkpoint for test set evaluation...\")\n    model.load_state_dict(torch.load(best_ckpt_path, map_location=DEVICE))\n \n    te_loss, te_acc, te_auc, te_f1, te_labels, te_probs, te_preds = run_epoch(\n        model, test_loader, criterion, phase=\"test\"\n    )\n \n    print(f\"\\n{'='*54}\")\n    print(f\"  TEST RESULTS  ({run_name})\")\n    print(f\"{'='*54}\")\n    print(f\"  Loss        : {te_loss:.4f}\")\n    print(f\"  Accuracy    : {te_acc * 100:.2f}%\")\n    print(f\"  AUC-ROC     : {te_auc:.4f}\")\n    print(f\"  F1 (binary) : {te_f1:.4f}\")\n    print(f\"{'─'*54}\")\n    print(classification_report(\n        te_labels, te_preds,\n        target_names=[\"Normal\", \"Pneumonia\"],\n        digits=4,\n    ))\n \n    # ── Save plots ──\n    _plot_training_curves(history, run_name)\n    _plot_confusion_matrix(te_labels, te_preds, run_name)\n \n    return {\n        \"run_name\"   : run_name,\n        \"history\"    : history,\n        \"test_loss\"  : te_loss,\n        \"test_acc\"   : te_acc,\n        \"test_auc\"   : te_auc,\n        \"test_f1\"    : te_f1,\n        \"test_labels\": te_labels,\n        \"test_probs\" : te_probs,\n        \"test_preds\" : te_preds,\n    }\n \n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:40:12.055958Z","iopub.execute_input":"2026-04-14T08:40:12.056664Z","iopub.status.idle":"2026-04-14T08:40:12.069108Z","shell.execute_reply.started":"2026-04-14T08:40:12.056631Z","shell.execute_reply":"2026-04-14T08:40:12.068333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n \n    # 1. Build label dataframe\n    df = build_label_df(LABEL_CSV)\n \n    # 2. Split data\n    train_df, val_df, test_df = make_splits(df)\n \n    # 3. Run A — no augmentation\n    results_baseline = train_and_evaluate(\n        train_df, val_df, test_df,\n        augment=False, run_name=\"no_augmentation\",\n    )\n \n    # 4. Run B — with augmentation\n    results_augmented = train_and_evaluate(\n        train_df, val_df, test_df,\n        augment=True, run_name=\"augmented\",\n    )\n \n    # 5. Side-by-side comparison\n    print(f\"\\n{'='*66}\")\n    print(\"  FINAL COMPARISON  (best checkpoint evaluated on held-out test set)\")\n    print(f\"{'='*66}\")\n    print(f\"  {'Run':<24}  {'Loss':>6}  {'Acc':>7}  {'AUC':>7}  {'F1':>7}\")\n    print(f\"  {'─'*58}\")\n    for r in [results_baseline, results_augmented]:\n        print(\n            f\"  {r['run_name']:<24}  \"\n            f\"{r['test_loss']:>6.4f}  \"\n            f\"{r['test_acc'] * 100:>6.2f}%  \"\n            f\"{r['test_auc']:>7.4f}  \"\n            f\"{r['test_f1']:>7.4f}\"\n        )\n    print(f\"{'='*66}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-14T08:40:17.132149Z","iopub.execute_input":"2026-04-14T08:40:17.132836Z","iopub.status.idle":"2026-04-14T11:40:44.831310Z","shell.execute_reply.started":"2026-04-14T08:40:17.132801Z","shell.execute_reply":"2026-04-14T11:40:44.829882Z"}},"outputs":[],"execution_count":null}]}