{"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":"ec2b5d67","cell_type":"markdown","source":"# Diabetic Retinopathy Stage Detection – APTOS 2019\n**Module:** Computer Vision  |  **Dataset:** [APTOS 2019 Blindness Detection (Kaggle)](https://www.kaggle.com/c/aptos2019-blindness-detection)\n\n**Pipeline**\n```\nFundus image → Crop black border → Resize 224 → Median denoise → CLAHE (contrast)\n             → Unsharp mask (edge enhancement) → Circle mask\n             → Augmentation (train only) → EfficientNetB0 (ImageNet, transfer learning)\n             → 5-class softmax (stage 0-4) → DR / No-DR decision (stage > 0)\n             → Evaluation (acc, P, R, F1, QWK, confusion matrix, ROC) → Grad-CAM explanation\n```\n**Run on Kaggle:** Accelerator = GPU (T4/P100), Internet = ON. Then *Run All*.","metadata":{}},{"id":"ccb8bae0","cell_type":"markdown","source":"## 0. Setup & configuration","metadata":{}},{"id":"026b257b","cell_type":"code","source":"import os, glob, json, random, time\nimport numpy as np, pandas as pd\nimport matplotlib.pyplot as plt, seaborn as sns\nimport cv2\nimport tensorflow as tf, keras\nfrom keras import layers\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import (classification_report, confusion_matrix, cohen_kappa_score,\n                             accuracy_score, precision_recall_fscore_support,\n                             roc_auc_score, roc_curve)\nfrom joblib import Parallel, delayed\n\n# ---------------- configuration ----------------\nSEED        = 42\nIMG_SIZE    = 224\nBATCH       = 32\nHEAD_EPOCHS = 5      # stage 1: train only the new classifier head\nFT_EPOCHS   = 25     # stage 2: fine-tune the whole network (early stopping will stop earlier)\nHEAD_LR     = 1e-3\nFT_LR       = 1e-4\nWEIGHTS     = \"imagenet\"   # transfer learning from ImageNet\nRUN_BASELINE = True  # also train a quick MobileNetV2 baseline for comparison (~5 min)\nOUT = \"/kaggle/working\" if os.path.exists(\"/kaggle/working\") else \"./outputs\"\nORG_DIR = \"/tmp/organized\"   # organised copy of the dataset: split/stage/image.png\nos.makedirs(f\"{OUT}/figures\", exist_ok=True)\n\n\nBASE = os.environ.get(\"DR_BASE\") or os.path.dirname(\n    [p for p in glob.glob(\"/kaggle/input/**/train.csv\", recursive=True) if \"aptos\" in p.lower()][0])\nTRAIN_DIR = os.path.join(BASE, \"train_images\")\n\nSTAGES = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative DR\"]\n\ndef seed_everything(s=SEED):\n    random.seed(s); np.random.seed(s); tf.random.set_seed(s); os.environ[\"PYTHONHASHSEED\"] = str(s)\nseed_everything()\n\ndef savefig(name):\n    plt.savefig(f\"{OUT}/figures/{name}.png\", dpi=150, bbox_inches=\"tight\")\n\nprint(\"TensorFlow\", tf.__version__, \"| Keras\", keras.__version__)\nprint(\"GPU:\", tf.config.list_physical_devices(\"GPU\"))\nprint(\"Data:\", BASE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:12:27.281326Z","iopub.execute_input":"2026-09-26T13:12:27.281755Z","iopub.status.idle":"2026-09-26T13:12:35.818166Z","shell.execute_reply.started":"2026-09-26T13:12:27.281725Z","shell.execute_reply":"2026-09-26T13:12:35.817419Z"}},"outputs":[],"execution_count":null},{"id":"39d99c7b","cell_type":"markdown","source":"## 1. Exploratory Data Analysis","metadata":{}},{"id":"f76093ba","cell_type":"code","source":"df = pd.read_csv(os.path.join(BASE, \"train.csv\"))\ndf[\"path\"] = TRAIN_DIR + \"/\" + df[\"id_code\"] + \".png\"\nprint(df.shape); display(df.head())\n\ncounts = df[\"diagnosis\"].value_counts().sort_index()\npd.DataFrame({\"stage\": STAGES, \"images\": counts.values,\n              \"percent\": (counts.values / len(df) * 100).round(1)})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:12:45.281304Z","iopub.execute_input":"2026-09-26T13:12:45.281703Z","iopub.status.idle":"2026-09-26T13:12:45.396237Z","shell.execute_reply.started":"2026-09-26T13:12:45.281673Z","shell.execute_reply":"2026-09-26T13:12:45.395411Z"}},"outputs":[],"execution_count":null},{"id":"defbee16","cell_type":"code","source":"plt.figure(figsize=(8, 4))\nbars = plt.bar(STAGES, counts.values, color=sns.color_palette(\"crest\", 5))\nplt.bar_label(bars); plt.ylabel(\"Images\"); plt.title(\"APTOS 2019 – class distribution\")\nsavefig(\"01_class_distribution\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:12:50.87752Z","iopub.execute_input":"2026-09-26T13:12:50.878191Z","iopub.status.idle":"2026-09-26T13:12:51.252373Z","shell.execute_reply.started":"2026-09-26T13:12:50.878161Z","shell.execute_reply":"2026-09-26T13:12:51.251641Z"}},"outputs":[],"execution_count":null},{"id":"3781c0d7","cell_type":"code","source":"fig, ax = plt.subplots(5, 4, figsize=(12, 15))\nfor c in range(5):\n    for j, p in enumerate(df[df.diagnosis == c].sample(4, random_state=SEED).path):\n        ax[c, j].imshow(cv2.cvtColor(cv2.imread(p), cv2.COLOR_BGR2RGB)); ax[c, j].axis(\"off\")\n        ax[c, j].set_title(f\"{c}: {STAGES[c]}\", fontsize=9)\nplt.suptitle(\"Raw samples per stage\"); plt.tight_layout(); savefig(\"02_samples_per_stage\"); plt.show()\n\nsizes = np.array([cv2.imread(p).shape[:2] for p in df.path.sample(min(300, len(df)), random_state=SEED)])\nprint(\"Height range:\", sizes[:, 0].min(), \"-\", sizes[:, 0].max(),\n      \"| Width range:\", sizes[:, 1].min(), \"-\", sizes[:, 1].max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:12:56.73061Z","iopub.execute_input":"2026-09-26T13:12:56.731202Z","iopub.status.idle":"2026-09-26T13:14:23.925545Z","shell.execute_reply.started":"2026-09-26T13:12:56.731167Z","shell.execute_reply":"2026-09-26T13:14:23.924762Z"}},"outputs":[],"execution_count":null},{"id":"e54d2a58","cell_type":"markdown","source":"## 2. Image preprocessing\nSame code as `src/preprocess.py` in the GitHub repo (the Gradio app uses the identical pipeline).","metadata":{}},{"id":"0cf2bc33","cell_type":"code","source":"def crop_black_border(img, tol=7):\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY); mask = gray > tol\n    if not mask.any(): return img\n    r = np.where(mask.any(1))[0]; c = np.where(mask.any(0))[0]\n    return img[r[ 0]:r[-1] + 1, c[0]:c[-1] + 1]\n\ndef apply_circle_mask(img):\n    h, w = img.shape[:2]; m = np.zeros((h, w), np.uint8)\n    cv2.circle(m, (w // 2, h // 2), int(min(h, w) * 0.49), 255, -1)\n    return cv2.bitwise_and(img, img, mask=m)\n\ndef apply_clahe(img, clip=2.0, grid=8):\n    l, a, b = cv2.split(cv2.cvtColor(img, cv2.COLOR_RGB2LAB))\n    l = cv2.createCLAHE(clipLimit=clip, tileGridSize=(grid, grid)).apply(l)\n    return cv2.cvtColor(cv2.merge((l, a, b)), cv2.COLOR_LAB2RGB)\n\ndef unsharp_mask(img, sigma=2.0, amount=1.0):\n    return cv2.addWeighted(img, 1 + amount, cv2.GaussianBlur(img, (0, 0), sigma), -amount, 0)\n\ndef ben_graham(img, sigma=10):\n    return cv2.addWeighted(img, 4, cv2.GaussianBlur(img, (0, 0), sigma), -4, 128)\n\ndef preprocess_image(img, size=IMG_SIZE):\n    img = crop_black_border(img)\n    img = cv2.resize(img, (size, size), interpolation=cv2.INTER_AREA)\n    img = cv2.medianBlur(img, 3)          # noise removal\n    img = apply_clahe(img)                # contrast enhancement\n    img = unsharp_mask(img)               # edge enhancement\n    return apply_circle_mask(img)\n\ndef load_and_preprocess(path, size=IMG_SIZE):\n    return preprocess_image(cv2.cvtColor(cv2.imread(path), cv2.COLOR_BGR2RGB), size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:27:47.3737Z","iopub.execute_input":"2026-09-26T13:27:47.374463Z","iopub.status.idle":"2026-09-26T13:27:47.383879Z","shell.execute_reply.started":"2026-09-26T13:27:47.374433Z","shell.execute_reply":"2026-09-26T13:27:47.383294Z"}},"outputs":[],"execution_count":null},{"id":"5cae132e","cell_type":"code","source":"# Step-by-step visual evidence of every preprocessing stage (for the report)\np = df[df.diagnosis == 2].path.iloc[0]\nraw = cv2.cvtColor(cv2.imread(p), cv2.COLOR_BGR2RGB)\ns1 = crop_black_border(raw)\ns2 = cv2.resize(s1, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_AREA)\ns3 = cv2.medianBlur(s2, 3)\ns4 = apply_clahe(s3)\ns5 = unsharp_mask(s4)\ns6 = apply_circle_mask(s5)\nedges = cv2.Canny(cv2.cvtColor(s6, cv2.COLOR_RGB2GRAY), 40, 120)\nsteps = [(raw, \"1. Raw\"), (s1, \"2. Cropped\"), (s2, \"3. Resized 224\"), (s3, \"4. Median denoise\"),\n         (s4, \"5. CLAHE contrast\"), (s5, \"6. Unsharp (edges)\"), (s6, \"7. Final (circle mask)\"),\n         (ben_graham(s2), \"Alt: Ben Graham\"), (edges, \"Canny edges of final\")]\nfig, ax = plt.subplots(1, len(steps), figsize=(24, 3.4))\nfor a, (im, t) in zip(ax, steps):\n    a.imshow(im, cmap=\"gray\" if im.ndim == 2 else None); a.set_title(t, fontsize=9); a.axis(\"off\")\nplt.tight_layout(); savefig(\"03_preprocessing_steps\"); plt.show()\n\n# Histogram of the green channel before/after CLAHE (shows the contrast improvement)\nplt.figure(figsize=(7, 3))\nplt.hist(s3[..., 1].ravel(), 64, alpha=.6, label=\"before CLAHE\")\nplt.hist(s4[..., 1].ravel(), 64, alpha=.6, label=\"after CLAHE\")\nplt.legend(); plt.title(\"Green-channel histogram\"); savefig(\"04_clahe_histogram\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:28:04.22023Z","iopub.execute_input":"2026-09-26T13:28:04.221067Z","iopub.status.idle":"2026-09-26T13:28:08.466977Z","shell.execute_reply.started":"2026-09-26T13:28:04.221038Z","shell.execute_reply":"2026-09-26T13:28:08.466082Z"}},"outputs":[],"execution_count":null},{"id":"8d23eadc","cell_type":"code","source":"# Preprocess every image ONCE and cache as a uint8 array (makes training fast)\ncache = f\"{OUT}/X_{IMG_SIZE}.npy\"\nt0 = time.time()\nif os.path.exists(cache):\n    X = np.load(cache)\nelse:\n    X = np.stack(Parallel(n_jobs=-1)(delayed(load_and_preprocess)(p) for p in df.path))\n    np.save(cache, X)\ny = df[\"diagnosis\"].values\nprint(X.shape, X.dtype, f\"{time.time() - t0:.0f}s\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:28:25.176545Z","iopub.execute_input":"2026-09-26T13:28:25.177211Z","iopub.status.idle":"2026-09-26T13:32:32.89347Z","shell.execute_reply.started":"2026-09-26T13:28:25.177179Z","shell.execute_reply":"2026-09-26T13:32:32.892726Z"}},"outputs":[],"execution_count":null},{"id":"176c56ca","cell_type":"markdown","source":"## 3. Train / validation / test split (stratified 70 / 15 / 15)","metadata":{}},{"id":"1df3e82e","cell_type":"code","source":"idx = np.arange(len(df))\ntr_idx, tmp_idx = train_test_split(idx, test_size=0.30, stratify=y, random_state=SEED)\nva_idx, te_idx = train_test_split(tmp_idx, test_size=0.50, stratify=y[tmp_idx], random_state=SEED)\nsplit = pd.DataFrame({n: np.bincount(y[i], minlength=5) for n, i in\n                      [(\"train\", tr_idx), (\"val\", va_idx), (\"test\", te_idx)]}, index=STAGES)\ndisplay(split)\nsplit.plot.bar(figsize=(8, 3.5), rot=0, title=\"Class distribution per split\")\nsavefig(\"05_split_distribution\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:39:09.733084Z","iopub.execute_input":"2026-09-26T13:39:09.733574Z","iopub.status.idle":"2026-09-26T13:39:10.102421Z","shell.execute_reply.started":"2026-09-26T13:39:09.733544Z","shell.execute_reply":"2026-09-26T13:39:10.101794Z"}},"outputs":[],"execution_count":null},{"id":"7509d5ab","cell_type":"markdown","source":"## 4. Data augmentation & class balancing\n* **Augmentation (train only, on-the-fly):** flips, rotation, zoom, brightness, contrast.\n  Fundus images have no fixed orientation, so flips/rotations are label-preserving.\n* **Balancing:** the dataset is highly imbalanced (No DR ≈ 49 %, Severe ≈ 5 %).\n  We use **class weights** in the loss so minority stages are penalised more when misclassified,\n  plus **label smoothing** to reduce over-confidence on noisy labels.","metadata":{}},{"id":"980857ca","cell_type":"code","source":"augment = keras.Sequential([\n    layers.RandomFlip(\"horizontal_and_vertical\"),\n    layers.RandomRotation(0.25, fill_mode=\"constant\"),\n    layers.RandomZoom((-0.1, 0.15), fill_mode=\"constant\"),\n    layers.RandomBrightness(0.15, value_range=(0, 255)),\n    layers.RandomContrast(0.15),\n], name=\"augmentation\")\n\ncw = compute_class_weight(\"balanced\", classes=np.arange(5), y=y[tr_idx])\nclass_weight = {i: float(w) for i, w in enumerate(cw)}\nprint(\"Class weights:\", {STAGES[k]: round(v, 2) for k, v in class_weight.items()})\n\ndef make_ds(indices, training):\n    ds = tf.data.Dataset.from_tensor_slices((X[indices], tf.one_hot(y[indices], 5)))\n    if training: ds = ds.shuffle(len(indices), seed=SEED)\n    ds = ds.map(lambda a, b: (tf.cast(a, tf.float32), b), num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.batch(BATCH)\n    if training:\n        ds = ds.map(lambda a, b: (augment(a, training=True), b), num_parallel_calls=tf.data.AUTOTUNE)\n    return ds.prefetch(tf.data.AUTOTUNE)\n\ntrain_ds, val_ds, test_ds = make_ds(tr_idx, True), make_ds(va_idx, False), make_ds(te_idx, False)\n\n# visual evidence of augmentation\nsample = tf.cast(X[tr_idx[:1]], tf.float32)\nfig, ax = plt.subplots(1, 6, figsize=(18, 3))\nax[0].imshow(X[tr_idx[0]]); ax[0].set_title(\"Original\"); ax[0].axis(\"off\")\nfor a in ax[1:]:\n    a.imshow(np.clip(augment(sample, training=True)[0].numpy(), 0, 255).astype(\"uint8\"))\n    a.set_title(\"Augmented\"); a.axis(\"off\")\nsavefig(\"06_augmentation_examples\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:39:16.770291Z","iopub.execute_input":"2026-09-26T13:39:16.7711Z","iopub.status.idle":"2026-09-26T13:39:24.740483Z","shell.execute_reply.started":"2026-09-26T13:39:16.771069Z","shell.execute_reply":"2026-09-26T13:39:24.73977Z"}},"outputs":[],"execution_count":null},{"id":"1f94da6c","cell_type":"markdown","source":"## 5. CNN architecture – EfficientNetB0 with transfer learning\n**Why EfficientNetB0?** Compound scaling gives high ImageNet accuracy with only ~4 M parameters\n(fast on a free GPU, small enough for cloud deployment). It already contains its own input\nrescaling/normalisation, so images are fed as 0–255 floats.\n\nHead: GlobalAveragePooling → BatchNorm → Dropout(0.4) → Dense(256, ReLU) → Dropout(0.3) → Dense(5, softmax)","metadata":{}},{"id":"9731ab84","cell_type":"code","source":"def build_model(backbone=\"efficientnet\", weights=WEIGHTS):\n    inp = keras.Input((IMG_SIZE, IMG_SIZE, 3), name=\"image\")\n    if backbone == \"efficientnet\":\n        base = keras.applications.EfficientNetB0(include_top=False, weights=weights, input_tensor=inp)\n    else:  # MobileNetV2 baseline expects inputs in [-1, 1]\n        x0 = layers.Rescaling(1 / 127.5, offset=-1)(inp)\n        base = keras.applications.MobileNetV2(include_top=False, weights=weights, input_tensor=x0)\n    x = layers.GlobalAveragePooling2D(name=\"gap\")(base.output)\n    x = layers.BatchNormalization()(x)\n    x = layers.Dropout(0.4)(x)\n    x = layers.Dense(256, activation=\"relu\")(x)\n    x = layers.Dropout(0.3)(x)\n    out = layers.Dense(5, activation=\"softmax\", name=\"stage\")(x)\n    return keras.Model(inp, out, name=f\"DR_{backbone}\"), base\n\ndef set_base_trainable(base, trainable):\n    for l in base.layers:\n        # keep BatchNorm frozen during fine-tuning (standard practice for small datasets)\n        l.trainable = trainable and not isinstance(l, layers.BatchNormalization)\n\ndef compile_model(m, lr):\n    m.compile(optimizer=keras.optimizers.AdamW(lr, weight_decay=1e-4),\n              loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.05),\n              metrics=[\"accuracy\"])\n\nmodel, base = build_model(\"efficientnet\")\nprint(f\"Total params: {model.count_params():,}  | backbone layers: {len(base.layers)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:40:05.09446Z","iopub.execute_input":"2026-09-26T13:40:05.09512Z","iopub.status.idle":"2026-09-26T13:40:06.672132Z","shell.execute_reply.started":"2026-09-26T13:40:05.095088Z","shell.execute_reply":"2026-09-26T13:40:06.671342Z"}},"outputs":[],"execution_count":null},{"id":"82b00c4f","cell_type":"markdown","source":"## 6. Training strategy (two-stage fine-tuning)\n1. **Stage 1 – head only:** backbone frozen, LR 1e-3, 5 epochs → the new layers learn without destroying ImageNet features.\n2. **Stage 2 – fine-tune:** backbone unfrozen (BatchNorm kept frozen), LR 1e-4 with AdamW.\n\nCallbacks: **EarlyStopping** (patience 6, restore best), **ReduceLROnPlateau** (×0.3, patience 2),\n**ModelCheckpoint** (best val_loss), **CSVLogger**.","metadata":{}},{"id":"d4687c63","cell_type":"code","source":"def callbacks(tag):\n    return [\n        keras.callbacks.EarlyStopping(monitor=\"val_loss\", patience=6, restore_best_weights=True, verbose=1),\n        keras.callbacks.ReduceLROnPlateau(monitor=\"val_loss\", factor=0.3, patience=2, min_lr=1e-7, verbose=1),\n        keras.callbacks.ModelCheckpoint(f\"{OUT}/{tag}_best.keras\", monitor=\"val_loss\", save_best_only=True),\n        keras.callbacks.CSVLogger(f\"{OUT}/{tag}_log.csv\", append=True),\n    ]\n\ndef train_two_stage(m, b, tag, head_epochs=HEAD_EPOCHS, ft_epochs=FT_EPOCHS):\n    t0 = time.time()\n    set_base_trainable(b, False); compile_model(m, HEAD_LR)\n    h1 = m.fit(train_ds, validation_data=val_ds, epochs=head_epochs,\n               class_weight=class_weight, callbacks=callbacks(tag), verbose=2)\n    set_base_trainable(b, True); compile_model(m, FT_LR)\n    h2 = m.fit(train_ds, validation_data=val_ds, epochs=head_epochs + ft_epochs,\n               initial_epoch=head_epochs, class_weight=class_weight,\n               callbacks=callbacks(tag), verbose=2)\n    hist = {k: h1.history[k] + h2.history[k] for k in h1.history if k in h2.history}\n    print(f\"{tag} training time: {(time.time() - t0) / 60:.1f} min\")\n    return hist, head_epochs\n\nhistory, switch_epoch = train_two_stage(model, base, \"efficientnet\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:40:22.014384Z","iopub.execute_input":"2026-09-26T13:40:22.014827Z","iopub.status.idle":"2026-09-26T13:49:19.201346Z","shell.execute_reply.started":"2026-09-26T13:40:22.014797Z","shell.execute_reply":"2026-09-26T13:49:19.20051Z"}},"outputs":[],"execution_count":null},{"id":"ff6ad05a","cell_type":"code","source":"fig, ax = plt.subplots(1, 2, figsize=(13, 4))\nfor a, k in zip(ax, [\"accuracy\", \"loss\"]):\n    a.plot(history[k], label=\"train\"); a.plot(history[\"val_\" + k], label=\"validation\")\n    a.axvline(switch_epoch - 0.5, ls=\"--\", c=\"gray\", label=\"start fine-tuning\")\n    a.set_title(f\"EfficientNetB0 – {k}\"); a.set_xlabel(\"epoch\"); a.legend()\nsavefig(\"07_accuracy_loss_curves\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:52:03.771401Z","iopub.execute_input":"2026-09-26T13:52:03.771846Z","iopub.status.idle":"2026-09-26T13:52:04.322411Z","shell.execute_reply.started":"2026-09-26T13:52:03.771814Z","shell.execute_reply":"2026-09-26T13:52:04.321701Z"}},"outputs":[],"execution_count":null},{"id":"f8df3175","cell_type":"markdown","source":"## 7. Evaluation on the held-out test set","metadata":{}},{"id":"0d727e7a","cell_type":"code","source":"def predict_tta(m, indices):\n    # Test-Time Augmentation: average predictions of original + horizontally flipped image\n    imgs = X[indices].astype(\"float32\")\n    p1 = m.predict(imgs, batch_size=BATCH, verbose=0)\n    p2 = m.predict(imgs[:, :, ::-1], batch_size=BATCH, verbose=0)\n    return (p1 + p2) / 2\n\nprobs = predict_tta(model, te_idx)\ny_true, y_pred = y[te_idx], probs.argmax(1)\n\nacc = accuracy_score(y_true, y_pred)\nqwk = cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\nP, R, F, _ = precision_recall_fscore_support(y_true, y_pred, average=\"macro\", zero_division=0)\nprint(f\"Accuracy {acc:.4f} | Macro P {P:.4f} R {R:.4f} F1 {F:.4f} | Quadratic Weighted Kappa {qwk:.4f}\\n\")\nprint(classification_report(y_true, y_pred, target_names=STAGES, digits=4, zero_division=0))\n\n# Binary task: DR present (stage 1-4) vs No DR\nyb_true, yb_prob = (y_true > 0).astype(int), 1 - probs[:, 0]\nyb_pred = (yb_prob > 0.5).astype(int)\nPb, Rb, Fb, _ = precision_recall_fscore_support(yb_true, yb_pred, average=\"binary\")\naucb = roc_auc_score(yb_true, yb_prob)\nprint(f\"DR vs No-DR → Acc {accuracy_score(yb_true, yb_pred):.4f} | Precision {Pb:.4f} | \"\n      f\"Recall (sensitivity) {Rb:.4f} | F1 {Fb:.4f} | ROC-AUC {aucb:.4f}\")\n\nmetrics = dict(accuracy=acc, macro_precision=P, macro_recall=R, macro_f1=F, qwk=qwk,\n               binary_accuracy=accuracy_score(yb_true, yb_pred), binary_precision=Pb,\n               binary_recall=Rb, binary_f1=Fb, binary_auc=aucb)\njson.dump({k: round(float(v), 4) for k, v in metrics.items()}, open(f\"{OUT}/metrics.json\", \"w\"), indent=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:52:12.237378Z","iopub.execute_input":"2026-09-26T13:52:12.237816Z","iopub.status.idle":"2026-09-26T13:52:34.754731Z","shell.execute_reply.started":"2026-09-26T13:52:12.237784Z","shell.execute_reply":"2026-09-26T13:52:34.754039Z"}},"outputs":[],"execution_count":null},{"id":"d3da7c61","cell_type":"code","source":"cm = confusion_matrix(y_true, y_pred)\nfig, ax = plt.subplots(1, 3, figsize=(20, 5))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=STAGES, yticklabels=STAGES, ax=ax[0])\nax[0].set_title(\"Confusion matrix (counts)\")\nsns.heatmap(cm / cm.sum(1, keepdims=True), annot=True, fmt=\".2f\", cmap=\"Blues\",\n            xticklabels=STAGES, yticklabels=STAGES, ax=ax[1])\nax[1].set_title(\"Confusion matrix (row-normalised = recall)\")\nfor a in ax[:2]: a.set_xlabel(\"Predicted\"); a.set_ylabel(\"True\")\nfpr, tpr, _ = roc_curve(yb_true, yb_prob)\nax[2].plot(fpr, tpr, label=f\"AUC = {aucb:.3f}\"); ax[2].plot([0, 1], [0, 1], \"--\", c=\"gray\")\nax[2].set_title(\"ROC – DR vs No DR\"); ax[2].set_xlabel(\"FPR\"); ax[2].set_ylabel(\"TPR\"); ax[2].legend()\nplt.tight_layout(); savefig(\"08_confusion_matrix_roc\"); plt.show()\n\nrep = pd.DataFrame(classification_report(y_true, y_pred, target_names=STAGES,\n                                         output_dict=True, zero_division=0)).T.iloc[:5]\nrep[[\"precision\", \"recall\", \"f1-score\"]].plot.bar(figsize=(9, 3.5), rot=0, ylim=(0, 1),\n                                                  title=\"Per-stage precision / recall / F1\")\nsavefig(\"09_per_class_metrics\"); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:58:08.058715Z","iopub.execute_input":"2026-09-26T13:58:08.059384Z","iopub.status.idle":"2026-09-26T13:58:09.525054Z","shell.execute_reply.started":"2026-09-26T13:58:08.059352Z","shell.execute_reply":"2026-09-26T13:58:09.524304Z"}},"outputs":[],"execution_count":null},{"id":"76e75620","cell_type":"markdown","source":"## 8. Explainability – Grad-CAM & error analysis\nGrad-CAM highlights the retinal regions that drove the prediction (should focus on lesions:\nhaemorrhages, exudates, neovascularisation) – important for clinical trust.","metadata":{}},{"id":"e2309bce","cell_type":"code","source":"grad_model = keras.Model(model.inputs, [model.get_layer(\"top_activation\").output, model.output])\n\ndef gradcam(img_uint8, class_idx=None):\n    x = tf.convert_to_tensor(img_uint8[None].astype(\"float32\"))\n    with tf.GradientTape() as tape:\n        conv, pred = grad_model(x, training=False)\n        class_idx = int(tf.argmax(pred[0])) if class_idx is None else class_idx\n        score = pred[:, class_idx]\n    g = tape.gradient(score, conv)\n    w = tf.reduce_mean(g, axis=(0, 1, 2))\n    cam = tf.nn.relu(tf.reduce_sum(conv[0] * w, -1)).numpy()\n    cam = cv2.resize(cam / (cam.max() + 1e-8), (IMG_SIZE, IMG_SIZE))\n    heat = cv2.cvtColor(cv2.applyColorMap(np.uint8(255 * cam), cv2.COLORMAP_JET), cv2.COLOR_BGR2RGB)\n    return cv2.addWeighted(img_uint8, 0.6, heat, 0.4, 0), class_idx\n\ndef show_cases(indices, title, fname):\n    if len(indices) == 0:\n        print(\"No cases for:\", title); return\n    fig, ax = plt.subplots(2, len(indices), figsize=(3.2 * len(indices), 6.5), squeeze=False)\n    for j, i in enumerate(indices):\n        overlay, pc = gradcam(X[i])\n        ax[0, j].imshow(X[i]); ax[0, j].set_title(f\"True: {STAGES[y[i]]}\", fontsize=9)\n        ax[1, j].imshow(overlay); ax[1, j].set_title(f\"Pred: {STAGES[pc]}\", fontsize=9)\n        ax[0, j].axis(\"off\"); ax[1, j].axis(\"off\")\n    plt.suptitle(title); plt.tight_layout(); savefig(fname); plt.show()\n\ncorrect = [te_idx[k] for k in range(len(te_idx)) if y_pred[k] == y_true[k] and y_true[k] > 0]\nif not correct:  # fall back to any correct prediction\n    correct = [te_idx[k] for k in range(len(te_idx)) if y_pred[k] == y_true[k]]\nwrong = [te_idx[k] for k in range(len(te_idx)) if y_pred[k] != y_true[k]]\nshow_cases(correct[:5], \"Grad-CAM – correctly classified DR cases\", \"10_gradcam_correct\")\nshow_cases(wrong[:5], \"Error analysis – misclassified cases\", \"11_gradcam_errors\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:58:25.104338Z","iopub.execute_input":"2026-09-26T13:58:25.104628Z","iopub.status.idle":"2026-09-26T13:58:36.737687Z","shell.execute_reply.started":"2026-09-26T13:58:25.104603Z","shell.execute_reply":"2026-09-26T13:58:36.736661Z"}},"outputs":[],"execution_count":null},{"id":"1bd4b464","cell_type":"markdown","source":"## 9. Experiment comparison – EfficientNetB0 vs MobileNetV2 baseline\nJustifies the architecture choice with evidence (same data, same augmentation, same callbacks).","metadata":{}},{"id":"b4783741","cell_type":"code","source":"results = [dict(model=\"EfficientNetB0\", params=model.count_params(), accuracy=acc, macro_f1=F, qwk=qwk)]\nif RUN_BASELINE:\n    bmodel, bbase = build_model(\"mobilenet\")\n    _hist_b, _ = train_two_stage(bmodel, bbase, \"mobilenet\", head_epochs=3, ft_epochs=10)\n    bp = predict_tta(bmodel, te_idx).argmax(1)\n    results.append(dict(model=\"MobileNetV2\", params=bmodel.count_params(),\n                        accuracy=accuracy_score(y_true, bp),\n                        macro_f1=precision_recall_fscore_support(y_true, bp, average=\"macro\", zero_division=0)[2],\n                        qwk=cohen_kappa_score(y_true, bp, weights=\"quadratic\")))\nres = pd.DataFrame(results).round(4); display(res)\nres.to_csv(f\"{OUT}/model_comparison.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T13:58:37.212178Z","iopub.execute_input":"2026-09-26T13:58:37.212792Z","iopub.status.idle":"2026-09-26T14:03:24.263431Z","shell.execute_reply.started":"2026-09-26T13:58:37.212759Z","shell.execute_reply":"2026-09-26T14:03:24.262792Z"}},"outputs":[],"execution_count":null},{"id":"c3e1ba30","cell_type":"markdown","source":"## 10. Save the model for the Gradio app (download these files → put in `models/` in VS Code)","metadata":{}},{"id":"786ac2e4","cell_type":"code","source":"model.save(f\"{OUT}/dr_efficientnetb0.keras\")\nmodel.save_weights(f\"{OUT}/dr_efficientnetb0.weights.h5\")\njson.dump({\"stages\": STAGES, \"img_size\": IMG_SIZE, \"tf\": tf.__version__, \"keras\": keras.__version__},\n          open(f\"{OUT}/model_info.json\", \"w\"), indent=2)\nos.system(f\"cd {OUT} && zip -qr figures.zip figures\")\nprint(\"Saved. Download from the Output panel: dr_efficientnetb0.keras, dr_efficientnetb0.weights.h5,\"\n      \" model_info.json, metrics.json, figures.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T14:05:50.887192Z","iopub.execute_input":"2026-09-26T14:05:50.88765Z","iopub.status.idle":"2026-09-26T14:05:52.859751Z","shell.execute_reply.started":"2026-09-26T14:05:50.887618Z","shell.execute_reply":"2026-09-26T14:05:52.859028Z"}},"outputs":[],"execution_count":null},{"id":"6dee5d5a-4622-4e23-a0b7-c906643f2d28","cell_type":"code","source":"# Copy 2 unseen TEST images per stage for the Gradio demo\nimport shutil\nex_dir = f\"{OUT}/examples\"\nos.makedirs(ex_dir, exist_ok=True)\nfor c in range(5):\n    for i in [i for i in te_idx if y[i] == c][:2]:\n        shutil.copy(df.path[i], f\"{ex_dir}/stage{c}_{STAGES[c].replace(' ', '_')}_{df.id_code[i]}.png\")\nos.system(f\"cd {OUT} && zip -qr examples.zip examples\")\nprint(sorted(os.listdir(ex_dir)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T17:29:54.568553Z","iopub.execute_input":"2026-09-26T17:29:54.569049Z","iopub.status.idle":"2026-09-26T17:29:55.835929Z","shell.execute_reply.started":"2026-09-26T17:29:54.569015Z","shell.execute_reply":"2026-09-26T17:29:55.835219Z"}},"outputs":[],"execution_count":null},{"id":"18214b55-b579-43fa-baab-98916f178a6b","cell_type":"code","source":"# Export the organised (preprocessed 224x224) dataset: split/stage/image.jpg\nimport shutil\nexp_dir = \"/tmp/organized_dataset\"\nshutil.rmtree(exp_dir, ignore_errors=True)\nfor split_name, ids in [(\"train\", tr_idx), (\"val\", va_idx), (\"test\", te_idx)]:\n    for i in ids:\n        folder = f\"{exp_dir}/{split_name}/{y[i]}_{STAGES[y[i]].replace(' ', '_')}\"\n        os.makedirs(folder, exist_ok=True)\n        cv2.imwrite(f\"{folder}/{df.id_code[i]}.jpg\", cv2.cvtColor(X[i], cv2.COLOR_RGB2BGR),\n                    [cv2.IMWRITE_JPEG_QUALITY, 95])\n\nfor split_name in [\"train\", \"val\", \"test\"]:\n    print(f\"{split_name}/\")\n    for c in sorted(os.listdir(f\"{exp_dir}/{split_name}\")):\n        print(f\"   ├── {c:<20} {len(os.listdir(f'{exp_dir}/{split_name}/{c}')):>5} images\")\n\nshutil.make_archive(f\"{OUT}/organized_dataset\", \"zip\", exp_dir)\nprint(\"zip size:\", round(os.path.getsize(f\"{OUT}/organized_dataset.zip\") / 1e6, 1), \"MB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T17:30:58.709732Z","iopub.execute_input":"2026-09-26T17:30:58.710513Z","iopub.status.idle":"2026-09-26T17:31:03.100233Z","shell.execute_reply.started":"2026-09-26T17:30:58.710481Z","shell.execute_reply":"2026-09-26T17:31:03.09948Z"}},"outputs":[],"execution_count":null},{"id":"e44a87ee-5cf8-463f-a900-c5490f6911da","cell_type":"code","source":"!pip install -q gradio\n\n# ===== RetinaScan UI – runs inside this notebook using the trained model in memory =====\nimport gradio as gr\n\n_grad_model = keras.Model(model.inputs, [model.get_layer(\"top_activation\").output, model.output])\n\nADVICE = {\n    0: \"No signs of diabetic retinopathy. Keep blood sugar, blood pressure and cholesterol controlled; routine eye screening every 12 months.\",\n    1: \"Mild non-proliferative DR (microaneurysms). Tighter diabetes control and a follow-up eye exam in 6–12 months.\",\n    2: \"Moderate non-proliferative DR. Refer to an ophthalmologist; follow-up every 3–6 months.\",\n    3: \"Severe non-proliferative DR – high risk of progression. Urgent ophthalmologist referral recommended.\",\n    4: \"Proliferative DR (new abnormal vessels). Sight-threatening – urgent specialist treatment (laser / anti-VEGF) usually required.\",\n}\nLAST = {\"stage\": None, \"conf\": None}\n\ndef _gradcam(img_uint8):\n    x = tf.convert_to_tensor(img_uint8[None].astype(\"float32\"))\n    with tf.GradientTape() as tape:\n        conv, pred = _grad_model(x, training=False)\n        idx = int(tf.argmax(pred[0]))\n        score = pred[:, idx]\n    g = tape.gradient(score, conv)\n    cam = tf.nn.relu(tf.reduce_sum(conv[0] * tf.reduce_mean(g, axis=(0, 1, 2)), -1)).numpy()\n    cam = cv2.resize(cam / (cam.max() + 1e-8), (IMG_SIZE, IMG_SIZE))\n    heat = cv2.cvtColor(cv2.applyColorMap(np.uint8(255 * cam), cv2.COLORMAP_JET), cv2.COLOR_BGR2RGB)\n    return cv2.addWeighted(img_uint8, 0.6, heat, 0.4, 0), pred[0].numpy()\n\ndef diagnose(image):\n    if image is None:\n        return None, None, {}, \"Please upload a retinal fundus image.\"\n    pre = preprocess_image(image)                 # same pipeline as training\n    overlay, probs = _gradcam(pre)\n    stage, conf = int(np.argmax(probs)), float(np.max(probs))\n    LAST.update(stage=stage, conf=conf)\n    verdict = \"🟢 **No DR detected**\" if stage == 0 else \"🔴 **Diabetic Retinopathy detected**\"\n    text = (f\"{verdict}\\n\\n**Predicted stage:** {stage} – {STAGES[stage]} (confidence {conf:.1%})\\n\\n\"\n            f\"**Probability of any DR:** {1 - probs[0]:.1%}\\n\\n**Guidance:** {ADVICE[stage]}\\n\\n\"\n            \"_Screening aid only – final diagnosis must be made by an eye-care professional._\")\n    return pre, overlay, {f\"{i} – {s}\": float(p) for i, (s, p) in enumerate(zip(STAGES, probs))}, text\n\ndef retinabot(message, history):\n    m, s = message.lower(), LAST[\"stage\"]\n    if any(k in m for k in [\"result\", \"my eye\", \"prediction\", \"diagnos\"]):\n        return (\"Upload an image in the Diagnose tab first.\" if s is None else\n                f\"Your image was classified as **stage {s} – {STAGES[s]}** ({LAST['conf']:.1%} confidence).\\n\\n{ADVICE[s]}\")\n    if any(k in m for k in [\"heat\", \"grad\", \"red\", \"colour\", \"color\"]):\n        return \"The Grad-CAM heat-map shows where the model looked. Red/yellow areas influenced the decision most – ideally lesions such as haemorrhages or exudates.\"\n    if any(k in m for k in [\"what is\", \"explain\", \"diabetic retinopathy\", \"stages\"]):\n        return \"Diabetic retinopathy is damage to retinal blood vessels caused by long-term high blood sugar. Stages: 0 No DR, 1 Mild, 2 Moderate, 3 Severe, 4 Proliferative DR.\"\n    if any(k in m for k in [\"treat\", \"next\", \"doctor\", \"what should\"]):\n        return ADVICE[s] if s is not None else \"Treatment depends on the stage – diabetes control early on; laser or anti-VEGF injections for advanced stages.\"\n    if any(k in m for k in [\"prevent\", \"avoid\", \"risk\"]):\n        return \"Control blood sugar, blood pressure and cholesterol, stop smoking, and have a dilated eye exam every year.\"\n    if any(k in m for k in [\"accura\", \"model\", \"trust\"]):\n        return \"EfficientNetB0 fine-tuned on APTOS 2019: 77.6% stage accuracy, QWK 0.871, 96.7% DR vs No-DR accuracy (97.1% sensitivity).\"\n    return \"Ask me: *what is my result*, *what does the heat-map show*, *what should I do next*, *what is diabetic retinopathy*, *how accurate is the model*.\"\n\nexamples = sorted(glob.glob(f\"{OUT}/examples/*.png\"))\nwith gr.Blocks(title=\"RetinaScan\") as demo:\n    gr.Markdown(\"# 👁️ RetinaScan – Diabetic Retinopathy Stage Detection\\nEfficientNetB0 + transfer learning · APTOS 2019 · Grad-CAM\")\n    with gr.Tab(\"Diagnose\"):\n        with gr.Row():\n            inp = gr.Image(type=\"numpy\", label=\"Upload fundus image\")\n            with gr.Column():\n                out_text = gr.Markdown()\n                out_label = gr.Label(num_top_classes=5, label=\"Stage probabilities\")\n        btn = gr.Button(\"Analyse\", variant=\"primary\")\n        with gr.Row():\n            out_pre = gr.Image(label=\"Preprocessed (CLAHE + edge enhancement)\")\n            out_cam = gr.Image(label=\"Grad-CAM heat-map\")\n        btn.click(diagnose, inp, [out_pre, out_cam, out_label, out_text])\n        if examples:\n            gr.Examples(examples, inp, label=\"Test images (not seen during training)\")\n    with gr.Tab(\"RetinaBot assistant\"):\n        gr.ChatInterface(retinabot, examples=[\"What is my result?\", \"What does the heat-map show?\",\n                                              \"What should I do next?\", \"How accurate is the model?\"])\n\ndemo.launch(share=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-26T17:35:49.41154Z","iopub.execute_input":"2026-09-26T17:35:49.412038Z","iopub.status.idle":"2026-09-26T17:36:13.163141Z","shell.execute_reply.started":"2026-09-26T17:35:49.412004Z","shell.execute_reply":"2026-09-26T17:36:13.162465Z"}},"outputs":[],"execution_count":null}]}