{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":14774,"sourceType":"competition"},{"sourceId":2819730,"sourceType":"datasetVersion"},{"sourceId":18382207,"sourceType":"datasetVersion"},{"sourceId":339937939,"sourceType":"kernelVersion"}],"isGpuEnabled":true,"isInternetEnabled":true,"language":"python","sourceType":"notebook"},"papermill":{"default_parameters":{},"duration":26484.621604,"end_time":"2026-08-13T03:08:36.874333+00:00","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-08-12T19:47:12.252729+00:00","version":"2.7.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"41518443","cell_type":"markdown","source":"## 0. Setup","metadata":{"papermill":{"duration":0.006099,"end_time":"2026-08-12T19:47:15.034614+00:00","exception":false,"start_time":"2026-08-12T19:47:15.028515+00:00","status":"completed"},"tags":[]}},{"id":"b04e01fc","cell_type":"code","source":"import importlib, subprocess, sys\ndef _ensure(pkg, pip_name=None):\n    try: importlib.import_module(pkg)\n    except ImportError:\n        subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pip_name or pkg], check=False)\n_ensure(\"timm\"); _ensure(\"albumentations\"); _ensure(\"seaborn\")","metadata":{"papermill":{"duration":19.828943,"end_time":"2026-08-12T19:47:34.869672+00:00","exception":false,"start_time":"2026-08-12T19:47:15.040729+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"583c25ff","cell_type":"code","source":"import os, sys, math, time, random, warnings, zipfile, hashlib\nfrom collections import Counter, defaultdict\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, WeightedRandomSampler\nimport timm, timm.data\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (accuracy_score, precision_recall_fscore_support, f1_score,\n                             cohen_kappa_score, confusion_matrix, classification_report,\n                             roc_auc_score, roc_curve)\nfrom sklearn.preprocessing import label_binarize\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings(\"ignore\")\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\n_CU = torch.cuda.is_available()\ndef amp_autocast(enabled=True):\n    try: return torch.amp.autocast(\"cuda\", enabled=enabled and _CU)\n    except Exception: return torch.cuda.amp.autocast(enabled=enabled and _CU)\ndef make_scaler(enabled=True):\n    try: return torch.amp.GradScaler(\"cuda\", enabled=enabled and _CU)\n    except Exception: return torch.cuda.amp.GradScaler(enabled=enabled and _CU)","metadata":{"papermill":{"duration":0.608581,"end_time":"2026-08-12T19:47:35.48482+00:00","exception":false,"start_time":"2026-08-12T19:47:34.876239+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"6040bd7d","cell_type":"markdown","source":"## 1. Configuration\n\n`RUN` is the only thing to change between runs. Each preset differs by one variable, so the\nresults form an ablation table — and every completed run becomes an ensemble member in section 15.","metadata":{"papermill":{"duration":0.006729,"end_time":"2026-08-12T19:47:35.498168+00:00","exception":false,"start_time":"2026-08-12T19:47:35.491439+00:00","status":"completed"},"tags":[]}},{"id":"1f33732b","cell_type":"code","source":"# ==================== PICK YOUR RUN ====================\nRUN = \"mlf1\"      # \"mlf0\" | \"mlf1\" | \"mlf2\" | \"mlf3\"\n# =======================================================\n\nPRESETS = {\n    #         kappa  preproc      freeze  seed\n    \"mlf0\":  (0.0,   \"rgb_clahe\", 0.5,    42),   # Run B architecture as-is (baseline)\n    \"mlf1\":  (0.5,   \"rgb_clahe\", 0.5,    42),   # + ordinal kappa loss   <-- start here\n    \"mlf2\":  (0.5,   \"rgb_clahe\", 0.3,    42),   # + less freezing (more capacity)\n    \"mlf3\":  (0.5,   \"rgb_clahe\", 0.5,    1337), # seed variant -> ensemble diversity\n}\nassert RUN in PRESETS\n_KW, _PRE, _FRZ, _SEED = PRESETS[RUN]\n\n\nclass CFG:\n    SEED = _SEED\n    IMG_SIZE = 384                       # ViT-B/16@384 position embeddings are baked in at 384\n\n    # ---- branches ----\n    EFF_NAME = \"tf_efficientnet_b5.ns_jft_in1k\"\n    EFF_DROP_PATH = 0.2\n    EFF_OUT_INDICES = (2, 3, 4)          # 3 CNN stages for multi-level fusion\n    VIT_NAME = \"vit_base_patch16_384.augreg_in21k_ft_in1k\"\n    VIT_DROP = 0.1\n    VIT_DROP_PATH = 0.1\n    VIT_HOOK_BLOCKS = (3, 7, 11)         # 3 ViT depths, tapped via forward hooks\n\n    # ---- fusion ----\n    PROJ_DIM = 256\n    XATTN_HEADS = 8\n    XATTN_FP32 = True                    # cross-attn outside autocast: NaN guard\n\n    # ---- task: COARSE head (Severe + PDR merged into one bucket) ----\n    NUM_CLASSES = 4\n    CLASS_NAMES = [\"No-DR\", \"Mild\", \"Moderate\", \"Severe+PDR\"]\n    GRADE_TO_LABEL = {0: 0, 1: 1, 2: 2, 3: 3, 4: 3}      # <-- 4 folds into 3\n\n    # ---- task: FINE reporting grade (what we are actually graded on) ----\n    FINE_NUM_CLASSES = 5\n    FINE_CLASS_NAMES = [\"No-DR\", \"Mild\", \"Moderate\", \"Severe\", \"PDR\"]\n    HI_GRADES = (3, 4)                                   # the two grades the expert splits\n\n    # ---- data ----\n    TEST_FRAC = 0.20\n    VAL_FRAC  = 0.10\n    MERGE_GRADES = (1, 2, 3, 4)\n    POOL_ALL_APTOS_CSVS = True\n    MD5_DEDUPE = True\n    MESSIDOR_CSV = None                  # manual override, see section 2.5 if auto-detect fails\n    MESSIDOR_IMG_COL = None\n    MESSIDOR_GRADE_COL = None\n\n    # ---- preprocessing ----\n    PREPROCESS = _PRE                    # \"rgb_clahe\" | \"green_clahe\"\n    CLAHE_CLIP = 2.0\n    CLAHE_GRID = 8\n    CROP_THRESHOLD = 7\n\n    # ---- loaders ----\n    BATCH_SIZE = 8\n    ACCUM_STEPS = 4                      # effective batch 32 (B5@384 OOM-safe)\n    NUM_WORKERS = 4\n    USE_SAMPLER = True\n    SAMPLER_POWER = 0.75\n\n    # ---- losses ----\n    FOCAL_GAMMA = 2.0\n    FOCAL_ALPHA = \"equal\"                # sampler is the sole imbalance corrector\n    AUX_WEIGHT = 0.3\n    KAPPA_WEIGHT = _KW\n\n    # ---- training ----\n    EPOCHS_HEAD = 4\n    LR_HEAD = 1e-3\n    EPOCHS_FT = 22\n    WARMUP_EPOCHS = 3\n    LR_FT_HEAD = 5e-4\n    LR_FT_EFF = 6e-5\n    LR_FT_VIT = 2e-5\n    LAYER_DECAY = 0.75                   # ViT layer-wise LR decay\n    WEIGHT_DECAY = 1e-4\n    GRAD_CLIP = 0.5                      # tightened from 1.0 (NaN guard)\n    PATIENCE = 6\n    EMA_DECAY = 0.9995\n    FREEZE_FRAC = _FRZ                   # Phase-2 partial backbone freeze\n    AMP = True\n    USE_TTA = True\n\n    # ---- Stage E: Severe-vs-PDR binary expert ----\n    EXPERT_NAME    = \"tf_efficientnet_b4.ns_jft_in1k\"  # different family depth from the B5 coarse net\n    EXPERT_EPOCHS  = 16\n    EXPERT_LR      = 2e-4\n    EXPERT_WD      = 1e-4\n    EXPERT_DROP    = 0.3                               # few hundred images -> regularise hard\n    EXPERT_DROPPATH= 0.2\n    EXPERT_HOLDOUT = 0.15                              # stratified slice of train 3/4 for ckpt selection\n    EXPERT_TTA     = True\n    EXPERT_GAMMA   = 1.5                               # focal gamma for the binary expert\n    USE_CHEAP_EXPERT = True                            # blend in a free MLP on the cached fusion taps\n\n    OUT = \"/kaggle/working\"\n    CKPT = f\"/kaggle/working/best_{RUN}_c4.pt\"\n\ncfg = CFG()\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)\nos.makedirs(cfg.OUT, exist_ok=True)\nassert len(cfg.CLASS_NAMES) == cfg.NUM_CLASSES\nassert max(cfg.GRADE_TO_LABEL.values()) == cfg.NUM_CLASSES - 1\nassert len(cfg.FINE_CLASS_NAMES) == cfg.FINE_NUM_CLASSES\n# the merged bucket is the LAST coarse class; Stage E splits exactly this one\nMERGED_LABEL = cfg.NUM_CLASSES - 1\nassert all(cfg.GRADE_TO_LABEL[g] == MERGED_LABEL for g in cfg.HI_GRADES), \\\n    \"HI_GRADES must all map to the merged coarse label\"\n\nprint(f\"===== {RUN.upper()} =====\")\nprint(f\"  kappa loss {cfg.KAPPA_WEIGHT} | preproc {cfg.PREPROCESS} | \"\n      f\"freeze {cfg.FREEZE_FRAC} | seed {cfg.SEED}\")\nprint(f\"  batch {cfg.BATCH_SIZE} x accum {cfg.ACCUM_STEPS} = eff {cfg.BATCH_SIZE*cfg.ACCUM_STEPS}\")\nprint(f\"  epochs {cfg.EPOCHS_HEAD} head + {cfg.EPOCHS_FT} ft (patience {cfg.PATIENCE})\")\nprint(f\"  COARSE task : {cfg.NUM_CLASSES}-class {cfg.CLASS_NAMES}\")\nprint(f\"  FINE report : {cfg.FINE_NUM_CLASSES}-class {cfg.FINE_CLASS_NAMES}\")\nprint(f\"  expert      : {cfg.EXPERT_NAME} | {cfg.EXPERT_EPOCHS} ep | cheap-blend {cfg.USE_CHEAP_EXPERT}\")","metadata":{"papermill":{"duration":0.029086,"end_time":"2026-08-12T19:47:35.533646+00:00","exception":false,"start_time":"2026-08-12T19:47:35.50456+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"4775da15","cell_type":"markdown","source":"**1.2 — Per-branch normalization stats, read from timm**","metadata":{"papermill":{"duration":0.006529,"end_time":"2026-08-12T19:47:35.546858+00:00","exception":false,"start_time":"2026-08-12T19:47:35.540329+00:00","status":"completed"},"tags":[]}},{"id":"ffab52e5","cell_type":"code","source":"def _stats(name, **kw):\n    m = timm.create_model(name, pretrained=False, **kw)\n    dc = timm.data.resolve_model_data_config(m); del m\n    return tuple(dc[\"mean\"]), tuple(dc[\"std\"]), tuple(dc[\"input_size\"])\n\nEFF_MEAN, EFF_STD, EFF_IN = _stats(cfg.EFF_NAME, num_classes=0)\nVIT_MEAN, VIT_STD, VIT_IN = _stats(cfg.VIT_NAME, num_classes=0)\nprint(f\"EffNet mean={EFF_MEAN} std={EFF_STD} native={EFF_IN}\")\nprint(f\"ViT    mean={VIT_MEAN} std={VIT_STD} native={VIT_IN}\")\nassert VIT_IN[1] == cfg.IMG_SIZE, f\"ViT requires {VIT_IN[1]}\"\nprint(\"\\nDifferent stats per branch -> normalization lives inside forward(), not in the transform.\")","metadata":{"papermill":{"duration":1.918387,"end_time":"2026-08-12T19:47:37.472169+00:00","exception":false,"start_time":"2026-08-12T19:47:35.553782+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"2c5a5b23","cell_type":"markdown","source":"## 2. Locate the datasets","metadata":{"papermill":{"duration":0.006437,"end_time":"2026-08-12T19:47:37.485221+00:00","exception":false,"start_time":"2026-08-12T19:47:37.478784+00:00","status":"completed"},"tags":[]}},{"id":"a82ce476","cell_type":"code","source":"INPUT_ROOT = \"/kaggle/input\" if os.path.isdir(\"/kaggle/input\") else \".\"\ndef _list_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\n\nall_csvs   = _list_files(INPUT_ROOT, (\".csv\",))\nall_images = _list_files(INPUT_ROOT, (\".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\"))\nall_zips   = _list_files(INPUT_ROOT, (\".zip\",))\nprint(f\"Scanned {INPUT_ROOT}: {len(all_csvs)} csv, {len(all_images)} images, {len(all_zips)} zips\\n\")\n_bf = defaultdict(int)\nfor p in all_images: _bf[os.path.relpath(p, INPUT_ROOT).split(os.sep)[0]] += 1\nprint(\"Images per input folder:\")\nfor k, v in sorted(_bf.items()): print(f\"  {k}: {v}\")\nif not _bf: print(\"WARNING: no images -> attach datasets via '+ Add Input'.\")","metadata":{"papermill":{"duration":13.427288,"end_time":"2026-08-12T19:47:50.919393+00:00","exception":false,"start_time":"2026-08-12T19:47:37.492105+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"91bf9450","cell_type":"markdown","source":"**2.2 — CSV inventory**\n\nThis is the cell that tells you whether Messidor-2 is present. If auto-detection later fails, read\nthe right column names off this output and paste them into `CFG.MESSIDOR_*`.","metadata":{"papermill":{"duration":0.006396,"end_time":"2026-08-12T19:47:50.932567+00:00","exception":false,"start_time":"2026-08-12T19:47:50.926171+00:00","status":"completed"},"tags":[]}},{"id":"54bbef6a","cell_type":"code","source":"print(\"=\" * 78); print(\"EVERY CSV FOUND\"); print(\"=\" * 78)\nfor p in all_csvs:\n    try:\n        d = pd.read_csv(p, nrows=2)\n        n = sum(1 for _ in open(p, encoding=\"utf-8\", errors=\"ignore\")) - 1\n        print(f\"\\n{p}\\n   rows~{n} | cols: {list(d.columns)}\")\n        if len(d): print(f\"   first: {d.iloc[0].to_dict()}\")\n    except Exception as e:\n        print(f\"\\n{p}\\n   [unreadable] {e}\")\nprint(\"\\n\" + \"=\" * 78)\n\ndef build_image_index(paths):\n    idx = {}\n    for p in paths: idx.setdefault(os.path.splitext(os.path.basename(p))[0].lower(), p)\n    return idx\nIMG_INDEX = build_image_index(all_images)\nIMG_FULL  = {os.path.basename(p).lower(): p for p in all_images}","metadata":{"papermill":{"duration":0.133773,"end_time":"2026-08-12T19:47:51.072721+00:00","exception":false,"start_time":"2026-08-12T19:47:50.938948+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"23c255dc","cell_type":"markdown","source":"### 2.3 — APTOS 2019 (pool every label CSV)","metadata":{"papermill":{"duration":0.00633,"end_time":"2026-08-12T19:47:51.085874+00:00","exception":false,"start_time":"2026-08-12T19:47:51.079544+00:00","status":"completed"},"tags":[]}},{"id":"7483c74b","cell_type":"code","source":"frames, srcs = [], []\nfor p in all_csvs:\n    if any(k in p.lower() for k in [\"messidor\", \"idrid\"]): continue\n    try: df = pd.read_csv(p)\n    except Exception: continue\n    df.columns = [c.strip().lower() for c in df.columns]\n    if not (any(\"id_code\" in c for c in df.columns) and any(\"diagnosis\" in c for c in df.columns)):\n        continue\n    ic = [c for c in df.columns if \"id_code\" in c][0]\n    dc = [c for c in df.columns if \"diagnosis\" in c][0]\n    sub = df[[ic, dc]].dropna(); sub.columns = [\"image_id\", \"grade\"]\n    if sub[\"grade\"].nunique() < 2: continue          # unlabelled submission file\n    frames.append(sub); srcs.append((p, len(sub)))\n    if not cfg.POOL_ALL_APTOS_CSVS: break\nassert frames, \"APTOS labels not found (need id_code + diagnosis).\"\nprint(\"APTOS label files used:\")\nfor p, n in srcs: print(f\"  {n:5d} rows <- {p}\")\n\naptos_df = pd.concat(frames, ignore_index=True)\nn0 = len(aptos_df)\naptos_df = aptos_df.drop_duplicates(subset=[\"image_id\"], keep=\"first\").reset_index(drop=True)\nprint(f\"\\n{n0} rows -> {len(aptos_df)} unique ids (dropped {n0-len(aptos_df)} dup ids)\")\naptos_df[\"grade\"] = aptos_df[\"grade\"].astype(int)\naptos_csv_set = {p for p, _ in srcs}\n\ndef _res_aptos(i):\n    s = str(i).strip().lower()\n    return (IMG_INDEX.get(s) or IMG_FULL.get(s) or IMG_FULL.get(s+\".png\")\n            or IMG_FULL.get(s+\".jpg\") or IMG_FULL.get(s+\".jpeg\"))\naptos_df[\"path\"] = aptos_df[\"image_id\"].map(_res_aptos)\nmiss = int(aptos_df[\"path\"].isna().sum())\naptos_df = aptos_df.dropna(subset=[\"path\"]).reset_index(drop=True)\naptos_df[\"source\"] = \"APTOS\"\nprint(f\"APTOS usable images: {len(aptos_df)} (dropped {miss} unmatched)\")\nif len(aptos_df) == 0: raise RuntimeError(\"APTOS images not found.\")\nif len(aptos_df) < 3500: print(f\"  [WARN] expected ~3662, matched {len(aptos_df)}\")\nprint(\"APTOS grades:\", dict(sorted(Counter(aptos_df[\"grade\"]).items())))","metadata":{"papermill":{"duration":0.060166,"end_time":"2026-08-12T19:47:51.152833+00:00","exception":false,"start_time":"2026-08-12T19:47:51.092667+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"20808a4d","cell_type":"markdown","source":"### 2.4 — IDRiD","metadata":{"papermill":{"duration":0.006821,"end_time":"2026-08-12T19:47:51.166473+00:00","exception":false,"start_time":"2026-08-12T19:47:51.159652+00:00","status":"completed"},"tags":[]}},{"id":"1abd5312","cell_type":"code","source":"def _pick(cols, keys, bad=()):\n    low = {c.lower().strip(): c for c in cols}\n    for k in keys:\n        for lc, orig in low.items():\n            if (k == lc or k in lc) and not any(b in lc for b in bad): return orig\n    return None\n\ndef load_idrid():\n    NAME = [\"image name\", \"image_name\", \"imagename\", \"image\", \"file\", \"filename\", \"name\"]\n    GRAD = [\"retinopathy grade\", \"retinopathy_grade\", \"grade\", \"diagnosis\", \"dr_grade\", \"class\", \"label\"]\n    hit, checked = None, []\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: df = pd.read_csv(p)\n        except Exception: continue\n        checked.append((p, list(df.columns)))\n        nc, gc = _pick(df.columns, NAME), _pick(df.columns, GRAD)\n        if nc and gc: hit = (p, df, nc, gc); break\n    if hit is None:\n        print(\"IDRiD: no matching csv.\")\n        for p, c in checked: print(f\"    {p} -> {c}\")\n        return None\n    p, df, nc, gc = hit\n    df = df[[nc, gc]].dropna(); df.columns = [\"image_id\", \"grade\"]\n    df[\"grade\"] = df[\"grade\"].astype(float).round().astype(int)\n    df = df[df[\"grade\"].between(0, 4)]\n    idx = build_image_index(all_images)\n    if not any(str(n).strip().lower() in idx for n in df[\"image_id\"].head(20)):\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                        dest = \"/kaggle/working/idrid_imgs\"; os.makedirs(dest, exist_ok=True)\n                        zf.extractall(dest); print(\"IDRiD: extracted\", z)\n                        all_images.extend(_list_files(dest, (\".png\", \".jpg\", \".jpeg\")))\n            except Exception as e: print(\"  zip error:\", e)\n        idx = build_image_index(all_images)\n    df[\"path\"] = df[\"image_id\"].map(lambda n: idx.get(str(n).strip().lower()))\n    m = int(df[\"path\"].isna().sum())\n    df = df.dropna(subset=[\"path\"]).reset_index(drop=True); df[\"source\"] = \"IDRiD\"\n    print(f\"IDRiD csv: {p}\\nIDRiD usable: {len(df)} (dropped {m})\")\n    print(\"IDRiD grades:\", dict(sorted(Counter(df['grade']).items())))\n    return df if len(df) else None\n\nidrid_df = load_idrid()","metadata":{"papermill":{"duration":0.044369,"end_time":"2026-08-12T19:47:51.217811+00:00","exception":false,"start_time":"2026-08-12T19:47:51.173442+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"63f0937e","cell_type":"markdown","source":"### 2.5 — Messidor-2\n\nRun A merged **zero** Messidor images; Run B got 514. This version tries several name-matching\nstrategies, reports which stage failed, and accepts a manual override from `CFG`.","metadata":{"papermill":{"duration":0.006746,"end_time":"2026-08-12T19:47:51.231646+00:00","exception":false,"start_time":"2026-08-12T19:47:51.2249+00:00","status":"completed"},"tags":[]}},{"id":"b9919528","cell_type":"code","source":"_MG = [\"adjudicated_dr_grade\", \"dr_grade\", \"retinopathy grade\", \"retinopathy_grade\",\n       \"dr_level\", \"diagnosis\", \"grade\", \"level\", \"label\", \"class\", \"severity\"]\n_MI = [\"image_id\", \"image id\", \"image name\", \"image_name\", \"imagename\", \"image\",\n       \"id_code\", \"filename\", \"file_name\", \"file\", \"img\", \"name\", \"id\"]\n\ndef _resolve_any(n):\n    s = str(n).strip()\n    if not s or s.lower() == \"nan\": return None\n    base = os.path.basename(s.replace(\"\\\\\", \"/\"))\n    stem = os.path.splitext(base)[0].lower()\n    cands = [base.lower(), stem, s.lower()] + [stem + e for e in\n             (\".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\")]\n    for c in cands:\n        if c in IMG_FULL:  return IMG_FULL[c]\n        if c in IMG_INDEX: return IMG_INDEX[c]\n    return None\n\ndef load_messidor():\n    if cfg.MESSIDOR_CSV:\n        print(f\"Messidor-2: manual override {cfg.MESSIDOR_CSV}\")\n        if not (cfg.MESSIDOR_IMG_COL and cfg.MESSIDOR_GRADE_COL):\n            print(\"  [ERR] also set MESSIDOR_IMG_COL / MESSIDOR_GRADE_COL\"); return None\n        cands = [(cfg.MESSIDOR_CSV, pd.read_csv(cfg.MESSIDOR_CSV),\n                  cfg.MESSIDOR_IMG_COL, cfg.MESSIDOR_GRADE_COL)]\n    else:\n        cands = []\n        order = ([p for p in all_csvs if \"messidor\" in p.lower()] +\n                 [p for p in all_csvs if \"messidor\" not in p.lower()])\n        for p in order:\n            if p in aptos_csv_set or \"idrid\" in p.lower(): continue\n            try: df = pd.read_csv(p)\n            except Exception: continue\n            if any(\"retinopathy grade\" in c.strip().lower() for c in df.columns): continue\n            ic = _pick(df.columns, _MI)\n            gc = _pick(df.columns, _MG, bad=(\"dme\", \"edema\", \"macular\", \"gradab\"))\n            if ic and gc and ic != gc: cands.append((p, df, ic, gc))\n        if not cands:\n            print(\"Messidor-2: NO csv matched.\")\n            print(\"  -> If a Messidor csv IS in the 2.2 inventory, set CFG.MESSIDOR_* and rerun.\")\n            print(\"  -> If not, the dataset is not attached. Add it under '+ Add Input'.\")\n            return None\n\n    for p, df, ic, gc in cands:\n        print(f\"\\nMessidor-2 candidate: {p}\\n  image='{ic}' grade='{gc}'\")\n        sub = df[[ic, gc]].dropna().copy(); sub.columns = [\"image_id\", \"grade\"]\n        try: sub[\"grade\"] = sub[\"grade\"].astype(float).round().astype(int)\n        except Exception: print(\"  grade not numeric -> next\"); continue\n        sub = sub[sub[\"grade\"].between(0, 4)]\n        if not len(sub): print(\"  no grades in 0-4 -> next\"); continue\n        idx = build_image_index(all_images)\n        if not any(str(n).strip().lower() in idx for n in sub[\"image_id\"].head(20)):\n            for z in all_zips:\n                try:\n                    with zipfile.ZipFile(z) as zf:\n                        if any(\"messidor\" in n.lower() for n in zf.namelist()):\n                            dest = \"/kaggle/working/messidor_imgs\"; os.makedirs(dest, exist_ok=True)\n                            zf.extractall(dest); print(\"  extracted\", z)\n                            all_images.extend(_list_files(dest, (\".png\", \".jpg\", \".jpeg\")))\n                            IMG_FULL.update({os.path.basename(q).lower(): q\n                                             for q in _list_files(dest, (\".png\", \".jpg\", \".jpeg\"))})\n                except Exception as e: print(\"  zip error:\", e)\n        sub[\"path\"] = sub[\"image_id\"].map(_resolve_any)\n        m = int(sub[\"path\"].isna().sum())\n        sub = sub.dropna(subset=[\"path\"]).reset_index(drop=True)\n        print(f\"  matched {len(sub)}, unmatched {m}\")\n        if not len(sub):\n            print(f\"  csv names look like {df[ic].head(3).tolist()}\")\n            print(f\"  files look like {[os.path.basename(q) for q in all_images[:3]]}\")\n            continue\n        sub[\"source\"] = \"Messidor-2\"\n        print(f\"  OK. grades: {dict(sorted(Counter(sub['grade']).items()))}\")\n        return sub\n    print(\"\\nMessidor-2: all candidates failed -> APTOS + IDRiD only.\")\n    return None\n\nmessidor_df = load_messidor()","metadata":{"papermill":{"duration":0.062877,"end_time":"2026-08-12T19:47:51.301318+00:00","exception":false,"start_time":"2026-08-12T19:47:51.238441+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"e24349f1","cell_type":"markdown","source":"## 3. Combine, dedupe, split — leakage-safe\n\nTwo guards, both carried from Run B:\n\n1. **MD5 content-hash dedupe across all three datasets before splitting.** APTOS/IDRiD/Messidor-2\n   are independent releases that occasionally re-host the same image; an exact duplicate landing in\n   both train and test is leakage.\n2. **All resampling happens inside the train DataLoader only**, strictly after the split, so no\n   oversampled row can cross a split boundary.","metadata":{"papermill":{"duration":0.006975,"end_time":"2026-08-12T19:47:51.315638+00:00","exception":false,"start_time":"2026-08-12T19:47:51.308663+00:00","status":"completed"},"tags":[]}},{"id":"3df72c6d","cell_type":"code","source":"frames = [aptos_df[[\"image_id\", \"grade\", \"path\", \"source\"]]]\nfor d in (idrid_df, messidor_df):\n    if d is not None and len(d): frames.append(d[[\"image_id\", \"grade\", \"path\", \"source\"]])\nall_df = pd.concat(frames, ignore_index=True)\nprint(\"Combined (pre-dedupe):\", len(all_df))\nprint(all_df.groupby(\"source\")[\"grade\"].value_counts().unstack(fill_value=0))\n\nif cfg.MD5_DEDUPE:\n    def _md5(p, chunk=1 << 16):\n        h = hashlib.md5()\n        try:\n            with open(p, \"rb\") as f:\n                while True:\n                    b = f.read(chunk)\n                    if not b: break\n                    h.update(b)\n            return h.hexdigest()\n        except Exception: return None\n    print(\"\\nHashing images for cross-dataset duplicate check...\")\n    all_df[\"_md5\"] = [_md5(p) for p in tqdm(all_df[\"path\"], leave=False)]\n    dup = all_df[\"_md5\"].duplicated(keep=\"first\") & all_df[\"_md5\"].notna()\n    if int(dup.sum()):\n        print(f\"[LEAKAGE GUARD] dropping {int(dup.sum())} exact duplicates: \"\n              f\"{all_df.loc[dup,'source'].value_counts().to_dict()}\")\n        all_df = all_df.loc[~dup].reset_index(drop=True)\n    else:\n        print(\"No exact duplicates across datasets.\")\n    all_df = all_df.drop(columns=[\"_md5\"])\nprint(\"Combined (post-dedupe):\", len(all_df))","metadata":{"papermill":{"duration":108.428857,"end_time":"2026-08-12T19:49:39.75128+00:00","exception":false,"start_time":"2026-08-12T19:47:51.322423+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"95d5562c","cell_type":"code","source":"aptos_only = all_df[all_df[\"source\"] == \"APTOS\"].reset_index(drop=True)\ntrain_ap, temp_ap = train_test_split(\n    aptos_only, test_size=cfg.VAL_FRAC + cfg.TEST_FRAC,\n    stratify=aptos_only[\"grade\"], random_state=cfg.SEED)\nval_ap, test_ap = train_test_split(\n    temp_ap, test_size=cfg.TEST_FRAC / (cfg.VAL_FRAC + cfg.TEST_FRAC),\n    stratify=temp_ap[\"grade\"], random_state=cfg.SEED)\n\nextra = all_df[(all_df[\"source\"] != \"APTOS\") & (all_df[\"grade\"].isin(cfg.MERGE_GRADES))]\ntrain_df = pd.concat([train_ap, extra], ignore_index=True) \\\n             .sample(frac=1, random_state=cfg.SEED).reset_index(drop=True)\nval_df, test_df = val_ap.reset_index(drop=True), test_ap.reset_index(drop=True)\n\ntp, vp, sp = set(train_df[\"path\"]), set(val_df[\"path\"]), set(test_df[\"path\"])\nassert not (tp & vp), \"LEAKAGE train/val\"\nassert not (tp & sp), \"LEAKAGE train/test\"\nassert not (vp & sp), \"LEAKAGE val/test\"\nprint(\"Leakage assertion passed: splits are path-disjoint.\\n\")\n\nfor d in (train_df, val_df, test_df):\n    un = sorted(set(d[\"grade\"].astype(int)) - set(cfg.GRADE_TO_LABEL))\n    assert not un, f\"unmapped grades {un}\"\n    d[\"label\"] = d[\"grade\"].astype(int).map(cfg.GRADE_TO_LABEL).astype(int)   # COARSE (4-class)\n    d[\"label_fine\"] = d[\"grade\"].astype(int)                                  # FINE   (5-class, for Stage E)\n    assert d[\"label\"].between(0, cfg.NUM_CLASSES - 1).all()\n    assert d[\"label_fine\"].between(0, cfg.FINE_NUM_CLASSES - 1).all()\n\ndef show(d, nm):\n    g = dict(sorted(Counter(d[\"grade\"]).items()))\n    l = {cfg.CLASS_NAMES[k]: v for k, v in sorted(Counter(d[\"label\"]).items())}\n    print(f\"{nm:6s} n={len(d):5d}\\n       grades {g}\\n       labels {l}\")\nprint(f\"===== SPLIT ({RUN}) =====\")\nprint(f\"Train {len(train_df)} = APTOS {len(train_ap)} + merged {len(extra)}\")\nshow(train_df, \"TRAIN\"); show(val_df, \"VAL\"); show(test_df, \"TEST\")\nprint(\"Sources in TRAIN:\", dict(Counter(train_df[\"source\"])))\nif len(extra) < 300:\n    print(f\"\\n[WARN] only {len(extra)} merged images. Run B got 514 - check Messidor-2 in 2.5.\")\nclass_count = np.bincount(train_df[\"label\"].values, minlength=cfg.NUM_CLASSES)\n\n# ---- how much of the problem Stage E is responsible for -------------------\nprint(\"\\n===== MERGE ACCOUNTING =====\")\nfor nm, d in ((\"TRAIN\", train_df), (\"VAL\", val_df), (\"TEST\", test_df)):\n    hi = d[\"label_fine\"].isin(cfg.HI_GRADES)\n    n3 = int((d[\"label_fine\"] == 3).sum()); n4 = int((d[\"label_fine\"] == 4).sum())\n    print(f\"  {nm:5s}: merged bucket = {int(hi.sum()):4d}/{len(d):5d} \"\n          f\"({100*hi.mean():5.2f}% of split)  ->  Severe {n3}  PDR {n4}\")\n_te_hi = test_df[\"label_fine\"].isin(cfg.HI_GRADES).mean()\nprint(f\"\\n[BUDGET] every 10 pts of expert error on the merged bucket costs at most \"\n      f\"{100*_te_hi*0.10:.2f} pts of final test accuracy.\")\n_mv = min(Counter(val_df[\"label\"]).values())\nprint(f\"\\n[NOTE] smallest VAL class = {_mv}. Treat sub-1.5-point swings as noise.\")","metadata":{"papermill":{"duration":0.048121,"end_time":"2026-08-12T19:49:39.806819+00:00","exception":false,"start_time":"2026-08-12T19:49:39.758698+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"a21a8a75","cell_type":"markdown","source":"## 4. Preprocessing — fundus crop + CLAHE","metadata":{"papermill":{"duration":0.00763,"end_time":"2026-08-12T19:49:39.821801+00:00","exception":false,"start_time":"2026-08-12T19:49:39.814171+00:00","status":"completed"},"tags":[]}},{"id":"43af4b25","cell_type":"code","source":"def crop_to_fundus(img, thr=None):\n    thr = cfg.CROP_THRESHOLD if thr is None else thr\n    g = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY); m = g > thr\n    if m.sum() < 100: return img\n    r = np.where(m.any(1))[0]; c = np.where(m.any(0))[0]\n    out = img[r[0]:r[-1]+1, c[0]:c[-1]+1]\n    return out if out.size else img\n\ndef _clahe(gray):\n    return cv2.createCLAHE(clipLimit=cfg.CLAHE_CLIP,\n                           tileGridSize=(cfg.CLAHE_GRID,)*2).apply(gray)\n\ndef preprocess(img):\n    img = crop_to_fundus(img)\n    if cfg.PREPROCESS == \"rgb_clahe\":\n        lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n        lab[:, :, 0] = _clahe(lab[:, :, 0])\n        return cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)\n    if cfg.PREPROCESS == \"green_clahe\":\n        g = _clahe(img[:, :, 1]); return np.stack([g, g, g], -1)\n    raise ValueError(cfg.PREPROCESS)\n\n_NN = dict(mean=(0., 0., 0.), std=(1., 1., 1.), max_pixel_value=255.0)\ntrain_tf = A.Compose([\n    A.Resize(cfg.IMG_SIZE, cfg.IMG_SIZE),\n    A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=30,\n                       border_mode=cv2.BORDER_CONSTANT, p=0.7),\n    A.RandomBrightnessContrast(0.15, 0.15, p=0.5),\n    A.Normalize(**_NN), ToTensorV2()])\neval_tf = A.Compose([A.Resize(cfg.IMG_SIZE, cfg.IMG_SIZE), A.Normalize(**_NN), ToTensorV2()])\n\nclass RetinaDS(Dataset):\n    def __init__(s, df, tf): s.df = df.reset_index(drop=True); s.tf = tf\n    def __len__(s): return len(s.df)\n    def __getitem__(s, i):\n        r = s.df.iloc[i]; im = cv2.imread(r[\"path\"])\n        im = (np.zeros((cfg.IMG_SIZE, cfg.IMG_SIZE, 3), np.uint8) if im is None\n              else cv2.cvtColor(im, cv2.COLOR_BGR2RGB))\n        return s.tf(image=preprocess(im))[\"image\"], torch.tensor(int(r[\"label\"]), dtype=torch.long)\n\nprint(f\"Preprocess: crop ON | mode {cfg.PREPROCESS} | transforms emit [0,1]\")","metadata":{"papermill":{"duration":0.028662,"end_time":"2026-08-12T19:49:39.857739+00:00","exception":false,"start_time":"2026-08-12T19:49:39.829077+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"52feb8f5","cell_type":"markdown","source":"**4.2 — Sanity check**","metadata":{"papermill":{"duration":0.007866,"end_time":"2026-08-12T19:49:39.874207+00:00","exception":false,"start_time":"2026-08-12T19:49:39.866341+00:00","status":"completed"},"tags":[]}},{"id":"2680572d","cell_type":"code","source":"_ds = RetinaDS(train_df, eval_tf)\nfig, ax = plt.subplots(2, 4, figsize=(14, 7))\nfor k in range(4):\n    r = train_df.iloc[k]\n    raw = cv2.cvtColor(cv2.imread(r[\"path\"]), cv2.COLOR_BGR2RGB)\n    t, lb = _ds[k]\n    ax[0, k].imshow(raw); ax[0, k].set_title(f\"raw | {r['source']} g{r['grade']}\", fontsize=9)\n    ax[1, k].imshow(np.clip(t.numpy().transpose(1, 2, 0), 0, 1))\n    ax[1, k].set_title(f\"{cfg.PREPROCESS} | {cfg.CLASS_NAMES[lb.item()]}\", fontsize=9)\n    for j in (0, 1): ax[j, k].axis(\"off\")\nplt.suptitle(f\"{RUN}: top = original, bottom = model input\"); plt.tight_layout(); plt.show()","metadata":{"papermill":{"duration":1.984123,"end_time":"2026-08-12T19:49:41.86549+00:00","exception":false,"start_time":"2026-08-12T19:49:39.881367+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"f99135c8","cell_type":"markdown","source":"## 5. Balancing & DataLoaders","metadata":{"papermill":{"duration":0.020194,"end_time":"2026-08-12T19:49:41.905795+00:00","exception":false,"start_time":"2026-08-12T19:49:41.885601+00:00","status":"completed"},"tags":[]}},{"id":"e7d051b2","cell_type":"code","source":"labels = train_df[\"label\"].values\ncw = 1.0 / np.maximum(class_count, 1) ** cfg.SAMPLER_POWER if cfg.USE_SAMPLER else np.ones(cfg.NUM_CLASSES)\nsw = cw[labels]\n_eff = class_count * cw; _eff = _eff / _eff.sum()\nprint(\"Train counts     :\", dict(zip(cfg.CLASS_NAMES, class_count.tolist())))\nprint(\"Raw proportions  :\", np.round(class_count/class_count.sum(), 3).tolist())\nprint(f\"Sampled (p={cfg.SAMPLER_POWER}) :\", np.round(_eff, 3).tolist())\n\ntrain_ds, val_ds, test_ds = RetinaDS(train_df, train_tf), RetinaDS(val_df, eval_tf), RetinaDS(test_df, eval_tf)\n_pw = cfg.NUM_WORKERS > 0\nif cfg.USE_SAMPLER:\n    smp = WeightedRandomSampler(torch.as_tensor(sw, dtype=torch.double), len(train_ds), replacement=True)\n    train_loader = DataLoader(train_ds, cfg.BATCH_SIZE, sampler=smp, num_workers=cfg.NUM_WORKERS,\n                              pin_memory=True, drop_last=True, persistent_workers=_pw)\nelse:\n    train_loader = DataLoader(train_ds, cfg.BATCH_SIZE, shuffle=True, num_workers=cfg.NUM_WORKERS,\n                              pin_memory=True, drop_last=True, persistent_workers=_pw)\nval_loader  = DataLoader(val_ds,  cfg.BATCH_SIZE, shuffle=False, num_workers=cfg.NUM_WORKERS,\n                         pin_memory=True, persistent_workers=_pw)\ntest_loader = DataLoader(test_ds, cfg.BATCH_SIZE, shuffle=False, num_workers=cfg.NUM_WORKERS,\n                         pin_memory=True, persistent_workers=_pw)\nprint(f\"Batches -> train {len(train_loader)} val {len(val_loader)} test {len(test_loader)}\")","metadata":{"papermill":{"duration":0.036179,"end_time":"2026-08-12T19:49:41.961349+00:00","exception":false,"start_time":"2026-08-12T19:49:41.92517+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"7b26ab3d","cell_type":"markdown","source":"## Helpers (losses, metrics, freeze utilities, EMA & train/eval loops)\n\nReused unchanged from the original notebook. The memory-probe block from the\noriginal freeze-helper cell is dropped here because no full model exists yet at\nthis point; Stage D builds the model and runs on real data instead.","metadata":{"papermill":{"duration":0.020097,"end_time":"2026-08-12T19:49:42.002123+00:00","exception":false,"start_time":"2026-08-12T19:49:41.982026+00:00","status":"completed"},"tags":[]}},{"id":"ec26045b","cell_type":"code","source":"def _norm_probs(p, C):\n    p = np.asarray(p, dtype=np.float64)\n    return p / np.maximum(p.sum(1, keepdims=True), 1e-12)\n\ndef ovr_auc(y, p, C):\n    Yb = label_binarize(np.asarray(y), classes=list(range(C)))\n    a = np.full(C, np.nan)\n    for c in range(C):\n        pos = Yb[:, c].sum()\n        if 0 < pos < len(Yb):\n            try: a[c] = roc_auc_score(Yb[:, c], p[:, c])\n            except Exception: pass\n    return a\n\ndef compute_metrics(y_true, y_pred, y_prob, C=None):\n    C = C or cfg.NUM_CLASSES; labs = list(range(C))\n    y_true = np.asarray(y_true); y_pred = np.asarray(y_pred); y_prob = _norm_probs(y_prob, C)\n    cm = confusion_matrix(y_true, y_pred, labels=labs)\n    pr, rc, f1, sup = precision_recall_fscore_support(y_true, y_pred, labels=labs, zero_division=0)\n    a = ovr_auc(y_true, y_prob, C)\n    spec = np.zeros(C); tot = cm.sum()\n    for c in range(C):\n        tp = cm[c, c]; fn = cm[c].sum()-tp; fp = cm[:, c].sum()-tp; tn = tot-tp-fn-fp\n        spec[c] = tn/max(tn+fp, 1)\n    w = sup/max(sup.sum(), 1); ok = ~np.isnan(a)\n    return dict(accuracy=accuracy_score(y_true, y_pred),\n                qwk=cohen_kappa_score(y_true, y_pred, weights=\"quadratic\", labels=labs),\n                precision_w=float(np.average(pr, weights=w)),\n                recall_w=float(np.average(rc, weights=w)),\n                f1_w=float(np.average(f1, weights=w)), f1_macro=float(f1.mean()),\n                specificity_macro=float(spec.mean()),\n                auc=float(np.average(a[ok], weights=w[ok])) if ok.any() else float(\"nan\"),\n                per_class=dict(precision=pr, recall=rc, f1=f1, spec=spec, auc=a, support=sup),\n                confusion_matrix=cm)\n\ndef print_metrics(m, title=\"\"):\n    if title: print(f\"--- {title} ---\")\n    for k, lb in [(\"accuracy\",\"Accuracy\"),(\"qwk\",\"QWK\"),(\"precision_w\",\"Precision (wtd)\"),\n                  (\"recall_w\",\"Recall (wtd)\"),(\"f1_w\",\"F1 (weighted)\"),(\"f1_macro\",\"F1 (macro)\"),\n                  (\"specificity_macro\",\"Specificity avg\"),(\"auc\",\"AUC (ovr, wtd)\")]:\n        print(f\"  {lb:16s}: {m[k]*100:.2f}%\")\n\ndef blended_score(m):\n    '''QWK alone under-penalises the Moderate-attractor error; accuracy and macro-F1 weighted back in.'''\n    return 0.4*m[\"qwk\"] + 0.3*m[\"f1_macro\"] + 0.3*m[\"accuracy\"]\n\n_y = np.array([0,1,2,3]*3); _p = np.eye(4)[_y]*0.7+0.1\n_m = compute_metrics(_y, _y, _p.astype(np.float16), 4)\nassert not np.isnan(_m[\"auc\"]), \"AUC nan - fix before training\"\nprint(f\"metric self-test acc={_m['accuracy']:.2f} qwk={_m['qwk']:.2f} auc={_m['auc']:.3f} OK\")","metadata":{"papermill":{"duration":0.057153,"end_time":"2026-08-12T19:49:42.079725+00:00","exception":false,"start_time":"2026-08-12T19:49:42.022572+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"d9e2725d","cell_type":"code","source":"class WeightedFocalLoss(nn.Module):\n    def __init__(self, alpha, gamma=2.0):\n        super().__init__()\n        self.register_buffer(\"alpha\", torch.as_tensor(alpha, dtype=torch.float)); self.gamma = gamma\n    def forward(self, logits, target):\n        logits = logits.float()\n        logp = F.log_softmax(logits, 1); p = logp.exp()\n        lt = logp.gather(1, target[:, None]).squeeze(1)\n        pt = p.gather(1, target[:, None]).squeeze(1)\n        at = self.alpha.to(logits.device).gather(0, target)\n        return (-at * (1 - pt).pow(self.gamma) * lt).mean()\n\nclass WeightedKappaLoss(nn.Module):\n    def __init__(self, C, eps=1e-6):\n        super().__init__(); self.C = C; self.eps = eps\n        i = torch.arange(C, dtype=torch.float32)\n        self.register_buffer(\"W\", ((i[:, None] - i[None, :]) ** 2) / (C - 1) ** 2)\n    def forward(self, logits, target):\n        p = F.softmax(logits.float(), 1); y = F.one_hot(target, self.C).float()\n        O = y.t() @ p; E = torch.outer(y.sum(0), p.sum(0)) / y.size(0)\n        return (self.W * O).sum() / ((self.W * E).sum() + self.eps)\n\nalpha = ([1.0]*cfg.NUM_CLASSES if cfg.FOCAL_ALPHA == \"equal\"\n         else ((1/np.maximum(class_count,1))/(1/np.maximum(class_count,1)).mean()).tolist()\n              if cfg.FOCAL_ALPHA == \"inverse\" else list(cfg.FOCAL_ALPHA))\nfocal = WeightedFocalLoss(alpha, cfg.FOCAL_GAMMA).to(device)\nkappa = WeightedKappaLoss(cfg.NUM_CLASSES).to(device)\n\ndef compute_loss(logits, aux_c, aux_v, y):\n    main = focal(logits, y)\n    tot = main + cfg.AUX_WEIGHT * (focal(aux_c, y) + focal(aux_v, y))\n    if cfg.KAPPA_WEIGHT > 0: tot = tot + cfg.KAPPA_WEIGHT * kappa(logits, y)\n    return tot, float(main)\n\nprint(f\"focal(g={cfg.FOCAL_GAMMA}, alpha={alpha}) + {cfg.AUX_WEIGHT}*aux \"\n      f\"+ {cfg.KAPPA_WEIGHT}*kappa\")\n_lg = torch.randn(16, cfg.NUM_CLASSES, device=device, requires_grad=True)\n_tg = torch.randint(0, cfg.NUM_CLASSES, (16,), device=device)\nassert float(kappa(F.one_hot(_tg, cfg.NUM_CLASSES).float()*20, _tg)) < 1e-4\nfocal(_lg, _tg).backward(); assert torch.isfinite(_lg.grad).all()\nprint(\"Loss self-tests passed.\")","metadata":{"papermill":{"duration":0.892726,"end_time":"2026-08-12T19:49:42.990884+00:00","exception":false,"start_time":"2026-08-12T19:49:42.098158+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"90949fcd","cell_type":"code","source":"def set_branches_trainable(m, eff_t, vit_t):\n    for p in m.eff.parameters(): p.requires_grad = eff_t\n    for p in m.vit.parameters(): p.requires_grad = vit_t\n    for p in m.fusion_parameters(): p.requires_grad = True\n\ndef set_partial_trainable(module, freeze_fraction):\n    ps = list(module.named_parameters()); cut = int(len(ps) * freeze_fraction)\n    for i, (_, p) in enumerate(ps): p.requires_grad = (i >= cut)\n    return sum(p.requires_grad for _, p in ps), len(ps)\n\ndef _split_wd(named, lr, wd):\n    dec, nod = [], []\n    for n, p in named:\n        if p.requires_grad: (nod if (p.ndim <= 1 or n.endswith(\".bias\")) else dec).append(p)\n    out = []\n    if dec: out.append(dict(params=dec, lr=lr, weight_decay=wd))\n    if nod: out.append(dict(params=nod, lr=lr, weight_decay=0.0))\n    return out\n\ndef vit_groups(vit, base_lr, wd, decay):\n    nb = len(vit.blocks); nl = nb + 1\n    def lid(n):\n        if n.startswith(\"patch_embed\") or n in (\"cls_token\", \"pos_embed\", \"reg_token\"): return 0\n        if n.startswith(\"blocks.\"): return int(n.split(\".\")[1]) + 1\n        return nl\n    g = {}\n    for n, p in vit.named_parameters():\n        if not p.requires_grad: continue\n        L = lid(n); nod = (p.ndim <= 1 or n.endswith(\".bias\")); k = (L, nod)\n        g.setdefault(k, dict(params=[], weight_decay=0.0 if nod else wd,\n                             lr=base_lr * (decay ** (nl - L))))[\"params\"].append(p)\n    return list(g.values())\n\ndef make_ft_groups(m):\n    fus_ids = {id(p) for p in m.fusion_parameters()}\n    g  = _split_wd(m.eff.named_parameters(), cfg.LR_FT_EFF, cfg.WEIGHT_DECAY); n_e = len(g)\n    g += vit_groups(m.vit, cfg.LR_FT_VIT, cfg.WEIGHT_DECAY, cfg.LAYER_DECAY); n_v = len(g) - n_e\n    g += _split_wd([(n, p) for n, p in m.named_parameters() if id(p) in fus_ids],\n                   cfg.LR_FT_HEAD, cfg.WEIGHT_DECAY)\n    g = [x for x in g if x[\"params\"]]\n    cov = sum(len(x[\"params\"]) for x in g)\n    tr  = sum(1 for p in m.parameters() if p.requires_grad)\n    assert cov == tr, f\"groups cover {cov}/{tr}\"\n    print(f\"  {len(g)} groups (eff {n_e}, vit {n_v}, fusion {len(g)-n_e-n_v}) | {cov}/{tr} tensors\")\n    print(f\"  lr {min(x['lr'] for x in g):.2e} .. {max(x['lr'] for x in g):.2e}\")\n    return g","metadata":{"papermill":{"duration":0.035165,"end_time":"2026-08-12T19:49:43.044571+00:00","exception":false,"start_time":"2026-08-12T19:49:43.009406+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"efee1688","cell_type":"code","source":"class ModelEMA:\n    def __init__(self, model, decay=0.9995):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone().float() for k, v in model.state_dict().items()}\n    @torch.no_grad()\n    def update(self, model):\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                if torch.isfinite(v).all():                  # NaN guard: skip poisoned tensors\n                    self.shadow[k].mul_(self.decay).add_(v.detach().float(), alpha=1-self.decay)\n            else:\n                self.shadow[k] = v.detach().clone()\n    def copy_to(self, model):\n        sd = model.state_dict()\n        model.load_state_dict({k: self.shadow[k].to(sd[k].dtype) for k in sd}, strict=True)\n\ndef fmt(s): s = int(s); return f\"{s//3600}h{(s%3600)//60:02d}m{s%60:02d}s\"\n\n@torch.no_grad()\ndef _fwd_probs(m, x, tta):\n    combos = [(False, False), (True, False), (False, True), (True, True)] if tta else [(False, False)]\n    acc = None\n    for fh, fv in combos:\n        z = x\n        if fh: z = torch.flip(z, [3])\n        if fv: z = torch.flip(z, [2])\n        with amp_autocast(cfg.AMP):\n            lg, _, _ = m(z)\n        p = F.softmax(lg.float(), 1)\n        acc = p if acc is None else acc + p\n    return acc / len(combos)\n\n@torch.no_grad()\ndef evaluate(m, loader, desc=\"val\", tta=False):\n    m.eval(); ys, ps = [], []\n    for x, y in tqdm(loader, desc=f\"[{desc}]\", leave=False):\n        ps.append(_fwd_probs(m, x.to(device, non_blocking=True), tta).cpu().numpy())\n        ys.append(y.numpy())\n    y = np.concatenate(ys); p = np.concatenate(ps)\n    met = compute_metrics(y, p.argmax(1), p, cfg.NUM_CLASSES)\n    return met, y, p\n\ndef train_one_epoch(m, loader, opt, scaler, ema, ep, tag):\n    m.train()\n    run, seen, corr, nskip = 0.0, 0, 0, 0\n    kloss = 0.0\n    acc_n = max(int(cfg.ACCUM_STEPS), 1)\n    t0 = time.time(); opt.zero_grad(set_to_none=True)\n    bar = tqdm(loader, desc=f\"{tag} ep{ep}\", leave=False)\n    for step, (x, y) in enumerate(bar):\n        x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)\n        with amp_autocast(cfg.AMP):\n            lg, ac, av = m(x)\n            loss, main = compute_loss(lg, ac, av, y)\n        # ---- NaN guard: drop the step instead of poisoning the weights ----\n        if not torch.isfinite(loss):\n            nskip += 1; opt.zero_grad(set_to_none=True); continue\n        scaler.scale(loss / acc_n).backward()\n        if (step + 1) % acc_n == 0 or (step + 1) == len(loader):\n            scaler.unscale_(opt)\n            gn = torch.nn.utils.clip_grad_norm_(\n                [p for p in m.parameters() if p.requires_grad], cfg.GRAD_CLIP)\n            if torch.isfinite(gn):\n                scaler.step(opt); scaler.update()\n                if ema is not None: ema.update(m)\n            else:\n                nskip += 1; scaler.update()\n            opt.zero_grad(set_to_none=True)\n        bs = x.size(0); seen += bs; run += main*bs; corr += int((lg.argmax(1) == y).sum())\n        if cfg.KAPPA_WEIGHT > 0: kloss += float(kappa(lg, y))*bs\n        bar.set_postfix(loss=f\"{run/max(seen,1):.4f}\", acc=f\"{corr/max(seen,1):.4f}\",\n                        kap=f\"{kloss/max(seen,1):.3f}\", skip=nskip,\n                        img_s=f\"{seen/max(time.time()-t0,1e-6):.1f}\")\n    return run/max(seen, 1), corr/max(seen, 1), time.time()-t0, nskip\n\ndef run_phase(m, opt, sched, n_ep, tag, history, best, ema, t0):\n    scaler = make_scaler(cfg.AMP); bad = 0\n    for ep in range(1, n_ep + 1):\n        tl, ta, et, nsk = train_one_epoch(m, train_loader, opt, scaler, ema, ep, tag)\n        eval_m = build_eval_copy(m, ema)\n        met, _, _ = evaluate(eval_m, val_loader, \"val\")\n        del eval_m\n        if sched: sched.step()\n        sc = blended_score(met)\n        history.append(dict(phase=tag, epoch=ep, train_loss=tl, train_acc=ta,\n                            val_acc=met[\"accuracy\"], val_qwk=met[\"qwk\"],\n                            val_f1m=met[\"f1_macro\"], score=sc, skipped=nsk,\n                            lr=max(g[\"lr\"] for g in opt.param_groups), sec=et))\n        star = \"\"\n        if sc > best[\"score\"]:\n            best.update(score=sc, qwk=met[\"qwk\"], acc=met[\"accuracy\"])\n            torch.save({k: v.cpu().clone() for k, v in\n                        (ema.shadow if ema else m.state_dict()).items()}, cfg.CKPT)\n            star = \"  <-- best (saved)\"; bad = 0\n        else:\n            bad += 1\n        warn = f\"  [!] {nsk} non-finite steps skipped\" if nsk else \"\"\n        print(f\"[{tag} {ep:02d}/{n_ep}] loss={tl:.4f} tr_acc={ta:.4f} | val_acc={met['accuracy']*100:.2f}% \"\n              f\"qwk={met['qwk']*100:.2f}% f1m={met['f1_macro']*100:.2f}% | score={sc:.4f} | \"\n              f\"{et:.0f}s (tot {fmt(time.time()-t0)}){star}{warn}\")\n        if bad >= cfg.PATIENCE:\n            print(f\"  early stop: no improvement in {cfg.PATIENCE} epochs\"); break\n    return history, best\n\ndef build_eval_copy(m, ema):\n    if ema is None: return m\n    ev = MultiLevelFusion().to(device)\n    ev.load_state_dict({k: v.to(dtype) for (k, v), dtype in\n                        zip(ema.shadow.items(), [p.dtype for p in m.state_dict().values()])},\n                       strict=True)\n    return ev.eval()","metadata":{"papermill":{"duration":0.044186,"end_time":"2026-08-12T19:49:43.1064+00:00","exception":false,"start_time":"2026-08-12T19:49:43.062214+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"b227599b","cell_type":"markdown","source":"## Stage A — train the two backbones once\n\nIndependent 4-class classifiers on the same splits/loaders defined above. Same\nper-branch normalisation, focal(+kappa) loss, EMA, and two-phase schedule idea,\nbut each backbone stands alone. Weights are cached to disk and only re-trained if\nthe files are missing (`FORCE_RETRAIN=True` to override).","metadata":{"papermill":{"duration":0.018467,"end_time":"2026-08-12T19:49:43.14324+00:00","exception":false,"start_time":"2026-08-12T19:49:43.124773+00:00","status":"completed"},"tags":[]}},{"id":"6c9125f5","cell_type":"code","source":"\nFORCE_RETRAIN   = False\n# Per-backbone schedules. ViT already trained well at 3e-5; EfficientNet-B4 was\n# under-fitting at that LR (it was still rising at epoch 12 with QWK barely positive),\n# so it gets a higher LR and more epochs. Keys are the `tag` passed to train_backbone.\nBACKBONE_CFG = {\n    \"EFF\": dict(lr=3e-4, epochs=35),   # B5 has more capacity; give it a longer schedule\n    \"VIT\": dict(lr=3e-5, epochs=20),   # more epochs than before; ViT was still improving\n}\n# Retrain ONLY EfficientNet by default (ViT's cached weights are good). Set\n# RETRAIN_TAGS = {\"EFF\",\"VIT\"} to redo both, or {} to reuse whatever is cached.\n# Both good checkpoints already exist from prior runs, so reuse both by default.\nRETRAIN_TAGS = {\"EFF\", \"VIT\"}   # NUM_CLASSES changed 5 -> 4, so every cached head is invalid;\n                                # so retrain both from scratch this run. Set to set() to reuse\n                                # once you have B5/384 checkpoints cached.\nEFF_CKPT = f\"{cfg.OUT}/backbone_eff_c4.pt\"\nVIT_CKPT = f\"{cfg.OUT}/backbone_vit_c4.pt\"\n\nclass Backbone(nn.Module):\n    \"\"\"Single backbone + linear head, with the branch's own normalisation baked in.\"\"\"\n    def __init__(self, name, mean, std):\n        super().__init__()\n        self.net = timm.create_model(name, pretrained=True, num_classes=cfg.NUM_CLASSES,\n                                     drop_rate=0.1, drop_path_rate=0.1)\n        self.register_buffer(\"mean\", torch.tensor(mean).view(1,3,1,1))\n        self.register_buffer(\"std\",  torch.tensor(std ).view(1,3,1,1))\n    def forward(self, x):\n        return self.net((x - self.mean) / self.std)\n\ndef train_backbone(tag, name, mean, std, ckpt):\n    force = FORCE_RETRAIN or (tag in RETRAIN_TAGS)\n    if os.path.exists(ckpt) and not force:\n        print(f\"[{tag}] found {ckpt} — skipping (in RETRAIN_TAGS or FORCE_RETRAIN=True to redo)\")\n        return\n    lr     = BACKBONE_CFG[tag][\"lr\"]\n    epochs = BACKBONE_CFG[tag][\"epochs\"]\n    print(f\"[{tag}] training  lr={lr}  epochs={epochs}\")\n    m = Backbone(name, mean, std).to(device)\n    opt = torch.optim.AdamW(m.parameters(), lr=lr, weight_decay=cfg.WEIGHT_DECAY)\n    sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    scaler = make_scaler(cfg.AMP)\n    ema = ModelEMA(m, cfg.EMA_DECAY)\n    focal_bb = WeightedFocalLoss(alpha, cfg.FOCAL_GAMMA).to(device)\n    best = -1.0\n    for ep in range(1, epochs+1):\n        m.train(); opt.zero_grad(set_to_none=True); seen=corr=0; run=0.0\n        for step,(x,y) in enumerate(tqdm(train_loader, desc=f\"{tag} ep{ep}\", leave=False)):\n            x,y = x.to(device,non_blocking=True), y.to(device,non_blocking=True)\n            with amp_autocast(cfg.AMP):\n                lg = m(x)\n                loss = focal_bb(lg, y)\n                if cfg.KAPPA_WEIGHT>0: loss = loss + cfg.KAPPA_WEIGHT*kappa(lg,y)\n            if not torch.isfinite(loss):\n                opt.zero_grad(set_to_none=True); continue\n            scaler.scale(loss/cfg.ACCUM_STEPS).backward()\n            if (step+1)%cfg.ACCUM_STEPS==0 or (step+1)==len(train_loader):\n                scaler.unscale_(opt)\n                gn = torch.nn.utils.clip_grad_norm_(m.parameters(), cfg.GRAD_CLIP)\n                if torch.isfinite(gn):\n                    scaler.step(opt); scaler.update(); ema.update(m)\n                else:\n                    scaler.update()\n                opt.zero_grad(set_to_none=True)\n            bs=x.size(0); seen+=bs; run+=float(loss)*bs; corr+=int((lg.argmax(1)==y).sum())\n        sch.step()\n        # eval with EMA weights\n        ev = Backbone(name, mean, std).to(device)\n        ev.load_state_dict({k:v.to(p.dtype) for (k,v),p in\n                            zip(ema.shadow.items(), ev.state_dict().values())}, strict=True)\n        ev.eval(); ys=[];ps=[]\n        with torch.no_grad():\n            for x,y in val_loader:\n                with amp_autocast(cfg.AMP):\n                    p = F.softmax(ev(x.to(device)).float(),1).cpu().numpy()\n                ps.append(p); ys.append(y.numpy())\n        met = compute_metrics(np.concatenate(ys), np.concatenate(ps).argmax(1),\n                              np.concatenate(ps), cfg.NUM_CLASSES)\n        sc = blended_score(met)\n        star=\"\"\n        if sc>best:\n            best=sc\n            torch.save({k:v.cpu().clone() for k,v in ema.shadow.items()}, ckpt)\n            star=\"  <-- saved\"\n        print(f\"[{tag} {ep:02d}] tr_loss={run/seen:.4f} tr_acc={corr/seen:.4f} \"\n              f\"| val_acc={met['accuracy']*100:.2f} qwk={met['qwk']*100:.2f} \"\n              f\"f1m={met['f1_macro']*100:.2f} score={sc:.4f}{star}\")\n        del ev\n    print(f\"[{tag}] done. best blended={best:.4f} -> {ckpt}\")\n\ntrain_backbone(\"EFF\", cfg.EFF_NAME, EFF_MEAN, EFF_STD, EFF_CKPT)\ntrain_backbone(\"VIT\", cfg.VIT_NAME, VIT_MEAN, VIT_STD, VIT_CKPT)\n","metadata":{"papermill":{"duration":16854.51707,"end_time":"2026-08-13T00:30:37.678588+00:00","exception":false,"start_time":"2026-08-12T19:49:43.161518+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"51fb7232","cell_type":"markdown","source":"### Stage A.1 — unfused baselines\n\nBefore any fusion, measure each backbone **on its own** on the test set. These two\nnumbers are the baseline the fused model must beat, and they're stored in\n`BASELINE` for the final comparison table.","metadata":{"papermill":{"duration":0.02571,"end_time":"2026-08-13T00:30:37.730124+00:00","exception":false,"start_time":"2026-08-13T00:30:37.704414+00:00","status":"completed"},"tags":[]}},{"id":"b43e481e","cell_type":"code","source":"\n@torch.no_grad()\ndef eval_backbone(tag, name, mean, std, ckpt):\n    m = Backbone(name, mean, std).to(device)\n    sd = torch.load(ckpt, map_location=device)\n    # ckpt holds EMA shadow of the wrapped Backbone; load into a fresh Backbone\n    m.load_state_dict({k: v.to(device) for k, v in sd.items()}, strict=False)\n    m.eval(); ys=[]; ps=[]\n    for x, y in tqdm(test_loader, desc=f\"[eval {tag}]\", leave=False):\n        with amp_autocast(cfg.AMP):\n            p = F.softmax(m(x.to(device)).float(), 1).cpu().numpy()\n        ps.append(p); ys.append(y.numpy())\n    met = compute_metrics(np.concatenate(ys), np.concatenate(ps).argmax(1),\n                          np.concatenate(ps), cfg.NUM_CLASSES)\n    print_metrics(met, f\"UNFUSED TEST — {tag} alone\")\n    return met\n\nBASELINE = {}\nBASELINE[\"EfficientNet-B4 (alone)\"] = eval_backbone(\"EFF\", cfg.EFF_NAME, EFF_MEAN, EFF_STD, EFF_CKPT)\nBASELINE[\"ViT-B/16 (alone)\"]        = eval_backbone(\"VIT\", cfg.VIT_NAME, VIT_MEAN, VIT_STD, VIT_CKPT)\nprint(\"\\nUnfused baselines (test):\")\nfor k, m in BASELINE.items():\n    print(f\"  {k:26s} acc={m['accuracy']*100:5.2f}  qwk={m['qwk']*100:5.2f}  macroF1={m['f1_macro']*100:5.2f}\")\n","metadata":{"papermill":{"duration":138.334707,"end_time":"2026-08-13T00:32:56.090579+00:00","exception":false,"start_time":"2026-08-13T00:30:37.755872+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"a69136f1","cell_type":"markdown","source":"## Stage B — cache pooled features at every candidate tap\n\nLoad the trained backbones, freeze them, and pass each image through once.\nFor the CNN we globally-average-pool every stage in `out_indices=(0..6)`.\nFor the ViT we hook every block `0..11` and take the CLS token. We also cache the\ncross-attention output (CNN-deepest query × ViT-deepest patch keys), which is the\none non-linear tap. All caches are `float16` tensors on disk keyed by split.","metadata":{"papermill":{"duration":0.025399,"end_time":"2026-08-13T00:32:56.141244+00:00","exception":false,"start_time":"2026-08-13T00:32:56.115845+00:00","status":"completed"},"tags":[]}},{"id":"19d54890","cell_type":"code","source":"\n# ---- Probe the ACTUAL number of taps each backbone exposes (never hardcode) ----\n# EfficientNet's feature-stage count varies by variant/timm version (B4 = 5 here,\n# indices 0..4), so we ask timm directly instead of assuming 7. Same for ViT depth.\n_probe_eff = timm.create_model(cfg.EFF_NAME, pretrained=False, features_only=True)\nCNN_STAGES = tuple(range(len(_probe_eff.feature_info.channels())))\n_probe_vit = timm.create_model(cfg.VIT_NAME, pretrained=False, num_classes=0)\nVIT_BLOCKS = tuple(range(len(_probe_vit.blocks)))\ndel _probe_eff, _probe_vit\nprint(f\"Detected {len(CNN_STAGES)} CNN feature stages -> indices {CNN_STAGES}\")\nprint(f\"Detected {len(VIT_BLOCKS)} ViT blocks           -> indices {VIT_BLOCKS}\")\n\nCACHE_DIR  = f\"{cfg.OUT}/feat_cache_c4\"   # 4-class backbones -> new cache namespace\nos.makedirs(CACHE_DIR, exist_ok=True)\n\nclass FeatureExtractor(nn.Module):\n    \"\"\"Frozen backbones exposing every CNN stage, every ViT-block CLS, and the xattn tap.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.eff = timm.create_model(cfg.EFF_NAME, pretrained=False, features_only=True,\n                                     out_indices=CNN_STAGES)\n        self.vit = timm.create_model(cfg.VIT_NAME, pretrained=False, num_classes=0)\n        for nm,mu,sd in ((\"eff\",EFF_MEAN,EFF_STD),(\"vit\",VIT_MEAN,VIT_STD)):\n            self.register_buffer(f\"{nm}_mean\", torch.tensor(mu).view(1,3,1,1))\n            self.register_buffer(f\"{nm}_std\",  torch.tensor(sd).view(1,3,1,1))\n        # cross-attention bridge (deepest CNN map -> deepest ViT patch tokens)\n        self.eff_ch = list(self.eff.feature_info.channels())\n        vit_d = self.vit.embed_dim\n        D = cfg.PROJ_DIM\n        self.q_proj  = nn.Linear(self.eff_ch[-1], D)\n        self.kv_proj = nn.Linear(vit_d, D)\n        self.xattn   = nn.MultiheadAttention(D, cfg.XATTN_HEADS, batch_first=True)\n        self.xnorm   = nn.LayerNorm(D)\n        self._taps = {}\n        for i in VIT_BLOCKS:\n            self.vit.blocks[i].register_forward_hook(self._mk(i))\n    def _mk(self, i):\n        def h(_m,_in,out): self._taps[i]=out\n        return h\n    def load_backbones(self):\n        \"\"\"Copy trained backbone weights into eff/vit (heads dropped).\"\"\"\n        eff_sd = torch.load(EFF_CKPT, map_location=\"cpu\")\n        vit_sd = torch.load(VIT_CKPT, map_location=\"cpu\")\n        # backbone was Backbone.net.* ; strip the wrapper + classifier\n        def strip(sd, drop_prefixes):\n            out={}\n            for k,v in sd.items():\n                if k.startswith(\"net.\"): k=k[4:]\n                if any(k.startswith(p) for p in drop_prefixes): continue\n                out[k]=v\n            return out\n        me = self.eff.load_state_dict(strip(eff_sd, (\"classifier\",\"fc\",\"head\")), strict=False)\n        mv = self.vit.load_state_dict(strip(vit_sd, (\"head\",)), strict=False)\n        print(\"eff missing/unexpected:\", len(me.missing_keys), len(me.unexpected_keys))\n        print(\"vit missing/unexpected:\", len(mv.missing_keys), len(mv.unexpected_keys))\n    @torch.no_grad()\n    def forward(self, x):\n        self._taps.clear()\n        maps = self.eff((x - self.eff_mean)/self.eff_std)          # list, one per CNN stage\n        _ = self.vit.forward_features((x - self.vit_mean)/self.vit_std)\n        cnn_vecs = {f\"cnn{ i}\": maps[j].mean((2,3)) for j,i in enumerate(CNN_STAGES)}\n        vit_vecs = {f\"vit{i}\": self._taps[i][:,0] for i in VIT_BLOCKS}\n        # cross-attention tap\n        last = maps[-1]; B,C,H,W = last.shape\n        q  = self.q_proj(last.flatten(2).transpose(1,2))\n        kv = self.kv_proj(self._taps[VIT_BLOCKS[-1]][:,1:])\n        with amp_autocast(False):\n            a,_ = self.xattn(q.float(), kv.float(), kv.float(), need_weights=False)\n        xv = self.xnorm(a.mean(1))                                  # [B, D]\n        feats = {**cnn_vecs, **vit_vecs, \"xattn\": xv}\n        return feats\n\nextractor = FeatureExtractor().to(device).eval()\nextractor.load_backbones()\nfor p in extractor.parameters(): p.requires_grad=False\nTAP_NAMES = [f\"cnn{i}\" for i in CNN_STAGES] + [f\"vit{i}\" for i in VIT_BLOCKS] + [\"xattn\"]\nprint(\"Candidate taps:\", TAP_NAMES)\n","metadata":{"papermill":{"duration":4.032946,"end_time":"2026-08-13T00:33:00.199973+00:00","exception":false,"start_time":"2026-08-13T00:32:56.167027+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"8050d208","cell_type":"code","source":"\n@torch.no_grad()\ndef cache_split(name, df, tf):\n    path = f\"{CACHE_DIR}/{name}.pt\"\n    # If any backbone was retrained this run, old cached features are stale -> rebuild.\n    cache_valid = (not FORCE_RETRAIN) and (len(RETRAIN_TAGS) == 0)\n    if os.path.exists(path) and cache_valid:\n        print(f\"[cache] {name}: found, loading\"); return torch.load(path)\n    ds = RetinaDS(df, tf)\n    ld = DataLoader(ds, cfg.BATCH_SIZE, shuffle=False, num_workers=cfg.NUM_WORKERS, pin_memory=True)\n    store = {t: [] for t in TAP_NAMES}; ys=[]\n    for x,y in tqdm(ld, desc=f\"[cache {name}]\", leave=False):\n        with amp_autocast(cfg.AMP):\n            feats = extractor(x.to(device, non_blocking=True))\n        for t in TAP_NAMES: store[t].append(feats[t].float().half().cpu())\n        ys.append(y)\n    out = {t: torch.cat(v) for t,v in store.items()}\n    out[\"y\"] = torch.cat(ys)\n    torch.save(out, path)\n    print(f\"[cache] {name}: {out['y'].shape[0]} rows, dims \" +\n          \", \".join(f\"{t}={out[t].shape[1]}\" for t in TAP_NAMES[:3]) + \" ...\")\n    return out\n\ncache_train = cache_split(\"train\", train_df, eval_tf)   # eval_tf: deterministic features\ncache_val   = cache_split(\"val\",   val_df,   eval_tf)\ncache_test  = cache_split(\"test\",  test_df,  eval_tf)\nDIMS = {t: cache_train[t].shape[1] for t in TAP_NAMES}\nprint(\"Cached. Per-tap dims:\", DIMS)\n","metadata":{"papermill":{"duration":348.712329,"end_time":"2026-08-13T00:38:48.939043+00:00","exception":false,"start_time":"2026-08-13T00:33:00.226714+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"3cf84ff2","cell_type":"markdown","source":"## Stage C — exhaustive top-k subset search\n\nFor a given subset of taps we train a small MLP head on the concatenated cached\nfeatures (seconds each) and score it on validation with `blended_score`. We first\nrank every single tap, then **exhaustively try all combinations of the best `TOPK`\ntaps** (all sizes from 1 up to `MAX_TAPS`). This finds genuine multi-tap fusions\nthat greedy forward-selection misses when a single tap already scores well.\nEach subset is scored as the mean over a few seeds to damp head-training noise.","metadata":{"papermill":{"duration":0.027808,"end_time":"2026-08-13T00:38:48.994056+00:00","exception":false,"start_time":"2026-08-13T00:38:48.966248+00:00","status":"completed"},"tags":[]}},{"id":"b1748f49","cell_type":"code","source":"\nHEAD_EPOCHS = 40\nHEAD_LR     = 3e-3\nSEEDS       = (0, 1, 2)   # average each subset over these seeds (noise control)\nTOPK        = 7           # exhaustively combine the best TOPK single taps\nMAX_TAPS    = 5           # largest subset size to try (B5 gives stronger CNN taps)\n\ndef make_head(in_dim):\n    return nn.Sequential(\n        nn.Linear(in_dim, 512), nn.BatchNorm1d(512), nn.Dropout(0.4), nn.Tanh(),\n        nn.Linear(512, 128),    nn.BatchNorm1d(128), nn.Dropout(0.3), nn.Tanh(),\n        nn.Linear(128, cfg.NUM_CLASSES)).to(device)\n\ndef _cat(cache, taps):\n    return torch.cat([cache[t].float() for t in taps], 1)\n\ndef eval_subset(taps, seed=0, return_test=False):\n    \"\"\"Train a head on cached train features for `taps`; return val blended score.\"\"\"\n    torch.manual_seed(seed)\n    Xtr, ytr = _cat(cache_train, taps).to(device), cache_train[\"y\"].to(device)\n    Xva, yva = _cat(cache_val,   taps).to(device), cache_val[\"y\"].to(device)\n    head = make_head(Xtr.shape[1])\n    opt = torch.optim.AdamW(head.parameters(), lr=HEAD_LR, weight_decay=cfg.WEIGHT_DECAY)\n    sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=HEAD_EPOCHS)\n    # class-weighted sampling via loss weights (cheap, cache is in-memory)\n    w = torch.tensor((1.0/np.maximum(class_count,1)**cfg.SAMPLER_POWER), dtype=torch.float, device=device)\n    w = w/w.mean()\n    bs = 512\n    n = Xtr.shape[0]\n    best_sc, best_state = -1.0, None\n    for ep in range(HEAD_EPOCHS):\n        head.train(); perm = torch.randperm(n, device=device)\n        for i in range(0, n, bs):\n            idx = perm[i:i+bs]\n            lg = head(Xtr[idx])\n            loss = F.cross_entropy(lg, ytr[idx], weight=w)\n            if cfg.KAPPA_WEIGHT>0: loss = loss + cfg.KAPPA_WEIGHT*kappa(lg, ytr[idx])\n            opt.zero_grad(set_to_none=True); loss.backward(); opt.step()\n        sch.step()\n        head.eval()\n        with torch.no_grad():\n            pv = F.softmax(head(Xva).float(),1).cpu().numpy()\n        m = compute_metrics(yva.cpu().numpy(), pv.argmax(1), pv, cfg.NUM_CLASSES)\n        sc = blended_score(m)\n        if sc>best_sc:\n            best_sc=sc; best_state={k:v.detach().clone() for k,v in head.state_dict().items()}\n    if return_test:\n        head.load_state_dict(best_state); head.eval()\n        with torch.no_grad():\n            pv = F.softmax(head(Xva).float(),1).cpu().numpy()\n            Xte = _cat(cache_test, taps).to(device)\n            pt = F.softmax(head(Xte).float(),1).cpu().numpy()\n        return best_sc, pv, pt\n    return best_sc\n\ndef score_subset(taps):\n    \"\"\"Mean val blended score over SEEDS (noise-robust ranking signal).\"\"\"\n    return float(np.mean([eval_subset(taps, seed=s) for s in SEEDS]))\n\n# sanity: score each single tap (multi-seed)\nprint(\"Single-tap val blended scores (mean over seeds):\")\nsingles = {t: score_subset([t]) for t in tqdm(TAP_NAMES, desc=\"[singles]\", leave=False)}\nfor t,s in sorted(singles.items(), key=lambda kv:-kv[1]):\n    print(f\"  {t:7s} {s:.4f}\")\n","metadata":{"papermill":{"duration":106.596308,"end_time":"2026-08-13T00:40:35.617531+00:00","exception":false,"start_time":"2026-08-13T00:38:49.021223+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"439616a7","cell_type":"code","source":"\nfrom itertools import combinations\n\n# candidate pool = the TOPK best single taps\nranked = [t for t,_ in sorted(singles.items(), key=lambda kv:-kv[1])]\npool = ranked[:TOPK]\nprint(f\"Exhaustive search over subsets of top-{TOPK} taps: {pool}\")\n\nresults = []\nsubsets = []\nfor r in range(1, MAX_TAPS+1):\n    subsets += [list(c) for c in combinations(pool, r)]\nprint(f\"Evaluating {len(subsets)} subsets (sizes 1..{MAX_TAPS})...\\n\")\n\nfor taps in tqdm(subsets, desc=\"[subsets]\"):\n    sc = score_subset(taps)\n    results.append({\"taps\": taps, \"k\": len(taps), \"score\": sc})\n\nsearch_hist = pd.DataFrame(results).sort_values(\"score\", ascending=False).reset_index(drop=True)\nbest_row = search_hist.iloc[0]\nbest_taps = list(best_row[\"taps\"])\nbest_val_score = float(best_row[\"score\"])\n\nprint(\"\\n===== SEARCH RESULT =====\")\nprint(\"Selected taps :\", best_taps)\nprint(f\"Val blended   : {best_val_score:.4f}\")\nprint(\"\\nTop 10 subsets:\")\nshow = search_hist.head(10).copy()\nshow[\"taps\"] = show[\"taps\"].apply(lambda t: \"+\".join(t))\ndisplay(show[[\"taps\",\"k\",\"score\"]].round(4))\n# compare best multi-tap vs best single tap so any fusion gain is explicit\nbest_single = search_hist[search_hist.k==1].iloc[0]\nprint(f\"\\nBest single tap : {best_single['taps']}  score={best_single['score']:.4f}\")\nprint(f\"Best subset     : {best_taps}  score={best_val_score:.4f}  \"\n      f\"(fusion gain {best_val_score-best_single['score']:+.4f})\")\nif len(best_taps)==1:\n    print(\"NOTE: best subset is a single tap — multi-tap fusion did not help on val. \"\n          \"Raise TOPK/MAX_TAPS or strengthen the weaker backbone.\")\nsearch_hist.assign(taps=search_hist[\"taps\"].apply(lambda t: \"+\".join(t))\n                   ).to_csv(f\"{cfg.OUT}/fusion_search_history.csv\", index=False)\n\n# proxy test metrics for the winning subset (head-on-frozen-features)\n_, pv_best, pt_best = eval_subset(best_taps, seed=0, return_test=True)\nproxy = compute_metrics(cache_test[\"y\"].numpy(), pt_best.argmax(1), pt_best, cfg.NUM_CLASSES)\nprint_metrics(proxy, \"PROXY TEST (frozen-feature head, winning taps)\")\n","metadata":{"papermill":{"duration":454.839682,"end_time":"2026-08-13T00:48:10.482853+00:00","exception":false,"start_time":"2026-08-13T00:40:35.643171+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"58a52747","cell_type":"markdown","source":"## Stage D — confirm the winner end-to-end\n\nMap the selected taps back to the original `MultiLevelFusion` design and run the\nreal two-phase fine-tune with unfrozen backbones. CNN taps become\n`EFF_OUT_INDICES`, ViT taps become `VIT_HOOK_BLOCKS`, and the cross-attention\nbridge is kept only if `xattn` was selected. **This is the reported result.**","metadata":{"papermill":{"duration":0.029615,"end_time":"2026-08-13T00:48:10.540599+00:00","exception":false,"start_time":"2026-08-13T00:48:10.510984+00:00","status":"completed"},"tags":[]}},{"id":"14748800","cell_type":"code","source":"\n# --- translate winning taps into architecture config ---\nsel_cnn = sorted(int(t[3:]) for t in best_taps if t.startswith(\"cnn\"))\nsel_vit = sorted(int(t[3:]) for t in best_taps if t.startswith(\"vit\"))\nUSE_XATTN = \"xattn\" in best_taps\n# guardrails: xattn needs the deepest CNN + deepest ViT tap present; ensure >=1 each\nif not sel_cnn: sel_cnn = [max(CNN_STAGES)]\nif not sel_vit: sel_vit = [max(VIT_BLOCKS)]\n# safety net: never pass an out-of-range index to timm (guards against manual edits)\nsel_cnn = sorted(set(i for i in sel_cnn if i in CNN_STAGES))\nsel_vit = sorted(set(i for i in sel_vit if i in VIT_BLOCKS))\nprint(f\"Winning architecture -> CNN stages {tuple(sel_cnn)} | ViT blocks {tuple(sel_vit)} \"\n      f\"| cross-attention {USE_XATTN}\")\n\ncfg.EFF_OUT_INDICES = tuple(sel_cnn)\ncfg.VIT_HOOK_BLOCKS = tuple(sel_vit)\n","metadata":{"papermill":{"duration":0.039022,"end_time":"2026-08-13T00:48:10.607918+00:00","exception":false,"start_time":"2026-08-13T00:48:10.568896+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"1d7121fb","cell_type":"code","source":"\nclass SearchedFusion(nn.Module):\n    \"\"\"MultiLevelFusion, but with the searched tap set and optional cross-attention.\"\"\"\n    def __init__(self, use_xattn=True):\n        super().__init__()\n        self.use_xattn = use_xattn\n        self.eff = timm.create_model(cfg.EFF_NAME, pretrained=True, features_only=True,\n                                     out_indices=cfg.EFF_OUT_INDICES, drop_path_rate=cfg.EFF_DROP_PATH)\n        eff_ch = list(self.eff.feature_info.channels())\n        self.vit = timm.create_model(cfg.VIT_NAME, pretrained=True, num_classes=0,\n                                     drop_rate=cfg.VIT_DROP, drop_path_rate=cfg.VIT_DROP_PATH)\n        vit_d = self.vit.embed_dim; D = cfg.PROJ_DIM\n        for nm,mu,sd in ((\"eff\",EFF_MEAN,EFF_STD),(\"vit\",VIT_MEAN,VIT_STD)):\n            self.register_buffer(f\"{nm}_mean\", torch.tensor(mu).view(1,3,1,1))\n            self.register_buffer(f\"{nm}_std\",  torch.tensor(sd).view(1,3,1,1))\n        self.proj_c = nn.ModuleList([nn.Linear(c, D) for c in eff_ch])\n        self.proj_v = nn.ModuleList([nn.Linear(vit_d, D) for _ in cfg.VIT_HOOK_BLOCKS])\n        if use_xattn:\n            self.q_proj  = nn.Linear(eff_ch[-1], D)\n            self.kv_proj = nn.Linear(vit_d, D)\n            self.xattn   = nn.MultiheadAttention(D, cfg.XATTN_HEADS, batch_first=True)\n            self.xnorm   = nn.LayerNorm(D)\n        n_parts = len(eff_ch) + len(cfg.VIT_HOOK_BLOCKS) + (1 if use_xattn else 0)\n        fused = D * n_parts\n        self.head = nn.Sequential(\n            nn.Linear(fused,512), nn.BatchNorm1d(512), nn.Dropout(0.4), nn.Tanh(),\n            nn.Linear(512,128),   nn.BatchNorm1d(128), nn.Dropout(0.3), nn.Tanh(),\n            nn.Linear(128, cfg.NUM_CLASSES))\n        self.aux_c = nn.Linear(eff_ch[-1], cfg.NUM_CLASSES)\n        self.aux_v = nn.Linear(vit_d, cfg.NUM_CLASSES)\n        self._taps = {}\n        for i in cfg.VIT_HOOK_BLOCKS:\n            self.vit.blocks[i].register_forward_hook(self._mk(i))\n        self.fused_dim, self.eff_ch = fused, eff_ch\n    def _mk(self,i):\n        def h(_m,_in,out): self._taps[i]=out\n        return h\n    def fusion_modules(self):\n        mods=[self.proj_c,self.proj_v,self.head,self.aux_c,self.aux_v]\n        if self.use_xattn: mods += [self.q_proj,self.kv_proj,self.xattn,self.xnorm]\n        return mods\n    def fusion_parameters(self):\n        for m in self.fusion_modules(): yield from m.parameters()\n    def features(self, x):\n        maps = self.eff((x - self.eff_mean)/self.eff_std)\n        self._taps.clear()\n        cls_final = self.vit.forward_features((x - self.vit_mean)/self.vit_std)\n        if cls_final.ndim==3: cls_final = cls_final[:,0]\n        vcnn = [pj(m.mean((2,3))) for pj,m in zip(self.proj_c, maps)]\n        vvit = [pj(self._taps[i][:,0]) for pj,i in zip(self.proj_v, cfg.VIT_HOOK_BLOCKS)]\n        parts = vcnn + vvit\n        if self.use_xattn:\n            last = maps[-1]; B,C,H,W = last.shape\n            q  = self.q_proj(last.flatten(2).transpose(1,2))\n            kv = self.kv_proj(self._taps[cfg.VIT_HOOK_BLOCKS[-1]][:,1:])\n            if cfg.XATTN_FP32:\n                with amp_autocast(False):\n                    a,_ = self.xattn(q.float(),kv.float(),kv.float(),need_weights=False)\n            else:\n                a,_ = self.xattn(q,kv,kv,need_weights=False)\n            parts.append(self.xnorm(a.to(q.dtype).mean(1)))\n        return torch.cat(parts,1), maps[-1].mean((2,3)), cls_final\n    def forward(self, x):\n        fused,pc,pv = self.features(x)\n        return self.head(fused), self.aux_c(pc), self.aux_v(pv)\n\n# Rebind the training helpers (they call MultiLevelFusion by name) to the searched class.\nMultiLevelFusion = SearchedFusion\ndef _mk_model(): return SearchedFusion(use_xattn=USE_XATTN).to(device)\n# build_eval_copy in the original constructs MultiLevelFusion(); patch it to pass use_xattn\ndef build_eval_copy(m, ema):\n    if ema is None: return m\n    ev = SearchedFusion(use_xattn=USE_XATTN).to(device)\n    ev.load_state_dict({k: v.to(d) for (k,v),d in\n                        zip(ema.shadow.items(), [p.dtype for p in m.state_dict().values()])},\n                       strict=True)\n    return ev.eval()\n\nmodel = _mk_model()\nn = sum(p.numel() for p in model.parameters())\nprint(f\"SearchedFusion | {n/1e6:.1f}M params | CNN {cfg.EFF_OUT_INDICES} \"\n      f\"| ViT {cfg.VIT_HOOK_BLOCKS} | xattn {USE_XATTN}\")\n","metadata":{"papermill":{"duration":22.752363,"end_time":"2026-08-13T00:48:33.387903+00:00","exception":false,"start_time":"2026-08-13T00:48:10.63554+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"150c6785","cell_type":"code","source":"\n# ---- Phase 1: fusion modules only ----\ncfg.CKPT = f\"{cfg.OUT}/best_searched_c4.pt\"\nhistory, best = [], {\"score\":-1.0,\"qwk\":-1.0,\"acc\":-1.0}\nt0 = time.time()\nset_branches_trainable(model, False, False)\nema = ModelEMA(model, cfg.EMA_DECAY)\nopt_h = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad],\n                          lr=cfg.LR_HEAD, weight_decay=cfg.WEIGHT_DECAY)\nsch_h = torch.optim.lr_scheduler.CosineAnnealingLR(opt_h, T_max=cfg.EPOCHS_HEAD)\nprint(f\"===== PHASE 1: fusion head ({cfg.EPOCHS_HEAD} ep) =====\")\nhistory, best = run_phase(model, opt_h, sch_h, cfg.EPOCHS_HEAD, \"HEAD\", history, best, ema, t0)\n","metadata":{"papermill":{"duration":1243.393163,"end_time":"2026-08-13T01:09:16.808988+00:00","exception":false,"start_time":"2026-08-13T00:48:33.415825+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"5e7d6213","cell_type":"code","source":"\n# ---- Phase 2: fine-tune ----\n# The proxy (frozen-feature head) was *beating* the end-to-end run, i.e. unfreezing\n# the backbones overfits the small val set and destroys good features. So by default\n# we KEEP THE BACKBONES FROZEN in Phase 2 and only keep training the fusion modules\n# (more epochs, with augmentation) — this recovers the proxy-level score.\n# Set PHASE2_UNFREEZE=True to go back to partial-unfreeze fine-tuning.\nPHASE2_UNFREEZE = False\n\nif PHASE2_UNFREEZE:\n    ne,te = set_partial_trainable(model.eff, cfg.FREEZE_FRAC)\n    nv,tv = set_partial_trainable(model.vit, cfg.FREEZE_FRAC)\n    for p in model.fusion_parameters(): p.requires_grad = True\n    print(f\"Partial freeze @ {cfg.FREEZE_FRAC}: eff {ne}/{te} | vit {nv}/{tv}\")\n    opt_f = torch.optim.AdamW(make_ft_groups(model), lr=cfg.LR_FT_HEAD)\nelse:\n    set_branches_trainable(model, False, False)          # keep both backbones frozen\n    for p in model.fusion_parameters(): p.requires_grad = True\n    print(\"Phase 2: backbones FROZEN — training fusion modules only (proxy-safe).\")\n    opt_f = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad],\n                              lr=cfg.LR_HEAD, weight_decay=cfg.WEIGHT_DECAY)\n\n_w = torch.optim.lr_scheduler.LinearLR(opt_f, start_factor=0.1, total_iters=cfg.WARMUP_EPOCHS)\n_c = torch.optim.lr_scheduler.CosineAnnealingLR(opt_f, T_max=max(cfg.EPOCHS_FT-cfg.WARMUP_EPOCHS,1))\nsch_f = torch.optim.lr_scheduler.SequentialLR(opt_f, [_w,_c], milestones=[cfg.WARMUP_EPOCHS])\nprint(f\"===== PHASE 2: fine-tune ({cfg.EPOCHS_FT} ep) =====\")\nhistory, best = run_phase(model, opt_f, sch_f, cfg.EPOCHS_FT, \"FT\", history, best, ema, t0)\nprint(f\"\\nDone in {fmt(time.time()-t0)} | best blended={best['score']:.4f} \"\n      f\"(val qwk={best['qwk']*100:.2f}% acc={best['acc']*100:.2f}%)\")\n","metadata":{"papermill":{"duration":6834.195543,"end_time":"2026-08-13T03:03:11.036124+00:00","exception":false,"start_time":"2026-08-13T01:09:16.840581+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"f25f3df5","cell_type":"code","source":"\n# ---- Final test evaluation: searched fusion, plus a TTA probability-average ensemble ----\neval_model = SearchedFusion(use_xattn=USE_XATTN).to(device)\neval_model.load_state_dict(torch.load(cfg.CKPT, map_location=device), strict=True)\neval_model.eval()\nmet, yt, ypr = evaluate(eval_model, test_loader, \"test\", tta=cfg.USE_TTA)\nprint_metrics(met, \"COARSE TEST (4-class) — searched fusion, end-to-end + TTA\")\nprint(\"\\nConfusion matrix:\\n\", met[\"confusion_matrix\"])\n\n# ---- Fix D: late-fusion ensemble = average TTA softmax of fusion + ViT-alone + EFF-alone\n@torch.no_grad()\ndef backbone_probs(name, mean, std, ckpt, tta=cfg.USE_TTA):\n    m = Backbone(name, mean, std).to(device)\n    m.load_state_dict({k:v.to(device) for k,v in torch.load(ckpt, map_location=device).items()}, strict=False)\n    m.eval(); ps=[]\n    combos = [(False,False),(True,False),(False,True),(True,True)] if tta else [(False,False)]\n    for x,_ in tqdm(test_loader, desc=f\"[ens {name[:8]}]\", leave=False):\n        x=x.to(device); acc=None\n        for fh,fv in combos:\n            z=x\n            if fh: z=torch.flip(z,[3])\n            if fv: z=torch.flip(z,[2])\n            with amp_autocast(cfg.AMP):\n                p=F.softmax(m(z).float(),1)\n            acc=p if acc is None else acc+p\n        ps.append((acc/len(combos)).cpu().numpy())\n    return np.concatenate(ps)\n\np_fus = ypr                                                   # fusion TTA probs (from evaluate)\np_vit = backbone_probs(cfg.VIT_NAME, VIT_MEAN, VIT_STD, VIT_CKPT)\np_eff = backbone_probs(cfg.EFF_NAME, EFF_MEAN, EFF_STD, EFF_CKPT)\n\n# tune ensemble weights on VALIDATION only (never on test), then apply to test\n@torch.no_grad()\ndef val_probs_fusion():\n    _, yv, pv = evaluate(eval_model, val_loader, \"val\", tta=cfg.USE_TTA); return yv, pv\nyv, pv_fus = val_probs_fusion()\n@torch.no_grad()\ndef _val_bb(name,mean,std,ckpt,tta=cfg.USE_TTA):\n    # TTA here too: the weights tuned on val are applied to TTA'd test probs, so the\n    # two sides of the tuning must be produced the same way or the weights are biased.\n    m=Backbone(name,mean,std).to(device)\n    m.load_state_dict({k:v.to(device) for k,v in torch.load(ckpt,map_location=device).items()},strict=False)\n    m.eval(); ps=[]\n    combos = [(False,False),(True,False),(False,True),(True,True)] if tta else [(False,False)]\n    for x,_ in tqdm(val_loader, desc=f\"[val ens {name[:8]}]\", leave=False):\n        x=x.to(device); acc=None\n        for fh,fv in combos:\n            z=x\n            if fh: z=torch.flip(z,[3])\n            if fv: z=torch.flip(z,[2])\n            with amp_autocast(cfg.AMP):\n                p=F.softmax(m(z).float(),1)\n            acc=p if acc is None else acc+p\n        ps.append((acc/len(combos)).detach().cpu().numpy())\n    return np.concatenate(ps)\npv_vit=_val_bb(cfg.VIT_NAME,VIT_MEAN,VIT_STD,VIT_CKPT)\npv_eff=_val_bb(cfg.EFF_NAME,EFF_MEAN,EFF_STD,EFF_CKPT)\n\nbest_w, best_s = (1.,0.,0.), -1.0\nfor wf in np.linspace(0,1,11):\n    for wv in np.linspace(0,1-wf,int((1-wf)*10)+1):\n        we=1-wf-wv\n        pe=wf*pv_fus+wv*pv_vit+we*pv_eff\n        s=blended_score(compute_metrics(yv, pe.argmax(1), pe, cfg.NUM_CLASSES))\n        if s>best_s: best_s, best_w = s, (wf,wv,we)\nwf,wv,we = best_w\nprint(f\"\\nEnsemble weights (tuned on val): fusion={wf:.2f} vit={wv:.2f} eff={we:.2f}\")\np_ens = wf*p_fus + wv*p_vit + we*p_eff\nens = compute_metrics(yt, p_ens.argmax(1), p_ens, cfg.NUM_CLASSES)\nprint_metrics(ens, \"COARSE TEST (4-class) — weighted ensemble (fusion + ViT + EFF, TTA)\")\n","metadata":{"papermill":{"duration":322.343438,"end_time":"2026-08-13T03:08:33.415273+00:00","exception":false,"start_time":"2026-08-13T03:03:11.071835+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"c963eda8","cell_type":"code","source":"# ---- Class-wise accuracy (per-class recall = diagonal / row-sum of confusion matrix) ----\ndef print_class_wise_accuracy(met, class_names, title=\"\"):\n    if title: print(f\"--- {title} ---\")\n    cm = np.asarray(met[\"confusion_matrix\"])\n    support = cm.sum(1)\n    correct = cm.diagonal()\n    for name, c, s in zip(class_names, correct, support):\n        a = c / max(int(s), 1)\n        print(f\"  {name:12s}: {a*100:6.2f}%  ({int(c)}/{int(s)})\")\n    macro = float(np.mean(correct / np.maximum(support, 1)))\n    print(f\"  {'Macro avg':12s}: {macro*100:6.2f}%\")\n\nprint_class_wise_accuracy(met, cfg.CLASS_NAMES,\n                           \"COARSE (4-class) — searched fusion, TTA — class-wise accuracy\")\nprint()\nprint_class_wise_accuracy(ens, cfg.CLASS_NAMES,\n                           \"COARSE (4-class) — weighted ensemble, TTA — class-wise accuracy\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"a998e2a0","cell_type":"code","source":"\ncomp = pd.DataFrame([\n    {\"model\": \"EfficientNet-B4 (unfused)\", \"Acc%\":100*BASELINE[\"EfficientNet-B4 (alone)\"][\"accuracy\"],\n     \"QWK%\":100*BASELINE[\"EfficientNet-B4 (alone)\"][\"qwk\"], \"macroF1%\":100*BASELINE[\"EfficientNet-B4 (alone)\"][\"f1_macro\"]},\n    {\"model\": \"ViT-B/16 (unfused)\", \"Acc%\":100*BASELINE[\"ViT-B/16 (alone)\"][\"accuracy\"],\n     \"QWK%\":100*BASELINE[\"ViT-B/16 (alone)\"][\"qwk\"], \"macroF1%\":100*BASELINE[\"ViT-B/16 (alone)\"][\"f1_macro\"]},\n    {\"model\": \"searched fusion (proxy head)\", \"Acc%\":100*proxy[\"accuracy\"],\n     \"QWK%\":100*proxy[\"qwk\"], \"macroF1%\":100*proxy[\"f1_macro\"]},\n    {\"model\": \"searched fusion (end-to-end, TTA)\", \"Acc%\":100*met[\"accuracy\"],\n     \"QWK%\":100*met[\"qwk\"], \"macroF1%\":100*met[\"f1_macro\"]},\n    {\"model\": \"weighted ensemble (TTA)\", \"Acc%\":100*ens[\"accuracy\"],\n     \"QWK%\":100*ens[\"qwk\"], \"macroF1%\":100*ens[\"f1_macro\"]},\n]).set_index(\"model\").round(2)\n\n# headline = whichever of fusion / ensemble scores best on test accuracy\ncands = {\"searched fusion (end-to-end, TTA)\": met, \"weighted ensemble (TTA)\": ens}\nbest_name = max(cands, key=lambda k: cands[k][\"accuracy\"])\nbest_met  = cands[best_name]\n\n# plain-English description of exactly what got fused\ncnn_list = \", \".join(f\"stage {i}\" for i in sel_cnn) or \"none\"\nvit_list = \", \".join(f\"block {i}\" for i in sel_vit) or \"none\"\nbest_base = max(BASELINE.values(), key=lambda m: m[\"accuracy\"])\ndelta_acc = 100*(best_met[\"accuracy\"] - best_base[\"accuracy\"])\ndelta_qwk = 100*(best_met[\"qwk\"]      - best_base[\"qwk\"])\n\nprint(\"=\"*70)\nprint(\"WHAT WAS FUSED\")\nprint(\"=\"*70)\nprint(f\"  CNN (EfficientNet-B4) taps : {cnn_list}\")\nprint(f\"  ViT (ViT-B/16) taps        : {vit_list}\")\nprint(f\"  Cross-attention bridge     : {'YES' if USE_XATTN else 'no'}\")\nprint(f\"  Raw selected tap names      : {best_taps}\")\nprint()\nprint(\"COARSE RESULT (4-class) vs UNFUSED BASELINE (test set)\")\nprint(\"-\"*70)\ndisplay(comp)\nprint(f\"\\n>>> BEST COARSE MODEL: {best_name} — \"\n      f\"Acc {best_met['accuracy']*100:.2f}%, QWK {best_met['qwk']*100:.2f}%, \"\n      f\"macroF1 {best_met['f1_macro']*100:.2f}%\")\nprint(f\">>> vs best single backbone: Acc {delta_acc:+.2f} pts, QWK {delta_qwk:+.2f} pts.\")\nif delta_acc <= 0 and delta_qwk <= 0:\n    print(\"NOTE: neither fusion nor ensemble beat the best single backbone on test. \"\n          \"Strengthen the weaker backbone or widen TOPK/MAX_TAPS in the search.\")\n\ncomp.to_csv(f\"{cfg.OUT}/searched_fusion_result.csv\")\nnp.save(f\"{cfg.OUT}/coarse_test_probs.npy\", ypr)\nnp.save(f\"{cfg.OUT}/coarse_ensemble_test_probs.npy\", p_ens)\nnp.save(f\"{cfg.OUT}/coarse_test_labels.npy\", yt)\nprint(\"\\nSaved: searched_fusion_result.csv, coarse_test_probs.npy, \"\n      \"coarse_ensemble_test_probs.npy, coarse_test_labels.npy\")\n","metadata":{"papermill":{"duration":0.065576,"end_time":"2026-08-13T03:08:33.516831+00:00","exception":false,"start_time":"2026-08-13T03:08:33.451255+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"e9ff648c","cell_type":"markdown","source":"## Notes\n\n- **Report Stage D**, not the proxy. The greedy search on frozen features is only\n  a cheap ranking signal; the end-to-end confirmation number is the defensible one.\n- **Determinism.** Features are cached with `eval_tf` (no augmentation) so a subset's\n  score is reproducible. The head trains fast; bump `HEAD_EPOCHS` if scores look noisy.\n- **If fusion doesn't beat the best single tap**, raise `TOPK` (widen the candidate\n  pool) or `MAX_TAPS` (allow larger combos), and make sure both backbones are strong —\n  fusing a weak backbone with a strong one usually collapses to the strong one's tap.\n- **Cost.** Stage A (train both backbones) is the only expensive part and runs once;\n  everything after reads the cache. Delete `feat_cache/` or set `FORCE_RETRAIN=True`\n  to rebuild.\n- **Don't tune against the test number** — the search selects on validation only; test\n  is touched once, at the end.\n","metadata":{"papermill":{"duration":0.038229,"end_time":"2026-08-13T03:08:33.591891+00:00","exception":false,"start_time":"2026-08-13T03:08:33.553662+00:00","status":"completed"},"tags":[]}},{"id":"51a782f2-1a98-4c72-bb60-d89cfdb922bb","cell_type":"markdown","source":"---\n# Stage E — Severe vs PDR expert, and hierarchical recomposition\n\nEverything above produced `p_coarse`, a 4-class distribution whose last bucket is\n`Severe+PDR`. Stage E splits that bucket.\n\n**E.0** pick the coarse model (fusion vs ensemble) **on validation**.\n**E.1** build the grade-3/4 subsets; hold out a stratified slice of *train* for expert checkpoint selection.\n**E.2** fine-tune a dedicated binary EfficientNet-B4 expert (heavier augmentation, balanced sampler, EMA, TTA).\n**E.3** train a free MLP expert on the cached fusion taps as a decorrelated second opinion.\n**E.4** blend the two experts with a weight chosen on the expert hold-out.\n**E.5** recompose to 5 classes; `(mode, tau)` chosen on validation.\n**E.6** final 5-class test report.\n","metadata":{}},{"id":"d5541ebf-6f64-414a-a2f7-38c5417e5232","cell_type":"code","source":"# ===== E.0 — pick the coarse model on VALIDATION (never on test) =====\n# Everything from here on consumes exactly two arrays: coarse val probs and coarse\n# test probs. Which upstream model produced them is decided here, once.\n\nassert (yv == val_df[\"label\"].values).all(),  \"val loader order != val_df order\"\nassert (yt == test_df[\"label\"].values).all(), \"test loader order != test_df order\"\n\npv_ens = wf * pv_fus + wv * pv_vit + we * pv_eff        # same weights tuned in Stage D\n\nCOARSE_CANDIDATES = {\n    \"searched fusion (TTA)\":   (pv_fus, p_fus),\n    \"weighted ensemble (TTA)\": (pv_ens, p_ens),\n}\n_rows = []\nfor nm, (pvx, ptx) in COARSE_CANDIDATES.items():\n    m = compute_metrics(yv, pvx.argmax(1), pvx, cfg.NUM_CLASSES)\n    _rows.append({\"coarse model\": nm, \"val Acc%\": 100*m[\"accuracy\"],\n                  \"val QWK%\": 100*m[\"qwk\"], \"val macroF1%\": 100*m[\"f1_macro\"]})\n_sel = pd.DataFrame(_rows).set_index(\"coarse model\").round(2)\ndisplay(_sel)\n\nCOARSE_NAME = max(COARSE_CANDIDATES,\n                  key=lambda k: compute_metrics(yv, COARSE_CANDIDATES[k][0].argmax(1),\n                                                COARSE_CANDIDATES[k][0],\n                                                cfg.NUM_CLASSES)[\"accuracy\"])\npv_coarse, pt_coarse = COARSE_CANDIDATES[COARSE_NAME]\ncoarse_test_met = compute_metrics(yt, pt_coarse.argmax(1), pt_coarse, cfg.NUM_CLASSES)\n\n# fine-grained ground truth (the thing we are actually scored on)\nyv_fine = val_df[\"label_fine\"].values.astype(int)\nyt_fine = test_df[\"label_fine\"].values.astype(int)\n\nprint(f\"\\nCoarse model selected on val : {COARSE_NAME}\")\nprint(f\"Coarse TEST accuracy (4-class): {100*coarse_test_met['accuracy']:.2f}%\")\nprint(f\"  ^ this is the ORACLE CEILING for the 5-class number: a perfect Severe/PDR\")\nprint(f\"    expert reproduces it exactly, and any expert error subtracts from it.\")\n\n_hi_pred = (pt_coarse.argmax(1) == MERGED_LABEL)\nprint(f\"\\nTest rows routed to the expert : {int(_hi_pred.sum())} \"\n      f\"({100*_hi_pred.mean():.2f}% of test)\")\nprint(f\"  of those, truly grade 3/4     : {int(np.isin(yt_fine[_hi_pred], cfg.HI_GRADES).sum())}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"3f2a350d-efb0-4066-a934-0f66ebafdc25","cell_type":"code","source":"# ===== E.1 — grade-3/4 subsets + expert hold-out =====\n# `pos` keeps the row's ORIGINAL positional index inside train/val/test_df, which is\n# also its row index in the Stage-B feature caches (cached with shuffle=False).\n\ndef hi_subset(df):\n    idx = np.where(df[\"label_fine\"].isin(cfg.HI_GRADES).values)[0]\n    d = df.iloc[idx].copy().reset_index(drop=True)\n    d[\"pos\"] = idx\n    d[\"label\"] = (d[\"label_fine\"].astype(int) == 4).astype(int)   # 0 = Severe, 1 = PDR\n    return d\n\ntrain_hi = hi_subset(train_df)\nval_hi   = hi_subset(val_df)\ntest_hi  = hi_subset(test_df)\n\neh_tr, eh_va = train_test_split(train_hi, test_size=cfg.EXPERT_HOLDOUT,\n                                stratify=train_hi[\"label\"], random_state=cfg.SEED)\neh_tr = eh_tr.reset_index(drop=True); eh_va = eh_va.reset_index(drop=True)\n\ndef _bc(d, nm):\n    c = Counter(d[\"label\"]); s = dict(Counter(d[\"source\"]))\n    print(f\"  {nm:9s} n={len(d):4d}  Severe={c.get(0,0):4d}  PDR={c.get(1,0):4d}  sources={s}\")\n\nprint(\"===== EXPERT DATA (grade 3 vs grade 4) =====\")\n_bc(train_hi, \"train_hi\"); _bc(eh_tr, \"  eh_tr\"); _bc(eh_va, \"  eh_va\")\n_bc(val_hi, \"val_hi\"); _bc(test_hi, \"test_hi\")\n\nassert not (set(eh_tr[\"path\"]) & set(eh_va[\"path\"])), \"expert holdout leak\"\nassert not (set(train_hi[\"path\"]) & set(val_hi[\"path\"])), \"train/val leak\"\nassert not (set(train_hi[\"path\"]) & set(test_hi[\"path\"])), \"train/test leak\"\nprint(\"\\nExpert split assertions passed.\")\n\n# Heavier augmentation than the coarse net: this is a few hundred images, and the\n# Severe/PDR cue (neovascularisation, laser scars) is rotation- and flip-invariant.\nexpert_train_tf = A.Compose([\n    A.Resize(cfg.IMG_SIZE, cfg.IMG_SIZE),\n    A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.08, scale_limit=0.15, rotate_limit=180,\n                       border_mode=cv2.BORDER_CONSTANT, p=0.9),\n    A.RandomBrightnessContrast(0.20, 0.20, p=0.7),\n    A.Normalize(**_NN), ToTensorV2()])\n\n_pw_e = cfg.NUM_WORKERS > 0\neh_tr_ds = RetinaDS(eh_tr, expert_train_tf)\n_c = np.bincount(eh_tr[\"label\"].values, minlength=2)\n_w = (1.0 / np.maximum(_c, 1))[eh_tr[\"label\"].values]\neh_tr_loader = DataLoader(eh_tr_ds, cfg.BATCH_SIZE,\n                          sampler=WeightedRandomSampler(torch.as_tensor(_w, dtype=torch.double),\n                                                        len(eh_tr_ds), replacement=True),\n                          num_workers=cfg.NUM_WORKERS, pin_memory=True,\n                          drop_last=True, persistent_workers=_pw_e)\neh_va_loader = DataLoader(RetinaDS(eh_va, eval_tf), cfg.BATCH_SIZE, shuffle=False,\n                          num_workers=cfg.NUM_WORKERS, pin_memory=True)\nprint(f\"Expert batches: train {len(eh_tr_loader)} | holdout {len(eh_va_loader)} \"\n      f\"| balanced sampler on counts {_c.tolist()}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"b29d91d6-2a9c-4965-bc26-c0282da95bfd","cell_type":"code","source":"# ===== E.2 — fine-tune the binary Severe-vs-PDR expert =====\nEXPERT_CKPT = f\"{cfg.OUT}/expert_sev_pdr.pt\"\nFORCE_RETRAIN_EXPERT = True          # set False to reuse a cached expert\n\nEXP_MEAN, EXP_STD, EXP_IN = _stats(cfg.EXPERT_NAME, num_classes=0)\nprint(f\"Expert backbone {cfg.EXPERT_NAME} | mean={EXP_MEAN} | native={EXP_IN} \"\n      f\"| running at {cfg.IMG_SIZE}\")\n\nclass BinaryExpert(nn.Module):\n    \"\"\"Single backbone, 2-way head, branch normalisation baked into forward().\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.net = timm.create_model(cfg.EXPERT_NAME, pretrained=True, num_classes=2,\n                                     drop_rate=cfg.EXPERT_DROP,\n                                     drop_path_rate=cfg.EXPERT_DROPPATH)\n        self.register_buffer(\"mean\", torch.tensor(EXP_MEAN).view(1, 3, 1, 1))\n        self.register_buffer(\"std\",  torch.tensor(EXP_STD ).view(1, 3, 1, 1))\n    def forward(self, x):\n        return self.net((x - self.mean) / self.std)\n\n@torch.no_grad()\ndef expert_probs(model, loader, tta=True, desc=\"expert\"):\n    \"\"\"P(PDR) for every row of `loader`, flip-TTA averaged.\"\"\"\n    model.eval()\n    combos = [(False, False), (True, False), (False, True), (True, True)] if tta else [(False, False)]\n    out = []\n    for x, _ in tqdm(loader, desc=f\"[{desc}]\", leave=False):\n        x = x.to(device, non_blocking=True); acc = None\n        for fh, fv in combos:\n            z = x\n            if fh: z = torch.flip(z, [3])\n            if fv: z = torch.flip(z, [2])\n            with amp_autocast(cfg.AMP):\n                p = F.softmax(model(z).float(), 1)\n            acc = p if acc is None else acc + p\n        out.append((acc / len(combos))[:, 1].cpu().numpy())\n    return np.concatenate(out)\n\ndef _bin_report(y, q, thr=0.5):\n    pred = (q >= thr).astype(int)\n    acc = float((pred == y).mean())\n    bacc = float(np.nanmean([ (pred[y==c]==c).mean() if (y==c).any() else np.nan for c in (0,1) ]))\n    try: auc = float(roc_auc_score(y, q)) if len(set(y.tolist())) > 1 else float(\"nan\")\n    except Exception: auc = float(\"nan\")\n    return acc, bacc, auc\n\ndef train_expert():\n    if os.path.exists(EXPERT_CKPT) and not FORCE_RETRAIN_EXPERT:\n        print(f\"[expert] found {EXPERT_CKPT} — reusing\"); return\n    m = BinaryExpert().to(device)\n    opt = torch.optim.AdamW(m.parameters(), lr=cfg.EXPERT_LR, weight_decay=cfg.EXPERT_WD)\n    sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=cfg.EXPERT_EPOCHS)\n    scaler = make_scaler(cfg.AMP)\n    ema = ModelEMA(m, 0.999)                       # shorter horizon: few steps per epoch\n    # sampler already balances the classes -> keep focal alpha flat (no double correction)\n    floss = WeightedFocalLoss([1.0, 1.0], cfg.EXPERT_GAMMA).to(device)\n    yva = eh_va[\"label\"].values.astype(int)\n    best = -1.0\n    print(f\"===== EXPERT: {cfg.EXPERT_NAME} | {cfg.EXPERT_EPOCHS} ep | lr {cfg.EXPERT_LR} =====\")\n    t0 = time.time()\n    for ep in range(1, cfg.EXPERT_EPOCHS + 1):\n        m.train(); opt.zero_grad(set_to_none=True)\n        run = seen = corr = 0.0; nskip = 0\n        bar = tqdm(eh_tr_loader, desc=f\"EXPERT ep{ep}\", leave=False)\n        for step, (x, y) in enumerate(bar):\n            x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)\n            with amp_autocast(cfg.AMP):\n                lg = m(x); loss = floss(lg, y)\n            if not torch.isfinite(loss):\n                nskip += 1; opt.zero_grad(set_to_none=True); continue\n            scaler.scale(loss / cfg.ACCUM_STEPS).backward()\n            if (step + 1) % cfg.ACCUM_STEPS == 0 or (step + 1) == len(eh_tr_loader):\n                scaler.unscale_(opt)\n                gn = torch.nn.utils.clip_grad_norm_(m.parameters(), cfg.GRAD_CLIP)\n                if torch.isfinite(gn):\n                    scaler.step(opt); scaler.update(); ema.update(m)\n                else:\n                    nskip += 1; scaler.update()\n                opt.zero_grad(set_to_none=True)\n            bs = x.size(0); seen += bs; run += float(loss) * bs\n            corr += int((lg.argmax(1) == y).sum())\n            bar.set_postfix(loss=f\"{run/max(seen,1):.4f}\", acc=f\"{corr/max(seen,1):.4f}\", skip=nskip)\n        sch.step()\n        ev = BinaryExpert().to(device); ema.copy_to(ev)\n        q = expert_probs(ev, eh_va_loader, tta=False, desc=\"eh_va\")\n        acc, bacc, auc = _bin_report(yva, q)\n        sc = 0.5 * bacc + 0.5 * (0.0 if np.isnan(auc) else auc)\n        star = \"\"\n        if sc > best:\n            best = sc\n            torch.save({k: v.cpu().clone() for k, v in ema.shadow.items()}, EXPERT_CKPT)\n            star = \"  <-- saved\"\n        print(f\"[EXPERT {ep:02d}/{cfg.EXPERT_EPOCHS}] loss={run/max(seen,1):.4f} \"\n              f\"tr_acc={corr/max(seen,1):.4f} | ho_acc={acc*100:.2f}% bacc={bacc*100:.2f}% \"\n              f\"auc={auc*100:.2f}% | score={sc:.4f} ({fmt(time.time()-t0)}){star}\")\n        del ev\n    print(f\"[expert] done in {fmt(time.time()-t0)} | best holdout score {best:.4f}\")\n\ntrain_expert()\n\nexpert = BinaryExpert().to(device)\nexpert.load_state_dict({k: v.to(device) for k, v in\n                        torch.load(EXPERT_CKPT, map_location=device).items()}, strict=True)\nexpert.eval()\nprint(\"Expert loaded.\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"7ace81c6-9294-4682-9387-5e451941a980","cell_type":"code","source":"# ===== E.2b — expert predictions on the FULL val / test sets =====\n# The expert must score every row, not just the true 3/4 rows: at inference we do\n# not know the grade, we only know which rows the coarse model routed to the bucket.\n\nq_val_cnn  = expert_probs(expert, val_loader,  tta=cfg.EXPERT_TTA, desc=\"expert val\")\nq_test_cnn = expert_probs(expert, test_loader, tta=cfg.EXPERT_TTA, desc=\"expert test\")\nq_ehva_cnn = expert_probs(expert, eh_va_loader, tta=cfg.EXPERT_TTA, desc=\"expert eh_va\")\n\n# standalone diagnostics on the true grade-3/4 rows only\nfor nm, q, d in ((\"VAL \", q_val_cnn, val_hi), (\"TEST\", q_test_cnn, test_hi)):\n    a, b, u = _bin_report(d[\"label\"].values.astype(int), q[d[\"pos\"].values])\n    print(f\"[CNN expert | true 3/4 rows only] {nm}: acc={a*100:.2f}%  \"\n          f\"balanced-acc={b*100:.2f}%  auc={u*100:.2f}%  (n={len(d)})\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"d4f1d01b-ef38-4bb3-855c-d8eebabf488c","cell_type":"code","source":"# ===== E.3 — free second expert: MLP on the cached fusion taps =====\n# The CNN expert was fine-tuned only on 3/4 images; this one reads features from\n# backbones trained on the WHOLE 4-class problem. Different training signal ->\n# decorrelated errors, which is where ensemble gain actually comes from.\n\ndef make_bin_head(in_dim):\n    return nn.Sequential(\n        nn.Linear(in_dim, 256), nn.LayerNorm(256), nn.Dropout(0.4), nn.Tanh(),\n        nn.Linear(256, 64),     nn.LayerNorm(64),  nn.Dropout(0.3), nn.Tanh(),\n        nn.Linear(64, 2)).to(device)\n\ndef train_cheap_expert(taps, seeds=(0, 1, 2, 3, 4), epochs=80, lr=2e-3, bs=128):\n    \"\"\"Returns (q_ehva, q_val, q_test) averaged over seeds. Costs seconds.\"\"\"\n    Xtr_all = _cat(cache_train, taps).float()\n    Xva_all = _cat(cache_val,   taps).float().to(device)\n    Xte_all = _cat(cache_test,  taps).float().to(device)\n    Xtr = Xtr_all[eh_tr[\"pos\"].values].to(device)\n    ytr = torch.tensor(eh_tr[\"label\"].values, dtype=torch.long, device=device)\n    Xho = Xtr_all[eh_va[\"pos\"].values].to(device)\n    yho = eh_va[\"label\"].values.astype(int)\n    cnt = np.bincount(eh_tr[\"label\"].values, minlength=2)\n    w = torch.tensor((1.0 / np.maximum(cnt, 1)), dtype=torch.float, device=device)\n    w = w / w.mean()\n    acc_ho = acc_va = acc_te = None\n    for sd in seeds:\n        torch.manual_seed(sd)\n        head = make_bin_head(Xtr.shape[1])\n        opt = torch.optim.AdamW(head.parameters(), lr=lr, weight_decay=cfg.WEIGHT_DECAY)\n        sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n        best_sc, best_state = -1.0, None\n        n = Xtr.shape[0]\n        for ep in range(epochs):\n            head.train(); perm = torch.randperm(n, device=device)\n            for i in range(0, n, bs):\n                idx = perm[i:i+bs]\n                if idx.numel() < 2: continue\n                loss = F.cross_entropy(head(Xtr[idx]), ytr[idx], weight=w, label_smoothing=0.05)\n                opt.zero_grad(set_to_none=True); loss.backward(); opt.step()\n            sch.step()\n            head.eval()\n            with torch.no_grad():\n                qh = F.softmax(head(Xho).float(), 1)[:, 1].cpu().numpy()\n            _, b, u = _bin_report(yho, qh)\n            sc = 0.5 * b + 0.5 * (0.0 if np.isnan(u) else u)\n            if sc > best_sc:\n                best_sc = sc\n                best_state = {k: v.detach().clone() for k, v in head.state_dict().items()}\n        head.load_state_dict(best_state); head.eval()\n        with torch.no_grad():\n            qh = F.softmax(head(Xho).float(), 1)[:, 1].cpu().numpy()\n            qv = F.softmax(head(Xva_all).float(), 1)[:, 1].cpu().numpy()\n            qt = F.softmax(head(Xte_all).float(), 1)[:, 1].cpu().numpy()\n        acc_ho = qh if acc_ho is None else acc_ho + qh\n        acc_va = qv if acc_va is None else acc_va + qv\n        acc_te = qt if acc_te is None else acc_te + qt\n    k = len(seeds)\n    return acc_ho / k, acc_va / k, acc_te / k\n\nif cfg.USE_CHEAP_EXPERT:\n    q_ehva_mlp, q_val_mlp, q_test_mlp = train_cheap_expert(best_taps)\n    a, b, u = _bin_report(eh_va[\"label\"].values.astype(int), q_ehva_mlp)\n    print(f\"[MLP expert] holdout: acc={a*100:.2f}%  bacc={b*100:.2f}%  auc={u*100:.2f}%\")\n    for nm, q, d in ((\"VAL \", q_val_mlp, val_hi), (\"TEST\", q_test_mlp, test_hi)):\n        a, b, u = _bin_report(d[\"label\"].values.astype(int), q[d[\"pos\"].values])\n        print(f\"[MLP expert | true 3/4 rows only] {nm}: acc={a*100:.2f}%  \"\n              f\"bacc={b*100:.2f}%  auc={u*100:.2f}%\")\nelse:\n    q_ehva_mlp = q_val_mlp = q_test_mlp = None\n    print(\"Cheap expert disabled (cfg.USE_CHEAP_EXPERT=False).\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"dee4c148-09d7-435d-b88e-ce3480cf90e5","cell_type":"code","source":"# ===== E.4 — blend the two experts (weight chosen on the expert HOLD-OUT) =====\n# The hold-out is a slice of TRAIN, so this spends no validation signal — val is\n# reserved entirely for the (mode, tau) decision in E.5.\n\nif q_ehva_mlp is None:\n    W_EXPERT = 1.0\n    print(\"Single expert: w_cnn = 1.00\")\nelse:\n    yho = eh_va[\"label\"].values.astype(int)\n    rows = []\n    for wc in np.linspace(0, 1, 21):\n        qh = wc * q_ehva_cnn + (1 - wc) * q_ehva_mlp\n        a, b, u = _bin_report(yho, qh)\n        rows.append({\"w_cnn\": wc, \"acc\": a, \"bacc\": b, \"auc\": u,\n                     \"score\": 0.5*b + 0.5*(0.0 if np.isnan(u) else u)})\n    blend_hist = pd.DataFrame(rows)\n    W_EXPERT = float(blend_hist.sort_values(\"score\", ascending=False).iloc[0][\"w_cnn\"])\n    display(blend_hist.round(4).sort_values(\"score\", ascending=False).head(6))\n    print(f\"Expert blend chosen on holdout: w_cnn={W_EXPERT:.2f}  w_mlp={1-W_EXPERT:.2f}\")\n\ndef _blend(qc, qm):\n    return qc if qm is None else W_EXPERT * qc + (1 - W_EXPERT) * qm\n\nq_val  = _blend(q_val_cnn,  q_val_mlp)\nq_test = _blend(q_test_cnn, q_test_mlp)\n\nfor nm, q, d in ((\"VAL \", q_val, val_hi), (\"TEST\", q_test, test_hi)):\n    a, b, u = _bin_report(d[\"label\"].values.astype(int), q[d[\"pos\"].values])\n    print(f\"[BLENDED expert | true 3/4 rows] {nm}: acc={a*100:.2f}%  \"\n          f\"bacc={b*100:.2f}%  auc={u*100:.2f}%  (n={len(d)})\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"4c92dae2-10c1-47a2-a671-b2b0aeaf6d4c","cell_type":"code","source":"# ===== E.5 — recompose 4-class + expert -> 5-class; tune (mode, tau) on VAL =====\n#   hard : keep the coarse argmax, split only the rows it sent to the merged bucket.\n#          Coarse decisions for classes 0-2 are untouched -> acc5 <= acc4 exactly.\n#   prob : full factorisation P(3)=P(merged)(1-q), P(4)=P(merged)q, then argmax over\n#          all five. Can also MOVE a row out of the bucket when the split halves both\n#          drop below P(Moderate) -- occasionally recovers Moderate rows, sometimes\n#          loses Severe ones. Which wins is data-dependent, so val decides.\n# tau is applied as a logit shift so it means the same thing in both modes.\n\ndef _logit(p, eps=1e-6):\n    p = np.clip(np.asarray(p, dtype=np.float64), eps, 1 - eps)\n    return np.log(p / (1 - p))\n\ndef _sig(z): return 1.0 / (1.0 + np.exp(-z))\n\ndef expand5(coarse_p, q, mode=\"hard\", tau=0.5):\n    coarse_p = np.asarray(coarse_p, dtype=np.float64)\n    qs = _sig(_logit(q) - _logit(tau))                     # tau -> 0.5 after the shift\n    P = np.zeros((coarse_p.shape[0], cfg.FINE_NUM_CLASSES), dtype=np.float64)\n    P[:, :MERGED_LABEL] = coarse_p[:, :MERGED_LABEL]\n    P[:, 3] = coarse_p[:, MERGED_LABEL] * (1.0 - qs)\n    P[:, 4] = coarse_p[:, MERGED_LABEL] * qs\n    if mode == \"prob\":\n        pred = P.argmax(1)\n    else:\n        pred = coarse_p.argmax(1).astype(int).copy()\n        m = pred == MERGED_LABEL\n        pred[m] = np.where(qs[m] >= 0.5, 4, 3)\n    return pred.astype(int), P\n\nTAUS = np.round(np.linspace(0.05, 0.95, 37), 4)\nrows = []\nfor mode in (\"hard\", \"prob\"):\n    for tau in TAUS:\n        pr, P = expand5(pv_coarse, q_val, mode, tau)\n        m = compute_metrics(yv_fine, pr, P, cfg.FINE_NUM_CLASSES)\n        rows.append({\"mode\": mode, \"tau\": float(tau), \"val Acc\": m[\"accuracy\"],\n                     \"val QWK\": m[\"qwk\"], \"val macroF1\": m[\"f1_macro\"]})\ntune = pd.DataFrame(rows)\n# primary criterion = 5-class VAL ACCURACY (the BTP success metric); QWK breaks ties\ntune = tune.sort_values([\"val Acc\", \"val QWK\"], ascending=False).reset_index(drop=True)\nBEST_MODE = tune.iloc[0][\"mode\"]; BEST_TAU = float(tune.iloc[0][\"tau\"])\ndisplay(tune.head(10).round(4))\nprint(f\"\\nSelected on VAL -> mode={BEST_MODE}  tau={BEST_TAU:.3f}  \"\n      f\"(val acc5 {100*tune.iloc[0]['val Acc']:.2f}%, qwk {100*tune.iloc[0]['val QWK']:.2f}%)\")\n\n# tau sensitivity, so a knife-edge optimum is visible rather than trusted\n_h = tune[tune[\"mode\"] == BEST_MODE].sort_values(\"tau\")\nplt.figure(figsize=(7, 3.2))\nplt.plot(_h[\"tau\"], 100*_h[\"val Acc\"], marker=\"o\", ms=3, label=\"val acc5 %\")\nplt.plot(_h[\"tau\"], 100*_h[\"val QWK\"], marker=\"s\", ms=3, label=\"val QWK %\")\nplt.axvline(BEST_TAU, ls=\"--\", c=\"k\", lw=1, label=f\"tau*={BEST_TAU:.2f}\")\nplt.xlabel(\"tau  (decision threshold on P(PDR))\"); plt.ylabel(\"%\")\nplt.title(f\"Threshold sensitivity — mode={BEST_MODE}\"); plt.legend(fontsize=8)\nplt.grid(alpha=0.3); plt.tight_layout(); plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"b101af59-9b85-46f6-9e45-48a7bc6bee38","cell_type":"code","source":"# ===== E.6 — FINAL 5-class test report =====\npred5_test, P5_test = expand5(pt_coarse, q_test, BEST_MODE, BEST_TAU)\nfinal = compute_metrics(yt_fine, pred5_test, P5_test, cfg.FINE_NUM_CLASSES)\n\nprint(\"=\" * 74)\nprint(f\"FINAL — hierarchical 5-class  ({COARSE_NAME}  +  Severe/PDR expert)\")\nprint(\"=\" * 74)\nprint_metrics(final, \"TEST (5-class, the reported number)\")\nprint(\"\\nConfusion matrix (rows = true 0..4):\\n\", final[\"confusion_matrix\"])\nprint()\nprint_class_wise_accuracy(final, cfg.FINE_CLASS_NAMES, \"TEST (5-class) — class-wise accuracy\")\n\n# ---- where the accuracy went: coarse ceiling -> expert cost -> final ----\ncoarse_acc = coarse_test_met[\"accuracy\"]\ncoarse_from_final = np.where(np.isin(pred5_test, (3, 4)), MERGED_LABEL, pred5_test)\ncoarse_ok = coarse_from_final == yt\n\nhi_true = np.isin(yt_fine, cfg.HI_GRADES)\nhi_routed_ok = coarse_ok & hi_true                       # correctly reached the bucket\nexpert_calls = int(hi_routed_ok.sum())\nexpert_right = int((pred5_test[hi_routed_ok] == yt_fine[hi_routed_ok]).sum())\ncost = (expert_calls - expert_right) / len(yt_fine)\n\nprint(\"\\n\" + \"-\" * 74)\nprint(\"ACCURACY DECOMPOSITION (test)\")\nprint(\"-\" * 74)\nprint(f\"  coarse 4-class accuracy (oracle ceiling) : {100*coarse_acc:6.2f}%\")\nprint(f\"  rows correctly routed to the bucket      : {expert_calls}\")\nprint(f\"  of those, expert got the grade right     : {expert_right}\"\n      f\"  ({100*expert_right/max(expert_calls,1):.2f}%)\")\nprint(f\"  accuracy paid to expert error            : {100*cost:6.2f} pts\")\nprint(f\"  final 5-class accuracy                   : {100*final['accuracy']:6.2f}%\")\n\nTARGET = 0.90\nprint(\"\\n\" + \"=\" * 74)\nif final[\"accuracy\"] >= TARGET:\n    print(f\">>> TARGET CLEARED: {100*final['accuracy']:.2f}% >= {100*TARGET:.0f}%  \"\n          f\"(QWK {100*final['qwk']:.2f}%)\")\nelse:\n    gap = TARGET - final[\"accuracy\"]\n    need_coarse = TARGET + cost\n    print(f\">>> SHORT OF TARGET by {100*gap:.2f} pts  \"\n          f\"(got {100*final['accuracy']:.2f}%, QWK {100*final['qwk']:.2f}%)\")\n    print(f\"    Diagnosis — which lever to pull next:\")\n    if coarse_acc < TARGET:\n        print(f\"      * COARSE is binding: 4-class acc {100*coarse_acc:.2f}% is already below \"\n              f\"{100*TARGET:.0f}%, so no expert can reach the target.\")\n        print(f\"        -> widen TOPK/MAX_TAPS, add a decorrelated third backbone to the \"\n              f\"ensemble, or train the backbones longer. The expert is not the problem.\")\n    else:\n        print(f\"      * EXPERT is binding: coarse ceiling is {100*coarse_acc:.2f}% but expert \"\n              f\"error costs {100*cost:.2f} pts.\")\n        print(f\"        -> the coarse model would need {100*need_coarse:.2f}% to absorb the \"\n              f\"current expert error; alternatively raise expert accuracy from \"\n              f\"{100*expert_right/max(expert_calls,1):.1f}% to \"\n              f\"{100*(1 - (coarse_acc-TARGET)*len(yt_fine)/max(expert_calls,1)):.1f}%.\")\nprint(\"=\" * 74)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"98d742e1-cf82-4fca-8791-84a1d2ebf7fe","cell_type":"code","source":"# ===== E.7 — summary table + artifacts =====\n# A flat 5-class reference from the same coarse probabilities is impossible (the\n# coarse net never had a 5th logit), so the honest comparisons are: coarse ceiling,\n# hierarchical-with-tau=0.5 (untuned), and the tuned hierarchical result.\n\npred_untuned, P_untuned = expand5(pt_coarse, q_test, \"hard\", 0.5)\nuntuned = compute_metrics(yt_fine, pred_untuned, P_untuned, cfg.FINE_NUM_CLASSES)\n\nsummary = pd.DataFrame([\n    {\"model\": f\"COARSE 4-class ({COARSE_NAME})  [oracle ceiling]\",\n     \"Acc%\": 100*coarse_test_met[\"accuracy\"], \"QWK%\": 100*coarse_test_met[\"qwk\"],\n     \"macroF1%\": 100*coarse_test_met[\"f1_macro\"]},\n    {\"model\": \"hierarchical 5-class, tau=0.5 (untuned)\",\n     \"Acc%\": 100*untuned[\"accuracy\"], \"QWK%\": 100*untuned[\"qwk\"],\n     \"macroF1%\": 100*untuned[\"f1_macro\"]},\n    {\"model\": f\"hierarchical 5-class, mode={BEST_MODE} tau={BEST_TAU:.2f}  [REPORTED]\",\n     \"Acc%\": 100*final[\"accuracy\"], \"QWK%\": 100*final[\"qwk\"],\n     \"macroF1%\": 100*final[\"f1_macro\"]},\n]).set_index(\"model\").round(2)\ndisplay(summary)\n\nsummary.to_csv(f\"{cfg.OUT}/round10_hierarchical_result.csv\")\ntune.to_csv(f\"{cfg.OUT}/round10_tau_tuning.csv\", index=False)\nnp.save(f\"{cfg.OUT}/round10_test_probs5.npy\", P5_test)\nnp.save(f\"{cfg.OUT}/round10_test_pred5.npy\", pred5_test)\nnp.save(f\"{cfg.OUT}/round10_test_labels5.npy\", yt_fine)\nnp.save(f\"{cfg.OUT}/round10_test_q_pdr.npy\", q_test)\n\ncmf = final[\"confusion_matrix\"]\nplt.figure(figsize=(5.4, 4.4))\nsns.heatmap(cmf, annot=True, fmt=\"d\", cmap=\"Blues\", cbar=False,\n            xticklabels=cfg.FINE_CLASS_NAMES, yticklabels=cfg.FINE_CLASS_NAMES)\nplt.xlabel(\"predicted\"); plt.ylabel(\"true\")\nplt.title(f\"Round 10 — hierarchical 5-class\\nAcc {100*final['accuracy']:.2f}%  \"\n          f\"QWK {100*final['qwk']:.2f}%\")\nplt.tight_layout(); plt.show()\n\nprint(\"Saved: round10_hierarchical_result.csv, round10_tau_tuning.csv, \"\n      \"round10_test_probs5.npy, round10_test_pred5.npy, round10_test_labels5.npy, \"\n      \"round10_test_q_pdr.npy\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"a9fb59d8-9547-4f28-83cb-a6aca8df1913","cell_type":"markdown","source":"## Notes on Round 10\n\n- **`acc4` is a hard ceiling.** In `hard` mode nothing can push the 5-class number\n  above the coarse 4-class number — the split only ever converts a correct coarse\n  prediction into a correct-or-wrong fine one. If E.6 reports the coarse model as\n  the binding constraint, tuning the expert further is wasted effort.\n- **Why the merge should raise `acc4` in the first place.** Every 3<->4 cell in the\n  old 5-class confusion matrix collapses onto the diagonal, and focal alpha no\n  longer has to trade Severe against PDR — which is the seesaw that kept costing\n  Moderate recall in rounds 6-7. Read the coarse class-wise output in the Stage D\n  cell and compare Moderate against the Round 5-7 numbers: that delta is the merge\n  paying off, separately from anything the expert does.\n- **Two experts, one blend.** The CNN expert only ever sees grade-3/4 images; the\n  MLP expert reads features from backbones trained on all four classes. They fail\n  on different images, which is the whole reason for blending — same principle that\n  made cross-family ensembling work in earlier rounds. The blend weight comes from\n  a slice of *train*, so it costs no validation signal.\n- **`tau` on ~35 validation rows is noisy.** The sensitivity plot in E.5 exists for\n  exactly this reason: if the curve is a spike rather than a plateau, prefer\n  `tau=0.5` and treat the tuned number as optimistic. A flat top means the choice\n  is real.\n- **Compute.** Stage E adds roughly 20-30 min on a P100 (a few hundred images x 16\n  epochs at 384, plus three TTA inference passes). If the session is tight, cut\n  `BACKBONE_CFG[\"EFF\"][\"epochs\"]` before cutting expert epochs — the expert is the\n  cheap half of this notebook.\n- **If the expert underperforms**, the most likely cause is too few grade-3/4\n  training images. Check the E.1 source counts: if IDRiD/Messidor contributed\n  little, `cfg.MERGE_GRADES` and the Messidor detection in section 2.5 are worth\n  revisiting before touching architecture.\n","metadata":{}}]}