{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"496efc26","cell_type":"markdown","source":"## 1. Setup and configuration\nAll tunable settings live in one cell so experiments are reproducible. Set `QUICK_RUN = True` to check the whole pipeline in a few minutes before the full run.","metadata":{}},{"id":"ee29286b","cell_type":"code","source":"# ---------- Imports ----------\nimport os, glob, json, time, random, math, warnings\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"2\"          # hide TensorFlow info logs\n\nimport numpy as np\nimport pandas as pd\nimport cv2                                         # OpenCV: image processing\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nfrom joblib import Parallel, delayed               # parallel image loading\nimport gc, psutil, shutil                          # memory housekeeping, file copies\n\nimport tensorflow as tf\nimport keras\nfrom keras import layers, regularizers\n\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, accuracy_score,\n                             precision_recall_fscore_support, f1_score, cohen_kappa_score,\n                             roc_auc_score, roc_curve)\n\n# ---------- Configuration ----------\nSEED        = 42\nIMG_SIZE    = 300          # v2: 300 px keeps small lesions (microaneurysms) visible\nBATCH_SIZE  = 16           # fits EfficientNetB3 at 300 px in a 16 GB T4 GPU\nNUM_CLASSES = 5\nCLASS_NAMES = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative DR\"]\n\nQUICK_RUN   = os.environ.get(\"DR_QUICK\", \"0\") == \"1\"   # True = tiny smoke test\nRUN_BACKBONE_COMP = True   # Experiment 1: compare CNN architectures\nRUN_HP_SEARCH     = True   # Experiment 2: learning-rate x dropout grid\nRUN_PREPROC_EXP   = True   # Experiment 3: preprocessing variants, full training\nRUN_BALANCE_EXP   = True   # Experiment 4: class-balancing methods, full training\n\nEPOCHS_SEARCH = 1 if QUICK_RUN else 5    # frozen-backbone runs (Experiments 1-2)\nEPOCHS_COMPACT_HEAD, EPOCHS_COMPACT_FINE = (1, 1) if QUICK_RUN else (3, 10)   # Experiments 3-4\nEPOCHS_HEAD   = 1 if QUICK_RUN else 8    # final phase 1: train classifier head only\nEPOCHS_FINE   = 2 if QUICK_RUN else 30   # final phase 2: fine-tune (early stopping decides)\nFINE_TUNE_FRACTION = 1.0                 # v2: fine-tune the whole backbone (BatchNorm stays frozen)\nFINE_TUNE_LR  = 2e-5\nBALANCE_METHOD = \"balanced\"              # default; Experiment 4 picks the best of sqrt/balanced/oversample\nSELECT_METRIC = \"qwk\"                    # model-selection metric on the VALIDATION set (official APTOS metric)\n\n# Defaults; the experiments below overwrite them with the best settings found on the validation set\nBEST_BACKBONE, BEST_LR, BEST_DROPOUT, PREPROCESS_MODE = \"EfficientNetB3\", 1e-3, 0.3, \"clahe_unsharp\"\n\n# Re-use a model trained in an earlier run of this notebook (attach that run's output via \"Add Input\").\n# Skips the ~2.5 h of experiments and training; evaluation, Grad-CAM and export are redone.\nREUSE_TRAINED_MODEL = True\nRUN_ENSEMBLE = True                      # train extra models and combine them (section 12b)\nHIRES_SIZE = 380                         # Version 4: resolution of the new members (EfficientNetB3/B4 native sizes 300/380)\nBATCH_SIZE_HIRES = 8                     # smaller batches so 380 px models fit in GPU memory\nENSEMBLE_MEMBERS = [(\"EfficientNetB3\", 11, HIRES_SIZE), (\"EfficientNetB4\", 13, HIRES_SIZE)]   # (backbone, seed, image size)\n\n# ImageNet weights need Internet ON in Kaggle (Settings -> Internet)\nWEIGHTS = None if os.environ.get(\"DR_NO_PRETRAINED\") == \"1\" else \"imagenet\"\n\nOUT_DIR = \"/kaggle/working\" if os.path.isdir(\"/kaggle/working\") else \"./outputs\"\nFIG_DIR = os.path.join(OUT_DIR, \"figures\")\nos.makedirs(FIG_DIR, exist_ok=True)\n\n# ---------- Reproducibility ----------\ndef set_seed(seed=SEED):\n    random.seed(seed); np.random.seed(seed); tf.random.set_seed(seed)\n    keras.utils.set_random_seed(seed)\nset_seed()\n\ndef ram_gb():\n    # Current RAM used by this notebook (TensorFlow keeps some memory after each model is built)\n    return round(psutil.Process().memory_info().rss / 1e9, 1)\n\ndef save_fig(name):\n    # Save every figure as PNG so it can go straight into the report\n    path = os.path.join(FIG_DIR, name)\n    plt.savefig(path, dpi=150, bbox_inches=\"tight\")\n    print(\"Saved figure ->\", path)\n\nprint(\"TensorFlow:\", tf.__version__, \"| Keras:\", keras.__version__)\nprint(\"GPUs:\", tf.config.list_physical_devices(\"GPU\"))\nprint(\"Quick run:\", QUICK_RUN, \"| Output folder:\", OUT_DIR)","metadata":{},"outputs":[],"execution_count":null},{"id":"6b096e73","cell_type":"markdown","source":"## 2. Locate and explore the dataset\nThe code searches `/kaggle/input` for `train.csv` + `train_images/`, so it works whichever way the dataset was attached.","metadata":{}},{"id":"294b7747","cell_type":"code","source":"def find_dataset():\n    # Look for train.csv next to a train_images folder anywhere under the input root\n    root = os.environ.get(\"DR_DATA_ROOT\", \"/kaggle/input\")\n    for csv_path in glob.glob(os.path.join(root, \"**\", \"train.csv\"), recursive=True):\n        img_dir = os.path.join(os.path.dirname(csv_path), \"train_images\")\n        if os.path.isdir(img_dir):\n            return csv_path, img_dir\n    raise FileNotFoundError(\"train.csv + train_images not found. Add the APTOS 2019 dataset via 'Add Input'.\")\n\nCSV_PATH, IMG_DIR = find_dataset()\ndf = pd.read_csv(CSV_PATH)                       # columns: id_code, diagnosis\ndf[\"path\"] = df[\"id_code\"].apply(lambda i: os.path.join(IMG_DIR, f\"{i}.png\"))\ndf = df[df[\"path\"].apply(os.path.exists)].reset_index(drop=True)\nif QUICK_RUN:                                    # small stratified subset for smoke tests\n    df = pd.concat([g.sample(min(len(g), 40), random_state=SEED) for _, g in df.groupby(\"diagnosis\")]).reset_index(drop=True)\n\nprint(\"CSV:\", CSV_PATH)\nprint(\"Images:\", len(df))\ndf.head()","metadata":{},"outputs":[],"execution_count":null},{"id":"402799b3","cell_type":"code","source":"# ---------- Class distribution ----------\ncounts = df[\"diagnosis\"].value_counts().sort_index()\ndist = pd.DataFrame({\"stage\": CLASS_NAMES, \"count\": counts.values,\n                     \"percent\": (counts.values / counts.sum() * 100).round(1)})\ndisplay(dist)\nprint(f\"Imbalance ratio (largest / smallest class): {counts.max() / counts.min():.1f} : 1\")\n\nplt.figure(figsize=(8, 4))\nax = sns.barplot(x=CLASS_NAMES, y=counts.values, palette=\"viridis\")\nfor p, c in zip(ax.patches, counts.values):\n    ax.annotate(str(c), (p.get_x() + p.get_width() / 2, p.get_height()), ha=\"center\", va=\"bottom\")\nplt.title(\"Class distribution (APTOS 2019)\"); plt.ylabel(\"Number of images\"); plt.xticks(rotation=15)\nsave_fig(\"01_class_distribution.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"8684eb59","cell_type":"code","source":"# ---------- Image resolution statistics (PIL reads only the header, so this is fast) ----------\nsizes = np.array([Image.open(p).size for p in df[\"path\"].sample(min(400, len(df)), random_state=SEED)])\nprint(f\"Width  range: {sizes[:,0].min()} - {sizes[:,0].max()} px\")\nprint(f\"Height range: {sizes[:,1].min()} - {sizes[:,1].max()} px\")\nprint(\"Distinct resolutions in sample:\", len({tuple(s) for s in sizes}))\n\nplt.figure(figsize=(6, 4))\nplt.scatter(sizes[:, 0], sizes[:, 1], alpha=0.4)\nplt.xlabel(\"Width (px)\"); plt.ylabel(\"Height (px)\"); plt.title(\"Original image resolutions (sample)\")\nsave_fig(\"02_resolutions.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"84120ca8","cell_type":"code","source":"# ---------- Sample raw images per class ----------\ndef read_rgb(path):\n    img = cv2.imread(path)                       # OpenCV loads BGR\n    return cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\nfig, axes = plt.subplots(NUM_CLASSES, 4, figsize=(12, 15))\nfor c in range(NUM_CLASSES):\n    paths = df[df[\"diagnosis\"] == c][\"path\"].sample(4, random_state=SEED, replace=True).values\n    for j, p in enumerate(paths):\n        axes[c, j].imshow(read_rgb(p)); axes[c, j].axis(\"off\")\n        if j == 0: axes[c, j].set_title(f\"{c}: {CLASS_NAMES[c]}\", loc=\"left\", fontsize=12)\nplt.suptitle(\"Raw fundus images by stage\", fontsize=14); plt.tight_layout()\nsave_fig(\"03_raw_samples.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"4d818f0e","cell_type":"markdown","source":"## 3. Image preprocessing\nRaw images vary in resolution, black border size, brightness and camera. The pipeline:\n\n1. **Crop black borders** — remove the uninformative dark frame around the retina.\n2. **Pad to square + resize to 300×300** — keeps the retina round and matches the CNN input (`INTER_AREA` for clean downscaling).\n3. **Noise removal** — 3×3 median filter removes salt-and-pepper sensor noise while keeping edges.\n4. **Contrast enhancement (CLAHE)** — Contrast Limited Adaptive Histogram Equalisation on the L channel of LAB colour space; boosts local contrast of microaneurysms and exudates without shifting colours.\n5. **Edge enhancement (unsharp masking)** — `1.5·img − 0.5·blur` sharpens vessel and lesion boundaries.\n6. **Circular mask** — sets everything outside the retina to black, removing edge artefacts.\n\nAn alternative, **Ben Graham's local-average subtraction** (`4·img − 4·GaussianBlur + 128`, used by the Kaggle competition winner), is also implemented and compared with full training in Experiment 3.\n\nLoading and resizing (the slow part) happens **once**; the enhancement steps run on the cached 300×300 images.","metadata":{}},{"id":"d300fa8c","cell_type":"code","source":"def crop_black_borders(img, tol=7):\n    # Keep only rows/columns that contain pixels brighter than `tol`\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    mask = gray > tol\n    if mask.sum() == 0:\n        return img\n    rows, cols = np.where(mask.any(1))[0], np.where(mask.any(0))[0]\n    return img[rows[0]:rows[-1] + 1, cols[0]:cols[-1] + 1]\n\ndef pad_to_square(img):\n    # Pad with black so the retina stays circular after resizing\n    h, w = img.shape[:2]; s = max(h, w)\n    top, left = (s - h) // 2, (s - w) // 2\n    return cv2.copyMakeBorder(img, top, s - h - top, left, s - w - left, cv2.BORDER_CONSTANT, value=0)\n\ndef load_and_resize(path, size=IMG_SIZE):\n    # Stage A (slow, done once): read -> crop borders -> square -> resize\n    img = read_rgb(path)\n    img = pad_to_square(crop_black_borders(img))\n    return cv2.resize(img, (size, size), interpolation=cv2.INTER_AREA)\n\ndef denoise(img):\n    return cv2.medianBlur(img, 3)\n\ndef apply_clahe(img, clip=2.0, grid=8):\n    # CLAHE on lightness only, so colour information (e.g. red haemorrhages) is preserved\n    lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n    l, a, b = cv2.split(lab)\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=0.5):\n    # Edge enhancement: add back the high-frequency detail (img - blur)\n    blur = cv2.GaussianBlur(img, (0, 0), sigma)\n    return cv2.addWeighted(img, 1 + amount, blur, -amount, 0)\n\ndef ben_graham(img, sigma=None):\n    # Subtract local average colour -> highlights lesions, normalises lighting\n    sigma = sigma or img.shape[0] / 30\n    return cv2.addWeighted(img, 4, cv2.GaussianBlur(img, (0, 0), sigma), -4, 128)\n\ndef circular_mask(img, scale=0.96):\n    h, w = img.shape[:2]\n    mask = np.zeros((h, w), np.uint8)\n    cv2.circle(mask, (w // 2, h // 2), int(min(h, w) / 2 * scale), 1, -1)\n    return img * mask[..., None]\n\ndef enhance(img, mode=\"clahe_unsharp\"):\n    # Stage B (fast): enhancement on a cached IMG_SIZE x IMG_SIZE image\n    if mode == \"raw\":\n        return img\n    img = denoise(img)\n    if mode == \"clahe_unsharp\":\n        img = unsharp_mask(apply_clahe(img))\n    elif mode == \"ben_graham\":\n        img = ben_graham(img)\n    else:\n        raise ValueError(mode)\n    return circular_mask(img)\n\ndef preprocess_fundus(path, mode=\"clahe_unsharp\"):\n    # Full pipeline for a single image (used for inference and in the app)\n    return enhance(load_and_resize(path), mode)","metadata":{},"outputs":[],"execution_count":null},{"id":"8383a1b3","cell_type":"code","source":"# ---------- Visual evidence: every preprocessing step on one image ----------\nsample_path = df[df[\"diagnosis\"] == 2][\"path\"].iloc[0]\norig = read_rgb(sample_path)\ns1 = pad_to_square(crop_black_borders(orig))\ns2 = cv2.resize(s1, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_AREA)\ns3 = denoise(s2)\ns4 = apply_clahe(s3)\ns5 = unsharp_mask(s4)\ns6 = circular_mask(s5)\nbg = circular_mask(ben_graham(s3))\nedges = cv2.Canny(cv2.cvtColor(s6, cv2.COLOR_RGB2GRAY), 40, 120)   # edge map for illustration\n\nsteps = [(orig, \"1. Original\"), (s1, \"2. Border crop + square\"), (s2, f\"3. Resize {IMG_SIZE}x{IMG_SIZE}\"),\n         (s3, \"4. Median denoise\"), (s4, \"5. CLAHE contrast\"), (s5, \"6. Unsharp edge enhance\"),\n         (s6, \"7. Circular mask (final)\"), (bg, \"Alt: Ben Graham\"), (edges, \"Canny edges of final\")]\nfig, axes = plt.subplots(3, 3, figsize=(12, 12))\nfor ax, (im, t) in zip(axes.ravel(), steps):\n    ax.imshow(im, cmap=\"gray\" if im.ndim == 2 else None); ax.set_title(t); ax.axis(\"off\")\nplt.suptitle(\"Preprocessing pipeline, step by step\", fontsize=14); plt.tight_layout()\nsave_fig(\"04_preprocessing_steps.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"150ea96c","cell_type":"code","source":"# ---------- Quantitative evidence: contrast and sharpness before vs after ----------\ndef contrast_rms(img):  return cv2.cvtColor(img, cv2.COLOR_RGB2GRAY).std()\ndef sharpness(img):     return cv2.Laplacian(cv2.cvtColor(img, cv2.COLOR_RGB2GRAY), cv2.CV_64F).var()\n\nrows = []\nfor p in df[\"path\"].sample(min(100, len(df)), random_state=SEED):\n    base = load_and_resize(p)\n    for mode in [\"raw\", \"clahe_unsharp\", \"ben_graham\"]:\n        im = enhance(base, mode)\n        rows.append({\"mode\": mode, \"RMS contrast\": contrast_rms(im), \"Laplacian sharpness\": sharpness(im)})\nquality = pd.DataFrame(rows).groupby(\"mode\").mean().round(2)\ndisplay(quality)","metadata":{},"outputs":[],"execution_count":null},{"id":"a0635d5b","cell_type":"code","source":"# ---------- Histogram of intensities before vs after CLAHE ----------\nplt.figure(figsize=(8, 4))\nplt.hist(cv2.cvtColor(s3, cv2.COLOR_RGB2GRAY).ravel(), bins=64, range=(1, 255), alpha=0.6, label=\"Before CLAHE\")\nplt.hist(cv2.cvtColor(s4, cv2.COLOR_RGB2GRAY).ravel(), bins=64, range=(1, 255), alpha=0.6, label=\"After CLAHE\")\nplt.legend(); plt.title(\"Grey-level histogram (background excluded)\"); plt.xlabel(\"Intensity\"); plt.ylabel(\"Pixels\")\nsave_fig(\"05_clahe_histogram.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"b705b16b","cell_type":"code","source":"# ---------- Stage A for the whole dataset (parallel, once) ----------\nt0 = time.time()\nX_base = np.stack(Parallel(n_jobs=-1)(delayed(load_and_resize)(p) for p in df[\"path\"]))\ny_all = df[\"diagnosis\"].values.astype(int)\nprint(f\"Loaded {X_base.shape} in {time.time() - t0:.0f}s  ({X_base.nbytes / 1e6:.0f} MB)\")\n\ndef build_set(mode):\n    # Stage B for the whole dataset for a given enhancement mode\n    return np.stack([enhance(im, mode) for im in X_base])","metadata":{},"outputs":[],"execution_count":null},{"id":"af506b88","cell_type":"markdown","source":"## 4. Stratified train / validation / test split\n70 % train, 15 % validation, 15 % test, stratified so every split keeps the class ratio. The split is made on image indices **before** any augmentation or oversampling, so no augmented copy of a test image can leak into training.","metadata":{}},{"id":"9c3c4329","cell_type":"code","source":"idx = np.arange(len(df))\ntrain_idx, temp_idx = train_test_split(idx, test_size=0.30, stratify=y_all, random_state=SEED)\nval_idx, test_idx   = train_test_split(temp_idx, test_size=0.50, stratify=y_all[temp_idx], random_state=SEED)\n\nsplit_table = pd.DataFrame({name: np.bincount(y_all[ix], minlength=NUM_CLASSES)\n                            for name, ix in [(\"train\", train_idx), (\"val\", val_idx), (\"test\", test_idx)]},\n                           index=CLASS_NAMES)\nsplit_table.loc[\"Total\"] = split_table.sum()\ndisplay(split_table)\n\nsplit_table.drop(\"Total\").plot(kind=\"bar\", figsize=(9, 4))\nplt.title(\"Images per class in each split\"); plt.ylabel(\"Images\"); plt.xticks(rotation=15)\nsave_fig(\"06_split_distribution.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"b9a5b331","cell_type":"markdown","source":"## 5. Data augmentation and class balancing\n**Augmentation** (training set only, applied on the fly so every epoch sees new variants):\n- Random horizontal + vertical flips and 360° rotation — a retina has no fixed orientation, so this is label-preserving.\n- Random zoom (±15 %) and small translations — simulate different camera framing.\n- Random brightness and contrast (±15 %) — simulate different cameras and lighting.\n\nGeometric transforms fill with black, matching the fundus background.\n\n**Class balancing:** the dataset is heavily imbalanced (No DR ≈ 49 %, Severe ≈ 5 %). Three strategies are implemented and compared with full training in Experiment 4:\n- `sqrt` — class weights `sqrt(n_samples / (n_classes · n_class))`: a milder correction that protects overall accuracy\n- `balanced` — class weights `n_samples / (n_classes · n_class)`: errors on rare stages cost as much as on common ones\n- `oversample` — minority images repeated until every class matches the largest (each repeat is augmented differently)\n\n**Data pipeline:** images stay in one NumPy array; `tf.data` streams batches of *indices* and gathers the pixels, so memory use stays flat whatever the balancing method.","metadata":{}},{"id":"8949169d","cell_type":"code","source":"augmenter = keras.Sequential([\n    layers.RandomFlip(\"horizontal_and_vertical\"),\n    layers.RandomRotation(0.5, fill_mode=\"constant\"),                 # up to ±180°\n    layers.RandomZoom((-0.15, 0.15), fill_mode=\"constant\"),\n    layers.RandomTranslation(0.05, 0.05, fill_mode=\"constant\"),\n    layers.RandomBrightness(0.15, value_range=(0, 255)),\n    layers.RandomContrast(0.15),\n], name=\"augmentation\")\n\n# ---------- Class weights ----------\ncw = compute_class_weight(\"balanced\", classes=np.arange(NUM_CLASSES), y=y_all[train_idx])\nWEIGHT_SETS = {\"balanced\": cw, \"sqrt\": np.sqrt(cw) / np.sqrt(cw).mean()}\nprint(pd.DataFrame(WEIGHT_SETS, index=CLASS_NAMES).round(2))\n\ndef oversample(indices, labels):\n    # Repeat minority-class indices until every class matches the largest class\n    counts = np.bincount(labels[indices], minlength=NUM_CLASSES)\n    rng = np.random.default_rng(SEED); out = []\n    for c in range(NUM_CLASSES):\n        ci = indices[labels[indices] == c]\n        out.append(np.concatenate([ci, rng.choice(ci, counts.max() - len(ci), replace=True)]))\n    return rng.permutation(np.concatenate(out))\n\nY_ONEHOT = keras.utils.to_categorical(y_all, NUM_CLASSES).astype(\"float32\")\n\ndef make_dataset(X, idx, training, balance=None, flip=None):\n    # tf.data pipeline over image INDICES: (image, one-hot label[, sample weight]) batches\n    # balance: sqrt / balanced / oversample (training only; anything else = no balancing); flip: None / \"lr\" / \"ud\" / \"both\" (TTA)\n    balance = balance or BALANCE_METHOD\n    idx = np.asarray(idx)\n    if training and balance == \"oversample\":\n        idx = oversample(idx, y_all)\n    use_w = training and balance in WEIGHT_SETS\n    w_all = WEIGHT_SETS[balance][y_all].astype(\"float32\") if use_w else None\n\n    def gather(ib):\n        ib = ib.astype(np.int64)\n        out = [X[ib].astype(np.float32), Y_ONEHOT[ib]]\n        if use_w: out.append(w_all[ib])\n        return tuple(out)\n\n    ds = tf.data.Dataset.from_tensor_slices(idx)\n    if training:\n        ds = ds.shuffle(len(idx), seed=SEED, reshuffle_each_iteration=True)\n    S = X.shape[1]                                   # image size of this array (300 or 380)\n    ds = ds.batch(BATCH_SIZE if S <= IMG_SIZE else BATCH_SIZE_HIRES)\n    types = [tf.float32, tf.float32] + ([tf.float32] if use_w else [])\n    def load(ib):\n        out = tf.numpy_function(gather, [ib], types)\n        out[0].set_shape([None, S, S, 3]); out[1].set_shape([None, NUM_CLASSES])\n        if use_w: out[2].set_shape([None])\n        return tuple(out)\n    ds = ds.map(load, num_parallel_calls=tf.data.AUTOTUNE)\n    if training:   # augmentation runs on whole batches (each image gets its own random transform)\n        ds = ds.map(lambda img, *rest: (tf.clip_by_value(augmenter(img, training=True), 0, 255), *rest),\n                    num_parallel_calls=tf.data.AUTOTUNE)\n    if flip in (\"lr\", \"both\"): ds = ds.map(lambda img, *rest: (tf.reverse(img, [2]), *rest))\n    if flip in (\"ud\", \"both\"): ds = ds.map(lambda img, *rest: (tf.reverse(img, [1]), *rest))\n    return ds.prefetch(tf.data.AUTOTUNE)\n\n# Class distribution the model effectively \"sees\" under each balancing method\nraw_counts = np.bincount(y_all[train_idx], minlength=NUM_CLASSES)\neff = pd.DataFrame({\"no balancing\": raw_counts,\n                    \"sqrt weights\": raw_counts * WEIGHT_SETS[\"sqrt\"],\n                    \"balanced weights\": raw_counts * WEIGHT_SETS[\"balanced\"],\n                    \"oversample\": np.bincount(y_all[oversample(train_idx, y_all)], minlength=NUM_CLASSES)},\n                   index=CLASS_NAMES).round()\ndisplay(eff)\neff.plot(kind=\"bar\", figsize=(10, 4))\nplt.title(\"Effective training samples per class under each balancing method\"); plt.ylabel(\"Effective samples\")\nplt.xticks(rotation=15)\nsave_fig(\"07_class_balancing.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"bd13fa81","cell_type":"code","source":"# ---------- Visual evidence: 8 augmentations of one preprocessed image ----------\ndemo = enhance(X_base[train_idx[0]], \"clahe_unsharp\").astype(\"float32\")\nfig, axes = plt.subplots(2, 5, figsize=(15, 6))  # last panel left blank\naxes[0, 0].imshow(demo.astype(\"uint8\")); axes[0, 0].set_title(\"Preprocessed original\")\nfor ax in axes.ravel()[1:9]:\n    aug = tf.clip_by_value(augmenter(demo[None], training=True)[0], 0, 255).numpy().astype(\"uint8\")\n    ax.imshow(aug); ax.set_title(\"Augmented\")\nfor ax in axes.ravel(): ax.axis(\"off\")\nplt.suptitle(\"Data augmentation examples\", fontsize=14); plt.tight_layout()\nsave_fig(\"08_augmentation_examples.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"b8030bcc","cell_type":"markdown","source":"## 6. Model builder, training loop and validation-QWK callback\nEvery model = **ImageNet-pretrained backbone** (no top) → Global Average Pooling → Dropout → Dense(256, ReLU, L2) → Dropout → Dense(5, softmax).\n\n- EfficientNetB0/B3 include their own input scaling, so they take raw 0–255 pixels.\n- ResNet50V2 and MobileNetV2 expect pixels in [−1, 1], so a `Rescaling` layer is added.\n- The backbone is called with `training=False`, so its BatchNorm statistics stay fixed even after unfreezing (the recommended Keras fine-tuning practice for small datasets).\n- Loss: categorical cross-entropy with label smoothing 0.05 (reduces over-confidence; neighbouring stages are visually similar).\n- **Validation QWK callback:** after every epoch the model is scored on the validation set with quadratic weighted kappa; early stopping and checkpointing keep the epoch with the best QWK (not just the lowest loss).\n- `train_full()` runs the complete two-phase schedule, so Experiments 3–4 and the final model use exactly the same procedure.","metadata":{}},{"id":"d0f8f48d","cell_type":"code","source":"# name -> (Keras constructor, needs [-1, 1] input rescaling?)\nBACKBONES = {\n    \"EfficientNetB0\": (keras.applications.EfficientNetB0, False),\n    \"EfficientNetB3\": (keras.applications.EfficientNetB3, False),\n    \"EfficientNetB4\": (keras.applications.EfficientNetB4, False),\n    \"ResNet50V2\":     (keras.applications.ResNet50V2,     True),\n    \"MobileNetV2\":    (keras.applications.MobileNetV2,    True),\n}\n\ndef build_model(backbone=\"EfficientNetB0\", dropout=0.3, dense_units=256, lr=1e-3, img_size=None):\n    img_size = img_size or IMG_SIZE\n    ctor, rescale = BACKBONES[backbone]\n    inputs = keras.Input((img_size, img_size, 3), name=\"image\")\n    x = layers.Rescaling(1 / 127.5, offset=-1, name=\"rescale\")(inputs) if rescale else inputs\n    base = ctor(include_top=False, weights=WEIGHTS, input_shape=(img_size, img_size, 3))\n    base.trainable = False                                   # phase 1: frozen feature extractor\n    x = base(x, training=False)\n    x = layers.GlobalAveragePooling2D(name=\"gap\")(x)\n    x = layers.Dropout(dropout, name=\"drop1\")(x)\n    x = layers.Dense(dense_units, activation=\"relu\", kernel_regularizer=regularizers.l2(1e-4), name=\"fc\")(x)\n    x = layers.Dropout(dropout, name=\"drop2\")(x)\n    outputs = layers.Dense(NUM_CLASSES, activation=\"softmax\", name=\"stage\")(x)\n    model = keras.Model(inputs, outputs, name=f\"DR_{backbone}\")\n    compile_model(model, lr)\n    return model, base\n\ndef compile_model(model, lr):\n    model.compile(optimizer=keras.optimizers.Adam(lr),\n                  loss=keras.losses.CategoricalCrossentropy(label_smoothing=0.05),\n                  metrics=[\"accuracy\"])\n\ndef metrics(true, pred):\n    return {\"accuracy\": accuracy_score(true, pred),\n            \"macro_f1\": f1_score(true, pred, average=\"macro\"),\n            \"qwk\": cohen_kappa_score(true, pred, weights=\"quadratic\")}\n\ndef predict_probs(model, X, idx, tta=False):\n    # Class probabilities; with TTA, average over original + 3 flipped views\n    views = [None, \"lr\", \"ud\", \"both\"] if tta else [None]\n    return np.mean([model.predict(make_dataset(X, idx, False, flip=f), verbose=0) for f in views], axis=0)\n\ndef evaluate_split(model, X, idx, tta=False):\n    # Returns probabilities and headline metrics on a split\n    probs = predict_probs(model, X, idx, tta)\n    return probs, metrics(y_all[idx], probs.argmax(1))\n\nclass ValQWK(keras.callbacks.Callback):\n    # Adds val_qwk / val_macro_f1 to the logs after each epoch\n    def __init__(self, X):\n        super().__init__(); self.X = X\n    def on_epoch_end(self, epoch, logs=None):\n        probs = self.model.predict(make_dataset(self.X, val_idx, False), verbose=0)\n        m = metrics(y_all[val_idx], probs.argmax(1))\n        logs[\"val_qwk\"], logs[\"val_macro_f1\"] = m[\"qwk\"], m[\"macro_f1\"]\n        print(f\" - val_qwk: {m['qwk']:.4f} - val_macro_f1: {m['macro_f1']:.4f}\")\n\ndef unfreeze(base, fraction):\n    # Unfreeze the top `fraction` of backbone layers; BatchNorm layers stay frozen\n    base.trainable = True\n    n_freeze = int(len(base.layers) * (1 - fraction))\n    for i, layer in enumerate(base.layers):\n        layer.trainable = i >= n_freeze and not isinstance(layer, layers.BatchNormalization)\n    return sum(l.trainable for l in base.layers)\n\ndef train_full(X, backbone, lr, dropout, balance, head_epochs, fine_epochs, tag=None, verbose=0, seed=SEED):\n    # Complete two-phase transfer-learning run. Returns (model, base, history_phase1, history_phase2)\n    set_seed(seed); keras.backend.clear_session()\n    model, base = build_model(backbone, dropout=dropout, lr=lr, img_size=X.shape[1])\n    train_ds = make_dataset(X, train_idx, True, balance=balance)\n    val_ds = make_dataset(X, val_idx, False)\n    def cbs(phase, patience):\n        c = [ValQWK(X),\n             keras.callbacks.EarlyStopping(monitor=\"val_qwk\", mode=\"max\", patience=patience,\n                                           restore_best_weights=True, verbose=verbose),\n             keras.callbacks.ReduceLROnPlateau(monitor=\"val_loss\", factor=0.3, patience=2, min_lr=1e-7, verbose=verbose)]\n        if tag:   # final run: keep a checkpoint and a CSV log for the report\n            c += [keras.callbacks.ModelCheckpoint(os.path.join(OUT_DIR, f\"{tag}_best.keras\"), monitor=\"val_qwk\",\n                                                  mode=\"max\", save_best_only=True),\n                  keras.callbacks.CSVLogger(os.path.join(OUT_DIR, f\"training_log_{phase}.csv\"))]\n        return c\n    h1 = model.fit(train_ds, validation_data=val_ds, epochs=head_epochs, callbacks=cbs(\"phase1\", 3), verbose=verbose)\n    n = unfreeze(base, FINE_TUNE_FRACTION)\n    if verbose: print(f\"Phase 2: {n} / {len(base.layers)} backbone layers trainable\")\n    compile_model(model, FINE_TUNE_LR)          # re-compile is required after changing trainable flags\n    h2 = model.fit(train_ds, validation_data=val_ds, epochs=fine_epochs, callbacks=cbs(\"phase2\", 6), verbose=verbose)\n    return model, base, h1, h2\n\ndef quick_train(X, backbone=\"EfficientNetB0\", lr=1e-3, dropout=0.3, epochs=EPOCHS_SEARCH):\n    # Short frozen-backbone run used for Experiments 1-2 (same seed -> fair comparison)\n    set_seed(); keras.backend.clear_session()\n    model, _ = build_model(backbone, dropout=dropout, lr=lr)\n    t0 = time.time()\n    model.fit(make_dataset(X, train_idx, True), validation_data=make_dataset(X, val_idx, False),\n              epochs=epochs, verbose=0)\n    _, m = evaluate_split(model, X, val_idx)\n    m.update(train_time_s=round(time.time() - t0), params_M=round(model.count_params() / 1e6, 2))\n    return m","metadata":{},"outputs":[],"execution_count":null},{"id":"0929427b","cell_type":"markdown","source":"### Re-using a previously trained model (optional)\nIf `REUSE_TRAINED_MODEL = True` and the output of an earlier run is attached as an input, the settings that run selected are loaded from its `results.json`, the four experiments are skipped (their tables and figures stay in that earlier version), and the saved model is loaded instead of retrained. The data split is identical because every seed is fixed.","metadata":{}},{"id":"bc973de6","cell_type":"code","source":"PREV_DIR = None\nif REUSE_TRAINED_MODEL:\n    root = os.environ.get(\"DR_DATA_ROOT\", \"/kaggle/input\")\n    hits = [h for h in glob.glob(os.path.join(root, \"**\", \"dr_model.keras\"), recursive=True)\n            if os.path.exists(os.path.join(os.path.dirname(h), \"results.json\"))]\n    if hits:\n        PREV_DIR = os.path.dirname(hits[0])\n        sel = json.load(open(os.path.join(PREV_DIR, \"results.json\")))[\"selected\"]\n        BEST_BACKBONE, PREPROCESS_MODE, BALANCE_METHOD = sel[\"backbone\"], sel[\"preprocess\"], sel[\"balancing\"]\n        BEST_LR, BEST_DROPOUT = sel[\"lr\"], sel[\"dropout\"]\n        RUN_BACKBONE_COMP = RUN_HP_SEARCH = RUN_PREPROC_EXP = RUN_BALANCE_EXP = False\nprint(\"Re-using trained model from:\", PREV_DIR or \"none (full training will run)\")\nprint(f\"Settings: {BEST_BACKBONE} | preprocessing={PREPROCESS_MODE} | balancing={BALANCE_METHOD} \"\n      f\"| lr={BEST_LR} | dropout={BEST_DROPOUT}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"c7b29778","cell_type":"markdown","source":"## 7. Experiment 1 — CNN backbone comparison\nFour ImageNet architectures (EfficientNetB0, EfficientNetB3, ResNet50V2, MobileNetV2), trained identically with frozen weights on unenhanced images (a neutral baseline). Selection criterion: validation **quadratic weighted kappa (QWK)**, the official APTOS metric, which penalises predictions further from the true stage more heavily. The test set is never used for any decision.","metadata":{}},{"id":"e5504c0d","cell_type":"code","source":"X_raw = build_set(\"raw\")\nif RUN_BACKBONE_COMP:\n    comp = {}\n    for name in BACKBONES:\n        comp[name] = quick_train(X_raw, backbone=name); gc.collect()\n        print(name, comp[name], \"| RAM GB:\", ram_gb())\n    comp_df = pd.DataFrame(comp).T.round(4)\n    display(comp_df)\n    BEST_BACKBONE = comp_df[SELECT_METRIC].idxmax()\n    comp_df[[\"accuracy\", \"macro_f1\", \"qwk\"]].plot(kind=\"bar\", figsize=(9, 4), rot=0)\n    plt.title(\"Experiment 1: backbone comparison (validation set, frozen backbone)\"); plt.ylim(0, 1)\n    save_fig(\"09_backbone_comparison.png\"); plt.show()\nprint(\"Selected backbone:\", BEST_BACKBONE)","metadata":{},"outputs":[],"execution_count":null},{"id":"161fc95c","cell_type":"markdown","source":"## 8. Experiment 2 — Hyperparameter tuning\nGrid search over the head learning rate and dropout rate for the selected backbone.","metadata":{}},{"id":"12d13a24","cell_type":"code","source":"if RUN_HP_SEARCH:\n    grid = [(lr, d) for lr in [1e-3, 3e-4] for d in [0.2, 0.4]]\n    hp_rows = []\n    for lr, d in grid:\n        m = quick_train(X_raw, backbone=BEST_BACKBONE, lr=lr, dropout=d)\n        hp_rows.append({\"learning_rate\": lr, \"dropout\": d, **m}); gc.collect()\n        print(hp_rows[-1], \"| RAM GB:\", ram_gb())\n    hp_df = pd.DataFrame(hp_rows).round(4)\n    display(hp_df)\n    best = hp_df.loc[hp_df[SELECT_METRIC].idxmax()]\n    BEST_LR, BEST_DROPOUT = float(best[\"learning_rate\"]), float(best[\"dropout\"])\n    pivot = hp_df.pivot(index=\"dropout\", columns=\"learning_rate\", values=SELECT_METRIC)\n    plt.figure(figsize=(5, 4)); sns.heatmap(pivot, annot=True, fmt=\".3f\", cmap=\"Blues\")\n    plt.title(f\"Experiment 2: validation {SELECT_METRIC.upper()} (learning rate x dropout)\")\n    save_fig(\"10_hyperparameter_grid.png\"); plt.show()\nprint(f\"Selected: lr={BEST_LR}, dropout={BEST_DROPOUT}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"2c2da75d","cell_type":"markdown","source":"## 9. Experiment 3 — Preprocessing comparison (full two-phase training)\nVersion 1 compared preprocessing with only 5 frozen epochs, which is too short to be reliable. Here every variant gets the **complete** two-phase schedule (frozen head, then full fine-tuning, early stopping on validation QWK), so the comparison reflects the final model.","metadata":{}},{"id":"d6d94e38","cell_type":"code","source":"if RUN_PREPROC_EXP:\n    pre_rows = {}\n    for mode in [\"clahe_unsharp\", \"ben_graham\", \"raw\"]:          # on a tie, the first (enhanced) mode wins\n        X_mode = X_raw if mode == \"raw\" else build_set(mode)\n        t0 = time.time()\n        m_, _, _, _ = train_full(X_mode, BEST_BACKBONE, BEST_LR, BEST_DROPOUT, BALANCE_METHOD,\n                                 EPOCHS_COMPACT_HEAD, EPOCHS_COMPACT_FINE)\n        _, pre_rows[mode] = evaluate_split(m_, X_mode, val_idx)\n        pre_rows[mode][\"train_time_min\"] = round((time.time() - t0) / 60, 1)\n        del m_, X_mode; gc.collect()\n        print(mode, pre_rows[mode], \"| RAM GB:\", ram_gb())\n    pre_df = pd.DataFrame(pre_rows).T.round(4)\n    display(pre_df)\n    PREPROCESS_MODE = pre_df[SELECT_METRIC].idxmax()\n    pre_df[[\"accuracy\", \"macro_f1\", \"qwk\"]].plot(kind=\"bar\", figsize=(8, 4), rot=0)\n    plt.title(\"Experiment 3: preprocessing with full training (validation set)\"); plt.ylim(0, 1)\n    save_fig(\"11_preprocessing_comparison.png\"); plt.show()\nprint(\"Selected preprocessing:\", PREPROCESS_MODE)\nX_all = X_raw if PREPROCESS_MODE == \"raw\" else build_set(PREPROCESS_MODE)   # final preprocessed dataset","metadata":{},"outputs":[],"execution_count":null},{"id":"2b88fc42","cell_type":"markdown","source":"## 10. Experiment 4 — Class-balancing comparison (full two-phase training)\nSame backbone, hyperparameters and preprocessing; only the imbalance strategy changes.","metadata":{}},{"id":"a18a1924","cell_type":"code","source":"if RUN_BALANCE_EXP:\n    bal_rows = {}\n    for bal in [\"balanced\", \"sqrt\", \"oversample\"]:\n        t0 = time.time()\n        m_, _, _, _ = train_full(X_all, BEST_BACKBONE, BEST_LR, BEST_DROPOUT, bal,\n                                 EPOCHS_COMPACT_HEAD, EPOCHS_COMPACT_FINE)\n        _, bal_rows[bal] = evaluate_split(m_, X_all, val_idx)\n        bal_rows[bal][\"train_time_min\"] = round((time.time() - t0) / 60, 1)\n        del m_; gc.collect()\n        print(bal, bal_rows[bal], \"| RAM GB:\", ram_gb())\n    bal_df = pd.DataFrame(bal_rows).T.round(4)\n    display(bal_df)\n    BALANCE_METHOD = bal_df[SELECT_METRIC].idxmax()\n    bal_df[[\"accuracy\", \"macro_f1\", \"qwk\"]].plot(kind=\"bar\", figsize=(8, 4), rot=0)\n    plt.title(\"Experiment 4: class-balancing methods (validation set)\"); plt.ylim(0, 1)\n    save_fig(\"12_balancing_comparison.png\"); plt.show()\nprint(\"Selected balancing:\", BALANCE_METHOD)","metadata":{},"outputs":[],"execution_count":null},{"id":"c1624b1f","cell_type":"markdown","source":"## 11. Final training — two-phase transfer learning\n**Phase 1 (feature extraction):** backbone frozen, only the new head learns (higher learning rate).\n**Phase 2 (fine-tuning):** the whole backbone is unfrozen (BatchNorm layers stay frozen) and trained with a very small learning rate, so the ImageNet features adapt to retinal lesions without being destroyed.\n\n**Callbacks / overfitting control**\n- `ValQWK` — scores the validation set after each epoch\n- `EarlyStopping` on validation QWK (patience 3 / 6, restores best weights)\n- `ReduceLROnPlateau` (×0.3 when val loss stalls for 2 epochs) — learning-rate scheduling\n- `ModelCheckpoint` (keeps the best-QWK model on disk) and `CSVLogger` (training log for the report)\n- plus Dropout, L2 weight decay, label smoothing and augmentation.","metadata":{}},{"id":"e24904b0","cell_type":"code","source":"print(f\"Final model: {BEST_BACKBONE} | preprocessing={PREPROCESS_MODE} | balancing={BALANCE_METHOD} \"\n      f\"| lr={BEST_LR} | dropout={BEST_DROPOUT} | img={IMG_SIZE}px\")\nclass LoggedHistory:\n    # Minimal stand-in for a Keras History object, rebuilt from a CSVLogger file\n    def __init__(self, csv_path):\n        self.history = pd.read_csv(csv_path).to_dict(orient=\"list\")\n\nif PREV_DIR:\n    model = keras.models.load_model(os.path.join(PREV_DIR, \"dr_model.keras\"))\n    base = next(l for l in model.layers if isinstance(l, keras.Model))\n    hist1 = LoggedHistory(os.path.join(PREV_DIR, \"training_log_phase1.csv\"))\n    hist2 = LoggedHistory(os.path.join(PREV_DIR, \"training_log_phase2.csv\"))\n    for f in [\"training_log_phase1.csv\", \"training_log_phase2.csv\"]:          # keep the logs with this run's outputs\n        shutil.copy(os.path.join(PREV_DIR, f), os.path.join(OUT_DIR, f))\n    print(\"Loaded trained model and its training logs from\", PREV_DIR)\nelse:\n    t0 = time.time()\n    model, base, hist1, hist2 = train_full(X_all, BEST_BACKBONE, BEST_LR, BEST_DROPOUT, BALANCE_METHOD,\n                                           EPOCHS_HEAD, EPOCHS_FINE, tag=\"dr_model\", verbose=2)\n    print(f\"Total training time: {(time.time() - t0) / 60:.1f} min | RAM GB: {ram_gb()}\")\nmodel.summary(show_trainable=True)","metadata":{},"outputs":[],"execution_count":null},{"id":"22a4303b","cell_type":"code","source":"# ---------- Accuracy and loss curves (both phases) ----------\nH = {k: hist1.history[k] + hist2.history[k] for k in [\"accuracy\", \"val_accuracy\", \"loss\", \"val_loss\", \"val_qwk\"]}\nsplit = len(hist1.history[\"loss\"])\nepochs_axis = np.arange(1, len(H[\"loss\"]) + 1)\n\nfig, axes = plt.subplots(1, 3, figsize=(20, 5))\nfor ax, metric in zip(axes[:2], [\"accuracy\", \"loss\"]):\n    ax.plot(epochs_axis, H[metric], \"o-\", label=f\"Train {metric}\")\n    ax.plot(epochs_axis, H[f\"val_{metric}\"], \"s-\", label=f\"Validation {metric}\")\naxes[2].plot(epochs_axis, H[\"val_qwk\"], \"s-\", color=\"tab:green\", label=\"Validation QWK\")\nfor ax, t in zip(axes, [\"accuracy\", \"loss\", \"QWK\"]):\n    ax.axvline(split + 0.5, color=\"grey\", ls=\"--\", label=\"Start fine-tuning\")\n    ax.set_xlabel(\"Epoch\"); ax.set_ylabel(t); ax.set_title(f\"Training vs validation {t}\")\n    ax.legend(); ax.grid(alpha=0.3)\nsave_fig(\"13_accuracy_loss_qwk_curves.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"5b52fb72","cell_type":"markdown","source":"## 12. Inference optimisation — TTA and stage thresholds\nTwo cheap improvements that need no retraining; both are decided on the **validation** set only:\n1. **Test-time augmentation (TTA):** each image is predicted 4 times (original, horizontal flip, vertical flip, both) and the probabilities are averaged, which smooths out orientation-specific errors.\n2. **Optimised stage thresholds:** DR stages are ordered, so each image gets an expected grade `E = Σ k·p(k)` (0–4). Four cut-points turn `E` into a stage; they are tuned by coordinate search to maximise validation QWK.\n\n**Selection rule:** the option with the best *average* of accuracy, macro F1 and QWK on the validation set is used. QWK alone is not enough: in the first run of this notebook, thresholds tuned for QWK raised QWK but pushed borderline Mild cases into Moderate, cutting Mild recall and overall accuracy. A screening tool must also find the early (Mild) stage, so all three metrics count.","metadata":{}},{"id":"cf82e343","cell_type":"code","source":"def expected_grade(probs):\n    return probs @ np.arange(NUM_CLASSES)\n\ndef apply_thresholds(score, th):\n    return np.digitize(score, th)\n\ndef optimise_thresholds(score, true, metric=SELECT_METRIC, passes=3):\n    # Coordinate search: move one cut-point at a time to the value that maximises the metric\n    th = np.array([0.5, 1.5, 2.5, 3.5])\n    for _ in range(passes):\n        for i in range(len(th)):\n            lo = th[i - 1] + 0.01 if i > 0 else 0.0\n            hi = th[i + 1] - 0.01 if i < len(th) - 1 else NUM_CLASSES - 1.0\n            cands = np.linspace(lo, hi, 60)\n            vals = [metrics(true, apply_thresholds(score, np.r_[th[:i], c, th[i + 1:]]))[metric] for c in cands]\n            th[i] = cands[int(np.argmax(vals))]\n    return th\n\nval_true = y_all[val_idx]\nval_plain = predict_probs(model, X_all, val_idx, tta=False)\nval_tta   = predict_probs(model, X_all, val_idx, tta=True)\nTHRESHOLDS = optimise_thresholds(expected_grade(val_tta), val_true)\n\noptions = {\n    \"argmax\":              (False, \"argmax\", metrics(val_true, val_plain.argmax(1))),\n    \"TTA + argmax\":        (True,  \"argmax\", metrics(val_true, val_tta.argmax(1))),\n    \"TTA + thresholds\":    (True,  \"thresholds\", metrics(val_true, apply_thresholds(expected_grade(val_tta), THRESHOLDS))),\n}\nopt_df = pd.DataFrame({k: v[2] for k, v in options.items()}).T.round(4)\nopt_df[\"mean_score\"] = opt_df[[\"accuracy\", \"macro_f1\", \"qwk\"]].mean(axis=1).round(4)\nprint(\"Validation set:\"); display(opt_df)\nBEST_INFERENCE = opt_df[\"mean_score\"].idxmax()\nUSE_TTA, DECISION = options[BEST_INFERENCE][0], options[BEST_INFERENCE][1]\nprint(f\"Selected inference: {BEST_INFERENCE} | thresholds = {np.round(THRESHOLDS, 3).tolist()}\")\n\ndef decide(probs):\n    return apply_thresholds(expected_grade(probs), THRESHOLDS) if DECISION == \"thresholds\" else probs.argmax(1)","metadata":{},"outputs":[],"execution_count":null},{"id":"4efba0ca","cell_type":"markdown","source":"## 12b. Version 4 — High-resolution ensemble\nAn ensemble averages the class probabilities of several independently trained models. Different models make different mistakes, so the average is usually more accurate and more stable than any single model.\n\n**Members**\n1. **EfficientNetB0, 300 px** (Version 2 model, loaded).\n2. **EfficientNetB3, 300 px** (Version 3 member, loaded).\n3. **EfficientNetB3, 380 px** (new): higher resolution keeps small lesions such as microaneurysms and tiny haemorrhages visible, which matters most for the Mild and Severe stages.\n4. **EfficientNetB4, 380 px** (new): a deeper backbone whose native input size is 380 px.\n\nThe new members are trained with exactly the same two-phase schedule, preprocessing and balancing as the Version 2 model; only the backbone and the image size differ. The 380 px images are produced with the same crop, pad and resize steps.\n\n**Selection (validation set only):** every non-empty combination of members, each with and without TTA, is scored with the same rule as section 12 (mean of accuracy, macro F1 and QWK). The Version 3 ensemble and the single model are among the candidates, so a new member is only used if it genuinely helps on validation. The test set is not used for this choice.","metadata":{}},{"id":"0f6bc701","cell_type":"code","source":"import itertools\nUSE_ENSEMBLE, ENSEMBLE_CHOICE = False, None\nens_models, ens_X = {\"B0 (Version 2)\": model}, {\"B0 (Version 2)\": X_all}\nif RUN_ENSEMBLE:\n    # 1. Load members saved by earlier runs (Version 3 ensemble) from the attached input\n    root = os.environ.get(\"DR_DATA_ROOT\", \"/kaggle/input\")\n    for cfg_path in glob.glob(os.path.join(root, \"**\", \"dr_config.json\"), recursive=True):\n        ens_cfg = json.load(open(cfg_path)).get(\"ensemble\") or {}\n        for mem in ens_cfg.get(\"members\", []):\n            if mem[\"file\"] == \"dr_model.keras\" or mem[\"name\"] in ens_models:\n                continue\n            path = os.path.join(os.path.dirname(cfg_path), mem[\"file\"])\n            if os.path.exists(path):\n                ens_models[mem[\"name\"]] = keras.models.load_model(path)\n                ens_X[mem[\"name\"]] = X_all\n                print(\"Loaded earlier member:\", mem[\"name\"], \"from\", path)\n\n    # 2. High-resolution copy of the dataset (same crop/pad/resize, then the selected preprocessing)\n    t0 = time.time()\n    X_hr = np.stack(Parallel(n_jobs=-1)(delayed(load_and_resize)(p, HIRES_SIZE) for p in df[\"path\"]))\n    if PREPROCESS_MODE != \"raw\":\n        X_hr = np.stack([enhance(im, PREPROCESS_MODE) for im in X_hr])\n    print(f\"High-resolution set {X_hr.shape} built in {time.time() - t0:.0f}s | RAM GB: {ram_gb()}\")\n\n    # 3. Train the new members\n    for bb, sd, size in ENSEMBLE_MEMBERS:\n        name = f\"{bb.replace('EfficientNet', '')} {size}px (seed {sd})\"\n        Xm = X_hr if size == HIRES_SIZE else X_all\n        t0 = time.time()\n        m_, _, _, _ = train_full(Xm, bb, BEST_LR, BEST_DROPOUT, BALANCE_METHOD, EPOCHS_HEAD, EPOCHS_FINE,\n                                 verbose=2, seed=sd)\n        ens_models[name], ens_X[name] = m_, Xm\n        gc.collect()\n        print(f\"Trained {name} in {(time.time() - t0) / 60:.1f} min | RAM GB: {ram_gb()}\")\n\n    # 4. Probabilities of every member on validation and test, with and without TTA\n    P = {}\n    for name, m_ in ens_models.items():\n        for tta in (False, True):\n            P[(name, tta, \"val\")]  = predict_probs(m_, ens_X[name], val_idx, tta=tta)\n            P[(name, tta, \"test\")] = predict_probs(m_, ens_X[name], test_idx, tta=tta)\n\n    member_df = pd.DataFrame({name: metrics(val_true, P[(name, False, \"val\")].argmax(1)) for name in ens_models}).T.round(4)\n    print(\"Each member on the validation set (no TTA):\"); display(member_df)\n\n    rows, names = [], list(ens_models)\n    for k in range(1, len(names) + 1):\n        for combo in itertools.combinations(names, k):\n            for tta in (False, True):\n                pv = np.mean([P[(n, tta, \"val\")] for n in combo], axis=0)\n                m = metrics(val_true, pv.argmax(1))\n                rows.append({\"members\": \" + \".join(combo), \"tta\": tta, **m,\n                             \"mean_score\": np.mean([m[\"accuracy\"], m[\"macro_f1\"], m[\"qwk\"]])})\n    ens_df = pd.DataFrame(rows).sort_values(\"mean_score\", ascending=False).round(4).reset_index(drop=True)\n    print(\"All combinations on the validation set (best first):\"); display(ens_df.head(10))\n\n    best = ens_df.iloc[0]\n    single_best = float(opt_df[\"mean_score\"].max())\n    if best[\"mean_score\"] > single_best:\n        USE_ENSEMBLE = True\n        ENSEMBLE_CHOICE = {\"members\": best[\"members\"].split(\" + \"), \"tta\": bool(best[\"tta\"])}\n    print(f\"Best ensemble option (validation mean score {best['mean_score']:.4f}) vs single model option \"\n          f\"'{BEST_INFERENCE}' ({single_best:.4f}) -> using {'ENSEMBLE' if USE_ENSEMBLE else 'single model'}\")\n\n    top = ens_df.head(8).copy()\n    top[\"label\"] = top[\"members\"] + np.where(top[\"tta\"], \" + TTA\", \"\")\n    top.set_index(\"label\")[[\"accuracy\", \"macro_f1\", \"qwk\"]].plot(kind=\"barh\", figsize=(10, 6)).invert_yaxis()\n    plt.xlim(0.5, 1); plt.title(\"Ensemble options on the validation set (top 8)\")\n    save_fig(\"14b_ensemble_selection.png\"); plt.show()\n\ndef ensemble_probs(split):\n    return np.mean([P[(n, ENSEMBLE_CHOICE[\"tta\"], split)] for n in ENSEMBLE_CHOICE[\"members\"]], axis=0)","metadata":{},"outputs":[],"execution_count":null},{"id":"ab097fea","cell_type":"markdown","source":"## 13. Evaluation on the held-out test set\nThe test set was never used for training or for any selection above, so these numbers estimate real-world performance. The table also shows how much each inference step adds on the test set (for transparency only; the choice was made on validation).","metadata":{}},{"id":"d7ddfeb4","cell_type":"code","source":"y_true = y_all[test_idx]\ntest_plain = predict_probs(model, X_all, test_idx, tta=False)\ntest_tta   = predict_probs(model, X_all, test_idx, tta=True)\nstep_df = pd.DataFrame({\n    \"argmax\":           metrics(y_true, test_plain.argmax(1)),\n    \"TTA + argmax\":     metrics(y_true, test_tta.argmax(1)),\n    \"TTA + thresholds\": metrics(y_true, apply_thresholds(expected_grade(test_tta), THRESHOLDS)),\n}).T.round(4)\nif RUN_ENSEMBLE:\n    ens_label = \"Ensemble: \" + \" + \".join(ENSEMBLE_CHOICE[\"members\"] if USE_ENSEMBLE else [ens_df.iloc[0][\"members\"]])\n    ens_test = ensemble_probs(\"test\") if USE_ENSEMBLE else np.mean(\n        [P[(n, bool(ens_df.iloc[0][\"tta\"]), \"test\")] for n in ens_df.iloc[0][\"members\"].split(\" + \")], axis=0)\n    step_df.loc[\"Ensemble (best on validation)\"] = pd.Series(metrics(y_true, ens_test.argmax(1))).round(4)\nprint(\"Test set, by inference method:\"); display(step_df)\nstep_df[[\"accuracy\", \"macro_f1\", \"qwk\"]].plot(kind=\"bar\", figsize=(8, 4), rot=0); plt.ylim(0, 1)\nplt.title(\"Effect of TTA and threshold optimisation (test set)\")\nsave_fig(\"14_inference_optimisation.png\"); plt.show()\n\nif USE_ENSEMBLE:\n    test_probs, DECISION = ensemble_probs(\"test\"), \"argmax\"\nelse:\n    test_probs = test_tta if USE_TTA else test_plain\ny_pred = decide(test_probs)\nprint(\"Final prediction method:\", (\"ensemble of \" + \", \".join(ENSEMBLE_CHOICE[\"members\"]) +\n      (\" + TTA\" if ENSEMBLE_CHOICE[\"tta\"] else \"\")) if USE_ENSEMBLE else BEST_INFERENCE)\ntest_summary = metrics(y_true, y_pred)\n\nprint(\"=== Multi-class (stage) results ===\")\nfor k, v in test_summary.items(): print(f\"{k:>10}: {v:.4f}\")\nreport = classification_report(y_true, y_pred, target_names=CLASS_NAMES, digits=4, output_dict=True, zero_division=0)\nreport_df = pd.DataFrame(report).T.round(4)\ndisplay(report_df)\nreport_df.to_csv(os.path.join(OUT_DIR, \"classification_report.csv\"))","metadata":{},"outputs":[],"execution_count":null},{"id":"b98655ae","cell_type":"code","source":"# ---------- Confusion matrices (counts and row-normalised) ----------\ncm = confusion_matrix(y_true, y_pred, labels=range(NUM_CLASSES))\ncm_norm = cm / cm.sum(1, keepdims=True)\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES, ax=axes[0])\nsns.heatmap(cm_norm, annot=True, fmt=\".2f\", cmap=\"Blues\", xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES, ax=axes[1])\nfor ax, t in zip(axes, [\"Confusion matrix (counts)\", \"Confusion matrix (recall per class)\"]):\n    ax.set_title(t); ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\")\nplt.tight_layout(); save_fig(\"15_confusion_matrix.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"ccab0995","cell_type":"code","source":"# ---------- Per-class precision / recall / F1 ----------\np, r, f, s = precision_recall_fscore_support(y_true, y_pred, labels=range(NUM_CLASSES), zero_division=0)\nprf = pd.DataFrame({\"Precision\": p, \"Recall\": r, \"F1-score\": f}, index=CLASS_NAMES)\nprf.plot(kind=\"bar\", figsize=(10, 4), rot=15); plt.ylim(0, 1.05)\nplt.title(\"Per-class precision, recall and F1 (test set)\"); plt.grid(axis=\"y\", alpha=0.3)\nsave_fig(\"16_per_class_prf.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"a9e4968a","cell_type":"code","source":"# ---------- Binary screening view: DR (stage 1-4) vs No DR (stage 0) ----------\nyb_true = (y_true > 0).astype(int)\nprob_dr = 1 - test_probs[:, 0]                     # P(any DR)\nyb_pred = (prob_dr >= 0.5).astype(int)\ntn, fp, fn, tp = confusion_matrix(yb_true, yb_pred, labels=[0, 1]).ravel()\nbinary = {\"accuracy\": (tp + tn) / (tp + tn + fp + fn),\n          \"sensitivity (recall)\": tp / (tp + fn), \"specificity\": tn / (tn + fp),\n          \"precision\": tp / max(tp + fp, 1), \"f1\": 2 * tp / max(2 * tp + fp + fn, 1),\n          \"roc_auc\": roc_auc_score(yb_true, prob_dr)}\nprint(\"=== DR vs No DR (screening) ===\")\nfor k, v in binary.items(): print(f\"{k:>22}: {v:.4f}\")\n\nfig, axes = plt.subplots(1, 2, figsize=(13, 5))\nfpr, tpr, _ = roc_curve(yb_true, prob_dr)\naxes[0].plot(fpr, tpr, label=f\"DR vs No DR (AUC={binary['roc_auc']:.3f})\", lw=2)\nfor c in range(NUM_CLASSES):                        # one-vs-rest ROC for each stage\n    if len(np.unique(y_true == c)) == 2:\n        fc, tc, _ = roc_curve(y_true == c, test_probs[:, c])\n        axes[0].plot(fc, tc, alpha=0.7, label=f\"{CLASS_NAMES[c]} (AUC={roc_auc_score(y_true == c, test_probs[:, c]):.3f})\")\naxes[0].plot([0, 1], [0, 1], \"k--\"); axes[0].set_xlabel(\"False positive rate\"); axes[0].set_ylabel(\"True positive rate\")\naxes[0].set_title(\"ROC curves (test set)\"); axes[0].legend(fontsize=8)\nsns.heatmap([[tn, fp], [fn, tp]], annot=True, fmt=\"d\", cmap=\"Greens\", ax=axes[1],\n            xticklabels=[\"No DR\", \"DR\"], yticklabels=[\"No DR\", \"DR\"])\naxes[1].set_title(\"Binary confusion matrix\"); axes[1].set_xlabel(\"Predicted\"); axes[1].set_ylabel(\"True\")\nplt.tight_layout(); save_fig(\"17_roc_binary.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"17269173","cell_type":"markdown","source":"## 14. Error analysis\nDistinguishes *near misses* (off by one stage — clinically less serious, and graders often disagree here too) from *severe errors* (≥ 2 stages apart).","metadata":{}},{"id":"75d5d786","cell_type":"code","source":"diff = np.abs(y_true - y_pred)\nerr = pd.Series({\"Correct\": (diff == 0).mean(), \"Off by 1 stage\": (diff == 1).mean(), \"Off by >=2 stages\": (diff >= 2).mean()})\ndisplay((err * 100).round(1).to_frame(\"% of test images\"))\n\npairs = pd.DataFrame([(CLASS_NAMES[t], CLASS_NAMES[p_], cm[t, p_]) for t in range(NUM_CLASSES)\n                      for p_ in range(NUM_CLASSES) if t != p_ and cm[t, p_] > 0],\n                     columns=[\"True\", \"Predicted\", \"Count\"]).sort_values(\"Count\", ascending=False)\nprint(\"Most frequent confusions:\"); display(pairs.head(6))\n\n# Show the most confident mistakes\nwrong = np.where(diff > 0)[0]\nwrong = wrong[np.argsort(-test_probs[wrong].max(1))][:8]\nif len(wrong):\n    fig, axes = plt.subplots(1, len(wrong), figsize=(3 * len(wrong), 3.6))\n    for ax, i in zip(np.atleast_1d(axes), wrong):\n        ax.imshow(X_all[test_idx[i]]); ax.axis(\"off\")\n        ax.set_title(f\"True: {CLASS_NAMES[y_true[i]]}\\nPred: {CLASS_NAMES[y_pred[i]]} ({test_probs[i].max():.2f})\", fontsize=9)\n    plt.suptitle(\"Most confident misclassifications\"); plt.tight_layout()\n    save_fig(\"18_misclassified.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"293935a0","cell_type":"markdown","source":"## 15. Explainability — Grad-CAM\nGrad-CAM highlights the image regions that most influenced the prediction (gradients of the predicted class w.r.t. the last convolutional feature map). Heat on lesions (haemorrhages, exudates, new vessels) rather than the image border shows the model learned clinically meaningful features.","metadata":{}},{"id":"2aed2651","cell_type":"code","source":"def grad_cam(model, img_uint8):\n    # Run the model layer by layer so we can watch the backbone's output feature map\n    base_layer = next(l for l in model.layers if isinstance(l, keras.Model))\n    pos = model.layers.index(base_layer)\n    x = tf.convert_to_tensor(img_uint8[None].astype(\"float32\"))\n    with tf.GradientTape() as tape:\n        h = x\n        for layer in model.layers[1:pos]:           # e.g. Rescaling\n            h = layer(h)\n        conv = base_layer(h, training=False)        # (1, h, w, C) last feature map\n        tape.watch(conv)\n        h = conv\n        for layer in model.layers[pos + 1:]:        # GAP -> Dense ... -> softmax\n            h = layer(h, training=False)\n        cls = int(tf.argmax(h[0]))\n        score = h[:, cls]\n    grads = tape.gradient(score, conv)\n    weights = tf.reduce_mean(grads, axis=(0, 1, 2))                  # importance of each channel\n    cam = tf.nn.relu(tf.reduce_sum(conv[0] * weights, axis=-1)).numpy()\n    cam = cv2.resize(cam / (cam.max() + 1e-8), (img_uint8.shape[1], img_uint8.shape[0]))\n    heat = cv2.cvtColor(cv2.applyColorMap(np.uint8(255 * cam), cv2.COLORMAP_JET), cv2.COLOR_BGR2RGB)\n    overlay = cv2.addWeighted(img_uint8, 0.6, heat, 0.4, 0)\n    return overlay, cls, h.numpy()[0]\n\nfig, axes = plt.subplots(2, NUM_CLASSES, figsize=(4 * NUM_CLASSES, 8))\nfor c in range(NUM_CLASSES):\n    cand = test_idx[y_all[test_idx] == c]\n    if len(cand) == 0: continue\n    img = X_all[cand[0]]\n    overlay, pc, pr = grad_cam(model, img)\n    axes[0, c].imshow(img); axes[0, c].set_title(f\"True: {CLASS_NAMES[c]}\")\n    axes[1, c].imshow(overlay); axes[1, c].set_title(f\"Most activated: {CLASS_NAMES[pc]} ({pr[pc]:.2f})\")\n    axes[0, c].axis(\"off\"); axes[1, c].axis(\"off\")\nplt.suptitle(\"Grad-CAM: where the model looks\", fontsize=14); plt.tight_layout()\nsave_fig(\"19_gradcam.png\"); plt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"2b7b3beb","cell_type":"markdown","source":"## 16. Save the model and results for the prototype app\nDownload these files from the **Output** panel (right side) → `/kaggle/working`:\n- `dr_model.keras` — trained model\n- `dr_config.json` — preprocessing mode, image size, class names, TTA / threshold settings, TF version (the app reads this)\n- `results.json`, `classification_report.csv`, `training_log_*.csv`, `figures/*.png`","metadata":{}},{"id":"3ad0ab3e","cell_type":"code","source":"MODEL_PATH = os.path.join(OUT_DIR, \"dr_model.keras\")\nmodel.save(MODEL_PATH)\nENSEMBLE_FILES = []\nif USE_ENSEMBLE:\n    for i, n in enumerate(ENSEMBLE_CHOICE[\"members\"]):\n        fname = \"dr_model.keras\" if n == \"B0 (Version 2)\" else f\"dr_model_member{i}.keras\"\n        if fname != \"dr_model.keras\":\n            ens_models[n].save(os.path.join(OUT_DIR, fname))\n        ENSEMBLE_FILES.append({\"name\": n, \"file\": fname, \"img_size\": int(ens_X[n].shape[1])})\n\nconfig = {\"img_size\": IMG_SIZE, \"preprocess_mode\": PREPROCESS_MODE, \"backbone\": BEST_BACKBONE,\n          \"class_names\": CLASS_NAMES, \"tta\": bool(USE_TTA), \"decision\": DECISION,\n          \"thresholds\": [float(t) for t in THRESHOLDS],\n          \"ensemble\": ({\"members\": ENSEMBLE_FILES, \"tta\": ENSEMBLE_CHOICE[\"tta\"]} if USE_ENSEMBLE else None),\n          \"tensorflow_version\": tf.__version__, \"keras_version\": keras.__version__}\njson.dump(config, open(os.path.join(OUT_DIR, \"dr_config.json\"), \"w\"), indent=2)\n\nresults = {\"test_stage_metrics\": test_summary, \"test_binary_metrics\": binary,\n           \"test_by_inference_method\": step_df.to_dict(orient=\"index\"),\n           \"validation_by_inference_method\": opt_df.to_dict(orient=\"index\"),\n           \"reused_model_from\": PREV_DIR,\n           \"ensemble_validation\": (ens_df.to_dict(orient=\"records\") if RUN_ENSEMBLE else None),\n           \"ensemble_used\": ENSEMBLE_CHOICE if USE_ENSEMBLE else None,\n           \"selected\": {\"preprocess\": PREPROCESS_MODE, \"backbone\": BEST_BACKBONE, \"lr\": BEST_LR, \"dropout\": BEST_DROPOUT,\n                        \"balancing\": BALANCE_METHOD, \"img_size\": IMG_SIZE, \"inference\": BEST_INFERENCE,\n                        \"thresholds\": [float(t) for t in THRESHOLDS]},\n           \"epochs_trained\": {\"phase1\": len(hist1.history[\"loss\"]), \"phase2\": len(hist2.history[\"loss\"])}}\njson.dump(results, open(os.path.join(OUT_DIR, \"results.json\"), \"w\"), indent=2, default=float)\nprint(json.dumps(results, indent=2, default=float))\nprint(\"\\nSaved:\", MODEL_PATH, \"| size:\", round(os.path.getsize(MODEL_PATH) / 1e6, 1), \"MB\")","metadata":{},"outputs":[],"execution_count":null},{"id":"ea4f2b2e","cell_type":"code","source":"# ---------- Sanity check: reload the saved model and predict one image from disk ----------\nreloaded = keras.models.load_model(MODEL_PATH)\np = df[\"path\"].iloc[test_idx[0]]\nimg = preprocess_fundus(p, PREPROCESS_MODE)\nbatch = np.stack([img, img[:, ::-1], img[::-1], img[::-1, ::-1]]) if USE_TTA else img[None]\nprobs = reloaded.predict(batch.astype(\"float32\"), verbose=0).mean(0)\nprint(\"Image:\", os.path.basename(p), \"| true:\", CLASS_NAMES[y_all[test_idx[0]]])\nfor name, pr in zip(CLASS_NAMES, probs): print(f\"  {name:<18} {pr:.3f}\")\nprint(\"Prediction (Version 2 model alone):\", CLASS_NAMES[int(probs.argmax())], \"| DR present:\", probs[1:].sum() >= 0.5)","metadata":{},"outputs":[],"execution_count":null},{"id":"d9756142","cell_type":"code","source":"# ---------- Export demo images (unseen test images) + one zip with everything ----------\nEX_DIR = os.path.join(OUT_DIR, \"examples\"); os.makedirs(EX_DIR, exist_ok=True)\nfor c in range(NUM_CLASSES):\n    for i in test_idx[y_all[test_idx] == c][:2]:             # 2 test images per stage\n        src = df[\"path\"].iloc[i]\n        shutil.copy(src, os.path.join(EX_DIR, f\"stage{c}_{os.path.basename(src)}\"))\n\nfor f in glob.glob(os.path.join(OUT_DIR, \"*_best.keras\")):   # duplicate of dr_model.keras, keep the zip small\n    os.remove(f)\ntmp_zip = shutil.make_archive(\"/tmp/dr_outputs\", \"zip\", OUT_DIR)\nshutil.move(tmp_zip, os.path.join(OUT_DIR, \"dr_outputs.zip\"))\nprint(\"Download dr_outputs.zip from the Output panel. Contents:\")\nprint(sorted(os.listdir(OUT_DIR)))","metadata":{},"outputs":[],"execution_count":null}]}