{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"2e2d1dac","cell_type":"markdown","source":"# Diabetic Retinopathy Stage Detection — CNN + Transfer Learning\n**Module:** Computer Vision · **Dataset:** APTOS 2019 Blindness Detection (Kaggle) · **Task:** 5-grade classification\n\nNotebook layout follows the marking rubric: (1) problem & data → (2) preprocessing → (3) augmentation & balancing →\n(4) architecture & transfer learning → (5) training strategy → (6) evaluation → (7) Grad-CAM / impact.\nAll logic lives in the commented modules under `src/`; this notebook only calls them and produces the figures for the report.\n\n**Run on a GPU** (Kaggle: Settings → Accelerator → GPU T4, Internet ON, add the APTOS competition data and this project as a dataset).","metadata":{}},{"id":"3e475c31-3376-478d-8c2b-5065c7e6ee12","cell_type":"code","source":"import sys, os, glob, warnings\nwarnings.filterwarnings(\"ignore\")\nhits = glob.glob(\"/kaggle/input/**/src/config.py\", recursive=True)\nif not hits:\n    raise SystemExit(\"Cannot find src/ - check the dataset is added under Input\")\nPROJECT_DIR = os.path.dirname(os.path.dirname(hits[0]))\nsys.path.insert(0, PROJECT_DIR)\nprint(\"Using project:\", PROJECT_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:29:08.348455Z","iopub.execute_input":"2026-09-29T19:29:08.348767Z","iopub.status.idle":"2026-09-29T19:29:08.378738Z","shell.execute_reply.started":"2026-09-29T19:29:08.348734Z","shell.execute_reply":"2026-09-29T19:29:08.377989Z"}},"outputs":[],"execution_count":null},{"id":"67371e78","cell_type":"code","source":" import os\nimport numpy as np, pandas as pd, torch, matplotlib.pyplot as plt\nfrom src.config import Config, CLASS_NAMES\nfrom src.utils import set_seed, get_device, save_json, ensure_dir\nfrom src import visualise as V, evaluate as E\nfrom src.pipeline import prepare_data, run_experiment, variant, sweep\nfrom src.dataset import load_dataframe, stratified_split, make_sampler, class_counts\nfrom src.model import build_model, gradcam_layer, count_params\nfrom src.gradcam import GradCAM, overlay\nfrom src.augment import get_eval_transform\nfrom PIL import Image\n\n# ---- paths ----\ncfg = Config(\n    data_dir=\"/kaggle/input/competitions/aptos2019-blindness-detection\",\n    cache_dir=\"/tmp/preprocessed\",\n    output_dir=\"/kaggle/working/outputs\",\n)\nFIG = ensure_dir(os.path.join(cfg.output_dir, \"figures\"))\nset_seed(cfg.seed); device = get_device()\nprint(\"device:\", device, \"|\", torch.cuda.get_device_name(0) if torch.cuda.is_available() else \"CPU (too slow for training!)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:57:15.389417Z","iopub.execute_input":"2026-09-29T19:57:15.390249Z","iopub.status.idle":"2026-09-29T19:57:15.399826Z","shell.execute_reply.started":"2026-09-29T19:57:15.390219Z","shell.execute_reply":"2026-09-29T19:57:15.398782Z"}},"outputs":[],"execution_count":null},{"id":"cc5f1ee5","cell_type":"markdown","source":"## 1. Problem understanding & dataset\nDiabetic retinopathy (DR) is damage to the retinal blood vessels caused by diabetes and a leading cause of preventable blindness.\nIt is graded 0–4 (none, mild, moderate, severe, proliferative). Screening needs specialists to read fundus photographs, which is slow\nand scarce in many regions, so automated grading can support triage. The dataset is the APTOS 2019 training set (3,662 labelled\nfundus images, 5 grades).","metadata":{}},{"id":"a81c304a","cell_type":"code","source":"df = load_dataframe(cfg)\nprint(len(df), \"images\"); display(df[\"diagnosis\"].value_counts().sort_index().rename(index=lambda i: CLASS_NAMES[i]).to_frame(\"count\").assign(percent=lambda d: (100*d[\"count\"]/d[\"count\"].sum()).round(1)))\nprint(\"\\nRaw image resolution statistics (300 random images):\"); display(V.image_size_stats(df))\nV.plot_samples(df, 4, path=f\"{FIG}/01_samples.png\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:29:08.41732Z","iopub.execute_input":"2026-09-29T19:29:08.417683Z","iopub.status.idle":"2026-09-29T19:29:29.417236Z","shell.execute_reply.started":"2026-09-29T19:29:08.417654Z","shell.execute_reply":"2026-09-29T19:29:29.416145Z"}},"outputs":[],"execution_count":null},{"id":"42935040","cell_type":"markdown","source":"### 1.1 Leakage-free stratified split (70 / 15 / 15)\nThe file list is split **before** any augmentation or balancing, stratified on the grade so rare classes appear in every split.","metadata":{}},{"id":"70ce3560","cell_type":"code","source":"train_df, val_df, test_df = stratified_split(df, cfg.val_frac, cfg.test_frac, cfg.seed)\nprint(len(train_df), len(val_df), len(test_df))\nassert not (set(train_df.id_code) & set(val_df.id_code)) and not (set(train_df.id_code) & set(test_df.id_code)) and not (set(val_df.id_code) & set(test_df.id_code))\nV.plot_class_distribution({\"train\": train_df, \"val\": val_df, \"test\": test_df}, f\"{FIG}/02_class_distribution.png\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:29:29.419588Z","iopub.execute_input":"2026-09-29T19:29:29.42026Z","iopub.status.idle":"2026-09-29T19:29:30.043106Z","shell.execute_reply.started":"2026-09-29T19:29:29.420215Z","shell.execute_reply":"2026-09-29T19:29:30.042348Z"}},"outputs":[],"execution_count":null},{"id":"3cb0c239","cell_type":"markdown","source":"## 2. Preprocessing\nCrop black borders → square pad + resize → non-local-means denoise → CLAHE on LAB lightness → unsharp masking (edge enhancement) → circular mask.\nWe also compare against `none`, `clahe` and Ben Graham's method, both visually, with objective image-quality numbers, and (section 5) by training accuracy.","metadata":{}},{"id":"40f4ebb1","cell_type":"code","source":"sample_path = df.iloc[0][\"path\"]\nV.plot_preprocessing_steps(sample_path, cfg.img_size, f\"{FIG}/03_preprocessing_steps.png\"); plt.show()\nV.plot_before_after(df, cfg.img_size, cfg.preprocess_mode, n=4, path=f\"{FIG}/04_before_after.png\"); plt.show()\nV.plot_mode_comparison(sample_path, cfg.img_size, f\"{FIG}/05_mode_comparison.png\"); plt.show()\nq = V.preprocessing_quality_table(df, cfg.img_size, n=100); q.to_csv(f\"{cfg.output_dir}/preprocessing_quality.csv\"); display(q)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:29:30.044225Z","iopub.execute_input":"2026-09-29T19:29:30.044586Z","iopub.status.idle":"2026-09-29T19:30:40.39774Z","shell.execute_reply.started":"2026-09-29T19:29:30.044538Z","shell.execute_reply":"2026-09-29T19:30:40.396706Z"}},"outputs":[],"execution_count":null},{"id":"8062a193","cell_type":"markdown","source":"## 3. Augmentation & class balancing\nFlip / rotate / mild zoom / brightness-contrast jitter — applied **on the training set only**. Class imbalance\n(grade 0 ≈ 49 %, grade 3 ≈ 5 %) is handled by a class-balanced `WeightedRandomSampler` (each epoch sees all grades about equally often,\nwith different random augmentations each time) — optionally combined with class-weighted loss (compared in section 5).","metadata":{}},{"id":"6c570973","cell_type":"code","source":"prepare_data(cfg)   # runs the preprocessing pipeline once for all images and caches results\nfrom src.pipeline import cache_path\nex = os.path.join(cache_path(cfg), test_df.iloc[0][\"id_code\"] + \".png\")\nV.plot_augmentations(ex, cfg.img_size, n=8, path=f\"{FIG}/06_augmentations.png\"); plt.show()\nV.plot_balance_effect(train_df, make_sampler(train_df), f\"{FIG}/07_balancing.png\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:30:40.399062Z","iopub.execute_input":"2026-09-29T19:30:40.399375Z","iopub.status.idle":"2026-09-29T19:40:03.386136Z","shell.execute_reply.started":"2026-09-29T19:30:40.39935Z","shell.execute_reply":"2026-09-29T19:40:03.385113Z"}},"outputs":[],"execution_count":null},{"id":"338aafab","cell_type":"markdown","source":"## 4. Model: CNN + transfer learning\nImageNet-pre-trained backbone + new head (Dropout → Linear(5)). **Phase 1:** backbone frozen, train head only.\n**Phase 2:** unfreeze all, discriminative learning rates (backbone 1e-4, head 3e-4), AdamW, ReduceLROnPlateau,\nlabel smoothing, gradient clipping, mixed precision, early stopping on validation macro-F1.","metadata":{}},{"id":"ffc1877b","cell_type":"code","source":"for name in [\"efficientnet_b0\", \"resnet50\", \"densenet121\"]:\n    m = build_model(name, 5, pretrained=False); t, _ = count_params(m); print(f\"{name:18s} {t/1e6:5.1f} M parameters\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:40:03.387285Z","iopub.execute_input":"2026-09-29T19:40:03.387653Z","iopub.status.idle":"2026-09-29T19:40:03.99639Z","shell.execute_reply.started":"2026-09-29T19:40:03.387591Z","shell.execute_reply":"2026-09-29T19:40:03.995642Z"}},"outputs":[],"execution_count":null},{"id":"86af674d","cell_type":"markdown","source":"## 5. Experimental design (validation set only)\nGreedy search: (a) preprocessing mode → (b) imbalance strategy → (c) backbone → (d) learning rate.\nEvery run is short (`QUICK_EPOCHS`) and judged on **validation** macro-F1; the test set is not used until section 6.\nSet `RUN_EXPERIMENTS = False` to skip and go straight to training with the defaults in `Config`.","metadata":{}},{"id":"221fdfba","cell_type":"code","source":"RUN_EXPERIMENTS = False   # the 12 experiments were already run; results are saved in experiments.csv\nQUICK = variant(cfg, head_epochs=2, finetune_epochs=6, patience=3)\ncache = {\"df\": df}\ntables = []\nif RUN_EXPERIMENTS:\n    t, best = sweep(QUICK, \"preprocess_mode\", [\"none\", \"clahe_unsharp\", \"ben_graham\"], cache);      tables.append(t); QUICK.preprocess_mode = best\n    t, best = sweep(QUICK, \"balance\", [\"none\", \"class_weights\", \"sampler\"], cache);                 tables.append(t); QUICK.balance = best\n    t, best = sweep(QUICK, \"model_name\", [\"efficientnet_b0\", \"resnet50\", \"densenet121\"], cache);    tables.append(t); QUICK.model_name = best\n    t, best = sweep(QUICK, \"lr_backbone\", [3e-5, 1e-4, 3e-4], cache);                               tables.append(t); QUICK.lr_backbone = best\n    exp = pd.concat(tables, ignore_index=True); exp.to_csv(f\"{cfg.output_dir}/experiments.csv\", index=False)\n    display(exp.round(4))\n    cfg = variant(cfg, preprocess_mode=QUICK.preprocess_mode, balance=QUICK.balance, model_name=QUICK.model_name, lr_backbone=QUICK.lr_backbone)\nelse:\n    # settings chosen by the earlier validation experiments\n    cfg = variant(cfg, preprocess_mode=\"none\", balance=\"none\", model_name=\"densenet121\", lr_backbone=1e-4)\nprint(\"Selected configuration:\", {k: getattr(cfg, k) for k in [\"preprocess_mode\", \"balance\", \"model_name\", \"lr_backbone\"]})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:40:04.000242Z","iopub.execute_input":"2026-09-29T19:40:04.000573Z","iopub.status.idle":"2026-09-29T19:40:04.008597Z","shell.execute_reply.started":"2026-09-29T19:40:04.000549Z","shell.execute_reply":"2026-09-29T19:40:04.007826Z"}},"outputs":[],"execution_count":null},{"id":"bbbb1983","cell_type":"markdown","source":"## 6. Final training with the selected configuration","metadata":{}},{"id":"fc6f3759","cell_type":"code","source":"train_df, val_df, test_df, loaders, image_dir = prepare_data(cfg, df)\nhist, model, best = run_experiment(cfg, train_df, val_df, loaders, \"final\")\nprint(\"best epoch:\", int(best[\"epoch\"]), \"| val macro-F1: %.3f\" % best[\"val_f1\"])\nE.plot_history(hist, f\"{FIG}/08_training_curves.png\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:40:04.044138Z","iopub.execute_input":"2026-09-29T19:40:04.044467Z","iopub.status.idle":"2026-09-29T19:45:26.467129Z","shell.execute_reply.started":"2026-09-29T19:40:04.044445Z","shell.execute_reply":"2026-09-29T19:45:26.465899Z"}},"outputs":[],"execution_count":null},{"id":"1ce6bc9f","cell_type":"markdown","source":"## 7. Evaluation on the held-out test set\nAccuracy, precision, recall, F1 (macro / weighted / per class), quadratic weighted kappa, AUC, confusion matrices,\nand the clinically relevant binary \"referable DR\" (grade ≥ 2) sensitivity / specificity, followed by error analysis.","metadata":{}},{"id":"5c8cd6b0","cell_type":"code","source":"y_true, y_pred, probs = E.predict(model, loaders[2], device)\nmetrics = E.compute_metrics(y_true, y_pred, probs); save_json(metrics, f\"{cfg.output_dir}/final/test_metrics.json\")\ndisplay(pd.Series(metrics).round(4).to_frame(\"test\"))\nrep = E.per_class_report(y_true, y_pred); rep.to_csv(f\"{cfg.output_dir}/final/per_class_report.csv\"); display(rep)\nE.plot_confusion(y_true, y_pred, f\"{FIG}/09_confusion.png\"); plt.show()\nE.plot_metric_bars(rep, f\"{FIG}/10_per_class_metrics.png\"); plt.show()\nE.plot_roc(y_true, probs, f\"{FIG}/11_roc.png\"); plt.show()\nE.plot_worst_errors(loaders[2].dataset, y_true, y_pred, probs, k=8, path=f\"{FIG}/12_errors.png\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:45:26.472485Z","iopub.execute_input":"2026-09-29T19:45:26.472789Z","iopub.status.idle":"2026-09-29T19:45:31.338285Z","shell.execute_reply.started":"2026-09-29T19:45:26.472747Z","shell.execute_reply":"2026-09-29T19:45:31.335484Z"}},"outputs":[],"execution_count":null},{"id":"0ac0c0ea","cell_type":"markdown","source":"## 8. Explainability (Grad-CAM)\nHeat-maps show where the network looked. In a clinical setting this lets a doctor verify that the model focuses on lesions\n(haemorrhages, exudates, neovascularisation) rather than image artefacts.","metadata":{}},{"id":"66e371c5","cell_type":"code","source":"gc = GradCAM(model.to(device), gradcam_layer(model, cfg.model_name)); tf = get_eval_transform(cfg.img_size)\nds = loaders[2].dataset\nfig, ax = plt.subplots(2, 5, figsize=(17, 7))\nfor c in range(5):\n    idxs = np.where((y_true == c) & (y_pred == c))[0]\n    if len(idxs) == 0: idxs = np.where(y_true == c)[0]          # class never predicted correctly -> show any example\n    i = idxs[0]; img = Image.open(ds.path_of(i)).convert(\"RGB\")\n    cam, k, p = gc(tf(img).unsqueeze(0).to(device))\n    ax[0, c].imshow(img); ax[0, c].set_title(f\"true: {CLASS_NAMES[c]}\", fontsize=9); ax[0, c].axis(\"off\")\n    ax[1, c].imshow(overlay(np.array(img.resize((cfg.img_size, cfg.img_size))), cam)); ax[1, c].set_title(f\"pred: {CLASS_NAMES[k]} ({p[k]:.2f})\", fontsize=9); ax[1, c].axis(\"off\")\nplt.tight_layout(); plt.savefig(f\"{FIG}/13_gradcam.png\", dpi=140, bbox_inches=\"tight\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:45:31.380884Z","iopub.execute_input":"2026-09-29T19:45:31.381198Z","iopub.status.idle":"2026-09-29T19:45:34.101337Z","shell.execute_reply.started":"2026-09-29T19:45:31.381171Z","shell.execute_reply":"2026-09-29T19:45:34.099082Z"}},"outputs":[],"execution_count":null},{"id":"b689f9c7","cell_type":"markdown","source":"## 9. Export for the demo app\n`outputs/final/best_model.pt` is loaded by `app.py` (Gradio) — record this app in operation for the required demo video.","metadata":{}},{"id":"f81b71ff","cell_type":"code","source":"print(os.listdir(f\"{cfg.output_dir}/final\")); print(\"Figures:\", sorted(os.listdir(FIG)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-29T19:45:34.11753Z","iopub.execute_input":"2026-09-29T19:45:34.117839Z","iopub.status.idle":"2026-09-29T19:45:34.123492Z"}},"outputs":[],"execution_count":null},{"id":"c4ebc2db-8fe6-4e40-8f12-d9368bb53dac","cell_type":"code","source":"def plot_history_wide(hist, path=None, smooth=3):\n    fig, ax = plt.subplots(1, 3, figsize=(16, 4.4))\n    for a, (k, title, lim) in zip(ax[:2], [(\"loss\", \"Loss\", (0, 1.4)), (\"acc\", \"Accuracy\", (0.5, 1.0))]):\n        for col, name, c in [(f\"train_{k}\", \"train\", \"tab:blue\"), (f\"val_{k}\", \"validation\", \"tab:orange\")]:\n            a.plot(hist[\"epoch\"], hist[col], \"o\", color=c, alpha=.35, ms=4)\n            a.plot(hist[\"epoch\"], hist[col].rolling(smooth, min_periods=1).mean(), \"-\", color=c, lw=2, label=name + \" (smoothed)\")\n        fine = hist.index[hist[\"phase\"] == \"finetune\"]\n        if len(fine) and fine[0] > 0:\n            a.axvline(hist.loc[fine[0], \"epoch\"] - 0.5, color=\"grey\", ls=\"--\", lw=1)\n        a.set_ylim(*lim); a.set_title(f\"{title} curve\"); a.set_xlabel(\"epoch\"); a.set_ylabel(title); a.legend(); a.grid(alpha=.3)\n    for col, name, c in [(\"val_f1\", \"val macro-F1\", \"tab:blue\"), (\"val_qwk\", \"val QWK\", \"tab:orange\")]:\n        ax[2].plot(hist[\"epoch\"], hist[col], \"o\", color=c, alpha=.35, ms=4)\n        ax[2].plot(hist[\"epoch\"], hist[col].rolling(smooth, min_periods=1).mean(), \"-\", color=c, lw=2, label=name + \" (smoothed)\")\n    ax[2].set_ylim(0, 1.0); ax[2].set_title(\"Validation F1 / kappa\"); ax[2].set_xlabel(\"epoch\"); ax[2].legend(); ax[2].grid(alpha=.3)\n    plt.tight_layout()\n    if path: plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    return fig\n\nplot_history_wide(hist, f\"{FIG}/08_training_curves.png\"); plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"da1d7e24-4823-4c9a-9daf-71e39f1c5762","cell_type":"code","source":"# Extra run: same final setup, but WITH contrast + edge enhancement (CLAHE + unsharp mask)\ncfg_enh = variant(cfg, preprocess_mode=\"clahe_unsharp\")\ntr2, va2, te2, loaders2, _ = prepare_data(cfg_enh, df)          # preprocesses + caches enhanced images\nhist2, model2, best2 = run_experiment(cfg_enh, tr2, va2, loaders2, \"final_clahe\")\nprint(\"best epoch:\", int(best2[\"epoch\"]), \"| val macro-F1: %.3f\" % best2[\"val_f1\"])\n\nyt, yp, pr = E.predict(model2, loaders2[2], device)\nm2 = E.compute_metrics(yt, yp, pr); save_json(m2, f\"{cfg.output_dir}/final_clahe/test_metrics.json\")\ncmp = pd.DataFrame({\"none (final)\": metrics, \"clahe_unsharp\": m2}).round(4)\ncmp.to_csv(f\"{cfg.output_dir}/final_comparison.csv\"); display(cmp)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}