{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"ead1b032","cell_type":"markdown","source":"# Robustness and Reliability of Deep Learning Models for Diabetic Retinopathy Grading on Fundus Photographs\n**Self-contained Kaggle notebook (T4 x2).** Dataset: APTOS 2019 Blindness Detection (add via *Add Input*).\n\nThis notebook evaluates two architectures (one CNN, one Transformer) for 5-class DR grading and stress-tests their **reliability**:\n1. **Benchmark** - accuracy, macro-F1, quadratic-weighted kappa (QWK), per-class sensitivity, referable-DR sensitivity/specificity, AUC, confusion matrix (with bootstrap 95% CIs)\n2. **Cross-condition generalisation** - train on high-quality images, test on low-quality (and reverse); generalisation gap\n3. **Robustness to image degradation** - 5 corruptions x 5 severities; accuracy/QWK curves **plus error-direction analysis** (under- vs over-grading) and **referable-DR sensitivity**\n4. **Calibration** - Expected Calibration Error (ECE); reliability diagram\n5. **Explainability** - Grad-CAM overlays + **energy-on-retina ratio** (a computable faithfulness proxy for fundus images)\n\n> **Run order:** keep `QUICK_TEST = True` for the first pass to validate the pipeline end-to-end (a few minutes). Once it completes cleanly, set `QUICK_TEST = False` and *Run All* for the full results.\n","metadata":{}},{"id":"cd33b04c","cell_type":"markdown","source":"## 1. Configuration","metadata":{}},{"id":"1b3cdcbe","cell_type":"code","source":"# ============================ CONFIG ============================\nimport os, random, numpy as np\n\nQUICK_TEST = False          # <-- True: fast pipeline check. Set False for full run.\n\nIMG_SIZE   = 224\nBATCH      = 64            # 32 per GPU across 2x T4 via DataParallel\nEPOCHS     = 10\nLR         = 3e-4\nWEIGHT_DECAY = 1e-4\nSEED       = 0\nNUM_CLASSES = 5            # DR grades 0..4\nMODELS     = [\"efficientnet_b0\", \"vit_small_patch16_224\"]  # CNN + Transformer\nGRADCAM_MODEL = \"efficientnet_b0\"   # Grad-CAM computed on the CNN\n\nCORRUPTIONS = [\"gaussian_blur\", \"brightness\", \"contrast\", \"gaussian_noise\", \"jpeg\"]\nSEVERITIES  = [1, 2, 3, 4, 5]\nN_BOOTSTRAP = 2000\n\nRUN_BENCHMARK = True; RUN_CROSSCOND = True; RUN_ROBUSTNESS = True\nRUN_CALIBRATION = True; RUN_XAI = True\n\nAPTOS_DIR = \"/kaggle/input/competitions/aptos2019-blindness-detection\"\nTRAIN_CSV = os.path.join(APTOS_DIR, \"train.csv\")\nTRAIN_IMG = os.path.join(APTOS_DIR, \"train_images\")\nOUT = \"/kaggle/working/outputs\"; os.makedirs(OUT, exist_ok=True)\n\nif QUICK_TEST:\n    EPOCHS = 2; N_BOOTSTRAP = 200; SEVERITIES = [1, 3, 5]\n\ndef set_seed(s=SEED):\n    random.seed(s); np.random.seed(s)\n    import torch; torch.manual_seed(s); torch.cuda.manual_seed_all(s)\nset_seed()\nprint(\"QUICK_TEST =\", QUICK_TEST, \"| MODELS =\", MODELS, \"| EPOCHS =\", EPOCHS)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:51.355359Z","iopub.execute_input":"2026-06-06T06:48:51.355698Z","iopub.status.idle":"2026-06-06T06:48:51.365975Z","shell.execute_reply.started":"2026-06-06T06:48:51.355669Z","shell.execute_reply":"2026-06-06T06:48:51.365363Z"}},"outputs":[],"execution_count":null},{"id":"97845480-8822-430f-b11d-ea1f17df4840","cell_type":"code","source":"import torch\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"Device count:\", torch.cuda.device_count())\nprint(\"Current device:\", torch.cuda.current_device() if torch.cuda.is_available() else \"CPU\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:51.367655Z","iopub.execute_input":"2026-06-06T06:48:51.367947Z","iopub.status.idle":"2026-06-06T06:48:51.386254Z","shell.execute_reply.started":"2026-06-06T06:48:51.367927Z","shell.execute_reply":"2026-06-06T06:48:51.385410Z"}},"outputs":[],"execution_count":null},{"id":"5d4a8def","cell_type":"markdown","source":"## 2. Environment and imports","metadata":{}},{"id":"edbec7fd","cell_type":"code","source":"import sys, subprocess\ndef pipq(pkg): subprocess.run([sys.executable,\"-m\",\"pip\",\"install\",\"-q\",pkg], check=False)\ntry: import timm\nexcept Exception: pipq(\"timm\"); import timm\ntry: from pytorch_grad_cam import GradCAM\nexcept Exception: pipq(\"grad-cam\")\nimport torch, torch.nn as nn\nimport pandas as pd, matplotlib.pyplot as plt\nfrom PIL import Image, ImageEnhance, ImageFilter\nimport io, cv2, time, warnings\nwarnings.filterwarnings(\"ignore\")\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (accuracy_score, f1_score, cohen_kappa_score,\n                             confusion_matrix, roc_auc_score, recall_score)\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nNGPU = 1\nprint(\"Torch\", torch.__version__, \"| device\", DEVICE, \"| GPUs\", NGPU, \"| timm\", timm.__version__)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:51.387221Z","iopub.execute_input":"2026-06-06T06:48:51.387523Z","iopub.status.idle":"2026-06-06T06:48:51.403436Z","shell.execute_reply.started":"2026-06-06T06:48:51.387491Z","shell.execute_reply":"2026-06-06T06:48:51.402752Z"}},"outputs":[],"execution_count":null},{"id":"d5dc7534-baad-4042-a9c0-d4cfd51da22c","cell_type":"code","source":"import os\nprint(os.path.exists(\"/kaggle/input\"))\nprint(os.listdir(\"/kaggle/input\") if os.path.exists(\"/kaggle/input\") else \"no input dir\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:51.404343Z","iopub.execute_input":"2026-06-06T06:48:51.404879Z","iopub.status.idle":"2026-06-06T06:48:51.419625Z","shell.execute_reply.started":"2026-06-06T06:48:51.404851Z","shell.execute_reply":"2026-06-06T06:48:51.418927Z"}},"outputs":[],"execution_count":null},{"id":"2a437b80-460d-4027-b57a-6d22a0a1f970","cell_type":"code","source":"import os\nfor root, dirs, files in os.walk(\"/kaggle/input/competitions\"):\n    print(root)\n    if files:\n        print(\"  files:\", files[:5], \"...\" if len(files)>5 else \"\")\n    if root.count(\"/\") > 5:  # don't go too deep\n        dirs.clear()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:51.421192Z","iopub.execute_input":"2026-06-06T06:48:51.421422Z","iopub.status.idle":"2026-06-06T06:48:54.138606Z","shell.execute_reply.started":"2026-06-06T06:48:51.421403Z","shell.execute_reply":"2026-06-06T06:48:54.137934Z"}},"outputs":[],"execution_count":null},{"id":"c1d31d3e","cell_type":"markdown","source":"## 3. Load APTOS labels and inspect class distribution","metadata":{}},{"id":"350c1c6c","cell_type":"code","source":"assert os.path.exists(TRAIN_CSV), f\"APTOS not found at {APTOS_DIR}. Add 'aptos2019-blindness-detection' via Add Input.\"\ndf = pd.read_csv(TRAIN_CSV)\ndf[\"path\"] = df[\"id_code\"].apply(lambda c: os.path.join(TRAIN_IMG, c + \".png\"))\ndf = df[df[\"path\"].apply(os.path.exists)].reset_index(drop=True)\nprint(\"Total labelled images:\", len(df))\ndist = df[\"diagnosis\"].value_counts().sort_index(); print(dist.to_string())\ngrade_names = {0:\"No DR\",1:\"Mild\",2:\"Moderate\",3:\"Severe\",4:\"Proliferative\"}\npd.DataFrame({\"grade\":list(grade_names),\"name\":list(grade_names.values()),\n              \"count\":[int(dist.get(k,0)) for k in grade_names]}).to_csv(\n              os.path.join(OUT,\"table_class_distribution.csv\"), index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:54.139496Z","iopub.execute_input":"2026-06-06T06:48:54.139856Z","iopub.status.idle":"2026-06-06T06:48:54.193675Z","shell.execute_reply.started":"2026-06-06T06:48:54.139826Z","shell.execute_reply":"2026-06-06T06:48:54.193092Z"}},"outputs":[],"execution_count":null},{"id":"e8eea3a2","cell_type":"markdown","source":"## 4. Image-quality score for the cross-condition domain split\nSharpness via variance of the Laplacian (a standard blur metric). The median splits data into **high-** and **low-quality** domains, a clinically meaningful acquisition shift analogous to high-end vs low-cost cameras.","metadata":{}},{"id":"9379cfc4","cell_type":"code","source":"def laplacian_var(path, size=256):\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    if img is None: return 0.0\n    img = cv2.resize(img,(size,size)); return float(cv2.Laplacian(img,cv2.CV_64F).var())\nsample = df.sample(n=min(len(df), 400 if QUICK_TEST else len(df)), random_state=SEED).copy()\nsample[\"sharp\"] = sample[\"path\"].apply(laplacian_var)\nthr = float(sample[\"sharp\"].median()); print(\"Sharpness median threshold =\", round(thr,1))\nif QUICK_TEST:\n    df[\"sharp\"] = sample.set_index(\"id_code\")[\"sharp\"].reindex(df[\"id_code\"]).fillna(thr).values\nelse:\n    df[\"sharp\"] = df[\"path\"].apply(laplacian_var)\ndf[\"quality\"] = np.where(df[\"sharp\"]>=thr, \"high\", \"low\")\nprint(df[\"quality\"].value_counts().to_string())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:48:54.195043Z","iopub.execute_input":"2026-06-06T06:48:54.195234Z","iopub.status.idle":"2026-06-06T06:49:32.026646Z","shell.execute_reply.started":"2026-06-06T06:48:54.195216Z","shell.execute_reply":"2026-06-06T06:49:32.025906Z"}},"outputs":[],"execution_count":null},{"id":"53e8de1b","cell_type":"markdown","source":"## 5. Stratified splits (70/15/15) plus quality-domain subsets","metadata":{}},{"id":"f5ebefbc","cell_type":"code","source":"def stratified_split(frame, seed=SEED):\n    tr, tmp = train_test_split(frame, test_size=0.30, stratify=frame[\"diagnosis\"], random_state=seed)\n    va, te  = train_test_split(tmp, test_size=0.50, stratify=tmp[\"diagnosis\"], random_state=seed)\n    return tr.reset_index(drop=True), va.reset_index(drop=True), te.reset_index(drop=True)\ntrain_df, val_df, test_df = stratified_split(df)\nprint(\"Full ->\", len(train_df), len(val_df), len(test_df))\ndom = {}\nfor q in [\"high\",\"low\"]:\n    sub = df[df[\"quality\"]==q]\n    tr,va,te = stratified_split(sub); dom[q]=(tr,va,te)\n    print(f\"{q:4s} ->\", len(tr), len(va), len(te))\npd.DataFrame({\"partition\":[\"full_train\",\"full_val\",\"full_test\",\"high_train\",\"high_val\",\"high_test\",\"low_train\",\"low_val\",\"low_test\"],\n  \"images\":[len(train_df),len(val_df),len(test_df),len(dom['high'][0]),len(dom['high'][1]),len(dom['high'][2]),\n            len(dom['low'][0]),len(dom['low'][1]),len(dom['low'][2])]}).to_csv(\n            os.path.join(OUT,\"table_dataset_summary.csv\"), index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:49:32.027483Z","iopub.execute_input":"2026-06-06T06:49:32.027831Z","iopub.status.idle":"2026-06-06T06:49:32.057061Z","shell.execute_reply.started":"2026-06-06T06:49:32.027808Z","shell.execute_reply":"2026-06-06T06:49:32.056454Z"}},"outputs":[],"execution_count":null},{"id":"90c1acd8","cell_type":"markdown","source":"## 6. Dataset, transforms, and corruption functions","metadata":{}},{"id":"9fc91b4b","cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nMEAN=[0.485,0.456,0.406]; STD=[0.229,0.224,0.225]\ntrain_tf = T.Compose([T.RandomResizedCrop(IMG_SIZE, scale=(0.8,1.0)), T.RandomHorizontalFlip(),\n                      T.ToTensor(), T.Normalize(MEAN,STD)])\neval_tf  = T.Compose([T.Resize(int(IMG_SIZE*1.14)), T.CenterCrop(IMG_SIZE),\n                      T.ToTensor(), T.Normalize(MEAN,STD)])\ndef corrupt(img, ctype, s):\n    if ctype==\"gaussian_blur\": return img.filter(ImageFilter.GaussianBlur([0.5,1.0,2.0,3.0,4.0][s-1]))\n    if ctype==\"brightness\": return ImageEnhance.Brightness(img).enhance([1.2,1.4,1.6,1.8,2.1][s-1])\n    if ctype==\"contrast\":   return ImageEnhance.Contrast(img).enhance([0.85,0.7,0.55,0.4,0.3][s-1])\n    if ctype==\"gaussian_noise\":\n        a=np.array(img).astype(np.float32)+np.random.normal(0,[8,16,28,42,60][s-1],np.array(img).shape)\n        return Image.fromarray(np.clip(a,0,255).astype(np.uint8))\n    if ctype==\"jpeg\":\n        b=io.BytesIO(); img.save(b,\"JPEG\",quality=[60,45,32,22,14][s-1]); b.seek(0)\n        return Image.open(b).convert(\"RGB\")\n    return img\nclass RetinaDS(Dataset):\n    def __init__(self, frame, tf, ctype=None, sev=0):\n        self.f=frame.reset_index(drop=True); self.tf=tf; self.ctype=ctype; self.sev=sev\n    def __len__(self): return len(self.f)\n    def __getitem__(self, i):\n        r=self.f.iloc[i]; img=Image.open(r[\"path\"]).convert(\"RGB\")\n        if self.ctype is not None and self.sev>0: img=corrupt(img,self.ctype,self.sev)\n        return self.tf(img), int(r[\"diagnosis\"])\ndef loader(frame, tf, shuffle=False, ctype=None, sev=0):\n    return DataLoader(RetinaDS(frame,tf,ctype,sev), batch_size=BATCH, shuffle=shuffle,\n                      num_workers=2, pin_memory=True)\nprint(\"data utils ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:49:32.057859Z","iopub.execute_input":"2026-06-06T06:49:32.058119Z","iopub.status.idle":"2026-06-06T06:49:32.069401Z","shell.execute_reply.started":"2026-06-06T06:49:32.058092Z","shell.execute_reply":"2026-06-06T06:49:32.068605Z"}},"outputs":[],"execution_count":null},{"id":"0f692df1","cell_type":"markdown","source":"## 7. Model builder, metrics, and helpers","metadata":{}},{"id":"9b01d792","cell_type":"code","source":"def build_model(name): return timm.create_model(name, pretrained=True, num_classes=NUM_CLASSES)\ndef to_parallel(m):\n    m=m.to(DEVICE)\n    if NGPU>1: m=nn.DataParallel(m)\n    return m\n@torch.no_grad()\ndef predict(model, ld):\n    model.eval(); P=[]; Y=[]\n    for x,y in ld:\n        x=x.to(DEVICE, non_blocking=True)\n        with torch.cuda.amp.autocast(): out=model(x)\n        P.append(torch.softmax(out,1).float().cpu().numpy()); Y.append(y.numpy())\n    return np.concatenate(P), np.concatenate(Y)\ndef ece_score(probs, y, n_bins=10):\n    conf=probs.max(1); pred=probs.argmax(1); acc=(pred==y).astype(float)\n    bins=np.linspace(0,1,n_bins+1); e=0.0; n=len(y)\n    for i in range(n_bins):\n        m=(conf>bins[i])&(conf<=bins[i+1])\n        if m.sum()>0: e+=m.sum()/n*abs(acc[m].mean()-conf[m].mean())\n    return float(e)\ndef all_metrics(probs, y):\n    pred=probs.argmax(1); labels=list(range(NUM_CLASSES)); out={}\n    out[\"accuracy\"]=accuracy_score(y,pred)\n    out[\"macro_f1\"]=f1_score(y,pred,average=\"macro\",labels=labels,zero_division=0)\n    out[\"qwk\"]=cohen_kappa_score(y,pred,weights=\"quadratic\",labels=labels)\n    rec=recall_score(y,pred,average=None,labels=labels,zero_division=0)\n    for k in labels: out[f\"recall_{k}\"]=float(rec[k])\n    yb=(y>=2).astype(int); pb=(pred>=2).astype(int)\n    tp=int(((pb==1)&(yb==1)).sum()); fn=int(((pb==0)&(yb==1)).sum())\n    tn=int(((pb==0)&(yb==0)).sum()); fp=int(((pb==1)&(yb==0)).sum())\n    out[\"referable_sens\"]=tp/(tp+fn+1e-9); out[\"referable_spec\"]=tn/(tn+fp+1e-9)\n    try:\n        present=[c for c in labels if (y==c).sum()>0]\n        if len(present)>1:\n            yoh=np.eye(NUM_CLASSES)[y][:,present]\n            out[\"auc_macro\"]=roc_auc_score(yoh, probs[:,present], average=\"macro\", multi_class=\"ovr\")\n        else: out[\"auc_macro\"]=float(\"nan\")\n    except Exception: out[\"auc_macro\"]=float(\"nan\")\n    out[\"ece\"]=ece_score(probs,y)\n    out[\"mean_signed_error\"]=float((pred-y).mean())\n    return out\ndef bootstrap_ci(probs, y, fn, n=N_BOOTSTRAP):\n    rng=np.random.RandomState(SEED); N=len(y); vals=[]\n    for _ in range(n):\n        idx=rng.randint(0,N,N)\n        try: vals.append(fn(probs[idx], y[idx]))\n        except Exception: pass\n    if not vals: return float(\"nan\"), float(\"nan\")\n    lo,hi=np.percentile(vals,[2.5,97.5]); return float(lo),float(hi)\nprint(\"model + metrics ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:49:32.070467Z","iopub.execute_input":"2026-06-06T06:49:32.070803Z","iopub.status.idle":"2026-06-06T06:49:32.089115Z","shell.execute_reply.started":"2026-06-06T06:49:32.070781Z","shell.execute_reply":"2026-06-06T06:49:32.088495Z"}},"outputs":[],"execution_count":null},{"id":"a2de1dad","cell_type":"markdown","source":"## 8. Training loop (AMP + DataParallel + class-weighted loss; best checkpoint by val QWK)","metadata":{}},{"id":"02f00a8b","cell_type":"code","source":"def train_model(name, tr_df, va_df, tag):\n    set_seed()\n    model = to_parallel(build_model(name))\n    counts=tr_df[\"diagnosis\"].value_counts().reindex(range(NUM_CLASSES)).fillna(0).values\n    w=torch.tensor(counts.sum()/(counts+1e-6), dtype=torch.float32); w=w/w.mean()\n    crit=nn.CrossEntropyLoss(weight=w.to(DEVICE))\n    opt=torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    sched=torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=EPOCHS)\n    scaler=torch.cuda.amp.GradScaler()\n    tr=loader(tr_df, train_tf, shuffle=True); va=loader(va_df, eval_tf)\n    best=-1; best_state=None\n    for ep in range(EPOCHS):\n        model.train(); t0=time.time()\n        for x,y in tr:\n            x=x.to(DEVICE,non_blocking=True); y=y.to(DEVICE,non_blocking=True)\n            opt.zero_grad()\n            with torch.cuda.amp.autocast(): loss=crit(model(x), y)\n            scaler.scale(loss).backward(); scaler.step(opt); scaler.update()\n        sched.step()\n        p,yt=predict(model, va); q=cohen_kappa_score(yt,p.argmax(1),weights=\"quadratic\",labels=list(range(NUM_CLASSES)))\n        if q>best: best=q; best_state={k:v.detach().cpu().clone() for k,v in model.state_dict().items()}\n        print(f\"  [{tag}] epoch {ep+1}/{EPOCHS} val_QWK={q:.4f} ({time.time()-t0:.0f}s)\")\n    if best_state is not None: model.load_state_dict(best_state)\n    return model\nprint(\"trainer ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:49:32.090073Z","iopub.execute_input":"2026-06-06T06:49:32.090329Z","iopub.status.idle":"2026-06-06T06:49:32.106959Z","shell.execute_reply.started":"2026-06-06T06:49:32.090309Z","shell.execute_reply":"2026-06-06T06:49:32.106116Z"}},"outputs":[],"execution_count":null},{"id":"3865d1a7","cell_type":"markdown","source":"## 9. Experiment 1 - Benchmark","metadata":{}},{"id":"09f44eb8","cell_type":"code","source":"benchmark_models={}; rows=[]\nif RUN_BENCHMARK:\n    for name in MODELS:\n        print(\"Training\", name, \"...\")\n        m=train_model(name, train_df, val_df, name); benchmark_models[name]=m\n        probs,y=predict(m, loader(test_df, eval_tf)); met=all_metrics(probs,y)\n        met[\"acc_lo\"],met[\"acc_hi\"]=bootstrap_ci(probs,y, lambda p,t: accuracy_score(t,p.argmax(1)))\n        met[\"qwk_lo\"],met[\"qwk_hi\"]=bootstrap_ci(probs,y, lambda p,t: cohen_kappa_score(t,p.argmax(1),weights=\"quadratic\",labels=list(range(NUM_CLASSES))))\n        met[\"model\"]=name; rows.append(met)\n        np.save(os.path.join(OUT,f\"probs_{name}.npy\"), probs); np.save(os.path.join(OUT,\"ytest.npy\"), y)\n    bench=pd.DataFrame(rows); bench.to_csv(os.path.join(OUT,\"table_benchmark.csv\"), index=False)\n    print(bench[[\"model\",\"accuracy\",\"macro_f1\",\"qwk\",\"referable_sens\",\"referable_spec\",\"auc_macro\",\"ece\"]].to_string(index=False))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T06:49:32.108855Z","iopub.execute_input":"2026-06-06T06:49:32.109081Z","iopub.status.idle":"2026-06-06T07:04:08.627157Z","shell.execute_reply.started":"2026-06-06T06:49:32.109062Z","shell.execute_reply":"2026-06-06T07:04:08.626192Z"}},"outputs":[],"execution_count":null},{"id":"b7d4e0b2","cell_type":"markdown","source":"### Confusion matrix figure","metadata":{}},{"id":"2fb1d22d","cell_type":"code","source":"if RUN_BENCHMARK:\n    y=np.load(os.path.join(OUT,\"ytest.npy\"))\n    fig,axes=plt.subplots(1,len(MODELS),figsize=(5*len(MODELS),4.2))\n    if len(MODELS)==1: axes=[axes]\n    for ax,name in zip(axes,MODELS):\n        probs=np.load(os.path.join(OUT,f\"probs_{name}.npy\"))\n        cm=confusion_matrix(y,probs.argmax(1),labels=list(range(NUM_CLASSES)))\n        cmn=cm/cm.sum(1,keepdims=True).clip(min=1)\n        ax.imshow(cmn,cmap=\"Blues\",vmin=0,vmax=1); ax.set_title(name)\n        ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\"); ax.set_xticks(range(5)); ax.set_yticks(range(5))\n        for i in range(5):\n            for j in range(5):\n                ax.text(j,i,f\"{cmn[i,j]:.2f}\",ha=\"center\",va=\"center\",\n                        color=\"white\" if cmn[i,j]>0.5 else \"black\",fontsize=8)\n    plt.tight_layout(); plt.savefig(os.path.join(OUT,\"fig_confusion.png\"),dpi=300,bbox_inches=\"tight\")\n    plt.savefig(os.path.join(OUT,\"fig_confusion.pdf\"),bbox_inches=\"tight\"); plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:04:08.628667Z","iopub.execute_input":"2026-06-06T07:04:08.628968Z","iopub.status.idle":"2026-06-06T07:04:10.605419Z","shell.execute_reply.started":"2026-06-06T07:04:08.628941Z","shell.execute_reply":"2026-06-06T07:04:10.604771Z"}},"outputs":[],"execution_count":null},{"id":"68c37bd9","cell_type":"markdown","source":"## 10. Experiment 2 - Cross-condition generalisation (high vs low image quality)","metadata":{}},{"id":"d808760e","cell_type":"code","source":"cross_rows=[]\nif RUN_CROSSCOND:\n    for name in MODELS:\n        for trq in [\"high\",\"low\"]:\n            tr,va,_=dom[trq]; m=train_model(name, tr, va, f\"{name}-{trq}\")\n            for teq in [\"high\",\"low\"]:\n                _,_,te=dom[teq]; probs,y=predict(m, loader(te, eval_tf)); met=all_metrics(probs,y)\n                cross_rows.append({\"model\":name,\"train_q\":trq,\"test_q\":teq,\n                    \"setting\":\"in_domain\" if trq==teq else \"cross_domain\",\n                    \"accuracy\":met[\"accuracy\"],\"qwk\":met[\"qwk\"],\"referable_sens\":met[\"referable_sens\"]})\n            del m; torch.cuda.empty_cache()\n    cross=pd.DataFrame(cross_rows); cross.to_csv(os.path.join(OUT,\"table_crosscondition.csv\"),index=False)\n    gaps=[]\n    for name in MODELS:\n        for tq in [\"high\",\"low\"]:\n            ind=cross[(cross.model==name)&(cross.train_q==tq)&(cross.setting==\"in_domain\")][\"qwk\"].values\n            crd=cross[(cross.model==name)&(cross.train_q==tq)&(cross.setting==\"cross_domain\")][\"qwk\"].values\n            if len(ind) and len(crd):\n                gaps.append({\"model\":name,\"train_q\":tq,\"in_domain_qwk\":float(ind[0]),\n                             \"cross_domain_qwk\":float(crd[0]),\"gap\":float(ind[0]-crd[0])})\n    pd.DataFrame(gaps).to_csv(os.path.join(OUT,\"table_generalisation_gap.csv\"),index=False)\n    print(cross.to_string(index=False)); print(); print(pd.DataFrame(gaps).to_string(index=False))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:04:10.606436Z","iopub.execute_input":"2026-06-06T07:04:10.606877Z","iopub.status.idle":"2026-06-06T07:19:57.829808Z","shell.execute_reply.started":"2026-06-06T07:04:10.606848Z","shell.execute_reply":"2026-06-06T07:19:57.829053Z"}},"outputs":[],"execution_count":null},{"id":"77c7ccd4","cell_type":"markdown","source":"## 11. Experiment 3 - Robustness to image degradation (centerpiece)\nPer corruption/severity: accuracy, QWK, referable-DR sensitivity, ECE, and the **mean signed grading error** (negative = dangerous under-grading; positive = safer over-grading).","metadata":{}},{"id":"eba5b299","cell_type":"code","source":"rob_rows=[]\nif RUN_ROBUSTNESS and RUN_BENCHMARK:\n    name=MODELS[0]; m=benchmark_models[name]\n    p,y=predict(m, loader(test_df, eval_tf)); base=all_metrics(p,y)\n    keep=[\"accuracy\",\"qwk\",\"referable_sens\",\"referable_spec\",\"ece\",\"mean_signed_error\"]\n    rob_rows.append({\"corruption\":\"clean\",\"severity\":0,**{k:base[k] for k in keep}})\n    for c in CORRUPTIONS:\n        for s in SEVERITIES:\n            p,y=predict(m, loader(test_df, eval_tf, ctype=c, sev=s)); mt=all_metrics(p,y)\n            rob_rows.append({\"corruption\":c,\"severity\":s,**{k:mt[k] for k in keep}})\n            print(f\"  {c:15s} sev{s}: acc={mt['accuracy']:.3f} qwk={mt['qwk']:.3f} \"\n                  f\"refSens={mt['referable_sens']:.3f} signedErr={mt['mean_signed_error']:+.3f}\")\n    pd.DataFrame(rob_rows).to_csv(os.path.join(OUT,\"table_robustness.csv\"),index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:19:57.831123Z","iopub.execute_input":"2026-06-06T07:19:57.831512Z","iopub.status.idle":"2026-06-06T07:38:51.593527Z","shell.execute_reply.started":"2026-06-06T07:19:57.831480Z","shell.execute_reply":"2026-06-06T07:38:51.592312Z"}},"outputs":[],"execution_count":null},{"id":"7dd2db91","cell_type":"markdown","source":"### Robustness figures","metadata":{}},{"id":"3b42d3e1","cell_type":"code","source":"if RUN_ROBUSTNESS:\n    rob=pd.read_csv(os.path.join(OUT,\"table_robustness.csv\")); base=rob[rob.corruption==\"clean\"].iloc[0]\n    metrics=[(\"accuracy\",\"Accuracy\"),(\"qwk\",\"Quadratic-weighted kappa\"),\n             (\"referable_sens\",\"Referable-DR sensitivity\"),(\"mean_signed_error\",\"Mean signed grading error\")]\n    colors={\"gaussian_blur\":\"#1f77b4\",\"brightness\":\"#ff7f0e\",\"contrast\":\"#2ca02c\",\n            \"gaussian_noise\":\"#9467bd\",\"jpeg\":\"#d62728\"}\n    fig,axes=plt.subplots(1,4,figsize=(18,4.2))\n    for ax,(mk,mn) in zip(axes,metrics):\n        for c in CORRUPTIONS:\n            sub=rob[rob.corruption==c].sort_values(\"severity\")\n            ax.plot(sub[\"severity\"],sub[mk],marker=\"o\",label=c.replace(\"_\",\" \"),color=colors[c])\n        ax.axhline(base[mk],ls=\"--\",color=\"grey\",label=\"clean\"); ax.set_title(mn)\n        ax.set_xlabel(\"Severity\"); ax.grid(alpha=0.25)\n    axes[0].set_ylabel(\"Score\")\n    h,l=axes[0].get_legend_handles_labels()\n    fig.legend(h,l,loc=\"lower center\",ncol=6,frameon=False,bbox_to_anchor=(0.5,-0.04))\n    fig.tight_layout(rect=[0,0.05,1,1])\n    fig.savefig(os.path.join(OUT,\"fig_robustness_curves.png\"),dpi=300,bbox_inches=\"tight\")\n    fig.savefig(os.path.join(OUT,\"fig_robustness_curves.pdf\"),bbox_inches=\"tight\"); plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:38:51.595216Z","iopub.execute_input":"2026-06-06T07:38:51.595703Z","iopub.status.idle":"2026-06-06T07:38:53.538914Z","shell.execute_reply.started":"2026-06-06T07:38:51.595671Z","shell.execute_reply":"2026-06-06T07:38:53.538244Z"}},"outputs":[],"execution_count":null},{"id":"29f15760","cell_type":"markdown","source":"## 12. Experiment 4 - Calibration (reliability diagram)","metadata":{}},{"id":"e00235c6","cell_type":"code","source":"if RUN_CALIBRATION and RUN_BENCHMARK:\n    name=MODELS[0]; probs=np.load(os.path.join(OUT,f\"probs_{name}.npy\")); y=np.load(os.path.join(OUT,\"ytest.npy\"))\n    conf=probs.max(1); pred=probs.argmax(1); acc=(pred==y).astype(float); bins=np.linspace(0,1,11); xs=[]; ys=[]\n    for i in range(10):\n        msk=(conf>bins[i])&(conf<=bins[i+1])\n        if msk.sum()>0: xs.append(conf[msk].mean()); ys.append(acc[msk].mean())\n    ece=ece_score(probs,y)\n    plt.figure(figsize=(5,5)); plt.plot([0,1],[0,1],\"--\",color=\"grey\",label=\"perfect\")\n    plt.plot(xs,ys,marker=\"o\",label=f\"{name} (ECE={ece:.3f})\")\n    plt.xlabel(\"Predicted confidence\"); plt.ylabel(\"Empirical accuracy\"); plt.legend()\n    plt.title(\"Reliability diagram\"); plt.tight_layout()\n    plt.savefig(os.path.join(OUT,\"fig_calibration.png\"),dpi=300,bbox_inches=\"tight\")\n    plt.savefig(os.path.join(OUT,\"fig_calibration.pdf\"),bbox_inches=\"tight\"); plt.show()\n    pd.DataFrame({\"model\":[name],\"ECE\":[ece],\"n\":[len(y)]}).to_csv(os.path.join(OUT,\"table_calibration.csv\"),index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:38:53.539832Z","iopub.execute_input":"2026-06-06T07:38:53.540164Z","iopub.status.idle":"2026-06-06T07:38:54.093392Z","shell.execute_reply.started":"2026-06-06T07:38:53.540142Z","shell.execute_reply":"2026-06-06T07:38:54.092849Z"}},"outputs":[],"execution_count":null},{"id":"1dce1804","cell_type":"markdown","source":"## 13. Experiment 5 - Explainability (Grad-CAM + energy-on-retina)\nFundus images have a dark circular background; a faithful model should concentrate activation **inside the retinal disc**. We report the fraction of Grad-CAM energy on the retina plus qualitative overlays.","metadata":{}},{"id":"6d2e6e3d","cell_type":"code","source":"if RUN_XAI and RUN_BENCHMARK:\n    from pytorch_grad_cam import GradCAM\n    from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\n    name=GRADCAM_MODEL; base_model=build_model(name).to(DEVICE).eval()\n    sd=benchmark_models[name].state_dict(); sd={k.replace(\"module.\",\"\"):v for k,v in sd.items()}\n    base_model.load_state_dict(sd)\n    target_layer=None\n    for mod in base_model.modules():\n        if isinstance(mod, nn.Conv2d): target_layer=mod\n    cam=GradCAM(model=base_model, target_layers=[target_layer])\n    inv=T.Normalize([-m/s for m,s in zip(MEAN,STD)],[1/s for s in STD])\n    samp=test_df.sample(n=min(6,len(test_df)), random_state=SEED).reset_index(drop=True)\n    energies=[]; fig,axes=plt.subplots(2,len(samp),figsize=(3.2*len(samp),6.2))\n    for j in range(len(samp)):\n        r=samp.iloc[j]; pil=Image.open(r[\"path\"]).convert(\"RGB\")\n        x=eval_tf(pil).unsqueeze(0).to(DEVICE)\n        g=cam(input_tensor=x, targets=[ClassifierOutputTarget(int(r[\"diagnosis\"]))])[0]\n        rgb=inv(x[0].cpu()).permute(1,2,0).numpy().clip(0,1)\n        mask=(rgb.mean(2)>0.10).astype(np.float32)\n        energies.append({\"image\":r[\"id_code\"],\"energy_on_retina\":float((g*mask).sum()/(g.sum()+1e-9)),\"grade\":int(r[\"diagnosis\"])})\n        axes[0,j].imshow(rgb); axes[0,j].set_title(f\"grade {int(r['diagnosis'])}\"); axes[0,j].axis(\"off\")\n        axes[1,j].imshow(rgb); axes[1,j].imshow(g,cmap=\"jet\",alpha=0.5); axes[1,j].set_title(\"Grad-CAM\"); axes[1,j].axis(\"off\")\n    plt.tight_layout(); plt.savefig(os.path.join(OUT,\"fig_xai.png\"),dpi=300,bbox_inches=\"tight\")\n    plt.savefig(os.path.join(OUT,\"fig_xai.pdf\"),bbox_inches=\"tight\"); plt.show()\n    xdf=pd.DataFrame(energies); xdf.to_csv(os.path.join(OUT,\"table_xai_energy.csv\"),index=False)\n    print(\"mean energy-on-retina:\", round(xdf[\"energy_on_retina\"].mean(),3))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:38:54.094348Z","iopub.execute_input":"2026-06-06T07:38:54.094633Z","iopub.status.idle":"2026-06-06T07:38:59.068291Z","shell.execute_reply.started":"2026-06-06T07:38:54.094601Z","shell.execute_reply":"2026-06-06T07:38:59.067373Z"}},"outputs":[],"execution_count":null},{"id":"d1591f66","cell_type":"markdown","source":"## 14. Package results","metadata":{}},{"id":"5db25d23","cell_type":"code","source":"import shutil\nshutil.make_archive(\"/kaggle/working/aptos_results\",\"zip\",OUT)\nprint(\"Saved /kaggle/working/aptos_results.zip\"); print(\"Files:\")\nfor f in sorted(os.listdir(OUT)): print(\"  \", f)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:38:59.069396Z","iopub.execute_input":"2026-06-06T07:38:59.069748Z","iopub.status.idle":"2026-06-06T07:38:59.209779Z","shell.execute_reply.started":"2026-06-06T07:38:59.069725Z","shell.execute_reply":"2026-06-06T07:38:59.209138Z"}},"outputs":[],"execution_count":null},{"id":"e591631b","cell_type":"markdown","source":"## 15. Full text dump (paste this back)","metadata":{}},{"id":"f9edf7bc","cell_type":"code","source":"import glob\nfor csv in sorted(glob.glob(os.path.join(OUT,\"*.csv\"))):\n    print(\"=\"*70); print(os.path.basename(csv)); print(\"=\"*70)\n    print(pd.read_csv(csv).to_string(index=False)); print()\nprint(\"FIGURES:\", [os.path.basename(f) for f in sorted(glob.glob(os.path.join(OUT,'*.png')))])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-06T07:38:59.210531Z","iopub.execute_input":"2026-06-06T07:38:59.210738Z","iopub.status.idle":"2026-06-06T07:38:59.242590Z","shell.execute_reply.started":"2026-06-06T07:38:59.210718Z","shell.execute_reply":"2026-06-06T07:38:59.241974Z"}},"outputs":[],"execution_count":null}]}