{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"3fc9d656-a1f4-4fab-beca-65267c746f96","cell_type":"markdown","source":"# Binary DR Detection (No-DR vs DR) — dual-branch **EfficientNet-B4 + ResNet50 + SE**\n\nTarget: **beat 98.50% / 99.46% sens / 97.51% spec** (Shakibania et al. 2024, Table 4).\n\n**What this notebook does differently from the from-scratch runs (best so far 97.41 ± 0.14%):**\n\n| Lever | Why it is here |\n|---|---|\n| **ImageNet-pretrained backbones** | The paper's 98.50% uses transfer learning. From scratch peaked at 97.41%. This is the single biggest difference — see the constraint warning in section 1. |\n| **5-fold CV + out-of-fold threshold** | The threshold has been tuned on 366 val images and drifted 0.57→0.74 between runs. Pooled OOF gives ~2,900 images to tune it on. |\n| **Auxiliary 5-class grade head** | Measured best config in the earlier work. `1 − P(grade 0)` is a second opinion from the same forward pass, fused with a val-tuned weight. |\n| **Pseudo-labelling the 1,928 unlabelled APTOS images** | Same hospital and cameras as the test set; the two best earlier single runs were both students. |\n| **EMA + 4-way flip TTA** | Both measured to help previously. |\n\nTest set: 20% of APTOS, fixed by `SPLIT_SEED`, identical in every fold, never used for any choice.\n**Kaggle: GPU T4, Internet ON.** `SESSION_MAX_HOURS = 9.5` leaves headroom under the 12 h kill.","metadata":{}},{"id":"89c68dcf-9880-4783-a5c1-b7b9972292f9","cell_type":"markdown","source":"## 0. Setup","metadata":{}},{"id":"6b729c6a-4289-40c2-90cb-5aab1207f090","cell_type":"code","source":"import importlib, subprocess, sys\ndef _ensure(pkg, pip_name=None):\n    try: importlib.import_module(pkg)\n    except ImportError: subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pip_name or pkg], check=False)\n_ensure(\"timm\"); _ensure(\"albumentations\"); _ensure(\"seaborn\")\n\nimport os, sys, math, time, json, glob, random, hashlib, zipfile, itertools, warnings\nfrom collections import Counter, defaultdict\nfrom pathlib import Path\nfrom multiprocessing import Pool\nimport numpy as np, pandas as pd, cv2\nimport matplotlib.pyplot as plt, seaborn as sns\nimport torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm, timm.data\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nfrom sklearn.metrics import (accuracy_score, roc_auc_score, roc_curve, confusion_matrix,\n                             average_precision_score, cohen_kappa_score, matthews_corrcoef, brier_score_loss)\nfrom tqdm.auto import tqdm\nwarnings.filterwarnings(\"ignore\")\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"torch\", torch.__version__, \"| timm\", timm.__version__, \"| device:\", device)\nif torch.cuda.is_available(): print(\"GPU:\", torch.cuda.get_device_name(0))\n_CU = torch.cuda.is_available()\ndef amp_autocast(e=True):\n    try: return torch.amp.autocast(\"cuda\", enabled=e and _CU)\n    except Exception: return torch.cuda.amp.autocast(enabled=e and _CU)\ndef make_scaler(e=True):\n    try: return torch.amp.GradScaler(\"cuda\", enabled=e and _CU)\n    except Exception: return torch.cuda.amp.GradScaler(enabled=e and _CU)\nT0 = time.time()","metadata":{},"outputs":[],"execution_count":null},{"id":"19c3d01b-f41a-4f9b-a353-97f8bc572eb1","cell_type":"markdown","source":"## 1. Configuration\n\n> **Constraint check.** `PRETRAINED = True` uses ImageNet weights. If the \"no pretrained weights anywhere\"\n> rule still applies to this deliverable, set it to `False` — but expect ~97.4%, not 98%+, and report it as\n> *from-scratch vs the paper's transfer learning*, which is a fair and more interesting comparison.\n> The notebook prints which mode it ran in and stamps it into every results file.","metadata":{}},{"id":"f06ce30e-72da-4315-a888-9783b55a8c74","cell_type":"code","source":"# ==================== PICK YOUR RUN ====================\nFOLD   = 0        # 0..4 -> one session each, then run section 10\nPSEUDO = False    # True for folds 1..4 AFTER fold 0 finished (its checkpoint becomes the teacher)\n# =======================================================\n\nclass CFG:\n    PRETRAINED   = True          # <-- see the constraint note above\n    SEED         = [42, 1337, 2024, 7, 99][FOLD]\n    SPLIT_SEED   = 42            # fixes the test set for every fold - never change\n    FOLD, N_FOLDS = FOLD, 5\n    TEST_FRAC    = 0.20\n\n    IMG_SIZE     = 384\n    PREPROCESS   = \"ben\"         # \"ben\" | \"rgb_clahe\"\n    CROP_THRESHOLD = 7\n\n    EFF_NAME, EFF_FB = \"tf_efficientnet_b4.ns_jft_in1k\", \"efficientnet_b4\"\n    RES_NAME, RES_FB = \"resnet50.a1_in1k\", \"resnet50\"\n    SE_REDUCTION = 8\n    DROPOUT      = 0.4\n    GRADE_AUX_W  = 0.4           # auxiliary 5-class head; 0 disables it\n    GRADE_FUSE   = True\n\n    DATA_POLICY  = \"selective\"   # \"selective\" = external grades 1,3,4 -> train | \"full\" = all grades\n    MERGE_GRADES = (1, 3, 4)\n\n    BATCH_SIZE   = 16\n    ACCUM_STEPS  = 2             # effective batch 32\n    NUM_WORKERS  = 4\n    EPOCHS       = 18            # transfer learning converges far faster than the from-scratch runs\n    WARMUP_EPOCHS = 1\n    LR_BB        = 1e-4\n    LR_HEAD      = 1e-3\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTH = 0.03\n    GRAD_CLIP    = 2.0\n    PATIENCE     = 6\n    EMA_DECAY    = 0.998\n    AMP          = True\n    USE_TTA      = True\n\n    PSEUDO         = PSEUDO\n    PSEUDO_CONF    = 0.97\n    PSEUDO_TEACHER = \"run_K0/best.pt\"\n\n    SESSION_MAX_HOURS = 9.5      # hard budget; Kaggle kills at 12 h\n    OUT = \"/kaggle/working\"\n\ncfg = CFG()\nRUN = f\"K{cfg.FOLD}\" + (\"_pseudo\" if cfg.PSEUDO else \"\")\nRUN_DIR = Path(cfg.OUT) / f\"run_{RUN}\"; RUN_DIR.mkdir(parents=True, exist_ok=True)\ndef set_seed(s):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s); torch.cuda.manual_seed_all(s)\nset_seed(cfg.SEED)\ntorch.backends.cudnn.benchmark = True\nPAPER = dict(accuracy=98.50, sensitivity=99.46, specificity=97.51, precision=97.61, auc=98.00)\njson.dump({k: v for k, v in vars(CFG).items() if not k.startswith(\"_\")} | {\"RUN\": RUN},\n          open(RUN_DIR / \"config.json\", \"w\"), indent=2, default=str)\nprint(f\"RUN={RUN} | fold {cfg.FOLD}/{cfg.N_FOLDS} | seed {cfg.SEED} | pseudo={cfg.PSEUDO}\")\nprint(\"INITIALISATION:\", \"ImageNet-pretrained (matches the paper's setting)\" if cfg.PRETRAINED\n      else \"RANDOM - from scratch (stricter than the paper; expect ~97.4%)\")\n\ndef _stats(n, fb):\n    try: m = timm.create_model(n, pretrained=False, num_classes=0)\n    except Exception: m = timm.create_model(fb, pretrained=False, num_classes=0)\n    d = timm.data.resolve_model_data_config(m); del m\n    return tuple(d[\"mean\"]), tuple(d[\"std\"])\nEFF_MEAN, EFF_STD = _stats(cfg.EFF_NAME, cfg.EFF_FB)\nRES_MEAN, RES_STD = _stats(cfg.RES_NAME, cfg.RES_FB)\ndef budget_left(): return cfg.SESSION_MAX_HOURS - (time.time() - T0) / 3600","metadata":{},"outputs":[],"execution_count":null},{"id":"73c146c6-f637-4364-93d7-d16627fd22ab","cell_type":"markdown","source":"## 2. Dataset discovery","metadata":{}},{"id":"e8f940ce-fc9c-4834-bb53-810c9341fb58","cell_type":"code","source":"INPUT_ROOT = \"/kaggle/input\" if os.path.isdir(\"/kaggle/input\") else \".\"\ndef _files(root, exts):\n    out = []\n    for dp, _, fns in os.walk(root, followlinks=True):\n        for fn in fns:\n            if fn.lower().endswith(exts): out.append(os.path.join(dp, fn))\n    return out\nall_csvs = _files(INPUT_ROOT, (\".csv\",))\nall_imgs = _files(INPUT_ROOT, (\".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\"))\nall_zips = _files(INPUT_ROOT, (\".zip\",))\nprint(f\"Scanned {INPUT_ROOT}: {len(all_csvs)} csv, {len(all_imgs)} images, {len(all_zips)} zips\\n\")\n_bf = defaultdict(int)\nfor p in all_imgs: _bf[os.path.relpath(p, INPUT_ROOT).split(os.sep)[0]] += 1\nfor k, v in sorted(_bf.items()): print(f\"  {k}: {v}\")\nif not _bf: print(\"WARNING: no images -> attach datasets via '+ Add Input'.\")\n\ndef build_index(paths):\n    d = {}\n    for p in paths: d.setdefault(os.path.splitext(os.path.basename(p))[0].lower(), p)\n    return d\nIMG_INDEX = build_index(all_imgs)\n\ndef _pick(cols, keys, bad=()):\n    low = {c.lower().strip(): c for c in cols}\n    for k in keys:\n        for lc, o in low.items():\n            if (k == lc or k in lc) and not any(b in lc for b in bad): return o\n    return None\n\ndef _resolve(df, source):\n    df = df.copy()\n    df[\"grade\"] = pd.to_numeric(df[\"grade\"], errors=\"coerce\")\n    df = df.dropna(subset=[\"grade\"]); df[\"grade\"] = df[\"grade\"].astype(int)\n    df = df[df[\"grade\"].between(0, 4)]\n    df[\"path\"] = df[\"image_id\"].map(lambda n: IMG_INDEX.get(os.path.splitext(os.path.basename(str(n).strip()))[0].lower()))\n    miss = int(df[\"path\"].isna().sum())\n    df = df.dropna(subset=[\"path\"]).drop_duplicates(\"path\").reset_index(drop=True)\n    df[\"source\"] = source\n    print(f\"{source}: {len(df)} usable ({miss} unmatched)  grades {dict(sorted(Counter(df.grade).items()))}\")\n    return df if len(df) else None","metadata":{},"outputs":[],"execution_count":null},{"id":"bc624596-22d5-4ed4-8c51-59304142beba","cell_type":"markdown","source":"### 2.1 — APTOS (labelled + unlabelled pool)","metadata":{}},{"id":"06b11f4b-e85b-4a21-8539-61608e889d9c","cell_type":"code","source":"frames = []\nfor p in all_csvs:\n    pl = p.lower()\n    if any(k in pl for k in [\"messidor\", \"idrid\", \"ddr\", \"eyepacs\", \"run_\", \"resized\"]) and \"aptos\" not in pl: continue\n    try: d = pd.read_csv(p)\n    except Exception: continue\n    d.columns = [c.strip().lower() for c in d.columns]\n    if not (any(\"id_code\" in c for c in d.columns) and any(\"diagnosis\" in c for c in d.columns)): continue\n    ic = [c for c in d.columns if \"id_code\" in c][0]; gc = [c for c in d.columns if \"diagnosis\" in c][0]\n    sub = d[[ic, gc]].dropna(); sub.columns = [\"image_id\", \"grade\"]\n    if sub[\"grade\"].nunique() < 2: continue          # unlabelled submission template\n    frames.append(sub); print(f\"  APTOS labels: {len(sub):5d} rows <- {p}\")\nassert frames, \"APTOS labels not found (need id_code + diagnosis).\"\naptos_df = _resolve(pd.concat(frames, ignore_index=True).drop_duplicates(\"image_id\"), \"APTOS\")\nif len(aptos_df) < 3500: print(f\"  [WARN] expected ~3662, matched {len(aptos_df)}\")\n\n# the 1,928 unlabelled competition test images, for pseudo-labelling only\nlab = set(aptos_df.path)\nUNLAB = sorted(p for p in all_imgs\n               if (\"aptos\" in p.lower() and \"test\" in p.lower() and p not in lab))\nprint(f\"\\nunlabelled APTOS pool: {len(UNLAB)} images\" + (\"\" if UNLAB else \"  (pseudo-labelling will be skipped)\"))","metadata":{},"outputs":[],"execution_count":null},{"id":"bdfcbf20-4724-4f2c-bd1f-cb84f97661da","cell_type":"markdown","source":"### 2.2 — IDRiD (both label CSVs) and Messidor-2","metadata":{}},{"id":"bb35292c-c4b7-4895-bd59-eec22aa3b236","cell_type":"code","source":"def load_idrid():\n    subs = []\n    for p in all_csvs:\n        if \"idrid\" not in p.lower() and \"disease grading\" not in os.path.basename(p).lower(): continue\n        try: d = pd.read_csv(p)\n        except Exception: continue\n        ic = _pick(d.columns, [\"image name\", \"image_name\", \"image\", \"file\", \"name\"])\n        gc = _pick(d.columns, [\"retinopathy grade\", \"retinopathy_grade\", \"grade\", \"diagnosis\", \"label\"])\n        if ic and gc:\n            s = d[[ic, gc]].dropna(); s.columns = [\"image_id\", \"grade\"]; subs.append(s)\n            print(f\"  IDRiD csv: {len(s):4d} rows <- {p}\")\n    if not subs:\n        for z in all_zips:\n            try:\n                with zipfile.ZipFile(z) as zf:\n                    if any(\"idrid\" in n.lower() for n in zf.namelist()):\n                        dst = \"/kaggle/working/idrid_x\"; os.makedirs(dst, exist_ok=True); zf.extractall(dst)\n                        all_imgs.extend(_files(dst, (\".png\", \".jpg\", \".jpeg\"))); IMG_INDEX.update(build_index(all_imgs))\n                        print(\"  extracted\", z)\n            except Exception as e: print(\"  zip:\", e)\n        return None\n    df = _resolve(pd.concat(subs).drop_duplicates(\"image_id\"), \"IDRiD\")\n    if df is not None and len(df) < 500:\n        print(f\"  [WARN] {len(df)} of ~516 - is the Testing-labels csv attached?\")\n    return df\n\ndef load_messidor():\n    best = None\n    for p in all_csvs:\n        if \"messidor\" not in p.lower(): continue\n        try: d = pd.read_csv(p)\n        except Exception: continue\n        ic = _pick(d.columns, [\"image_id\", \"image name\", \"image\", \"file\", \"id\"])\n        gc = _pick(d.columns, [\"adjudicated_dr_grade\", \"dr_grade\", \"retinopathy grade\", \"grade\", \"diagnosis\"])\n        if not (ic and gc): continue\n        g = [c for c in d.columns if c.strip().lower() == \"adjudicated_gradable\"]\n        if g: d = d[d[g[0]] == 1]\n        s = d[[ic, gc]].dropna(); s.columns = [\"image_id\", \"grade\"]\n        print(f\"  Messidor-2 csv: {len(s):4d} rows <- {p}\")\n        if best is None or len(s) > len(best): best = s\n    return _resolve(best, \"Messidor-2\") if best is not None else None\n\nidrid_df, messidor_df = load_idrid(), load_messidor()","metadata":{},"outputs":[],"execution_count":null},{"id":"1b7684a1-69b1-4e06-816f-f371de1dd4c8","cell_type":"markdown","source":"## 3. Combine, dedupe, split\n\nTwo guards: exact-duplicate removal by MD5 across datasets (APTOS is kept, the external copy is dropped),\nand a path-disjointness assertion on the final splits. The test set is carved out first with `SPLIT_SEED`\nand stratified on the **5-class** grade, so every fold and every run evaluates on identical images.","metadata":{}},{"id":"cc4f0ad5-42b8-4266-ab51-930a5187b396","cell_type":"code","source":"parts = [aptos_df[[\"image_id\", \"grade\", \"path\", \"source\"]]]\nfor d in (idrid_df, messidor_df):\n    if d is not None and len(d): parts.append(d[[\"image_id\", \"grade\", \"path\", \"source\"]])\nall_df = pd.concat(parts, ignore_index=True)\nprint(\"pre-dedupe:\", len(all_df))\nprint(all_df.groupby(\"source\")[\"grade\"].value_counts().unstack(fill_value=0))\n\ndef _md5(p):\n    h = hashlib.md5()\n    try:\n        with open(p, \"rb\") as f:\n            for b in iter(lambda: f.read(1 << 16), b\"\"): h.update(b)\n        return h.hexdigest()\n    except Exception: return None\nall_df[\"_h\"] = [_md5(p) for p in tqdm(all_df.path, leave=False, desc=\"md5\")]\ndup = all_df[\"_h\"].duplicated(keep=\"first\") & all_df[\"_h\"].notna()\nif dup.sum(): print(f\"[LEAKAGE GUARD] dropping {int(dup.sum())} duplicates: {all_df.loc[dup,'source'].value_counts().to_dict()}\")\nall_df = all_df.loc[~dup].drop(columns=\"_h\").reset_index(drop=True)\nall_df[\"y\"] = (all_df[\"grade\"] > 0).astype(int)\nprint(\"post-dedupe:\", len(all_df))\n\naptos = all_df[all_df.source == \"APTOS\"].reset_index(drop=True)\ntrainval, test_df = train_test_split(aptos, test_size=cfg.TEST_FRAC,\n                                     stratify=aptos[\"grade\"], random_state=cfg.SPLIT_SEED)\nskf = StratifiedKFold(n_splits=cfg.N_FOLDS, shuffle=True, random_state=cfg.SPLIT_SEED)\ntr_i, va_i = list(skf.split(trainval, trainval[\"grade\"]))[cfg.FOLD]\ntrain_df, val_df = trainval.iloc[tr_i].copy(), trainval.iloc[va_i].copy()\n\next = all_df[all_df.source != \"APTOS\"]\nif cfg.DATA_POLICY == \"selective\": ext = ext[ext[\"grade\"].isin(cfg.MERGE_GRADES)]\ntrain_df = pd.concat([train_df, ext], ignore_index=True).sample(frac=1, random_state=cfg.SPLIT_SEED).reset_index(drop=True)\nval_df, test_df = val_df.reset_index(drop=True), test_df.reset_index(drop=True)\n\nassert not set(train_df.path) & set(val_df.path), \"LEAKAGE train/val\"\nassert not set(train_df.path) & set(test_df.path), \"LEAKAGE train/test\"\nassert not set(val_df.path) & set(test_df.path), \"LEAKAGE val/test\"\nprint(\"\\nLeakage assertion passed.\")\nfor n, d in [(\"TRAIN\", train_df), (\"VAL\", val_df), (\"TEST\", test_df)]:\n    print(f\"{n:5s} n={len(d):5d}  No-DR {int((d.y==0).sum()):4d}  DR {int((d.y==1).sum()):4d}  \"\n          f\"prevalence {d.y.mean():.3f}  grades {dict(sorted(Counter(d.grade).items()))}\")\nprint(f\"\\nTEST fingerprint (identical in every fold): \"\n      f\"{hashlib.md5(''.join(sorted(os.path.basename(p) for p in test_df.path)).encode()).hexdigest()[:16]}\")\nprint(f\"train/test prevalence gap {abs(train_df.y.mean()-test_df.y.mean()):.3f}\"\n      \"  <- if this is large, try DATA_POLICY='full'\")\nfor n, d in [(\"train\", train_df), (\"val\", val_df), (\"test\", test_df)]: d.to_csv(RUN_DIR / f\"split_{n}.csv\", index=False)","metadata":{},"outputs":[],"execution_count":null},{"id":"ccb210e0-a8de-4c36-9566-dc95bd5e1a32","cell_type":"markdown","source":"## 4. Preprocessing cache (done once, reused by every epoch)","metadata":{}},{"id":"38aeb9df-0528-415d-81e7-90c31cee1abf","cell_type":"code","source":"S = cfg.IMG_SIZE\nCACHE = Path(\"/kaggle/temp\" if os.path.isdir(\"/kaggle/temp\") else cfg.OUT) / f\"cache_{cfg.PREPROCESS}{S}\"\nCACHE.mkdir(parents=True, exist_ok=True)\n\ndef crop_fundus(img, thr):\n    m = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY) > thr\n    if m.sum() < 1000: return img\n    r, c = np.where(m.any(1))[0], np.where(m.any(0))[0]\n    o = img[r[0]:r[-1] + 1, c[0]:c[-1] + 1]\n    return o if o.size else img\n\ndef pad_square(img):\n    h, w = img.shape[:2]; d = max(h, w); t, l = (d - h) // 2, (d - w) // 2\n    return cv2.copyMakeBorder(img, t, d - h - t, l, d - w - l, cv2.BORDER_CONSTANT, value=0)\n\ndef enhance(img):\n    if cfg.PREPROCESS == \"ben\":\n        img = cv2.addWeighted(img, 4, cv2.GaussianBlur(img, (0, 0), S / 30), -4, 128)\n        m = np.zeros(img.shape[:2], np.uint8); cv2.circle(m, (S // 2, S // 2), int(S * 0.47), 1, -1)\n        return img * m[..., None]\n    lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n    lab[:, :, 0] = cv2.createCLAHE(2.0, (8, 8)).apply(lab[:, :, 0])\n    return cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)\n\ndef cache_one(p):\n    dst = CACHE / (hashlib.md5(p.encode()).hexdigest()[:20] + \".jpg\")\n    if not dst.exists():\n        im = cv2.imread(p)\n        if im is None: return None\n        im = cv2.cvtColor(im, cv2.COLOR_BGR2RGB)\n        im = cv2.resize(pad_square(crop_fundus(im, cfg.CROP_THRESHOLD)), (S, S), interpolation=cv2.INTER_AREA)\n        cv2.imwrite(str(dst), cv2.cvtColor(enhance(im), cv2.COLOR_RGB2BGR), [cv2.IMWRITE_JPEG_QUALITY, 95])\n    return str(dst)\n\ncore = sorted(set(train_df.path) | set(val_df.path) | set(test_df.path))\ntodo = core + ([p for p in UNLAB] if (cfg.PSEUDO and UNLAB) else [])\nwith Pool(os.cpu_count()) as pool: cmap = dict(zip(todo, pool.map(cache_one, todo, chunksize=16)))\nbad = [p for p in core if cmap.get(p) is None]\nassert not bad, f\"{len(bad)} unreadable images e.g. {bad[:3]}\"\nfor d in (train_df, val_df, test_df): d[\"cache\"] = d[\"path\"].map(cmap)\nUNLAB_C = [cmap[p] for p in UNLAB if cmap.get(p)] if (cfg.PSEUDO and UNLAB) else []\nprint(f\"cached {len(todo)} images in {(time.time()-T0)/60:.1f} min\")\n\nfig, ax = plt.subplots(1, 4, figsize=(14, 3.6))\nfor a, r in zip(ax, train_df.sample(4, random_state=0).itertuples()):\n    a.imshow(cv2.cvtColor(cv2.imread(r.cache), cv2.COLOR_BGR2RGB)); a.axis(\"off\")\n    a.set_title(f\"{'DR' if r.y else 'No-DR'} (grade {r.grade})\", fontsize=9)\nplt.suptitle(f\"model input — {cfg.PREPROCESS}\"); plt.tight_layout(); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"f649e8ae-b2cb-420e-afb9-99ef8e710a69","cell_type":"markdown","source":"## 5. Augmentation and loaders","metadata":{}},{"id":"faed7ca7-47d3-4b9c-b193-a16beb6926a5","cell_type":"code","source":"_N = dict(mean=(0., 0., 0.), std=(1., 1., 1.), max_pixel_value=255.0)\ntrain_tf = A.Compose([\n    A.Affine(rotate=(-20, 20), scale=(0.85, 1.2), translate_percent=(-0.1, 0.1), shear=(-10, 10), p=0.9),\n    A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5),\n    A.RandomBrightnessContrast(0.15, 0.15, p=0.5),\n    A.OneOf([A.GaussianBlur(blur_limit=(3, 5)), A.Sharpen(), A.Emboss()], p=0.25),\n    A.Normalize(**_N), ToTensorV2()])\neval_tf = A.Compose([A.Normalize(**_N), ToTensorV2()])\n\nclass DS(Dataset):\n    def __init__(s, df, tf, soft=None):\n        s.p = df[\"cache\"].values; s.y = df[\"y\"].values.astype(np.float32)\n        s.g = df[\"grade\"].values.astype(np.int64); s.tf = tf\n        s.soft = soft if soft is not None else np.zeros(len(s.p), np.float32)\n    def __len__(s): return len(s.p)\n    def __getitem__(s, i):\n        im = cv2.cvtColor(cv2.imread(s.p[i]), cv2.COLOR_BGR2RGB)\n        return s.tf(image=im)[\"image\"], torch.tensor(s.y[i]), torch.tensor(s.g[i]), torch.tensor(s.soft[i])\n\nkw = dict(num_workers=cfg.NUM_WORKERS, pin_memory=True, persistent_workers=cfg.NUM_WORKERS > 0)\ndef loaders(tr):\n    return (DataLoader(DS(tr, train_tf), cfg.BATCH_SIZE, shuffle=True, drop_last=True, **kw),\n            DataLoader(DS(val_df, eval_tf), cfg.BATCH_SIZE * 2, shuffle=False, **kw),\n            DataLoader(DS(test_df, eval_tf), cfg.BATCH_SIZE * 2, shuffle=False, **kw))\ntrain_loader, val_loader, test_loader = loaders(train_df)\nprint(f\"batches -> train {len(train_loader)} val {len(val_loader)} test {len(test_loader)}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"1de1a7f0-4cec-4c89-8b8a-bbe8be06d2fd","cell_type":"markdown","source":"## 6. Model\n\n```\nx ─┬─ norm_eff → EfficientNet-B4 → SE → GAP → f1 ─┐\n   └─ norm_res → ResNet50        → SE → GAP → f2 ─┴─ concat → BN → dropout ─┬─ binary logit\n                                                                            └─ 5-class grade logits (aux)\n```","metadata":{}},{"id":"91746130-3d85-4943-a0c0-be655c6d6692","cell_type":"code","source":"class SE(nn.Module):\n    def __init__(s, c, r=8):\n        super().__init__(); s.fc1, s.fc2 = nn.Linear(c, c // r), nn.Linear(c // r, c)\n    def forward(s, x):\n        w = torch.sigmoid(s.fc2(F.relu(s.fc1(x.mean((2, 3))))))\n        return x + x * w[:, :, None, None]\n\ndef bb(name, fb, pre):\n    try: return timm.create_model(name, pretrained=pre, num_classes=0)\n    except Exception as e:\n        print(f\"{name} unavailable ({e}) -> {fb}\"); return timm.create_model(fb, pretrained=pre, num_classes=0)\n\nclass DualSE(nn.Module):\n    def __init__(s):\n        super().__init__()\n        s.eff = bb(cfg.EFF_NAME, cfg.EFF_FB, cfg.PRETRAINED)\n        s.res = bb(cfg.RES_NAME, cfg.RES_FB, cfg.PRETRAINED)\n        ce, cr = s.eff.num_features, s.res.num_features\n        s.se_e, s.se_r = SE(ce, cfg.SE_REDUCTION), SE(cr, cfg.SE_REDUCTION)\n        s.bn, s.drop = nn.BatchNorm1d(ce + cr), nn.Dropout(cfg.DROPOUT)\n        s.head = nn.Linear(ce + cr, 1)\n        s.grade = nn.Linear(ce + cr, 5) if cfg.GRADE_AUX_W > 0 else None\n        for t, m, sd in [(\"e\", EFF_MEAN, EFF_STD), (\"r\", RES_MEAN, RES_STD)]:\n            s.register_buffer(f\"{t}_m\", torch.tensor(m).view(1, 3, 1, 1))\n            s.register_buffer(f\"{t}_s\", torch.tensor(sd).view(1, 3, 1, 1))\n    def forward(s, x):\n        f1 = s.se_e(s.eff.forward_features((x - s.e_m) / s.e_s)).mean((2, 3))\n        f2 = s.se_r(s.res.forward_features((x - s.r_m) / s.r_s)).mean((2, 3))\n        f = s.drop(s.bn(torch.cat([f1, f2], 1)))\n        return s.head(f).squeeze(1), (s.grade(f) if s.grade is not None else None)\n\nmodel = DualSE().to(device).to(memory_format=torch.channels_last)\nwith torch.no_grad():\n    _a, _b = model.eval()(torch.rand(2, 3, S, S, device=device))\nassert _a.shape == (2,) and (_b is None or _b.shape == (2, 5))\nprint(f\"params {sum(p.numel() for p in model.parameters())/1e6:.1f}M | pretrained={cfg.PRETRAINED}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"847fcf35-4612-4e97-9d04-7c65f91f9a2d","cell_type":"markdown","source":"## 7. Metrics and threshold selection","metadata":{}},{"id":"468ec7af-4182-4a18-b9fe-3cc26d5c8a47","cell_type":"code","source":"def binary_metrics(y, p, thr):\n    y = np.asarray(y).astype(int); pred = (np.asarray(p) >= thr).astype(int)\n    tn, fp, fn, tp = confusion_matrix(y, pred, labels=[0, 1]).ravel()\n    sens, spec = tp / max(tp + fn, 1), tn / max(tn + fp, 1)\n    prec = tp / max(tp + fp, 1)\n    return dict(accuracy=100 * (tp + tn) / len(y), sensitivity=100 * sens, specificity=100 * spec,\n                precision=100 * prec, npv=100 * tn / max(tn + fn, 1),\n                f1=100 * 2 * prec * sens / max(prec + sens, 1e-9), balanced=100 * (sens + spec) / 2,\n                auc=100 * roc_auc_score(y, p), ap=100 * average_precision_score(y, p),\n                kappa=100 * cohen_kappa_score(y, pred), mcc=100 * matthews_corrcoef(y, pred),\n                brier=brier_score_loss(y, p), thr=thr, TP=int(tp), TN=int(tn), FP=int(fp), FN=int(fn))\n\ndef best_threshold(y, p, grid=None):\n    \"\"\"Accuracy-maximising threshold, taking the CENTRE of the widest optimal plateau.\n    Youden's J drifted to 0.69 and cost accuracy in earlier runs; plateau-centre is stabler.\"\"\"\n    grid = np.arange(0.05, 0.96, 0.0025) if grid is None else grid\n    acc = np.array([accuracy_score(y, (p >= t).astype(int)) for t in grid])\n    best = acc.max(); idx = np.where(acc >= best - 1e-12)[0]\n    runs, cur = [], [idx[0]]\n    for i in idx[1:]:\n        if i == cur[-1] + 1: cur.append(i)\n        else: runs.append(cur); cur = [i]\n    runs.append(cur); longest = max(runs, key=len)\n    return float(grid[longest[len(longest) // 2]])\n\ndef best_fuse_weight(y, p_bin, p_grade):\n    ws = np.arange(0, 1.01, 0.05)\n    sc = [roc_auc_score(y, (1 - w) * p_bin + w * p_grade) + accuracy_score(\n              y, ((1 - w) * p_bin + w * p_grade >= best_threshold(y, (1 - w) * p_bin + w * p_grade)).astype(int))\n          for w in ws]\n    return float(ws[int(np.argmax(sc))])\n\ndef table(rows, title, note=None):\n    df = pd.DataFrame(rows).T\n    print(\"\\n\" + \"=\" * 104); print(title); print(\"=\" * 104)\n    print(df.round(2).to_string())\n    if note: print(\"  \" + note)\n    return df","metadata":{},"outputs":[],"execution_count":null},{"id":"fa106279-5178-42ef-a3f6-c56aeb7ba197","cell_type":"markdown","source":"## 8. Training\n\nAdamW (backbones 1e-4, SE/heads 1e-3), warmup→cosine, AMP, gradient accumulation, EMA. The checkpoint is\nchosen by `0.5·val AUC + 0.5·val accuracy` — AUC alone picked the wrong variant in an earlier run.\nResumes from `last.pt`; stops when the session budget runs out.","metadata":{}},{"id":"0802b62d-0bed-451f-856f-3a506f6ed482","cell_type":"code","source":"import copy\nclass EMA:\n    def __init__(s, m, d):\n        s.d = d; s.m = copy.deepcopy(m).eval()\n        for p in s.m.parameters(): p.requires_grad_(False)\n    @torch.no_grad()\n    def update(s, m):\n        for e, v in zip(s.m.state_dict().values(), m.state_dict().values()):\n            if v.dtype.is_floating_point:\n                if torch.isfinite(v).all(): e.mul_(s.d).add_(v.detach(), alpha=1 - s.d)\n            else: e.copy_(v)\n\nFLIPS = [None, [3], [2], [2, 3]]\n@torch.no_grad()\ndef predict(net, loader, tta=False):\n    net.eval(); P, G, Y, GR = [], [], [], []\n    for x, y, g, _ in loader:\n        x = x.to(device, non_blocking=True).to(memory_format=torch.channels_last)\n        with amp_autocast(cfg.AMP):\n            ps, gs = [], []\n            for f in (FLIPS if tta else FLIPS[:1]):\n                lo, gl = net(torch.flip(x, f) if f else x)\n                ps.append(torch.sigmoid(lo.float()))\n                if gl is not None: gs.append(1 - gl.float().softmax(1)[:, 0])\n        P.append((sum(ps) / len(ps)).cpu()); Y.append(y); GR.append(g)\n        G.append((sum(gs) / len(gs)).cpu() if gs else torch.zeros(len(y)))\n    return (torch.cat(P).numpy().astype(np.float64), torch.cat(Y).numpy().astype(int),\n            torch.cat(GR).numpy().astype(int), torch.cat(G).numpy().astype(np.float64))\n\ndef train(model, train_loader, tag, epochs):\n    bbp = [p for n, p in model.named_parameters() if n.startswith((\"eff.\", \"res.\"))]\n    hdp = [p for n, p in model.named_parameters() if not n.startswith((\"eff.\", \"res.\"))]\n    opt = torch.optim.AdamW([{\"params\": bbp, \"lr\": cfg.LR_BB}, {\"params\": hdp, \"lr\": cfg.LR_HEAD}],\n                            weight_decay=cfg.WEIGHT_DECAY)\n    steps = epochs * len(train_loader) // cfg.ACCUM_STEPS\n    wu = max(1, cfg.WARMUP_EPOCHS * len(train_loader) // cfg.ACCUM_STEPS)\n    sch = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: (s + 1) / wu if s < wu else\n              0.01 + 0.99 * 0.5 * (1 + math.cos(math.pi * (s - wu) / max(1, steps - wu))))\n    scaler = make_scaler(cfg.AMP); ema = EMA(model, cfg.EMA_DECAY)\n    bce = nn.BCEWithLogitsLoss(); ce = nn.CrossEntropyLoss(label_smoothing=cfg.LABEL_SMOOTH)\n    ls = cfg.LABEL_SMOOTH\n    ck_last, ck_best = RUN_DIR / f\"{tag}_last.pt\", RUN_DIR / f\"{tag}_best.pt\"\n    start, best, bad, hist = 0, -1.0, 0, []\n    if ck_last.exists():\n        d = torch.load(ck_last, map_location=device, weights_only=False)\n        model.load_state_dict(d[\"m\"]); ema.m.load_state_dict(d[\"e\"]); opt.load_state_dict(d[\"o\"])\n        sch.load_state_dict(d[\"s\"]); scaler.load_state_dict(d[\"sc\"])\n        start, best, bad, hist = d[\"ep\"] + 1, d[\"best\"], d[\"bad\"], d[\"hist\"]\n        print(f\"resumed {tag} at epoch {start} (best {best:.4f})\")\n    for ep in range(start, epochs):\n        t = time.time(); model.train(); tl = tn = 0; nan = 0\n        opt.zero_grad(set_to_none=True)\n        for i, (x, y, g, soft) in enumerate(tqdm(train_loader, leave=False, desc=f\"{tag} ep{ep}\")):\n            x = x.to(device, non_blocking=True).to(memory_format=torch.channels_last)\n            y, g, soft = y.to(device), g.to(device), soft.to(device)\n            tgt = torch.where(soft > 0, soft, y).clamp(ls, 1 - ls)\n            with amp_autocast(cfg.AMP):\n                lo, gl = model(x)\n                loss = bce(lo.float(), tgt)\n                if gl is not None:\n                    real = soft <= 0\n                    if real.any(): loss = loss + cfg.GRADE_AUX_W * ce(gl.float()[real], g[real])\n            if not torch.isfinite(loss): nan += 1; opt.zero_grad(set_to_none=True); continue\n            scaler.scale(loss / cfg.ACCUM_STEPS).backward()\n            if (i + 1) % cfg.ACCUM_STEPS == 0:\n                scaler.unscale_(opt); nn.utils.clip_grad_norm_(model.parameters(), cfg.GRAD_CLIP)\n                scaler.step(opt); scaler.update(); opt.zero_grad(set_to_none=True); sch.step(); ema.update(model)\n            tl += float(loss) * len(y); tn += len(y)\n        row = {}\n        for nm, net in [(\"raw\", model), (\"ema\", ema.m)]:\n            p, yv, _, _ = predict(net, val_loader)\n            thr = best_threshold(yv, p); m = binary_metrics(yv, p, thr)\n            row[nm] = 0.5 * m[\"auc\"] + 0.5 * m[\"accuracy\"]; row[nm + \"_m\"] = m\n        which = \"ema\" if row[\"ema\"] >= row[\"raw\"] else \"raw\"\n        sc = row[which]; m = row[which + \"_m\"]\n        hist.append(dict(ep=ep, loss=tl / max(tn, 1), used=which, val_acc=m[\"accuracy\"], val_auc=m[\"auc\"],\n                         val_sens=m[\"sensitivity\"], val_spec=m[\"specificity\"], thr=m[\"thr\"],\n                         sec=time.time() - t, nan=nan))\n        pd.DataFrame(hist).to_csv(RUN_DIR / f\"{tag}_history.csv\", index=False)\n        if sc > best:\n            best, bad = sc, 0\n            torch.save({\"state\": (ema.m if which == \"ema\" else model).state_dict(), \"which\": which, \"ep\": ep}, ck_best)\n        else: bad += 1\n        print(f\"ep {ep:02d} | loss {hist[-1]['loss']:.4f} | val acc {m['accuracy']:.2f} auc {m['auc']:.2f} \"\n              f\"sens {m['sensitivity']:.2f} spec {m['specificity']:.2f} thr {m['thr']:.3f} [{which}] \"\n              f\"| {hist[-1]['sec']:.0f}s {'*' if bad == 0 else ''}\")\n        torch.save(dict(m=model.state_dict(), e=ema.m.state_dict(), o=opt.state_dict(), s=sch.state_dict(),\n                        sc=scaler.state_dict(), ep=ep, best=best, bad=bad, hist=hist), ck_last)\n        if bad >= cfg.PATIENCE: print(\"early stop\"); break\n        if budget_left() < 0.6: print(\"session budget - rerun to resume\"); break\n    return ck_best\n\nCK = train(model, train_loader, \"m\", cfg.EPOCHS)\nmodel.load_state_dict(torch.load(CK, map_location=device)[\"state\"])","metadata":{},"outputs":[],"execution_count":null},{"id":"931009bd-1994-4aee-9c33-1c83965549ed","cell_type":"markdown","source":"## 8.5 Pseudo-labelling (optional, `PSEUDO = True`)\n\nThe trained model labels the 1,928 unlabelled APTOS images, keeps only predictions beyond ±`PSEUDO_CONF`,\nand fine-tunes for a few more epochs on labelled + confident-pseudo data. Those images come from the same\nhospital and cameras as the test set, and are disjoint from it.","metadata":{}},{"id":"e84d4438-891a-42dc-90fb-a8331d532055","cell_type":"code","source":"if cfg.PSEUDO and UNLAB_C:\n    class _P(Dataset):\n        def __init__(s, paths): s.p = paths\n        def __len__(s): return len(s.p)\n        def __getitem__(s, i):\n            im = cv2.cvtColor(cv2.imread(s.p[i]), cv2.COLOR_BGR2RGB)\n            return eval_tf(image=im)[\"image\"], torch.tensor(0.), torch.tensor(0), torch.tensor(0.)\n    pl = DataLoader(_P(UNLAB_C), cfg.BATCH_SIZE * 2, shuffle=False, **kw)\n    pp, _, _, _ = predict(model, pl, tta=True)\n    keep = (pp >= cfg.PSEUDO_CONF) | (pp <= 1 - cfg.PSEUDO_CONF)\n    print(f\"pseudo: keeping {int(keep.sum())}/{len(pp)} at conf {cfg.PSEUDO_CONF} \"\n          f\"({int((pp[keep]>=0.5).sum())} DR / {int((pp[keep]<0.5).sum())} No-DR)\")\n    if keep.sum() > 50:\n        ps = pd.DataFrame(dict(cache=np.array(UNLAB_C)[keep], path=np.array(UNLAB_C)[keep],\n                               y=(pp[keep] >= 0.5).astype(int), grade=-1, source=\"pseudo\",\n                               image_id=\"\", _soft=pp[keep]))\n        aug = pd.concat([train_df.assign(_soft=0.0), ps], ignore_index=True).sample(frac=1, random_state=cfg.SEED)\n        pl2 = DataLoader(DS(aug, train_tf, soft=aug[\"_soft\"].values.astype(np.float32)),\n                         cfg.BATCH_SIZE, shuffle=True, drop_last=True, **kw)\n        print(f\"student trains on {len(aug)} images (was {len(train_df)})\")\n        CK = train(model, pl2, \"student\", max(6, cfg.EPOCHS // 2))\n        model.load_state_dict(torch.load(CK, map_location=device)[\"state\"])\nelse:\n    print(\"pseudo-labelling skipped\")","metadata":{},"outputs":[],"execution_count":null},{"id":"ae3101f1-6645-4fa0-ace2-914200715e96","cell_type":"markdown","source":"## 9. Evaluation\n\nFour inference variants (`raw`/`ema` × TTA are already collapsed into the selected checkpoint, so here it is\nthe binary head, the grade head, and their val-tuned fusion). Each gets its **own threshold tuned on validation**;\nthe variant is picked by `0.5·val AUC + 0.5·val acc`; the choice is applied to test exactly once.","metadata":{}},{"id":"c9a0a0b1-32c9-43c0-891a-f7dd6383a537","cell_type":"code","source":"vp, vy, vg, vpg = predict(model, val_loader, tta=cfg.USE_TTA)\ntp_, ty, tg, tpg = predict(model, test_loader, tta=cfg.USE_TTA)\n\ncands = {\"binary head\": (vp, tp_)}\nif cfg.GRADE_AUX_W > 0 and vpg.any():\n    w = best_fuse_weight(vy, vp, vpg)\n    cands[\"grade head (1-P(g0))\"] = (vpg, tpg)\n    cands[f\"fused w={w:.2f}\"] = ((1 - w) * vp + w * vpg, (1 - w) * tp_ + w * tpg)\n\nvrows, trows, sel = {}, {}, {}\nfor k, (a, b) in cands.items():\n    thr = best_threshold(vy, a); vm = binary_metrics(vy, a, thr)\n    vrows[k], trows[k] = vm, binary_metrics(ty, b, thr)\n    sel[k] = 0.5 * vm[\"auc\"] + 0.5 * vm[\"accuracy\"]\nbest_k = max(sel, key=sel.get)\ntable(vrows, \"TABLE 1 - VALIDATION (each variant at its own val-tuned threshold)\",\n      \"used only to choose the variant and the threshold\")\ntrows[\"PAPER (Shakibania 2024)\"] = {**PAPER, \"thr\": np.nan}\ntable(trows, f\"TABLE 2 - TEST (n={len(ty)}) - selected variant: '{best_k}'\",\n      \"this is the number to report\")\nm = trows[best_k]\nprint(f\"\\nSELECTED '{best_k}' | acc {m['accuracy']:.2f}  sens {m['sensitivity']:.2f}  spec {m['specificity']:.2f}\"\n      f\"  AUC {m['auc']:.2f}  errors {m['FP']+m['FN']}/{len(ty)}\")\nprint(f\"PAPER            | acc {PAPER['accuracy']:.2f}  sens {PAPER['sensitivity']:.2f}  \"\n      f\"spec {PAPER['specificity']:.2f}  AUC {PAPER['auc']:.2f}  errors ~{round(len(ty)*(1-PAPER['accuracy']/100))}/{len(ty)}\")\n\nvbest, tbest = cands[best_k]\nnp.savez(RUN_DIR / f\"preds_{RUN}.npz\", val_prob=vbest, val_y=vy, val_grade=vg,\n         test_prob=tbest, test_y=ty, test_grade=tg, val_thr=vrows[best_k][\"thr\"],\n         variant=best_k, pretrained=cfg.PRETRAINED, fold=cfg.FOLD)\njson.dump({k: {kk: float(vv) for kk, vv in v.items()} for k, v in trows.items() if k != \"PAPER (Shakibania 2024)\"},\n          open(RUN_DIR / \"test_metrics.json\", \"w\"), indent=2)\n\ncm = confusion_matrix(ty, (tbest >= vrows[best_k][\"thr\"]).astype(int), labels=[0, 1])\nfig, ax = plt.subplots(1, 3, figsize=(17, 4.6))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=[\"No-DR\", \"DR\"], yticklabels=[\"No-DR\", \"DR\"], ax=ax[0], cbar=False)\nax[0].set(title=f\"Test confusion — {m['accuracy']:.2f}%\", xlabel=\"Predicted\", ylabel=\"True\")\nfpr, tpr, _ = roc_curve(ty, tbest); ax[1].plot(fpr, tpr, label=f\"AUC {m['auc']:.2f}\")\nax[1].plot([0, 1], [0, 1], \"k--\", lw=.8); ax[1].set(title=\"ROC (test)\", xlabel=\"FPR\", ylabel=\"TPR\"); ax[1].legend()\ngr = pd.DataFrame({\"grade\": tg, \"correct\": (tbest >= vrows[best_k][\"thr\"]).astype(int) == ty}) \\\n       .groupby(\"grade\")[\"correct\"].agg([\"mean\", \"size\"])\nax[2].bar(gr.index.astype(str), gr[\"mean\"], color=\"steelblue\"); ax[2].set_ylim(0, 1.05)\nax[2].set(title=\"Accuracy by true ICDR grade\", xlabel=\"grade\", ylabel=\"accuracy\")\nfor i, (v, n_) in enumerate(zip(gr[\"mean\"], gr[\"size\"])): ax[2].text(i, v + .02, f\"n={n_}\", ha=\"center\", fontsize=8)\nplt.tight_layout(); plt.savefig(RUN_DIR / \"test_report.png\", dpi=120); plt.show()\nprint(\"\\nper-grade accuracy (which DR stages are missed):\"); print(gr.round(4).to_string())","metadata":{},"outputs":[],"execution_count":null},{"id":"128b5ab3-9e49-4336-bc24-4e2ae8a11b02","cell_type":"markdown","source":"## 10. Out-of-fold ensemble\n\nRun after folds 0–4 are done: attach each fold's committed output as an input, then run **section 0, section 7\nand this cell only**. The threshold is tuned on the pooled out-of-fold predictions (~2,900 images instead of 580),\nwhich is the largest honest validation signal available, and no test image is used to choose it.","metadata":{}},{"id":"3c5bd59e-39ac-4128-aa95-719f45d9f463","cell_type":"code","source":"files = sorted(set(glob.glob(\"/kaggle/working/run_*/preds_*.npz\") +\n                   glob.glob(\"/kaggle/input/**/run_*/preds_*.npz\", recursive=True)))\nprint(f\"{len(files)} prediction files:\", [os.path.basename(f) for f in files])\nif len(files) < 2:\n    print(\"need at least 2 folds - run FOLD = 0..4 first\")\nelse:\n    vps, vys, tps, meta = [], [], [], []\n    ty = tg = None\n    for f in files:\n        d = np.load(f, allow_pickle=True)\n        vps.append(d[\"val_prob\"]); vys.append(d[\"val_y\"]); tps.append(d[\"test_prob\"])\n        if ty is None: ty, tg = d[\"test_y\"], d[\"test_grade\"]\n        else: assert len(d[\"test_y\"]) == len(ty) and (d[\"test_y\"] == ty).all(), f\"{f}: different test set\"\n        meta.append(dict(file=os.path.basename(f), fold=int(d[\"fold\"]), variant=str(d[\"variant\"]),\n                         pretrained=bool(d[\"pretrained\"]), val_thr=float(d[\"val_thr\"])))\n    print(pd.DataFrame(meta).to_string(index=False))\n    oof_p, oof_y = np.concatenate(vps), np.concatenate(vys)\n    ov = np.concatenate([[i] * len(v) for i, v in enumerate(vys)])\n    if len(set(map(len, vys))) and len(oof_y) > 1.5 * len(vys[0]):\n        print(f\"\\nOOF pool: {len(oof_y)} images from {len(files)} folds (vs {len(vys[0])} for one fold)\")\n    thr_oof = best_threshold(oof_y, oof_p)\n    test_ens = np.mean(tps, 0)\n\n    rows = {}\n    for mt, t_i in zip(meta, tps): rows[f\"fold {mt['fold']}\"] = binary_metrics(ty, t_i, mt[\"val_thr\"])\n    rows[\"ENSEMBLE @ mean fold thr\"] = binary_metrics(ty, test_ens, float(np.mean([m[\"val_thr\"] for m in meta])))\n    rows[\"ENSEMBLE @ 0.50\"] = binary_metrics(ty, test_ens, 0.5)\n    rows[f\"ENSEMBLE @ OOF thr ({thr_oof:.3f})\"] = binary_metrics(ty, test_ens, thr_oof)\n    rows[\"PAPER (Shakibania 2024)\"] = {**PAPER, \"thr\": np.nan}\n    df = table(rows, f\"OUT-OF-FOLD ENSEMBLE - test n={len(ty)}\",\n               \"report the OOF-thr row: the largest honest validation signal, and no test image chose it\")\n    main = rows[f\"ENSEMBLE @ OOF thr ({thr_oof:.3f})\"]\n    accs = [rows[f\"fold {m['fold']}\"][\"accuracy\"] for m in meta]\n    print(f\"\\nmembers: {np.mean(accs):.2f} +/- {np.std(accs):.2f} %   ensemble: {main['accuracy']:.2f} %\")\n    print(f\"errors {main['FP']+main['FN']}/{len(ty)} vs the paper's ~{round(len(ty)*(1-PAPER['accuracy']/100))}/{len(ty)}\"\n          f\"   ({'AHEAD' if main['accuracy'] > PAPER['accuracy'] else 'behind'} on accuracy, \"\n          f\"{'AHEAD' if main['auc'] > PAPER['auc'] else 'behind'} on AUC)\")\n    se = (main[\"accuracy\"] / 100 * (1 - main[\"accuracy\"] / 100) / len(ty)) ** .5 * 100\n    print(f\"95% CI for this accuracy: {main['accuracy']-1.96*se:.2f} - {main['accuracy']+1.96*se:.2f} \"\n          f\"(+/-{1.96*se:.2f} pp) - a sub-1-point gap on {len(ty)} images is not a significant win\")\n    ENS = Path(\"/kaggle/working/ensemble\"); ENS.mkdir(exist_ok=True)\n    df.to_csv(ENS / \"oof_ensemble.csv\")\n    np.savez(ENS / \"oof_ensemble.npz\", test_prob=test_ens, test_y=ty, test_grade=tg, oof_thr=thr_oof)","metadata":{},"outputs":[],"execution_count":null}]}