{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"a5fbfca3-f89e-465e-bf03-2bdc18cecd39","cell_type":"markdown","source":"# ResNet50 + ViT Hybrid Backbone Benchmark\n### Per-dataset performance check (Aptos, DDR, IDRiD, Eyepacs, Messidor)\n\nThis notebook implements the **backbone selection stage** of your pipeline (the `Select best backbone Model` box in your diagram).\n\nIt builds a **hybrid feature encoder = ResNet50 + ViT-B/16**, then trains and evaluates it **separately on each of the 5 hospital datasets** (H1–H5). The datasets are **never mixed** — each one gets its own split, its own freshly-initialized classifier head, and its own metrics. At the end, all 5 results are collected into one comparison table/chart.\n\nThis version is wired to the **exact Kaggle paths you already have added** (confirmed via a full `os.walk` + CSV preview you ran):\n- **Aptos** → `datasets/mariaherrerot/aptos2019/` (pre-split train/val/test, `id_code,diagnosis`)\n- **DDR** → `datasets/mariaherrerot/ddrdataset/` (`DR_grading.csv`, `id_code,diagnosis`)\n- **IDRiD** → `datasets/mariaherrerot/idrid-dataset/` (`idrid_labels.csv`, `id_code,diagnosis`)\n- **Eyepacs** → `datasets/dreamer07/eyepacs/` (`trainLabels.csv/trainLabels.csv`, `image,level`)\n- **Messidor** → `datasets/hanhan2010/messidor/` (pre-split train/test, columns `Image,Id,Risk of macular edema`)\n\n**One thing to double check before trusting the results:** the Messidor CSV you have doesn't contain a `diagnosis` column — only `Image`, `Id`, and `Risk of macular edema`. This notebook assumes `Id` is the DR severity grade (Messidor's real scale is 0–3, i.e. 4 classes, not 5 like the others). Section 2.5 below prints the value counts of that column — if they don't look like a clean 0–3 grade, this particular Messidor re-upload may not be the right one for DR grading and should be swapped out.\n\n**No path should need editing anymore** — just run Section 1 → 2 → 2.5 (to sanity-check) → the rest, top to bottom.\n","metadata":{}},{"id":"7a8323c3-e49f-442b-a7c9-38f046e60514","cell_type":"markdown","source":"## 1. Environment setup (Kaggle)\n\nKaggle notebooks already ship with `torch`, `torchvision`, `scikit-learn`, `pandas`, `matplotlib`, `seaborn`, `tqdm` pre-installed, so no `pip install` is needed.\n\n1. **Settings (right sidebar) → Accelerator → GPU T4 x2 (or P100)**.\n2. **Settings → Internet → On** (needed once, to download ImageNet-pretrained ResNet50/ViT weights).\n3. Make sure all 5 datasets are added via **Add Input**.","metadata":{}},{"id":"cfee612d-af2f-4462-986a-4a247a36f14b","cell_type":"code","source":"# Lists exactly what Kaggle mounted under /kaggle/input -- copy these exact names\n# into CONFIG[\"datasets\"][...][\"root_folder\"] in Section 2.\nimport os\nfor d in sorted(os.listdir('/kaggle/input')):\n    print(d)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"1a2ab0f3-6298-491d-969b-21a34d8f787f","cell_type":"code","source":"import os\nimport copy\nimport time\nimport random\nimport json\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (accuracy_score, precision_recall_fscore_support,\n                              cohen_kappa_score, confusion_matrix, roc_auc_score)\nfrom PIL import Image\nfrom tqdm.auto import tqdm\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\nseed_everything(42)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a6c49880-2bf1-4700-ba77-c3944ac67565","cell_type":"markdown","source":"## 2. Dataset configuration (already filled in)\n\nAll 5 `root_folder` values below are the exact paths confirmed from your `os.walk` output — nothing to edit here. Run Section 2.5 right after this once to sanity-check everything still resolves (Kaggle occasionally changes a dataset's mount path if you re-add it).\n","metadata":{}},{"id":"a2c58992-08b8-4953-99d9-a767f44f0cf1","cell_type":"code","source":"INPUT_ROOT = \"/kaggle/input\"\n\nCONFIG = {\n    \"img_size\": 224,          # required by ViT-B/16 (patch16, 224 input)\n    \"batch_size\": 16,\n    \"num_epochs\": 20,\n    \"lr\": 1e-4,\n    \"weight_decay\": 1e-4,\n    \"val_split\": 0.15,        # only used for datasets that don't already ship a val split\n    \"test_split\": 0.15,       # only used for datasets that don't already ship a test split\n    \"patience\": 5,\n    \"num_workers\": 2,\n    \"output_dir\": \"/kaggle/working/backbone_results\",\n\n    \"datasets\": {\n\n        # -------- H1: APTOS-2019 (mariaherrerot/aptos2019) — pre-split train/val/test --------\n        \"Aptos\": {\n            \"type\": \"presplit\",\n            \"root_folder\": \"datasets/mariaherrerot/aptos2019\",\n            \"train_csv\": \"{root}/train_1.csv\",\n            \"train_dir\": \"{root}/train_images/train_images\",\n            \"val_csv\": \"{root}/valid.csv\",\n            \"val_dir\": \"{root}/val_images/val_images\",\n            \"test_csv\": \"{root}/test.csv\",\n            \"test_dir\": \"{root}/test_images/test_images\",\n            \"id_col\": \"id_code\",\n            \"label_col\": \"diagnosis\",\n            \"ext\": \".png\",            # id_code has no extension in this csv -> we append .png\n            \"num_classes\": 5,\n        },\n\n        # -------- H2: DDR (mariaherrerot/ddrdataset) --------\n        \"DDR\": {\n            \"type\": \"csv_singlesplit\",   # single folder + csv -> we create our own 70/15/15 split\n            \"root_folder\": \"datasets/mariaherrerot/ddrdataset\",\n            \"csv_path\": \"{root}/DR_grading.csv\",\n            \"image_dir\": \"{root}/DR_grading/DR_grading\",\n            \"id_col\": \"id_code\",\n            \"label_col\": \"diagnosis\",\n            \"ext\": \"\",                # id_code already ends in .jpg, e.g. '20170413102628830.jpg'\n            \"num_classes\": 5,\n        },\n\n        # -------- H3: IDRiD (mariaherrerot/idrid-dataset) --------\n        \"IDRiD\": {\n            \"type\": \"csv_singlesplit\",\n            \"root_folder\": \"datasets/mariaherrerot/idrid-dataset\",\n            \"csv_path\": \"{root}/idrid_labels.csv\",\n            \"image_dir\": \"{root}/Imagenes/Imagenes\",\n            \"id_col\": \"id_code\",\n            \"label_col\": \"diagnosis\",\n            \"ext\": \".jpg\",             # id_code has no extension, e.g. 'IDRiD_001' -> 'IDRiD_001.jpg'\n            \"num_classes\": 5,\n        },\n\n        # -------- H4: Eyepacs (dreamer07/eyepacs) --------\n        \"Eyepacs\": {\n            \"type\": \"csv_singlesplit\",\n            \"root_folder\": \"datasets/dreamer07/eyepacs\",\n            \"csv_path\": \"{root}/trainLabels.csv/trainLabels.csv\",   # yes, a folder AND a file both named trainLabels.csv\n            \"image_dir\": \"{root}/data/data\",\n            \"id_col\": \"image\",\n            \"label_col\": \"level\",\n            \"ext\": \".jpeg\",            # image column has no extension, e.g. '10_left' -> '10_left.jpeg'\n            \"num_classes\": 5,\n        },\n\n        # -------- H5: Messidor (hanhan2010/messidor) — pre-split train/test --------\n        # NOTE: this particular re-upload's CSV has no 'diagnosis' column. Its only\n        # numeric column besides risk-of-edema is 'Id', which we use as the DR grade.\n        # Messidor's ORIGINAL grading scale is 0-3 (4 classes), not 0-4 like the others --\n        # double check this by running: pd.read_csv(train_csv)['Id'].value_counts()\n        # If the values don't look like a 0-3 grade after all, this dataset may need to\n        # be swapped for a different Messidor re-upload that ships a real 'diagnosis' column.\n        \"Messidor\": {\n            \"type\": \"presplit\",\n            \"root_folder\": \"datasets/hanhan2010/messidor\",\n            \"train_csv\": \"{root}/Messidor/train.csv\",\n            \"train_dir\": \"{root}/Messidor/train\",\n            \"val_csv\": None,          # no separate val split shipped -> carved out of train automatically\n            \"val_dir\": None,\n            \"test_csv\": \"{root}/Messidor/test.csv\",\n            \"test_dir\": \"{root}/Messidor/test\",\n            \"id_col\": \"Image\",\n            \"label_col\": \"Id\",\n            \"ext\": \"\",                 # Image column already ends in .tif\n            \"num_classes\": 4,\n        },\n    },\n}\n\nos.makedirs(CONFIG[\"output_dir\"], exist_ok=True)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"76a743af-36c9-47d0-81d7-1d3b39b66141","cell_type":"markdown","source":"## 2.5. Verify every path BEFORE training (run this every time you edit Section 2)","metadata":{}},{"id":"ed5c3a6d-0daa-4958-94ab-e71ef48d9be7","cell_type":"code","source":"def resolve(cfg, key):\n    val = cfg.get(key)\n    if val is None:\n        return None\n    root = os.path.join(INPUT_ROOT, cfg[\"root_folder\"])\n    return val.format(root=root)\n\n\ndef find_candidate_csvs(root_folder, max_results=10):\n    root = os.path.join(INPUT_ROOT, root_folder)\n    hits = []\n    if not os.path.isdir(root):\n        return hits\n    for dirpath, _, filenames in os.walk(root):\n        for fn in filenames:\n            if fn.lower().endswith(\".csv\"):\n                hits.append(os.path.join(dirpath, fn))\n                if len(hits) >= max_results:\n                    return hits\n    return hits\n\n\ndef verify_dataset(name, cfg):\n    print(\"=\" * 70)\n    print(name)\n    print(\"=\" * 70)\n\n    root = os.path.join(INPUT_ROOT, cfg[\"root_folder\"])\n    if cfg[\"root_folder\"] == \"REPLACE_ME\" or not os.path.isdir(root):\n        print(f\"  [FAIL] root_folder '{cfg['root_folder']}' not found under /kaggle/input.\")\n        print(\"         Run the '!ls /kaggle/input/' cell above and fix root_folder.\")\n        return\n\n    print(f\"  [OK] root resolved: {root}\")\n\n    csv_keys = [k for k in (\"csv_path\", \"train_csv\", \"val_csv\", \"test_csv\") if k in cfg]\n    dir_keys = [k for k in (\"image_dir\", \"train_dir\", \"val_dir\", \"test_dir\") if k in cfg]\n\n    any_csv_missing = False\n    for k in csv_keys:\n        path = resolve(cfg, k)\n        if path is None:\n            continue\n        if os.path.exists(path):\n            df = pd.read_csv(path)\n            print(f\"  [OK] {k}: {path}\")\n            print(f\"       columns: {list(df.columns)}  (rows: {len(df)})\")\n            print(f\"       head:\\n{df.head(2).to_string(index=False)}\")\n            if cfg.get(\"label_col\") in df.columns:\n                print(f\"       label_col '{cfg['label_col']}' value counts:\\n{df[cfg['label_col']].value_counts().sort_index().to_string()}\")\n        else:\n            any_csv_missing = True\n            print(f\"  [FAIL] {k} not found at: {path}\")\n\n    if any_csv_missing:\n        candidates = find_candidate_csvs(cfg[\"root_folder\"])\n        if candidates:\n            print(\"  -> CSV files that DO exist under this dataset (pick the right one and update the path):\")\n            for c in candidates:\n                print(f\"       {c}\")\n        else:\n            print(\"  -> No CSV files found anywhere under this dataset root.\")\n\n    for k in dir_keys:\n        path = resolve(cfg, k)\n        if path is None:\n            continue\n        if os.path.isdir(path):\n            n = sum(1 for f in os.listdir(path) if f.lower().endswith((\".jpg\", \".jpeg\", \".png\", \".tif\", \".tiff\")))\n            print(f\"  [OK] {k}: {path}  ({n} image files found)\")\n        else:\n            print(f\"  [FAIL] {k} not found at: {path}\")\n    print()\n\n\nfor name, cfg in CONFIG[\"datasets\"].items():\n    verify_dataset(name, cfg)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"3348a8bb-8676-488c-b271-6e5be6eb1e3e","cell_type":"markdown","source":"## 3. Dataset class, sample builders, and transforms","metadata":{}},{"id":"e614a323-bd9a-4a89-b7f3-cf37eb61a1ce","cell_type":"code","source":"class RetinaDataset(Dataset):\n    \"\"\"Generic (filepath, label) dataset.\"\"\"\n\n    def __init__(self, samples, transform=None):\n        self.samples = samples\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        path, label = self.samples[idx]\n        img = Image.open(path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        return img, label\n\n\ndef samples_from_csv(csv_path, image_dir, id_col, label_col, ext):\n    \"\"\"Reads a labels CSV + image folder into a list of (filepath, label).\"\"\"\n    if not csv_path or not os.path.exists(csv_path):\n        return []\n    if not image_dir or not os.path.isdir(image_dir):\n        return []\n    df = pd.read_csv(csv_path)\n    samples = []\n    valid_ext = (\".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\")\n    for _, row in df.iterrows():\n        img_id = str(row[id_col]).strip()\n        if ext and not img_id.lower().endswith(valid_ext):\n            img_id = img_id + ext\n        fpath = os.path.join(image_dir, img_id)\n        if os.path.exists(fpath):\n            samples.append((fpath, int(row[label_col])))\n    return samples\n\n\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\ndef get_transforms(img_size):\n    train_tf = transforms.Compose([\n        transforms.Resize((img_size, img_size)),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.RandomVerticalFlip(p=0.2),\n        transforms.RandomRotation(20),\n        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n        transforms.ToTensor(),\n        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n    eval_tf = transforms.Compose([\n        transforms.Resize((img_size, img_size)),\n        transforms.ToTensor(),\n        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n    return train_tf, eval_tf\n\n\ndef make_dataloaders(name, cfg, config):\n    train_tf, eval_tf = get_transforms(config[\"img_size\"])\n\n    if cfg[\"type\"] == \"presplit\":\n        # Dataset already ships its own train / (val) / test split -- use it as-is,\n        # never re-mixed with other hospitals' data.\n        train_samples = samples_from_csv(resolve(cfg, \"train_csv\"), resolve(cfg, \"train_dir\"),\n                                          cfg[\"id_col\"], cfg[\"label_col\"], cfg[\"ext\"])\n        if len(train_samples) == 0:\n            print(f\"[{name}] no train samples resolved -> skipping this dataset.\")\n            return None\n\n        test_samples = samples_from_csv(resolve(cfg, \"test_csv\"), resolve(cfg, \"test_dir\"),\n                                         cfg[\"id_col\"], cfg[\"label_col\"], cfg[\"ext\"])\n\n        if cfg.get(\"val_csv\") and cfg.get(\"val_dir\"):\n            val_samples = samples_from_csv(resolve(cfg, \"val_csv\"), resolve(cfg, \"val_dir\"),\n                                            cfg[\"id_col\"], cfg[\"label_col\"], cfg[\"ext\"])\n        else:\n            # No shipped val split -> carve one out of train (stratified)\n            labels = [s[1] for s in train_samples]\n            train_samples, val_samples = train_test_split(\n                train_samples, test_size=config[\"val_split\"], stratify=labels, random_state=42)\n\n        if len(test_samples) == 0:\n            # No usable test csv/dir -> carve test out of train as well\n            labels = [s[1] for s in train_samples]\n            train_samples, test_samples = train_test_split(\n                train_samples, test_size=config[\"test_split\"], stratify=labels, random_state=42)\n\n    elif cfg[\"type\"] == \"csv_singlesplit\":\n        # One folder + one csv -> we create our own stratified 70/15/15 split.\n        samples = samples_from_csv(resolve(cfg, \"csv_path\"), resolve(cfg, \"image_dir\"),\n                                    cfg[\"id_col\"], cfg[\"label_col\"], cfg[\"ext\"])\n        if len(samples) == 0:\n            print(f\"[{name}] no samples resolved -> skipping this dataset.\")\n            return None\n        labels = [s[1] for s in samples]\n        train_samples, temp_samples, train_labels, temp_labels = train_test_split(\n            samples, labels, test_size=config[\"val_split\"] + config[\"test_split\"],\n            stratify=labels, random_state=42)\n        rel_test = config[\"test_split\"] / (config[\"val_split\"] + config[\"test_split\"])\n        val_samples, test_samples = train_test_split(\n            temp_samples, test_size=rel_test, stratify=temp_labels, random_state=42)\n\n    elif cfg[\"type\"] == \"imagefolder\":\n        root = Path(resolve(cfg, \"root_dir\") if \"root_dir\" in cfg else os.path.join(INPUT_ROOT, cfg[\"root_folder\"]))\n        if not root.exists():\n            print(f\"[{name}] folder not found at {root} -> skipping this dataset.\")\n            return None\n        samples = []\n        class_dirs = sorted([d for d in root.iterdir() if d.is_dir()], key=lambda d: d.name)\n        for label, cdir in enumerate(class_dirs):\n            for fpath in cdir.glob(\"*\"):\n                if fpath.suffix.lower() in [\".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\"]:\n                    samples.append((str(fpath), label))\n        if len(samples) == 0:\n            print(f\"[{name}] no samples found -> skipping this dataset.\")\n            return None\n        labels = [s[1] for s in samples]\n        train_samples, temp_samples, train_labels, temp_labels = train_test_split(\n            samples, labels, test_size=config[\"val_split\"] + config[\"test_split\"],\n            stratify=labels, random_state=42)\n        rel_test = config[\"test_split\"] / (config[\"val_split\"] + config[\"test_split\"])\n        val_samples, test_samples = train_test_split(\n            temp_samples, test_size=rel_test, stratify=temp_labels, random_state=42)\n\n    else:\n        raise ValueError(f\"Unknown dataset type: {cfg['type']}\")\n\n    train_ds = RetinaDataset(train_samples, train_tf)\n    val_ds = RetinaDataset(val_samples, eval_tf)\n    test_ds = RetinaDataset(test_samples, eval_tf)\n\n    loaders = {\n        \"train\": DataLoader(train_ds, batch_size=config[\"batch_size\"], shuffle=True,\n                             num_workers=config[\"num_workers\"], pin_memory=True),\n        \"val\": DataLoader(val_ds, batch_size=config[\"batch_size\"], shuffle=False,\n                           num_workers=config[\"num_workers\"], pin_memory=True),\n        \"test\": DataLoader(test_ds, batch_size=config[\"batch_size\"], shuffle=False,\n                            num_workers=config[\"num_workers\"], pin_memory=True),\n    }\n    print(f\"[{name}] samples -> train={len(train_ds)}  val={len(val_ds)}  test={len(test_ds)}\")\n    return loaders\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c996ec3c-5059-46fd-a9e1-5d582f4ad023","cell_type":"markdown","source":"## 4. Hybrid feature encoder: ResNet50 + ViT-B/16\n\n- **ResNet50 branch**: convolutional, strong at local texture (microaneurysms, hemorrhages, exudates).\n- **ViT-B/16 branch**: transformer, strong at global/long-range structure (vessel layout, optic disc position).\n- Their pooled features (2048-d + 768-d) are concatenated → fused → classified.\n","metadata":{}},{"id":"1306634f-1bd4-441c-b01a-57923d95720f","cell_type":"code","source":"class ResNetViTHybrid(nn.Module):\n    def __init__(self, num_classes, dropout=0.3, freeze_backbones=False):\n        super().__init__()\n\n        resnet = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\n        self.resnet_features = nn.Sequential(*list(resnet.children())[:-1])  # -> (B, 2048, 1, 1)\n        resnet_out_dim = 2048\n\n        vit = models.vit_b_16(weights=models.ViT_B_16_Weights.IMAGENET1K_V1)\n        self.vit_backbone = vit\n        self.vit_backbone.heads = nn.Identity()  # -> (B, 768) CLS embedding\n        vit_out_dim = 768\n\n        if freeze_backbones:\n            for p in self.resnet_features.parameters():\n                p.requires_grad = False\n            for p in self.vit_backbone.parameters():\n                p.requires_grad = False\n\n        fused_dim = resnet_out_dim + vit_out_dim\n        self.classifier = nn.Sequential(\n            nn.LayerNorm(fused_dim),\n            nn.Dropout(dropout),\n            nn.Linear(fused_dim, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout),\n            nn.Linear(512, num_classes),\n        )\n\n    def forward(self, x):\n        r = self.resnet_features(x).flatten(1)   # (B, 2048)\n        v = self.vit_backbone(x)                  # (B, 768)\n        fused = torch.cat([r, v], dim=1)\n        return self.classifier(fused)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"aa154c40-dfde-4029-b45e-31875c8cb7ba","cell_type":"markdown","source":"## 5. Training and evaluation loops","metadata":{}},{"id":"77867400-e783-41ec-9fb2-677b4b0c5e3a","cell_type":"code","source":"def train_one_dataset(name, loaders, num_classes, config):\n    model = ResNetViTHybrid(num_classes=num_classes).to(DEVICE)\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=config[\"lr\"], weight_decay=config[\"weight_decay\"])\n    scheduler = CosineAnnealingLR(optimizer, T_max=config[\"num_epochs\"])\n\n    best_val_kappa = -1.0\n    best_state = None\n    epochs_no_improve = 0\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_acc\": [], \"val_kappa\": []}\n\n    for epoch in range(config[\"num_epochs\"]):\n        model.train()\n        running_loss = 0.0\n        for imgs, labels in tqdm(loaders[\"train\"], desc=f\"[{name}] epoch {epoch+1}/{config['num_epochs']} (train)\", leave=False):\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item() * imgs.size(0)\n        train_loss = running_loss / len(loaders[\"train\"].dataset)\n\n        model.eval()\n        val_loss = 0.0\n        all_preds, all_labels = [], []\n        with torch.no_grad():\n            for imgs, labels in loaders[\"val\"]:\n                imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n                outputs = model(imgs)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item() * imgs.size(0)\n                preds = outputs.argmax(dim=1)\n                all_preds.extend(preds.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n        val_loss /= len(loaders[\"val\"].dataset)\n        val_acc = accuracy_score(all_labels, all_preds)\n        val_kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n        scheduler.step()\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_acc\"].append(val_acc)\n        history[\"val_kappa\"].append(val_kappa)\n\n        print(f\"[{name}] epoch {epoch+1}: train_loss={train_loss:.4f} val_loss={val_loss:.4f} \"\n              f\"val_acc={val_acc:.4f} val_kappa={val_kappa:.4f}\")\n\n        if val_kappa > best_val_kappa:\n            best_val_kappa = val_kappa\n            best_state = copy.deepcopy(model.state_dict())\n            epochs_no_improve = 0\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= config[\"patience\"]:\n                print(f\"[{name}] early stopping at epoch {epoch+1}\")\n                break\n\n    model.load_state_dict(best_state)\n    return model, history\n\n\ndef evaluate_on_test(name, model, loaders, num_classes):\n    model.eval()\n    all_preds, all_labels, all_probs = [], [], []\n    with torch.no_grad():\n        for imgs, labels in loaders[\"test\"]:\n            imgs = imgs.to(DEVICE)\n            outputs = model(imgs)\n            probs = torch.softmax(outputs, dim=1).cpu().numpy()\n            preds = outputs.argmax(dim=1).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.numpy())\n            all_probs.extend(probs)\n\n    all_preds = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    all_probs = np.array(all_probs)\n\n    acc = accuracy_score(all_labels, all_preds)\n    precision, recall, f1, _ = precision_recall_fscore_support(\n        all_labels, all_preds, average=\"macro\", zero_division=0)\n    kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n    cm = confusion_matrix(all_labels, all_preds, labels=list(range(num_classes)))\n\n    try:\n        auc = roc_auc_score(all_labels, all_probs, multi_class=\"ovr\", average=\"macro\")\n    except ValueError:\n        auc = float(\"nan\")\n\n    metrics = {\n        \"dataset\": name,\n        \"accuracy\": acc,\n        \"precision_macro\": precision,\n        \"recall_macro\": recall,\n        \"f1_macro\": f1,\n        \"quadratic_kappa\": kappa,\n        \"auc_macro_ovr\": auc,\n    }\n\n    fig, ax = plt.subplots(figsize=(5, 4))\n    sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", ax=ax)\n    ax.set_title(f\"Confusion Matrix — {name}\")\n    ax.set_xlabel(\"Predicted\")\n    ax.set_ylabel(\"Actual\")\n    plt.tight_layout()\n    plt.savefig(os.path.join(CONFIG[\"output_dir\"], f\"confusion_matrix_{name}.png\"))\n    plt.show()\n\n    return metrics\n","metadata":{},"outputs":[],"execution_count":null},{"id":"11d7da80-067d-4af7-94c3-43e0521e3eb2","cell_type":"markdown","source":"## 6. Run each dataset separately\n\nOne dataset at a time — its own split, its own model instance, its own metrics. Nothing is combined. If `verify_dataset` printed a `[FAIL]` for something above, fix that before running this cell, or that dataset will simply be skipped.\n","metadata":{}},{"id":"e3f3d8aa-6ceb-46bf-817b-517352ca1723","cell_type":"code","source":"all_results = []\nall_histories = {}\n\nfor name, cfg in CONFIG[\"datasets\"].items():\n    print(\"=\" * 70)\n    print(f\"Running backbone benchmark on: {name}\")\n    print(\"=\" * 70)\n\n    loaders = make_dataloaders(name, cfg, CONFIG)\n    if loaders is None:\n        continue  # dataset path not resolved yet -> skip cleanly\n\n    model, history = train_one_dataset(name, loaders, cfg[\"num_classes\"], CONFIG)\n    metrics = evaluate_on_test(name, model, loaders, cfg[\"num_classes\"])\n    all_results.append(metrics)\n    all_histories[name] = history\n\n    torch.save(model.state_dict(), os.path.join(CONFIG[\"output_dir\"], f\"resnet_vit_{name}.pt\"))\n    with open(os.path.join(CONFIG[\"output_dir\"], f\"history_{name}.json\"), \"w\") as f:\n        json.dump(history, f, indent=2)\n\n    del model\n    torch.cuda.empty_cache()\n\nprint(\"\\nFinished. Datasets actually evaluated:\", [r['dataset'] for r in all_results])\n","metadata":{},"outputs":[],"execution_count":null},{"id":"8e40acf1-0548-4c66-b277-f017e1b4a8c8","cell_type":"markdown","source":"## 7. Compare backbone performance across the 5 datasets","metadata":{}},{"id":"2f2acdcd-1cf2-4d4b-8547-43631fb7be53","cell_type":"code","source":"results_df = pd.DataFrame(all_results).set_index(\"dataset\")\nresults_df = results_df.round(4)\nresults_df.to_csv(os.path.join(CONFIG[\"output_dir\"], \"backbone_comparison_all_datasets.csv\"))\nresults_df\n","metadata":{},"outputs":[],"execution_count":null},{"id":"dbb12801-b232-47d1-adae-ddb4e1a402aa","cell_type":"code","source":"metrics_to_plot = [\"accuracy\", \"f1_macro\", \"quadratic_kappa\", \"auc_macro_ovr\"]\nfig, axes = plt.subplots(1, len(metrics_to_plot), figsize=(20, 4))\nfor ax, metric in zip(axes, metrics_to_plot):\n    results_df[metric].plot(kind=\"bar\", ax=ax, color=\"steelblue\")\n    ax.set_title(metric)\n    ax.set_ylim(0, 1)\n    ax.tick_params(axis='x', rotation=45)\nplt.suptitle(\"ResNet50 + ViT hybrid — performance per dataset (evaluated separately)\")\nplt.tight_layout()\nplt.savefig(os.path.join(CONFIG[\"output_dir\"], \"backbone_comparison_chart.png\"))\nplt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"752fd7d3-b532-4704-9ab4-40e6c3811497","cell_type":"code","source":"fig, ax = plt.subplots(figsize=(8, 5))\nfor name, hist in all_histories.items():\n    ax.plot(hist[\"val_kappa\"], label=name)\nax.set_xlabel(\"Epoch\")\nax.set_ylabel(\"Validation Quadratic Kappa\")\nax.set_title(\"Validation kappa curves per dataset\")\nax.legend()\nplt.tight_layout()\nplt.savefig(os.path.join(CONFIG[\"output_dir\"], \"val_kappa_curves.png\"))\nplt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"d18b00c2-c6a1-420e-a802-cdbcdb40dc60","cell_type":"markdown","source":"## 8. Reading the results\n\n- **`accuracy`** / **`f1_macro`**: standard classification quality; macro-F1 matters more than accuracy since DR grades are usually imbalanced.\n- **`quadratic_kappa`**: the metric most DR-grading papers optimize for — penalizes far-off misgrades more than adjacent ones.\n- **`auc_macro_ovr`**: one-vs-rest AUC, robust to class imbalance.\n- Compare across H1–H5: a backbone that performs well on one hospital but poorly on another is a sign of **domain shift**, exactly what your federated `pFedH2A` stage is meant to handle later.\n\nAll artifacts (per-dataset model weights, training histories, confusion matrices, comparison CSV/chart) are saved under `CONFIG[\"output_dir\"]`.\n","metadata":{}}]}