{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":13451,"datasetId":654585,"databundleVersionId":1188070,"isSourceIdPinned":false}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"header","cell_type":"markdown","source":"# Module 2 — Model Training & Evaluation\n## RSNA Intracranial Hemorrhage Detection\n\n**Pipeline:**\n- ResNet-50 (CNN baseline)\n- ViT-B/16 (Vision Transformer)\n- Binary Cross-Entropy Loss + Adam Optimizer\n- 5-Fold Stratified Cross-Validation\n- Metrics: AUC-ROC, F1, Sensitivity, Specificity\n- Saves best model weights per fold + final metrics report\n\n> **Note:** Enable GPU accelerator in Kaggle (Settings → Accelerator → GPU P100/T4) before running.","metadata":{}},{"id":"cell-install","cell_type":"markdown","source":"## Cell 1 — Install Dependencies","metadata":{}},{"id":"install","cell_type":"code","source":"!pip install pydicom nibabel SimpleITK timm -q","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-imports","cell_type":"markdown","source":"## Cell 2 — Imports","metadata":{}},{"id":"imports","cell_type":"code","source":"import os\nimport copy\nimport time\nimport json\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport pydicom\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler, Subset\nfrom torchvision import transforms, models\n\nimport timm  # for ViT-B/16\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import (\n    roc_auc_score, f1_score, confusion_matrix,\n    classification_report, roc_curve\n)\n\n# Reproducibility\nSEED = 42\ntorch.manual_seed(SEED)\nnp.random.seed(SEED)\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: {DEVICE}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-config","cell_type":"markdown","source":"## Cell 3 — Configuration","metadata":{}},{"id":"config","cell_type":"code","source":"# ─── PATHS (same as Module 1) ───────────────────────────────────────────────\nBASE        = '/kaggle/input/rsna-intracranial-hemorrhage-detection'\nTRAIN_DIR   = os.path.join(BASE, 'stage_2_train')\nCSV_PATH    = os.path.join(BASE, 'stage_2_train.csv')\nOUTPUT_DIR  = '/kaggle/working/module2_outputs'\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\n# ─── HYPERPARAMETERS ────────────────────────────────────────────────────────\nCFG = {\n    'img_size'    : 224,\n    'batch_size'  : 32,\n    'num_workers' : 2,\n    'num_epochs'  : 10,       # increase to 20-30 for full training\n    'lr'          : 1e-4,\n    'weight_decay': 1e-5,\n    'n_folds'     : 5,\n    'patience'    : 3,        # early stopping patience\n    'models'      : ['resnet50', 'vit'],\n}\n\nprint(\"Config loaded:\", CFG)","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-preproc","cell_type":"markdown","source":"## Cell 4 — Preprocessing Helpers (from Module 1)","metadata":{}},{"id":"preproc","cell_type":"code","source":"def apply_windowing(image, window_center=40, window_width=80):\n    \"\"\"Apply brain CT windowing to focus on hemorrhage-relevant HU range.\"\"\"\n    lower = window_center - (window_width // 2)\n    upper = window_center + (window_width // 2)\n    image = np.clip(image, lower, upper)\n    image = (image - lower) / window_width\n    return image.astype(np.float32)\n\n\ndef load_dicom_slice(filepath):\n    \"\"\"Load a DICOM file and return a windowed float32 array.\"\"\"\n    dcm   = pydicom.dcmread(filepath)\n    image = dcm.pixel_array.astype(np.float32)\n    if hasattr(dcm, 'RescaleSlope'):\n        image = image * dcm.RescaleSlope + dcm.RescaleIntercept\n    image = apply_windowing(image)\n    return image","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-dataset","cell_type":"markdown","source":"## Cell 5 — Dataset Class","metadata":{}},{"id":"dataset","cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df        = df.reset_index(drop=True)\n        self.img_dir   = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row      = self.df.iloc[idx]\n        img_path = os.path.join(self.img_dir, row['ID'] + '.dcm')\n\n        # Load & convert to RGB PIL Image\n        image = load_dicom_slice(img_path)\n        image = Image.fromarray((image * 255).astype(np.uint8)).convert('RGB')\n\n        if self.transform:\n            image = self.transform(image)\n\n        label = torch.tensor(row['label'], dtype=torch.float32)\n        return image, label","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-dataprep","cell_type":"markdown","source":"## Cell 6 — Load & Prepare DataFrame","metadata":{}},{"id":"dataprep","cell_type":"code","source":"df = pd.read_csv(CSV_PATH)\nprint(\"Raw CSV shape:\", df.shape)\n\n# Parse ID and subtype\ndf[['ID', 'subtype']] = df['ID'].str.rsplit('_', n=1, expand=True)\n\n# Binary label: any hemorrhage present?\ndf_binary = df.groupby('ID')['Label'].max().reset_index()\ndf_binary.columns = ['ID', 'label']\n\nprint(\"Binary dataset shape:\", df_binary.shape)\nprint(\"\\nClass distribution:\")\nprint(df_binary['label'].value_counts())\nprint(f\"\\nPositive rate: {df_binary['label'].mean():.3f}\")\n\n# Optional: subsample for fast testing (comment out for full training)\n# df_binary = df_binary.sample(n=5000, random_state=SEED).reset_index(drop=True)\n\ndf_binary.head()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-transforms","cell_type":"markdown","source":"## Cell 7 — Transforms","metadata":{}},{"id":"transforms","cell_type":"code","source":"MEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\nIMG  = CFG['img_size']\n\ntrain_transforms = transforms.Compose([\n    transforms.Resize((IMG, IMG)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=MEAN, std=STD)\n])\n\nval_transforms = transforms.Compose([\n    transforms.Resize((IMG, IMG)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=MEAN, std=STD)\n])\n\nprint(\"Transforms defined.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-models","cell_type":"markdown","source":"## Cell 8 — Model Builders","metadata":{}},{"id":"model-builders","cell_type":"code","source":"def build_resnet50():\n    \"\"\"ResNet-50 pretrained on ImageNet, final FC → binary output.\"\"\"\n    model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)\n    in_features = model.fc.in_features\n    model.fc = nn.Sequential(\n        nn.Dropout(0.5),\n        nn.Linear(in_features, 1)\n    )\n    return model\n\n\ndef build_vit():\n    \"\"\"\n    ViT-B/16 via timm library, pretrained on ImageNet-21k.\n    Head replaced with Dropout + Linear → 1.\n    \"\"\"\n    model = timm.create_model(\n        'vit_base_patch16_224',\n        pretrained=True,\n        num_classes=0   # remove default head\n    )\n    embed_dim = model.embed_dim  # 768 for ViT-B\n    model.head = nn.Sequential(\n        nn.LayerNorm(embed_dim),\n        nn.Dropout(0.3),\n        nn.Linear(embed_dim, 1)\n    )\n    return model\n\n\ndef count_params(model):\n    total     = sum(p.numel() for p in model.parameters())\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Total params    : {total:,}\")\n    print(f\"  Trainable params: {trainable:,}\")\n\n\nprint(\"=== ResNet-50 ===\")\ncount_params(build_resnet50())\nprint(\"\\n=== ViT-B/16 ===\")\ncount_params(build_vit())","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-train-fn","cell_type":"markdown","source":"## Cell 9 — Training & Evaluation Functions","metadata":{}},{"id":"train-fn","cell_type":"code","source":"def train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_labels = [], []\n\n    for images, labels in loader:\n        images, labels = images.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images).squeeze(1)   # (B,)\n        loss    = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n        preds = torch.sigmoid(outputs).detach().cpu().numpy()\n        all_preds.extend(preds)\n        all_labels.extend(labels.cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_auc  = roc_auc_score(all_labels, all_preds)\n    return epoch_loss, epoch_auc\n\n\n@torch.no_grad()\ndef evaluate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds, all_labels = [], []\n\n    for images, labels in loader:\n        images, labels = images.to(device), labels.to(device)\n        outputs = model(images).squeeze(1)\n        loss    = criterion(outputs, labels)\n\n        running_loss += loss.item() * images.size(0)\n        preds = torch.sigmoid(outputs).cpu().numpy()\n        all_preds.extend(preds)\n        all_labels.extend(labels.cpu().numpy())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_auc  = roc_auc_score(all_labels, all_preds)\n    return epoch_loss, epoch_auc, np.array(all_preds), np.array(all_labels)\n\n\ndef compute_metrics(labels, probs, threshold=0.5):\n    \"\"\"Return dict with AUC-ROC, F1, Sensitivity, Specificity.\"\"\"\n    preds = (probs >= threshold).astype(int)\n    auc   = roc_auc_score(labels, probs)\n    f1    = f1_score(labels, preds)\n    cm    = confusion_matrix(labels, preds)\n    tn, fp, fn, tp = cm.ravel()\n    sensitivity = tp / (tp + fn + 1e-8)\n    specificity = tn / (tn + fp + 1e-8)\n    return {\n        'AUC-ROC'    : round(auc, 4),\n        'F1'         : round(f1, 4),\n        'Sensitivity': round(sensitivity, 4),\n        'Specificity': round(specificity, 4),\n        'TP': int(tp), 'FP': int(fp),\n        'TN': int(tn), 'FN': int(fn),\n    }\n\n\ndef get_sampler(labels):\n    \"\"\"WeightedRandomSampler to handle class imbalance.\"\"\"\n    class_counts  = np.bincount(labels.astype(int))\n    class_weights = 1.0 / class_counts\n    sample_weights = class_weights[labels.astype(int)]\n    return WeightedRandomSampler(\n        weights     = torch.DoubleTensor(sample_weights),\n        num_samples = len(sample_weights),\n        replacement = True\n    )\n\nprint(\"Training functions defined.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-cv-loop","cell_type":"markdown","source":"## Cell 10 — 5-Fold Cross-Validation Training Loop","metadata":{}},{"id":"cv-loop","cell_type":"code","source":"def run_cv(model_name, df_binary, train_dir, cfg, device):\n    \"\"\"\n    Run 5-fold stratified cross-validation for a given model.\n\n    Parameters\n    ----------\n    model_name : 'resnet50' or 'vit'\n    Returns    : fold_results (list of dicts), oof_probs, oof_labels\n    \"\"\"\n    skf          = StratifiedKFold(n_splits=cfg['n_folds'], shuffle=True, random_state=SEED)\n    labels_all   = df_binary['label'].values\n    fold_results = []\n    oof_probs    = np.zeros(len(df_binary))\n    oof_labels   = np.zeros(len(df_binary))\n\n    criterion = nn.BCEWithLogitsLoss()\n\n    for fold, (train_idx, val_idx) in enumerate(skf.split(df_binary, labels_all)):\n        print(f\"\\n{'='*60}\")\n        print(f\"  MODEL: {model_name.upper()}  |  FOLD {fold+1}/{cfg['n_folds']}\")\n        print(f\"{'='*60}\")\n\n        # ── Datasets ───────────────────────────────────────────────\n        train_df = df_binary.iloc[train_idx].reset_index(drop=True)\n        val_df   = df_binary.iloc[val_idx].reset_index(drop=True)\n\n        train_ds = RSNADataset(train_df, train_dir, transform=train_transforms)\n        val_ds   = RSNADataset(val_df,   train_dir, transform=val_transforms)\n\n        sampler      = get_sampler(train_df['label'].values)\n        train_loader = DataLoader(\n            train_ds, batch_size=cfg['batch_size'],\n            sampler=sampler, num_workers=cfg['num_workers'],\n            pin_memory=True\n        )\n        val_loader = DataLoader(\n            val_ds, batch_size=cfg['batch_size']*2,\n            shuffle=False, num_workers=cfg['num_workers'],\n            pin_memory=True\n        )\n\n        # ── Model ──────────────────────────────────────────────────\n        model = build_resnet50() if model_name == 'resnet50' else build_vit()\n        model = model.to(device)\n\n        optimizer = optim.Adam(\n            model.parameters(),\n            lr=cfg['lr'],\n            weight_decay=cfg['weight_decay']\n        )\n        scheduler = optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=cfg['num_epochs']\n        )\n\n        # ── Training loop ──────────────────────────────────────────\n        best_auc      = 0.0\n        best_weights  = None\n        patience_ctr  = 0\n        history       = {'train_loss': [], 'train_auc': [],\n                         'val_loss'  : [], 'val_auc'  : []}\n\n        for epoch in range(cfg['num_epochs']):\n            t0 = time.time()\n            tr_loss, tr_auc = train_one_epoch(\n                model, train_loader, criterion, optimizer, device\n            )\n            vl_loss, vl_auc, vl_probs, vl_labels = evaluate(\n                model, val_loader, criterion, device\n            )\n            scheduler.step()\n\n            history['train_loss'].append(tr_loss)\n            history['train_auc'].append(tr_auc)\n            history['val_loss'].append(vl_loss)\n            history['val_auc'].append(vl_auc)\n\n            elapsed = time.time() - t0\n            print(f\"  Epoch {epoch+1:02d}/{cfg['num_epochs']} \"\n                  f\"| TR loss: {tr_loss:.4f} AUC: {tr_auc:.4f} \"\n                  f\"| VAL loss: {vl_loss:.4f} AUC: {vl_auc:.4f} \"\n                  f\"| {elapsed:.1f}s\")\n\n            # Early stopping & best model checkpoint\n            if vl_auc > best_auc:\n                best_auc     = vl_auc\n                best_weights = copy.deepcopy(model.state_dict())\n                patience_ctr = 0\n                ckpt_path    = os.path.join(\n                    OUTPUT_DIR,\n                    f'{model_name}_fold{fold+1}_best.pth'\n                )\n                torch.save(best_weights, ckpt_path)\n                print(f\"  ✓ Saved best model → {ckpt_path} (AUC={best_auc:.4f})\")\n            else:\n                patience_ctr += 1\n                if patience_ctr >= cfg['patience']:\n                    print(f\"  Early stopping triggered at epoch {epoch+1}\")\n                    break\n\n        # ── Final evaluation with best weights ─────────────────────\n        model.load_state_dict(best_weights)\n        _, _, final_probs, final_labels = evaluate(\n            model, val_loader, criterion, device\n        )\n\n        # Store OOF predictions\n        oof_probs[val_idx]  = final_probs\n        oof_labels[val_idx] = final_labels\n\n        # Compute & store fold metrics\n        metrics = compute_metrics(final_labels, final_probs)\n        metrics['fold']    = fold + 1\n        metrics['history'] = history\n        fold_results.append(metrics)\n\n        print(f\"\\n  Fold {fold+1} Results:\")\n        for k, v in metrics.items():\n            if k not in ('fold', 'history'):\n                print(f\"    {k}: {v}\")\n\n        del model\n        torch.cuda.empty_cache()\n\n    return fold_results, oof_probs, oof_labels\n\n\nprint(\"CV function defined.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-run-resnet","cell_type":"markdown","source":"## Cell 11 — Train ResNet-50","metadata":{}},{"id":"run-resnet","cell_type":"code","source":"print(\"Starting ResNet-50 training...\\n\")\nresnet_results, resnet_oof_probs, resnet_oof_labels = run_cv(\n    model_name  = 'resnet50',\n    df_binary   = df_binary,\n    train_dir   = TRAIN_DIR,\n    cfg         = CFG,\n    device      = DEVICE\n)","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-run-vit","cell_type":"markdown","source":"## Cell 12 — Train ViT-B/16","metadata":{}},{"id":"run-vit","cell_type":"code","source":"print(\"Starting ViT-B/16 training...\\n\")\nvit_results, vit_oof_probs, vit_oof_labels = run_cv(\n    model_name = 'vit',\n    df_binary  = df_binary,\n    train_dir  = TRAIN_DIR,\n    cfg        = CFG,\n    device     = DEVICE\n)","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-metrics","cell_type":"markdown","source":"## Cell 13 — Aggregate Metrics & Report","metadata":{}},{"id":"metrics","cell_type":"code","source":"def aggregate_results(fold_results, model_name):\n    \"\"\"Print and return mean ± std across folds.\"\"\"\n    keys    = ['AUC-ROC', 'F1', 'Sensitivity', 'Specificity']\n    summary = {}\n\n    print(f\"\\n{'='*50}\")\n    print(f\"  {model_name} — Cross-Validation Summary\")\n    print(f\"{'='*50}\")\n\n    for k in keys:\n        vals = [r[k] for r in fold_results]\n        mean = np.mean(vals)\n        std  = np.std(vals)\n        summary[k] = {'mean': round(mean, 4), 'std': round(std, 4)}\n        print(f\"  {k:<15}: {mean:.4f} ± {std:.4f}\")\n\n    return summary\n\n\nresnet_summary = aggregate_results(resnet_results, 'ResNet-50')\nvit_summary    = aggregate_results(vit_results,    'ViT-B/16')\n\n# OOF (out-of-fold) aggregate metrics\nprint(\"\\n--- OOF Metrics (full dataset) ---\")\nprint(\"ResNet-50 OOF:\", compute_metrics(resnet_oof_labels, resnet_oof_probs))\nprint(\"ViT-B/16  OOF:\", compute_metrics(vit_oof_labels,    vit_oof_probs))","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-plots","cell_type":"markdown","source":"## Cell 14 — Visualisations","metadata":{}},{"id":"plots","cell_type":"code","source":"# ── 1. Training curves (last fold of each model) ───────────────────────────\nfig, axes = plt.subplots(2, 2, figsize=(16, 10))\nfig.suptitle('Training Curves — Last Fold', fontsize=14)\n\nfor col, (name, results) in enumerate(\n    [('ResNet-50', resnet_results), ('ViT-B/16', vit_results)]\n):\n    hist = results[-1]['history']\n    epochs = range(1, len(hist['train_loss']) + 1)\n\n    axes[0, col].plot(epochs, hist['train_loss'], 'b-o', label='Train')\n    axes[0, col].plot(epochs, hist['val_loss'],   'r-o', label='Val')\n    axes[0, col].set_title(f'{name} — Loss')\n    axes[0, col].set_xlabel('Epoch')\n    axes[0, col].set_ylabel('BCE Loss')\n    axes[0, col].legend()\n    axes[0, col].grid(True)\n\n    axes[1, col].plot(epochs, hist['train_auc'], 'b-o', label='Train')\n    axes[1, col].plot(epochs, hist['val_auc'],   'r-o', label='Val')\n    axes[1, col].set_title(f'{name} — AUC-ROC')\n    axes[1, col].set_xlabel('Epoch')\n    axes[1, col].set_ylabel('AUC-ROC')\n    axes[1, col].legend()\n    axes[1, col].grid(True)\n\nplt.tight_layout()\nplt.savefig(os.path.join(OUTPUT_DIR, 'training_curves.png'), dpi=150)\nplt.show()\n\n\n# ── 2. ROC Curves ──────────────────────────────────────────────────────────\nfig, ax = plt.subplots(1, 1, figsize=(8, 6))\n\nfor name, probs, labels in [\n    ('ResNet-50', resnet_oof_probs, resnet_oof_labels),\n    ('ViT-B/16',  vit_oof_probs,    vit_oof_labels)\n]:\n    fpr, tpr, _ = roc_curve(labels, probs)\n    auc_val     = roc_auc_score(labels, probs)\n    ax.plot(fpr, tpr, lw=2, label=f'{name} (AUC={auc_val:.4f})')\n\nax.plot([0, 1], [0, 1], 'k--', lw=1)\nax.set_xlabel('False Positive Rate')\nax.set_ylabel('True Positive Rate')\nax.set_title('OOF ROC Curves')\nax.legend(loc='lower right')\nax.grid(True)\nplt.tight_layout()\nplt.savefig(os.path.join(OUTPUT_DIR, 'roc_curves.png'), dpi=150)\nplt.show()\n\n\n# ── 3. Confusion Matrices ──────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 2, figsize=(12, 5))\nfig.suptitle('OOF Confusion Matrices (threshold=0.5)')\n\nfor ax, (name, probs, labels) in zip(axes, [\n    ('ResNet-50', resnet_oof_probs, resnet_oof_labels),\n    ('ViT-B/16',  vit_oof_probs,    vit_oof_labels)\n]):\n    preds = (probs >= 0.5).astype(int)\n    cm    = confusion_matrix(labels, preds)\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=ax,\n                xticklabels=['No Bleed', 'Bleed'],\n                yticklabels=['No Bleed', 'Bleed'])\n    ax.set_title(name)\n    ax.set_xlabel('Predicted')\n    ax.set_ylabel('Actual')\n\nplt.tight_layout()\nplt.savefig(os.path.join(OUTPUT_DIR, 'confusion_matrices.png'), dpi=150)\nplt.show()\n\n\n# ── 4. Per-fold AUC comparison ─────────────────────────────────────────────\nfig, ax = plt.subplots(figsize=(9, 5))\nfolds = [r['fold'] for r in resnet_results]\nrn_aucs = [r['AUC-ROC'] for r in resnet_results]\nvt_aucs = [r['AUC-ROC'] for r in vit_results]\n\nx = np.arange(len(folds))\nw = 0.35\nax.bar(x - w/2, rn_aucs, w, label='ResNet-50', color='steelblue')\nax.bar(x + w/2, vt_aucs, w, label='ViT-B/16',  color='coral')\nax.set_xticks(x)\nax.set_xticklabels([f'Fold {f}' for f in folds])\nax.set_ylabel('AUC-ROC')\nax.set_title('Per-Fold AUC Comparison')\nax.legend()\nax.set_ylim(0.5, 1.0)\nax.grid(axis='y', alpha=0.4)\nplt.tight_layout()\nplt.savefig(os.path.join(OUTPUT_DIR, 'fold_auc_comparison.png'), dpi=150)\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-report","cell_type":"markdown","source":"## Cell 15 — Save Full Metrics Report (JSON + CSV)","metadata":{}},{"id":"report","cell_type":"code","source":"# ── Build report dict ──────────────────────────────────────────────────────\nreport = {\n    'ResNet-50': {\n        'fold_metrics' : [\n            {k: v for k, v in r.items() if k != 'history'}\n            for r in resnet_results\n        ],\n        'cv_summary'   : resnet_summary,\n        'oof_metrics'  : compute_metrics(resnet_oof_labels, resnet_oof_probs),\n    },\n    'ViT-B/16': {\n        'fold_metrics' : [\n            {k: v for k, v in r.items() if k != 'history'}\n            for r in vit_results\n        ],\n        'cv_summary'   : vit_summary,\n        'oof_metrics'  : compute_metrics(vit_oof_labels, vit_oof_probs),\n    }\n}\n\n# ── Save JSON ──────────────────────────────────────────────────────────────\njson_path = os.path.join(OUTPUT_DIR, 'metrics_report.json')\nwith open(json_path, 'w') as f:\n    json.dump(report, f, indent=2)\nprint(f\"JSON report saved → {json_path}\")\n\n# ── Save CSV summary ───────────────────────────────────────────────────────\nrows = []\nfor model_name, data in report.items():\n    for fold_m in data['fold_metrics']:\n        row = {'model': model_name}\n        row.update(fold_m)\n        rows.append(row)\n\ndf_report = pd.DataFrame(rows)\ncsv_path  = os.path.join(OUTPUT_DIR, 'metrics_report.csv')\ndf_report.to_csv(csv_path, index=False)\nprint(f\"CSV report  saved → {csv_path}\")\n\nprint(\"\\n=== FINAL REPORT ===\")\nprint(df_report.to_string(index=False))","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-ensemble","cell_type":"markdown","source":"## Cell 16 — Optional: Ensemble (Average Probabilities)","metadata":{}},{"id":"ensemble","cell_type":"code","source":"# Simple averaging ensemble of both models' OOF predictions\nensemble_probs = (resnet_oof_probs + vit_oof_probs) / 2\nensemble_labels = resnet_oof_labels  # same ground truth\n\nens_metrics = compute_metrics(ensemble_labels, ensemble_probs)\nprint(\"\\n=== Ensemble (ResNet-50 + ViT-B/16) OOF Metrics ===\")\nfor k, v in ens_metrics.items():\n    print(f\"  {k}: {v}\")\n\n# Add ensemble to report\nreport['Ensemble'] = {'oof_metrics': ens_metrics}\nwith open(json_path, 'w') as f:\n    json.dump(report, f, indent=2)\nprint(\"\\nEnsemble metrics added to JSON report.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"cell-outputs","cell_type":"markdown","source":"## Cell 17 — List Output Files","metadata":{}},{"id":"list-outputs","cell_type":"code","source":"print(\"Files saved in\", OUTPUT_DIR)\nfor f in sorted(os.listdir(OUTPUT_DIR)):\n    fp   = os.path.join(OUTPUT_DIR, f)\n    size = os.path.getsize(fp) / 1024\n    print(f\"  {f:<45} {size:.1f} KB\")","metadata":{},"outputs":[],"execution_count":null}]}