{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"accelerator":"GPU"},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"ece531c7-2188-4ddc-9630-2c7fb300ccf2","cell_type":"markdown","source":"# Diabetic Retinopathy — Binary Detection (No DR vs DR)\n## Custom from-scratch dual-branch network (NO transfer learning)\n\n**Reference baseline:** Shakibania et al., *Dual branch deep learning network for detection and\nstage grading of diabetic retinopathy*, Biomed. Signal Process. Control 93 (2024) 106168.\nTheir binary result: **98.50% accuracy / 99.46% sensitivity / 97.51% specificity**, obtained with\n**ImageNet-pretrained** ResNet50 + EfficientNetB0 branches.\n\n**This notebook deviates deliberately in one way:** every weight is randomly initialised and learned\nfrom the fundus data only. No `pretrained=True`, no `timm`, no downloaded checkpoints. The\narchitecture is written from scratch in this notebook.\n\n**What is kept from the paper (so the comparison is fair):**\n* Same data sources: APTOS 2019 (primary) + IDRiD + Messidor-2.\n* Same protocol: APTOS split 70 / 10 / 20 stratified on the 5 ICDR grades; test set is APTOS-only.\n* Same \"selective merge\" idea: external data added to **train only**.\n* Same metric set (accuracy, precision, sensitivity, specificity, F1, AUC, kappa) plus extras.\n\n**What is added to make a scratch model competitive:**\n* Fundus-specific preprocessing (border crop, circular mask, Ben Graham illumination correction, CLAHE).\n* A two-branch design where one branch is deliberately kept high-resolution / dilated to preserve\n  microaneurysm-scale evidence, and a cross-attention bridge between the branches.\n* Deep supervision, SE blocks, GeM pooling, stochastic depth, mixup, EMA weights, flip TTA.\n* Validation-tuned decision threshold (never tuned on test).\n\n**Output:** for every experiment you get val + test metrics for all four inference variants\n(raw / EMA weights x TTA off / on), confusion matrices, ROC + PR curves, per-ICDR-grade error\nanalysis, and an accumulating `results.csv` so runs from several Kaggle accounts can be merged.","metadata":{}},{"id":"3ace52ee-6f2f-46dd-8cee-605c3e110e92","cell_type":"markdown","source":"## 1. Configuration\n\nEdit `RUN_EXPERIMENTS` to choose what runs this session. One experiment ≈ 50-70 min on a P100\nat 384 px / 45 epochs. Set `QUICK_TEST = True` first to verify the whole pipeline in ~5 minutes.","metadata":{}},{"id":"5bd344bb-f1ae-4993-b9a0-f0585cf3fece","cell_type":"code","source":"import os, sys, json, math, time, random, glob, hashlib, warnings, itertools, textwrap\nwarnings.filterwarnings(\"ignore\")\n\n# ---------------------------------------------------------------- global knobs\nSEED            = 42\nIMG_SIZE        = 384          # 384 or 448; 448 helps Mild but needs a smaller batch\nCACHE_SIZE      = 512          # images cached at this size (crop+resize done once)\nBATCH_SIZE      = 20 if IMG_SIZE <= 384 else 12\nGRAD_ACCUM      = 1            # effective batch = BATCH_SIZE * GRAD_ACCUM\n# --- ONE plain cosine schedule. Warm restarts were tried and dropped: the snapshot\n#     ensemble scored 96.73 against 97.14 for a plain ema+tta, because the early-cycle\n#     snapshot (val AUC 98.07) dragged the average down. Simpler is better here.\nSCHEDULE_CYCLES = [100]        # best epoch has landed at 82-96 every time, so 100 is ample\nCYCLE_LR_SCALE  = [1.0]\nEPOCHS          = sum(SCHEDULE_CYCLES)\n\n# --- progressive resizing: OFF. It never fired in practice (the time estimate refused it),\n#     and an untested feature is worse than no feature. To test 448 properly, set\n#     IMG_SIZE = 448 at the top and run it as its own experiment.\nPROGRESSIVE_SIZE  = None\nPROGRESSIVE_AT    = 10**9\nPROGRESSIVE_BATCH = 12\nTTA_MODE          = \"flip5\"    # \"flip\" (4 views) | \"flip5\" (4 flips + centre zoom)\n\n# --- variant selection. Validation AUC alone picked 'raw' (val AUC 99.09) over 'ema+tta'\n#     (98.68) in run 1, yet on test ema+tta scored 97.14 vs 96.86. AUC is threshold-free and\n#     does not reflect the metric actually reported, so selection now blends val accuracy\n#     with val AUC, and also considers an average of the top-3 variants.\n# --- AUXILIARY 5-CLASS HEAD. The binary label throws away the ICDR grade you already have.\n#     \"moderate\" and \"PDR\" are both just DR=1 to the binary loss, so the network never learns\n#     that they are different, even though the ordering is exactly what separates a borderline\n#     Mild case from a healthy eye. A second head predicting the 5 grades gives a richer\n#     training signal at almost no parameter cost, and at inference 1 - P(grade 0) is a SECOND\n#     opinion on the same question that can be averaged with the binary head for free.\nGRADE_AUX_W  = 0.4             # weight of the grade loss; 0 disables the head\nGRADE_FUSE   = True            # also score a blend of the binary head and the grade head\nGRADE_MODE   = \"softmax\"       # \"softmax\" | \"ordinal\"\n#   REVERTED to softmax on evidence. Ordinal encoding is the better idea in principle -\n#   the ICDR grades are ordered - but measured on this data it LOST: the identical\n#   384px/ben config scored 97.41 +/- 0.14 with softmax and 96.86 with ordinal. Theory\n#   loses to measurement. \"ordinal\" is kept as an option, not a default.\n#   ordinal: 4 sigmoid outputs for P(grade>0), P(grade>1), P(grade>2), P(grade>3). This\n#   respects the fact that the ICDR grades are ORDERED - softmax treats Mild and PDR as\n#   equally unrelated to No-DR, which is plainly false. It also means the FIRST output is\n#   literally \"is there any DR\", i.e. a second, ordering-aware estimate of the exact\n#   quantity being reported, rather than something derived from a 5-way softmax.\nFUSE_WEIGHT  = 0.5             # fixed. \"auto\" sweeps on validation, but the sweep overfits\n#   a 366-image val set - it beat a fixed 0.5 in only 35% of simulated trials, and the runs\n#   that used it scored below the fixed-0.5 runs.\n\n# --- K-FOLD CROSS-VALIDATION ENSEMBLE.\n#     Every run so far trained on the SAME 70% of APTOS and validated on the SAME 10%, so\n#     three seeds saw identical data and made near-identical mistakes - which is exactly why\n#     the seed ensembles kept failing. With K folds each model trains on a DIFFERENT subset,\n#     so the members genuinely disagree, and the ensemble covers 100% of the non-test data\n#     instead of 70%. This is the standard route to the last point of accuracy and the one\n#     structural lever never tried here. Test set stays fixed and untouched.\nN_FOLDS = 5\n\n# --- SEMI-SUPERVISED PSEUDO-LABELLING ---------------------------------------------------\n# The APTOS competition ships 3,662 LABELLED training images and 1,928 UNLABELLED test\n# images. Everything so far used only the labelled half. Those 1,928 come from the same\n# hospital and cameras as your evaluation set, so they are the single best source of extra\n# in-distribution data available - a 63% increase in training images.\n#\n# The method: train a teacher, have it label the unlabelled pool, keep ONLY the predictions\n# it is very confident about, and train a student on labelled + confident-pseudo data. With\n# a teacher at 99.4 ROC-AUC, predictions beyond +/-0.97 are right almost every time, so the\n# added noise is small relative to the added data. This is the standard route to the last\n# point of accuracy in exactly this setting, and it is legitimate here: no external weights,\n# no external labels, and the pseudo-labelled images are the competition's own test split -\n# entirely disjoint from the 733 images this project evaluates on.\nPSEUDO_ENABLED = True\nPSEUDO_CONF    = 0.97          # keep only p >= 0.97 (DR) or p <= 0.03 (no DR)\nPSEUDO_TEACHER = \"K0_fold0_ema.pt\"      # checkpoint name; searched for automatically\n\nSELECT_MODE = \"blend\"          # \"blend\" (acc+AUC) | \"val_auc\" | \"val_acc\"\nUSE_TOP3_AVG = True            # add a top-3-by-val-score probability average as a candidate\nBASE_LR         = 3e-4\nWARMUP_EPOCHS   = 5\n\n# --- regularisation, dialled DOWN because the model underfits (train 92.8 / val 92.9).\n#     If a future run shows train_acc near 100 with val_auc off its peak, raise these back.\nWEIGHT_DECAY    = 0.02         # was 0.05\nLABEL_SMOOTH    = 0.03         # was 0.05\nMIXUP_ALPHA     = 0.2\nMIXUP_PROB      = 0.35         # was 0.50\nDROP_PATH       = 0.05         # was 0.10\nHEAD_DROPOUT    = 0.25         # was 0.30\n\nWIDTH_MULT      = 1.0          # 1.25 widens every branch ~55% more params (use only if a long\n                               # run plateaus with train_acc still under 98)\n\nAUX_LOSS_W      = 0.3          # deep-supervision weight on each branch head\nEMA_DECAY       = 0.999\nGRAD_CLIP       = 1.0\nNUM_WORKERS     = 2\nEARLY_STOP      = 30           # longer schedules need a longer fuse\nSESSION_MAX_HOURS = 11.0       # HARD budget for the whole notebook (Kaggle kills at 12 h)\nEVAL_RESERVE_MIN  = 15         # minutes held back per experiment for TTA + tables\nRESUME          = True         # continue from {name}_last.pt if the session died\nTHRESH_OBJECTIVE = \"accuracy\"  # \"accuracy\" | \"balanced\" | \"youden\" | \"fixed\" (=0.50)\nQUICK_TEST      = False        # True -> tiny subset, 2 epochs (smoke test)\n\n# ---------------------------------------------------------------- paths\nOUT_DIR   = \"/kaggle/working\" if os.path.isdir(\"/kaggle/working\") else \"./out\"\nCACHE_DIR = \"/kaggle/temp/dr_cache\" if os.path.isdir(\"/kaggle/temp\") else \"./dr_cache\"\nCKPT_DIR  = os.path.join(OUT_DIR, \"ckpt\")\nfor d in (OUT_DIR, CACHE_DIR, CKPT_DIR):\n    os.makedirs(d, exist_ok=True)\n\n# Leave as None for auto-discovery under /kaggle/input. Override only if discovery fails.\nDATASET_PATHS = {\"aptos\": None, \"idrid\": None, \"messidor2\": None}\n\n# ---------------------------------------------------------------- experiment grid\n# data    : \"aptos_only\" | \"selective\" (external grades 1,3,4 -> train) | \"full\" (all external grades)\n# preproc : \"rgb\" | \"ben\" | \"clahe\" | \"ben_clahe\" | \"green_clahe\"\n# arch    : \"dual\" (cross-attention) | \"dual_concat\" | \"context_only\" | \"lesion_only\"\n# balance : \"none\" | \"sampler\" | \"pos_weight\"     (never combine sampler + pos_weight)\n# Per-experiment overrides (all optional): epochs, img_size, batch, lr_scale, mixup_prob,\n# init_from (path to a checkpoint to WARM-START from).\nDEFAULTS = dict(data=\"selective\", preproc=\"ben\", arch=\"dual\", balance=\"none\", seed=SEED)\n\nEXPERIMENTS = [\n    dict(name=\"E1_aptos_rgb_dual\",        data=\"aptos_only\", preproc=\"rgb\"),\n    dict(name=\"E2_aptos_ben_dual\",        data=\"aptos_only\", preproc=\"ben\"),\n    dict(name=\"E3_selective_ben_dual\",    data=\"selective\",  preproc=\"ben\"),\n    dict(name=\"E4_full_ben_dual\",         data=\"full\",       preproc=\"ben\",       balance=\"sampler\"),\n    dict(name=\"E5_selective_benclahe\",    data=\"selective\",  preproc=\"ben_clahe\"),\n    dict(name=\"E6_selective_greenclahe\",  data=\"selective\",  preproc=\"green_clahe\"),\n    dict(name=\"E7_ablate_context_only\",   data=\"selective\",  preproc=\"ben\", arch=\"context_only\"),\n    dict(name=\"E8_ablate_lesion_only\",    data=\"selective\",  preproc=\"ben\", arch=\"lesion_only\"),\n    dict(name=\"E9_ablate_no_crossattn\",   data=\"selective\",  preproc=\"ben\", arch=\"dual_concat\"),\n    dict(name=\"E10_selective_ben_s1337\",  data=\"selective\",  preproc=\"ben\", seed=1337),\n    dict(name=\"E11_selective_ben_s2024\",  data=\"selective\",  preproc=\"ben\", seed=2024),\n    dict(name=\"E12_selective_sampler\",    data=\"selective\",  preproc=\"ben\", balance=\"sampler\"),\n    dict(name=\"E13_selective_posweight\",  data=\"selective\",  preproc=\"ben\", balance=\"pos_weight\"),\n    dict(name=\"E14_full_benclahe\",        data=\"full\",       preproc=\"ben_clahe\", balance=\"sampler\"),\n    # --- A/B: DOES THE FULL MERGE HELP *BINARY*? --------------------------------------\n    # The paper tested this (Table 3) and selective won - but for 5-CLASS grading, where\n    # 1,151 extra No-DR images just swamp the rare classes. Binary has a different problem:\n    #\n    #   selective  train prevalence 0.591 positive   test prevalence 0.508\n    #   full       train prevalence 0.489 positive   test prevalence 0.508\n    #\n    # Training at 59% positive and testing at 51% is a prevalence mismatch, and it shows in\n    # the runs: the tuned threshold keeps drifting above 0.5 (0.637 in one run) because the\n    # model over-predicts DR. The full merge happens to fix that. So the paper's conclusion\n    # may not transfer to binary, and it costs one session to find out instead of assuming.\n    # Identical fold, identical seed, identical everything except the data policy.\n    dict(name=\"K0_fold0\", data=\"selective\", preproc=\"ben\", seed=42,   fold=0),\n    # --- SEMI-SUPERVISED: teacher K0_fold0 labels the 1,928 unlabelled images, then these\n    #     students train on labelled + confident-pseudo data (~4,800 -> ~6,600 images).\n    dict(name=\"S1_student_f1\", data=\"selective\", preproc=\"ben\", seed=1337, fold=1, pseudo=True),\n    dict(name=\"S2_student_f2\", data=\"selective\", preproc=\"ben\", seed=2024, fold=2, pseudo=True),\n    dict(name=\"K1_fold1\", data=\"selective\", preproc=\"ben\", seed=1337, fold=1),\n    dict(name=\"K2_fold2\", data=\"selective\", preproc=\"ben\", seed=2024, fold=2),\n    dict(name=\"K3_fold3\", data=\"selective\", preproc=\"ben\", seed=7,    fold=3),\n    dict(name=\"K4_fold4\", data=\"selective\", preproc=\"ben\", seed=99,   fold=4),\n    # data policy A/B, if you want to settle selective vs full for binary\n    dict(name=\"AB_fold0_full\",      data=\"full\",      preproc=\"ben\", seed=42, fold=0),\n\n    # --- THE CURRENT BEST PLAN --------------------------------------------------------\n    # Three runs that DIFFER, instead of three seeds of one config. Rationale: the\n    # same-config seed ensemble failed twice (97.14 vs 97.54 for its best member) because\n    # three identical setups make almost identical mistakes. Varying resolution and\n    # preprocessing produces members that disagree, which is what an ensemble needs.\n    # Each is also a legitimate ablation row on its own.\n    dict(name=\"G1_384_ben\",       preproc=\"ben\",       seed=42,   img_size=384, batch=20),\n    dict(name=\"G2_448_ben\",       preproc=\"ben\",       seed=1337, img_size=448, batch=12),\n    dict(name=\"G3_384_benclahe\",  preproc=\"ben_clahe\", seed=2024, img_size=384, batch=20),\n\n    # --- FINE-TUNE AT 448 FROM YOUR OWN 384px CHECKPOINTS (needs an attached output) ---\n    # This is the cheap version of \"train at 448\". Instead of 3 h from random init, it\n    # warm-starts from a model you already trained and adapts it to the higher resolution\n    # in ~25 epochs (~45 min). Still no transfer learning: the starting weights are YOUR\n    # weights, learned from fundus images only - nothing external is imported.\n    # Bonus: a 448px child and its 384px parent disagree in different ways, so they\n    # ENSEMBLE well - unlike three same-resolution seeds, which were too correlated.\n    dict(name=\"F1_ft448_s42\",   seed=42,   img_size=448, batch=12, epochs=25,\n         lr_scale=0.15, mixup_prob=0.15,\n         init_from=os.path.join(CKPT_DIR, \"E3_selective_ben_dual_ema.pt\")),\n    dict(name=\"F2_ft448_s1337\", seed=1337, img_size=448, batch=12, epochs=25,\n         lr_scale=0.15, mixup_prob=0.15,\n         init_from=os.path.join(CKPT_DIR, \"E10_selective_ben_s1337_ema.pt\")),\n    dict(name=\"F3_ft448_s2024\", seed=2024, img_size=448, batch=12, epochs=25,\n         lr_scale=0.15, mixup_prob=0.15,\n         init_from=os.path.join(CKPT_DIR, \"E11_selective_ben_s2024_ema.pt\")),\n]\n\n# <<< CHOOSE WHAT TO RUN THIS SESSION >>>\n# One 110-epoch run costs ~3.2 h including evaluation, so THREE seeds fit in an 11 h session\n# (~9.6 h). Three seeds give a real variance estimate and a stronger cross-seed ensemble than\n# the two-seed attempt, which tied its best member.\n# BACK TO THE CONFIGURATION THAT ACTUALLY SCORED HIGHEST: 384px, ben, softmax grade head,\n# fixed 0.5 fusion, 110 epochs, selective merge. Three seeds -> 97.41 +/- 0.14.\n# Everything tried since then was measured and lost:\n#   ordinal grade head      96.86  (vs 97.41 for the same config with softmax)\n#   448px                   97.00\n#   ben_clahe preprocessing 97.14\n#   swept fuse weight       used by all three of the above\n#   warm restarts/snapshots 96.73  (vs 97.14 for a plain ema+tta)\n#   cross-seed ensembling   97.00  (vs 97.54 for its best member)\n# SESSION 1 (this one): the A/B, ~6.8 h. Two runs differing ONLY in the data policy.\n# SESSION 2: set data= on K1..K4 to whichever won, run them, then run the ensemble cell -\n#            it picks up every preds_*.npz automatically, including the A/B winner.\n# TEACHER -> PSEUDO-LABEL -> TWO STUDENTS.  ~2.9 h + 5 min + 2 x 3.6 h = ~10.2 h.\n# The teacher is a normal run; pseudo-labelling happens automatically before the first\n# student; students train on ~35% more images than any run so far.\nRUN_EXPERIMENTS = [\"K0_fold0\", \"S1_student_f1\", \"S2_student_f2\"]\n\nEXPERIMENTS = [{**DEFAULTS, **e} for e in EXPERIMENTS]\nEXP_BY_NAME = {e[\"name\"]: e for e in EXPERIMENTS}\n\nif QUICK_TEST:\n    SCHEDULE_CYCLES, CYCLE_LR_SCALE = [2], [1.0]\n    EPOCHS, WARMUP_EPOCHS, EARLY_STOP = 2, 0, 99\n\nSESSION_T0 = time.time()\n\n\ndef session_left_h():\n    return SESSION_MAX_HOURS - (time.time() - SESSION_T0) / 3600.0\n\n\nprint(f\"Output dir : {OUT_DIR}\")\nprint(f\"Cache dir  : {CACHE_DIR}\")\nprint(f\"Running    : {RUN_EXPERIMENTS}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"777a4948-ed35-4a41-83f7-c02494e1ff6b","cell_type":"markdown","source":"## 2. Imports & environment","metadata":{}},{"id":"e3cb9920-c168-4446-9c13-1d624a74b843","cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (accuracy_score, precision_score, recall_score, f1_score,\n                             roc_auc_score, average_precision_score, confusion_matrix,\n                             cohen_kappa_score, matthews_corrcoef, roc_curve,\n                             precision_recall_curve, balanced_accuracy_score, brier_score_loss)\n\ncv2.setNumThreads(0)  # avoid thread thrash inside DataLoader workers\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntorch.backends.cudnn.benchmark = True\n\ndef set_seed(s):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s); torch.cuda.manual_seed_all(s)\n\nset_seed(SEED)\n\n# AMP compatibility shim (torch <2.4 vs >=2.4)\ntry:\n    from torch.amp import autocast as _autocast, GradScaler as _GradScaler\n    def amp_autocast(): return _autocast(\"cuda\", enabled=torch.cuda.is_available())\n    def make_scaler(): return _GradScaler(\"cuda\", enabled=torch.cuda.is_available())\nexcept Exception:\n    from torch.cuda.amp import autocast as _autocast, GradScaler as _GradScaler\n    def amp_autocast(): return _autocast(enabled=torch.cuda.is_available())\n    def make_scaler(): return _GradScaler(enabled=torch.cuda.is_available())\n\n\n\n\n\n\n# Full metric names, printed as ROWS so nothing wraps in the Kaggle log.\nMETRIC_LABELS = [\n    (\"accuracy\",     \"Accuracy (%)\"),\n    (\"sensitivity\",  \"Sensitivity / Recall / TPR (%)\"),\n    (\"specificity\",  \"Specificity / TNR (%)\"),\n    (\"precision\",    \"Precision / PPV (%)\"),\n    (\"npv\",          \"Neg. predictive value (%)\"),\n    (\"f1\",           \"F1-score (%)\"),\n    (\"balanced_acc\", \"Balanced accuracy (%)\"),\n    (\"auc\",          \"ROC-AUC (%)\"),\n    (\"pr_auc\",       \"PR-AUC / Avg precision (%)\"),\n    (\"kappa\",        \"Cohen's kappa (%)\"),\n    (\"mcc\",          \"Matthews corr. coef. (%)\"),\n    (\"brier\",        \"Brier score (lower=better)\"),\n    (\"threshold\",    \"Decision threshold\"),\n    (\"tp\",           \"True positives  (DR called DR)\"),\n    (\"tn\",           \"True negatives  (healthy ok)\"),\n    (\"fp\",           \"False positives (false alarm)\"),\n    (\"fn\",           \"False negatives (MISSED DR)\"),\n]\nINT_METRICS = {\"tp\", \"tn\", \"fp\", \"fn\"}\n\npd.set_option(\"display.width\", 200)\npd.set_option(\"display.max_columns\", 40)\n\n\ndef metrics_table(named):\n    \"\"\"dict{name -> metrics dict} -> DataFrame with metrics as ROWS, runs as COLUMNS.\n\n    Metrics go down the page because there are 17 of them and only a handful of runs;\n    the transpose is what stops the Kaggle log wrapping columns into an unreadable mess.\n    \"\"\"\n    data = {}\n    for key, label in METRIC_LABELS:\n        data[label] = {name: m[key] for name, m in named.items()}\n    df = pd.DataFrame(data).T\n    df.index.name = \"metric\"\n    return df\n\n\ndef fmt_cell(v):\n    if isinstance(v, (int, np.integer)):\n        return f\"{v:d}\"\n    if isinstance(v, float) and float(v).is_integer() and abs(v) < 1e6:\n        return f\"{int(v):d}\"\n    return f\"{v:.4f}\" if abs(v) < 1 else f\"{v:.2f}\"\n\n\ndef show_table(title, df, note=None):\n    print()\n    print(title)\n    print(\"=\" * min(100, max(len(title), 60)))\n    print(df.to_string(float_format=fmt_cell, na_rep=\"-\"))\n    if note:\n        for line in textwrap.wrap(note, 96):\n            print(f\"  {line}\")\n\n\nprint(\"torch\", torch.__version__, \"| albumentations\", A.__version__, \"| device\", DEVICE)\nif torch.cuda.is_available():\n    print(\"gpu:\", torch.cuda.get_device_name(0))","metadata":{},"outputs":[],"execution_count":null},{"id":"fbfa0f18-a847-40df-a04d-5792ba5fe20a","cell_type":"markdown","source":"## 3. Dataset discovery\n\nKaggle mounts inputs in several different layouts (`/kaggle/input/<slug>`, or\n`/kaggle/input/competitions/<slug>` and `/kaggle/input/datasets/<owner>/<slug>`), and mirrors\nof the same dataset nest their images differently. So instead of walking down from named\nfolders, this scans `/kaggle/input` **once, globally**, builds a flat index of every image, and\nmatches CSV rows against it. Layout stops mattering.\n\nGuards that earned their place:\n* **Messidor-2 `adjudicated_dme`** (0-1, macular edema) looks like a label but is the wrong\n  column. Anything mentioning dme / macular / edema / risk / gradable is refused.\n* **IDRiD filename collision** — the train and test folders both contain `IDRiD_001.jpg`.\n  Candidate paths are scored against a train/test hint taken from the CSV filename.\n* **Zipped datasets** are extracted automatically, which is the usual reason a dataset folder\n  looks empty.\n\nThe inventory printed below is your ground truth. **Read it before training.**","metadata":{}},{"id":"630bdd23-9cce-46f8-80b7-366c99d922f3","cell_type":"code","source":"INPUT_ROOT = \"/kaggle/input\" if os.path.isdir(\"/kaggle/input\") else \".\"\nEXTRACT_ZIPS = True\n\nIMG_EXTS = {\".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\", \".bmp\"}\nID_PATTERNS = [\"id_code\", \"image_id\", \"image name\", \"image_name\", \"imagename\",\n               \"image\", \"filename\", \"file_name\", \"file\", \"name\", \"id\"]\nGRADE_PATTERNS = [\"adjudicated_dr_grade\", \"retinopathy grade\", \"retinopathy_grade\",\n                  \"dr_grade\", \"dr grade\", \"diagnosis\", \"level\", \"grade\", \"label\", \"class\"]\n# substring match -> column is definitely not the DR grade\nBAD_SUBSTR = [\"dme\", \"macular\", \"edema\", \"risk\", \"gradable\", \"quality\", \"ungradable\",\n              \"optic\", \"fovea\", \"coord\"]\n# exact match only (as substrings these would wrongly kill e.g. \"retinopathy grade\")\nBAD_EXACT = {\"x\", \"y\", \"row\", \"col\", \"index\", \"unnamed: 0\"}\n\n\ndef scan(root, exts):\n    out = []\n    for dp, _, fns in os.walk(root, followlinks=True):\n        for fn in fns:\n            if os.path.splitext(fn)[1].lower() in exts:\n                out.append(os.path.join(dp, fn))\n    return sorted(out)\n\n\ndef rescan_images():\n    global ALL_IMAGES, STEM_IDX, BASE_IDX\n    ALL_IMAGES = scan(INPUT_ROOT, IMG_EXTS) + scan(UNZIP_DIR, IMG_EXTS)\n    STEM_IDX, BASE_IDX = {}, {}\n    for p in ALL_IMAGES:\n        b = os.path.basename(p).lower()\n        STEM_IDX.setdefault(os.path.splitext(b)[0], []).append(p)\n        BASE_IDX.setdefault(b, []).append(p)\n\n\nUNZIP_DIR = os.path.join(CACHE_DIR, \"_unzipped\")\nos.makedirs(UNZIP_DIR, exist_ok=True)\n\nALL_CSVS = scan(INPUT_ROOT, {\".csv\"})\nALL_ZIPS = scan(INPUT_ROOT, {\".zip\"})\nrescan_images()\n\n# ---------------------------------------------------------------- inventory\nprint(f\"Scanned {INPUT_ROOT}: {len(ALL_CSVS)} csv, {len(ALL_IMAGES)} images, {len(ALL_ZIPS)} zips\\n\")\nfrom collections import defaultdict\n_per = defaultdict(int)\nfor p in ALL_IMAGES:\n    parts = os.path.relpath(p, INPUT_ROOT).split(os.sep)\n    _per[os.sep.join(parts[:3])] += 1\nprint(\"Images per folder:\")\nfor k, v in sorted(_per.items()):\n    print(f\"  {v:6d}  {k}\")\nif ALL_ZIPS:\n    print(\"\\nZips:\")\n    for z in ALL_ZIPS:\n        print(f\"  {z}\")\n\n# ---------------------------------------------------------------- unzip if needed\nif EXTRACT_ZIPS and ALL_ZIPS:\n    import zipfile\n    for z in ALL_ZIPS:\n        dest = os.path.join(UNZIP_DIR, os.path.splitext(os.path.basename(z))[0])\n        if os.path.isdir(dest) and scan(dest, IMG_EXTS):\n            continue\n        try:\n            with zipfile.ZipFile(z) as zf:\n                names = zf.namelist()\n                n_img = sum(1 for n in names\n                            if os.path.splitext(n)[1].lower() in IMG_EXTS)\n                if n_img < 20:\n                    continue\n                print(f\"\\nextracting {z}  ({n_img} images) -> {dest}\")\n                os.makedirs(dest, exist_ok=True)\n                zf.extractall(dest)\n        except Exception as e:\n            print(f\"  zip error on {z}: {e}\")\n    rescan_images()\n    ALL_CSVS = ALL_CSVS + scan(UNZIP_DIR, {\".csv\"})\n    print(f\"after extraction: {len(ALL_IMAGES)} images, {len(ALL_CSVS)} csv\")\n\nprint(\"\\n=== EVERY CSV FOUND ===\")\nfor p in ALL_CSVS:\n    try:\n        d = pd.read_csv(p, nrows=2)\n        print(f\"\\n{p}\\n   cols: {list(d.columns)}\")\n        if len(d):\n            print(f\"   first row: {d.iloc[0].to_dict()}\")\n    except Exception as e:\n        print(f\"\\n{p}\\n   [unreadable] {e}\")\nprint(\"=\" * 78)\n\n\n# ---------------------------------------------------------------- matching helpers\ndef pick_columns(df):\n    \"\"\"Return (id_col, grade_col) or (None, None).\"\"\"\n    norm = {c: \" \".join(str(c).strip().lower().split()) for c in df.columns}\n\n    id_col = None\n    for pat in ID_PATTERNS:                       # exact match first\n        for c, l in norm.items():\n            if l == pat:\n                id_col = c\n                break\n        if id_col:\n            break\n    if id_col is None:                            # then substring\n        for pat in ID_PATTERNS:\n            for c, l in norm.items():\n                if pat in l and df[c].dtype == object:\n                    id_col = c\n                    break\n            if id_col:\n                break\n    if id_col is None:\n        return None, None\n\n    cands = []\n    for c, l in norm.items():\n        if c == id_col or l in BAD_EXACT or any(b in l for b in BAD_SUBSTR):\n            continue\n        s = pd.to_numeric(df[c], errors=\"coerce\")\n        if s.notna().sum() < 0.8 * len(df):\n            continue\n        v = s.dropna()\n        if len(v) == 0 or v.min() < 0 or v.max() > 4 or v.nunique() < 2:\n            continue\n        if not np.allclose(v, v.round()):\n            continue\n        score = 0\n        for i, pat in enumerate(GRADE_PATTERNS):\n            if pat in l:\n                score = len(GRADE_PATTERNS) - i\n                break\n        cands.append((score, int(v.max()), c))\n    if not cands:\n        return id_col, None\n    cands.sort(reverse=True)\n    return id_col, cands[0][2]\n\n\ndef csv_source(path, cols):\n    \"\"\"Which dataset does this CSV belong to? Path keywords first, then column signature.\"\"\"\n    p = path.lower()\n    if \"idrid\" in p or \"disease grading\" in p:\n        return \"idrid\"\n    if \"messidor\" in p:\n        return \"messidor2\"\n    if \"aptos\" in p or \"blindness\" in p:\n        return \"aptos\"\n    if {\"id_code\", \"diagnosis\"} <= cols:\n        return \"aptos\"\n    if any(\"retinopathy grade\" in c or \"retinopathy_grade\" in c for c in cols):\n        return \"idrid\"\n    if any(\"adjudicated\" in c for c in cols):\n        return \"messidor2\"\n    return None\n\n\nSOURCE_KEYS = {\"aptos\": [\"aptos\", \"blindness\"], \"idrid\": [\"idrid\"],\n               \"messidor2\": [\"messidor\"]}\n\n\ndef resolve_image(raw, prefer=()):\n    \"\"\"Find the file for a CSV id. `prefer` biases which candidate wins on collisions.\"\"\"\n    s = str(raw).strip()\n    if not s or s.lower() == \"nan\":\n        return None\n    base = os.path.basename(s.replace(\"\\\\\", \"/\")).lower()\n    stem = os.path.splitext(base)[0]\n    cands, seen = [], set()\n    for key, idx in ((base, BASE_IDX), (stem, STEM_IDX)):\n        for p in idx.get(key, []):\n            if p not in seen:\n                seen.add(p); cands.append(p)\n    for e in IMG_EXTS:\n        for p in BASE_IDX.get(stem + e, []):\n            if p not in seen:\n                seen.add(p); cands.append(p)\n    if not cands:\n        return None\n    if len(cands) == 1:\n        return cands[0]\n    scored = [(sum(1 for k in prefer if k and k in c.lower()), c) for c in cands]\n    scored.sort(key=lambda t: -t[0])\n    return scored[0][1]\n\n\n# ---------------------------------------------------------------- build the table\nframes = []\nfor csv in ALL_CSVS:\n    try:\n        df = pd.read_csv(csv)\n    except Exception:\n        continue\n    if len(df) == 0:\n        continue\n    cols = {\" \".join(str(c).strip().lower().split()) for c in df.columns}\n    src = csv_source(csv, cols)\n    if src is None:\n        continue\n    if DATASET_PATHS.get(src) and not csv.startswith(DATASET_PATHS[src]):\n        continue                                   # manual override pins the folder\n    id_col, grade_col = pick_columns(df)\n    if id_col is None or grade_col is None:\n        continue\n\n    for c in df.columns:                           # Messidor-2 gradability\n        if \"gradable\" in str(c).lower():\n            g = pd.to_numeric(df[c], errors=\"coerce\")\n            df = df[g.fillna(1) > 0]\n            break\n\n    low = os.path.basename(csv).lower()\n    hint = \"train\" if \"train\" in low else (\"test\" if \"test\" in low else None)\n    prefer = tuple(SOURCE_KEYS[src] + ([hint] if hint else []))\n\n    rows = []\n    for _, r in df.iterrows():\n        g = pd.to_numeric(r[grade_col], errors=\"coerce\")\n        if pd.isna(g) or not (0 <= g <= 4):\n            continue\n        p = resolve_image(r[id_col], prefer)\n        if p:\n            rows.append((p, int(round(g))))\n    print(f\"  [{src:9s}] id='{id_col}' grade='{grade_col}' -> {len(rows):5d}/{len(df):5d} matched\"\n          f\"   ({os.path.basename(csv)})\")\n    if rows:\n        sub = pd.DataFrame(rows, columns=[\"image_path\", \"grade\"])\n        sub[\"source\"] = src\n        frames.append(sub)\n\nassert frames, (\"No labelled data matched. Read the CSV inventory above: check that a CSV with a \"\n                \"0-4 grade column exists and that its ids match real filenames.\")\nDATA = (pd.concat(frames, ignore_index=True)\n        .drop_duplicates(\"image_path\").reset_index(drop=True))\n\nfor src in [\"aptos\", \"idrid\", \"messidor2\"]:\n    if not (DATA.source == src).any():\n        print(f\"  [{src}] MISSING — experiments needing it will fall back to what is present.\")\nassert (DATA.source == \"aptos\").any(), \"APTOS 2019 is required (it defines the test set).\"\n\n# --- unlabelled pool: APTOS images with no label row anywhere -------------------------\n_labelled = set(DATA.image_path)\nUNLABELLED = [p for p in ALL_IMAGES\n              if p not in _labelled and (\"aptos\" in p.lower() or \"blindness\" in p.lower())]\nprint(f\"\\nUnlabelled APTOS pool: {len(UNLABELLED)} images \"\n      f\"(the competition test split - no labels, never evaluated on)\")\n\nprint(\"\\n=== Grade distribution per source (compare with paper Fig. 3) ===\")\nprint(pd.crosstab(DATA.source, DATA.grade, margins=True))\nprint(\"\\nExpected roughly: aptos 3662 | idrid 516 | messidor2 1744\")","metadata":{},"outputs":[],"execution_count":null},{"id":"9b56f7fa-2d30-4593-86a6-0b65a3e856b7","cell_type":"markdown","source":"## 4. Splits\n\nAPTOS is split 70 / 10 / 20 stratified on the **5-class** grade, exactly as in Table 1 of the paper,\nso the binary test set matches the published protocol. External data is only ever added to train,\nso there is no cross-dataset leakage into val/test and no image appears twice.","metadata":{}},{"id":"cf1bd37e-aa1f-437c-ab55-c9ac1375493e","cell_type":"code","source":"aptos = DATA[DATA.source == \"aptos\"].reset_index(drop=True)\next   = DATA[DATA.source != \"aptos\"].reset_index(drop=True)\n\ntr_a, tmp = train_test_split(aptos, test_size=0.30, stratify=aptos.grade, random_state=SEED)\nval_a, te_a = train_test_split(tmp, test_size=2 / 3, stratify=tmp.grade, random_state=SEED)\n\ntr_a  = tr_a.assign(split=\"train\")\nval_a = val_a.assign(split=\"val\")\nte_a  = te_a.assign(split=\"test\")\n\nSPLITS = pd.concat([tr_a, val_a, te_a], ignore_index=True)\nSPLITS[\"label\"] = (SPLITS.grade > 0).astype(int)\next = ext.assign(split=\"train\")\next[\"label\"] = (ext.grade > 0).astype(int)\n\nprint(\"APTOS split sizes (5-class):\")\nprint(pd.crosstab(SPLITS.split, SPLITS.grade, margins=True))\nprint(\"\\nAPTOS split sizes (binary):\")\nprint(pd.crosstab(SPLITS.split, SPLITS.label, margins=True))\nprint(\"\\nExternal pool available for training:\")\nprint(pd.crosstab(ext.source, ext.grade, margins=True))\n\n# integrity checks\nassert SPLITS.image_path.duplicated().sum() == 0\nassert len(set(SPLITS.image_path) & set(ext.image_path)) == 0\nSPLITS.to_csv(os.path.join(OUT_DIR, \"aptos_splits.csv\"), index=False)\n\n\nfrom sklearn.model_selection import StratifiedKFold\n\nPOOL = SPLITS[SPLITS.split != \"test\"].reset_index(drop=True)   # the 80% that is not test\nTEST_FIXED = SPLITS[SPLITS.split == \"test\"].reset_index(drop=True)\nFOLD_IDX = list(StratifiedKFold(N_FOLDS, shuffle=True, random_state=SEED)\n                .split(POOL, POOL.grade))\nprint(f\"\\n{N_FOLDS}-fold split of the non-test pool ({len(POOL)} images): \"\n      f\"train {len(FOLD_IDX[0][0])}, val {len(FOLD_IDX[0][1])} per fold; \"\n      f\"test held out at {len(TEST_FIXED)}\")\n\n\nPSEUDO_DF = None          # filled in by generate_pseudo_labels()\n\n\ndef _add_pseudo(tr):\n    if PSEUDO_DF is None or len(PSEUDO_DF) == 0:\n        return tr\n    return pd.concat([tr, PSEUDO_DF], ignore_index=True)\n\n\ndef make_frames(data_mode, fold=None, pseudo=False):\n    \"\"\"Return (train_df, val_df, test_df).\n\n    fold=None  -> the original fixed 70/10/20 split.\n    fold=k     -> fold k of the K-fold split over the non-test 80%. Test is identical\n                  in both cases, so results stay directly comparable.\n    \"\"\"\n    if fold is not None:\n        tr_i, va_i = FOLD_IDX[int(fold)]\n        tr = POOL.iloc[tr_i].copy()\n        va = POOL.iloc[va_i].copy()\n        te = TEST_FIXED.copy()\n        if data_mode == \"selective\":\n            tr = pd.concat([tr, ext[ext.grade.isin([1, 3, 4])]], ignore_index=True)\n        elif data_mode == \"full\":\n            tr = pd.concat([tr, ext], ignore_index=True)\n        if pseudo:\n            tr = _add_pseudo(tr)\n        if QUICK_TEST:\n            tr, va, te = (tr.sample(min(400, len(tr)), random_state=SEED),\n                          va.sample(min(120, len(va)), random_state=SEED),\n                          te.sample(min(160, len(te)), random_state=SEED))\n        return (tr.reset_index(drop=True), va.reset_index(drop=True),\n                te.reset_index(drop=True))\n\n    tr = SPLITS[SPLITS.split == \"train\"].copy()\n    if data_mode == \"selective\":\n        add = ext[ext.grade.isin([1, 3, 4])]\n        tr = pd.concat([tr, add], ignore_index=True)\n    elif data_mode == \"full\":\n        tr = pd.concat([tr, ext], ignore_index=True)\n    elif data_mode != \"aptos_only\":\n        raise ValueError(data_mode)\n    if pseudo:\n        tr = _add_pseudo(tr)\n    va = SPLITS[SPLITS.split == \"val\"].copy()\n    te = SPLITS[SPLITS.split == \"test\"].copy()\n    if QUICK_TEST:\n        tr = tr.sample(min(400, len(tr)), random_state=SEED)\n        va = va.sample(min(120, len(va)), random_state=SEED)\n        te = te.sample(min(160, len(te)), random_state=SEED)\n    return tr.reset_index(drop=True), va.reset_index(drop=True), te.reset_index(drop=True)\n\n\nfor m in [\"aptos_only\", \"selective\", \"full\"]:\n    t, v, s = make_frames(m)\n    print(f\"{m:12s} train={len(t):5d} (pos {t.label.mean():.3f})  val={len(v)}  test={len(s)}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"c41cd6fd-3c6b-4782-99f2-cf4184ca2b80","cell_type":"markdown","source":"## 5. Preprocessing & cache\n\nDone once and reused by every experiment:\n1. threshold away the black border, 2. pad to square, 3. resize to 448, 4. circular mask.\n\nThe *variant* (Ben Graham / CLAHE / green) is applied on the fly from the cache, so switching\npreprocessing costs nothing.","metadata":{}},{"id":"ac4085c5-7dc7-4dde-9be3-9b607d3c51da","cell_type":"code","source":"from concurrent.futures import ThreadPoolExecutor\n\nCIRCLE_MASK = np.zeros((CACHE_SIZE, CACHE_SIZE), np.uint8)\ncv2.circle(CIRCLE_MASK, (CACHE_SIZE // 2, CACHE_SIZE // 2), CACHE_SIZE // 2, 1, -1)\n\n\ndef crop_and_resize(path, size=CACHE_SIZE):\n    img = cv2.imread(path, cv2.IMREAD_COLOR)\n    if img is None:\n        return None\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    thr = max(7, int(gray.mean() * 0.10))\n    mask = gray > thr\n    if mask.sum() > 100:\n        rows, cols = np.where(mask.any(1))[0], np.where(mask.any(0))[0]\n        img = img[rows[0]:rows[-1] + 1, cols[0]:cols[-1] + 1]\n    h, w = img.shape[:2]\n    s = max(h, w)\n    canvas = np.zeros((s, s, 3), np.uint8)\n    canvas[(s - h) // 2:(s - h) // 2 + h, (s - w) // 2:(s - w) // 2 + w] = img\n    img = cv2.resize(canvas, (size, size), interpolation=cv2.INTER_AREA)\n    return img * CIRCLE_MASK[..., None]\n\n\ndef cache_name(path):\n    return os.path.join(CACHE_DIR, hashlib.md5(path.encode()).hexdigest()[:20] + \".png\")\n\n\ndef build_cache(paths):\n    todo = [p for p in paths if not os.path.exists(cache_name(p))]\n    print(f\"cache: {len(paths) - len(todo)} present, {len(todo)} to build\")\n    if not todo:\n        return\n    fails = []\n\n    def work(p):\n        img = crop_and_resize(p)\n        if img is None:\n            fails.append(p)\n            return\n        cv2.imwrite(cache_name(p), cv2.cvtColor(img, cv2.COLOR_RGB2BGR))\n\n    t0 = time.time()\n    with ThreadPoolExecutor(max_workers=8) as exe:\n        for i, _ in enumerate(exe.map(work, todo)):\n            if (i + 1) % 1000 == 0:\n                print(f\"  {i+1}/{len(todo)}  ({time.time()-t0:.0f}s)\")\n    print(f\"cache built in {time.time()-t0:.0f}s; {len(fails)} unreadable\")\n    if fails:\n        print(\"  e.g.\", fails[:3])\n\n\nALL_PATHS = pd.concat([SPLITS, ext], ignore_index=True).image_path.unique().tolist()\nbuild_cache(ALL_PATHS)\n\nCACHE_OK = {p for p in ALL_PATHS if os.path.exists(cache_name(p))}\nSPLITS = SPLITS[SPLITS.image_path.isin(CACHE_OK)].reset_index(drop=True)\next = ext[ext.image_path.isin(CACHE_OK)].reset_index(drop=True)\nprint(\"usable images:\", len(CACHE_OK))\n\n\ndef apply_preproc(img, mode):\n    \"\"\"img: uint8 RGB at CACHE_SIZE, already circle-masked.\"\"\"\n    if mode == \"rgb\":\n        return img\n    out = img\n    if mode.startswith(\"ben\"):\n        sigma = img.shape[0] / 30.0\n        bg = cv2.GaussianBlur(out, (0, 0), sigma)\n        out = cv2.addWeighted(out, 4, bg, -4, 128)\n        out = (out * CIRCLE_MASK[..., None]).astype(np.uint8)\n    if mode == \"green_clahe\":\n        g = img[:, :, 1]\n        cl = cv2.createCLAHE(clipLimit=2.5, tileGridSize=(8, 8)).apply(g)\n        out = np.stack([cl, cl, cl], -1) * CIRCLE_MASK[..., None]\n        return out.astype(np.uint8)\n    if mode in (\"clahe\", \"ben_clahe\"):\n        lab = cv2.cvtColor(out, cv2.COLOR_RGB2LAB)\n        lab[:, :, 0] = cv2.createCLAHE(clipLimit=2.5, tileGridSize=(8, 8)).apply(lab[:, :, 0])\n        out = cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)\n        out = (out * CIRCLE_MASK[..., None]).astype(np.uint8)\n    return out\n\n\n# visual sanity check\n_p = SPLITS.image_path.iloc[0]\n_base = cv2.cvtColor(cv2.imread(cache_name(_p)), cv2.COLOR_BGR2RGB)\nmodes = [\"rgb\", \"ben\", \"clahe\", \"ben_clahe\", \"green_clahe\"]\nfig, ax = plt.subplots(1, len(modes), figsize=(3 * len(modes), 3.2))\nfor a, m in zip(ax, modes):\n    a.imshow(apply_preproc(_base, m)); a.set_title(m); a.axis(\"off\")\nplt.tight_layout(); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"78498f61-7493-4845-8cba-0b2a7f871805","cell_type":"markdown","source":"## 6. Augmentation & Dataset\n\nAugmentation is heavier than a fine-tuning recipe would use, because a randomly-initialised\nnetwork on ~2.5-4.5k images will otherwise memorise the training set within a dozen epochs.\nArgument names are probed defensively — albumentations renamed several of them between 1.x and 2.x.","metadata":{}},{"id":"e4e15505-82ef-4638-ab05-96d719d0ad9d","cell_type":"code","source":"def safe_coarse_dropout(p=0.3):\n    try:\n        return A.CoarseDropout(num_holes_range=(1, 6), hole_height_range=(0.02, 0.10),\n                               hole_width_range=(0.02, 0.10), fill=0, p=p)\n    except TypeError:\n        return A.CoarseDropout(max_holes=6, max_height=int(IMG_SIZE * 0.10),\n                               max_width=int(IMG_SIZE * 0.10), fill_value=0, p=p)\n\n\ndef build_transforms(train, size=None):\n    size = size or IMG_SIZE\n    norm = [A.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)), ToTensorV2()]\n    if not train:\n        return A.Compose([A.Resize(size, size)] + norm)\n    return A.Compose([\n        A.Resize(size, size),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.Affine(scale=(0.88, 1.12), translate_percent=(-0.05, 0.05),\n                 rotate=(-180, 180), p=0.8),\n        A.RandomBrightnessContrast(p=0.5),\n        A.HueSaturationValue(p=0.3),\n        A.OneOf([A.GaussNoise(p=1.0), A.Blur(blur_limit=3, p=1.0), A.Sharpen(p=1.0)], p=0.3),\n        safe_coarse_dropout(0.3),\n    ] + norm)\n\n\nclass FundusDataset(Dataset):\n    def __init__(self, df, preproc, train, size=None):\n        self.paths = df.image_path.values\n        self.labels = df.label.values.astype(np.float32)\n        self.grades = df.grade.values.astype(np.int64)\n        self.preproc = preproc\n        self.tf = build_transforms(train, size)\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, i):\n        img = cv2.imread(cache_name(self.paths[i]), cv2.IMREAD_COLOR)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = apply_preproc(img, self.preproc)\n        x = self.tf(image=img)[\"image\"]\n        return x, self.labels[i], self.grades[i]\n\n\ndef build_loaders(cfg, size=None, batch=None):\n    size = size or IMG_SIZE\n    batch = batch or BATCH_SIZE\n    tr, va, te = make_frames(cfg[\"data\"], cfg.get(\"fold\"), cfg.get(\"pseudo\", False))\n    train_ds = FundusDataset(tr, cfg[\"preproc\"], True, size)\n    sampler, shuffle = None, True\n    if cfg[\"balance\"] == \"sampler\":\n        cnt = np.bincount(tr.label.values.astype(int), minlength=2)\n        w = (1.0 / cnt)[tr.label.values.astype(int)]\n        sampler = WeightedRandomSampler(torch.as_tensor(w, dtype=torch.double), len(w), True)\n        shuffle = False\n    common = dict(num_workers=NUM_WORKERS, pin_memory=True,\n                  persistent_workers=NUM_WORKERS > 0)\n    train_ld = DataLoader(train_ds, batch, shuffle=shuffle, sampler=sampler,\n                          drop_last=True, **common)\n    val_ld = DataLoader(FundusDataset(va, cfg[\"preproc\"], False, size), batch * 2,\n                        shuffle=False, **common)\n    test_ld = DataLoader(FundusDataset(te, cfg[\"preproc\"], False, size), batch * 2,\n                         shuffle=False, **common)\n    return train_ld, val_ld, test_ld, tr, va, te","metadata":{},"outputs":[],"execution_count":null},{"id":"072c7d31-d8b9-4873-b79b-8f67c2584d67","cell_type":"markdown","source":"## 6b. Preflight checks (automatic)\n\nRuns in seconds and **raises** rather than warns. It verifies the things that would\notherwise waste a session silently: that folds are disjoint, that every fold sees the\nsame held-out test set, that no image leaks between train and val/test, and that the\nqueued experiments actually differ from each other. If this cell passes, the run is\nwired correctly and you do not need to watch it.","metadata":{}},{"id":"228d0276-a120-48ee-aa61-97a2b5f3d761","cell_type":"code","source":"def preflight():\n    problems = []\n\n    # 1. fold integrity\n    fold_tests = []\n    for k in range(N_FOLDS):\n        tr, va, te = make_frames(\"selective\", fold=k)\n        overlap_v = set(tr.image_path) & set(va.image_path)\n        overlap_t = set(tr.image_path) & set(te.image_path)\n        if overlap_v:\n            problems.append(f\"fold {k}: {len(overlap_v)} images in BOTH train and val\")\n        if overlap_t:\n            problems.append(f\"fold {k}: {len(overlap_t)} images in BOTH train and test\")\n        fold_tests.append(frozenset(te.image_path))\n    if len(set(fold_tests)) != 1:\n        problems.append(\"folds do not share one fixed test set - results are not comparable\")\n\n    # 2. every fold's val slice is different\n    vals = [frozenset(make_frames(\"selective\", fold=k)[1].image_path) for k in range(N_FOLDS)]\n    if len(set(vals)) != N_FOLDS:\n        problems.append(\"two folds have identical validation sets\")\n\n    # 3. the queued experiments really do differ\n    rows, fps = [], {}\n    for name in RUN_EXPERIMENTS:\n        if name not in EXP_BY_NAME:\n            problems.append(f\"unknown experiment '{name}'\")\n            continue\n        c = EXP_BY_NAME[name]\n        tr, va, te = make_frames(c[\"data\"], c.get(\"fold\"), c.get(\"pseudo\", False))\n        fps[name] = (len(tr), round(float(tr.label.mean()), 4),\n                     hash(frozenset(tr.image_path)) & 0xFFFF)\n        rows.append({\"experiment\": name, \"data\": c[\"data\"], \"fold\": c.get(\"fold\", \"-\"),\n                     \"seed\": c[\"seed\"], \"img\": c.get(\"img_size\", IMG_SIZE),\n                     \"pseudo\": \"yes\" if c.get(\"pseudo\") else \"no\",\n                     \"train\": len(tr), \"train prevalence\": round(float(tr.label.mean()), 3),\n                     \"val\": len(va), \"test\": len(te)})\n    if rows:\n        show_table(\"PREFLIGHT - what each queued experiment will actually train on\",\n                   pd.DataFrame(rows).set_index(\"experiment\"))\n    same = [(a, b) for i, a in enumerate(fps) for b in list(fps)[i + 1:]\n            if fps[a] == fps[b] and EXP_BY_NAME[a].get(\"data\") != EXP_BY_NAME[b].get(\"data\")]\n    for a, b in same:\n        problems.append(f\"'{a}' and '{b}' differ in config but build the SAME training set \"\n                        f\"- the comparison would be meaningless\")\n\n    # 4. test set is the canonical APTOS 20%\n    _, _, te0 = make_frames(\"selective\", fold=0)\n    if not QUICK_TEST and len(te0) != len(TEST_FIXED):\n        problems.append(f\"test set is {len(te0)}, expected {len(TEST_FIXED)}\")\n    if (te0.source != \"aptos\").any():\n        problems.append(\"test set contains non-APTOS images\")\n\n    if problems:\n        raise RuntimeError(\"PREFLIGHT FAILED:\\n  - \" + \"\\n  - \".join(problems))\n    print(f\"\\nPREFLIGHT PASSED: {N_FOLDS} disjoint folds, one fixed \"\n          f\"{len(te0)}-image APTOS test set, no leakage, queued runs differ as intended.\")\n\n\npreflight()","metadata":{},"outputs":[],"execution_count":null},{"id":"f895e024-acc9-4aea-8148-cad48fd5f15a","cell_type":"markdown","source":"## 7. The network (written from scratch)\n\n```\n                     Stem  (384 -> 96, 64ch)\n                       |\n         +-------------+--------------+\n         |                            |\n  Context branch                Lesion branch\n  4 stages, stride-heavy        3 stages, dilated, stays wider\n  96->48->24->12->6             96->48->24->12\n  captures vessel arcade,       preserves microaneurysm /\n  optic disc, global colour     small-haemorrhage evidence\n         |                            |\n    1x1 -> 256                   (256 ch)\n    36 tokens                    144 tokens\n         +------ cross-attention -----+      (each branch queries the other)\n                       |\n    GeM(ctx) | GeM(les) | mean(ctx') | mean(les')   -> 1024-d\n                       |\n                 BN -> Dropout -> 256 -> SiLU -> 1 logit\n```\n\nAuxiliary heads sit directly on each branch's pooled features (deep supervision) — without them a\nscratch two-branch model tends to collapse onto whichever branch converges first.\n\nEvery block: 3x3 conv -> BN -> SiLU -> 3x3 conv -> BN -> Squeeze-Excite -> stochastic-depth residual.\nThe second BN's gamma is zero-initialised so each block starts as identity, which matters a lot when\nthere is no pretrained initialisation to fall back on.","metadata":{}},{"id":"9b8078b3-d710-4b76-81ab-ba20498f726d","cell_type":"code","source":"class DropPath(nn.Module):\n    def __init__(self, p=0.0):\n        super().__init__(); self.p = p\n\n    def forward(self, x):\n        if self.p == 0.0 or not self.training:\n            return x\n        keep = 1.0 - self.p\n        mask = x.new_empty(x.shape[0], *([1] * (x.dim() - 1))).bernoulli_(keep)\n        return x * mask / keep\n\n\nclass SE(nn.Module):\n    def __init__(self, c, r=8):\n        super().__init__()\n        h = max(8, c // r)\n        self.fc1, self.fc2 = nn.Conv2d(c, h, 1), nn.Conv2d(h, c, 1)\n\n    def forward(self, x):\n        s = F.adaptive_avg_pool2d(x, 1)\n        s = torch.sigmoid(self.fc2(F.silu(self.fc1(s))))\n        return x * s\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3.0, eps=1e-6):\n        super().__init__()\n        self.p = nn.Parameter(torch.ones(1) * p); self.eps = eps\n\n    def forward(self, x):\n        x = x.float().clamp(min=self.eps).pow(self.p)   # fp32: pow() is unstable in fp16\n        return F.avg_pool2d(x, (x.size(-2), x.size(-1))).pow(1.0 / self.p).flatten(1)\n\n\nclass ResSEBlock(nn.Module):\n    def __init__(self, cin, cout, stride=1, dilation=1, dp=0.0):\n        super().__init__()\n        pad = dilation\n        self.conv1 = nn.Conv2d(cin, cout, 3, stride, pad, dilation=dilation, bias=False)\n        self.bn1 = nn.BatchNorm2d(cout)\n        self.conv2 = nn.Conv2d(cout, cout, 3, 1, pad, dilation=dilation, bias=False)\n        self.bn2 = nn.BatchNorm2d(cout)\n        self.se = SE(cout)\n        self.dp = DropPath(dp)\n        self.short = None\n        if stride != 1 or cin != cout:\n            self.short = nn.Sequential(nn.Conv2d(cin, cout, 1, stride, bias=False),\n                                       nn.BatchNorm2d(cout))\n\n    def forward(self, x):\n        idt = x if self.short is None else self.short(x)\n        o = F.silu(self.bn1(self.conv1(x)))\n        o = self.se(self.bn2(self.conv2(o)))\n        return F.silu(idt + self.dp(o))\n\n\ndef stage(cin, cout, n, stride, dilation, dps):\n    blocks = [ResSEBlock(cin, cout, stride, 1, dps[0])]\n    blocks += [ResSEBlock(cout, cout, 1, dilation, dps[i]) for i in range(1, n)]\n    return nn.Sequential(*blocks)\n\n\ndef _w(c):\n    \"\"\"Round a channel count to a multiple of 8 under WIDTH_MULT.\"\"\"\n    return max(16, int(round(c * WIDTH_MULT / 8) * 8))\n\n\nclass Stem(nn.Module):\n    def __init__(self, w=None):\n        super().__init__()\n        w = w or _w(32)\n        self.out_ch = w * 2\n        self.net = nn.Sequential(\n            nn.Conv2d(3, w, 3, 2, 1, bias=False), nn.BatchNorm2d(w), nn.SiLU(inplace=True),\n            nn.Conv2d(w, w, 3, 1, 1, bias=False), nn.BatchNorm2d(w), nn.SiLU(inplace=True),\n            nn.Conv2d(w, w * 2, 3, 2, 1, bias=False), nn.BatchNorm2d(w * 2), nn.SiLU(inplace=True))\n\n    def forward(self, x):\n        return self.net(x)\n\n\nclass ContextBranch(nn.Module):\n    \"\"\"Stride-heavy: global retinal structure.\"\"\"\n\n    def __init__(self, cin, dp=DROP_PATH):\n        super().__init__()\n        c = [_w(96), _w(160), _w(256), _w(384)]\n        self.out_ch = c[3]\n        d = np.linspace(0, dp, 9).tolist()\n        self.s1 = stage(cin, c[0], 2, 2, 1, d[0:2])\n        self.s2 = stage(c[0], c[1], 2, 2, 1, d[2:4])\n        self.s3 = stage(c[1], c[2], 3, 2, 1, d[4:7])\n        self.s4 = stage(c[2], c[3], 2, 2, 1, d[7:9])\n\n    def forward(self, x):\n        return self.s4(self.s3(self.s2(self.s1(x))))\n\n\nclass LesionBranch(nn.Module):\n    \"\"\"Dilated, fewer channels: keeps small-lesion detail alive.\"\"\"\n\n    def __init__(self, cin, dp=DROP_PATH):\n        super().__init__()\n        c = [_w(96), _w(160), _w(256)]\n        self.out_ch = c[2]\n        d = np.linspace(0, dp, 6).tolist()\n        self.l1 = stage(cin, c[0], 2, 2, 1, d[0:2])\n        self.l2 = stage(c[0], c[1], 2, 2, 2, d[2:4])\n        self.l3 = stage(c[1], c[2], 2, 2, 2, d[4:6])\n\n    def forward(self, x):\n        return self.l3(self.l2(self.l1(x)))\n\n\nclass CrossAttn(nn.Module):\n    def __init__(self, d=256, heads=8, drop=0.1):\n        super().__init__()\n        self.nq, self.nk = nn.LayerNorm(d), nn.LayerNorm(d)\n        self.attn = nn.MultiheadAttention(d, heads, dropout=drop, batch_first=True)\n        self.ff = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d * 2), nn.GELU(),\n                                nn.Dropout(drop), nn.Linear(d * 2, d))\n\n    def forward(self, q, kv):\n        k = self.nk(kv)\n        h, _ = self.attn(self.nq(q), k, k, need_weights=False)\n        q = q + h\n        return q + self.ff(q)\n\n\nclass DRNet(nn.Module):\n    def __init__(self, arch=\"dual\", d=256, drop=HEAD_DROPOUT):\n        super().__init__()\n        self.arch = arch\n        self.stem = Stem()\n        self.use_ctx = arch in (\"dual\", \"dual_concat\", \"context_only\")\n        self.use_les = arch in (\"dual\", \"dual_concat\", \"lesion_only\")\n        feat = 0\n        if self.use_ctx:\n            self.ctx = ContextBranch(self.stem.out_ch)\n            self.ctx_proj = nn.Conv2d(self.ctx.out_ch, d, 1)\n            self.ctx_pool = GeM()\n            feat += d\n        if self.use_les:\n            self.les = LesionBranch(self.stem.out_ch)\n            self.les_proj = nn.Conv2d(self.les.out_ch, d, 1)\n            self.les_pool = GeM()\n            feat += d\n        self.cross = arch == \"dual\"\n        if self.cross:\n            self.c2l = CrossAttn(d)\n            self.l2c = CrossAttn(d)\n            feat += 2 * d\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(feat), nn.Dropout(drop),\n            nn.Linear(feat, 256), nn.SiLU(inplace=True), nn.Dropout(drop * 0.66),\n            nn.Linear(256, 1))\n        self.aux_ctx = nn.Linear(d, 1) if self.use_ctx else None\n        self.aux_les = nn.Linear(d, 1) if self.use_les else None\n        # ICDR grade head: shares every feature, adds ~4-5k params\n        n_out = 4 if GRADE_MODE == \"ordinal\" else 5\n        self.grade_head = nn.Linear(feat, n_out) if GRADE_AUX_W > 0 else None\n        self._init()\n\n    def _init(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n            elif isinstance(m, (nn.BatchNorm2d, nn.BatchNorm1d)):\n                nn.init.ones_(m.weight); nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.Linear):\n                nn.init.trunc_normal_(m.weight, std=0.02)\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n        for m in self.modules():                       # identity-start residuals\n            if isinstance(m, ResSEBlock):\n                nn.init.zeros_(m.bn2.weight)\n\n    def forward(self, x):\n        s = self.stem(x)\n        parts, aux = [], []\n        cf = lf = None\n        if self.use_ctx:\n            cf = self.ctx_proj(self.ctx(s))\n            cp = self.ctx_pool(cf)\n            parts.append(cp); aux.append(self.aux_ctx(cp))\n        if self.use_les:\n            lf = self.les_proj(self.les(s))\n            lp = self.les_pool(lf)\n            parts.append(lp); aux.append(self.aux_les(lp))\n        if self.cross:\n            B, C, H, W = cf.shape\n            ct = cf.flatten(2).transpose(1, 2)          # B, 36, 256\n            lt = lf.flatten(2).transpose(1, 2)          # B, 144, 256\n            parts.append(self.c2l(lt, ct).mean(1))      # lesion queries context\n            parts.append(self.l2c(ct, lt).mean(1))      # context queries lesion\n        z = torch.cat(parts, 1)\n        grade = self.grade_head(z) if self.grade_head is not None else None\n        return self.head(z).squeeze(1), [a.squeeze(1) for a in aux], grade\n\n\ndef count_params(m):\n    return sum(p.numel() for p in m.parameters() if p.requires_grad)\n\n\nfor a in [\"dual\", \"dual_concat\", \"context_only\", \"lesion_only\"]:\n    _m = DRNet(a)\n    print(f\"{a:14s} {count_params(_m)/1e6:6.2f} M params\")\ndel _m\n\n# shape check\n_m = DRNet(\"dual\").eval()\nwith torch.no_grad():\n    _o, _a, _g = _m(torch.randn(2, 3, IMG_SIZE, IMG_SIZE))\nprint(\"forward ok: binary\", tuple(_o.shape), \"| aux\", [tuple(x.shape) for x in _a],\n      \"| grade\", tuple(_g.shape) if _g is not None else None)\ndel _m","metadata":{},"outputs":[],"execution_count":null},{"id":"49afb525-c3d8-49a8-8168-c075b380b41d","cell_type":"markdown","source":"## 8. Loss, EMA, metrics","metadata":{}},{"id":"3a369a62-8ac8-4c89-95d2-73ca09ae457b","cell_type":"code","source":"class EMA:\n    def __init__(self, model, decay=EMA_DECAY):\n        self.decay = decay\n        self.shadow = {k: v.detach().clone().float() for k, v in model.state_dict().items()}\n        self.step = 0\n\n    @torch.no_grad()\n    def update(self, model):\n        self.step += 1\n        d = min(self.decay, (1 + self.step) / (10 + self.step))   # fast early, slow later\n        for k, v in model.state_dict().items():\n            if v.dtype.is_floating_point:\n                self.shadow[k].mul_(d).add_(v.detach().float(), alpha=1 - d)\n            else:\n                self.shadow[k] = v.detach().clone().float()\n\n    def state_dict(self, model):\n        out = {}\n        for k, v in model.state_dict().items():\n            out[k] = self.shadow[k].to(v.dtype) if k in self.shadow else v\n        return out\n\n\ndef smooth(y, eps=LABEL_SMOOTH):\n    return y * (1 - 2 * eps) + eps\n\n\ndef mixup_batch(x, targets, alpha=MIXUP_ALPHA):\n    \"\"\"Mixes the image and EVERY target with the same lambda and permutation.\"\"\"\n    lam = float(np.random.beta(alpha, alpha))\n    idx = torch.randperm(x.size(0), device=x.device)\n    return lam * x + (1 - lam) * x[idx], [lam * t + (1 - lam) * t[idx] for t in targets]\n\n\ndef soft_ce(logits, target):\n    \"\"\"Cross-entropy against soft (mixed) one-hot targets.\"\"\"\n    return -(target * F.log_softmax(logits.float(), dim=1)).sum(1).mean()\n\n\ndef grade_targets(gr):\n    \"\"\"Ordinal: [g>0, g>1, g>2, g>3]. Softmax: one-hot over 5 classes.\"\"\"\n    gr = gr.clamp(0, 4)\n    if GRADE_MODE == \"ordinal\":\n        ks = torch.arange(4, device=gr.device).unsqueeze(0)\n        return (gr.unsqueeze(1) > ks).float()\n    return F.one_hot(gr, 5).float()\n\n\ndef grade_loss(logits, target):\n    if GRADE_MODE == \"ordinal\":\n        return F.binary_cross_entropy_with_logits(logits.float(), target)\n    return soft_ce(logits, target)\n\n\ndef grade_to_dr_prob(logits):\n    \"\"\"The grade head's own answer to 'is there any DR'.\"\"\"\n    if GRADE_MODE == \"ordinal\":\n        return torch.sigmoid(logits.float()[:, 0])          # P(grade > 0), directly\n    return 1.0 - torch.softmax(logits.float(), 1)[:, 0]\n\n\ndef best_fuse_weight(vy, vp, vpg):\n    \"\"\"Blend weight for binary vs grade head, swept on VALIDATION only.\"\"\"\n    if FUSE_WEIGHT != \"auto\":\n        return float(FUSE_WEIGHT)\n    best_w, best_s = 0.0, -1.0\n    for w in np.arange(0.0, 1.01, 0.05):\n        p = (1 - w) * vp + w * vpg\n        m = binary_metrics(vy, p, best_threshold(vy, p))\n        sc = 0.5 * m[\"auc\"] + 0.5 * m[\"accuracy\"]\n        if sc > best_s:\n            best_w, best_s = float(w), sc\n    return best_w\n\n\ndef binary_metrics(y, prob, thr=0.5):\n    pred = (prob >= thr).astype(int)\n    tn, fp, fn, tp = confusion_matrix(y, pred, labels=[0, 1]).ravel()\n    return {\n        \"threshold\":   float(thr),\n        \"accuracy\":    accuracy_score(y, pred) * 100,\n        \"balanced_acc\": balanced_accuracy_score(y, pred) * 100,\n        \"precision\":   precision_score(y, pred, zero_division=0) * 100,\n        \"sensitivity\": recall_score(y, pred, zero_division=0) * 100,\n        \"specificity\": (tn / (tn + fp) * 100) if (tn + fp) else 0.0,\n        \"npv\":         (tn / (tn + fn) * 100) if (tn + fn) else 0.0,\n        \"f1\":          f1_score(y, pred, zero_division=0) * 100,\n        \"auc\":         roc_auc_score(y, prob) * 100,\n        \"pr_auc\":      average_precision_score(y, prob) * 100,\n        \"kappa\":       cohen_kappa_score(y, pred) * 100,\n        \"mcc\":         matthews_corrcoef(y, pred) * 100,\n        \"brier\":       brier_score_loss(y, prob),\n        \"tp\": int(tp), \"tn\": int(tn), \"fp\": int(fp), \"fn\": int(fn),\n    }\n\n\ndef best_threshold(y, prob, objective=None):\n    \"\"\"Tuned on VALIDATION only.\n\n    Two things this gets right that a naive sweep does not:\n    * Youden's J maximises sensitivity+specificity, which is NOT accuracy - on a\n      prevalence-shifted training set it drifts well above 0.5 and costs accuracy.\n    * With only a few hundred validation images many thresholds tie for best, and a\n      plain argmax lands on whichever noise spike comes first. Taking the CENTRE of\n      the optimal plateau is far more stable on the test set.\n    \"\"\"\n    objective = objective or THRESH_OBJECTIVE\n    if objective == \"fixed\":\n        return 0.5\n    if objective == \"youden\":\n        fpr, tpr, thr = roc_curve(y, prob)\n        return float(thr[int(np.argmax(tpr - fpr))])\n    score = balanced_accuracy_score if objective == \"balanced\" else accuracy_score\n    grid = np.round(np.arange(0.05, 0.955, 0.005), 3)\n    vals = np.array([score(y, (prob >= t).astype(int)) for t in grid])\n    plateau = grid[vals >= vals.max() - 1e-9]        # every threshold tying for best\n    return float(np.median(plateau))\n\n\ndef bce_np(y, prob, eps=1e-7):\n    \"\"\"Validation loss from probabilities - no extra forward pass.\"\"\"\n    p = np.clip(prob, eps, 1 - eps)\n    return float(-(y * np.log(p) + (1 - y) * np.log(1 - p)).mean())","metadata":{},"outputs":[],"execution_count":null},{"id":"7a3847ee-480f-4830-abd7-229043e11eb3","cell_type":"markdown","source":"## 9. Train / evaluate\n\nEvery epoch logs train loss, train accuracy, validation loss, and validation accuracy + AUC for\n**both** the raw and EMA weights, as an aligned table. `*` marks an epoch that improved the best\nvalidation AUC (the checkpointing criterion).\n\nTrain accuracy is measured only on batches where mixup did **not** fire — on a mixed batch the\nhard label is meaningless, so counting it would report a fictional number.","metadata":{}},{"id":"cd39d93e-7948-4dbc-b912-d2e440eee3f1","cell_type":"code","source":"@torch.no_grad()\ndef predict(model, loader, tta=False):\n    model.eval()\n    probs, ys, gs, gprobs = [], [], [], []\n    for x, y, g in loader:\n        x = x.to(DEVICE, non_blocking=True).contiguous(memory_format=torch.channels_last)\n        views = [x]\n        if tta:\n            views += [torch.flip(x, [3]), torch.flip(x, [2]), torch.flip(x, [2, 3])]\n            if TTA_MODE == \"flip5\":                      # centre zoom: lesions get bigger\n                h, w = x.shape[-2:]\n                crop = x[:, :, h // 8:h - h // 8, w // 8:w - w // 8]\n                views.append(F.interpolate(crop, size=(h, w), mode=\"bilinear\",\n                                           align_corners=False))\n        acc_b, acc_g, n = 0, 0, len(views)\n        with amp_autocast():\n            for v in views:\n                logit, _, glog = model(v)\n                acc_b = acc_b + torch.sigmoid(logit.float())\n                if glog is not None:\n                    acc_g = acc_g + grade_to_dr_prob(glog)\n        probs.append((acc_b / n).cpu().numpy())\n        gprobs.append((acc_g / n).cpu().numpy() if torch.is_tensor(acc_g)\n                      else np.zeros(len(y)))\n        ys.append(y.numpy()); gs.append(g.numpy())\n    return (np.concatenate(probs), np.concatenate(ys), np.concatenate(gs),\n            np.concatenate(gprobs))\n\n\nEPOCH_HDR = (f\"{'epoch':>6} {'lr':>9} {'tr_loss':>9} {'tr_acc':>7} {'va_loss':>9} \"\n             f\"{'va_acc':>7} {'va_auc':>7} {'va_acc*':>8} {'va_auc*':>8} {'best':>5} {'sec':>5}\")\n\n\ndef resolve_ckpt(path):\n    \"\"\"Find a checkpoint by name anywhere under /kaggle/input or /kaggle/working.\n\n    /kaggle/working does NOT persist between sessions, so a warm-start checkpoint from a\n    previous run lives in that run's OUTPUT, which has to be attached as an input dataset.\n    This searches for the file by basename so the exact mount path does not matter.\n    \"\"\"\n    if not path:\n        return None\n    if os.path.exists(path):\n        return path\n    want = os.path.basename(path)\n    for root in [\"/kaggle/input\", \"/kaggle/working\", OUT_DIR]:\n        if not os.path.isdir(root):\n            continue\n        for dp, _, fns in os.walk(root):\n            if want in fns:\n                found = os.path.join(dp, want)\n                print(f\"  found checkpoint at {found}\")\n                return found\n    return None\n\n\ndef load_compatible(model, path):\n    \"\"\"Load only the tensors whose shapes match. Lets a 448px fine-tune start from a 384px\n    checkpoint, and survives a changed grade-head width (softmax 5 -> ordinal 4).\"\"\"\n    sd = torch.load(path, map_location=DEVICE)\n    own = model.state_dict()\n    ok = {k: v for k, v in sd.items() if k in own and own[k].shape == v.shape}\n    skipped = [k for k in own if k not in ok]\n    model.load_state_dict(ok, strict=False)\n    print(f\"  warm-start from {os.path.basename(path)}: loaded {len(ok)}/{len(own)} tensors\"\n          + (f\", re-initialised {len(skipped)} ({skipped[:3]}{'...' if len(skipped) > 3 else ''})\"\n             if skipped else \"\"))\n    return model\n\n\ndef train_one(cfg, verbose=True, deadline_ts=None):\n    set_seed(cfg[\"seed\"])\n    n_epochs = int(cfg.get(\"epochs\", EPOCHS))\n    size = int(cfg.get(\"img_size\", IMG_SIZE))\n    batch = int(cfg.get(\"batch\", BATCH_SIZE))\n    lr = BASE_LR * float(cfg.get(\"lr_scale\", 1.0))\n    mix_p = float(cfg.get(\"mixup_prob\", MIXUP_PROB))\n    warm_ep = min(WARMUP_EPOCHS, max(1, n_epochs // 10))\n    train_ld, val_ld, test_ld, tr_df, va_df, te_df = build_loaders(cfg, size=size, batch=batch)\n\n    model = DRNet(cfg[\"arch\"]).to(DEVICE).to(memory_format=torch.channels_last)\n    if cfg.get(\"init_from\"):\n        found = resolve_ckpt(cfg[\"init_from\"])\n        if found:\n            load_compatible(model, found)\n        else:\n            raise FileNotFoundError(\n                f\"init_from checkpoint not found: {cfg['init_from']}\\n\"\n                f\"  A warm-start run is pointless without it - 25 epochs from scratch is\\n\"\n                f\"  far worse than the 110-epoch model you already have.\\n\"\n                f\"  Fix: attach the previous run's OUTPUT as an input dataset (Add Input ->\\n\"\n                f\"  Your Work -> Notebooks -> pick that version), then re-run. The file is\\n\"\n                f\"  searched for by name, so the mount path does not matter.\")\n    ema = EMA(model)\n    ema_model = DRNet(cfg[\"arch\"]).to(DEVICE).to(memory_format=torch.channels_last)\n\n    pos_weight = None\n    if cfg[\"balance\"] == \"pos_weight\":\n        n = tr_df.label.values\n        pos_weight = torch.tensor([(n == 0).sum() / max(1, (n == 1).sum())],\n                                  dtype=torch.float32, device=DEVICE)\n    crit = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\n\n    decay, no_decay = [], []\n    for _, prm in model.named_parameters():\n        (no_decay if prm.ndim <= 1 else decay).append(prm)\n    opt = torch.optim.AdamW([{\"params\": decay, \"weight_decay\": WEIGHT_DECAY},\n                             {\"params\": no_decay, \"weight_decay\": 0.0}], lr=lr)\n\n    opt_steps = max(1, math.ceil(len(train_ld) / GRAD_ACCUM))\n    warm = warm_ep * opt_steps\n    cycles = SCHEDULE_CYCLES if n_epochs == EPOCHS else [n_epochs]\n    bounds, acc_e = [], 0                       # cycle boundaries in EPOCHS\n    for c in cycles:\n        acc_e += c\n        bounds.append(acc_e)\n\n    def lr_fn(st):\n        \"\"\"Cosine warm restarts: each cycle anneals to ~0, then jumps back up (scaled).\"\"\"\n        if st < warm:\n            return (st + 1) / max(1, warm)\n        ep_f = st / opt_steps                   # fractional epoch\n        lo = 0\n        for ci, hi in enumerate(bounds):\n            if ep_f < hi or ci == len(bounds) - 1:\n                span = max(1e-9, hi - lo)\n                q = min(1.0, max(0.0, (ep_f - lo) / span))\n                sc = CYCLE_LR_SCALE[ci] if ci < len(CYCLE_LR_SCALE) else 1.0\n                return sc * (0.01 + 0.99 * 0.5 * (1 + math.cos(math.pi * q)))\n            lo = hi\n        return 0.01\n\n    sched = torch.optim.lr_scheduler.LambdaLR(opt, lr_fn)\n    scaler = make_scaler()\n\n    raw_ck = os.path.join(CKPT_DIR, f\"{cfg['name']}_raw.pt\")\n    ema_ck = os.path.join(CKPT_DIR, f\"{cfg['name']}_ema.pt\")\n    last_ck = os.path.join(CKPT_DIR, f\"{cfg['name']}_last.pt\")\n\n    snap_dir = os.path.join(CKPT_DIR, f\"{cfg['name']}_snaps\")\n    os.makedirs(snap_dir, exist_ok=True)\n    best_auc, best_epoch, hist, start_ep = -1.0, -1, [], 0\n    cur_size, cur_batch = size, batch\n    if RESUME and os.path.exists(last_ck):\n        try:\n            ck = torch.load(last_ck, map_location=DEVICE, weights_only=False)\n            model.load_state_dict(ck[\"model\"]); opt.load_state_dict(ck[\"opt\"])\n            sched.load_state_dict(ck[\"sched\"]); scaler.load_state_dict(ck[\"scaler\"])\n            ema.shadow = {k: v.to(DEVICE) for k, v in ck[\"ema\"].items()}\n            ema.step = ck[\"ema_step\"]\n            best_auc, best_epoch, hist, start_ep = (ck[\"best_auc\"], ck[\"best_epoch\"],\n                                                    ck[\"hist\"], ck[\"epoch\"])\n            cur_size = ck.get(\"size\", IMG_SIZE)\n            if cur_size != IMG_SIZE:\n                cur_batch = PROGRESSIVE_BATCH\n                train_ld, val_ld, test_ld, tr_df, va_df, te_df = build_loaders(\n                    cfg, size=cur_size, batch=cur_batch)\n            print(f\"RESUMED {cfg['name']} from epoch {start_ep} (best val AUC {best_auc:.2f})\")\n        except Exception as e:\n            print(f\"resume failed ({e}) - starting fresh\")\n\n    print(f\"train={len(tr_df)} (pos {tr_df.label.mean():.3f}) | val={len(va_df)} | \"\n          f\"test={len(te_df)} | params={count_params(model)/1e6:.2f}M | \"\n          f\"img={IMG_SIZE} | batch={BATCH_SIZE}x{GRAD_ACCUM} | epochs={EPOCHS}\")\n    if verbose:\n        print(EPOCH_HDR)\n        print(\"-\" * len(EPOCH_HDR))\n\n    if deadline_ts:\n        print(f\"time budget: {(deadline_ts - time.time())/3600:.2f} h for this experiment \"\n              f\"(session has {session_left_h():.2f} h left)\")\n    for ep in range(start_ep, n_epochs):\n        # ---- progressive resize: switch to the larger input for the final cycle ----\n        if (PROGRESSIVE_SIZE and cur_size != PROGRESSIVE_SIZE and ep >= PROGRESSIVE_AT):\n            fits = True\n            if deadline_ts and hist:\n                scale = (PROGRESSIVE_SIZE / cur_size) ** 2\n                est = (hist[-1][\"secs\"] * scale) * max(0, n_epochs - ep)  # epochs LEFT\n                fits = time.time() + est < deadline_ts\n            if fits:\n                cur_size, cur_batch = PROGRESSIVE_SIZE, PROGRESSIVE_BATCH\n                train_ld, val_ld, test_ld, tr_df, va_df, te_df = build_loaders(\n                    cfg, size=cur_size, batch=cur_batch)\n                print(f\"  -> progressive resize: now training at {cur_size}px \"\n                      f\"(batch {cur_batch}) from epoch {ep+1}\")\n            else:\n                print(f\"  -> skipping {PROGRESSIVE_SIZE}px phase: would not fit the time budget\")\n                globals()[\"PROGRESSIVE_SIZE\"] = None\n\n        model.train()\n        t0 = time.time()\n        run_loss = run_n = clean_ok = clean_n = 0\n        opt.zero_grad(set_to_none=True)\n\n        for it, (x, y, gr) in enumerate(train_ld):\n            x = x.to(DEVICE, non_blocking=True).contiguous(memory_format=torch.channels_last)\n            y_hard = y.to(DEVICE, non_blocking=True).float()\n            y_t = smooth(y_hard)\n            g_t = grade_targets(gr.to(DEVICE, non_blocking=True))\n            mixed = MIXUP_ALPHA > 0 and np.random.rand() < mix_p\n            if mixed:\n                x, (y_t, g_t) = mixup_batch(x, [y_t, g_t])\n\n            with amp_autocast():\n                logit, aux, glog = model(x)\n                loss = crit(logit, y_t)\n                for a in aux:\n                    loss = loss + AUX_LOSS_W * crit(a, y_t)\n                if glog is not None:\n                    loss = loss + GRADE_AUX_W * grade_loss(glog, g_t)\n            scaler.scale(loss / GRAD_ACCUM).backward()\n\n            if (it + 1) % GRAD_ACCUM == 0 or (it + 1) == len(train_ld):\n                scaler.unscale_(opt)\n                nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                scaler.step(opt); scaler.update()\n                opt.zero_grad(set_to_none=True)\n                sched.step(); ema.update(model)\n\n            run_loss += loss.item() * x.size(0); run_n += x.size(0)\n            if not mixed:                       # hard labels only meaningful un-mixed\n                with torch.no_grad():\n                    pred = (torch.sigmoid(logit.float()) >= 0.5).float()\n                    clean_ok += (pred == y_hard).sum().item(); clean_n += x.size(0)\n\n        vp, vy, _, _ = predict(model, val_ld)\n        ema_model.load_state_dict(ema.state_dict(model))\n        ep_, ey_, _, _ = predict(ema_model, val_ld)\n\n        row = dict(\n            epoch=ep + 1, lr=sched.get_last_lr()[0],\n            train_loss=run_loss / max(1, run_n),\n            train_acc=100.0 * clean_ok / max(1, clean_n),\n            val_loss=bce_np(vy, vp),\n            val_acc_raw=accuracy_score(vy, (vp >= 0.5).astype(int)) * 100,\n            val_auc_raw=roc_auc_score(vy, vp) * 100,\n            val_loss_ema=bce_np(ey_, ep_),\n            val_acc_ema=accuracy_score(ey_, (ep_ >= 0.5).astype(int)) * 100,\n            val_auc_ema=roc_auc_score(ey_, ep_) * 100,\n            secs=time.time() - t0)\n        hist.append(row)\n\n        improved = max(row[\"val_auc_raw\"], row[\"val_auc_ema\"]) > best_auc\n        if improved:\n            best_auc = max(row[\"val_auc_raw\"], row[\"val_auc_ema\"])\n            best_epoch = ep + 1\n            torch.save(model.state_dict(), raw_ck)\n            torch.save(ema.state_dict(model), ema_ck)\n\n        if (ep + 1) in bounds:\n            snap = os.path.join(snap_dir, f\"cycle{bounds.index(ep+1)}.pt\")\n            torch.save(ema.state_dict(model), snap)\n            print(f\"  -> snapshot saved at end of cycle {bounds.index(ep+1)+1}: \"\n                  f\"{os.path.basename(snap)} (val AUC ema {row['val_auc_ema']:.2f})\")\n\n        torch.save(dict(model=model.state_dict(), opt=opt.state_dict(),\n                        sched=sched.state_dict(), scaler=scaler.state_dict(),\n                        ema={k: v.cpu() for k, v in ema.shadow.items()}, ema_step=ema.step,\n                        best_auc=best_auc, best_epoch=best_epoch, hist=hist,\n                        epoch=ep + 1, size=cur_size), last_ck)\n\n        if verbose:\n            print(f\"{row['epoch']:>6} {row['lr']:>9.2e} {row['train_loss']:>9.4f} \"\n                  f\"{row['train_acc']:>7.2f} {row['val_loss']:>9.4f} \"\n                  f\"{row['val_acc_raw']:>7.2f} {row['val_auc_raw']:>7.2f} \"\n                  f\"{row['val_acc_ema']:>8.2f} {row['val_auc_ema']:>8.2f} \"\n                  f\"{'  *' if improved else '   ':>5} {row['secs']:>5.0f}\")\n\n        if deadline_ts:\n            recent = [r[\"secs\"] for r in hist[-3:]]\n            est_next = (sum(recent) / len(recent)) * 1.15 + 20      # +15% slack, +eval overhead\n            if time.time() + est_next > deadline_ts:\n                print(f\"time guard: stopping at epoch {ep+1} - the next epoch \"\n                      f\"(~{est_next/60:.1f} min) would not fit in the remaining budget. \"\n                      f\"Re-run with RESUME=True to continue from here.\")\n                break\n        if ep + 1 - best_epoch >= EARLY_STOP:\n            print(f\"early stop: no val-AUC gain for {EARLY_STOP} epochs\")\n            break\n\n    del ema_model\n    torch.cuda.empty_cache()\n    h = pd.DataFrame(hist)\n    print(f\"\\nbest epoch {best_epoch}/{len(h)} | best val AUC {best_auc:.2f} | \"\n          f\"final train_acc {h.train_acc.iloc[-1]:.2f} | final val_acc {h.val_acc_raw.iloc[-1]:.2f}\")\n\n    # ---- automatic fit diagnosis: tells you what to change next ----\n    tr_acc, last_auc = h.train_acc.iloc[-1], h.val_auc_raw.iloc[-1]\n    tail = h.val_auc_raw.tail(min(10, len(h)))\n    if tr_acc > 99.0 and last_auc < best_auc - 0.5:\n        print(\"  DIAGNOSIS overfitting -> raise DROP_PATH / MIXUP_PROB / WEIGHT_DECAY, \"\n              \"or switch data to 'full' for more images\")\n    elif tr_acc < 97.0 and best_epoch >= 0.85 * len(h):\n        print(\"  DIAGNOSIS underfitting AND still improving at the end -> raise EPOCHS\")\n    elif tr_acc < 97.0:\n        print(\"  DIAGNOSIS underfitting but plateaued -> raise capacity or IMG_SIZE, \"\n              \"or lower regularisation further\")\n    elif len(tail) >= 5 and (tail.max() - tail.min()) < 0.3:\n        print(\"  DIAGNOSIS converged -> next gains come from IMG_SIZE 448 or an ensemble\")\n    return dict(cfg=cfg, hist=hist, best_epoch=best_epoch, best_val_auc=best_auc,\n                raw_ck=raw_ck, ema_ck=ema_ck, loaders=(val_ld, test_ld),\n                snaps=sorted(glob.glob(os.path.join(snap_dir, \"*.pt\"))),\n                img_size=size, frames=(tr_df, va_df, te_df))","metadata":{},"outputs":[],"execution_count":null},{"id":"39ada5c6-fb0e-444a-a640-5085d29aa1f3","cell_type":"markdown","source":"## 9b. Pseudo-labelling the unlabelled APTOS pool\n\nRuns automatically before the first student experiment. Predicts the 1,928 unlabelled\nimages with flip TTA, keeps only the confident ones, and reports how many survive and how\nthey are split. The 5-class head supplies a grade for each kept image so the auxiliary\nloss stays meaningful.\n\n**Sanity check to read in the output:** the kept fraction should be roughly 80-95%, and\nthe DR / no-DR balance should be in the same ballpark as APTOS itself (about 50/50). A\nwildly skewed split means the teacher is over-predicting one class and the pseudo-labels\nwould poison the student rather than help it.","metadata":{}},{"id":"c80692d9-14d4-46e4-bd78-c74dc03ba37f","cell_type":"code","source":"class _PathDS(Dataset):\n    def __init__(self, paths, preproc, size):\n        self.paths, self.preproc = paths, preproc\n        self.tf = build_transforms(False, size)\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, i):\n        img = cv2.cvtColor(cv2.imread(cache_name(self.paths[i]), cv2.IMREAD_COLOR),\n                           cv2.COLOR_BGR2RGB)\n        return self.tf(image=apply_preproc(img, self.preproc))[\"image\"], i\n\n\ndef generate_pseudo_labels(teacher_name=PSEUDO_TEACHER, preproc=\"ben\", arch=\"dual\",\n                           size=None, conf=PSEUDO_CONF):\n    global PSEUDO_DF\n    if not PSEUDO_ENABLED or not UNLABELLED:\n        print(\"pseudo-labelling disabled or no unlabelled images found\")\n        return None\n    ck = resolve_ckpt(os.path.join(CKPT_DIR, teacher_name))\n    if ck is None:\n        print(f\"teacher '{teacher_name}' not available yet - run the teacher experiment first\")\n        return None\n\n    size = size or IMG_SIZE\n    paths = [p for p in UNLABELLED if os.path.exists(cache_name(p))]\n    missing = [p for p in UNLABELLED if not os.path.exists(cache_name(p))]\n    if missing:\n        print(f\"caching {len(missing)} unlabelled images first\")\n        build_cache(missing)\n        paths = [p for p in UNLABELLED if os.path.exists(cache_name(p))]\n\n    model = DRNet(arch).to(DEVICE).to(memory_format=torch.channels_last)\n    model.load_state_dict(torch.load(ck, map_location=DEVICE))\n    model.eval()\n    ld = DataLoader(_PathDS(paths, preproc, size), BATCH_SIZE * 2, shuffle=False,\n                    num_workers=NUM_WORKERS, pin_memory=True)\n\n    probs, grades = [], []\n    t0 = time.time()\n    with torch.no_grad():\n        for x, _ in ld:\n            x = x.to(DEVICE, non_blocking=True).contiguous(memory_format=torch.channels_last)\n            views = [x, torch.flip(x, [3]), torch.flip(x, [2]), torch.flip(x, [2, 3])]\n            pb, gl = 0, 0\n            with amp_autocast():\n                for v in views:\n                    lo, _, gg = model(v)\n                    pb = pb + torch.sigmoid(lo.float())\n                    if gg is not None:\n                        gl = gl + torch.softmax(gg.float(), 1)\n            probs.append((pb / len(views)).cpu().numpy())\n            grades.append((gl / len(views)).argmax(1).cpu().numpy()\n                          if torch.is_tensor(gl) else np.zeros(x.size(0), int))\n    del model\n    torch.cuda.empty_cache()\n    probs, grades = np.concatenate(probs), np.concatenate(grades)\n\n    keep = (probs >= conf) | (probs <= 1 - conf)\n    lab = (probs[keep] >= 0.5).astype(int)\n    gr = grades[keep]\n    gr = np.where(lab == 0, 0, np.clip(np.where(gr == 0, 2, gr), 1, 4))   # keep grade consistent\n    PSEUDO_DF = pd.DataFrame({\"image_path\": np.array(paths)[keep], \"grade\": gr,\n                              \"source\": \"pseudo\", \"split\": \"train\", \"label\": lab})\n    PSEUDO_DF.to_csv(os.path.join(OUT_DIR, \"pseudo_labels.csv\"), index=False)\n\n    print(f\"\\npseudo-labelled {len(paths)} unlabelled images in {time.time()-t0:.0f}s\")\n    print(f\"  kept {keep.sum()} ({100*keep.mean():.1f}%) at confidence >= {conf}\")\n    print(f\"  of those: {(lab==0).sum()} no-DR / {(lab==1).sum()} DR \"\n          f\"({100*lab.mean():.1f}% positive)\")\n    print(f\"  discarded {int((~keep).sum())} uncertain images\")\n    if keep.mean() < 0.5:\n        print(\"  [warn] fewer than half kept - teacher is unsure; consider lowering \"\n              \"PSEUDO_CONF or using a better teacher\")\n    if not 0.25 < lab.mean() < 0.75:\n        print(\"  [warn] pseudo-label balance is skewed - the teacher may be biased; \"\n              \"check before training a student on this\")\n    return PSEUDO_DF","metadata":{},"outputs":[],"execution_count":null},{"id":"6761196b-bb13-4169-81b5-31ae88ebecba","cell_type":"markdown","source":"## 10. Full evaluation of one experiment\n\nEverything below is tabular. Four inference variants (`raw` / `ema` x TTA off / on) are scored on\nvalidation and test. The reported variant is the one with the best **validation** AUC, and its\nthreshold is tuned on **validation** only.\n\nThe oracle-threshold row is printed as a diagnostic ceiling — the gap between it and the\nval-tuned row tells you how much of your error is threshold placement rather than the model.\n**Never report the oracle row.**","metadata":{}},{"id":"56a91055-18de-4162-95c6-1e97005a178f","cell_type":"code","source":"PAPER = {\"accuracy\": 98.50, \"sensitivity\": 99.46, \"specificity\": 97.51, \"precision\": 97.61}\n\n\ndef evaluate_experiment(run):\n    cfg = run[\"cfg\"]\n    val_ld, test_ld = run[\"loaders\"]\n    cache, val_rows, test_rows = {}, {}, {}\n\n    def score(key, vp, vy, vg, tp_, ty, tg):\n        thr = best_threshold(vy, vp)\n        vm = binary_metrics(vy, vp, thr)\n        if SELECT_MODE == \"val_auc\":\n            sel = vm[\"auc\"]\n        elif SELECT_MODE == \"val_acc\":\n            sel = vm[\"accuracy\"]\n        else:                                   # blend: AUC ranks, accuracy is what we report\n            sel = 0.5 * vm[\"auc\"] + 0.5 * vm[\"accuracy\"]\n        cache[key] = dict(val=(vp, vy, vg), test=(tp_, ty, tg), thr=thr,\n                          val_auc=vm[\"auc\"], sel=sel)\n        val_rows[key] = vm\n        test_rows[key] = binary_metrics(ty, tp_, thr)\n\n    for wname, ck in [(\"raw\", run[\"raw_ck\"]), (\"ema\", run[\"ema_ck\"])]:\n        m = DRNet(cfg[\"arch\"]).to(DEVICE).to(memory_format=torch.channels_last)\n        m.load_state_dict(torch.load(ck, map_location=DEVICE))\n        for tta in [False, True]:\n            vp, vy, vg, vpg = predict(m, val_ld, tta)\n            tp_, ty, tg, tpg = predict(m, test_ld, tta)\n            tag = f\"{wname}{'+tta' if tta else ''}\"\n            score(tag, vp, vy, vg, tp_, ty, tg)\n            # Fusing the binary head with 1 - P(grade 0) costs nothing: both come out of\n            # the SAME forward pass, and the two heads make partly different mistakes.\n            if GRADE_FUSE and GRADE_AUX_W > 0 and vpg.any():\n                w = best_fuse_weight(vy, vp, vpg)          # swept on validation only\n                score(f\"{tag}+fused(w{w:.2f})\", (1 - w) * vp + w * vpg, vy, vg,\n                      (1 - w) * tp_ + w * tpg, ty, tg)\n                score(f\"{tag}+gradeonly\", vpg, vy, vg, tpg, ty, tg)\n        del m\n        torch.cuda.empty_cache()\n\n    # ---- snapshot ensemble: the payoff from cosine warm restarts ----------------------\n    # Each cycle ends in a different low-LR basin, so the snapshots make different mistakes.\n    # Averaging them approximates a multi-seed ensemble at the cost of a single run.\n    snaps = run.get(\"snaps\", [])\n    if len(snaps) >= 2:\n        sv, st = [], []\n        for sp in snaps:\n            m = DRNet(cfg[\"arch\"]).to(DEVICE).to(memory_format=torch.channels_last)\n            m.load_state_dict(torch.load(sp, map_location=DEVICE))\n            vpi, vy, vg, _ = predict(m, val_ld, True)\n            tpi, ty, tg, _ = predict(m, test_ld, True)\n            sv.append(vpi); st.append(tpi)\n            score(f\"snap{os.path.basename(sp)[5]}+tta\", vpi, vy, vg, tpi, ty, tg)\n            del m\n            torch.cuda.empty_cache()\n        # Only average snapshots that are actually competitive. Blindly averaging cost 0.4\n        # points last time, because the first-cycle snapshot was far weaker than the rest.\n        aucs = [cache[f\"snap{os.path.basename(sp)[5]}+tta\"][\"val_auc\"] for sp in snaps]\n        keep = [i for i, a in enumerate(aucs) if a >= max(aucs) - 0.75]\n        if len(keep) >= 2:\n            score(\"SNAPSHOT-ENS+tta\", np.mean([sv[i] for i in keep], 0), vy, vg,\n                  np.mean([st[i] for i in keep], 0), ty, tg)\n        else:\n            print(f\"  snapshot ensemble skipped: only {len(keep)} snapshot(s) within \"\n                  f\"0.75 val-AUC of the best\")\n\n    # Average the best few variants instead of betting everything on one. Validation gaps\n    # between variants are smaller than the noise on a ~366-image set, so single-variant\n    # selection is close to a coin flip.\n    if USE_TOP3_AVG and len(cache) >= 3:\n        top3 = sorted(cache, key=lambda k: -cache[k][\"sel\"])[:3]\n        score(\"TOP3-AVG\",\n              np.mean([cache[k][\"val\"][0] for k in top3], 0), vy, vg,\n              np.mean([cache[k][\"test\"][0] for k in top3], 0), ty, tg)\n        print(f\"  TOP3-AVG built from: {', '.join(top3)}\")\n\n    best_key = max(cache, key=lambda k: cache[k][\"sel\"])\n    tp_, ty, tg = cache[best_key][\"test\"]\n    thr = cache[best_key][\"thr\"]\n\n    print(\"\\n\" + \"=\" * 96)\n    print(f\"EXPERIMENT  {cfg['name']}\")\n    print(f\"  data={cfg['data']}  preproc={cfg['preproc']}  arch={cfg['arch']}  \"\n          f\"balance={cfg['balance']}  seed={cfg['seed']}\")\n    print(f\"  best epoch {run['best_epoch']}   best val AUC {run['best_val_auc']:.2f}   \"\n          f\"final input {run.get('img_size', IMG_SIZE)}px   snapshots {len(run.get('snaps', []))}\")\n    print(f\"  selected variant '{best_key}' by {SELECT_MODE}   threshold {thr:.3f} \"\n          f\"({THRESH_OBJECTIVE}, tuned on val)\")\n    print(\"=\" * 96)\n\n    show_table(\"TABLE 1 - validation, all four inference variants (each at its own val-tuned thr)\",\n               metrics_table(val_rows))\n    show_table(\"TABLE 2 - test, all four inference variants (val-tuned thr, honest numbers)\",\n               metrics_table(test_rows),\n               note=f\"variant '{best_key}' is the reported one (highest val AUC)\")\n\n    sweep = {\"@0.50\": binary_metrics(ty, tp_, 0.5),\n             f\"@val-tuned ({thr:.3f})\": binary_metrics(ty, tp_, thr),\n             \"@oracle (ceiling)\": binary_metrics(ty, tp_, best_threshold(ty, tp_))}\n    show_table(f\"TABLE 3 - threshold sensitivity on test, variant '{best_key}'\",\n               metrics_table(sweep),\n               note=\"oracle is tuned ON TEST - a ceiling for diagnosis, never a reported result\")\n\n    head = test_rows[best_key]\n    at_half = binary_metrics(ty, tp_, 0.5)\n    cmp_rows = {f\"this work @ val-tuned {thr:.3f}\": head,\n                \"this work @ fixed 0.50\": at_half}\n\n    # --- sensitivity-matched comparison ---------------------------------------------\n    # A screening tool is judged at a clinically useful recall, not at whatever threshold\n    # happens to maximise accuracy. With a higher AUC than the baseline, matching their\n    # sensitivity should still leave better specificity - that is the fair comparison.\n    grid = np.round(np.arange(0.005, 0.995, 0.002), 3)\n    sens = np.array([recall_score(ty, (tp_ >= t).astype(int), zero_division=0) for t in grid])\n    for target, tag in [(0.9946, \"sens matched to paper (99.46%)\"), (0.98, \"sens >= 98%\")]:\n        ok = grid[sens >= target]\n        if len(ok):\n            cmp_rows[f\"this work @ {tag}\"] = binary_metrics(ty, tp_, float(ok.max()))\n        else:\n            print(f\"  [note] sensitivity {target*100:.2f}% unreachable on this test set\")\n    cmp_rows[\"PAPER (ImageNet transfer)\"] = {**PAPER, \"auc\": 98.00,\n                                             **{k: np.nan for k in\n                                                [\"npv\", \"f1\", \"balanced_acc\", \"pr_auc\", \"kappa\",\n                                                 \"mcc\", \"brier\", \"threshold\", \"tp\", \"tn\", \"fp\", \"fn\"]}}\n    show_table(\"TABLE 4 - headline vs the published baseline, at several operating points\",\n               metrics_table(cmp_rows),\n               note=\"The paper starts from pretrained ResNet50 + EfficientNetB0; this model starts \"\n                    \"from random init. Report ONE operating point and say which. For a screening \"\n                    \"application the sensitivity-matched row is the fair comparison, because a \"\n                    \"missed case costs far more than a false alarm.\")\n\n    pred = (tp_ >= thr).astype(int)\n    rows = []\n    for g in sorted(np.unique(tg)):\n        msk = tg == g\n        ok = int((pred[msk] == ty[msk]).sum())\n        rows.append({\"ICDR grade\": int(g),\n                     \"Stage\": [\"No DR\", \"Mild NPDR\", \"Moderate NPDR\", \"Severe NPDR\", \"PDR\"][int(g)],\n                     \"Images in test\": int(msk.sum()), \"Correct\": ok,\n                     \"Errors\": int(msk.sum()) - ok,\n                     \"Accuracy (%)\": 100.0 * ok / max(1, msk.sum()),\n                     \"Mean predicted P(DR)\": float(tp_[msk].mean())})\n    grade_df = pd.DataFrame(rows).set_index(\"ICDR grade\")\n    show_table(\"TABLE 5 - per-ICDR-grade behaviour on the binary test set\", grade_df,\n               note=\"grade 0 errors are false positives; grades 1-4 errors are missed disease\")\n\n    tr_df = run[\"frames\"][0]\n    prev_tr, prev_te = float(tr_df.label.mean()), float((ty == 1).mean())\n    print(f\"\\n  prevalence: train {prev_tr:.3f} positive vs test {prev_te:.3f} \"\n          f\"(gap {prev_tr - prev_te:+.3f}) | chosen threshold {thr:.3f}\")\n    print(\"  a large positive gap pushes the tuned threshold above 0.5 and costs sensitivity\")\n\n    h = pd.DataFrame(run[\"hist\"])\n    marks = [c for c in [1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 60, 80, 100] if c <= len(h)]\n    if len(h) not in marks:\n        marks.append(len(h))\n    traj = h[h.epoch.isin(marks)].set_index(\"epoch\")[\n        [\"train_loss\", \"train_acc\", \"val_loss\", \"val_acc_raw\",\n         \"val_auc_raw\", \"val_acc_ema\", \"val_auc_ema\"]]\n    traj.columns = [\"Train loss\", \"Train acc (%)\", \"Val loss\", \"Val acc raw (%)\",\n                    \"Val AUC raw (%)\", \"Val acc EMA (%)\", \"Val AUC EMA (%)\"]\n    traj.index.name = \"Epoch\"\n    show_table(\"TABLE 6 - training trajectory (sampled epochs)\", traj.T,\n               note=\"train accuracy well below 100 means the model is still underfit\")\n\n    # ---------------- figures ----------------\n    fig, ax = plt.subplots(1, 3, figsize=(16, 4.4))\n    cm = confusion_matrix(ty, pred, labels=[0, 1])\n    ax[0].imshow(cm, cmap=\"Blues\")\n    for i in range(2):\n        for j in range(2):\n            ax[0].text(j, i, cm[i, j], ha=\"center\", va=\"center\", fontsize=15,\n                       color=\"white\" if cm[i, j] > cm.max() / 2 else \"black\")\n    ax[0].set_xticks([0, 1], [\"No DR\", \"DR\"]); ax[0].set_yticks([0, 1], [\"No DR\", \"DR\"])\n    ax[0].set_xlabel(\"Predicted\"); ax[0].set_ylabel(\"Actual\")\n    ax[0].set_title(f\"Confusion matrix (test) - acc {head['accuracy']:.2f}%\")\n    fpr, tpr, _ = roc_curve(ty, tp_)\n    ax[1].plot(fpr, tpr, label=f\"AUC={head['auc']:.2f}\")\n    ax[1].plot([0, 1], [0, 1], \"--\", c=\"gray\"); ax[1].legend()\n    ax[1].set_title(\"ROC (test)\"); ax[1].set_xlabel(\"FPR\"); ax[1].set_ylabel(\"TPR\")\n    pr, rc, _ = precision_recall_curve(ty, tp_)\n    ax[2].plot(rc, pr, label=f\"AP={head['pr_auc']:.2f}\"); ax[2].legend()\n    ax[2].set_title(\"Precision-Recall (test)\"); ax[2].set_xlabel(\"Recall\"); ax[2].set_ylabel(\"Precision\")\n    plt.suptitle(cfg[\"name\"]); plt.tight_layout()\n    plt.savefig(os.path.join(OUT_DIR, f\"curves_{cfg['name']}.png\"), dpi=120); plt.show()\n\n    fig, ax = plt.subplots(1, 3, figsize=(16, 3.8))\n    ax[0].plot(h.epoch, h.train_loss, label=\"train\")\n    ax[0].plot(h.epoch, h.val_loss, label=\"val\")\n    ax[0].set_title(\"loss\"); ax[0].set_xlabel(\"epoch\"); ax[0].legend()\n    ax[1].plot(h.epoch, h.train_acc, label=\"train\")\n    ax[1].plot(h.epoch, h.val_acc_raw, label=\"val raw\")\n    ax[1].plot(h.epoch, h.val_acc_ema, label=\"val ema\")\n    ax[1].set_title(\"accuracy @0.5 (%)\"); ax[1].set_xlabel(\"epoch\"); ax[1].legend()\n    ax[2].plot(h.epoch, h.val_auc_raw, label=\"raw\")\n    ax[2].plot(h.epoch, h.val_auc_ema, label=\"ema\")\n    ax[2].axvline(run[\"best_epoch\"], color=\"crimson\", ls=\"--\", lw=1, label=\"checkpoint\")\n    ax[2].set_title(\"validation AUC (%)\"); ax[2].set_xlabel(\"epoch\"); ax[2].legend()\n    plt.suptitle(f\"{cfg['name']} - training trajectory\"); plt.tight_layout()\n    plt.savefig(os.path.join(OUT_DIR, f\"train_{cfg['name']}.png\"), dpi=120); plt.show()\n\n    # ---------------- persist ----------------\n    np.savez(os.path.join(OUT_DIR, f\"preds_{cfg['name']}.npz\"),\n             val_prob=cache[best_key][\"val\"][0], val_y=cache[best_key][\"val\"][1],\n             test_prob=tp_, test_y=ty, test_grade=tg,\n             val_thr=thr, best_variant=best_key, val_auc=cache[best_key][\"val_auc\"])\n    h.to_csv(os.path.join(OUT_DIR, f\"history_{cfg['name']}.csv\"), index=False)\n    grade_df.to_csv(os.path.join(OUT_DIR, f\"grade_{cfg['name']}.csv\"))\n    metrics_table(test_rows).to_csv(os.path.join(OUT_DIR, f\"variants_{cfg['name']}.csv\"))\n    with open(os.path.join(OUT_DIR, f\"metrics_{cfg['name']}.json\"), \"w\") as f:\n        json.dump({\"cfg\": cfg, \"best_epoch\": run[\"best_epoch\"], \"best_variant\": best_key,\n                   \"val\": val_rows, \"test\": test_rows}, f, indent=2, default=float)\n\n    row = {\"experiment\": cfg[\"name\"], **{k: cfg[k] for k in\n           [\"data\", \"preproc\", \"arch\", \"balance\", \"seed\"]},\n           \"img_size\": cfg.get(\"img_size\", IMG_SIZE),\n           \"fold\": cfg.get(\"fold\", -1),\n           \"init_from\": os.path.basename(cfg.get(\"init_from\", \"\") or \"scratch\"),\n           \"variant\": best_key, \"epochs_run\": len(h), \"best_epoch\": run[\"best_epoch\"],\n           \"val_auc\": cache[best_key][\"val_auc\"],\n           \"train_loss\": h.train_loss.iloc[-1], \"train_acc\": h.train_acc.iloc[-1],\n           \"val_loss\": h.val_loss.iloc[-1],\n           **{f\"test_{k}\": v for k, v in head.items()}}\n    res_path = os.path.join(OUT_DIR, \"results.csv\")\n    old = pd.read_csv(res_path) if os.path.exists(res_path) else pd.DataFrame()\n    pd.concat([old, pd.DataFrame([row])], ignore_index=True) \\\n      .drop_duplicates(subset=[\"experiment\"], keep=\"last\").to_csv(res_path, index=False)\n    return dict(val=val_rows, test=test_rows, grades=grade_df, best=best_key)","metadata":{},"outputs":[],"execution_count":null},{"id":"d90ee38f-03e0-4585-b32a-10a15448b610","cell_type":"markdown","source":"## 11. Run\n\nThe session budget is divided evenly across the experiments queued in `RUN_EXPERIMENTS`, with\n`EVAL_RESERVE_MIN` held back per experiment for TTA and the tables. Training stops an epoch\nearly rather than starting one that would not fit, so **the notebook cannot overrun the 12 h\nlimit**. Anything cut short resumes from its checkpoint when you re-run this cell.","metadata":{}},{"id":"9b019d1e-2682-4f67-aa7d-6cb902cb528d","cell_type":"code","source":"queue = list(RUN_EXPERIMENTS)\nfor n in queue:\n    assert n in EXP_BY_NAME, f\"unknown experiment {n}; options: {list(EXP_BY_NAME)}\"\nprint(f\"queued: {queue}\\nsession budget {SESSION_MAX_HOURS:.1f} h \"\n      f\"({session_left_h():.2f} h left), {EVAL_RESERVE_MIN} min reserved per experiment\\n\")\n\ncompleted = {}\nfor i, name in enumerate(queue):\n    remaining = len(queue) - i\n    left_h = session_left_h()\n    share_h = left_h / remaining - EVAL_RESERVE_MIN / 60.0\n    if share_h < 0.15:\n        print(f\"\\nSKIPPING {name}: only {left_h:.2f} h left for {remaining} experiment(s). \"\n              f\"Re-run the notebook in a fresh session - RESUME=True will continue it.\")\n        continue\n\n    cfg = EXP_BY_NAME[name]\n    if cfg.get(\"pseudo\") and PSEUDO_DF is None:\n        print(f\"\\n{'='*96}\\n# PSEUDO-LABELLING before {name}\\n{'='*96}\")\n        generate_pseudo_labels(preproc=cfg[\"preproc\"], arch=cfg[\"arch\"],\n                               size=cfg.get(\"img_size\", IMG_SIZE))\n        if PSEUDO_DF is None:\n            print(f\"SKIPPING {name}: no pseudo-labels available \"\n                  f\"(the teacher '{PSEUDO_TEACHER}' has not been trained yet)\")\n            continue\n    if cfg.get(\"init_from\") and not resolve_ckpt(cfg[\"init_from\"]):\n        print(f\"\\nSKIPPING {name}: warm-start checkpoint \"\n              f\"'{os.path.basename(cfg['init_from'])}' is not available in this session.\\n\"\n              f\"  Attach the previous run's output as an input dataset, or run the \"\n              f\"from-scratch experiments instead.\")\n        continue\n    print(f\"\\n{'#'*96}\\n# TRAINING {name}  ({i+1}/{len(queue)})\\n{'#'*96}\")\n    print(json.dumps({k: str(v) for k, v in cfg.items()}, indent=2))\n    t0 = time.time()\n    run = train_one(cfg, deadline_ts=time.time() + share_h * 3600)\n    print(f\"trained in {(time.time()-t0)/60:.1f} min | session used \"\n          f\"{SESSION_MAX_HOURS - session_left_h():.2f} h of {SESSION_MAX_HOURS:.1f} h\")\n    completed[name] = evaluate_experiment(run)\n    del run\n    torch.cuda.empty_cache()\n\nprint(f\"\\nDONE. {len(completed)}/{len(queue)} finished. \"\n      f\"Session used {SESSION_MAX_HOURS - session_left_h():.2f} h.\")\nif len(completed) < len(queue):\n    print(\"Re-run in a new session to finish the queue (checkpoints are in /kaggle/working/ckpt).\")","metadata":{},"outputs":[],"execution_count":null},{"id":"1a30b98f-8d63-4bb3-b3e0-c3ce1dda8411","cell_type":"markdown","source":"## 12. Comparison across everything tried\n\n`results.csv` accumulates across sessions. To merge runs from several Kaggle accounts: download\neach account's `results.csv` and `preds_*.npz`, upload them as one dataset, and copy them into\n`/kaggle/working` before running this cell.","metadata":{}},{"id":"8082e20b-bab4-4a17-ae38-75553126328c","cell_type":"code","source":"res_path = os.path.join(OUT_DIR, \"results.csv\")\nif os.path.exists(res_path):\n    res = pd.read_csv(res_path).sort_values(\"test_accuracy\", ascending=False)\n\n    setup = [\"experiment\", \"data\", \"preproc\", \"arch\", \"balance\", \"seed\",\n             \"variant\", \"epochs_run\", \"best_epoch\"]\n    show_table(\"TABLE A - experiment setups\",\n               res[[c for c in setup if c in res.columns]].set_index(\"experiment\"))\n\n    perf = {\"test_accuracy\": \"Accuracy (%)\", \"test_sensitivity\": \"Sensitivity (%)\",\n            \"test_specificity\": \"Specificity (%)\", \"test_precision\": \"Precision (%)\",\n            \"test_f1\": \"F1 (%)\", \"test_auc\": \"ROC-AUC (%)\"}\n    have = {k: v for k, v in perf.items() if k in res.columns}\n    perf_df = res[[\"experiment\"] + list(have)].set_index(\"experiment\")\n    perf_df.columns = list(have.values())\n    perf_df.loc[\"PAPER (transfer learning)\"] = [\n        PAPER.get(k.split()[0].lower().replace(\"roc-auc\", \"auc\"), np.nan) for k in have.values()]\n    perf_df.loc[\"PAPER (transfer learning)\", \"Accuracy (%)\"] = PAPER[\"accuracy\"]\n    perf_df.loc[\"PAPER (transfer learning)\", \"Sensitivity (%)\"] = PAPER[\"sensitivity\"]\n    perf_df.loc[\"PAPER (transfer learning)\", \"Specificity (%)\"] = PAPER[\"specificity\"]\n    perf_df.loc[\"PAPER (transfer learning)\", \"Precision (%)\"] = PAPER[\"precision\"]\n    if \"ROC-AUC (%)\" in perf_df.columns:\n        perf_df.loc[\"PAPER (transfer learning)\", \"ROC-AUC (%)\"] = 98.00\n    show_table(\"TABLE B - test performance, ranked (paper baseline appended)\", perf_df,\n               note=\"all rows use a validation-tuned threshold; the paper row is quoted from \"\n                    \"Table 4 of the article\")\n\n    fit = [\"train_loss\", \"train_acc\", \"val_loss\", \"val_auc\"]\n    fit_df = res[[\"experiment\"] + [c for c in fit if c in res.columns]].set_index(\"experiment\")\n    fit_df.columns = [\"Train loss\", \"Train acc (%)\", \"Val loss\", \"Val AUC (%)\"][:len(fit_df.columns)]\n    show_table(\"TABLE C - fit diagnostics (final epoch)\", fit_df,\n               note=\"train acc near 100 with val AUC below its peak = overfitting; \"\n                    \"train acc below ~98 = still underfit, give it more epochs\")\n\n    if len(res) > 1:\n        fig, ax = plt.subplots(figsize=(11, 0.45 * len(res) + 2.2))\n        y = np.arange(len(res))\n        ax.barh(y, res.test_accuracy, color=\"#4C72B0\")\n        ax.set_yticks(y, res.experiment)\n        ax.axvline(PAPER[\"accuracy\"], color=\"crimson\", ls=\"--\",\n                   label=f\"paper (transfer learning) {PAPER['accuracy']}\")\n        for i, (a, u) in enumerate(zip(res.test_accuracy, res.test_auc)):\n            ax.text(a + 0.1, i, f\"{a:.2f} (AUC {u:.1f})\", va=\"center\", fontsize=8)\n        ax.set_xlim(min(80, res.test_accuracy.min() - 3), 101)\n        ax.set_xlabel(\"test accuracy (%)\"); ax.legend(loc=\"lower right\"); ax.invert_yaxis()\n        plt.tight_layout(); plt.savefig(os.path.join(OUT_DIR, \"comparison.png\"), dpi=120); plt.show()\nelse:\n    print(\"no results yet - run at least one experiment\")","metadata":{},"outputs":[],"execution_count":null},{"id":"6789d541-a59b-40f8-9822-d105939ff925","cell_type":"markdown","source":"## 12b. Seed variance\n\nA single run on a 733-image test set carries roughly +/- 0.7% of sampling noise, so one number\nis not a result. Repeat runs of the same config with different seeds and report **mean +/- std**.\nThis is also the table an examiner asks for when you claim to beat or approach a baseline.","metadata":{}},{"id":"33bed7d0-ab53-46e6-8e18-d4a8ac3ec7fe","cell_type":"code","source":"if os.path.exists(res_path):\n    res = pd.read_csv(res_path)\n    grp = [\"data\", \"preproc\", \"arch\", \"balance\"]\n    met = [\"test_accuracy\", \"test_sensitivity\", \"test_specificity\", \"test_auc\", \"test_f1\"]\n    met = [m for m in met if m in res.columns]\n    agg = res.groupby(grp)[met].agg([\"mean\", \"std\", \"count\"])\n    multi = agg[(met[0], \"count\")] > 1\n    if multi.any():\n        out = pd.DataFrame(index=agg.index[multi])\n        for m in met:\n            mu, sd = agg.loc[multi, (m, \"mean\")], agg.loc[multi, (m, \"std\")].fillna(0)\n            out[m.replace(\"test_\", \"\")] = [f\"{a:.2f} +/- {b:.2f}\" for a, b in zip(mu, sd)]\n        out[\"n_seeds\"] = agg.loc[multi, (met[0], \"count\")].astype(int)\n        print(\"\\nSEED VARIANCE - configs with more than one seed\")\n        print(\"-\" * 46)\n        print(out.to_string())\n        print(\"\\n  quote the mean +/- std in the thesis, not the best single seed\")\n    else:\n        print(\"only one seed per config so far - run E10 and E11 for a variance estimate\")\nelse:\n    print(\"no results yet\")","metadata":{},"outputs":[],"execution_count":null},{"id":"54d3192e-f657-4172-97d0-5df8edd90482","cell_type":"markdown","source":"## 13. Ensemble of every experiment present\n\nAll experiments share one fixed APTOS test set, so their probability vectors are directly\naverageable. The ensemble threshold is tuned on validation only.","metadata":{}},{"id":"a8d3bcac-ce9c-457b-87d5-c70e50cd2f4b","cell_type":"code","source":"files = sorted(glob.glob(os.path.join(OUT_DIR, \"preds_*.npz\")))\nprint(f\"{len(files)} prediction files found\")\nif len(files) >= 2:\n    vps, tps, ws, names = [], [], [], []\n    vy = ty = tg = None\n    for f in files:\n        d = np.load(f, allow_pickle=True)\n        vps.append(d[\"val_prob\"]); tps.append(d[\"test_prob\"])\n        ws.append(float(d[\"val_auc\"])); names.append(os.path.basename(f)[6:-4])\n        vy, ty, tg = d[\"val_y\"], d[\"test_y\"], d[\"test_grade\"]\n\n    rows = {}\n    for f, tp_i in zip(names, tps):\n        d = np.load(os.path.join(OUT_DIR, f\"preds_{f}.npz\"), allow_pickle=True)\n        rows[f] = binary_metrics(ty, tp_i, float(d[\"val_thr\"]))\n\n    for label, w in [(\"ensemble (uniform)\", np.ones(len(files))),\n                     (\"ensemble (val-AUC weighted)\", np.array(ws) - min(ws) + 1e-6)]:\n        w = w / w.sum()\n        vp = np.average(np.stack(vps), axis=0, weights=w)\n        tp_ = np.average(np.stack(tps), axis=0, weights=w)\n        thr = best_threshold(vy, vp)\n        rows[label] = binary_metrics(ty, tp_, thr)\n        np.savez(os.path.join(OUT_DIR, f\"{label.split()[1].strip('()')}_ensemble.npz\"),\n                 test_prob=tp_, test_y=ty, test_grade=tg, members=np.array(names))\n\n    show_table(\"TABLE D - members vs ensembles (test, val-tuned thresholds)\",\n               metrics_table(rows),\n               note=\"an ensemble that does not beat its best member usually means the members are too similar\")\nelse:\n    print(\"run at least two experiments to ensemble\")","metadata":{},"outputs":[],"execution_count":null},{"id":"a5852a03-3c79-4406-8697-6b41a008eb71","cell_type":"markdown","source":"## 13c. Out-of-fold ensemble\n\n**This is the one change with a solid reason to expect a gain, and it costs no extra training.**\n\nEvery threshold so far was tuned on 366 validation images. That is small enough that the\ntuned threshold has repeatedly been *worse* on test than a plain 0.50 - you saw it directly\n(Table 3 of an earlier run: 96.86 at 0.50 vs 97.00 at a tuned 0.637, and elsewhere the\nreverse). It is noise-fitting either way.\n\nThe K folds have **disjoint** validation slices, so stacking them gives one out-of-fold set\nseveral times larger. A threshold estimated there is far more stable, and it is still\nstrictly out-of-sample: no test image is involved at any point.","metadata":{}},{"id":"5011d9f2-c42f-450f-801c-22858841b81d","cell_type":"code","source":"def oof_ensemble():\n    files = sorted(glob.glob(os.path.join(OUT_DIR, \"preds_K*.npz\")))\n    if len(files) < 2:\n        print(f\"need at least 2 fold runs, found {len(files)}\")\n        return\n    vps, vys, tps = [], [], []\n    ty = tg = None\n    for f in files:\n        d = np.load(f, allow_pickle=True)\n        vps.append(d[\"val_prob\"]); vys.append(d[\"val_y\"]); tps.append(d[\"test_prob\"])\n        ty, tg = d[\"test_y\"], d[\"test_grade\"]\n\n    oof_p, oof_y = np.concatenate(vps), np.concatenate(vys)\n    thr_oof = best_threshold(oof_y, oof_p)\n    test_ens = np.mean(tps, 0)\n\n    rows = {}\n    for f, tp_i in zip(files, tps):\n        d = np.load(f, allow_pickle=True)\n        rows[os.path.basename(f)[6:-4]] = binary_metrics(ty, tp_i, float(d[\"val_thr\"]))\n    rows[\"ENSEMBLE @ each fold's own thr\"] = binary_metrics(\n        ty, test_ens, float(np.mean([np.load(f, allow_pickle=True)[\"val_thr\"] for f in files])))\n    rows[\"ENSEMBLE @ 0.50\"] = binary_metrics(ty, test_ens, 0.5)\n    rows[f\"ENSEMBLE @ OOF thr ({thr_oof:.3f})\"] = binary_metrics(ty, test_ens, thr_oof)\n\n    show_table(f\"OUT-OF-FOLD ENSEMBLE - threshold from {len(oof_y)} pooled OOF images \"\n               f\"(vs 366 before)\", metrics_table(rows),\n               note=\"report the OOF-threshold row: the largest honest validation signal \"\n                    \"available, and no test image was used to choose it\")\n\n    best = max(rows, key=lambda k: rows[k][\"accuracy\"])\n    print(f\"\\n  best row here: {best} at {rows[best]['accuracy']:.2f}% accuracy\")\n    np.savez(os.path.join(OUT_DIR, \"oof_ensemble.npz\"), test_prob=test_ens, test_y=ty,\n             test_grade=tg, oof_thr=thr_oof, members=np.array([os.path.basename(f) for f in files]))\n    res = pd.read_csv(res_path) if os.path.exists(res_path) else pd.DataFrame()\n    row = {\"experiment\": \"OOF-ENSEMBLE\", \"data\": \"selective\", \"preproc\": \"ben\", \"arch\": \"dual\",\n           \"balance\": \"none\", \"seed\": -1, \"variant\": \"fold-ensemble\",\n           **{f\"test_{k}\": v for k, v in rows[f\"ENSEMBLE @ OOF thr ({thr_oof:.3f})\"].items()}}\n    pd.concat([res, pd.DataFrame([row])], ignore_index=True) \\\n      .drop_duplicates(\"experiment\", keep=\"last\").to_csv(res_path, index=False)\n\n\noof_ensemble()","metadata":{},"outputs":[],"execution_count":null},{"id":"f2642e22-3c8d-4346-ba65-8f8aef6e3aed","cell_type":"markdown","source":"## 13b. Verdict (automatic)\n\nReads `results.csv` and states, in one line, what the A/B decided and what to do next.\nYou do not need to interpret the tables above to act on the result.","metadata":{}},{"id":"c4fb7f86-0923-4fc2-8fe1-ff8e230f1a8d","cell_type":"code","source":"def verdict():\n    if not os.path.exists(res_path):\n        print(\"no results yet\")\n        return\n    r = pd.read_csv(res_path).drop_duplicates(\"experiment\", keep=\"last\")\n    got = {n: r[r.experiment == n] for n in [\"AB_fold0_selective\", \"AB_fold0_full\"]}\n    if any(len(v) == 0 for v in got.values()):\n        have = [n for n, v in got.items() if len(v)]\n        print(f\"A/B incomplete - have {have or 'neither run'}. \"\n              f\"Re-run the notebook; RESUME=True continues an interrupted run.\")\n        return\n\n    sel = got[\"AB_fold0_selective\"].iloc[0]\n    ful = got[\"AB_fold0_full\"].iloc[0]\n    tbl = pd.DataFrame({\n        \"selective (paper's recipe)\": {\"accuracy\": sel.test_accuracy,\n                                       \"sensitivity\": sel.test_sensitivity,\n                                       \"specificity\": sel.test_specificity,\n                                       \"ROC-AUC\": sel.test_auc,\n                                       \"threshold\": sel.test_threshold},\n        \"full (all external grades)\": {\"accuracy\": ful.test_accuracy,\n                                       \"sensitivity\": ful.test_sensitivity,\n                                       \"specificity\": ful.test_specificity,\n                                       \"ROC-AUC\": ful.test_auc,\n                                       \"threshold\": ful.test_threshold}})\n    tbl[\"difference\"] = tbl.iloc[:, 1] - tbl.iloc[:, 0]\n    show_table(\"VERDICT - selective vs full merge, identical fold and seed\", tbl)\n\n    gap = ful.test_accuracy - sel.test_accuracy\n    NOISE = 0.30          # roughly 2x the +/-0.14 seed std measured on this test set\n    if abs(gap) < NOISE:\n        winner, why = \"selective\", (f\"the {gap:+.2f} difference is inside the noise band \"\n                                    f\"(+/-{NOISE}), so there is no evidence to deviate from \"\n                                    f\"the paper's recipe\")\n    elif gap > 0:\n        winner, why = \"full\", (f\"full is ahead by {gap:+.2f}, beyond the +/-{NOISE} noise \"\n                               f\"band - the paper's 5-class finding does NOT transfer to binary\")\n    else:\n        winner, why = \"selective\", (f\"selective is ahead by {-gap:.2f} - the paper's Table 3 \"\n                                    f\"finding holds for binary too\")\n    print(f\"\\n  WINNER: {winner}  ({why})\")\n    print(f\"\\n  NEXT: set data=\\\"{winner}\\\" on K1_fold1..K4_fold4 in the config cell,\")\n    print(f\"        set RUN_EXPERIMENTS = [\\\"K1_fold1\\\", \\\"K2_fold2\\\", \\\"K3_fold3\\\"],\")\n    print(f\"        and run. Then run section 13 to ensemble every fold.\")\n    print(f\"\\n  Also report this comparison as an ablation row - it is a result either way.\")\n\n\nverdict()","metadata":{},"outputs":[],"execution_count":null},{"id":"5a5dc80f-0691-42bc-b67b-9f3b50351077","cell_type":"markdown","source":"## 13d. Master results table - every method tried\n\nThe complete record, including the things that did **not** work. A thesis that reports its\nfailed variants is far more convincing than one that presents only the winner, and every\nrow below is a real measurement on the same fixed 733-image APTOS test set.","metadata":{}},{"id":"83dc6cea-10ba-494a-9485-79a285c004a0","cell_type":"code","source":"# Measured earlier in this project. Edit only if you re-measure one of them.\nHISTORY = [\n    # method                                        acc     sens    spec    auc   n  note\n    (\"Baseline: 45 epochs, Youden threshold\",       94.41,  91.67,  97.23, 98.11, 1, \"first working run\"),\n    (\"+ accuracy-based threshold\",                  95.23,  95.70,  94.74, 98.32, 1, \"threshold fix\"),\n    (\"+ 100 epochs\",                                96.86,  96.51,  97.23, 99.46, 1, \"was underfit\"),\n    (\"+ 150 epochs\",                                97.00,  96.24,  97.78, 99.46, 1, \"diminishing\"),\n    (\"+ grade head, 110 ep, 3 seeds  [BEST]\",       97.41,  97.31,  97.51, 99.42, 3, \"97.41 +/- 0.14\"),\n    (\"- warm restarts + snapshot ensemble\",         96.73,  np.nan, np.nan, np.nan, 1, \"WORSE than ema+tta\"),\n    (\"- ordinal grade head\",                        96.86,  np.nan, np.nan, np.nan, 1, \"WORSE than softmax\"),\n    (\"- 448 px input\",                              97.00,  np.nan, np.nan, np.nan, 1, \"no gain\"),\n    (\"- ben_clahe preprocessing\",                   97.14,  np.nan, np.nan, np.nan, 1, \"no gain\"),\n    (\"- cross-seed ensemble (3 identical seeds)\",   97.14,  np.nan, np.nan, np.nan, 1, \"below best member\"),\n    (\"- 448 px fine-tune\",                          np.nan, np.nan, np.nan, np.nan, 0, \"invalid: ckpt missing\"),\n    (\"PAPER (ImageNet transfer learning)\",          98.50,  99.46,  97.51, 98.00, 1, \"Shakibania 2024\"),\n]\n\n\ndef master_table():\n    rows = []\n    for name, a, se, sp, au, n, note in HISTORY:\n        rows.append({\"method\": name, \"accuracy\": a, \"sensitivity\": se, \"specificity\": sp,\n                     \"ROC-AUC\": au, \"runs\": n, \"note\": note})\n    if os.path.exists(res_path):\n        r = pd.read_csv(res_path).drop_duplicates(\"experiment\", keep=\"last\")\n        for _, x in r.iterrows():\n            if str(x.experiment).startswith((\"K\", \"OOF\", \"AB\")):\n                rows.append({\"method\": f\"THIS SESSION: {x.experiment}\",\n                             \"accuracy\": x.get(\"test_accuracy\", np.nan),\n                             \"sensitivity\": x.get(\"test_sensitivity\", np.nan),\n                             \"specificity\": x.get(\"test_specificity\", np.nan),\n                             \"ROC-AUC\": x.get(\"test_auc\", np.nan),\n                             \"runs\": 1, \"note\": \"k-fold / ensemble\"})\n    df = pd.DataFrame(rows).set_index(\"method\")\n    show_table(\"MASTER TABLE - every method tried, same 733-image APTOS test set\", df,\n               note=\"'-' prefixed rows were tried and REJECTED on measurement. All from-scratch \"\n                    \"except the final paper row, which uses ImageNet-pretrained backbones.\")\n\n    ours = df.drop(index=[i for i in df.index if \"PAPER\" in i])\n    best = ours.accuracy.idxmax()\n    print(f\"\\n  best from-scratch result: {best}  ->  {ours.accuracy.max():.2f}% accuracy\")\n    print(f\"  best ROC-AUC: {ours['ROC-AUC'].max():.2f} vs the paper's 98.00 \"\n          f\"(AUC is threshold-free, so it is the fairest single comparison)\")\n    df.to_csv(os.path.join(OUT_DIR, \"master_results.csv\"))\n    print(f\"  saved -> {os.path.join(OUT_DIR, 'master_results.csv')}\")\n\n\nmaster_table()","metadata":{},"outputs":[],"execution_count":null},{"id":"ee80406f-d72f-46e1-a0c0-3f9bb1630188","cell_type":"markdown","source":"## 14. Notes for the write-up\n\n**Reporting honestly**\n* The paper's 98.50% uses ImageNet initialisation. A scratch model on the same protocol is a\n  *different* claim: comparable binary screening performance with no external pretraining.\n  Report the gap; it is the interesting result, not a failure.\n* Pick ONE operating point (Table 4 gives both) and say which. On a ~366-image validation set,\n  threshold tuning buys little over a fixed 0.50 and can lose to it.\n* Never report the oracle row of Table 3. It is tuned on test and exists only to show how much\n  of your error is threshold placement.\n* Table 5 is the clinically useful figure. Binary accuracy hides which grades are missed.\n\n**Reading the fit diagnosis printed after training**\n* *overfitting* -> raise `DROP_PATH`, `MIXUP_PROB`, `WEIGHT_DECAY`, or use `data=\"full\"`.\n* *underfitting AND still improving* -> raise `EPOCHS`. This was the 45-epoch verdict\n  (peak at epoch 40, train accuracy 92.8 vs val 92.9 - the model had not finished fitting),\n  which is why the defaults here are 100 epochs with lighter regularisation.\n* *underfitting but plateaued* -> raise capacity or `IMG_SIZE`; lower regularisation further.\n* *converged* -> remaining gains come from `IMG_SIZE = 448` or an ensemble.\n\n**Session management**\n* `RESUME = True` continues from `{name}_last.pt` if a session dies. `MAX_HOURS` stops cleanly\n  before Kaggle's 12 h limit; re-run the same cell to pick up where it left off.\n* Checkpoints live in `/kaggle/working/ckpt` and persist as run output.\n* Strip cell outputs before committing - Kaggle rejects notebooks over 1 MB.\n\n**Suggested order of experiments**\n1. `E3` at the new defaults (100 epochs) - establishes whether budget was the whole story.\n2. `E10`, `E11` (seeds 1337, 2024), then the ensemble cell. Multi-seed averaging is the most\n   reliable remaining gain and also gives you a variance estimate for the thesis.\n3. `E1`/`E2` (`aptos_only`) - tells you whether the preprocessed Messidor-2 mirror helps at all.\n4. `E7`/`E8`/`E9` - the architecture ablation. Required for the write-up regardless of outcome.\n5. `IMG_SIZE = 448`, `BATCH_SIZE = 12` - last, because it is the slowest.\n\n**Caveat to state in the methods section:** the Messidor-2 mirror in use is already preprocessed\nby its publisher, so your pipeline runs on top of someone else's. That is a domain difference\nbetween training images (external, preprocessed) and test images (APTOS, raw).","metadata":{}}]}