{"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"},"accelerator":"GPU","colab":{"gpuType":"T4","provenance":[]}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"86dd4dca","cell_type":"markdown","source":"# Diabetic Retinopathy Stage Detection — EfficientNetB3 Transfer Learning\n**BSc (Hons) Computer Science — Computer Vision Coursework (CW1)**\n\nThis notebook detects diabetic retinopathy (DR) and grades its stage from colour fundus photographs using ImageNet transfer learning (EfficientNetB3), with Grad-CAM explainability, embedding-based similar-case retrieval, a multi-agent decision pipeline and a Gradio interface.\n\n**Dataset:** [APTOS 2019 Blindness Detection](https://www.kaggle.com/competitions/aptos2019-blindness-detection) (Kaggle competition). Only the **3,662 labelled images** in `train_images/` are used (`train.csv`: `id_code`, `diagnosis`); the competition's `test_images/` have no public labels and are ignored. Class counts are 1,805 No DR · 370 Mild · 999 Moderate · 193 Severe · 295 Proliferative, a **9.4× imbalance** between the largest and smallest class (recomputed in Section 1.9). All images come from **one source** (Aravind Eye Hospital, India), so results may not transfer to other cameras or populations, and there are **no patient IDs**. No DR images are also mostly much lower-resolution than the diseased ones (Section 1.10b), which points to a different camera and is a possible shortcut.\n\n**Classes (ICDR 0–4 scale):** 0 No DR · 1 Mild NPDR · 2 Moderate NPDR · 3 Severe NPDR · 4 Proliferative DR (PDR). The 5-stage output is also collapsed into *DR vs No DR* and *referable DR (stage ≥ 2)* screening results in Section 7.4d.\n\n### How to run\n- **Kaggle (recommended):** join the APTOS 2019 competition and accept its rules, then in the editor use **Add Input → Competitions → APTOS 2019 Blindness Detection** (mounted at `/kaggle/input/aptos2019-blindness-detection/`). Turn on the GPU and Internet (needed for ImageNet weights, `pip` and the repo clone), then use **Save Version → Save & Run All** so the run continues in the background.\n- **Google Colab / local:** upload `kaggle.json` when asked; Section 1.5 runs `kaggle competitions download -c aptos2019-blindness-detection` and extracts only `train.csv` and `train_images/`. You can also point the `APTOS_DATA_DIR` environment variable at an existing copy. Enable a GPU.\n\n**Integrity rules.** Section 4.1 builds a **duplicate-aware stratified 70/15/15 split** (perceptual hashing + union-find; whole duplicate groups placed by a seeded, label-only split search, with groups of more than 50 images kept in training) and asserts that no image or near-duplicate group appears in two splits. The **test split is used once**, in the Section 7 final evaluation; early stopping, pilot selection, TTA choice and QWK thresholds all use validation only. Every reported number is computed by the notebook.\n\n**Defaults (anti-overfitting):** 300×300 input, dropout 0.5, Phase 1 LR 1e-3 (≤8 epochs), Phase 2 LR 1e-5 with AdamW weight decay 1e-4 (≤25 epochs), top-120 fine-tuning with BatchNorm frozen, early stopping and checkpointing on **validation loss** (patience 4), balanced class weights on. With `RUN_HYPERPARAMETER_TUNING=True`, pilot trials (hyperparameters, balancing, preprocessing/augmentation ablations, five ImageNet backbones and a from-scratch CNN) train on the full training split and are compared on the full validation split; the best selectable trial then trains the final model. Pilot metrics are screening evidence, not final results.\n\n---\n## Pipeline Sections\n1. Setup & Data Acquisition\n2. Preprocessing (incl. step-by-step gallery, histograms, edge detection and morphology)\n3. Data Augmentation & Class Balancing\n4. Duplicate-Aware Train / Validation / Test Split\n5. CNN Architecture & Transfer Learning (EfficientNetB3)\n6. Training Strategy (pilot trials, ablations, backbone comparison, two-phase fine-tuning)\n7. Evaluation (accuracy, QWK, precision/recall/F1, confusion matrix, ROC, DR/referable-DR screening)\n8. Explainability — Grad-CAM\n9. Innovation A — Embedding-Based Similar-Case Retrieval\n10. Innovation B — Multi-Agent Clinical Decision Pipeline\n11. UI — Gradio Interface\n12. Exploratory — U-Net Pseudo-Mask Demonstration\n13. Research Evidence — Error Analysis & Reproducibility\n---","metadata":{}},{"id":"4a7879456702438abe5457a64ca716cd","cell_type":"markdown","source":"## Section 1: Setup & Data Acquisition\n**Goal:** Install dependencies, configure all project constants in one place (Config class),\nverify GPU, locate or download the APTOS 2019 competition data, and load/validate the labels —\nincluding corrupt-image filtering, a class-distribution summary and a resolution audit.","metadata":{"id":"4a7879456702438abe5457a64ca716cd","language":"markdown"}},{"id":"fccfba0a","cell_type":"markdown","source":"## Project Mindmap\n\nThe mindmap gives a high-level overview of the project pipeline, from dataset acquisition to deployment, and was used during planning to scope each deliverable.\n\n![RetinaGuard AI project mindmap](https://raw.githubusercontent.com/ShazzySal/ComputerVision_CW/main/report_images/mindmap.png)\n\n> **Figure:** Project mindmap — the central node branches into nine phases: Dataset Acquisition, Preprocessing & Enhancement, Augmentation & Balancing, Model & Transfer Learning, Training Strategy, Evaluation Metrics, Explainability Innovations (Grad-CAM, CBR, U-Net), Multi-Agent Decision Pipeline, and UI & Deployment.","metadata":{}},{"id":"b5f3ecf5","cell_type":"markdown","source":"## Notebook Artifact Guide: Section-by-Section Outputs\n\nThis guide maps the notebook cells to the files and visible outputs they are designed to produce. The notebook cells have not all been executed in this workspace, so these are expected outputs, not confirmation that every file currently exists. Most saved files go under the run-specific `Config.REPORTS_DIR`; classifier and U-Net weights go under `Config.CHECKPOINT_DIR`. In Colab or Kaggle, these are run folders, not automatically the app's root `report_images/` or `checkpoints/` folders.\n\n| Notebook section | Graphs, images, CSVs, and other outputs |\n|---|---|\n| Section 1: Setup & Data Acquisition | `class_distribution.png` graph; `resolution_by_class.csv`, `resolution_by_class_summary.csv` and `resolution_by_class.png` (native resolution vs DR stage, Section 1.10b). Corrupt-image and image-dimension checks are also printed in notebook output. `report_images/mindmap.png` is referenced as a project figure; it is not generated by the training cells. |\n| Section 2: Preprocessing | `preprocessing_class*.png` per-class before/after examples, `preprocessing_comparison.png`, `preprocessing_step_gallery.png` (each filter + histograms) and `edge_morphology_analysis.png` (Sobel, Canny, opening, top-hat, black-hat). The cells also display preprocessing examples inline. |\n| Section 3: Data Augmentation & Class Balancing | `augmentation_examples.png` (one row per DR stage) and `augmentation_class0.png`-`augmentation_class4.png` (three strips per stage); augmentation previews and calculated class weights are also displayed in the notebook. |\n| Section 4: Train / Validation / Test Split | `train_split.csv`, `validation_split.csv`, `test_split.csv` (also copied to the project's `splits/` folder), `duplicate_groups.csv`, `deduplication_log.csv` (including the chosen split-search seed and distance), `duplicate_group_sizes.csv`, `duplicate_group_sizes.png`, `largest_duplicate_groups.png`, `duplicate_examples.png` and `dataset_audit_summary.csv` (image and duplicate-group overlap counts); dataset audit and split-count tables; `preprocessing_quality_ablation.csv`, `preprocessing_quality_summary.csv`, `split_class_counts.csv`, `class_index_check.csv` (label -> class name -> count per split), `imbalance_ratio_summary.csv`, and `class_weighting_effect.csv` (balanced and sqrt weights); plots `preprocessing_contrast_sharpness.png`, `preprocessing_ablation_examples.png`, `split_class_distribution.png`, and `class_weighting_effect.png`. Some tables and checks are inline only. |\n| Section 5: CNN Architecture & Transfer Learning | EfficientNetB3 architecture/model summary and configuration are printed inline. No separate graph or CSV is saved by this section. |\n| Section 6: Training Strategy | `checkpoints/best_phase1.weights.h5` and `checkpoints/best_phase2.weights.h5`; `phase1_training_history.csv`, `phase2_training_history.csv`, and `experiment_log.csv`; `overfitting_diagnostics.csv`; `learning_rate_history.csv` and `lr_curve_final.png`; when pilot tuning is enabled, `hyperparameter_tuning_results.csv`, `model_comparison_table.csv`, `comparison_trials_results.csv` (balancing, ablations, backbones and scratch CNN vs baseline), `pilot_comparison.png` (validation QWK and macro-F1 by trial group), `pilot_recall_heatmap.png`, and per trial `pilots/<trial_id>/history.csv`, `training_curves.png` and `validation_confusion_matrix.png`; phase learning-curve PNGs (accuracy and loss with augmented-train, clean-train-subset and validation lines, plus val_qwk) under the reports folder; phase history CSVs include `clean_train_loss`, `clean_train_accuracy`, `val_qwk` and `learning_rate`. Training logs and model summaries also appear inline. |\n| Section 7: Evaluation | `test_predictions.npz`, `model_vs_majority_baseline.csv`, `test_loss_and_metrics.csv` (single test-loss number) and `test_decision_rule_summary.csv` (argmax vs validation-QWK thresholds: accuracy / QWK / macro-F1); `classification_report.csv`, `confusion_matrix.png`, `top_confusion_pairs.csv`, `classification_report_argmax.csv`, and `confusion_matrix_argmax.png`; `roc_auc_curves.png` and `roc_auc_summary.csv`; `per_class_recall_comparison.csv` and `per_class_recall_comparison.png`; `screening_metrics.csv` and `screening_confusion_matrices.png` (DR vs No DR, referable DR); validation calibration/class-weight evidence CSVs. Metric tables and error interpretations are also displayed inline. |\n| Section 8: Explainability — Grad-CAM | `gradcam_correct_and_incorrect_validation.png`, showing validation images, heatmaps, and overlays. Quadrant descriptions and example details are printed inline. |\n| Section 9: Similar-Case Retrieval | `artifacts/embeddings.npz` with the 256-D reference embeddings, labels, and file paths; `similar_cases_demo.png`; `retrieval_validation_metrics.csv`, `retrieval_validation_by_stage.csv`, and `retrieval_validation_comparison.png`. Retrieved cases and metric tables are also shown inline. |\n| Section 10: Multi-Agent Decision Pipeline | Agent classes and pipeline outputs are demonstrated in notebook output. This section does not save a separate report image or CSV. |\n| Section 11: Optional Gradio Interface | Builds the interface in memory and prints launch/status messages. The optional launch is commented out, so running all cells does not itself create a report file or start the web app. |\n| Section 12: U-Net Pseudo-Mask Demonstration | `checkpoints/unet_pseudomask.weights.h5` and `three_layer_explainability_stack.png`; pseudo-mask examples, training metrics, and comparisons are displayed inline. The masks are synthetic targets, not expert annotations. The saved checkpoint name differs from the deployed app's configured `unet_lesion_best.weights.h5`. |\n| Section 13: Research evidence | `report_images/error_analysis/error_cases.csv`, `report_images/error_analysis/error_analysis_summary.json`, and `report_images/reproducibility_manifest.json`; split CSVs are referenced by the manifest. This section creates analysis/evidence files and records configured ablation variants; model-level ablations are run by the Section 6 pilot runner. |\n\nThe exact run folder depends on the environment selected in the configuration cell. Keep artifacts from one run together so metrics, split manifests, checkpoints, and embeddings are not mixed across runs.","metadata":{}},{"id":"f666667388764c358f7f0ce6ea351a5f","cell_type":"code","source":"# ============================================================\n# SECTION 1.1 — INSTALL ONLY MISSING DEPENDENCIES\n# Keep Kaggle's preinstalled TensorFlow, NumPy, and other packages\n# intact; install only distributions whose import is unavailable.\n# ============================================================\n\nimport importlib.util\nimport subprocess\nimport sys\n\n_required_packages = {\n    \"kaggle\": \"kaggle\",\n    \"sklearn\": \"scikit-learn\",\n    \"seaborn\": \"seaborn\",\n    \"PIL\": \"pillow\",\n    \"cv2\": \"opencv-python-headless\",\n    \"gradio\": \"gradio\",\n    \"matplotlib\": \"matplotlib\",\n    \"tqdm\": \"tqdm\",\n    \"pandas\": \"pandas\",\n    \"numpy\": \"numpy\",\n    \"tensorflow\": \"tensorflow\",\n    \"imagehash\": \"imagehash\",  # perceptual hashing for the duplicate-aware split\n}\n_missing_packages = [\n    package\n    for module, package in _required_packages.items()\n    if importlib.util.find_spec(module) is None\n]\nif _missing_packages:\n    subprocess.check_call(\n        [sys.executable, \"-m\", \"pip\", \"install\", \"-q\", *_missing_packages]\n    )\nelse:\n    print(\"[Dependencies] All required packages are already available.\")","metadata":{"id":"f666667388764c358f7f0ce6ea351a5f","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:49:43.271163Z","iopub.execute_input":"2026-09-30T15:49:43.271353Z","iopub.status.idle":"2026-09-30T15:49:43.296683Z","shell.execute_reply.started":"2026-09-30T15:49:43.27133Z","shell.execute_reply":"2026-09-30T15:49:43.29595Z"}},"outputs":[],"execution_count":null},{"id":"b6bd785238a04513ad5068cac1ae33f5","cell_type":"code","source":"# ============================================================\n# SECTION 1.2 — GLOBAL CONFIGURATION (single source of truth)\n# All magic numbers and tunable hyperparameters live here.\n# Modify Config values rather than hunting through code.\n#\n# Dataset: APTOS 2019 Blindness Detection (Kaggle competition).\n# Only the 3,662 labelled images in train_images/ are used; the\n# competition's test_images/ folder has no public labels and is ignored.\n# ============================================================\n\nimport os\nimport random\nimport datetime\nfrom pathlib import Path\nfrom typing import List, Optional\nimport numpy as np\nimport tensorflow as tf\n\n_IN_KAGGLE = os.path.isdir(\"/kaggle/input\") and os.path.isdir(\"/kaggle/working\")\n_IN_COLAB = False\n_DRIVE_MOUNTED = False\nif not _IN_KAGGLE:\n    try:\n        from google.colab import drive\n        _IN_COLAB = True\n        try:\n            drive.mount(\"/content/drive\")\n            _DRIVE_MOUNTED = True\n        except (NotImplementedError, OSError, ValueError) as exc:\n            print(f\"[Environment] Google Drive mount unavailable; using Colab temporary storage. ({exc})\")\n    except ImportError:\n        pass\n\n_PROJECT_ROOT = os.getcwd()\n_APTOS_COMPETITION = \"aptos2019-blindness-detection\"\n\n\ndef _is_aptos_root(candidate: Path) -> bool:\n    \"\"\"True when a folder holds the APTOS train.csv and train_images/ folder.\"\"\"\n    return (candidate / \"train.csv\").is_file() and (candidate / \"train_images\").is_dir()\n\n\ndef _find_aptos_root() -> Optional[str]:\n    \"\"\"\n    Locate the APTOS 2019 dataset root for the current environment.\n\n    Order: APTOS_DATA_DIR environment variable, the standard Kaggle\n    competition mount, any other attached Kaggle input with the same\n    layout, then local/Colab folders. Returns None if nothing is found.\n    \"\"\"\n    candidates = []\n    if os.environ.get(\"APTOS_DATA_DIR\"):\n        candidates.append(Path(os.environ[\"APTOS_DATA_DIR\"]))\n    if _IN_KAGGLE:\n        kaggle_input = Path(\"/kaggle/input\")\n        candidates.append(kaggle_input / _APTOS_COMPETITION)\n        candidates.append(kaggle_input / \"competitions\" / _APTOS_COMPETITION)\n        candidates.extend(path.parent for path in kaggle_input.glob(\"*/train.csv\"))\n        candidates.extend(path.parent for path in kaggle_input.glob(\"*/*/train.csv\"))\n    else:\n        candidates.append(Path(_PROJECT_ROOT) / \"data\" / _APTOS_COMPETITION)\n        candidates.append(Path(_PROJECT_ROOT) / \"data\" / \"aptos2019\")\n        candidates.append(Path(\"/content\") / _APTOS_COMPETITION)\n    for candidate in candidates:\n        if _is_aptos_root(candidate):\n            return str(candidate)\n    return None\n\n\n_DATA_DIR = _find_aptos_root()\nif _DATA_DIR is None:\n    if _IN_KAGGLE:\n        raise FileNotFoundError(\n            \"The APTOS 2019 Blindness Detection data is not attached. In the Kaggle editor use \"\n            \"Add Input -> Competitions -> 'APTOS 2019 Blindness Detection' (join the competition and \"\n            \"accept its rules first), then rerun this cell. Expected /kaggle/input/\"\n            f\"{_APTOS_COMPETITION}/train.csv and train_images/.\"\n        )\n    # Colab / local: Section 1.5 downloads the competition files into this folder.\n    _DATA_DIR = str(Path(\"/content\" if _IN_COLAB else _PROJECT_ROOT) / \"data\" / _APTOS_COMPETITION)\nprint(f\"[Dataset] APTOS 2019 root: {_DATA_DIR}\")\n\n\n\nclass Config:\n    \"\"\"Central configuration for data, training, and output paths.\"\"\"\n\n    SEED: int = 42\n\n    # Dataset paths (APTOS 2019 Blindness Detection, labelled training images only)\n    KAGGLE_COMPETITION: str = _APTOS_COMPETITION\n    DATA_DIR: str = _DATA_DIR\n    LABELS_CSV: str = os.path.join(_DATA_DIR, \"train.csv\")          # id_code, diagnosis\n    TRAIN_IMAGES_DIR: str = os.path.join(_DATA_DIR, \"train_images\")  # <id_code>.png\n    EXPECTED_IMAGE_COUNT: int = 3662  # published size of the labelled APTOS training set\n\n    # Image preprocessing\n    IMG_SIZE: int = 300\n    BEN_GRAHAM_SIGMA: int = 10\n    BEN_GRAHAM_ALPHA: float = 4.0\n    BEN_GRAHAM_BETA: float = -4.0\n    BEN_GRAHAM_GAMMA: float = 128\n\n    # Class labels (ICDR 0-4)\n    CLASS_NAMES: list = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative DR\"]\n    NUM_CLASSES: int = 5\n\n    # Duplicate-aware split (Section 4.1)\n    TRAIN_RATIO: float = 0.70\n    VAL_RATIO: float = 0.15\n    TEST_RATIO: float = 0.15\n    PHASH_SIZE: int = 256                 # images are cropped and resized to this before hashing\n    DUPLICATE_HAMMING_THRESHOLD: int = 4  # 64-bit pHash distance <= 4 counts as a near-duplicate\n\n    # Model and training hyperparameters (anti-overfitting defaults)\n    BATCH_SIZE: int = 16\n    PHASE1_EPOCHS: int = 8\n    PHASE2_EPOCHS: int = 25\n    PHASE1_LR: float = 1e-3\n    PHASE2_LR: float = 1e-5\n    WEIGHT_DECAY: float = 1e-4      # AdamW decoupled weight decay in Phase 2\n    PHASE2_LR_SCHEDULE: str = \"plateau\"  # \"plateau\" (ReduceLROnPlateau) or \"cosine\" (warm-up + cosine)\n    LR_WARMUP_EPOCHS: int = 2            # cosine mode: linear warm-up from 10% of PHASE2_LR\n    LABEL_SMOOTHING: float = 0.1\n    DROPOUT_RATE: float = 0.5\n    DENSE_UNITS: int = 256\n    UNFREEZE_TOP_N: int = 120       # fine-tune only the top 120 EfficientNetB3 layers (385 in Kaggle's Keras 3.13)\n    USE_CLASS_WEIGHTS: bool = True  # class weights from the training split only\n    CLASS_WEIGHT_MODE: str = \"balanced\"  # \"balanced\", \"sqrt\" (softer) or \"none\"; must agree with USE_CLASS_WEIGHTS\n    USE_OVERSAMPLING: bool = False  # on-the-fly class-balanced sampling (never with class weights)\n\n    # Callbacks (all monitor validation loss; val_qwk is logged but not used for stopping)\n    ES_PATIENCE: int = 4\n    RLROP_FACTOR: float = 0.5\n    RLROP_PATIENCE: int = 2\n    CLEAN_TRAIN_SUBSET_SIZE: int = 500  # fixed training subset for clean-train diagnostic curves\n\n    # Pilot runner (Section 6.2b)\n    # Skipped trials stay defined in 6.2b and are written to the tuning CSV as \"skipped_by_config\".\n    # They were dropped after an earlier validation-only pilot showed they underperformed the\n    # baseline or duplicated other evidence, to fit the compute budget. Remove an ID to re-enable it.\n    PILOT_SKIP_TRIALS: List[str] = [\n        \"phase2_lr_3e5\", \"unfreeze_top240\", \"dropout_040\", \"ablation_all_filters\", \"backbone_vgg16\",\n    ]\n    PILOT_TRIAL_LIMIT: Optional[int] = None         # e.g. 3 for a memory test: only the first N non-skipped trials\n    PILOT_TIME_BUDGET_HOURS: Optional[float] = 5.0  # no new pilot trial starts after this many hours\n    # Set to an earlier run ID (e.g. \"20260930_114110\") to reuse retinatrace_runs/run_<ID>\n    # and its completed pilot trials instead of creating a new run folder.\n    RESUME_RUN_ID: Optional[str] = None\n\n    # Governance threshold\n    CONFIDENCE_THRESHOLD: float = 0.70\n\n    # Output paths (set below from the run folder)\n    CHECKPOINT_DIR: str = \"\"\n    REPORTS_DIR: str = \"\"\n    EMBEDDINGS_PATH: str = \"\"\n\n\n# Run folder: a new timestamped folder, or the earlier one named by Config.RESUME_RUN_ID.\n_RUN_ID = Config.RESUME_RUN_ID or datetime.datetime.now(datetime.timezone.utc).strftime(\"%Y%m%d_%H%M%S\")\nif _DRIVE_MOUNTED:\n    _RUN_ROOT = os.path.join(\"/content/drive/MyDrive/colab_run_evidence\", f\"run_{_RUN_ID}\")\nelif _IN_KAGGLE:\n    _RUN_ROOT = os.path.join(\"/kaggle/working/retinatrace_runs\", f\"run_{_RUN_ID}\")\nelif _IN_COLAB:\n    _RUN_ROOT = os.path.join(\"/content/retinatrace_runs\", f\"run_{_RUN_ID}\")\nelse:\n    _RUN_ROOT = os.path.join(_PROJECT_ROOT, \"retinatrace_runs\", f\"run_{_RUN_ID}\")\nif Config.RESUME_RUN_ID and not os.path.isdir(_RUN_ROOT):\n    raise FileNotFoundError(\n        f\"Config.RESUME_RUN_ID={Config.RESUME_RUN_ID!r}, but {_RUN_ROOT} does not exist. \"\n        \"Check the ID, or set RESUME_RUN_ID = None to start a new run.\"\n    )\n\n_CHECKPOINT_DIR = os.path.join(_RUN_ROOT, \"checkpoints\")\n_REPORTS_DIR = os.path.join(_RUN_ROOT, \"report_images\")\n_EMBEDDINGS_PATH = os.path.join(_RUN_ROOT, \"artifacts\", \"embeddings.npz\")\nConfig.CHECKPOINT_DIR = _CHECKPOINT_DIR\nConfig.REPORTS_DIR = _REPORTS_DIR\nConfig.EMBEDDINGS_PATH = _EMBEDDINGS_PATH\nif Config.RESUME_RUN_ID:\n    print(f\"[Resume] Reusing run folder: {_RUN_ROOT}\")\nif _IN_COLAB:\n    print(f\"Colab run ID: {_RUN_ID}\")\n    print(f\"Run evidence folder: {_RUN_ROOT}\")\nelif _IN_KAGGLE:\n    print(f\"Kaggle output folder: {_RUN_ROOT}\")\n\n\ndef set_all_seeds(seed: int = Config.SEED) -> None:\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    print(f\"[Config] All random seeds set to {seed}.\")\n\n\nset_all_seeds()\n","metadata":{"id":"b6bd785238a04513ad5068cac1ae33f5","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:49:46.372387Z","iopub.execute_input":"2026-09-30T15:49:46.373041Z","iopub.status.idle":"2026-09-30T15:50:03.048141Z","shell.execute_reply.started":"2026-09-30T15:49:46.37301Z","shell.execute_reply":"2026-09-30T15:50:03.047354Z"}},"outputs":[],"execution_count":null},{"id":"cc0a71fb0b4e4a9fa38c7c067a62cb34","cell_type":"code","source":"# ============================================================\n# SECTION 1.3 — IMPORT ALL LIBRARIES\n# Centralised imports make missing-dependency errors easy to spot\n# and prevent the notebook from failing midway through a long run.\n# ============================================================\n\nimport sys\nimport glob\nimport shutil\nimport warnings\nimport pathlib\nfrom typing import Any, Dict, List, Optional, Tuple, Union\n\nimport cv2\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image, UnidentifiedImageError\nimport imagehash\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split, StratifiedGroupKFold\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import classification_report, confusion_matrix, cohen_kappa_score\n\nfrom tensorflow import keras\nfrom tensorflow.keras import layers, Model\nfrom tensorflow.keras.applications import EfficientNetB3\nfrom tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau, ModelCheckpoint\n\nwarnings.filterwarnings(\"ignore\")\n\nfor _dir in [\n    Config.CHECKPOINT_DIR,\n    Config.REPORTS_DIR,\n    Config.DATA_DIR,\n    os.path.dirname(Config.EMBEDDINGS_PATH),\n]:\n    os.makedirs(_dir, exist_ok=True)\n\nprint(f\"Python     : {sys.version}\")\nprint(f\"TensorFlow : {tf.__version__}\")\nprint(f\"Keras      : {keras.__version__}\")\nprint(\"All libraries imported successfully.\")","metadata":{"id":"cc0a71fb0b4e4a9fa38c7c067a62cb34","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:50:11.7088Z","iopub.execute_input":"2026-09-30T15:50:11.709429Z","iopub.status.idle":"2026-09-30T15:50:12.224194Z","shell.execute_reply.started":"2026-09-30T15:50:11.709399Z","shell.execute_reply":"2026-09-30T15:50:12.223479Z"}},"outputs":[],"execution_count":null},{"id":"fa68061255754f1daf885a9033740f40","cell_type":"code","source":"# ============================================================\n# SECTION 1.4 — GPU VERIFICATION\n# Training EfficientNetB3 on CPU, even for 3,662 images, is very slow.\n# We detect GPUs here and warn loudly if none is found.\n# In Colab: Runtime > Change runtime type > GPU (T4 or A100).\n# ============================================================\n\ndef verify_gpu() -> None:\n    \"\"\"\n    Detect and report available GPU devices.\n\n    Enables memory growth on each GPU to prevent TensorFlow from\n    grabbing all VRAM at startup, which can cause OOM errors when\n    sharing a Colab GPU with other processes.\n\n    Prints a warning (but does not raise) if no GPU is found, so\n    the notebook can still be used for small-batch testing on CPU.\n    \"\"\"\n    gpus = tf.config.list_physical_devices(\"GPU\")\n    if gpus:\n        for gpu in gpus:\n            # Incremental VRAM allocation rather than reserving all memory up-front\n            tf.config.experimental.set_memory_growth(gpu, True)\n        print(f\"[GPU] {len(gpus)} GPU(s) detected:\")\n        for g in gpus:\n            print(f\"      {g.name}\")\n        os.system(\"nvidia-smi --query-gpu=name,memory.total --format=csv,noheader 2>/dev/null\")\n    else:\n        print(\n            \"[WARNING] No GPU detected. Training will be very slow on CPU.\\n\"\n            \"          In Google Colab: Runtime > Change runtime type > GPU.\"\n        )\n\n\nverify_gpu()","metadata":{"id":"fa68061255754f1daf885a9033740f40","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:50:16.521201Z","iopub.execute_input":"2026-09-30T15:50:16.52197Z","iopub.status.idle":"2026-09-30T15:50:17.890485Z","shell.execute_reply.started":"2026-09-30T15:50:16.52194Z","shell.execute_reply":"2026-09-30T15:50:17.889798Z"}},"outputs":[],"execution_count":null},{"id":"bd82a8dca2d043adaa55f3a88a546b9e","cell_type":"code","source":"# ============================================================\n# SECTION 1.5 — KAGGLE COMPETITION DATA (APTOS 2019)\n# On Kaggle, attach the competition with Add Input -> Competitions ->\n# \"APTOS 2019 Blindness Detection\"; it is mounted read-only, so no token\n# is needed. On Colab/local runs the data is downloaded once with\n#   kaggle competitions download -c aptos2019-blindness-detection\n# which requires a kaggle.json token AND having accepted the competition\n# rules on kaggle.com. Only train.csv and train_images/ are extracted.\n# ============================================================\n\nimport zipfile\n\n\ndef setup_kaggle_credentials() -> None:\n    \"\"\"\n    Place kaggle.json in ~/.kaggle/kaggle.json with correct permissions.\n\n    The Kaggle CLI refuses to run if permissions on kaggle.json are too\n    permissive (it is a security token). We enforce 600 (owner read/write).\n\n    Raises:\n        FileNotFoundError: If kaggle.json cannot be found or uploaded.\n    \"\"\"\n    kaggle_dir = os.path.expanduser(\"~/.kaggle\")\n    os.makedirs(kaggle_dir, exist_ok=True)\n    target = os.path.join(kaggle_dir, \"kaggle.json\")\n\n    if os.path.exists(target):\n        print(\"[Kaggle] kaggle.json already in place — skipping upload.\")\n        return\n\n    try:\n        from google.colab import files  # type: ignore\n        print(\"[Kaggle] Please upload your kaggle.json when the dialog appears...\")\n        uploaded = files.upload()\n        if \"kaggle.json\" not in uploaded:\n            raise FileNotFoundError(\n                \"kaggle.json was not found in the uploaded files. \"\n                \"Download it from https://www.kaggle.com/settings > API > Create New Token.\"\n            )\n        shutil.move(\"kaggle.json\", target)\n    except ImportError:\n        if os.path.exists(\"kaggle.json\"):\n            shutil.copy(\"kaggle.json\", target)\n        else:\n            raise FileNotFoundError(\n                \"kaggle.json not found in the current directory. \"\n                \"Place it here to download the competition data with the Kaggle API.\"\n            )\n\n    os.chmod(target, 0o600)\n    print(f\"[Kaggle] Credentials saved to {target} with 600 permissions.\")\n\n\ndef download_aptos_competition(data_dir: str = Config.DATA_DIR) -> None:\n    \"\"\"\n    Download the APTOS 2019 competition archive and extract the labelled training data.\n\n    Only train.csv and train_images/ are extracted; test_images/ has no public\n    labels and is not used anywhere in this notebook.\n\n    Args:\n        data_dir: Destination folder (becomes Config.DATA_DIR).\n\n    Raises:\n        RuntimeError: If the Kaggle CLI fails (missing token or rules not accepted).\n    \"\"\"\n    os.makedirs(data_dir, exist_ok=True)\n    archive = os.path.join(data_dir, f\"{Config.KAGGLE_COMPETITION}.zip\")\n    if not os.path.exists(archive):\n        print(f\"[Dataset] Downloading competition '{Config.KAGGLE_COMPETITION}' — this may take several minutes...\")\n        exit_code = os.system(\n            f\"kaggle competitions download -c {Config.KAGGLE_COMPETITION} -p {data_dir} -q\"\n        )\n        if exit_code != 0 or not os.path.exists(archive):\n            raise RuntimeError(\n                f\"Kaggle download failed (exit code {exit_code}). Check kaggle.json and make sure you have \"\n                \"joined the APTOS 2019 competition and accepted its rules on kaggle.com.\"\n            )\n    with zipfile.ZipFile(archive) as bundle:\n        members = [\n            name for name in bundle.namelist()\n            if name == \"train.csv\" or name.startswith(\"train_images/\")\n        ]\n        bundle.extractall(data_dir, members=members)\n    print(f\"[Dataset] Extracted {len(members):,} files (train.csv + train_images/) -> {data_dir}\")\n\n\nif _is_aptos_root(Path(Config.DATA_DIR)):\n    print(f\"[Dataset] Using the APTOS 2019 data already available at {Config.DATA_DIR}; download skipped.\")\nelse:\n    setup_kaggle_credentials()\n    download_aptos_competition(Config.DATA_DIR)\nif not _is_aptos_root(Path(Config.DATA_DIR)):\n    raise FileNotFoundError(f\"APTOS train.csv/train_images/ not found under {Config.DATA_DIR}.\")\n","metadata":{"id":"bd82a8dca2d043adaa55f3a88a546b9e","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:50:22.303348Z","iopub.execute_input":"2026-09-30T15:50:22.303936Z","iopub.status.idle":"2026-09-30T15:50:22.317272Z","shell.execute_reply.started":"2026-09-30T15:50:22.303909Z","shell.execute_reply":"2026-09-30T15:50:22.316356Z"}},"outputs":[],"execution_count":null},{"id":"efca0f661d764c7fab994ad7adcedbb0","cell_type":"code","source":"# ============================================================\n# SECTION 1.6 — DISCOVER DATASET STRUCTURE\n# Print the top-level directory tree so we know exactly what Kaggle\n# delivered before attempting label loading.\n# ============================================================\n\ndef print_directory_tree(root: str, max_depth: int = 3, max_items: int = 20) -> None:\n    \"\"\"\n    Print a human-readable directory tree for quick structure inspection.\n\n    Args:\n        root:      Root directory to start from.\n        max_depth: Maximum folder depth to recurse into.\n        max_items: Max items shown per directory (prevents flooding output).\n    \"\"\"\n    root_path = pathlib.Path(root)\n    if not root_path.exists():\n        print(f\"[Tree] Path does not exist: {root}\")\n        return\n\n    def _recurse(path: pathlib.Path, depth: int, prefix: str) -> None:\n        if depth > max_depth:\n            return\n        try:\n            children = sorted(path.iterdir())\n        except PermissionError:\n            return\n        dirs  = [c for c in children if c.is_dir()]\n        files = [c for c in children if c.is_file()]\n        items = dirs + files\n        truncated = len(items) > max_items\n        items = items[:max_items]\n        for i, item in enumerate(items):\n            is_last = (i == len(items) - 1) and not truncated\n            connector = \"\\u2514\\u2500\\u2500 \" if is_last else \"\\u251c\\u2500\\u2500 \"\n            suffix = \"/\" if item.is_dir() else \"\"\n            print(f\"{prefix}{connector}{item.name}{suffix}\")\n            if item.is_dir():\n                ext = \"    \" if is_last else \"\\u2502   \"\n                _recurse(item, depth + 1, prefix + ext)\n        if truncated:\n            print(f\"{prefix}    ... (showing first {max_items} items)\")\n\n    print(f\"\\n[Tree] {root}/\")\n    _recurse(root_path, 1, \"\")\n\n\nprint_directory_tree(Config.DATA_DIR)","metadata":{"id":"efca0f661d764c7fab994ad7adcedbb0","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:50:32.164895Z","iopub.execute_input":"2026-09-30T15:50:32.165298Z","iopub.status.idle":"2026-09-30T15:50:51.585933Z","shell.execute_reply.started":"2026-09-30T15:50:32.165272Z","shell.execute_reply":"2026-09-30T15:50:51.58506Z"}},"outputs":[],"execution_count":null},{"id":"02b64a32fa8e49328dc67e8c37fbecbe","cell_type":"code","source":"# ============================================================\n# SECTION 1.7 — LABEL LOADING (APTOS 2019 train.csv)\n#\n# train.csv has two columns:\n#   id_code   -> image file name without extension (train_images/<id_code>.png)\n#   diagnosis -> ICDR grade 0-4 assigned by clinicians (0 No DR ... 4 PDR)\n#\n# APTOS 2019 is a single-source dataset (Aravind Eye Hospital, India),\n# already graded on the ICDR 0-4 scale, so no cross-dataset label mapping\n# is required. APTOS provides no patient identifiers.\n# ============================================================\n\ndef load_aptos_labels(\n    labels_csv: str = Config.LABELS_CSV,\n    image_dir: str = Config.TRAIN_IMAGES_DIR,\n) -> pd.DataFrame:\n    \"\"\"\n    Load the APTOS 2019 training labels and resolve each image path.\n\n    Validates that every label is an ICDR grade in {0,1,2,3,4}, drops rows\n    whose image file is missing (reporting how many), and warns if the row\n    count differs from the published 3,662 labelled images.\n\n    Args:\n        labels_csv: Path to the competition train.csv.\n        image_dir:  Folder containing <id_code>.png files.\n\n    Returns:\n        DataFrame with columns 'id_code', 'filepath' (absolute) and 'label' (int 0-4).\n    \"\"\"\n    raw_df = pd.read_csv(labels_csv)\n    missing_columns = {\"id_code\", \"diagnosis\"} - set(raw_df.columns)\n    if missing_columns:\n        raise ValueError(f\"{labels_csv} is missing columns {sorted(missing_columns)}; expected id_code, diagnosis.\")\n    print(f\"[Label] {len(raw_df):,} rows loaded from {labels_csv}\")\n\n    df = pd.DataFrame({\n        \"id_code\": raw_df[\"id_code\"].astype(str).str.strip(),\n        \"label\": pd.to_numeric(raw_df[\"diagnosis\"], errors=\"coerce\"),\n    })\n    bad_mask = ~df[\"label\"].isin({0, 1, 2, 3, 4})\n    if bad_mask.any():\n        raise ValueError(\n            f\"{int(bad_mask.sum())} rows have labels outside ICDR 0-4: \"\n            f\"{df.loc[bad_mask, 'label'].unique().tolist()[:10]}\"\n        )\n    if df[\"id_code\"].duplicated().any():\n        raise ValueError(\"train.csv contains repeated id_code values.\")\n    df[\"label\"] = df[\"label\"].astype(int)\n    df[\"filepath\"] = [\n        str((Path(image_dir) / f\"{id_code}.png\").resolve()) for id_code in df[\"id_code\"]\n    ]\n\n    exists_mask = df[\"filepath\"].map(os.path.isfile)\n    if not exists_mask.all():\n        print(f\"[Label] WARNING: {int((~exists_mask).sum())} labelled image(s) not found on disk. Removing.\")\n        df = df[exists_mask].reset_index(drop=True)\n    if len(df) != Config.EXPECTED_IMAGE_COUNT:\n        print(\n            f\"[Label] WARNING: {len(df):,} usable images, but the APTOS 2019 training set has \"\n            f\"{Config.EXPECTED_IMAGE_COUNT:,}. Check that the full competition data is attached.\"\n        )\n    print(f\"[Label] Final usable dataset size: {len(df):,} images.\")\n    return df[[\"id_code\", \"filepath\", \"label\"]]\n\n\ndf_labels = load_aptos_labels()\ndf_labels.head()\n","metadata":{"id":"02b64a32fa8e49328dc67e8c37fbecbe","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:51:07.073477Z","iopub.execute_input":"2026-09-30T15:51:07.073932Z","iopub.status.idle":"2026-09-30T15:51:11.207611Z","shell.execute_reply.started":"2026-09-30T15:51:07.073901Z","shell.execute_reply":"2026-09-30T15:51:11.206938Z"}},"outputs":[],"execution_count":null},{"id":"59560ab455c640ea8ea8ef2157330c27","cell_type":"code","source":"# ============================================================\n# SECTION 1.8 — CORRUPT IMAGE DETECTION\n# Corrupt images (truncated, zero-byte, unreadable) crash the pipeline\n# or silently produce garbage. PIL.Image.verify() reads only the header\n# and raises on corrupt files — much faster than full pixel decoding.\n# ============================================================\n\ndef filter_corrupt_images(\n    df: pd.DataFrame, filepath_col: str = \"filepath\"\n) -> pd.DataFrame:\n    \"\"\"\n    Remove rows referencing corrupt or unreadable image files.\n\n    Uses PIL.Image.verify() which reads image metadata without decoding\n    pixel data. This catches JPEG truncation and file-format errors that\n    OpenCV silently ignores (returning a black or partial image instead).\n\n    Args:\n        df:           DataFrame with a file path column.\n        filepath_col: Name of the file path column.\n\n    Returns:\n        DataFrame with corrupt-image rows removed.\n    \"\"\"\n    corrupt_indices = []\n\n    for idx, row in tqdm(df.iterrows(), total=len(df), desc=\"Checking image integrity\"):\n        try:\n            with Image.open(row[filepath_col]) as img:\n                img.verify()  # header-only check — does not decode pixel data\n        except (UnidentifiedImageError, OSError, SyntaxError):\n            # SyntaxError is raised by PIL for some truncated JPEG files\n            corrupt_indices.append(idx)\n\n    if corrupt_indices:\n        examples = df.loc[corrupt_indices, filepath_col].tolist()[:5]\n        print(f\"[Corrupt] {len(corrupt_indices)} corrupt image(s) found and removed. Examples:\")\n        for p in examples:\n            print(f\"  {p}\")\n        df = df.drop(index=corrupt_indices).reset_index(drop=True)\n    else:\n        print(\"[Corrupt] All images passed integrity check — dataset is clean.\")\n\n    return df\n\n\ndf_labels = filter_corrupt_images(df_labels)\nprint(f\"Dataset size after integrity check: {len(df_labels):,} images.\")","metadata":{"id":"59560ab455c640ea8ea8ef2157330c27","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:51:20.466956Z","iopub.execute_input":"2026-09-30T15:51:20.467359Z","iopub.status.idle":"2026-09-30T15:53:01.19731Z","shell.execute_reply.started":"2026-09-30T15:51:20.467334Z","shell.execute_reply":"2026-09-30T15:53:01.196403Z"}},"outputs":[],"execution_count":null},{"id":"1243fe34afe3417f9fcba5a8b1756388","cell_type":"code","source":"# ============================================================\n# SECTION 1.9 — CLASS DISTRIBUTION SUMMARY & VISUALISATION\n# Print a per-class count table and plot a bar chart using the\n# APTOS 2019 training labels.\n# ============================================================\n\ndef print_class_distribution(df: pd.DataFrame, label_col: str = \"label\") -> pd.DataFrame:\n    \"\"\"Compute and print a class distribution summary table.\"\"\"\n    counts = df[label_col].value_counts().sort_index()\n    total = counts.sum()\n\n    summary = pd.DataFrame({\n        \"class_id\": counts.index,\n        \"class_name\": [Config.CLASS_NAMES[i] for i in counts.index],\n        \"count\": counts.values,\n        \"percentage\": (counts.values / total * 100).round(2),\n    })\n\n    print(\"\\n\" + \"=\" * 58)\n    print(f\"  CLASS DISTRIBUTION  (total: {total:,} images)\")\n    print(\"=\" * 58)\n    print(summary.to_string(index=False))\n    print(\"-\" * 58)\n    ratio = counts.max() / counts.min()\n    print(f\"  Imbalance ratio (max / min class): {ratio:.1f}x\")\n    if ratio > 1:\n        print(f\"  [ACTION] Class weighting remains useful for the observed {ratio:.1f}x imbalance.\")\n    print(\"=\" * 58 + \"\\n\")\n    return summary\n\n\ndef plot_class_distribution(\n    df: pd.DataFrame,\n    label_col: str = \"label\",\n    save_path: Optional[str] = None,\n) -> None:\n    \"\"\"Plot and optionally save a class count bar chart.\"\"\"\n    counts = df[label_col].value_counts().reindex(range(Config.NUM_CLASSES), fill_value=0)\n\n    fig, ax = plt.subplots(figsize=(9, 5))\n    palette = sns.color_palette(\"Blues_d\", len(Config.CLASS_NAMES))\n    bars = ax.bar(\n        Config.CLASS_NAMES,\n        counts.values,\n        color=palette,\n        edgecolor=\"black\",\n        linewidth=0.7,\n    )\n    for bar, count in zip(bars, counts.values):\n        ax.text(\n            bar.get_x() + bar.get_width() / 2,\n            bar.get_height() + max(counts.max() * 0.01, 1),\n            f\"{count:,}\",\n            ha=\"center\", va=\"bottom\", fontsize=10, fontweight=\"bold\",\n        )\n\n    ax.set_title(\"Class Distribution — APTOS 2019 (labelled training images)\", fontsize=14, pad=15)\n    ax.set_xlabel(\"DR Stage (ICDR 0-4 Scale)\", fontsize=12)\n    ax.set_ylabel(\"Number of Images\", fontsize=12)\n    ax.set_ylim(0, counts.max() * 1.15)\n    plt.tight_layout()\n\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n        print(f\"[Plot] Saved -> {save_path}\")\n\n    plt.show()\n\n\nclass_summary = print_class_distribution(df_labels)\nplot_class_distribution(\n    df_labels,\n    save_path=os.path.join(Config.REPORTS_DIR, \"class_distribution.png\"),\n)","metadata":{"id":"1243fe34afe3417f9fcba5a8b1756388","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:53:21.041976Z","iopub.execute_input":"2026-09-30T15:53:21.042428Z","iopub.status.idle":"2026-09-30T15:53:21.460183Z","shell.execute_reply.started":"2026-09-30T15:53:21.042398Z","shell.execute_reply":"2026-09-30T15:53:21.4595Z"}},"outputs":[],"execution_count":null},{"id":"0316da4829484143a71dfe2b8de10cf7","cell_type":"code","source":"# ============================================================\n# SECTION 1.10 — IMAGE DIMENSION AUDIT\n# Report the native image resolution before resizing. APTOS 2019\n# images were captured on different cameras and vary widely in size.\n# We sample 500 images to report the range so the preprocessing section\n# is designed with accurate expectations.\n# ============================================================\n\ndef audit_image_dimensions(\n    df: pd.DataFrame,\n    sample_n: int = 500,\n    filepath_col: str = \"filepath\",\n) -> None:\n    \"\"\"\n    Sample images and report width/height statistics.\n\n    Reads only the image header (not full pixel data) via PIL, so\n    auditing 500 images takes only a few seconds regardless of resolution.\n\n    Args:\n        df:           DataFrame with filepath column.\n        sample_n:     Number of images to randomly sample.\n        filepath_col: Name of the file path column.\n    \"\"\"\n    sample = df.sample(min(sample_n, len(df)), random_state=Config.SEED)\n    widths, heights = [], []\n\n    for fp in tqdm(sample[filepath_col], desc=\"Auditing dimensions\"):\n        try:\n            with Image.open(fp) as img:\n                w, h = img.size  # PIL: (width, height)\n                widths.append(w)\n                heights.append(h)\n        except Exception:\n            pass  # already screened by corrupt filter\n\n    w = np.array(widths)\n    h = np.array(heights)\n    print(f\"\\n[Dimensions] Sampled {len(w)} images:\")\n    print(f\"  Width  — min:{w.min():6d}  max:{w.max():6d}  mean:{w.mean():.0f}\")\n    print(f\"  Height — min:{h.min():6d}  max:{h.max():6d}  mean:{h.mean():.0f}\")\n    print(f\"  All images will be resized to {Config.IMG_SIZE}x{Config.IMG_SIZE} in Section 2.\")\n\n\naudit_image_dimensions(df_labels)\n\n# ---- Section 1 complete ----\nprint(\"\\n\" + \"=\" * 60)\nprint(\"  SECTION 1 COMPLETE\")\nprint(f\"  Total images loaded and validated : {len(df_labels):,}\")\nprint(f\"  Label column                      : 'label'    (int, ICDR 0-4)\")\nprint(f\"  Image path column                 : 'filepath' (absolute path)\")\nprint(\"  Next -> Section 2: Preprocessing\")\nprint(\"=\" * 60)","metadata":{"id":"0316da4829484143a71dfe2b8de10cf7","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:53:36.976459Z","iopub.execute_input":"2026-09-30T15:53:36.976924Z","iopub.status.idle":"2026-09-30T15:53:37.536191Z","shell.execute_reply.started":"2026-09-30T15:53:36.976898Z","shell.execute_reply":"2026-09-30T15:53:37.535561Z"}},"outputs":[],"execution_count":null},{"id":"f6abd42e","cell_type":"code","source":"# ============================================================\n# SECTION 1.10b - RESOLUTION VS DR STAGE\n#\n# WHY: APTOS 2019 images were collected over time at several clinics\n# on different fundus cameras. If one camera contributes mostly one\n# DR stage, image resolution becomes\n# correlated with the label. A CNN can then learn a shortcut - \"large\n# or wide image => severe stage\" - instead of the retinal lesions.\n# This cell measures whether native resolution differs by stage.\n#\n# Descriptive only: it samples a COPY of df_labels, reads image\n# headers only, and never filters images or changes any split.\n# ============================================================\n\nfrom scipy.stats import kruskal\n\nRESOLUTION_SAMPLE_PER_CLASS: int = 200\nSTAGE_COLORS: List[str] = [\"#2E7D32\", \"#0284C7\", \"#D97706\", \"#EA580C\", \"#BE123C\"]\n\n\ndef sample_per_class(\n    df: pd.DataFrame,\n    n_per_class: int = RESOLUTION_SAMPLE_PER_CLASS,\n    label_col: str = \"label\",\n) -> pd.DataFrame:\n    \"\"\"\n    Draw up to n_per_class rows from every DR stage so small stages are represented.\n\n    Classes with fewer than n_per_class images contribute all of their rows.\n    The input DataFrame is not modified.\n\n    Args:\n        df:          DataFrame with filepath and label columns.\n        n_per_class: Maximum rows sampled per class.\n        label_col:   Name of the integer label column.\n\n    Returns:\n        A new DataFrame containing the sampled rows.\n    \"\"\"\n    parts = []\n    for class_id in range(Config.NUM_CLASSES):\n        class_rows = df[df[label_col] == class_id]\n        if len(class_rows):\n            parts.append(class_rows.sample(min(n_per_class, len(class_rows)), random_state=Config.SEED))\n    return pd.concat(parts, ignore_index=True) if parts else df.iloc[0:0].copy()\n\n\ndef read_image_resolutions(\n    sample: pd.DataFrame,\n    filepath_col: str = \"filepath\",\n    label_col: str = \"label\",\n) -> Tuple[pd.DataFrame, int]:\n    \"\"\"\n    Read width and height from each image header without decoding pixels.\n\n    Args:\n        sample:       DataFrame of sampled images.\n        filepath_col: Name of the file path column.\n        label_col:    Name of the integer label column.\n\n    Returns:\n        (per-image resolution table, number of unreadable files skipped)\n    \"\"\"\n    records = []\n    skipped = 0\n    for filepath, label in tqdm(\n        zip(sample[filepath_col], sample[label_col]),\n        total=len(sample),\n        desc=\"Reading image headers\",\n    ):\n        try:\n            with Image.open(filepath) as img:\n                width, height = img.size  # header only; pixels are not decoded\n        except (OSError, UnidentifiedImageError, ValueError):\n            skipped += 1\n            continue\n        records.append({\n            \"filepath\": filepath,\n            \"label\": int(label),\n            \"class_name\": Config.CLASS_NAMES[int(label)],\n            \"width\": int(width),\n            \"height\": int(height),\n            \"megapixels\": width * height / 1e6,\n            \"aspect_ratio\": width / height,\n        })\n    columns = [\"filepath\", \"label\", \"class_name\", \"width\", \"height\", \"megapixels\", \"aspect_ratio\"]\n    return pd.DataFrame(records, columns=columns), skipped\n\n\ndef summarise_resolution_by_class(resolution_df: pd.DataFrame) -> pd.DataFrame:\n    \"\"\"\n    Summarise native resolution per DR stage.\n\n    Args:\n        resolution_df: Per-image table from read_image_resolutions().\n\n    Returns:\n        One row per class with count, width/height, megapixel and aspect-ratio statistics.\n    \"\"\"\n    summary = (\n        resolution_df.groupby(\"label\")\n        .agg(\n            count=(\"filepath\", \"size\"),\n            median_width=(\"width\", \"median\"),\n            min_width=(\"width\", \"min\"),\n            max_width=(\"width\", \"max\"),\n            median_height=(\"height\", \"median\"),\n            median_megapixels=(\"megapixels\", \"median\"),\n            median_aspect_ratio=(\"aspect_ratio\", \"median\"),\n        )\n        .reindex(range(Config.NUM_CLASSES))\n    )\n    summary[\"count\"] = summary[\"count\"].fillna(0).astype(int)\n    summary.insert(0, \"class_name\", [Config.CLASS_NAMES[c] for c in summary.index])\n    summary.index.name = \"label\"\n    return summary.reset_index()\n\n\ndef plot_resolution_by_class(resolution_df: pd.DataFrame, save_path: str) -> None:\n    \"\"\"\n    Save a two-panel figure: (a) megapixels per DR stage, (b) width vs height by stage.\n\n    Args:\n        resolution_df: Per-image table from read_image_resolutions().\n        save_path:     PNG output path.\n    \"\"\"\n    present = [c for c in range(Config.NUM_CLASSES) if (resolution_df[\"label\"] == c).any()]\n    fig, (ax_box, ax_scatter) = plt.subplots(1, 2, figsize=(15, 6))\n\n    box = ax_box.boxplot(\n        [resolution_df.loc[resolution_df[\"label\"] == c, \"megapixels\"] for c in present],\n        patch_artist=True,\n        widths=0.6,\n        medianprops={\"color\": \"black\", \"linewidth\": 1.5},\n        flierprops={\"marker\": \"o\", \"markersize\": 3, \"alpha\": 0.5},\n    )\n    for patch, class_id in zip(box[\"boxes\"], present):\n        patch.set_facecolor(STAGE_COLORS[class_id])\n        patch.set_alpha(0.75)\n    ax_box.set_xticks(range(1, len(present) + 1), [Config.CLASS_NAMES[c] for c in present], rotation=15)\n    ax_box.set_title(\"(a) Native Resolution per DR Stage\", fontweight=\"bold\")\n    ax_box.set_xlabel(\"Diabetic retinopathy stage\")\n    ax_box.set_ylabel(\"Megapixels (width x height / 1e6)\")\n    ax_box.grid(axis=\"y\", alpha=0.3)\n\n    # Fundus images come in a few fixed sizes, so points from different stages often share\n    # exactly the same (width, height). Hollow rings of decreasing size per stage keep every\n    # stage visible as nested rings instead of the last-drawn stage hiding the others.\n    for class_id in present:\n        rows = resolution_df[resolution_df[\"label\"] == class_id]\n        ax_scatter.scatter(\n            rows[\"width\"], rows[\"height\"], s=190 - 35 * class_id, facecolors=\"none\",\n            edgecolors=STAGE_COLORS[class_id], linewidths=1.6, alpha=0.8,\n            label=Config.CLASS_NAMES[class_id],\n        )\n    ax_scatter.set_title(\"(b) Image Width vs Height by DR Stage\", fontweight=\"bold\")\n    ax_scatter.set_xlabel(\"Width (pixels)\")\n    ax_scatter.set_ylabel(\"Height (pixels)\")\n    ax_scatter.legend(title=\"DR stage (nested rings = same size)\")\n    ax_scatter.grid(alpha=0.3)\n\n    fig.suptitle(\n        f\"Resolution vs DR Stage (up to {RESOLUTION_SAMPLE_PER_CLASS} images per stage; \"\n        f\"all resized to {Config.IMG_SIZE}x{Config.IMG_SIZE} for training)\",\n        fontsize=12,\n    )\n    fig.tight_layout()\n    os.makedirs(os.path.dirname(save_path), exist_ok=True)\n    fig.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.show()\n\n\nresolution_sample = sample_per_class(df_labels)\nresolution_by_class, n_resolution_skipped = read_image_resolutions(resolution_sample)\nresolution_summary = summarise_resolution_by_class(resolution_by_class)\nprint(f\"\\n[Resolution] Read {len(resolution_by_class):,} image headers; skipped {n_resolution_skipped} unreadable file(s).\")\n\nresolution_csv_path = os.path.join(Config.REPORTS_DIR, \"resolution_by_class.csv\")\nresolution_summary_path = os.path.join(Config.REPORTS_DIR, \"resolution_by_class_summary.csv\")\nresolution_plot_path = os.path.join(Config.REPORTS_DIR, \"resolution_by_class.png\")\nresolution_by_class.to_csv(resolution_csv_path, index=False)\nresolution_summary.to_csv(resolution_summary_path, index=False)\n\nprint(\"\\n[Resolution] Per-class summary (native image size before resizing):\")\nprint(resolution_summary.to_string(\n    index=False,\n    formatters={\n        \"median_width\": \"{:.0f}\".format,\n        \"min_width\": \"{:.0f}\".format,\n        \"max_width\": \"{:.0f}\".format,\n        \"median_height\": \"{:.0f}\".format,\n        \"median_megapixels\": \"{:.2f}\".format,\n        \"median_aspect_ratio\": \"{:.3f}\".format,\n    },\n    na_rep=\"-\",\n))\n\nplot_resolution_by_class(resolution_by_class, resolution_plot_path)\n\n# Association check: does megapixel distribution differ between stages?\nmegapixel_groups = [\n    resolution_by_class.loc[resolution_by_class[\"label\"] == c, \"megapixels\"].to_numpy()\n    for c in range(Config.NUM_CLASSES)\n]\nmegapixel_groups = [group for group in megapixel_groups if len(group)]\ntry:\n    kruskal_stat, kruskal_p = kruskal(*megapixel_groups)\nexcept ValueError:\n    # scipy raises when every value is identical: no difference between stages at all.\n    kruskal_stat, kruskal_p = 0.0, 1.0\nclass_medians = resolution_summary[\"median_megapixels\"].dropna()\nmedian_ratio = float(class_medians.max() / class_medians.min())\n\nprint(f\"\\n[Resolution] Kruskal-Wallis on megapixels across {len(megapixel_groups)} stages: \"\n      f\"H = {kruskal_stat:.2f}, p-value = {kruskal_p:.3g}\")\nprint(f\"[Resolution] Largest / smallest per-class median megapixels: {median_ratio:.2f}x\")\nif kruskal_p < 0.05 and median_ratio > 1.5:\n    print(\n        \"[Resolution] Resolution differs noticeably between stages — possible source/camera bias; \"\n        \"the model could partly learn resolution instead of lesions. Resizing to \"\n        f\"Config.IMG_SIZE ({Config.IMG_SIZE}x{Config.IMG_SIZE}) reduces but may not remove this; \"\n        \"discuss as a limitation.\"\n    )\nelse:\n    print(\n        \"[Resolution] No strong resolution difference between stages; \"\n        \"resolution is unlikely to be a shortcut feature.\"\n    )\nprint(f\"[Resolution] Saved -> {resolution_csv_path}\")\nprint(f\"[Resolution] Saved -> {resolution_summary_path}\")\nprint(f\"[Resolution] Saved -> {resolution_plot_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:53:43.147715Z","iopub.execute_input":"2026-09-30T15:53:43.148158Z","iopub.status.idle":"2026-09-30T15:53:45.052269Z","shell.execute_reply.started":"2026-09-30T15:53:43.148132Z","shell.execute_reply":"2026-09-30T15:53:45.051554Z"}},"outputs":[],"execution_count":null},{"id":"b2c13e0e","cell_type":"markdown","source":"## Section 2: Preprocessing\n\n**Goal:** Transform raw fundus photographs into clean, normalised tensors at `Config.IMG_SIZE` (300×300 by default).\n\n**Default training pipeline** (applied to every image, in order):\n1. **Black-border crop** — fundus images are circular; the black padding carries no information and wastes model capacity.\n2. **Resize** to `Config.IMG_SIZE` square pixels (area interpolation when shrinking).\n3. **Ben Graham enhancement** — subtracts a heavily Gaussian-blurred copy of the image to remove uneven illumination and boost local lesion/vessel contrast.\n4. **Normalise** pixel values to [0, 1] (the model's input adapter rescales to what EfficientNet expects).\n\n**Optional filters** available in `core.preprocessing` and shown in the Section 2.5b gallery: bilateral **denoising**, **CLAHE** (contrast-limited adaptive histogram equalisation on the LAB L-channel) and **unsharp-mask edge enhancement**. They are off by default; Section 6 trains pilot models with them switched on (and with Ben Graham switched off) to measure whether each one actually helps.\n\nBefore/after figures, histograms and edge/morphology maps are saved to `Config.REPORTS_DIR`.","metadata":{}},{"id":"402b607557e547b8a1f180d6b9b117c5","cell_type":"code","source":"# ============================================================\n# SECTION 2.1 — IMAGE LOADING UTILITY\n# A single function that loads an image file into an RGB NumPy\n# array, used as the entry point for every preprocessing step.\n# ============================================================\n\ndef load_image_rgb(filepath: str) -> np.ndarray:\n    \"\"\"\n    Load an image from disk and convert to an RGB NumPy array.\n\n    Uses OpenCV (fast, handles many formats) and converts BGR -> RGB\n    so colours are correct for matplotlib and TensorFlow.\n\n    Args:\n        filepath: Absolute path to the image file.\n\n    Returns:\n        NumPy array of shape (H, W, 3), dtype uint8, RGB channel order.\n\n    Raises:\n        FileNotFoundError: If the file does not exist.\n        ValueError:        If OpenCV fails to decode the image.\n    \"\"\"\n    if not os.path.exists(filepath):\n        raise FileNotFoundError(f\"Image not found: {filepath}\")\n\n    img_bgr = cv2.imread(filepath)\n    if img_bgr is None:\n        raise ValueError(\n            f\"OpenCV could not decode image: {filepath}. \"\n            \"The file may be corrupt or in an unsupported format.\"\n        )\n    # OpenCV loads as BGR; convert to RGB for compatibility with PIL and TF\n    return cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)","metadata":{"id":"402b607557e547b8a1f180d6b9b117c5","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:54:03.10211Z","iopub.execute_input":"2026-09-30T15:54:03.102525Z","iopub.status.idle":"2026-09-30T15:54:03.108483Z","shell.execute_reply.started":"2026-09-30T15:54:03.102497Z","shell.execute_reply":"2026-09-30T15:54:03.107661Z"}},"outputs":[],"execution_count":null},{"id":"e54a262c64bf461eb0dbe1adddc8aba3","cell_type":"code","source":"# ============================================================\n# SECTION 2.2 — BLACK-BORDER CROPPING\n#\n# WHY: Fundus cameras capture a circular retinal image centred\n# on a rectangular sensor. The corners of the photograph are\n# pure black (pixel value ≈ 0). This padding:\n#   - Wastes model capacity on uninformative pixels\n#   - Can confuse augmentation (random crops may include mostly black)\n#   - Inflates the apparent image size\n#\n# APPROACH (Ben Graham / crop_image_from_gray style):\n#   1. Convert to grayscale.\n#   2. Apply a binary threshold at a low value (e.g. 7) to create a\n#      mask that is 1 where the retina is and 0 where it is black.\n#   3. Find the bounding box of the non-zero mask region.\n#   4. Crop the original colour image to that bounding box.\n#   5. If the crop is degenerate (all black image), return the original.\n# ============================================================\n\ndef crop_image_from_gray(\n    img: np.ndarray,\n    threshold: int = 7,\n    tol: int = 7,\n) -> np.ndarray:\n    \"\"\"\n    Remove black borders from a fundus photograph.\n\n    The retinal fundus image is circular; surrounding pixels are near-black\n    (value < threshold). We threshold the grayscale image to find the retinal\n    disc boundary and crop the colour image to that bounding box.\n\n    This is commonly called the \"crop_image_from_gray\" technique, popularised\n    by Ben Graham's APTOS competition solution.\n\n    Args:\n        img:       RGB NumPy array (H, W, 3), uint8.\n        threshold: Grayscale pixel value below which a pixel is considered\n                   background (black border). Default 7 is empirically robust\n                   across APTOS 2019 fundus photographs.\n        tol:       Additional tolerance in pixels to expand the crop slightly\n                   and avoid clipping the very edge of the optic disc.\n\n    Returns:\n        Cropped RGB NumPy array. Falls back to the original if the crop is\n        degenerate (e.g. an entirely black image).\n    \"\"\"\n    if img.ndim == 2:\n        # Already grayscale — work directly\n        mask = img > threshold\n        if not mask.any():\n            return img  # guard: entirely black image\n        row_mask = mask.any(axis=1)\n        col_mask = mask.any(axis=0)\n        rmin, rmax = np.where(row_mask)[0][[0, -1]]\n        cmin, cmax = np.where(col_mask)[0][[0, -1]]\n        # Add tolerance so we do not clip the disc edge\n        rmin = max(0, rmin - tol)\n        rmax = min(img.shape[0] - 1, rmax + tol)\n        cmin = max(0, cmin - tol)\n        cmax = min(img.shape[1] - 1, cmax + tol)\n        return img[rmin:rmax + 1, cmin:cmax + 1]\n\n    # For an RGB image, compute the grayscale mask from a single channel\n    # (green channel has best contrast for retinal vessels)\n    gray = img[:, :, 1]\n    mask = gray > threshold\n\n    if not mask.any():\n        return img  # guard: entirely black or near-black image\n\n    row_mask = mask.any(axis=1)\n    col_mask = mask.any(axis=0)\n    rmin, rmax = np.where(row_mask)[0][[0, -1]]\n    cmin, cmax = np.where(col_mask)[0][[0, -1]]\n\n    rmin = max(0, rmin - tol)\n    rmax = min(img.shape[0] - 1, rmax + tol)\n    cmin = max(0, cmin - tol)\n    cmax = min(img.shape[1] - 1, cmax + tol)\n\n    cropped = img[rmin:rmax + 1, cmin:cmax + 1]\n\n    # Sanity check: if the crop is nearly empty, return the original\n    if cropped.size == 0 or min(cropped.shape[:2]) < 10:\n        return img\n\n    return cropped","metadata":{"id":"e54a262c64bf461eb0dbe1adddc8aba3","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:54:06.842024Z","iopub.execute_input":"2026-09-30T15:54:06.842514Z","iopub.status.idle":"2026-09-30T15:54:06.851897Z","shell.execute_reply.started":"2026-09-30T15:54:06.842487Z","shell.execute_reply":"2026-09-30T15:54:06.850999Z"}},"outputs":[],"execution_count":null},{"id":"3c67306ceb384e51b3da39360104e78c","cell_type":"code","source":"# ============================================================\n# SECTION 2.3 — BEN GRAHAM CONTRAST ENHANCEMENT\n#\n# WHY: Fundus photographs suffer from uneven illumination — the\n# centre of the image is often brighter than the periphery due to\n# the ophthalmoscope light path. This hides subtle lesion contrast\n# (microaneurysms, haemorrhages, exudates) in darker regions.\n#\n# THE TECHNIQUE: Subtract a heavily blurred (low-frequency) version\n# of the image from the original, then re-centre the distribution:\n#\n#   enhanced = alpha * original + beta * GaussianBlur(original, sigma) + gamma\n#\n# With alpha=4, beta=-4, gamma=128:\n#   - The blurred image captures the slow illumination gradient\n#   - Subtracting it removes that gradient (local contrast normalisation)\n#   - gamma=128 re-centres pixel values near the middle of [0,255]\n#     so we do not clip large negative or positive residuals\n#   - alpha=4 amplifies fine-detail (vessel edges, lesion boundaries)\n#\n# This is NOT the same as CLAHE (which operates on histograms of tiles);\n# Ben Graham enhancement operates globally in the spatial domain and\n# specifically targets the illumination gradient rather than histogram spread.\n# It was used by the winning solution of the 2015 Kaggle DR competition.\n# ============================================================\n\ndef ben_graham_enhance(\n    img: np.ndarray,\n    sigma: int = Config.BEN_GRAHAM_SIGMA,\n    alpha: float = Config.BEN_GRAHAM_ALPHA,\n    beta: float  = Config.BEN_GRAHAM_BETA,\n    gamma: float = Config.BEN_GRAHAM_GAMMA,\n) -> np.ndarray:\n    \"\"\"\n    Apply Ben Graham-style contrast enhancement to a fundus image.\n\n    Subtracts a Gaussian-blurred version of the image from itself to\n    remove the slow illumination gradient present in fundus photography,\n    then re-centres pixel values with an additive bias.\n\n    Formula: enhanced = clip(alpha * img + beta * blur(img, sigma) + gamma)\n\n    Args:\n        img:   RGB uint8 NumPy array (H, W, 3).\n        sigma: Gaussian blur kernel sigma. Large sigma (default 10) captures\n               the broad illumination gradient across the whole image.\n        alpha: Weight on the original image. Value of 4 amplifies fine detail.\n        beta:  Weight on the blurred image (typically -alpha to subtract it).\n        gamma: Additive constant to re-centre the output around 128.\n\n    Returns:\n        Enhanced RGB uint8 NumPy array, same shape as input.\n    \"\"\"\n    # Build the blur kernel size from sigma: must be odd and positive\n    # ksize = 2 * (4 * sigma) + 1 is the standard \"full-width\" Gaussian kernel\n    ksize = int(2 * round(4 * sigma) + 1)\n\n    # Gaussian blur captures the low-frequency illumination pattern\n    blurred = cv2.GaussianBlur(img, (ksize, ksize), sigma)\n\n    # Weighted addition: original (amplified) minus blur (background), re-centred\n    enhanced = cv2.addWeighted(img, alpha, blurred, beta, gamma)\n\n    # clip is implicit in addWeighted for uint8, but we apply it explicitly\n    # to avoid artefacts if the input dtype differs\n    enhanced = np.clip(enhanced, 0, 255).astype(np.uint8)\n    return enhanced","metadata":{"id":"3c67306ceb384e51b3da39360104e78c","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:54:11.304051Z","iopub.execute_input":"2026-09-30T15:54:11.304461Z","iopub.status.idle":"2026-09-30T15:54:11.312144Z","shell.execute_reply.started":"2026-09-30T15:54:11.304435Z","shell.execute_reply":"2026-09-30T15:54:11.311157Z"}},"outputs":[],"execution_count":null},{"id":"96e31bed8d2a406b951dc96687876887","cell_type":"code","source":"# ============================================================\n# SECTION 2.4 — FULL PREPROCESSING PIPELINE (Unified with core/)\n# Reusable preprocessing pipeline imported from core.preprocessing\n# ensuring a single source of truth across research notebook,\n# standalone GPU training, and the clinical web application.\n# ============================================================\n\nimport subprocess\nimport sys\nimport inspect  # CHANGED: print the active wrapper source for resolution verification\nfrom pathlib import Path\n\n_PROJECT_ROOT = Path.cwd()\nif not (_PROJECT_ROOT / \"core\").is_dir():\n    if _IN_KAGGLE:\n        _PROJECT_ROOT = Path(\"/kaggle/working/ComputerVision_Coursework\")\n    else:\n        _PROJECT_ROOT = Path(\"/content/ComputerVision_Coursework\")\n    if not (_PROJECT_ROOT / \"core\").is_dir():\n        subprocess.run(\n            [\n                \"git\",\n                \"clone\",\n                \"--depth\",\n                \"1\",\n                \"--branch\",\n                \"main\",\n                \"https://github.com/ShaznaSalman/ComputerVision_Coursework_COBSCCOMP24.2P-019.git\",\n                str(_PROJECT_ROOT),\n            ],\n            check=True,\n        )\n\nif str(_PROJECT_ROOT) not in sys.path:\n    sys.path.insert(0, str(_PROJECT_ROOT))\n\nfrom core.preprocessing import (\n    crop_image_from_gray,\n    ben_graham_enhance,\n    preprocess_image as _core_preprocess_image,\n)\n\n\ndef preprocess_image(image_input, img_size=None, **kwargs):\n    target_size = Config.IMG_SIZE if img_size is None else img_size\n    return _core_preprocess_image(image_input, img_size=target_size, **kwargs)\n\n\nprint(\"[Preprocessing] Active notebook wrapper source:\")  # CHANGED: expose resolution forwarding for verification\nprint(inspect.getsource(preprocess_image))  # CHANGED: verify Config.IMG_SIZE is passed to core preprocessing\n\n_test_fp = df_labels[\"filepath\"].iloc[0]\n_test_out = preprocess_image(_test_fp)\nprint(\"Preprocessing check:\")\nprint(f\"  Project root: {_PROJECT_ROOT}\")\nprint(f\"  Input file  : {_test_fp}\")\nprint(f\"  Output shape: {_test_out.shape}\")\nprint(f\"  Dtype       : {_test_out.dtype}\")\nprint(f\"  Value range : [{_test_out.min():.3f}, {_test_out.max():.3f}]\")","metadata":{"id":"96e31bed8d2a406b951dc96687876887","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:54:16.945875Z","iopub.execute_input":"2026-09-30T15:54:16.946133Z","iopub.status.idle":"2026-09-30T15:54:18.877156Z","shell.execute_reply.started":"2026-09-30T15:54:16.946113Z","shell.execute_reply":"2026-09-30T15:54:18.876449Z"}},"outputs":[],"execution_count":null},{"id":"d7566a93c895487f896f1703e44b149c","cell_type":"code","source":"# ============================================================\n# SECTION 2.5 — BEFORE / AFTER COMPARISON VISUALISATION\n# Saves side-by-side comparison grids for the report.\n# We show: (a) original raw image, (b) after crop+resize,\n# (c) after Ben Graham enhancement. One grid per DR class.\n# ============================================================\n\ndef visualise_preprocessing(\n    df: pd.DataFrame,\n    n_per_class: int = 2,\n    save_dir: str = Config.REPORTS_DIR,\n) -> None:\n    \"\"\"\n    Generate and save before/after preprocessing comparison figures.\n\n    For each DR class, randomly samples n_per_class images and renders\n    a 3-column grid showing: raw, cropped+resized, and fully enhanced.\n    Figures are saved to save_dir for inclusion in the report.\n\n    Args:\n        df:           DataFrame with 'filepath' and 'label' columns.\n        n_per_class:  Number of example images per class to display.\n        save_dir:     Directory where PNG comparison files are saved.\n    \"\"\"\n    os.makedirs(save_dir, exist_ok=True)\n\n    for class_id, class_name in enumerate(Config.CLASS_NAMES):\n        class_df = df[df[\"label\"] == class_id]\n        if class_df.empty:\n            print(f\"[Preprocess] WARNING: No images found for class {class_id} ({class_name}). Skipping.\")\n            continue\n\n        # Sample up to n_per_class images (fewer if the class is small)\n        sample = class_df.sample(min(n_per_class, len(class_df)), random_state=Config.SEED)\n\n        n_cols = 3   # raw | cropped+resized | Ben Graham enhanced\n        n_rows = len(sample)\n        fig, axes = plt.subplots(\n            n_rows, n_cols,\n            figsize=(n_cols * 4, n_rows * 4),\n        )\n        # Ensure axes is always 2D even for n_rows=1\n        if n_rows == 1:\n            axes = axes[np.newaxis, :]\n\n        col_titles = [\"Original (raw)\", \"Cropped + Resized\", \"Ben Graham Enhanced\"]\n        for ax, title in zip(axes[0], col_titles):\n            ax.set_title(title, fontsize=10, fontweight=\"bold\")\n\n        for row_idx, (_, row) in enumerate(sample.iterrows()):\n            fp = row[\"filepath\"]\n\n            # Column 0: original raw image\n            raw = load_image_rgb(fp)\n            axes[row_idx, 0].imshow(raw)\n            axes[row_idx, 0].axis(\"off\")\n\n            # Column 1: cropped and resized (no Ben Graham)\n            cropped = crop_image_from_gray(raw)\n            resized = cv2.resize(\n                cropped,\n                (Config.IMG_SIZE, Config.IMG_SIZE),\n                interpolation=cv2.INTER_AREA if (\n                    cropped.shape[0] > Config.IMG_SIZE\n                ) else cv2.INTER_LINEAR,\n            )\n            axes[row_idx, 1].imshow(resized)\n            axes[row_idx, 1].axis(\"off\")\n\n            # Column 2: fully processed (Ben Graham applied to the resized image)\n            enhanced = ben_graham_enhance(resized)\n            axes[row_idx, 2].imshow(enhanced)\n            axes[row_idx, 2].axis(\"off\")\n\n        fig.suptitle(\n            f\"Preprocessing Steps — Class {class_id}: {class_name}\",\n            fontsize=13,\n            fontweight=\"bold\",\n            y=1.01,\n        )\n        plt.tight_layout()\n\n        save_path = os.path.join(save_dir, f\"preprocessing_class{class_id}.png\")\n        plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n        plt.show()\n        print(f\"[Preprocess] Saved comparison figure -> {save_path}\")\n\n\nvisualise_preprocessing(df_labels, n_per_class=2)\n\n# ---- Also save a combined single-row comparison (all 5 classes side by side) ----\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport cv2 as _cv2\nimport numpy as _np\n\n_fig_comp, _axes_comp = plt.subplots(2, 5, figsize=(20, 8))\n_class_colors = [\"#27ae60\", \"#f39c12\", \"#e67e22\", \"#e74c3c\", \"#8e44ad\"]\nfor _cls in range(5):\n    _cls_df = df_labels[df_labels[\"label\"] == _cls]\n    _sample_row = _cls_df.sample(1, random_state=Config.SEED + _cls).iloc[0]\n    _raw = load_image_rgb(_sample_row[\"filepath\"])\n    _cropped = crop_image_from_gray(_raw)\n    _resized = _cv2.resize(_cropped, (Config.IMG_SIZE, Config.IMG_SIZE))\n    _enhanced = ben_graham_enhance(_resized)\n\n    _axes_comp[0, _cls].imshow(_resized)\n    _axes_comp[0, _cls].axis(\"off\")\n    _axes_comp[0, _cls].set_title(\n        Config.CLASS_NAMES[_cls], color=_class_colors[_cls], fontsize=10, fontweight=\"bold\"\n    )\n    if _cls == 0:\n        _axes_comp[0, _cls].set_ylabel(\"Before Enhancement\", fontsize=9)\n\n    _axes_comp[1, _cls].imshow(_enhanced)\n    _axes_comp[1, _cls].axis(\"off\")\n    if _cls == 0:\n        _axes_comp[1, _cls].set_ylabel(\"After Ben Graham\", fontsize=9)\n\n_fig_comp.suptitle(\n    \"Ben Graham Preprocessing: Before vs After — All 5 DR Stages\",\n    fontsize=13, fontweight=\"bold\"\n)\nplt.tight_layout()\n_comp_path = os.path.join(Config.REPORTS_DIR, \"preprocessing_comparison.png\")\nplt.savefig(_comp_path, dpi=150, bbox_inches=\"tight\")\nplt.close(_fig_comp)\nprint(f\"[Preprocess] Saved combined comparison -> {_comp_path}\")\n","metadata":{"id":"d7566a93c895487f896f1703e44b149c","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:54:26.26253Z","iopub.execute_input":"2026-09-30T15:54:26.262957Z","iopub.status.idle":"2026-09-30T15:54:45.407892Z","shell.execute_reply.started":"2026-09-30T15:54:26.262927Z","shell.execute_reply":"2026-09-30T15:54:45.407174Z"}},"outputs":[],"execution_count":null},{"id":"2441bdaa","cell_type":"markdown","source":"### 2.5b Step-by-step preprocessing, histograms, edges and morphology\n\nThe next cell visualises every preprocessing operation on one image with the intensity histogram of the retinal disc under each step, then applies classical convolution (Sobel, Canny) and morphological operators (opening, top-hat, black-hat) to the green channel. These link the pipeline to histogram operations, convolution/edge detection and morphology, and show which retinal structures each filter emphasises.","metadata":{}},{"id":"d22a9076","cell_type":"code","source":"# ============================================================\n# SECTION 2.5b — STEP-BY-STEP PREPROCESSING GALLERY, HISTOGRAMS,\n#                EDGE DETECTION AND MORPHOLOGY\n#\n# Figure 1 follows one fundus image through every available filter:\n#   raw -> crop + resize -> bilateral denoise -> CLAHE -> unsharp-mask\n#   edge enhancement, alongside Ben Graham (the default training step).\n#   The row underneath shows the grey-level and green-channel histogram\n#   of the retinal disc for each step, so the effect of each operation on\n#   the intensity distribution is visible (histogram operations).\n#\n# Figure 2 applies classical convolution and morphology operators to the\n# green channel, which carries the most vessel/lesion contrast:\n#   Sobel gradient magnitude, Canny edges, morphological opening of the\n#   retina mask, black-hat (dark structures: vessels, haemorrhages,\n#   microaneurysms) and top-hat (bright structures: exudates).\n#\n# Denoise, CLAHE and edge enhancement are OFF in the default training\n# pipeline; Section 6 trains pilot models with them ON to test whether\n# they help (preprocessing ablation). All visuals are descriptive only;\n# nothing here is fitted to data, so there is no leakage.\n# ============================================================\n\nfrom collections import OrderedDict\nfrom core.preprocessing import apply_clahe, denoise_fundus, enhance_edges\n\n\ndef _resize_to_config(image: np.ndarray) -> np.ndarray:\n    interpolation = cv2.INTER_AREA if image.shape[0] > Config.IMG_SIZE else cv2.INTER_LINEAR\n    return cv2.resize(image, (Config.IMG_SIZE, Config.IMG_SIZE), interpolation=interpolation)\n\n\ndef build_preprocessing_stages(filepath: str) -> \"OrderedDict[str, np.ndarray]\":\n    \"\"\"Return each preprocessing stage of one image as uint8 RGB arrays.\"\"\"\n    raw = load_image_rgb(filepath)\n    resized = _resize_to_config(crop_image_from_gray(raw))\n    denoised = denoise_fundus(resized)\n    clahe_image = apply_clahe(denoised)\n    sharpened = enhance_edges(clahe_image)\n    return OrderedDict([\n        (\"1. Raw photograph\", raw),\n        (\"2. Crop borders + resize\", resized),\n        (\"3. Bilateral denoise\", denoised),\n        (\"4. CLAHE (LAB L-channel)\", clahe_image),\n        (\"5. Unsharp-mask edges\", sharpened),\n        (\"Ben Graham (training default)\", ben_graham_enhance(resized)),\n    ])\n\n\ndef retina_mask(image_rgb: np.ndarray, threshold: int = 10) -> np.ndarray:\n    \"\"\"Binary mask of the circular retinal disc, cleaned by morphological opening.\"\"\"\n    _, mask = cv2.threshold(image_rgb[:, :, 1], threshold, 255, cv2.THRESH_BINARY)\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (7, 7))\n    return cv2.morphologyEx(mask, cv2.MORPH_OPEN, kernel)\n\n\ndef plot_preprocessing_gallery(filepath: str, label: int, save_path: str) -> None:\n    stages = build_preprocessing_stages(filepath)\n    disc_mask = retina_mask(stages[\"2. Crop borders + resize\"]) > 0\n    raw_mask = retina_mask(stages[\"1. Raw photograph\"]) > 0\n\n    fig, axes = plt.subplots(2, len(stages), figsize=(4 * len(stages), 8),\n                             gridspec_kw={\"height_ratios\": [3, 2]})\n    for column, (title, image) in enumerate(stages.items()):\n        axes[0, column].imshow(image)\n        axes[0, column].set_title(title, fontsize=10, fontweight=\"bold\")\n        axes[0, column].axis(\"off\")\n\n        mask = raw_mask if column == 0 else disc_mask\n        grey = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)[mask]\n        green = image[:, :, 1][mask]\n        axes[1, column].hist(grey, bins=64, range=(0, 255), color=\"#555555\", alpha=0.75, label=\"Grey\")\n        axes[1, column].hist(green, bins=64, range=(0, 255), color=\"#2E8B57\", alpha=0.5, label=\"Green\")\n        axes[1, column].set_xlim(0, 255)\n        axes[1, column].set_yticks([])\n        axes[1, column].set_xlabel(f\"Intensity (std {grey.std():.1f})\", fontsize=9)\n        if column == 0:\n            axes[1, column].legend(fontsize=8)\n\n    fig.suptitle(\n        f\"Preprocessing steps and retinal-disc histograms — dataset grade {label} \"\n        f\"({Config.CLASS_NAMES[label]})\",\n        fontsize=13, fontweight=\"bold\",\n    )\n    fig.tight_layout()\n    fig.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"[Preprocess] Step gallery saved -> {save_path}\")\n\n\ndef plot_edge_and_morphology(filepaths_by_label: Dict[int, str], save_path: str) -> None:\n    columns = [\"Green channel (CLAHE)\", \"Sobel magnitude\", \"Canny edges\",\n               \"Retina mask (opened)\", \"Black-hat: dark lesions/vessels\", \"Top-hat: bright lesions\"]\n    fig, axes = plt.subplots(len(filepaths_by_label), len(columns),\n                             figsize=(3.6 * len(columns), 3.8 * len(filepaths_by_label)), squeeze=False)\n    ellipse_15 = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))\n\n    for row, (label, filepath) in enumerate(filepaths_by_label.items()):\n        resized = _resize_to_config(crop_image_from_gray(load_image_rgb(filepath)))\n        mask = retina_mask(resized)\n        inner_mask = cv2.erode(mask, ellipse_15)   # drop the strong disc-boundary edge\n        green = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8)).apply(resized[:, :, 1])\n        smoothed = cv2.GaussianBlur(green, (5, 5), 0)\n\n        sobel_x = cv2.Sobel(smoothed, cv2.CV_32F, 1, 0, ksize=3)\n        sobel_y = cv2.Sobel(smoothed, cv2.CV_32F, 0, 1, ksize=3)\n        sobel = cv2.magnitude(sobel_x, sobel_y)\n        sobel = cv2.normalize(sobel, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n\n        def _inside_disc(image: np.ndarray) -> np.ndarray:\n            return cv2.bitwise_and(image, image, mask=inner_mask)\n\n        sobel = _inside_disc(sobel)\n        canny = _inside_disc(cv2.Canny(smoothed, 30, 90))\n        black_hat = _inside_disc(cv2.morphologyEx(green, cv2.MORPH_BLACKHAT, ellipse_15))\n        top_hat = _inside_disc(cv2.morphologyEx(green, cv2.MORPH_TOPHAT, ellipse_15))\n\n        panels = [green, sobel, canny, mask, black_hat, top_hat]\n        cmaps = [\"gray\", \"magma\", \"gray\", \"gray\", \"inferno\", \"inferno\"]\n        for column, (panel, cmap) in enumerate(zip(panels, cmaps)):\n            axes[row, column].imshow(panel, cmap=cmap)\n            axes[row, column].axis(\"off\")\n            if row == 0:\n                axes[row, column].set_title(columns[column], fontsize=10, fontweight=\"bold\")\n        axes[row, 0].text(-12, Config.IMG_SIZE / 2, f\"Grade {label}\\n{Config.CLASS_NAMES[label]}\",\n                          ha=\"right\", va=\"center\", fontsize=10, fontweight=\"bold\")\n\n    fig.suptitle(\"Classical edge detection and morphology on the green channel\",\n                 fontsize=13, fontweight=\"bold\")\n    fig.tight_layout()\n    fig.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"[Preprocess] Edge/morphology figure saved -> {save_path}\")\n\n\n_gallery_row = df_labels[df_labels[\"label\"] == 2].sample(1, random_state=Config.SEED).iloc[0]\nplot_preprocessing_gallery(\n    _gallery_row[\"filepath\"],\n    int(_gallery_row[\"label\"]),\n    save_path=os.path.join(Config.REPORTS_DIR, \"preprocessing_step_gallery.png\"),\n)\n\n_edge_examples = {\n    label: df_labels[df_labels[\"label\"] == label].sample(1, random_state=Config.SEED + label).iloc[0][\"filepath\"]\n    for label in (0, 2, 4)\n    if (df_labels[\"label\"] == label).any()\n}\nplot_edge_and_morphology(\n    _edge_examples,\n    save_path=os.path.join(Config.REPORTS_DIR, \"edge_morphology_analysis.png\"),\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:54:55.940003Z","iopub.execute_input":"2026-09-30T15:54:55.940783Z","iopub.status.idle":"2026-09-30T15:55:04.429818Z","shell.execute_reply.started":"2026-09-30T15:54:55.940706Z","shell.execute_reply":"2026-09-30T15:55:04.428782Z"}},"outputs":[],"execution_count":null},{"id":"e0a6681c4998461c86121986088607b0","cell_type":"code","source":"# ============================================================\n# SECTION 2.6 — TF.DATA PIPELINE\n#\n# WHY tf.data: Loading and preprocessing images on-the-fly in\n# Python (with a plain Python generator) creates a CPU bottleneck\n# that starves the GPU. tf.data pipelines:\n#   - Use parallel map workers (num_parallel_calls) so multiple\n#     images are preprocessed concurrently on the CPU\n#   - Use prefetch so the next batch is ready before the GPU\n#     finishes the current one — zero GPU idle time\n#   - Cache decoded images to RAM on second-epoch replay if the\n#     dataset fits in memory (Section 4 caches preprocessed PNGs to disk instead)\n#\n# The pipeline wraps preprocess_image in a tf.py_function so our\n# existing NumPy/OpenCV preprocessing functions work seamlessly\n# inside the TF graph.\n# ============================================================\n\ndef _preprocess_tf(filepath: tf.Tensor, label: tf.Tensor) -> tuple:\n    \"\"\"\n    TensorFlow-compatible wrapper around preprocess_image.\n\n    Called by tf.data.Dataset.map(). Uses tf.py_function to run\n    our NumPy/OpenCV preprocessing inside the TF graph.\n\n    Args:\n        filepath: Scalar string tensor containing the image file path.\n        label:    Scalar int32 tensor containing the ICDR label.\n\n    Returns:\n        Tuple of:\n            img   — float32 tensor (IMG_SIZE, IMG_SIZE, 3)\n            label — int32 tensor (scalar)\n    \"\"\"\n    def _py_preprocess(fp_bytes):\n        # Decode the byte string tensor to a Python str\n        fp = fp_bytes.numpy().decode(\"utf-8\")\n        img = preprocess_image(fp)          # returns float32 (H, W, 3)\n        return img\n\n    img = tf.py_function(\n        func=_py_preprocess,\n        inp=[filepath],\n        Tout=tf.float32,\n    )\n    # Set the static shape so Keras knows the tensor dimensions\n    # (py_function outputs have unknown shape by default)\n    img.set_shape([Config.IMG_SIZE, Config.IMG_SIZE, 3])\n\n    # One-hot encode the label for categorical cross-entropy\n    label_oh = tf.one_hot(tf.cast(label, tf.int32), depth=Config.NUM_CLASSES)\n    return img, label_oh\n\n\ndef build_tf_dataset(\n    filepaths: List[str],\n    labels: List[int],\n    batch_size: int = Config.BATCH_SIZE,\n    shuffle: bool = True,\n    augment_fn=None,\n) -> tf.data.Dataset:\n    \"\"\"\n    Build a tf.data.Dataset from file paths and integer labels.\n\n    Args:\n        filepaths:  List of absolute image file paths.\n        labels:     Corresponding list of integer ICDR labels (0-4).\n        batch_size: Number of images per batch.\n        shuffle:    Whether to shuffle the dataset each epoch.\n                    Should be True for training, False for val/test.\n        augment_fn: Optional callable (image, label) -> (image, label)\n                    applied after preprocessing for training augmentation.\n                    Set in Section 3 once augmentation is defined.\n\n    Returns:\n        Batched, prefetched tf.data.Dataset ready for model.fit().\n    \"\"\"\n    # Build a dataset of (filepath, label) pairs\n    dataset = tf.data.Dataset.from_tensor_slices((filepaths, labels))\n\n    if shuffle:\n        # Buffer size = full dataset ensures true random shuffling\n        dataset = dataset.shuffle(\n            buffer_size=len(filepaths),\n            seed=Config.SEED,\n            reshuffle_each_iteration=True,  # re-shuffle every epoch\n        )\n\n    # Preprocess each image in parallel; AUTOTUNE lets TF choose the\n    # number of parallel workers based on available CPU cores\n    dataset = dataset.map(\n        _preprocess_tf,\n        num_parallel_calls=tf.data.AUTOTUNE,\n    )\n\n    # Apply augmentation if provided (training only)\n    if augment_fn is not None:\n        dataset = dataset.map(augment_fn, num_parallel_calls=tf.data.AUTOTUNE)\n\n    # Batch and prefetch — prefetch(AUTOTUNE) overlaps batch preparation\n    # with GPU computation, eliminating the data-loading bottleneck\n    dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n\n    return dataset\n\n\nprint(\"[Section 2] tf.data pipeline functions defined.\")\nprint(\"  build_tf_dataset() will be called in Section 4 after the train/val/test split.\")\nprint(\"  Augmentation will be wired in via the augment_fn argument in Section 3.\")","metadata":{"id":"e0a6681c4998461c86121986088607b0","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:13.805333Z","iopub.execute_input":"2026-09-30T15:55:13.805756Z","iopub.status.idle":"2026-09-30T15:55:13.815366Z","shell.execute_reply.started":"2026-09-30T15:55:13.805716Z","shell.execute_reply":"2026-09-30T15:55:13.814684Z"}},"outputs":[],"execution_count":null},{"id":"7e73e83c2a994589bfae3d97d25b8685","cell_type":"code","source":"# ============================================================\n# SECTION 2.7 — SECTION SUMMARY\n# Print preprocessing pipeline confirmation so we can verify\n# before moving on to augmentation.\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  SECTION 2 COMPLETE — Preprocessing\")\nprint(\"=\" * 60)\nprint(f\"  crop_image_from_gray()  : black border removal\")\nprint(f\"  ben_graham_enhance()    : local contrast normalisation\")\nprint(f\"  preprocess_image()      : full pipeline (load -> crop -> resize -> enhance -> normalise)\")\nprint(f\"  build_tf_dataset()      : tf.data pipeline with parallel map + prefetch\")\nprint(f\"\")\nprint(f\"  Target resolution : {Config.IMG_SIZE} x {Config.IMG_SIZE}\")\nprint(f\"  Output dtype      : float32, values in [0.0, 1.0]\")\nprint(f\"  Comparison images : saved to {Config.REPORTS_DIR}/preprocessing_class*.png\")\nprint(\"=\" * 60)\nprint(\"  Next -> Section 3: Data Augmentation & Class Balancing\")","metadata":{"id":"7e73e83c2a994589bfae3d97d25b8685","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:18.809566Z","iopub.execute_input":"2026-09-30T15:55:18.81011Z","iopub.status.idle":"2026-09-30T15:55:18.816245Z","shell.execute_reply.started":"2026-09-30T15:55:18.81008Z","shell.execute_reply":"2026-09-30T15:55:18.815305Z"}},"outputs":[],"execution_count":null},{"id":"6f53e300901f4f1e9bfea7c0c76ca3d0","cell_type":"code","source":"# ============================================================\n# SECTION 2.8 — TRAINING-SPLIT PREPROCESSING ABLATION SETUP\n# Defines deterministic transforms and image-statistic proxies.\n# The ablation runs after the saved split is loaded and samples train_df only.\n# ============================================================\n\nfrom core.preprocessing import ben_graham_enhance, crop_image_from_gray\n\n\ndef _load_rgb_for_ablation(filepath: str) -> np.ndarray:\n    image_bgr = cv2.imread(filepath)\n    if image_bgr is None:\n        raise ValueError(f\"Could not decode image: {filepath}\")\n    return cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)\n\n\ndef _bilateral_denoise_for_ablation(image: np.ndarray) -> np.ndarray:\n    return cv2.bilateralFilter(image, d=5, sigmaColor=20, sigmaSpace=20)\n\n\ndef _prepare_preprocessing_variant(filepath: str, variant: str) -> np.ndarray:\n    raw = _load_rgb_for_ablation(filepath)\n    resized_raw = cv2.resize(raw, (Config.IMG_SIZE, Config.IMG_SIZE), interpolation=cv2.INTER_AREA)\n    if variant == \"raw_resize\":\n        return resized_raw\n\n    cropped = crop_image_from_gray(raw)\n    resized = cv2.resize(\n        cropped,\n        (Config.IMG_SIZE, Config.IMG_SIZE),\n        interpolation=cv2.INTER_AREA if cropped.shape[0] > Config.IMG_SIZE else cv2.INTER_LINEAR,\n    )\n    if variant == \"crop_resize\":\n        return resized\n    if variant == \"crop_resize_ben_graham\":\n        return ben_graham_enhance(resized)\n    if variant == \"crop_resize_ben_graham_denoise\":\n        return ben_graham_enhance(_bilateral_denoise_for_ablation(resized))\n    raise ValueError(f\"Unknown preprocessing variant: {variant}\")\n\n\ndef _preprocessing_quality_metrics(image: np.ndarray) -> Dict[str, float]:\n    gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY).astype(np.float32) / 255.0\n    local_brightness = cv2.GaussianBlur(gray, (0, 0), sigmaX=15)\n    high_frequency_residual = gray - cv2.GaussianBlur(gray, (0, 0), sigmaX=1)\n    edges = cv2.Canny((gray * 255).astype(np.uint8), 50, 150)\n    laplacian = cv2.Laplacian(gray * 255.0, cv2.CV_32F)\n    return {\n        \"contrast_std\": float(gray.std()),\n        \"edge_density\": float((edges > 0).mean()),\n        \"brightness_variation\": float(local_brightness.std()),\n        \"signal_noise_proxy\": float(gray.std() / (high_frequency_residual.std() + 1e-8)),\n        \"sharpness_laplacian_variance\": float(laplacian.var()),\n    }\n\n\npreprocessing_variants = [\n    \"raw_resize\",\n    \"crop_resize\",\n    \"crop_resize_ben_graham\",\n    \"crop_resize_ben_graham_denoise\",\n]\nprint(\"[Preprocessing Ablation] Helpers ready; execution is deferred until saved train_df is loaded.\")","metadata":{"id":"6f53e300901f4f1e9bfea7c0c76ca3d0","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:24.806699Z","iopub.execute_input":"2026-09-30T15:55:24.807227Z","iopub.status.idle":"2026-09-30T15:55:24.817795Z","shell.execute_reply.started":"2026-09-30T15:55:24.8072Z","shell.execute_reply":"2026-09-30T15:55:24.816954Z"}},"outputs":[],"execution_count":null},{"id":"bd5fcfa5","cell_type":"markdown","source":"## Section 3: Data Augmentation & Class Balancing\n\n**Goal:** Address two separate but related problems:\n\n1. **Generalisation** — augmentation expands the training distribution so the model does not memorise orientation, framing or exposure artefacts.\n2. **Class imbalance** — APTOS 2019 is **9.4× imbalanced** (1,805 No DR vs 193 Severe). Augmentation keeps the observed class ratio, so on its own it does not change the loss's bias towards No DR.\n\n**Augmentation (training images only; validation and test are never augmented):**\n- Horizontal and vertical flips, and **full 360° rotation** — a fundus photograph has no canonical \"up\", so every orientation is plausible.\n- **Zoom ±15%** — simulates differences in camera field of view and framing.\n- **Brightness and contrast ±15%** — simulates exposure differences between cameras, kept small so lesions are not hidden or faked.\n- Strong colour/hue shifts are deliberately **not** used, because colour is diagnostic (e.g. red haemorrhages, yellow exudates).\n\n**Balancing:** class weights are computed from the **training split only** and are **on by default**. `Config.CLASS_WEIGHT_MODE` chooses between `\"balanced\"` (sklearn inverse-frequency, the default), `\"sqrt\"` (square root of the balanced weights, rescaled to mean 1 — a softer correction that over-weights the rare Severe stage less) and `\"none\"`; Section 4.4 prints both weightings and `class_weighting_effect.png` shows the loss share under each. Section 3.3b adds an optional **on-the-fly oversampling** mode (`tf.data.Dataset.sample_from_datasets`, one stream per class) that is tested as a separate pilot trial in Section 6 (`oversampling_instead`), alongside `class_weights_off` and `class_weights_sqrt` trials. Oversampling and class weights are never combined in the same run.","metadata":{}},{"id":"39ede36d9d0540a2bec9e93ed8e3f075","cell_type":"code","source":"# ============================================================\n# SECTION 3.1 — KERAS AUGMENTATION LAYER (graph-mode)\n#\n# WHY Keras augmentation layers instead of imgaug or albumentations:\n#   - They run inside tf.data; TensorFlow uses available devices,\n#     including a GPU when the runtime provides one.\n#   - They are part of the Keras functional API and behave correctly\n#     at inference time (turned off automatically when training=False).\n#   - They are seeded via tf.random, consistent with our global seed.\n#\n# Each transform is applied with a probability; not every image gets\n# every transform on every pass — this stochasticity is intentional.\n# ============================================================\n\ndef build_augmentation_layer() -> keras.Sequential:\n    \"\"\"\n    Build a Keras Sequential augmentation pipeline.\n\n    All layers only apply their transform during training\n    (training=True); during validation/testing they act as identity\n    functions. This is handled automatically by Keras.\n\n    Augmentations applied:\n    - Horizontal and vertical reflection for image-level ICDR grading.\n    - RandomRotation factor=1.0 for full-circle rotation.  # CHANGED: document full rotation range\n    - RandomZoom +/- 15 percent.  # CHANGED: document increased zoom range\n    - RandomBrightness and RandomContrast factor=0.15 over normalized [0, 1].\n\n    Returns:\n        A compiled keras.Sequential model (used as a callable layer).\n    \"\"\"\n    augmentation = keras.Sequential(\n        [\n            layers.RandomFlip(\"horizontal\", seed=3),\n            layers.RandomFlip(\"vertical\", seed=4),\n            layers.RandomRotation(factor=1.0, fill_mode=\"constant\", fill_value=0.5, seed=1),  # CHANGED: full rotation for orientation robustness\n            layers.RandomZoom(height_factor=(-0.15, 0.15), fill_mode=\"constant\", fill_value=0.5, seed=2),  # CHANGED: widen zoom range conservatively\n            layers.RandomBrightness(factor=0.15, value_range=(0.0, 1.0), seed=5),\n            layers.RandomContrast(factor=0.15, value_range=(0.0, 1.0), seed=6),\n        ],\n        name=\"augmentation_pipeline\",\n    )\n    return augmentation\n\n\n# Build the augmentation layer once and reuse it\naugmentation_layer = build_augmentation_layer()\nprint(\"[Augmentation] Layer built:\")\naugmentation_layer.summary()","metadata":{"id":"39ede36d9d0540a2bec9e93ed8e3f075","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:30.909473Z","iopub.execute_input":"2026-09-30T15:55:30.910167Z","iopub.status.idle":"2026-09-30T15:55:31.401438Z","shell.execute_reply.started":"2026-09-30T15:55:30.910127Z","shell.execute_reply":"2026-09-30T15:55:31.400847Z"}},"outputs":[],"execution_count":null},{"id":"a6becdc18cf64ae48a459b74f1138ef6","cell_type":"code","source":"# ============================================================\n# SECTION 3.2 — AUGMENTATION WRAPPER FOR tf.data\n# Wraps the Keras augmentation layer in a function signature\n# compatible with tf.data.Dataset.map().\n# ============================================================\n\ndef augment_fn(image: tf.Tensor, label: tf.Tensor) -> tuple:\n    \"\"\"\n    Apply augmentation to a single (image, label) pair.\n\n    Used as the augment_fn argument to build_tf_dataset() so augmentation\n    runs inside the tf.data pipeline. The training=True flag activates\n    the random transforms; during inference (training=False) no transform\n    is applied — this is handled automatically by Keras random layers.\n\n    Args:\n        image: float32 tensor (IMG_SIZE, IMG_SIZE, 3), values in [0, 1].\n        label: one-hot float32 tensor of shape (NUM_CLASSES,).\n\n    Returns:\n        Tuple (augmented_image, label) — label is unchanged.\n    \"\"\"\n    # The augmentation model is built for batched input; map() supplies one image.\n    image = tf.expand_dims(image, axis=0)\n    image = augmentation_layer(image, training=True)\n    image = tf.squeeze(image, axis=0)\n    return image, label\n\n\nprint(\"[Augmentation] augment_fn() defined and ready for build_tf_dataset().\")","metadata":{"id":"a6becdc18cf64ae48a459b74f1138ef6","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:38.062686Z","iopub.execute_input":"2026-09-30T15:55:38.063145Z","iopub.status.idle":"2026-09-30T15:55:38.070185Z","shell.execute_reply.started":"2026-09-30T15:55:38.063118Z","shell.execute_reply":"2026-09-30T15:55:38.06949Z"}},"outputs":[],"execution_count":null},{"id":"f0b1604a91c54499a26a4b506cd6368b","cell_type":"code","source":"# ============================================================\n# SECTION 3.3 — CLASS WEIGHT COMPUTATION\n#\n# Augmentation increases variety but preserves the observed class ratio.\n# APTOS 2019 is about 9.4x imbalanced (No DR vs Severe), so balanced\n# inverse-frequency weighting is ON by default (Config.USE_CLASS_WEIGHTS).\n# ============================================================\n\ndef compute_class_weights(labels: List[int]) -> Dict[int, float]:\n    \"\"\"Compute per-class loss weights using sklearn's balanced strategy.\"\"\"\n    unique_classes = sorted(set(labels))\n    weights_array = compute_class_weight(\n        class_weight=\"balanced\",\n        classes=np.array(unique_classes),\n        y=np.array(labels),\n    )\n    class_weight_dict = {\n        cls: float(weight) for cls, weight in zip(unique_classes, weights_array)\n    }\n\n    print(\"\\n[Class Weights] Computed using sklearn balanced strategy:\")\n    print(\"-\" * 45)\n    for cls, weight in class_weight_dict.items():\n        bar = \"#\" * int(weight * 3)\n        print(f\"  Class {cls} ({Config.CLASS_NAMES[cls]:20s}): {weight:6.3f}  {bar}\")\n    print(\"-\" * 45)\n    if Config.USE_CLASS_WEIGHTS:  # CHANGED: report whether Keras will apply computed weights\n        print(\"  These weights are passed to model.fit(class_weight=...) in Section 6.\")\n    else:  # CHANGED: avoid claiming disabled weights affect training\n        print(\"  Weights are computed for inspection only; class weighting is disabled for this run.\")\n    print()\n    return class_weight_dict\n\n\ndef compute_sqrt_class_weights(balanced_weights: Dict[int, float]) -> Dict[int, float]:\n    \"\"\"\n    Softer class weights: square root of the balanced weights, rescaled so the mean weight is 1.\n\n    The square root shrinks the ratio between the rarest and most common class\n    (e.g. 4x becomes 2x), so minority stages still count more without the\n    loss being dominated by a handful of Severe images.\n    \"\"\"\n    rooted = {cls: float(np.sqrt(weight)) for cls, weight in balanced_weights.items()}\n    mean_weight = float(np.mean(list(rooted.values())))\n    return {cls: weight / mean_weight for cls, weight in rooted.items()}\n\n\n# Preview only; the final weights are recomputed from training-split labels in Section 4.\n_all_labels = df_labels[\"label\"].tolist()\nclass_weights_preview = compute_class_weights(_all_labels)\nprint(\"[Note] Final class weights will be recomputed from training-split labels in Section 4.\")","metadata":{"id":"f0b1604a91c54499a26a4b506cd6368b","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:42.46907Z","iopub.execute_input":"2026-09-30T15:55:42.469492Z","iopub.status.idle":"2026-09-30T15:55:42.480613Z","shell.execute_reply.started":"2026-09-30T15:55:42.469465Z","shell.execute_reply":"2026-09-30T15:55:42.479773Z"}},"outputs":[],"execution_count":null},{"id":"ab3b7c99","cell_type":"code","source":"# ============================================================\n# SECTION 3.3b — OPTIONAL ON-THE-FLY OVERSAMPLING (tf.data)\n#\n# An alternative to class weights: instead of re-weighting the loss,\n# draw training batches with equal probability from each DR stage.\n# One tf.data stream is built per class and repeated indefinitely;\n# tf.data.Dataset.sample_from_datasets picks a class uniformly at random\n# for every example, so minority stages (Severe, Proliferative) are seen\n# as often as No DR. Images are re-used, never copied to disk, and\n# augmentation is applied after sampling so repeats look different.\n#\n# Oversampling and class weights are never combined in one run: both\n# correct the same imbalance and together would over-correct.\n# Only the training split is ever oversampled.\n# ============================================================\n\ndef build_oversampled_dataset(\n    filepaths: List[str],\n    labels: List[int],\n    load_fn,\n    seed: int = Config.SEED,\n) -> tf.data.Dataset:\n    \"\"\"\n    Return an infinite, class-balanced stream of (image, one-hot label) pairs.\n\n    Args:\n        filepaths: Image paths for the training split (e.g. cached PNGs).\n        labels:    Integer ICDR labels aligned with filepaths.\n        load_fn:   tf.data map function (path, label) -> (image, one_hot_label).\n        seed:      Seed for per-class shuffling and class sampling.\n\n    Returns:\n        Unbatched, infinite tf.data.Dataset; pass steps_per_epoch to model.fit().\n    \"\"\"\n    filepaths = np.asarray(filepaths)\n    labels = np.asarray(labels, dtype=np.int32)\n    per_class = []\n    for class_id in range(Config.NUM_CLASSES):\n        class_mask = labels == class_id\n        if not class_mask.any():\n            continue\n        class_ds = tf.data.Dataset.from_tensor_slices((filepaths[class_mask], labels[class_mask]))\n        class_ds = class_ds.shuffle(int(class_mask.sum()), seed=seed + class_id, reshuffle_each_iteration=True)\n        per_class.append(class_ds.repeat())\n    if not per_class:\n        raise ValueError(\"No training labels available for oversampling.\")\n    weights = [1.0 / len(per_class)] * len(per_class)\n    balanced = tf.data.Dataset.sample_from_datasets(per_class, weights=weights, seed=seed)\n    return balanced.map(load_fn, num_parallel_calls=tf.data.AUTOTUNE)\n\n\ndef check_balancing_choice(use_class_weights: bool, use_oversampling: bool) -> None:\n    \"\"\"Refuse to combine class weights with oversampling in the same training run.\"\"\"\n    if use_class_weights and use_oversampling:\n        raise ValueError(\"Class weights and oversampling must not be combined in one run; choose one.\")\n\n\ncheck_balancing_choice(Config.USE_CLASS_WEIGHTS, Config.USE_OVERSAMPLING)\nif Config.USE_CLASS_WEIGHTS != (Config.CLASS_WEIGHT_MODE != \"none\"):\n    raise ValueError(\"Config.USE_CLASS_WEIGHTS must be True exactly when CLASS_WEIGHT_MODE is not 'none'.\")\nprint(\"[Balancing] build_oversampled_dataset() ready \"\n      f\"(default run: class weights={Config.USE_CLASS_WEIGHTS}, oversampling={Config.USE_OVERSAMPLING}).\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:55:47.249447Z","iopub.execute_input":"2026-09-30T15:55:47.249962Z","iopub.status.idle":"2026-09-30T15:55:47.259077Z","shell.execute_reply.started":"2026-09-30T15:55:47.249933Z","shell.execute_reply":"2026-09-30T15:55:47.258204Z"}},"outputs":[],"execution_count":null},{"id":"1ec8a7b4782042a4a611ce5fb5e445dd","cell_type":"code","source":"# ============================================================\n# SECTION 3.4 — VISUALISE AUGMENTED SAMPLES\n# Show examples of augmented images for the report.\n# Helps verify that the augmentations look realistic (not so\n# extreme that they destroy diagnostic features).\n# ============================================================\n\ndef visualise_augmented_samples(\n    df: pd.DataFrame,\n    n_per_class: int = 1,\n    n_augmentations: int = 4,\n    save_path: str = None,\n    classes: Optional[List[int]] = None,\n    title: str = \"Original vs Augmented Fundus Images (every DR stage)\",\n) -> None:\n    \"\"\"\n    Display and save a grid of original and augmented images.\n\n    Samples n_per_class images from EVERY class (not a random subset of\n    the whole dataset), so each DR stage is guaranteed to appear. Each row\n    shows the preprocessed original and n_augmentations augmented versions.\n\n    Args:\n        df:              DataFrame with 'filepath' and 'label' columns.\n        n_per_class:     Number of source images per class.\n        n_augmentations: Number of augmented versions per source image.\n        save_path:       If provided, save the figure as a PNG.\n        classes:         Class ids to include (default: all classes).\n        title:           Figure title.\n    \"\"\"\n    classes = list(range(Config.NUM_CLASSES)) if classes is None else list(classes)\n    sample = pd.concat(\n        [\n            df[df[\"label\"] == cls].sample(\n                min(n_per_class, int((df[\"label\"] == cls).sum())),\n                random_state=Config.SEED,\n            )\n            for cls in classes\n        ],\n        ignore_index=True,\n    )\n    n_cols = n_augmentations + 1  # original + N augmented\n    n_rows = len(sample)\n    if n_rows == 0:\n        print(\"[Augmentation] No images available for the requested classes.\")\n        return\n\n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols * 3, n_rows * 3))\n    if n_rows == 1:\n        axes = axes[np.newaxis, :]\n\n    # Column header labels\n    col_labels = [\"Original\"] + [f\"Aug {i+1}\" for i in range(n_augmentations)]\n    for ax, lbl in zip(axes[0], col_labels):\n        ax.set_title(lbl, fontsize=9, fontweight=\"bold\")\n\n    for row_idx, (_, row) in enumerate(sample.iterrows()):\n        # Load and preprocess the original (no augmentation)\n        img_arr = preprocess_image(row[\"filepath\"])          # float32, [0,1]\n        img_tensor = tf.constant(img_arr)[tf.newaxis, ...]  # add batch dim\n\n        # Column 0: original preprocessed image\n        axes[row_idx, 0].imshow(img_arr)\n        axes[row_idx, 0].axis(\"off\")\n        # Row label as text: set_ylabel is hidden once the axis is turned off.\n        axes[row_idx, 0].text(\n            -0.08, 0.5, f\"Class {int(row['label'])}\\n{Config.CLASS_NAMES[int(row['label'])]}\",\n            transform=axes[row_idx, 0].transAxes, fontsize=9, fontweight=\"bold\",\n            ha=\"right\", va=\"center\",\n        )\n\n        # Columns 1..n_augmentations: different augmented versions\n        for aug_idx in range(n_augmentations):\n            # Apply the augmentation layer (training=True activates random ops)\n            aug_tensor = augmentation_layer(img_tensor, training=True)\n            aug_arr = aug_tensor[0].numpy()  # remove batch dim\n            # Clip to [0,1] in case brightness/contrast pushed values out of range\n            aug_arr = np.clip(aug_arr, 0.0, 1.0)\n            axes[row_idx, aug_idx + 1].imshow(aug_arr)\n            axes[row_idx, aug_idx + 1].axis(\"off\")\n\n    fig.suptitle(title, fontsize=13, fontweight=\"bold\", y=1.01)\n    plt.tight_layout()\n\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n        print(f\"[Augmentation] Grid saved -> {save_path}\")\n\n    plt.show()\n\n\n# One row per DR stage in the overview figure, so no stage can be missed.\nvisualise_augmented_samples(\n    df_labels,\n    n_per_class=1,\n    n_augmentations=4,\n    save_path=os.path.join(Config.REPORTS_DIR, \"augmentation_examples.png\"),\n)\n\n# Three source images (strips) per stage, saved as one figure per class.\nfor _aug_class in range(Config.NUM_CLASSES):\n    visualise_augmented_samples(\n        df_labels,\n        n_per_class=3,\n        n_augmentations=4,\n        classes=[_aug_class],\n        title=f\"Augmentation Examples - Class {_aug_class}: {Config.CLASS_NAMES[_aug_class]}\",\n        save_path=os.path.join(Config.REPORTS_DIR, f\"augmentation_class{_aug_class}.png\"),\n    )\n","metadata":{"id":"1ec8a7b4782042a4a611ce5fb5e445dd","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:56:00.750905Z","iopub.execute_input":"2026-09-30T15:56:00.751303Z","iopub.status.idle":"2026-09-30T15:56:30.442277Z","shell.execute_reply.started":"2026-09-30T15:56:00.751277Z","shell.execute_reply":"2026-09-30T15:56:30.440472Z"}},"outputs":[],"execution_count":null},{"id":"ef66f7115ba34ccba0a24fa6c7c201a1","cell_type":"code","source":"# ============================================================\n# SECTION 3.5 — OVERFITTING MITIGATION SUMMARY\n# ============================================================\n\nprint(\"[Section 3] Training regularization configuration:\")\nprint(\"  Augmentation  : existing rotation, flips, zoom, brightness, contrast\")\nprint(f\"  Class weights : {Config.USE_CLASS_WEIGHTS} (mode '{Config.CLASS_WEIGHT_MODE}', training split only)\")\nprint(f\"  Oversampling  : {Config.USE_OVERSAMPLING} (never combined with class weights)\")\nprint(f\"  Dropout       : {Config.DROPOUT_RATE}\")\nprint(f\"  Fine-tuning   : top {Config.UNFREEZE_TOP_N} layers at LR={Config.PHASE2_LR}\")\nprint(f\"  Early stopping: patience={Config.ES_PATIENCE} on val_loss (best weights restored)\")\nprint(f\"  LR reduction  : factor={Config.RLROP_FACTOR}, patience={Config.RLROP_PATIENCE} on val_loss\")\nprint(f\"  Phase 2 LR    : schedule '{Config.PHASE2_LR_SCHEDULE}' (plateau or warm-up + cosine)\")\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"  SECTION 3 COMPLETE — Augmentation & Class Balancing\")\nprint(\"=\" * 60)\nprint(\"  augmentation_layer       : Keras Sequential (6 random transforms)\")\nprint(\"  augment_fn()             : tf.data-compatible wrapper\")\nprint(\"  compute_class_weights()  : inverse-frequency loss weighting\")\nprint(\"  augmentation_examples.png saved to REPORTS_DIR\")\nprint(\"  Next -> Section 4: Train / Validation / Test Split\")\nprint(\"=\" * 60)","metadata":{"id":"ef66f7115ba34ccba0a24fa6c7c201a1","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:56:44.069011Z","iopub.execute_input":"2026-09-30T15:56:44.069267Z","iopub.status.idle":"2026-09-30T15:56:44.076767Z","shell.execute_reply.started":"2026-09-30T15:56:44.069246Z","shell.execute_reply":"2026-09-30T15:56:44.075882Z"}},"outputs":[],"execution_count":null},{"id":"4021a170","cell_type":"markdown","source":"## Section 4: Duplicate-Aware Train / Validation / Test Split\n\n**Goal:** Build three non-overlapping subsets and the final `tf.data` pipelines for each.\n\n- **≈70% Training** — used for weight updates.\n- **≈15% Validation** — used for early stopping (validation loss), learning-rate reduction, checkpoint selection, pilot-trial selection, TTA choice and QWK threshold calibration.\n- **≈15% Test** — held out until the final evaluation in Section 7 and used exactly once.\n\n**Duplicate control (Section 4.1):** APTOS 2019 contains repeated photographs of the same retina, sometimes with different grades. Every image is border-cropped, resized and given a 64-bit perceptual hash (`imagehash.phash`); images whose hashes differ by at most 4 bits are joined into one `duplicate_group` with union-find. Because union-find grouping is transitive, a few very large groups can form; groups of more than 50 images are therefore always placed in training. Every remaining group is placed whole, largest first, into the split (train, validation or test) that most needs that group's classes to reach 70/15/15 per class; the order of equal-sized groups is shuffled by a seed, 200 seeds (starting at 42) are tried, and the candidate closest to 70/15/15 overall and per class is kept. The search uses only labels and group IDs, never model results (`duplicate_group_sizes.csv`/`.png` and `largest_duplicate_groups.png` show the group sizes). Whole groups are always kept together, and the notebook **asserts** that validation and test each hold at least 13% of images with at least 20 images of every class, and that no image path, `patient_id` or `duplicate_group` appears in more than one split and that every near-duplicate pair shares a split. Groups whose copies carry conflicting grades are kept together in one split and listed in `duplicate_groups.csv`; nothing is deleted silently. Example pairs are saved to `duplicate_examples.png`.\n\n**Limitations:** APTOS provides **no patient identifiers**, so `patient_id = id_code`; two photographs of the same patient that are not near-duplicates (for example left and right eyes) can fall into different splits. The dataset is **single-source**, and the test split is small (550 images, with only **29 Severe** cases), so per-class test metrics — especially for Severe — have wide uncertainty.\n\nThe split is regenerated deterministically on every run (the chosen search seed is printed) and saved to `Config.REPORTS_DIR` and the project's `splits/` folder. Class weights are recomputed on training labels only, so validation and test labels never influence training.","metadata":{}},{"id":"6b8a1dd29722495ca21c80c77c20b876","cell_type":"code","source":"# ============================================================\n# SECTION 4.1 — DUPLICATE-AWARE STRATIFIED SPLIT (70 / 15 / 15)\n#\n# APTOS 2019 contains exact and near-duplicate photographs (the same\n# retina saved more than once, sometimes with different grades). If one\n# copy lands in training and another in test, the test score measures\n# memorisation, not generalisation. So:\n#   1. Every image is cropped (black border removed), resized to\n#      Config.PHASH_SIZE and given a 64-bit perceptual hash (imagehash.phash).\n#   2. Images whose hashes differ in <= Config.DUPLICATE_HAMMING_THRESHOLD\n#      bits are joined with union-find into one duplicate_group.\n#   3. Split search (labels and group IDs only, never model results):\n#      groups larger than FORCE_TRAIN_GROUP_SIZE images go to train;\n#      every other group is placed whole, largest first, into the split\n#      (train / validation / test) that most needs that group's classes to\n#      reach 70/15/15 per class. The order of equal-sized groups is shuffled\n#      by a seed; SPLIT_SEARCH_SEEDS seeds are tried and the candidate\n#      closest to 70/15/15 (overall and per class) is kept. Union-find\n#      chaining can create a few very large groups; forcing them into train\n#      stops a single group from unbalancing validation or test.\n# Groups whose members carry different grades are kept together in ONE\n# split and listed in duplicate_groups.csv; nothing is deleted silently.\n# APTOS has no patient identifiers, so patient_id = id_code; both eyes of\n# one person may still fall into different splits (documented limitation).\n# ============================================================\n\nimport imagehash\nfrom concurrent.futures import ThreadPoolExecutor\n\nSPLITS_DIR = Path(_PROJECT_ROOT) / \"splits\"\nSPLIT_SEARCH_SEEDS = 200         # candidate splits tried (seeds Config.SEED ... Config.SEED + 199)\nFORCE_TRAIN_GROUP_SIZE = 50      # duplicate groups larger than this always go to train\nSPLIT_TARGETS = {\"train\": Config.TRAIN_RATIO, \"validation\": Config.VAL_RATIO, \"test\": Config.TEST_RATIO}\nSPLIT_NAMES = list(SPLIT_TARGETS)\nMIN_EVAL_SHARE = 0.13            # validation and test must each hold at least 13% of images\nMIN_EVAL_CLASS_COUNT = 20        # every class needs at least 20 images in validation and in test\n\n\ndef compute_phash(filepath: str, size: int = Config.PHASH_SIZE) -> imagehash.ImageHash:\n    \"\"\"Perceptual hash of the border-cropped, resized fundus image.\"\"\"\n    image = crop_image_from_gray(load_image_rgb(filepath))\n    image = cv2.resize(image, (size, size), interpolation=cv2.INTER_AREA)\n    return imagehash.phash(Image.fromarray(image))\n\n\ndef hamming_distance_matrix(hashes: List[imagehash.ImageHash]) -> np.ndarray:\n    \"\"\"Pairwise Hamming distances between 64-bit perceptual hashes (N x N, int32).\"\"\"\n    bits = np.stack([h.hash.flatten() for h in hashes]).astype(np.int32)\n    return bits @ (1 - bits).T + (1 - bits) @ bits.T\n\n\ndef union_find_groups(n_items: int, pairs: np.ndarray) -> np.ndarray:\n    \"\"\"Connected components of the near-duplicate graph; returns a group id per item.\"\"\"\n    parent = np.arange(n_items)\n\n    def find(item: int) -> int:\n        while parent[item] != item:\n            parent[item] = parent[parent[item]]\n            item = parent[item]\n        return item\n\n    for first, second in pairs:\n        root_a, root_b = find(int(first)), find(int(second))\n        if root_a != root_b:\n            parent[max(root_a, root_b)] = min(root_a, root_b)\n    roots = np.array([find(item) for item in range(n_items)])\n    _, group_ids = np.unique(roots, return_inverse=True)  # consecutive ids in first-seen root order\n    return group_ids\n\n\n# ---- 1. Perceptual hashes (reused if this cell already hashed the same images) --\nsplit_source_df = df_labels.reset_index(drop=True).copy()\n_hash_cache = globals().get(\"_phash_by_filepath\", {})\nif not set(split_source_df[\"filepath\"]).issubset(_hash_cache):\n    with ThreadPoolExecutor(max_workers=min(os.cpu_count() or 1, 8)) as executor:\n        _new_hashes = list(tqdm(\n            executor.map(compute_phash, split_source_df[\"filepath\"]),\n            total=len(split_source_df),\n            desc=\"Perceptual hashing\",\n        ))\n    _hash_cache = dict(zip(split_source_df[\"filepath\"], _new_hashes))\n    _phash_by_filepath = _hash_cache\nelse:\n    print(\"[Duplicates] Reusing perceptual hashes computed earlier in this session.\")\nimage_hashes = [_hash_cache[path] for path in split_source_df[\"filepath\"]]\nsplit_source_df[\"phash\"] = [str(h) for h in image_hashes]\n\n# ---- 2. Near-duplicate graph + union-find ------------------------------------\ndistance_matrix = hamming_distance_matrix(image_hashes)\npair_i, pair_j = np.where(np.triu(distance_matrix <= Config.DUPLICATE_HAMMING_THRESHOLD, k=1))\nnear_duplicate_pairs = pd.DataFrame({\n    \"index_a\": pair_i,\n    \"index_b\": pair_j,\n    \"hamming_distance\": distance_matrix[pair_i, pair_j],\n})\nsplit_source_df[\"duplicate_group\"] = union_find_groups(\n    len(split_source_df), near_duplicate_pairs[[\"index_a\", \"index_b\"]].to_numpy()\n)\nsplit_source_df[\"patient_id\"] = split_source_df[\"id_code\"]  # APTOS has no patient IDs\n\ngroup_summary = split_source_df.groupby(\"duplicate_group\").agg(\n    n_images=(\"id_code\", \"size\"),\n    id_codes=(\"id_code\", lambda values: \";\".join(values)),\n    labels=(\"label\", lambda values: \";\".join(str(v) for v in values)),\n    n_distinct_labels=(\"label\", \"nunique\"),\n)\ngroup_summary[\"conflicting_labels\"] = group_summary[\"n_distinct_labels\"] > 1\nmulti_image_groups = group_summary[group_summary[\"n_images\"] > 1]\nconflicting_groups = multi_image_groups[multi_image_groups[\"conflicting_labels\"]]\nprint(f\"[Duplicates] pHash threshold: Hamming distance <= {Config.DUPLICATE_HAMMING_THRESHOLD} of 64 bits\")\nprint(f\"[Duplicates] Near-duplicate pairs found     : {len(near_duplicate_pairs):,}\")\nprint(f\"[Duplicates] Groups with more than one image: {len(multi_image_groups):,} \"\n      f\"({int(multi_image_groups['n_images'].sum()):,} images)\")\nprint(f\"[Duplicates] Groups with conflicting labels : {len(conflicting_groups):,} \"\n      f\"({int(conflicting_groups['n_images'].sum()):,} images) — kept together in one split, not deleted\")\n\n\n# ---- 3. Split search: whole groups placed greedily, scored on labels only -----\ndef split_distance(labels: np.ndarray, assignment: np.ndarray) -> float:\n    \"\"\"\n    Distance of a candidate split from the 70/15/15 target.\n\n    Sum of |overall share - target| over the three splits, plus the mean over\n    classes of sum_s |share of that class placed in split s - target_s|.\n    Uses labels and split names only; no model output is involved.\n    \"\"\"\n    size_term = sum(abs(np.mean(assignment == name) - target) for name, target in SPLIT_TARGETS.items())\n    class_terms = []\n    for class_id in np.unique(labels):\n        in_class = assignment[labels == class_id]\n        class_terms.append(sum(abs(np.mean(in_class == name) - target) for name, target in SPLIT_TARGETS.items()))\n    return float(size_term + np.mean(class_terms))\n\n\nall_labels = split_source_df[\"label\"].to_numpy()\nall_groups = split_source_df[\"duplicate_group\"].to_numpy()\nn_total = len(split_source_df)\ngroup_sizes = split_source_df[\"duplicate_group\"].map(split_source_df[\"duplicate_group\"].value_counts())\nforced_train_mask = (group_sizes > FORCE_TRAIN_GROUP_SIZE).to_numpy()\n\n# Pre-compute each free group's rows and per-class counts once.\n_free_rows = np.where(~forced_train_mask)[0]\n_free_group_ids, _free_inverse = np.unique(all_groups[_free_rows], return_inverse=True)\ngroup_rows = [_free_rows[_free_inverse == k] for k in range(len(_free_group_ids))]\ngroup_class_counts = np.stack([np.bincount(all_labels[rows], minlength=Config.NUM_CLASSES) for rows in group_rows])\ngroup_n = group_class_counts.sum(axis=1)\nsplit_ratios = np.array([SPLIT_TARGETS[name] for name in SPLIT_NAMES])\ntarget_counts = split_ratios[:, None] * np.bincount(all_labels, minlength=Config.NUM_CLASSES)[None, :]\nforced_counts = np.bincount(all_labels[forced_train_mask], minlength=Config.NUM_CLASSES)\n\n\ndef greedy_group_split(seed: int) -> np.ndarray:\n    \"\"\"Place whole groups, largest first, into the split that most needs their classes.\"\"\"\n    rng = np.random.default_rng(seed)\n    counts = np.zeros_like(target_counts)\n    counts[SPLIT_NAMES.index(\"train\")] += forced_counts\n    order = rng.permutation(len(group_rows))\n    order = order[np.argsort(-group_n[order], kind=\"stable\")]   # largest first, ties in seeded order\n    split_of_group = np.empty(len(group_rows), dtype=int)\n    for k in order:\n        relative_need = (target_counts - counts) / np.maximum(target_counts, 1.0)\n        best_split = int(np.argmax(relative_need @ group_class_counts[k]))\n        counts[best_split] += group_class_counts[k]\n        split_of_group[k] = best_split\n    assignment = np.full(n_total, \"train\", dtype=object)\n    for k, rows in enumerate(group_rows):\n        assignment[rows] = SPLIT_NAMES[split_of_group[k]]\n    return assignment\n\n\nprint(f\"[Split] {int(forced_train_mask.sum()):,} images in {split_source_df.loc[forced_train_mask, 'duplicate_group'].nunique()} \"\n      f\"groups larger than {FORCE_TRAIN_GROUP_SIZE} images are forced into train.\")\nprint(f\"[Split] Searching {SPLIT_SEARCH_SEEDS} seeds over {len(group_rows):,} remaining groups \"\n      f\"({len(_free_rows):,} images).\")\n\nbest_seed, best_distance, best_assignment = None, np.inf, None\nfor search_seed in tqdm(range(Config.SEED, Config.SEED + SPLIT_SEARCH_SEEDS), desc=\"Split search\"):\n    assignment = greedy_group_split(search_seed)\n    distance = split_distance(all_labels, assignment)\n    if distance < best_distance:\n        best_seed, best_distance, best_assignment = search_seed, distance, assignment\n\nsplit_source_df[\"split\"] = best_assignment\nprint(f\"[Split] Chosen search seed: {best_seed} | distance from 70/15/15 (sizes + per-class shares): {best_distance:.4f}\")\n\nmanifest_columns = [\"filepath\", \"label\", \"split\", \"patient_id\", \"duplicate_group\"]\ntrain_df = split_source_df[split_source_df[\"split\"] == \"train\"].reset_index(drop=True)\nval_df = split_source_df[split_source_df[\"split\"] == \"validation\"].reset_index(drop=True)\ntest_df = split_source_df[split_source_df[\"split\"] == \"test\"].reset_index(drop=True)\n\n# ---- 4. Integrity assertions: no image or duplicate group in two splits ------\n_split_frames = {\"train\": train_df, \"validation\": val_df, \"test\": test_df}\nassert sum(len(frame) for frame in _split_frames.values()) == len(split_source_df)\nassert split_source_df[\"filepath\"].is_unique, \"An image path appears more than once.\"\nfor _group_col in (\"filepath\", \"patient_id\", \"duplicate_group\"):\n    _groups = [set(frame[_group_col]) for frame in _split_frames.values()]\n    assert not (_groups[0] & _groups[1] or _groups[0] & _groups[2] or _groups[1] & _groups[2]), \\\n        f\"Split overlap found in {_group_col}\"\n    print(_group_col, \"overlap check passed\")\n_pair_splits = split_source_df[\"split\"].to_numpy()\nassert (_pair_splits[near_duplicate_pairs[\"index_a\"]] == _pair_splits[near_duplicate_pairs[\"index_b\"]]).all(), \\\n    \"A near-duplicate pair was split across train/validation/test.\"\nprint(\"near-duplicate pair check passed: every pair shares one split\")\nassert split_source_df.loc[forced_train_mask, \"split\"].eq(\"train\").all(), \"Large duplicate groups must stay in train.\"\nfor _name in (\"validation\", \"test\"):\n    _share = len(_split_frames[_name]) / n_total\n    assert _share >= MIN_EVAL_SHARE, f\"{_name} holds only {_share:.1%} of images (minimum {MIN_EVAL_SHARE:.0%}).\"\n    _class_counts = _split_frames[_name][\"label\"].value_counts().reindex(range(Config.NUM_CLASSES), fill_value=0)\n    assert (_class_counts >= MIN_EVAL_CLASS_COUNT).all(), \\\n        f\"{_name} has a class with fewer than {MIN_EVAL_CLASS_COUNT} images: {_class_counts.to_dict()}\"\nprint(f\"size checks passed: validation and test each >= {MIN_EVAL_SHARE:.0%} of images, \"\n      f\"every class >= {MIN_EVAL_CLASS_COUNT} images in both\")\n\n# ---- 5. Save manifests, duplicate groups and the dedup log --------------------\nSPLITS_DIR.mkdir(parents=True, exist_ok=True)\nfor _name, _frame in ((\"train_split.csv\", train_df), (\"validation_split.csv\", val_df), (\"test_split.csv\", test_df)):\n    _frame[manifest_columns].to_csv(Path(Config.REPORTS_DIR) / _name, index=False)\n    _frame[manifest_columns].to_csv(SPLITS_DIR / _name, index=False)\n\ncross_split_group_leaks = int(split_source_df.groupby(\"duplicate_group\")[\"split\"].nunique().gt(1).sum())\ngroup_split = split_source_df.groupby(\"duplicate_group\")[\"split\"].first()\nduplicate_groups_export = multi_image_groups.assign(split=group_split.reindex(multi_image_groups.index))\nduplicate_groups_export.reset_index().to_csv(Path(Config.REPORTS_DIR) / \"duplicate_groups.csv\", index=False)\ndeduplication_log = pd.DataFrame([\n    {\"metric\": \"total_images_evaluated\", \"value\": len(split_source_df)},\n    {\"metric\": \"phash_hamming_threshold\", \"value\": Config.DUPLICATE_HAMMING_THRESHOLD},\n    {\"metric\": \"largest_duplicate_group_size\", \"value\": int(group_summary[\"n_images\"].max())},\n    {\"metric\": \"groups_forced_into_train\", \"value\": int(split_source_df.loc[forced_train_mask, \"duplicate_group\"].nunique())},\n    {\"metric\": \"images_forced_into_train\", \"value\": int(forced_train_mask.sum())},\n    {\"metric\": \"split_search_method\", \"value\": \"greedy whole-group placement, largest first, seeded tie order\"},\n    {\"metric\": \"split_search_seeds_tried\", \"value\": SPLIT_SEARCH_SEEDS},\n    {\"metric\": \"split_search_chosen_seed\", \"value\": int(best_seed)},\n    {\"metric\": \"split_search_distance\", \"value\": round(best_distance, 6)},\n    {\"metric\": \"near_duplicate_pairs\", \"value\": len(near_duplicate_pairs)},\n    {\"metric\": \"duplicate_groups_with_more_than_one_image\", \"value\": len(multi_image_groups)},\n    {\"metric\": \"images_in_duplicate_groups\", \"value\": int(multi_image_groups[\"n_images\"].sum())},\n    {\"metric\": \"groups_with_conflicting_labels\", \"value\": len(conflicting_groups)},\n    {\"metric\": \"images_in_conflicting_groups\", \"value\": int(conflicting_groups[\"n_images\"].sum())},\n    {\"metric\": \"cross_split_group_leaks\", \"value\": cross_split_group_leaks},\n    {\"metric\": \"leakage_status\", \"value\": \"ZERO_LEAKAGE_CONFIRMED\" if cross_split_group_leaks == 0 else \"LEAKAGE_DETECTED\"},\n])\ndeduplication_log.to_csv(Path(Config.REPORTS_DIR) / \"deduplication_log.csv\", index=False)\n\ndf_labels = pd.concat([train_df, val_df, test_df], ignore_index=True)\nprint(\"Split sizes:\", {name: f\"{len(frame)} ({len(frame) / n_total:.1%})\" for name, frame in _split_frames.items()})\nprint(\"Class counts per split:\",\n      {name: frame[\"label\"].value_counts().sort_index().to_dict() for name, frame in _split_frames.items()})\nprint(f\"[Split] Manifests saved to {Config.REPORTS_DIR} and {SPLITS_DIR}\")","metadata":{"id":"6b8a1dd29722495ca21c80c77c20b876","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T15:56:50.345612Z","iopub.execute_input":"2026-09-30T15:56:50.34607Z","iopub.status.idle":"2026-09-30T15:59:38.952496Z","shell.execute_reply.started":"2026-09-30T15:56:50.346041Z","shell.execute_reply":"2026-09-30T15:59:38.951804Z"}},"outputs":[],"execution_count":null},{"id":"91d9120eedeb465892827a8fa61f63ac","cell_type":"code","source":"# SECTION 4.2 — DUPLICATE AUDIT AND EXAMPLE PAIRS\n# Shows the deduplication summary, the duplicate-group size distribution,\n# thumbnails of the three largest groups (true duplicates or union-find\n# chaining?), lists groups whose copies carry different grades (kept\n# together in one split), and saves 3-4 example near-duplicate pairs.\n_dedupe_path = Path(Config.REPORTS_DIR) / \"deduplication_log.csv\"\ndisplay(pd.read_csv(_dedupe_path))\n\n# ---- Group-size distribution -------------------------------------------------\ngroup_size_table = (\n    group_summary[\"n_images\"].value_counts().sort_index()\n    .rename_axis(\"group_size\").reset_index(name=\"n_groups\")\n)\ngroup_size_table[\"n_images\"] = group_size_table[\"group_size\"] * group_size_table[\"n_groups\"]\ngroup_size_table.to_csv(Path(Config.REPORTS_DIR) / \"duplicate_group_sizes.csv\", index=False)\nprint(\"[Duplicates] Group-size distribution (groups of 1 = unique images):\")\nprint(group_size_table.to_string(index=False))\n\nmulti_sizes = group_summary.loc[group_summary[\"n_images\"] > 1, \"n_images\"]\nfig, axis = plt.subplots(figsize=(9, 4.5))\nif len(multi_sizes):\n    axis.hist(multi_sizes, bins=np.arange(1.5, multi_sizes.max() + 1.5, max(1, int(multi_sizes.max() // 40))),\n              color=\"#367C8A\", edgecolor=\"white\")\naxis.axvline(FORCE_TRAIN_GROUP_SIZE + 0.5, color=\"#B91C1C\", linestyle=\"--\",\n             label=f\"> {FORCE_TRAIN_GROUP_SIZE} images: forced into train\")\naxis.set_yscale(\"log\")\naxis.set_xlabel(\"Images per duplicate group\")\naxis.set_ylabel(\"Number of groups (log scale)\")\naxis.set_title(\"Duplicate-group sizes (groups with more than one image)\")\naxis.legend()\nfig.tight_layout()\nfig.savefig(Path(Config.REPORTS_DIR) / \"duplicate_group_sizes.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\n\ndef plot_largest_groups(source: pd.DataFrame, n_groups: int = 3, per_group: int = 8, save_path: str = None) -> None:\n    \"\"\"Thumbnail montage of the largest duplicate groups, to check for chaining vs true duplicates.\"\"\"\n    largest = source[\"duplicate_group\"].value_counts().head(n_groups)\n    fig, axes = plt.subplots(len(largest), per_group, figsize=(1.9 * per_group, 2.3 * len(largest)), squeeze=False)\n    for row, (group_id, size) in enumerate(largest.items()):\n        members = source[source[\"duplicate_group\"] == group_id].head(per_group)\n        for column in range(per_group):\n            axis = axes[row, column]\n            axis.axis(\"off\")\n            if column < len(members):\n                record = members.iloc[column]\n                thumb = cv2.resize(crop_image_from_gray(load_image_rgb(record[\"filepath\"])), (160, 160),\n                                   interpolation=cv2.INTER_AREA)\n                axis.imshow(thumb)\n                axis.set_title(f\"{record['id_code'][:8]} g{record['label']}\", fontsize=7)\n        axes[row, 0].text(-0.15, 0.5, f\"group {group_id}\\n{size} images\", transform=axes[row, 0].transAxes,\n                          ha=\"right\", va=\"center\", fontsize=9, fontweight=\"bold\")\n    fig.suptitle(\"Largest duplicate groups (first images of each; g = grade)\", fontsize=11, fontweight=\"bold\")\n    fig.tight_layout()\n    if save_path:\n        fig.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.show()\n\n\nplot_largest_groups(split_source_df, save_path=os.path.join(Config.REPORTS_DIR, \"largest_duplicate_groups.png\"))\n\nif len(conflicting_groups):\n    print(\"[Duplicates] Groups with conflicting labels (kept together in one split, not deleted):\")\n    display(duplicate_groups_export[duplicate_groups_export[\"conflicting_labels\"]]\n            [[\"n_images\", \"id_codes\", \"labels\", \"split\"]].head(20))\nelse:\n    print(\"[Duplicates] No duplicate group has conflicting labels.\")\n\n\ndef plot_duplicate_examples(pairs: pd.DataFrame, source: pd.DataFrame, save_path: str, max_pairs: int = 4) -> None:\n    \"\"\"Plot near-duplicate pairs side by side, conflicting-label pairs first.\"\"\"\n    if pairs.empty:\n        print(\"[Duplicates] No near-duplicate pairs to plot.\")\n        return\n    ranked = pairs.assign(\n        conflict=[\n            source.at[a, \"label\"] != source.at[b, \"label\"]\n            for a, b in zip(pairs[\"index_a\"], pairs[\"index_b\"])\n        ]\n    ).sort_values([\"conflict\", \"hamming_distance\"], ascending=[False, True]).head(max_pairs)\n\n    fig, axes = plt.subplots(len(ranked), 2, figsize=(8, 4 * len(ranked)), squeeze=False)\n    for row_index, pair in enumerate(ranked.itertuples()):\n        for column, image_index in enumerate((pair.index_a, pair.index_b)):\n            record = source.loc[image_index]\n            image = cv2.resize(crop_image_from_gray(load_image_rgb(record[\"filepath\"])), (Config.PHASH_SIZE, Config.PHASH_SIZE))\n            axes[row_index, column].imshow(image)\n            axes[row_index, column].axis(\"off\")\n            axes[row_index, column].set_title(\n                f\"{record['id_code']}\\nGrade {record['label']} ({Config.CLASS_NAMES[record['label']]}) | {record['split']}\",\n                fontsize=9,\n            )\n        axes[row_index, 0].text(\n            1.05, -0.08,\n            f\"Hamming distance {pair.hamming_distance}\"\n            + (\" — CONFLICTING GRADES\" if pair.conflict else \"\"),\n            transform=axes[row_index, 0].transAxes, ha=\"center\", fontsize=9,\n            color=\"#B91C1C\" if pair.conflict else \"#334155\", fontweight=\"bold\",\n        )\n    fig.suptitle(\"Example near-duplicate pairs (grouped into the same split)\", fontsize=12, fontweight=\"bold\")\n    fig.tight_layout()\n    fig.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"[Duplicates] Example pairs saved -> {save_path}\")\n\n\nplot_duplicate_examples(\n    near_duplicate_pairs,\n    split_source_df,\n    save_path=os.path.join(Config.REPORTS_DIR, \"duplicate_examples.png\"),\n)\n","metadata":{"id":"91d9120eedeb465892827a8fa61f63ac","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:00:24.922856Z","iopub.execute_input":"2026-09-30T16:00:24.923245Z","iopub.status.idle":"2026-09-30T16:00:32.305228Z","shell.execute_reply.started":"2026-09-30T16:00:24.923219Z","shell.execute_reply":"2026-09-30T16:00:32.304247Z"}},"outputs":[],"execution_count":null},{"id":"2345c0ace018440fb1d9b45b96047a23","cell_type":"code","source":"# Confirm the generated split is complete and uses the three expected split names.\nassert len(train_df) + len(val_df) + len(test_df) == len(df_labels)\nassert set(train_df[\"split\"]) == {\"train\"}\nassert set(val_df[\"split\"]) == {\"validation\"}\nassert set(test_df[\"split\"]) == {\"test\"}\nprint(\"Split is complete; no post-training repartition is performed.\")","metadata":{"id":"2345c0ace018440fb1d9b45b96047a23","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:01:13.582779Z","iopub.execute_input":"2026-09-30T16:01:13.583619Z","iopub.status.idle":"2026-09-30T16:01:13.59115Z","shell.execute_reply.started":"2026-09-30T16:01:13.583581Z","shell.execute_reply":"2026-09-30T16:01:13.590062Z"}},"outputs":[],"execution_count":null},{"id":"9d1dcefc","cell_type":"code","source":"# ============================================================\n# SECTION 4.2b — TRAINING-ONLY PREPROCESSING QUALITY ABLATION\n# This descriptive image-statistics audit samples the training split only.\n# It does not fit preprocessing parameters or use validation/test images.\n# ============================================================\n\nsample_frames = []\nfor class_id in range(Config.NUM_CLASSES):\n    class_train = train_df[train_df[\"label\"] == class_id]\n    if class_train.empty:\n        continue\n    sample_frames.append(\n        class_train.sample(\n            min(20, len(class_train)),\n            random_state=Config.SEED + class_id,\n        )\n    )\n\nif not sample_frames:\n    raise ValueError(\"No labeled training images are available for preprocessing quality analysis.\")\n\nsample_rows = pd.concat(sample_frames, ignore_index=True)\nif not sample_rows[\"split\"].eq(\"train\").all():\n    raise AssertionError(\"Preprocessing quality sample contains a non-training record.\")\n\nquality_records = []\nfor _, row in tqdm(sample_rows.iterrows(), total=len(sample_rows), desc=\"Training-only preprocessing ablation\"):\n    for variant in preprocessing_variants:\n        variant_image = _prepare_preprocessing_variant(row[\"filepath\"], variant)\n        quality_records.append({\n            \"filepath\": row[\"filepath\"],\n            \"split\": row[\"split\"],\n            \"label\": int(row[\"label\"]),\n            \"variant\": variant,\n            **_preprocessing_quality_metrics(variant_image),\n        })\n\nquality_df = pd.DataFrame(quality_records)\nquality_metrics = [\n    \"contrast_std\",\n    \"edge_density\",\n    \"brightness_variation\",\n    \"signal_noise_proxy\",\n    \"sharpness_laplacian_variance\",\n]\nquality_summary = quality_df.groupby(\"variant\")[quality_metrics].agg([\"mean\", \"std\"]).round(6)\n\nquality_path = os.path.join(Config.REPORTS_DIR, \"preprocessing_quality_ablation.csv\")\nquality_df.to_csv(quality_path, index=False)\nquality_summary.to_csv(os.path.join(Config.REPORTS_DIR, \"preprocessing_quality_summary.csv\"))\nprint(\"[Training-only preprocessing ablation] Sample split counts:\", sample_rows[\"split\"].value_counts().to_dict())\nprint(\"[Training-only preprocessing ablation] Mean image-statistic proxies:\")\nprint(quality_df.groupby(\"variant\")[quality_metrics].mean().round(6).to_string())\nprint(f\"[Training-only preprocessing ablation] Saved -> {quality_path}\")\n\nfig, axes = plt.subplots(1, 2, figsize=(13, 5))\nvariant_order = preprocessing_variants\nsns.boxplot(data=quality_df, x=\"variant\", y=\"contrast_std\", order=variant_order, ax=axes[0], color=\"#3B82A0\")\naxes[0].set_title(\"Grayscale contrast\")\naxes[0].set_xlabel(\"Preprocessing variant\")\naxes[0].set_ylabel(\"Pixel standard deviation (0-1)\")\naxes[0].tick_params(axis=\"x\", rotation=20)\nsns.boxplot(data=quality_df, x=\"variant\", y=\"sharpness_laplacian_variance\", order=variant_order, ax=axes[1], color=\"#D28E45\")\naxes[1].set_title(\"Sharpness proxy\")\naxes[1].set_xlabel(\"Preprocessing variant\")\naxes[1].set_ylabel(\"Laplacian variance (grayscale 0-255)\")\naxes[1].tick_params(axis=\"x\", rotation=20)\nfig.suptitle(\"Training-only preprocessing statistics (descriptive, not model performance)\", fontsize=12)\nplt.tight_layout()\npreprocessing_metrics_path = os.path.join(Config.REPORTS_DIR, \"preprocessing_contrast_sharpness.png\")\nplt.savefig(preprocessing_metrics_path, dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(f\"[Training-only preprocessing ablation] Contrast/sharpness plot saved -> {preprocessing_metrics_path}\")\n\ncomparison_row = sample_rows.iloc[0]\ncomparison_images = [\n    _prepare_preprocessing_variant(comparison_row[\"filepath\"], variant)\n    for variant in preprocessing_variants\n]\nfig, axes = plt.subplots(1, len(preprocessing_variants), figsize=(16, 4))\nfor axis, variant, image in zip(axes, preprocessing_variants, comparison_images):\n    axis.imshow(image)\n    axis.set_title(variant.replace(\"_\", \"\\n\"), fontsize=9)\n    axis.axis(\"off\")\nplt.tight_layout()\ncomparison_path = os.path.join(Config.REPORTS_DIR, \"preprocessing_ablation_examples.png\")\nplt.savefig(comparison_path, dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(f\"[Training-only preprocessing ablation] Visual comparison saved -> {comparison_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:01:21.325931Z","iopub.execute_input":"2026-09-30T16:01:21.326214Z","iopub.status.idle":"2026-09-30T16:02:27.295982Z","shell.execute_reply.started":"2026-09-30T16:01:21.326193Z","shell.execute_reply":"2026-09-30T16:02:27.295242Z"}},"outputs":[],"execution_count":null},{"id":"d346922145bd45cfa82291b1b3398e9c","cell_type":"code","source":"# ============================================================\n# SECTION 4.3 — PER-CLASS SIZE CONFIRMATION\n# Print class counts for each split to confirm stratification\n# worked correctly. All three splits should have roughly the\n# same percentage breakdown as the full dataset.\n# ============================================================\n\ndef print_split_distribution(\n    train_df: pd.DataFrame,\n    val_df:   pd.DataFrame,\n    test_df:  pd.DataFrame,\n) -> None:\n    \"\"\"\n    Print per-class counts for all three splits side by side.\n\n    Allows visual verification that stratification preserved class\n    proportions. Large deviations from expected ratios would indicate\n    a problem with the stratified split.\n\n    Args:\n        train_df: Training split DataFrame.\n        val_df:   Validation split DataFrame.\n        test_df:  Test split DataFrame.\n    \"\"\"\n    all_classes = sorted(train_df[\"label\"].unique())\n\n    header = f\"{'Class':<5} {'Name':<22} {'Train':>7} {'Val':>7} {'Test':>7}  {'Train%':>7} {'Val%':>7} {'Test%':>7}\"\n    print(\"\\n\" + \"=\" * len(header))\n    print(\"  PER-CLASS SPLIT SIZES\")\n    print(\"=\" * len(header))\n    print(header)\n    print(\"-\" * len(header))\n\n    for cls in all_classes:\n        n_train = (train_df[\"label\"] == cls).sum()\n        n_val   = (val_df[\"label\"]   == cls).sum()\n        n_test  = (test_df[\"label\"]  == cls).sum()\n        n_total = n_train + n_val + n_test\n\n        pct_train = n_train / n_total * 100 if n_total else 0\n        pct_val   = n_val   / n_total * 100 if n_total else 0\n        pct_test  = n_test  / n_total * 100 if n_total else 0\n\n        print(\n            f\"  {cls:<5} {Config.CLASS_NAMES[cls]:<22} \"\n            f\"{n_train:>7,} {n_val:>7,} {n_test:>7,}  \"\n            f\"{pct_train:>6.1f}% {pct_val:>6.1f}% {pct_test:>6.1f}%\"\n        )\n\n    print(\"=\" * len(header))\n    print(\"  Expected: ~70% / ~15% / ~15% per class (stratification check)\")\n    print(\"=\" * len(header) + \"\\n\")\n\n\nprint_split_distribution(train_df, val_df, test_df)\n\n\n# ------------------------------------------------------------\n# CLASS-TO-INDEX CHECK TABLE (label -> class name -> count, per split)\n# Confirms every split uses the same integer->stage mapping, that all\n# labels are valid ICDR ids, and that each image's class folder (when the\n# path is .../<split>/<class_id>/<file>) agrees with its manifest label.\n# ------------------------------------------------------------\ndef _folder_label(filepath: str) -> Optional[int]:\n    parent = Path(str(filepath)).parent.name\n    return int(parent) if parent.isdigit() else None\n\n\n_class_index_rows = []\nfor _split_name, _frame in ((\"train\", train_df), (\"validation\", val_df), (\"test\", test_df)):\n    _invalid = sorted(set(_frame[\"label\"]) - set(range(Config.NUM_CLASSES)))\n    if _invalid:\n        raise ValueError(f\"{_split_name} split contains labels outside 0-{Config.NUM_CLASSES - 1}: {_invalid}\")\n    _folder_labels = _frame[\"filepath\"].map(_folder_label)\n    for _class_id in range(Config.NUM_CLASSES):\n        _mask = _frame[\"label\"] == _class_id\n        _checked = _folder_labels[_mask].dropna()\n        _class_index_rows.append({\n            \"split\": _split_name,\n            \"label_index\": _class_id,\n            \"class_name\": Config.CLASS_NAMES[_class_id],\n            \"image_count\": int(_mask.sum()),\n            \"percent_of_split\": float(_mask.mean() * 100),\n            \"folder_checked\": int(len(_checked)),\n            \"folder_mismatches\": int((_checked != _class_id).sum()),\n        })\n\nclass_index_check = pd.DataFrame(_class_index_rows)\nclass_index_check_path = os.path.join(Config.REPORTS_DIR, \"class_index_check.csv\")\nclass_index_check.to_csv(class_index_check_path, index=False)\n\nprint(\"[Class-Index Check] label -> class name -> count, per split:\")\nfor _split_name, _table in class_index_check.groupby(\"split\", sort=False):\n    print(f\"\\n  {_split_name.upper()} ({int(_table['image_count'].sum()):,} images)\")\n    print(f\"  {'Index':<6} {'Class name':<22} {'Count':>7} {'%':>7} {'Folder mismatches':>18}\")\n    for _, _row in _table.iterrows():\n        print(\n            f\"  {_row['label_index']:<6} {_row['class_name']:<22} {_row['image_count']:>7,} \"\n            f\"{_row['percent_of_split']:>6.1f}% {_row['folder_mismatches']:>18}\"\n        )\nif class_index_check[\"folder_mismatches\"].sum() > 0:\n    raise ValueError(\"Some image class folders disagree with their manifest labels; see class_index_check.csv.\")\nprint(f\"\\n[Class-Index Check] Mapping consistent across splits. Saved -> {class_index_check_path}\")\n","metadata":{"id":"d346922145bd45cfa82291b1b3398e9c","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:02:33.636221Z","iopub.execute_input":"2026-09-30T16:02:33.636991Z","iopub.status.idle":"2026-09-30T16:02:33.693511Z","shell.execute_reply.started":"2026-09-30T16:02:33.636962Z","shell.execute_reply":"2026-09-30T16:02:33.692674Z"}},"outputs":[],"execution_count":null},{"id":"bc102acb29f3417692176c7bacc63de2","cell_type":"code","source":"# ============================================================\n# SECTION 4.4 — RECOMPUTE CLASS WEIGHTS ON TRAINING SET ONLY\n#\n# WHY recompute here (not use the full-dataset preview from Section 3):\n#   class_weight must reflect only the training distribution.\n#   Using weights computed on the full dataset would incorporate\n#   val/test label counts into the training objective — a subtle\n#   form of data leakage. The difference is small (70% vs 100% of\n#   data) but it is the principled approach and correct by convention.\n#\n# Two weightings are computed from the training labels:\n#   balanced -> sklearn \"balanced\" inverse-frequency weights\n#   sqrt     -> square root of the balanced weights, rescaled so the mean\n#               weight is 1 (a softer correction; less over-weighting of\n#               the rare Severe stage)\n# Config.CLASS_WEIGHT_MODE (\"balanced\", \"sqrt\" or \"none\") picks which one\n# model.fit() receives; active_class_weights() is the single source used\n# by the pilot trials and the final fit.\n# ============================================================\n\n# Recompute class weights on training labels only\ntrain_labels = train_df[\"label\"].tolist()\nclass_weight_balanced = compute_class_weights(train_labels)\nclass_weight_sqrt = compute_sqrt_class_weights(class_weight_balanced)\nclass_weight_dict = class_weight_balanced  # balanced weights (kept for existing references)\n\n\ndef active_class_weights() -> Optional[Dict[int, float]]:\n    \"\"\"Class weights for model.fit() under the current Config (None when weighting is off).\"\"\"\n    if not Config.USE_CLASS_WEIGHTS or Config.CLASS_WEIGHT_MODE == \"none\":\n        return None\n    if Config.CLASS_WEIGHT_MODE == \"sqrt\":\n        return class_weight_sqrt\n    if Config.CLASS_WEIGHT_MODE == \"balanced\":\n        return class_weight_balanced\n    raise ValueError(f\"Unknown Config.CLASS_WEIGHT_MODE: {Config.CLASS_WEIGHT_MODE!r}\")\n\n\nprint(\"[Class Weights] Final weights (computed on training split only):\")\nprint(f\"  {'Class':<25} {'balanced':>9} {'sqrt':>9}\")\nfor cls in sorted(class_weight_balanced):\n    print(f\"  {cls} {Config.CLASS_NAMES[cls]:<23} {class_weight_balanced[cls]:>9.4f} {class_weight_sqrt[cls]:>9.4f}\")\nprint(f\"\\n  Mode used for training: CLASS_WEIGHT_MODE='{Config.CLASS_WEIGHT_MODE}', \"\n      f\"USE_CLASS_WEIGHTS={Config.USE_CLASS_WEIGHTS}.\")\nprint(\"  The class_weights_off, class_weights_sqrt and oversampling_instead pilot trials in Section 6 test the alternatives.\")\n","metadata":{"id":"bc102acb29f3417692176c7bacc63de2","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:02:57.679184Z","iopub.execute_input":"2026-09-30T16:02:57.679677Z","iopub.status.idle":"2026-09-30T16:02:57.689649Z","shell.execute_reply.started":"2026-09-30T16:02:57.679641Z","shell.execute_reply":"2026-09-30T16:02:57.689002Z"}},"outputs":[],"execution_count":null},{"id":"7c67a056","cell_type":"code","source":"# ============================================================\n# SECTION 4.4b — SPLIT AND CLASS-WEIGHT FIGURES\n# Split counts come from the unchanged manifests. Weight effects use\n# training labels only; no images are duplicated or oversampled.\n# ============================================================\n\nclass_ids = list(range(Config.NUM_CLASSES))\nclass_names = [Config.CLASS_NAMES[class_id] for class_id in class_ids]\nsplit_frames = {\n    \"Train\": train_df,\n    \"Validation\": val_df,\n    \"Test\": test_df,\n}\nsplit_counts = pd.DataFrame(\n    {\n        split_name: frame[\"label\"].value_counts().reindex(class_ids, fill_value=0).to_numpy()\n        for split_name, frame in split_frames.items()\n    },\n    index=class_names,\n)\nsplit_counts.index.name = \"class_name\"\nsplit_counts.to_csv(os.path.join(Config.REPORTS_DIR, \"split_class_counts.csv\"))\n\nfig, ax = plt.subplots(figsize=(11, 6))\nsplit_counts.plot(kind=\"bar\", ax=ax, color=[\"#367C8A\", \"#D49A45\", \"#7A8F55\"], width=0.78)\nax.set_title(\"Class Distribution Across the Saved Dataset Splits\")\nax.set_xlabel(\"Diabetic retinopathy stage\")\nax.set_ylabel(\"Number of images\")\nax.legend(title=\"Saved split\")\nax.tick_params(axis=\"x\", rotation=18)\nax.grid(axis=\"y\", alpha=0.25)\nfig.tight_layout()\nsplit_plot_path = os.path.join(Config.REPORTS_DIR, \"split_class_distribution.png\")\nfig.savefig(split_plot_path, dpi=160, bbox_inches=\"tight\")\nplt.show()\nprint(f\"[Split Figures] Saved manifest class counts -> {split_plot_path}\")\n\ntrain_counts = train_df[\"label\"].value_counts().reindex(class_ids, fill_value=0).astype(float)\n# Always plot both computed weightings so the figure shows what each would do,\n# whichever mode (balanced, sqrt or none) is used for training.\napplied_weights = pd.Series(class_weight_balanced, dtype=float).reindex(class_ids)\nsqrt_weights = pd.Series(class_weight_sqrt, dtype=float).reindex(class_ids)\nif applied_weights.isna().any() or sqrt_weights.isna().any():\n    raise ValueError(\"A training class is missing its computed class weight.\")\n\nunweighted_share = train_counts / train_counts.sum() * 100\nweighted_mass = train_counts * applied_weights\nweighted_share = weighted_mass / weighted_mass.sum() * 100\nsqrt_mass = train_counts * sqrt_weights\nsqrt_share = sqrt_mass / sqrt_mass.sum() * 100\nweighting_summary = pd.DataFrame({\n    \"class_id\": class_ids,\n    \"class_name\": class_names,\n    \"training_images\": train_counts.values.astype(int),\n    \"class_weight_balanced\": applied_weights.values,\n    \"class_weight_sqrt\": sqrt_weights.values,\n    \"class_weight_mode_in_training\": Config.CLASS_WEIGHT_MODE if Config.USE_CLASS_WEIGHTS else \"none\",\n    \"unweighted_loss_mass_percent\": unweighted_share.values,\n    \"balanced_loss_mass_percent\": weighted_share.values,\n    \"sqrt_loss_mass_percent\": sqrt_share.values,\n})\nweighting_summary.to_csv(os.path.join(Config.REPORTS_DIR, \"class_weighting_effect.csv\"), index=False)\n\n# Headline imbalance ratio (largest class / smallest class) before and after weighting.\nfull_counts = df_labels[\"label\"].value_counts().reindex(class_ids, fill_value=0).astype(float)\nimbalance_summary = pd.DataFrame([\n    {\"stage\": \"Full dataset (image counts)\", \"imbalance_ratio\": full_counts.max() / full_counts.min()},\n    {\"stage\": \"Training split (image counts)\", \"imbalance_ratio\": train_counts.max() / train_counts.min()},\n    {\"stage\": \"Training split with balanced class weights (effective loss mass)\",\n     \"imbalance_ratio\": weighted_mass.max() / weighted_mass.min()},\n    {\"stage\": \"Training split with sqrt class weights (effective loss mass)\",\n     \"imbalance_ratio\": sqrt_mass.max() / sqrt_mass.min()},\n])\nimbalance_summary[\"class_weight_mode_in_training\"] = Config.CLASS_WEIGHT_MODE if Config.USE_CLASS_WEIGHTS else \"none\"\nimbalance_summary.to_csv(os.path.join(Config.REPORTS_DIR, \"imbalance_ratio_summary.csv\"), index=False)\nraw_imbalance_ratio = float(imbalance_summary.loc[1, \"imbalance_ratio\"])\nweighted_imbalance_ratio = float(imbalance_summary.loc[2, \"imbalance_ratio\"])\nsqrt_imbalance_ratio = float(imbalance_summary.loc[3, \"imbalance_ratio\"])\nprint(\n    \"[Imbalance] Largest/smallest class ratio: \"\n    + \" -> \".join(f\"{ratio:.2f}x\" for ratio in imbalance_summary[\"imbalance_ratio\"])\n    + \"  (full dataset -> training split -> balanced weights -> sqrt weights)\"\n)\nif not Config.USE_CLASS_WEIGHTS:\n    print(\"[Imbalance] Class weights are disabled for this run, so training uses the unweighted training ratio.\")\nelse:\n    print(f\"[Imbalance] Training uses the '{Config.CLASS_WEIGHT_MODE}' weights.\")\n\npositions = np.arange(Config.NUM_CLASSES)\nbar_width = 0.27\nfig, ax = plt.subplots(figsize=(11, 6))\nax.bar(positions - bar_width, unweighted_share, bar_width, label=\"Without class weights\", color=\"#6E8792\")\nax.bar(positions, weighted_share, bar_width, label=\"Balanced class weights\", color=\"#D8784B\")\nax.bar(positions + bar_width, sqrt_share, bar_width, label=\"Sqrt class weights (softer)\", color=\"#7A8F55\")\nax.set_title(\n    \"Class Weights Rebalance Expected Training-Loss Contribution\\n\"\n    f\"Imbalance ratio {raw_imbalance_ratio:.2f}x -> {weighted_imbalance_ratio:.2f}x (balanced) / \"\n    f\"{sqrt_imbalance_ratio:.2f}x (sqrt)\"\n)\nax.set_xlabel(\"Diabetic retinopathy stage\")\nax.set_ylabel(\"Share of summed per-example loss weight (%)\")\nax.set_xticks(positions, class_names, rotation=18, ha=\"right\")\nax.legend()\nax.grid(axis=\"y\", alpha=0.25)\nax.text(\n    0.01,\n    -0.27,\n    \"Weighted contribution assumes equal per-example losses; it is not an image-count change or an accuracy result.\\n\"\n    + (f\"Training uses the '{Config.CLASS_WEIGHT_MODE}' weights (CLASS_WEIGHT_MODE); the other bars show the alternatives.\"\n       if Config.USE_CLASS_WEIGHTS\n       else \"Class weights are NOT applied during training (CLASS_WEIGHT_MODE='none'); weighted bars show their potential effect.\"),\n    transform=ax.transAxes,\n    fontsize=9,\n)\nfig.tight_layout()\nweighting_plot_path = os.path.join(Config.REPORTS_DIR, \"class_weighting_effect.png\")\nfig.savefig(weighting_plot_path, dpi=160, bbox_inches=\"tight\")\nplt.show()\nprint(f\"[Class-Weight Figures] Saved weighted-loss comparison -> {weighting_plot_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:03:27.659774Z","iopub.execute_input":"2026-09-30T16:03:27.660232Z","iopub.status.idle":"2026-09-30T16:03:28.58655Z","shell.execute_reply.started":"2026-09-30T16:03:27.66019Z","shell.execute_reply":"2026-09-30T16:03:28.585887Z"}},"outputs":[],"execution_count":null},{"id":"2e10eddc","cell_type":"code","source":"import hashlib\nfrom concurrent.futures import ThreadPoolExecutor\n\nif Path(\"/kaggle/temp\").exists():\n    _cache_root = Path(\"/kaggle/temp\")\nelif Path(\"/kaggle/working\").exists():\n    _cache_root = Path(\"/kaggle/working\")\nelse:\n    _cache_root = Path(_RUN_ROOT) / \"cache\"\nCACHE_DIR = _cache_root / \"preproc_cache\" / str(Config.IMG_SIZE)\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\n\n\ndef _cache_path(fp):\n    return str(CACHE_DIR / (hashlib.md5(fp.encode()).hexdigest() + \".png\"))\n\n\ndef _cache_one(fp):\n    out = _cache_path(fp)\n    if not os.path.exists(out):\n        img = preprocess_image(fp, img_size=Config.IMG_SIZE)\n        img8 = np.clip(img * 255.0, 0, 255).round().astype(np.uint8)\n        if not cv2.imwrite(out, cv2.cvtColor(img8, cv2.COLOR_RGB2BGR), [cv2.IMWRITE_PNG_COMPRESSION, 1]):\n            raise OSError(f\"Could not write cache file {out} (check disk space and path length).\")\n    return out\n\n\nall_paths = pd.concat([train_df, val_df, test_df], ignore_index=True)[\"filepath\"].tolist()\nwith ThreadPoolExecutor(max_workers=os.cpu_count()) as executor:\n    list(tqdm(executor.map(_cache_one, all_paths), total=len(all_paths), desc=\"Caching\"))\n\n\ndef build_cached_dataset(\n    filepaths,\n    labels,\n    batch_size=Config.BATCH_SIZE,\n    shuffle=True,\n    augment_fn=None,\n    oversample=False,\n):\n    \"\"\"\n    Build a batched tf.data pipeline from the preprocessed PNG cache.\n\n    With oversample=True (training split only) the stream is class-balanced and\n    infinite, so model.fit() must be given steps_per_epoch.\n    \"\"\"\n    cached = [_cache_path(fp) for fp in filepaths]\n\n    def _load(path, label):\n        image = tf.io.decode_png(tf.io.read_file(path), channels=3)\n        image = tf.image.convert_image_dtype(image, tf.float32)\n        image.set_shape([Config.IMG_SIZE, Config.IMG_SIZE, 3])\n        return image, tf.one_hot(tf.cast(label, tf.int32), Config.NUM_CLASSES)\n\n    if oversample:\n        dataset = build_oversampled_dataset(cached, labels, _load)\n    else:\n        dataset = tf.data.Dataset.from_tensor_slices((cached, labels))\n        if shuffle:\n            dataset = dataset.shuffle(\n                len(cached), seed=Config.SEED, reshuffle_each_iteration=True\n            )\n        dataset = dataset.map(_load, num_parallel_calls=tf.data.AUTOTUNE)\n    if augment_fn is not None:\n        dataset = dataset.map(augment_fn, num_parallel_calls=tf.data.AUTOTUNE)\n    return dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)\n\n\nprint(f\"[Cache] Preprocessed image cache ready at {CACHE_DIR}\")\n","metadata":{"language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:03:43.317984Z","iopub.execute_input":"2026-09-30T16:03:43.318371Z","iopub.status.idle":"2026-09-30T16:07:10.066478Z","shell.execute_reply.started":"2026-09-30T16:03:43.318345Z","shell.execute_reply":"2026-09-30T16:07:10.065756Z"}},"outputs":[],"execution_count":null},{"id":"faa206e133d94e67b908ec1d5d14120e","cell_type":"code","source":"# ============================================================\n# SECTION 4.5 — BUILD tf.data DATASETS\n# Wire together the preprocessing cache (Section 4),\n# augmentation (Section 3), and the split DataFrames to produce\n# three ready-to-use tf.data.Dataset objects.\n#\n# Key decisions:\n#   - TRAINING dataset: shuffle=True, augment_fn applied\n#   - VALIDATION dataset: shuffle=False, no augmentation\n#   - TEST dataset: shuffle=False, no augmentation\n# ============================================================\n\n# Extract file paths and labels as plain Python lists\ntrain_paths = train_df[\"filepath\"].tolist()\ntrain_labels_list = train_df[\"label\"].tolist()\n\nval_paths = val_df[\"filepath\"].tolist()\nval_labels_list = val_df[\"label\"].tolist()\n\ntest_paths = test_df[\"filepath\"].tolist()\ntest_labels_list = test_df[\"label\"].tolist()\n\ntrain_dataset = build_cached_dataset(\n    filepaths=train_paths,\n    labels=train_labels_list,\n    batch_size=Config.BATCH_SIZE,\n    shuffle=True,\n    augment_fn=augment_fn,\n    oversample=Config.USE_OVERSAMPLING,\n)\nval_dataset = build_cached_dataset(\n    filepaths=val_paths,\n    labels=val_labels_list,\n    batch_size=Config.BATCH_SIZE,\n    shuffle=False,\n    augment_fn=None,\n)\ntest_dataset = build_cached_dataset(\n    filepaths=test_paths,\n    labels=test_labels_list,\n    batch_size=Config.BATCH_SIZE,\n    shuffle=False,\n    augment_fn=None,\n)\n\nsteps_per_epoch = int(np.ceil(len(train_paths) / Config.BATCH_SIZE))  # required when oversampling\nvalidation_steps = len(val_paths) // Config.BATCH_SIZE\ntest_steps = len(test_paths) // Config.BATCH_SIZE\n\nprint(\"[Datasets] Cached tf.data pipelines built:\")\nprint(f\"  train_dataset  : {len(train_paths):,} images  | {steps_per_epoch} batches/epoch\")\nprint(f\"  val_dataset    : {len(val_paths):,} images  | {validation_steps} batches/epoch\")\nprint(f\"  test_dataset   : {len(test_paths):,} images  | {test_steps} batches (eval)\")\nprint(f\"  Batch size     : {Config.BATCH_SIZE}\")\nprint(\"  Augmentation   : train only\")\nprint(\"  Shuffle        : train only\")\nprint(f\"  Oversampling   : {Config.USE_OVERSAMPLING} (train only)\")","metadata":{"id":"faa206e133d94e67b908ec1d5d14120e","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:07:50.605621Z","iopub.execute_input":"2026-09-30T16:07:50.606477Z","iopub.status.idle":"2026-09-30T16:07:52.121285Z","shell.execute_reply.started":"2026-09-30T16:07:50.606445Z","shell.execute_reply":"2026-09-30T16:07:52.120633Z"}},"outputs":[],"execution_count":null},{"id":"6c5ec541d431479f9b84c8d4aed6fa8d","cell_type":"code","source":"# ============================================================\n# SECTION 4.6 — BATCH SHAPE VERIFICATION\n# Peek at one batch to confirm shapes and dtypes are correct\n# before we pass these datasets to the model.\n# ============================================================\n\ndef verify_dataset_batch(dataset: tf.data.Dataset, name: str) -> None:\n    \"\"\"\n    Inspect one batch from a tf.data.Dataset and print shape/dtype info.\n\n    Args:\n        dataset: A batched tf.data.Dataset (image, label) pairs.\n        name:    Human-readable name for the dataset (e.g. \"train\").\n    \"\"\"\n    for images, labels in dataset.take(1):\n        print(f\"[Batch Check] {name} dataset:\")\n        print(f\"  images shape : {images.shape}   dtype: {images.dtype}\")\n        print(f\"  labels shape : {labels.shape}   dtype: {labels.dtype}\")\n        print(f\"  images range : [{images.numpy().min():.3f}, {images.numpy().max():.3f}]\")\n        print(f\"  labels sample: {labels[0].numpy()}  (one-hot, sum={labels[0].numpy().sum():.0f})\")\n\n\nverify_dataset_batch(train_dataset, \"train\")\nverify_dataset_batch(val_dataset,   \"val\")\nverify_dataset_batch(test_dataset,  \"test\")\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"  SECTION 4 COMPLETE — Train / Val / Test Split\")\nprint(\"=\" * 60)\nprint(f\"  train_df       : {len(train_df):,} rows\")\nprint(f\"  val_df         : {len(val_df):,} rows\")\nprint(f\"  test_df        : {len(test_df):,} rows\")\nprint(f\"  class_weight_dict : computed on training labels\")\nprint(f\"  train_dataset, val_dataset, test_dataset : ready for model.fit()\")\nprint(\"  Next -> Section 5: Model — EfficientNetB3 Transfer Learning\")\nprint(\"=\" * 60)","metadata":{"id":"6c5ec541d431479f9b84c8d4aed6fa8d","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:08:46.753828Z","iopub.execute_input":"2026-09-30T16:08:46.75477Z","iopub.status.idle":"2026-09-30T16:08:47.239361Z","shell.execute_reply.started":"2026-09-30T16:08:46.754739Z","shell.execute_reply":"2026-09-30T16:08:47.238808Z"}},"outputs":[],"execution_count":null},{"id":"5f1031c7629b45bfba1b14987ee9268c","cell_type":"code","source":"# ============================================================\n# SECTION 4.7 — DATASET AUDIT & CROSS-SPLIT LEAKAGE EVIDENCE\n# ============================================================\n\ndef section_1_dataset_audit(df, train_df, val_df, test_df):\n    \"\"\"Save a complete, reproducible dataset and patient-level audit.\"\"\"\n    print(\"\\n\" + \"=\" * 90)\n    print(\"DATASET AUDIT\")\n    print(\"=\" * 90)\n\n    total_images = len(df)\n    class_counts = df[\"label\"].value_counts().sort_index()\n    train_counts = train_df[\"label\"].value_counts().sort_index()\n    val_counts = val_df[\"label\"].value_counts().sort_index()\n    test_counts = test_df[\"label\"].value_counts().sort_index()\n\n    # APTOS has no patient IDs (patient_id = id_code), so this is an image and\n    # duplicate-group overlap check, not a patient-level one.\n    train_images = set(train_df[\"patient_id\"])\n    val_images = set(val_df[\"patient_id\"])\n    test_images = set(test_df[\"patient_id\"])\n    train_groups, val_groups, test_groups = (\n        set(frame[\"duplicate_group\"]) for frame in (train_df, val_df, test_df)\n    )\n    overlaps = {\n        \"train_validation_image_overlap\": len(train_images & val_images),\n        \"train_test_image_overlap\": len(train_images & test_images),\n        \"validation_test_image_overlap\": len(val_images & test_images),\n        \"train_validation_duplicate_group_overlap\": len(train_groups & val_groups),\n        \"train_test_duplicate_group_overlap\": len(train_groups & test_groups),\n        \"validation_test_duplicate_group_overlap\": len(val_groups & test_groups),\n    }\n\n    print(f\"Total dataset images: {total_images:,}\")\n    print(\"\\nClass distribution (overall):\")\n    print(class_counts.to_string())\n    print(\"\\nTrain split distribution:\")\n    print(train_counts.to_string())\n    print(\"\\nValidation split distribution:\")\n    print(val_counts.to_string())\n    print(\"\\nTest split distribution:\")\n    print(test_counts.to_string())\n\n    print(\"\\nImage / duplicate-group overlap check (APTOS has no patient IDs):\")\n    print(f\"Train images: {len(train_images):,} in {len(train_groups):,} duplicate groups\")\n    print(f\"Validation images: {len(val_images):,} in {len(val_groups):,} duplicate groups\")\n    print(f\"Test images: {len(test_images):,} in {len(test_groups):,} duplicate groups\")\n    for metric, value in overlaps.items():\n        print(f\"{metric}: {value:,}\")\n\n    split_frames = {\"train\": train_df, \"validation\": val_df, \"test\": test_df}\n    audit_records = [\n        {\"metric\": \"total_images\", \"value\": total_images},\n        {\"metric\": \"train_images\", \"value\": len(train_df)},\n        {\"metric\": \"validation_images\", \"value\": len(val_df)},\n        {\"metric\": \"test_images\", \"value\": len(test_df)},\n        {\"metric\": \"train_percentage\", \"value\": len(train_df) / total_images * 100},\n        {\"metric\": \"validation_percentage\", \"value\": len(val_df) / total_images * 100},\n        {\"metric\": \"test_percentage\", \"value\": len(test_df) / total_images * 100},\n        {\"metric\": \"train_unique_images\", \"value\": len(train_images)},\n        {\"metric\": \"validation_unique_images\", \"value\": len(val_images)},\n        {\"metric\": \"test_unique_images\", \"value\": len(test_images)},\n        {\"metric\": \"train_duplicate_groups\", \"value\": len(train_groups)},\n        {\"metric\": \"validation_duplicate_groups\", \"value\": len(val_groups)},\n        {\"metric\": \"test_duplicate_groups\", \"value\": len(test_groups)},\n        *[{\"metric\": metric, \"value\": value} for metric, value in overlaps.items()],\n    ]\n\n    for class_id in sorted(class_counts.index):\n        audit_records.append({\n            \"metric\": f\"overall_class_{class_id}_count\",\n            \"value\": int(class_counts.get(class_id, 0)),\n        })\n        for split_name, split_frame in split_frames.items():\n            audit_records.append({\n                \"metric\": f\"{split_name}_class_{class_id}_count\",\n                \"value\": int(split_frame[\"label\"].eq(class_id).sum()),\n            })\n\n    audit_records.append({\n        \"metric\": \"image_and_duplicate_group_overlap_status\",\n        \"value\": \"ZERO_OVERLAP\" if max(overlaps.values()) == 0 else \"OVERLAP_DETECTED\",\n    })\n\n    summary_path = os.path.join(Config.REPORTS_DIR, \"dataset_audit_summary.csv\")\n    pd.DataFrame(audit_records).to_csv(summary_path, index=False)\n    print(f\"\\nDataset audit summary saved to: {summary_path}\")\n    print(\"\\n\" + \"=\" * 90)\n\n\nsection_1_dataset_audit(df_labels, train_df, val_df, test_df)\n","metadata":{"id":"5f1031c7629b45bfba1b14987ee9268c","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:09:11.011074Z","iopub.execute_input":"2026-09-30T16:09:11.011333Z","iopub.status.idle":"2026-09-30T16:09:11.03724Z","shell.execute_reply.started":"2026-09-30T16:09:11.011313Z","shell.execute_reply":"2026-09-30T16:09:11.036615Z"}},"outputs":[],"execution_count":null},{"id":"757533fc","cell_type":"markdown","source":"## Section 5: CNN Architecture & Transfer Learning — EfficientNetB3\n\n### Why EfficientNetB3?\n\nThe table lists published ImageNet classification figures as context for architecture selection. It is not a DR benchmark, so Section 6 trains every alternative backbone (and a CNN from scratch) on this project's own split with the same head, preprocessing and schedule, to check the choice on APTOS 2019.\n\n| Backbone | ImageNet-1K top-1 | Parameters | GFLOPs | Trade-off |\n|---|---:|---:|---:|---|\n| VGG-16 | 71.59% | 138.4M | 15.47 | Simple design, but very large classifier and compute cost |\n| ResNet-50 | 76.13% | 25.6M | 4.09 | Residual learning; higher compute and parameter cost |\n| DenseNet-121 | 74.98% | 8.0M | 2.87 | Dense connections reuse features; widely used in medical imaging |\n| MobileNetV2 | 71.88% | 3.5M | 0.30 | Very efficient, with lower capacity than B3 |\n| EfficientNetB0 | 77.1% | 5.3M | 0.39 | Same family as B3 at lower resolution and capacity |\n| **EfficientNetB3** | **81.6%** | **12.0M** | **1.8** | Compound-scaled balance of depth, width, resolution, and capacity |\n\nEfficientNetB3 was chosen as a practical capacity/compute trade-off for five-stage fundus grading. Published ImageNet figures motivate this choice but do not prove superiority on diabetic-retinopathy images. The ImageNet classifier top is removed; the backbone starts from ImageNet weights.\n\n### Transfer Learning, Classifier Head and Overfitting Controls\n\nAPTOS 2019 leaves only 2,561 training images after the split, so the defaults are chosen to limit overfitting: Phase 1 freezes the backbone and trains the new head; Phase 2 restores the best Phase 1 checkpoint and fine-tunes only the **top 120 of 384** EfficientNetB3 layers (BatchNorm stays frozen) with a low **1e-5** learning rate and **AdamW weight decay 1e-4**. Both phases stop on **validation loss** and keep the lowest-loss weights. The Phase 2 learning rate either drops on a validation-loss plateau (default) or follows a 2-epoch warm-up and cosine decay (`Config.PHASE2_LR_SCHEDULE=\"cosine\"`).\n\nThe head is **global average pooling → BatchNorm → Dense(256, ReLU) → Dropout(0.5) → Dense(5, softmax)**:\n\n- **Global average pooling** aggregates each feature channel over space while avoiding the large parameter count of flattening the feature map.\n- **BatchNorm** stabilizes the feature distribution entering the new classifier head.\n- **Dense(256, ReLU)** learns a compact DR-specific representation; it is also used for case-retrieval embeddings.\n- **Dropout 0.5** regularizes the head during training.\n- **Softmax(5)** returns probabilities for ICDR stages 0–4.\n\nThe input adapter rescales preprocessed float values from [0,1] back to [0,255] because Keras EfficientNet applies its own internal 1/255 rescaling. Each comparison backbone gets its own ImageNet adapter: Caffe-style BGR mean subtraction for ResNet-50 and VGG-16, [-1, 1] scaling for MobileNetV2, Torch-style mean/std normalisation for DenseNet-121, and ×255 for EfficientNetB0.\n\n```mermaid\nflowchart LR\n    A[Fundus image<br/>Config.IMG_SIZE x Config.IMG_SIZE x 3<br/>float values in 0 to 1] --> B[Rescaling x255 adapter]\n    B --> C[EfficientNetB3<br/>ImageNet weights<br/>include_top=False]\n    C --> D[Global average pooling]\n    D --> E[Batch normalization]\n    E --> F[Dense 256, ReLU]\n    F --> G[Dropout Config.DROPOUT_RATE]\n    G --> H[Dense 5, softmax]\n    H --> I[ICDR stage probabilities 0 to 4]\n    J[Phase 1<br/>freeze backbone<br/>Adam, Config.PHASE1_LR] -.-> C\n    K[Phase 2<br/>unfreeze top Config.UNFREEZE_TOP_N<br/>AdamW, Config.PHASE2_LR<br/>BatchNorm stays frozen] -.-> C\n```\n\n### Pilot Trials (Section 6)\n\nWith `RUN_HYPERPARAMETER_TUNING=True`, every pilot trial trains on the **full training split** and is scored on the **full validation split**, for up to **4 Phase 1 + 6 Phase 2 epochs** with validation-loss early stopping (patience 2), so the saved curves are long enough to show overfitting. Each trial saves its history, training curves (accuracy and loss for augmented training, a clean training subset and validation; val_qwk) and validation confusion matrix under `report_images/pilots/<trial_id>/`; all rows go to `hyperparameter_tuning_results.csv` and `model_comparison_table.csv`.\n\n| Pilot trial | Group | Change vs baseline |\n|---|---|---|\n| `baseline_300_d05_lr1e5_top120` | hyperparameter | Baseline: 300 px, batch 16, dropout 0.5, LR 1e-3 / 1e-5 (plateau), top-120 fine-tuning, weight decay 1e-4, balanced class weights |\n| `dropout_040` | hyperparameter | Dropout 0.5 → 0.4 |\n| `unfreeze_top60` | hyperparameter | Fine-tune 120 → 60 layers |\n| `unfreeze_top240` | hyperparameter | Fine-tune 120 → 240 layers |\n| `phase2_lr_3e5` | hyperparameter | Phase 2 LR 1e-5 → 3e-5 |\n| `resolution_224` | hyperparameter | Input 300 → 224 |\n| `resolution_380` | hyperparameter | Input 300 → 380, batch 16 → 8 |\n| `lr_schedule_cosine` | hyperparameter | Phase 2 LR plateau → 2-epoch warm-up + cosine decay |\n| `weight_decay_1e3` | hyperparameter | AdamW weight decay 1e-4 → 1e-3 |\n| `class_weights_off` | balancing | Class weights on → off |\n| `class_weights_sqrt` | balancing | Balanced → sqrt (softer) class weights |\n| `oversampling_instead` | balancing | Class weights → on-the-fly oversampling |\n| `ablation_no_ben_graham` | preprocessing ablation | Ben Graham off (crop + resize only) |\n| `ablation_clahe_instead` | preprocessing ablation | CLAHE instead of Ben Graham |\n| `ablation_all_filters` | preprocessing ablation | Denoise + CLAHE + unsharp edges before Ben Graham |\n| `ablation_no_augmentation` | augmentation ablation | Augmentation off |\n| `backbone_resnet50` | backbone | EfficientNetB3 → ResNet-50 |\n| `backbone_mobilenetv2` | backbone | EfficientNetB3 → MobileNetV2 |\n| `backbone_efficientnetb0` | backbone | EfficientNetB3 → EfficientNetB0 |\n| `backbone_densenet121` | backbone | EfficientNetB3 → DenseNet-121 |\n| `backbone_vgg16` | backbone | EfficientNetB3 → VGG-16 |\n| `scratch_cnn_4block` | scratch baseline | Small 4-block CNN, random initialisation, trained end to end (transfer-learning baseline) |\n\nThat is **22 defined pilot trials**. By default `Config.PILOT_SKIP_TRIALS` skips five of them (`phase2_lr_3e5`, `unfreeze_top240`, `dropout_040`, `ablation_all_filters`, `backbone_vgg16`): they were dropped after an earlier validation-only pilot showed they underperformed the baseline or duplicated other evidence, to fit the compute budget. They stay defined and appear in `hyperparameter_tuning_results.csv` as `skipped_by_config`, so they can be re-enabled. **17 trials run** — 9 selectable (6 hyperparameter, 3 balancing) and 8 comparison trials. The final configuration is selected by the **highest validation QWK (tie-break: lower validation loss)** among the **hyperparameter** and **balancing** rows only. The notebook also prints, without using them for selection, the selectable trial with the highest validation accuracy and the one with the smallest clean-train vs validation accuracy gap, so the report can discuss the trade-offs. Ablation, backbone and scratch rows are evidence: the deployed pipeline (Grad-CAM layer, embeddings, app) is built around EfficientNetB3 with Ben Graham preprocessing, so those rows are reported and discussed instead of silently changing the final architecture. The held-out test split is not used for any selection. Pilot runs are shorter than the final fit, so rankings may not perfectly predict full-run rankings.","metadata":{}},{"id":"caf542af846f478cb16b26ef6dc1d902","cell_type":"code","source":"# ============================================================\n# SECTION 5.1 — EFFICIENTNETB3 CLASSIFIER CONSTRUCTION\n# Recreate the classifier architecture used by the shared application model.\n# ============================================================\n\n_gpu_available = bool(tf.config.list_physical_devices(\"GPU\"))\nkeras.mixed_precision.set_global_policy(\"mixed_float16\" if _gpu_available else \"float32\")\nprint(f\"[Precision] Global policy: {keras.mixed_precision.global_policy().name}\")\n\nbase_model = EfficientNetB3(\n    include_top=False,\n    weights=\"imagenet\",\n    input_shape=(Config.IMG_SIZE, Config.IMG_SIZE, 3),\n)\nbase_model.trainable = False\n\ninputs = keras.Input(shape=(Config.IMG_SIZE, Config.IMG_SIZE, 3))\nx = layers.Rescaling(255.0, name=\"restore_efficientnet_input_range\")(inputs)\nx = base_model(x, training=False)\nx = layers.GlobalAveragePooling2D(name=\"gap\")(x)\nx = layers.BatchNormalization(name=\"head_bn\")(x)\nx = layers.Dense(Config.DENSE_UNITS, activation=\"relu\", name=\"head_dense\")(x)\nx = layers.Dropout(Config.DROPOUT_RATE, name=\"head_dropout\")(x)\noutputs = layers.Dense(\n    Config.NUM_CLASSES,\n    activation=\"softmax\",\n    dtype=\"float32\",\n    name=\"predictions\",\n)(x)\nmodel = keras.Model(inputs=inputs, outputs=outputs, name=\"RetinaGuard_EfficientNetB3\")\n\nprint(\"[Model] EfficientNetB3 initialized with ImageNet weights.\")\nprint(f\"[Model] Classifier head: GAP -> BatchNorm -> Dense({Config.DENSE_UNITS}) -> Dropout({Config.DROPOUT_RATE}) -> Dense({Config.NUM_CLASSES})\")\nprint(\"[Model] Base frozen for Phase 1.\")","metadata":{"id":"caf542af846f478cb16b26ef6dc1d902","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:09:42.146139Z","iopub.execute_input":"2026-09-30T16:09:42.147039Z","iopub.status.idle":"2026-09-30T16:09:44.485205Z","shell.execute_reply.started":"2026-09-30T16:09:42.147006Z","shell.execute_reply":"2026-09-30T16:09:44.484368Z"}},"outputs":[],"execution_count":null},{"id":"b3cd387fa8244798a380488fea538e91","cell_type":"code","source":"# ============================================================\n# SECTION 5.2 — COMPILE MODEL (PHASE 1 CONFIGURATION)\n#\n# Loss: CategoricalCrossentropy\n#   - Labels are one-hot (output of tf.one_hot in the dataset pipeline)\n#   - label_smoothing=0.1 softens the targets slightly (e.g. turns a\n#     hard [0,0,1,0,0] into [0.02, 0.02, 0.92, 0.02, 0.02]).\n#     This prevents the model from becoming over-confident on the\n#     training set and improves calibration on unseen data.\n#\n# Metric: SparseCategoricalAccuracy\n#   - We convert one-hot labels back to sparse integers for the metric\n#     so the accuracy printout is human-readable (not over one-hot vectors)\n#   - QWK is computed post-training in Section 7 (it is not a differentiable\n#     TF metric, so we implement it via sklearn after inference)\n# ============================================================\n\ndef compile_model_phase1(m: keras.Model, jit_compile: Union[bool, str] = \"auto\") -> None:\n    \"\"\"\n    Compile the model for Phase 1 training (frozen base, head only).\n\n    Uses Adam optimiser with Config.PHASE1_LR (1e-3). A higher learning\n    rate is safe in Phase 1 because only the randomly-initialised head\n    parameters are updated — the pretrained base weights are frozen.\n\n    Args:\n        m: The keras.Model to compile (modified in place).\n        jit_compile: Passed to Model.compile. The final model keeps \"auto\" (XLA\n            where supported; compiled once, so the speed-up is worth it). The\n            Section 6.2b pilots pass False, because an XLA-compiled cluster for\n            every pilot model accumulates host RAM across the trial loop.\n    \"\"\"\n    m.compile(\n        optimizer=keras.optimizers.Adam(learning_rate=Config.PHASE1_LR),\n        loss=keras.losses.CategoricalCrossentropy(\n            label_smoothing=Config.LABEL_SMOOTHING,  # soft targets for better calibration\n        ),\n        metrics=[\n            keras.metrics.CategoricalAccuracy(name=\"accuracy\"),\n            keras.metrics.AUC(name=\"auc\", multi_label=False),\n        ],\n        jit_compile=jit_compile,\n    )\n    print(f\"[Compile] Phase 1 — LR={Config.PHASE1_LR}, base frozen.\")\n\n\ncompile_model_phase1(model)","metadata":{"id":"b3cd387fa8244798a380488fea538e91","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:09:49.819249Z","iopub.execute_input":"2026-09-30T16:09:49.819649Z","iopub.status.idle":"2026-09-30T16:09:49.848887Z","shell.execute_reply.started":"2026-09-30T16:09:49.819622Z","shell.execute_reply":"2026-09-30T16:09:49.848039Z"}},"outputs":[],"execution_count":null},{"id":"96b18aa081834dab92f8f327ccaaa8a6","cell_type":"code","source":"# ============================================================\n# SECTION 5.3 — MODEL SUMMARY\n# Print the full model summary: layer names, output shapes,\n# parameter counts, and trainable vs frozen parameter totals.\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"  FULL MODEL SUMMARY\")\nprint(\"=\" * 70)\nmodel.summary(line_length=70, show_trainable=True)\n\n# Separately report base model stats for clarity\ntotal_params     = model.count_params()\ntrainable_params = sum(\n    tf.size(v).numpy() for v in model.trainable_variables\n)\nfrozen_params    = total_params - trainable_params\n\nprint(\"\\n\" + \"-\" * 70)\nprint(f\"  Total parameters     : {total_params:>12,}\")\nprint(f\"  Trainable (Phase 1)  : {trainable_params:>12,}  (head only)\")\nprint(f\"  Frozen (Phase 1)     : {frozen_params:>12,}  (EfficientNetB3 base)\")\nprint(f\"  Base model layers    : {len(base_model.layers):>12,}\")\nprint(f\"  Layers to unfreeze   : {Config.UNFREEZE_TOP_N:>12,}  (Phase 2)\")\nprint(\"-\" * 70)\nprint(f\"  [MULTI-STAGE NOTE] Final Dense layer outputs {Config.NUM_CLASSES} classes\")\nprint(f\"  covering all ICDR stages 0-4 (No DR through Proliferative DR).\")\nprint(f\"  Stage-level output; binary DR / referable-DR results are derived in Section 7.4d.\")\nprint(\"=\" * 70)","metadata":{"id":"96b18aa081834dab92f8f327ccaaa8a6","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:10:21.66094Z","iopub.execute_input":"2026-09-30T16:10:21.661199Z","iopub.status.idle":"2026-09-30T16:10:21.698681Z","shell.execute_reply.started":"2026-09-30T16:10:21.661179Z","shell.execute_reply":"2026-09-30T16:10:21.698065Z"}},"outputs":[],"execution_count":null},{"id":"8a94df4451564b259d218660a344f982","cell_type":"code","source":"# ============================================================\n# SECTION 5.4 — PHASE 2 UNFREEZE CONFIGURATION\n# Unfreeze the backbone while keeping BatchNorm layers frozen.\n# ============================================================\n\ndef configure_phase2(m: keras.Model, base: keras.Model, jit_compile: Union[bool, str] = \"auto\") -> None:\n    \"\"\"Enable fine-tuning for the backbone, keeping BatchNorm stable.\n\n    jit_compile is passed to Model.compile: the final model keeps \"auto\" (XLA\n    where supported), while the Section 6.2b pilots pass False so an XLA\n    cluster is not compiled and kept in host RAM for every pilot model.\n    \"\"\"\n    n_layers = len(base.layers)\n    if not 0 < Config.UNFREEZE_TOP_N <= n_layers:\n        raise ValueError(\n            f\"UNFREEZE_TOP_N must be between 1 and {n_layers}; \"\n            f\"received {Config.UNFREEZE_TOP_N}.\"\n        )\n\n    base.trainable = True\n    freeze_until = n_layers - Config.UNFREEZE_TOP_N\n\n    for i, layer in enumerate(base.layers):\n        if i < freeze_until or isinstance(layer, layers.BatchNormalization):\n            layer.trainable = False\n        else:\n            layer.trainable = True\n\n    newly_trainable = sum(\n        tf.size(v).numpy() for v in m.trainable_variables\n    )\n\n    m.compile(\n        optimizer=keras.optimizers.AdamW(\n            learning_rate=Config.PHASE2_LR,\n            weight_decay=Config.WEIGHT_DECAY,  # decoupled weight decay for fine-tuning\n        ),\n        loss=keras.losses.CategoricalCrossentropy(label_smoothing=Config.LABEL_SMOOTHING),\n        metrics=[\n            keras.metrics.CategoricalAccuracy(name=\"accuracy\"),\n            keras.metrics.AUC(name=\"auc\", multi_label=False),\n        ],\n        jit_compile=jit_compile,\n    )\n\n    print(f\"[Phase 2] Unfroze the top {Config.UNFREEZE_TOP_N} of {n_layers} backbone layers except BatchNorm.\")\n    print(f\"[Phase 2] Trainable parameters now: {newly_trainable:,}\")\n    print(f\"[Phase 2] Learning rate: {Config.PHASE2_LR}\")\n\n\nprint(\"[Section 5] configure_phase2() defined — BatchNorm layers remain frozen.\")\n\nprint(\"\\n\" + \"=\" * 60)\nprint(\"  SECTION 5 COMPLETE — Model Architecture\")\nprint(\"=\" * 60)\nprint(\"  model            : EfficientNetB3 + custom head (5-class)\")\nprint(\"  base_model       : EfficientNetB3 (frozen for Phase 1)\")\nprint(\"  compile_model_phase1() : Phase 1 compilation defined\")\nprint(\"  configure_phase2()     : Phase 2 transition defined\")\nprint(\"  Output           : softmax over 5 ICDR classes (multi-stage)\")\nprint(\"  Next -> Section 6: Training Strategy\")\nprint(\"=\" * 60)","metadata":{"id":"8a94df4451564b259d218660a344f982","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:10:28.442096Z","iopub.execute_input":"2026-09-30T16:10:28.442717Z","iopub.status.idle":"2026-09-30T16:10:28.451857Z","shell.execute_reply.started":"2026-09-30T16:10:28.442687Z","shell.execute_reply":"2026-09-30T16:10:28.451053Z"}},"outputs":[],"execution_count":null},{"id":"c862adf3","cell_type":"markdown","source":"## Section 6: Training Strategy\n\n**Goal:** Select a configuration with validation-only pilot trials, then train the selected model on the complete training split and honestly report train/validation behaviour.\n\n### Two-Phase Final Training\n\n**Phase 1 — Feature extraction:** The ImageNet-pretrained backbone is frozen while the five-class head trains with `Config.PHASE1_LR` for up to `Config.PHASE1_EPOCHS` (8) epochs.\n\n**Phase 2 — Fine-tuning:** The top `Config.UNFREEZE_TOP_N` (120) backbone layers are fine-tuned while BatchNorm remains frozen, using AdamW (`Config.PHASE2_LR` 1e-5, weight decay `Config.WEIGHT_DECAY` 1e-4) for up to `Config.PHASE2_EPOCHS` (25) epochs. With `Config.PHASE2_LR_SCHEDULE=\"plateau\"` (default) the learning rate halves after 2 epochs without a validation-loss improvement; with `\"cosine\"` it warms up linearly over 2 epochs from 10% of the peak, then follows a cosine down to 1% of the peak (no plateau reduction in that mode).\n\nIn both phases **EarlyStopping and ModelCheckpoint monitor validation loss** (mode min, patience 4, best weights restored). Two monitoring-only callbacks are never used for stopping, checkpointing or selection: `ValidationQWK` logs validation QWK as `val_qwk`, and `CleanTrainMetrics` logs `clean_train_loss` and `clean_train_accuracy` on a fixed, stratified subset of 500 training images (seed 42) scored without augmentation or class weights.\n\n### Pilot Selection, Balancing, Ablations and Backbone Comparison\n\nWhen `RUN_HYPERPARAMETER_TUNING=True`, all pilot trials use the full training and validation splits with up to 4 + 6 epochs and validation-loss early stopping (patience 2). Six hyperparameter trials (baseline, `unfreeze_top60`, `resolution_224`, `resolution_380`, `lr_schedule_cosine`, `weight_decay_1e3`) and three balancing trials (`class_weights_off`, `class_weights_sqrt`, `oversampling_instead`) are selectable; the winner (highest validation QWK, lower validation loss as tie-break) is used for the final full-data training run. With `RUN_COMPARISON_TRIALS=True`, the same runner also trains two preprocessing ablations (no Ben Graham, CLAHE instead), an augmentation ablation, four alternative ImageNet backbones (ResNet-50, MobileNetV2, EfficientNetB0, DenseNet-121) and a small CNN trained from scratch as the transfer-learning baseline. Set `RUN_COMPARISON_TRIALS=False` to skip those eight trials if GPU time is short. Pilot results and final full-run histories are kept separately so pilot scores are not presented as final model performance.\n\n**Compute limits.** Five further trials (`phase2_lr_3e5`, `unfreeze_top240`, `dropout_040`, `ablation_all_filters`, `backbone_vgg16`) are listed in `Config.PILOT_SKIP_TRIALS`: they were dropped after an earlier validation-only pilot showed they underperformed the baseline or duplicated other evidence, to fit the compute budget. They remain defined, are recorded as `skipped_by_config` and are left out of selection and of the 6.2c/6.2d tables and plots. `Config.PILOT_TIME_BUDGET_HOURS` (5 h) stops new trials from starting once the budget is used, recording the rest as `skipped_time_budget`; selection then uses the completed selectable trials. `Config.PILOT_TRIAL_LIMIT` runs only the first N non-skipped trials (for a quick memory test). Setting `Config.RESUME_RUN_ID` to an earlier run ID reuses that run folder: completed trials whose split fingerprint matches the current split are not retrained, and the selection is rebuilt from their saved rows with the same rule. Pilot models are compiled without XLA and fully released after each trial (process RAM and GPU memory are logged per trial) so a long pilot run does not exhaust host memory.\n\n### Overfitting Assessment\n\nEvery pilot row records the train–validation accuracy gap and the **clean-train–validation** accuracy gap at the lowest-validation-loss epoch, and how far validation loss rose after its minimum. After the final full-data fit, `overfitting_diagnostics.csv` records the same quantities for each phase. **Why training accuracy can sit below validation accuracy.** Keras' training accuracy is averaged over *augmented* batches (rotations, zoom, brightness/contrast changes), with **dropout** switching off half of the head's units, and — with class weights — a loss that emphasises the hard minority stages. Validation images are scored unaugmented, with dropout off and an unweighted loss, so the \"Train (augmented)\" curve is measured on a harder task and can legitimately lie *below* validation without meaning the model generalises better than it fits. The **\"Train (clean subset)\"** curve removes those differences: 500 fixed training images scored exactly like validation. The gap between the clean curve and validation is therefore the like-for-like overfitting signal; both are plotted next to the original augmented curve, unsmoothed.\n\nRead these metrics alongside the saved learning curves; describe overfitting only if the observed curves support it.\n\nDropout 0.5, partial fine-tuning (top 120 layers), frozen BatchNorm, AdamW weight decay, training-only augmentation, label smoothing and validation-loss early stopping are the implemented regularization controls. QWK threshold calibration and TTA selection use validation predictions only. The held-out test set is not used to select settings, thresholds, or checkpoints and is evaluated only after the final configuration is fixed.","metadata":{}},{"id":"10035df66a5249719c064dabedbc6114","cell_type":"code","source":"# ============================================================\n# SECTION 6.1 — EXPERIMENT LOG TABLE\n#\n# WHY: The coursework requires an \"experiments tried\" record.\n# We maintain a simple DataFrame that logs every training run's\n# hyperparameters and final results. This lets us compare Phase 1\n# vs Phase 2 (and any additional runs) at a glance.\n# ============================================================\n\n# Initialise the experiment log as a global DataFrame\n# Each row = one call to model.fit() with its configuration + results\nexperiment_log = pd.DataFrame(columns=[\n    \"phase\",\n    \"learning_rate\",\n    \"batch_size\",\n    \"epochs_planned\",\n    \"epochs_run\",        # actual epochs before early stopping\n    \"val_accuracy\",\n    \"val_loss\",\n    \"val_auc\",\n    \"val_qwk\",\n    \"notes\",\n])\n\n\ndef log_experiment(\n    phase: str,\n    lr: float,\n    batch_size: int,\n    epochs_planned: int,\n    history: keras.callbacks.History,\n    notes: str = \"\",\n) -> None:\n    \"\"\"\n    Append a completed training run's results to the experiment log.\n\n    Args:\n        phase:          Human-readable phase name (e.g. \"Phase 1 - Frozen Base\").\n        lr:             Learning rate used in this run.\n        batch_size:     Batch size used.\n        epochs_planned: Maximum epochs configured (before early stopping).\n        history:        The History object returned by model.fit().\n        notes:          Optional free-text notes about this run.\n    \"\"\"\n    global experiment_log\n\n    epochs_run   = len(history.history[\"val_loss\"])\n    best_val_acc = max(history.history.get(\"val_accuracy\", [0]))\n    best_val_loss = min(history.history[\"val_loss\"])\n    best_val_auc  = max(history.history.get(\"val_auc\", [0]))\n    best_val_qwk  = max(history.history.get(\"val_qwk\", [float(\"nan\")]))\n\n    new_row = pd.DataFrame([{\n        \"phase\":           phase,\n        \"learning_rate\":   lr,\n        \"batch_size\":      batch_size,\n        \"epochs_planned\":  epochs_planned,\n        \"epochs_run\":      epochs_run,\n        \"val_accuracy\":    round(best_val_acc,  4),\n        \"val_loss\":        round(best_val_loss, 4),\n        \"val_auc\":         round(best_val_auc,  4),\n        \"val_qwk\":         round(best_val_qwk,  4),\n        \"notes\":           notes,\n    }])\n\n    experiment_log = pd.concat([experiment_log, new_row], ignore_index=True)\n    print(f\"[Experiment Log] Logged: {phase}\")\n    print(experiment_log.to_string(index=False))\n\n\nprint(\"[Section 6] Experiment log initialised.\")","metadata":{"id":"10035df66a5249719c064dabedbc6114","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T16:10:41.067971Z","iopub.execute_input":"2026-09-30T16:10:41.068225Z","iopub.status.idle":"2026-09-30T16:10:41.07792Z","shell.execute_reply.started":"2026-09-30T16:10:41.068205Z","shell.execute_reply":"2026-09-30T16:10:41.076938Z"}},"outputs":[],"execution_count":null},{"id":"dae9d76dfc1f42f09e89fc286c4ddabb","cell_type":"code","source":"# ============================================================\n# SECTION 6.2 — CALLBACKS\n#\n# Final-fit stopping/checkpointing rule (run B):\n# EarlyStopping and ModelCheckpoint monitor validation QWK (mode max) and\n# restore the highest-QWK weights after STOP_PATIENCE epochs without\n# improvement. Run A monitored validation loss; under balanced class weights\n# and label smoothing, validation loss and QWK disagreed in both final phases\n# (loss chose epoch 1 of Phase 1 and epoch 3 of Phase 2 while QWK/accuracy\n# were still rising). This change was decided from VALIDATION evidence only,\n# before looking at run A's test results, and run B is the reported model.\n# Learning-rate control depends on the schedule:\n#   \"plateau\" -> ReduceLROnPlateau halves the LR after Config.RLROP_PATIENCE\n#                epochs without a validation-loss improvement (unchanged).\n#   \"cosine\"  -> LearningRateScheduler: linear warm-up over\n#                Config.LR_WARMUP_EPOCHS epochs from 10% of the peak LR,\n#                then cosine decay to 1% of the peak by the last epoch.\n#                (No ReduceLROnPlateau in this mode.)\n# ValidationQWK is first in the callback list, so val_qwk is in the epoch\n# logs before EarlyStopping / ModelCheckpoint read it.\n# Monitoring only:\n#   CleanTrainMetrics -> clean_train_loss / clean_train_accuracy on a fixed,\n#                        stratified 500-image TRAINING subset, evaluated\n#                        without augmentation, dropout or class weights.\n# Note: the Section 6.2b pilots use their own EarlyStopping (val_loss) and\n# are unaffected by this rule.\n# ============================================================\n\nimport math\n\nSTOP_MONITOR = \"val_qwk\"   # final-fit EarlyStopping / ModelCheckpoint metric (run B)\nSTOP_MODE = \"max\"\nSTOP_PATIENCE = 6          # QWK on 551 images is noisy, so allow a few more epochs than the val_loss rule (4)\n\n\nclass LearningRateLogger(keras.callbacks.Callback):\n    \"\"\"Record the optimizer learning rate each epoch so history CSVs always contain it.\"\"\"\n\n    def on_epoch_end(self, epoch, logs=None):\n        if logs is not None and \"learning_rate\" not in logs:\n            logs[\"learning_rate\"] = float(np.array(self.model.optimizer.learning_rate))\n\n\nclass ValidationQWK(keras.callbacks.Callback):\n    \"\"\"Log validation quadratic weighted kappa as 'val_qwk' after every epoch (used by the final-fit stopping rule).\"\"\"\n\n    def __init__(self, validation_dataset: tf.data.Dataset, validation_labels: np.ndarray):\n        super().__init__()\n        self.validation_dataset = validation_dataset\n        self.validation_labels = np.asarray(validation_labels, dtype=int)\n\n    def on_epoch_end(self, epoch, logs=None):\n        probabilities = self.model.predict(self.validation_dataset, verbose=0)\n        predictions = np.argmax(probabilities, axis=1)\n        if len(predictions) != len(self.validation_labels):\n            raise ValueError(\"ValidationQWK: prediction count does not match validation labels.\")\n        qwk = float(cohen_kappa_score(self.validation_labels, predictions, weights=\"quadratic\"))\n        if logs is not None:\n            logs[\"val_qwk\"] = qwk\n        print(f\" — val_qwk: {qwk:.4f}\")\n\n\nclass CleanTrainMetrics(keras.callbacks.Callback):\n    \"\"\"\n    Log clean_train_loss / clean_train_accuracy on a fixed training subset (monitoring only).\n\n    Keras' own training accuracy is measured on augmented batches, with dropout\n    active and (optionally) class-weighted loss, so it is not comparable with\n    validation. This callback re-scores a fixed stratified subset of TRAINING\n    images in inference mode, without augmentation or class weights, using the\n    same label-smoothed loss as val_loss — a like-for-like train/validation view.\n    \"\"\"\n\n    def __init__(self, clean_dataset: tf.data.Dataset, clean_labels: np.ndarray):\n        super().__init__()\n        self.clean_dataset = clean_dataset\n        self.clean_labels = np.asarray(clean_labels, dtype=int)\n        self.loss_fn = keras.losses.CategoricalCrossentropy(label_smoothing=Config.LABEL_SMOOTHING)\n\n    def on_epoch_end(self, epoch, logs=None):\n        probabilities = self.model.predict(self.clean_dataset, verbose=0)\n        if len(probabilities) != len(self.clean_labels):\n            raise ValueError(\"CleanTrainMetrics: prediction count does not match subset labels.\")\n        one_hot = tf.one_hot(self.clean_labels, Config.NUM_CLASSES)\n        clean_loss = float(self.loss_fn(one_hot, probabilities))\n        clean_accuracy = float(np.mean(np.argmax(probabilities, axis=1) == self.clean_labels))\n        if logs is not None:\n            logs[\"clean_train_loss\"] = clean_loss\n            logs[\"clean_train_accuracy\"] = clean_accuracy\n        print(f\" — clean_train_accuracy: {clean_accuracy:.4f} — clean_train_loss: {clean_loss:.4f}\")\n\n\ndef make_clean_train_subset(\n    train_frame: pd.DataFrame,\n    n_images: int = Config.CLEAN_TRAIN_SUBSET_SIZE,\n    seed: int = Config.SEED,\n) -> pd.DataFrame:\n    \"\"\"Fixed, stratified subset of the training split used by CleanTrainMetrics.\"\"\"\n    if len(train_frame) <= n_images:\n        return train_frame.reset_index(drop=True)\n    subset, _ = train_test_split(\n        train_frame, train_size=n_images, stratify=train_frame[\"label\"], random_state=seed\n    )\n    return subset.sort_index().reset_index(drop=True)\n\n\nCLEAN_TRAIN_SUBSET_DF = make_clean_train_subset(train_df)\nassert set(CLEAN_TRAIN_SUBSET_DF[\"split\"]) == {\"train\"}, \"The clean-train subset must come from the training split.\"\n# Legend explanation shared by the final-fit curves (6.4) and pilot evidence plots (6.2b);\n# the subset size is taken from the subset actually built.\nCURVE_LEGEND_NOTE = (\n    \"augmented = training batches with augmentation, dropout and class weights; \"\n    f\"clean subset = {len(CLEAN_TRAIN_SUBSET_DF):,} fixed training images scored like validation\"\n)\nprint(f\"[Callbacks] Clean-train subset: {len(CLEAN_TRAIN_SUBSET_DF):,} training images \"\n      f\"{CLEAN_TRAIN_SUBSET_DF['label'].value_counts().sort_index().to_dict()} (seed {Config.SEED})\")\n\n\ndef warmup_cosine_lr(epoch: int, total_epochs: int, peak_lr: float,\n                     warmup_epochs: int = Config.LR_WARMUP_EPOCHS) -> float:\n    \"\"\"\n    Learning rate for one epoch of the warm-up + cosine schedule.\n\n    Epochs 0..warmup-1 rise linearly from 10% of peak_lr towards peak_lr;\n    the remaining epochs follow a cosine from peak_lr down to 1% of peak_lr\n    at the final epoch.\n    \"\"\"\n    warmup_epochs = max(0, min(warmup_epochs, total_epochs - 1))\n    if epoch < warmup_epochs:\n        return float(peak_lr * (0.1 + 0.9 * epoch / warmup_epochs))\n    decay_epochs = max(1, total_epochs - warmup_epochs - 1)\n    progress = min(1.0, (epoch - warmup_epochs) / decay_epochs)\n    minimum_lr = 0.01 * peak_lr\n    return float(minimum_lr + 0.5 * (peak_lr - minimum_lr) * (1.0 + math.cos(math.pi * progress)))\n\n\ndef lr_schedule_callbacks(schedule: str, total_epochs: int, peak_lr: float) -> list:\n    \"\"\"Return the learning-rate callback for 'plateau' (ReduceLROnPlateau) or 'cosine' (warm-up + cosine).\"\"\"\n    if schedule == \"cosine\":\n        return [keras.callbacks.LearningRateScheduler(\n            lambda epoch, lr: warmup_cosine_lr(epoch, total_epochs, peak_lr), verbose=0\n        )]\n    if schedule == \"plateau\":\n        return [ReduceLROnPlateau(\n            monitor=\"val_loss\",\n            factor=Config.RLROP_FACTOR,\n            patience=Config.RLROP_PATIENCE,\n            min_lr=1e-7,\n            verbose=1,\n            mode=\"min\",\n        )]\n    raise ValueError(f\"Unknown learning-rate schedule: {schedule!r} (use 'plateau' or 'cosine').\")\n\n\ndef clean_train_callback(image_size: int, preprocessing: Optional[Dict[str, bool]] = None) -> CleanTrainMetrics:\n    \"\"\"CleanTrainMetrics on the fixed training subset, built at the current image size and preprocessing.\"\"\"\n    clean_dataset = _build_tuning_dataset(\n        CLEAN_TRAIN_SUBSET_DF, image_size, shuffle=False, augment=False, preprocessing=preprocessing,\n    )\n    return CleanTrainMetrics(clean_dataset, CLEAN_TRAIN_SUBSET_DF[\"label\"].to_numpy())\n\n\ndef build_callbacks(\n    phase_name: str,\n    validation_dataset: Optional[tf.data.Dataset] = None,\n    validation_labels: Optional[np.ndarray] = None,\n    lr_schedule: str = \"plateau\",\n    total_epochs: Optional[int] = None,\n    peak_lr: Optional[float] = None,\n) -> list:\n    \"\"\"Build val_qwk-driven stopping/checkpointing, the LR schedule, and monitoring callbacks.\"\"\"\n    os.makedirs(Config.CHECKPOINT_DIR, exist_ok=True)\n    checkpoint_path = os.path.join(\n        Config.CHECKPOINT_DIR,\n        f\"best_{phase_name}.weights.h5\",\n    )\n    validation_dataset = val_dataset if validation_dataset is None else validation_dataset\n    validation_labels = val_df[\"label\"].to_numpy() if validation_labels is None else validation_labels\n\n    callbacks = [\n        ValidationQWK(validation_dataset, validation_labels),  # must stay first: provides val_qwk\n        clean_train_callback(Config.IMG_SIZE),                 # monitoring only\n        LearningRateLogger(),  # logs the LR used for this epoch\n        EarlyStopping(\n            monitor=STOP_MONITOR,\n            patience=STOP_PATIENCE,\n            restore_best_weights=True,\n            verbose=1,\n            mode=STOP_MODE,\n        ),\n        *lr_schedule_callbacks(lr_schedule, total_epochs or Config.PHASE2_EPOCHS, peak_lr or Config.PHASE2_LR),\n        ModelCheckpoint(\n            filepath=checkpoint_path,\n            monitor=STOP_MONITOR,\n            save_best_only=True,\n            save_weights_only=True,\n            verbose=1,\n            mode=STOP_MODE,\n        ),\n    ]\n\n    print(f\"[Callbacks] Built for {phase_name}:\")\n    print(f\"  EarlyStopping     : patience={STOP_PATIENCE} ({STOP_MONITOR}, mode {STOP_MODE}, restore best weights)\")\n    if lr_schedule == \"cosine\":\n        print(f\"  LR schedule       : warm-up {Config.LR_WARMUP_EPOCHS} epochs + cosine to 1% of {peak_lr or Config.PHASE2_LR}\")\n    else:\n        print(f\"  ReduceLROnPlateau : factor={Config.RLROP_FACTOR}, patience={Config.RLROP_PATIENCE} (val_loss)\")\n    print(f\"  ModelCheckpoint   : -> {checkpoint_path} (highest {STOP_MONITOR})\")\n    print(\"  ValidationQWK     : val_qwk logged each epoch (drives stopping and checkpointing)\")\n    print(f\"  CleanTrainMetrics : clean_train_loss/accuracy on {len(CLEAN_TRAIN_SUBSET_DF)} training images (monitoring only)\")\n    print(\"  LearningRateLogger: learning_rate saved in the training history\")\n    return callbacks","metadata":{"id":"dae9d76dfc1f42f09e89fc286c4ddabb","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:00:50.689407Z","iopub.execute_input":"2026-09-30T20:00:50.690387Z","iopub.status.idle":"2026-09-30T20:00:50.719938Z","shell.execute_reply.started":"2026-09-30T20:00:50.690347Z","shell.execute_reply":"2026-09-30T20:00:50.719108Z"}},"outputs":[],"execution_count":null},{"id":"4aba2c52","cell_type":"code","source":"# Fingerprint the fixed manifests so every tuning row proves it used the same splits.\n_tuning_manifest_rows = pd.concat(\n    [\n        frame[[\"patient_id\", \"duplicate_group\", \"label\"]].assign(split=split_name)\n        for split_name, frame in (\n            (\"train\", train_df),\n            (\"validation\", val_df),\n            (\"test\", test_df),\n        )\n    ],\n    ignore_index=True,\n)\n_tuning_manifest_rows = _tuning_manifest_rows.astype(\"string\").fillna(\"\").sort_values(\n    [\"split\", \"patient_id\", \"duplicate_group\", \"label\"]\n)\n_split_fingerprint = hashlib.sha256(\n    _tuning_manifest_rows.to_csv(index=False, lineterminator=\"\\n\").encode(\"utf-8\")\n).hexdigest()\nprint(\"[Tuning] Fixed manifest fingerprint:\", _split_fingerprint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:01:02.315066Z","iopub.execute_input":"2026-09-30T20:01:02.315463Z","iopub.status.idle":"2026-09-30T20:01:02.34653Z","shell.execute_reply.started":"2026-09-30T20:01:02.315435Z","shell.execute_reply":"2026-09-30T20:01:02.34594Z"}},"outputs":[],"execution_count":null},{"id":"d9a3ba93","cell_type":"markdown","source":"### Pilot Trials and Final Fit\n\nThe pilot runner trains every configuration listed in Section 5 on the **full** training split and scores it on the **full** validation split (up to 4 Phase 1 + 6 Phase 2 epochs, validation-loss early stopping with patience 2). For each trial it saves `pilots/<trial_id>/history.csv`, a training-curve PNG (augmented-train, clean-train-subset and validation lines; val_qwk) and a validation confusion-matrix PNG, and records trainable parameters, epochs run, the train–validation and clean-train–validation accuracy gaps at the best epoch, the validation-loss rise after its minimum and the training time. It then selects among the completed hyperparameter and balancing trials (9 with the default skip list) by validation QWK, breaking ties with lower validation loss, and additionally prints (for discussion only) the trial with the highest validation accuracy and the one with the smallest clean-train gap. These are pilot results, not final model metrics.\n\nAfter selection, the notebook rebuilds the datasets at the winning resolution and batch size (and with oversampling, the chosen class-weight mode, LR schedule and weight decay if that trial won) and runs the normal Phase 1 and Phase 2 epoch limits on the complete training split. The resulting `best_phase1.weights.h5` and `best_phase2.weights.h5` are the final model checkpoints. The held-out test split is excluded from all of this and is evaluated only after this final fit. Set `RUN_HYPERPARAMETER_TUNING=False` to skip the pilots and use the configured single-run settings.","metadata":{}},{"id":"0f1960dd","cell_type":"markdown","source":"Each pilot output row also includes the reason for the change and the hypothesis being tested. These are experiment rationales, not assumptions that a setting will improve performance; the validation metrics determine the observed result.","metadata":{}},{"id":"a4d447e7","cell_type":"code","source":"# ============================================================\n# SECTION 6.2b — VALIDATION-ONLY PILOT TRIALS: HYPERPARAMETERS,\n#                BALANCING, ABLATIONS AND BACKBONE COMPARISON\n#\n# Every trial trains on the FULL training split and is scored on the FULL\n# validation split (TUNING_*_FRACTION = 1.0), for up to 4 Phase 1 + 6\n# Phase 2 epochs with val_loss early stopping (patience 2), so the curves\n# are long enough to show overfitting and the rows are directly comparable.\n#\n# Trial groups:\n#   hyperparameter          -> selectable; tunes the final EfficientNetB3\n#   balancing               -> selectable; class weights vs none vs oversampling\n#   preprocessing_ablation  -> evidence only; is Ben Graham actually helping?\n#   augmentation_ablation   -> evidence only; is augmentation actually helping?\n#   backbone                -> evidence only; other ImageNet backbones\n#   scratch_baseline        -> evidence only; small CNN with no pretraining\n#\n# Selection: highest validation QWK (tie-break: lower validation loss)\n# among the hyperparameter and balancing groups only. The deployed pipeline\n# (Grad-CAM layer, embeddings, Gradio app) is built around EfficientNetB3\n# with Ben Graham preprocessing, so ablation/backbone rows are evidence.\n# The held-out test split is never used to train, stop or rank any trial.\n#\n# Proof saved per trial in report_images/pilots/<trial_id>/:\n#   history.csv, training_curves.png (accuracy and loss with augmented-train,\n#   clean-train-subset and validation lines; val_qwk) and\n#   validation_confusion_matrix.png.\n# Trial keys batch_size, weight_decay, phase2_lr_schedule and\n# class_weight_mode are applied to Config, so a selected trial (e.g.\n# resolution_380 with batch 8) carries through to the final fit.\n#\n# Compute limits (Config, Section 1.2):\n#   PILOT_SKIP_TRIALS       -> kept defined but not trained (\"skipped_by_config\")\n#   PILOT_TRIAL_LIMIT       -> only the first N non-skipped trials (\"skipped_trial_limit\")\n#   PILOT_TIME_BUDGET_HOURS -> no new trial starts after the budget (\"skipped_time_budget\")\n#   RESUME_RUN_ID           -> reuse completed rows of hyperparameter_tuning_results.csv\n#                              whose split_fingerprint matches the current split\n# Only rows with status \"completed\" are selectable or shown in 6.2c/6.2d.\n# After every trial the model, datasets, callbacks, augmentation layer,\n# histories and predictions are deleted, Keras is cleared and process RAM\n# and GPU memory are recorded (ram_gb_after_trial, gpu_mem_gb_after_trial).\n# ============================================================\n\nimport gc\nimport time\nimport inspect\nimport psutil\nimport hashlib\nfrom concurrent.futures import ThreadPoolExecutor\nfrom pathlib import Path\n\nRUN_HYPERPARAMETER_TUNING = True\nRUN_COMPARISON_TRIALS = True   # set False to skip the ablation, backbone and scratch rows\nTUNING_TRAIN_FRACTION = 1.0\nTUNING_VALIDATION_FRACTION = 1.0\nTUNING_PHASE1_EPOCHS = 4\nTUNING_PHASE2_EPOCHS = 6\nPILOT_ES_PATIENCE = 2\nSELECTABLE_GROUPS = {\"hyperparameter\", \"balancing\"}\nPILOT_DIR = Path(Config.REPORTS_DIR) / \"pilots\"\n\n# Baseline settings shared by every trial unless a trial overrides them\n# (matches the Config defaults in Section 1.2).\nPILOT_BASELINE = {\n    \"image_size\": 300,\n    \"dropout_rate\": 0.5,\n    \"phase1_lr\": 1e-3,\n    \"phase2_lr\": 1e-5,\n    \"unfreeze_top_n\": 120,\n    \"backbone\": \"efficientnetb3\",\n    \"preprocessing\": {},        # {} = default pipeline: crop -> resize -> Ben Graham\n    \"augment\": True,\n    \"class_weight_mode\": \"balanced\",   # \"balanced\", \"sqrt\" or \"none\"\n    \"oversample\": False,\n    \"batch_size\": 16,\n    \"weight_decay\": 1e-4,\n    \"phase2_lr_schedule\": \"plateau\",   # \"plateau\" or \"cosine\"\n}\n\n\ndef _pilot_trial(trial_id: str, group: str, change: str, rationale: str, **overrides) -> Dict[str, Any]:\n    \"\"\"Build one pilot trial definition from the baseline plus overrides.\"\"\"\n    trial = {**PILOT_BASELINE, **overrides}\n    trial[\"class_weights\"] = trial[\"class_weight_mode\"] != \"none\"\n    check_balancing_choice(trial[\"class_weights\"], trial[\"oversample\"])\n    trial.update({\n        \"trial_id\": trial_id,\n        \"group\": group,\n        \"change_vs_baseline\": change,\n        \"trial_rationale\": rationale,\n    })\n    return trial\n\n\nTUNING_TRIALS = [\n    _pilot_trial(\n        \"baseline_300_d05_lr1e5_top120\", \"hyperparameter\", \"Baseline\",\n        \"Reference configuration: dropout 0.5, top-120 fine-tuning, Phase 2 LR 1e-5 with plateau \"\n        \"reduction, balanced class weights, weight decay 1e-4, batch 16.\",\n    ),\n    _pilot_trial(\n        \"dropout_040\", \"hyperparameter\", \"Dropout 0.5 -> 0.4\",\n        \"Test whether lighter head regularization learns more without widening the train-validation gap.\",\n        dropout_rate=0.4,\n    ),\n    _pilot_trial(\n        \"unfreeze_top60\", \"hyperparameter\", \"Fine-tune 120 -> 60 layers\",\n        \"Test whether adapting fewer late layers reduces overfitting on about 2,560 training images.\",\n        unfreeze_top_n=60,\n    ),\n    _pilot_trial(\n        \"unfreeze_top240\", \"hyperparameter\", \"Fine-tune 120 -> 240 layers\",\n        \"Test whether adapting more of the backbone helps despite the higher overfitting risk.\",\n        unfreeze_top_n=240,\n    ),\n    _pilot_trial(\n        \"phase2_lr_3e5\", \"hyperparameter\", \"Phase 2 LR 1e-5 -> 3e-5\",\n        \"Test whether a larger fine-tuning step adapts faster or destabilises the pretrained features.\",\n        phase2_lr=3e-5,\n    ),\n    _pilot_trial(\n        \"resolution_224\", \"hyperparameter\", \"Input resolution 300 -> 224\",\n        \"Test whether lower compute and memory cost can retain useful grading performance; \"\n        \"smaller inputs may lose fine retinal detail.\",\n        image_size=224,\n    ),\n    _pilot_trial(\n        \"resolution_380\", \"hyperparameter\", \"Input resolution 300 -> 380 (batch 16 -> 8)\",\n        \"Test whether finer retinal detail (microaneurysms) helps; batch size halved to fit GPU memory.\",\n        image_size=380, batch_size=8,\n    ),\n    _pilot_trial(\n        \"lr_schedule_cosine\", \"hyperparameter\", \"Phase 2 LR plateau -> warm-up + cosine\",\n        \"Test whether a short warm-up followed by smooth cosine decay fine-tunes more stably than \"\n        \"step reductions on a validation-loss plateau.\",\n        phase2_lr_schedule=\"cosine\",\n    ),\n    _pilot_trial(\n        \"weight_decay_1e3\", \"hyperparameter\", \"AdamW weight decay 1e-4 -> 1e-3\",\n        \"Test whether stronger weight decay in Phase 2 narrows the train-validation gap.\",\n        weight_decay=1e-3,\n    ),\n    _pilot_trial(\n        \"class_weights_off\", \"balancing\", \"Class weights ON -> OFF\",\n        \"Measure how much inverse-frequency weighting changes minority-stage recall and overall QWK.\",\n        class_weight_mode=\"none\",\n    ),\n    _pilot_trial(\n        \"class_weights_sqrt\", \"balancing\", \"Balanced -> sqrt (softer) class weights\",\n        \"Test whether a softer correction keeps minority-stage recall without over-weighting the \"\n        \"rare Severe images.\",\n        class_weight_mode=\"sqrt\",\n    ),\n    _pilot_trial(\n        \"oversampling_instead\", \"balancing\", \"Class weights -> on-the-fly oversampling\",\n        \"Balance batches by sampling each stage equally instead of re-weighting the loss.\",\n        class_weight_mode=\"none\", oversample=True,\n    ),\n]\n\nCOMPARISON_TRIALS = [\n    _pilot_trial(\n        \"ablation_no_ben_graham\", \"preprocessing_ablation\", \"Ben Graham ON -> OFF (crop + resize only)\",\n        \"Measure the model-level contribution of Ben Graham illumination normalisation.\",\n        preprocessing={\"apply_ben_graham\": False},\n    ),\n    _pilot_trial(\n        \"ablation_clahe_instead\", \"preprocessing_ablation\", \"Ben Graham -> CLAHE (LAB L-channel)\",\n        \"Compare histogram-based local contrast (CLAHE) against Gaussian background subtraction.\",\n        preprocessing={\"apply_ben_graham\": False, \"apply_clahe_enhancement\": True},\n    ),\n    _pilot_trial(\n        \"ablation_all_filters\", \"preprocessing_ablation\", \"Add bilateral denoise + CLAHE + unsharp edges before Ben Graham\",\n        \"Test whether explicit noise removal and edge enhancement add information beyond Ben Graham.\",\n        preprocessing={\"apply_denoise\": True, \"apply_clahe_enhancement\": True, \"apply_edge_enhancement\": True},\n    ),\n    _pilot_trial(\n        \"ablation_no_augmentation\", \"augmentation_ablation\", \"Augmentation ON -> OFF\",\n        \"Measure the contribution of training-only geometric and photometric augmentation.\",\n        augment=False,\n    ),\n    _pilot_trial(\n        \"backbone_resnet50\", \"backbone\", \"EfficientNetB3 -> ResNet50\",\n        \"Residual backbone with more parameters and FLOPs; tests whether the EfficientNet choice is justified.\",\n        backbone=\"resnet50\",\n    ),\n    _pilot_trial(\n        \"backbone_mobilenetv2\", \"backbone\", \"EfficientNetB3 -> MobileNetV2\",\n        \"Lightweight mobile backbone; quantifies the accuracy cost of a deployable low-compute model.\",\n        backbone=\"mobilenetv2\",\n    ),\n    _pilot_trial(\n        \"backbone_efficientnetb0\", \"backbone\", \"EfficientNetB3 -> EfficientNetB0\",\n        \"Smaller member of the same family; tests whether B3's extra capacity is needed on 3.6k images.\",\n        backbone=\"efficientnetb0\",\n    ),\n    _pilot_trial(\n        \"backbone_densenet121\", \"backbone\", \"EfficientNetB3 -> DenseNet121\",\n        \"Densely connected backbone widely used in medical imaging.\",\n        backbone=\"densenet121\",\n    ),\n    _pilot_trial(\n        \"backbone_vgg16\", \"backbone\", \"EfficientNetB3 -> VGG16\",\n        \"Classic plain convolutional backbone; a large, older architecture for reference.\",\n        backbone=\"vgg16\",\n    ),\n    _pilot_trial(\n        \"scratch_cnn_4block\", \"scratch_baseline\", \"EfficientNetB3 (ImageNet) -> small CNN from scratch\",\n        \"Transfer-learning baseline: 4 conv blocks with random initialisation, trained end to end for the \"\n        \"same total epoch budget, to show what ImageNet pretraining contributes.\",\n        backbone=\"scratch_cnn\",\n    ),\n]\n\n\ndef _preprocessing_tag(preprocessing: Optional[Dict[str, bool]]) -> str:\n    \"\"\"Cache-folder suffix for non-default preprocessing (empty for the default pipeline).\"\"\"\n    if not preprocessing:\n        return \"\"\n    return \"_\" + \"_\".join(\n        f\"{key.replace('apply_', '')}-{int(bool(value))}\"\n        for key, value in sorted(preprocessing.items())\n    )\n\n\ndef _tuning_cache_path(filepath: str, image_size: int, preprocessing: Optional[Dict[str, bool]] = None) -> Path:\n    cache_dir = Path(_cache_root) / \"preproc_cache\" / f\"{image_size}{_preprocessing_tag(preprocessing)}\"\n    return cache_dir / (hashlib.md5(str(filepath).encode(\"utf-8\")).hexdigest() + \".png\")\n\n\ndef _cache_tuning_images(\n    filepaths: List[str],\n    image_size: int,\n    preprocessing: Optional[Dict[str, bool]] = None,\n) -> List[str]:\n    \"\"\"Preprocess each image once per (size, preprocessing variant) and cache it as PNG.\"\"\"\n    preprocessing = preprocessing or {}\n    cache_dir = _tuning_cache_path(\"x\", image_size, preprocessing).parent\n    cache_dir.mkdir(parents=True, exist_ok=True)\n\n    def _cache_one(filepath: str) -> str:\n        output_path = _tuning_cache_path(filepath, image_size, preprocessing)\n        if not output_path.is_file():\n            image = preprocess_image(filepath, img_size=image_size, **preprocessing)\n            image_uint8 = np.clip(image * 255.0, 0, 255).round().astype(np.uint8)\n            written = cv2.imwrite(\n                str(output_path),\n                cv2.cvtColor(image_uint8, cv2.COLOR_RGB2BGR),\n                [cv2.IMWRITE_PNG_COMPRESSION, 1],\n            )\n            if not written:\n                raise OSError(f\"Could not write cache file {output_path} (check disk space and path length).\")\n        return str(output_path)\n\n    workers = min(os.cpu_count() or 1, 8)\n    with ThreadPoolExecutor(max_workers=workers) as executor:\n        return list(\n            tqdm(\n                executor.map(_cache_one, filepaths),\n                total=len(filepaths),\n                desc=f\"Caching {image_size}px{_preprocessing_tag(preprocessing)} images\",\n            )\n        )\n\n\ndef _build_tuning_dataset(\n    dataframe: pd.DataFrame,\n    image_size: int,\n    shuffle: bool,\n    augment: bool,\n    preprocessing: Optional[Dict[str, bool]] = None,\n    oversample: bool = False,\n) -> tf.data.Dataset:\n    \"\"\"Batched dataset from cached PNGs; oversample=True gives an infinite class-balanced stream.\"\"\"\n    cached_paths = _cache_tuning_images(dataframe[\"filepath\"].tolist(), image_size, preprocessing)\n    labels = dataframe[\"label\"].to_numpy(dtype=np.int32)\n\n    def _load_cached_image(path: tf.Tensor, label: tf.Tensor) -> tuple:\n        image = tf.io.decode_png(tf.io.read_file(path), channels=3)\n        image = tf.image.convert_image_dtype(image, tf.float32)\n        image.set_shape([image_size, image_size, 3])\n        label = tf.one_hot(tf.cast(label, tf.int32), Config.NUM_CLASSES)\n        return image, label\n\n    if oversample:\n        dataset = build_oversampled_dataset(cached_paths, labels, _load_cached_image)\n    else:\n        dataset = tf.data.Dataset.from_tensor_slices((cached_paths, labels))\n        if shuffle:\n            dataset = dataset.shuffle(\n                buffer_size=len(cached_paths),\n                seed=Config.SEED,\n                reshuffle_each_iteration=True,\n            )\n        dataset = dataset.map(_load_cached_image, num_parallel_calls=tf.data.AUTOTUNE)\n    if augment:\n        dataset = dataset.map(augment_fn, num_parallel_calls=tf.data.AUTOTUNE)\n    return dataset.batch(Config.BATCH_SIZE).prefetch(tf.data.AUTOTUNE)\n\n\ndef _stratified_pilot_sample(\n    dataframe: pd.DataFrame,\n    fraction: float,\n    seed: int,\n) -> pd.DataFrame:\n    if not 0 < fraction <= 1:\n        raise ValueError(\"Pilot sample fractions must be in (0, 1].\")\n    if fraction == 1:\n        return dataframe.reset_index(drop=True)\n    sampled_groups = []\n    for label, group in dataframe.groupby(\"label\", sort=True):\n        sample_count = min(len(group), max(1, int(np.ceil(len(group) * fraction))))\n        sampled_groups.append(group.sample(sample_count, random_state=seed + int(label)))\n    return pd.concat(sampled_groups).sort_index().reset_index(drop=True)\n\n\n# Backbone registry. Every backbone receives images in [0, 1]; the adapter\n# converts them to the input convention that backbone was pretrained with.\n_BACKBONE_CONSTRUCTORS = {\n    \"efficientnetb3\": keras.applications.EfficientNetB3,   # expects [0, 255], rescales internally\n    \"efficientnetb0\": keras.applications.EfficientNetB0,   # expects [0, 255], rescales internally\n    \"resnet50\": keras.applications.ResNet50,               # Caffe-style BGR mean-subtracted [0, 255]\n    \"vgg16\": keras.applications.VGG16,                     # Caffe-style BGR mean-subtracted [0, 255]\n    \"mobilenetv2\": keras.applications.MobileNetV2,         # expects [-1, 1]\n    \"densenet121\": keras.applications.DenseNet121,         # Torch-style ImageNet mean/std normalised\n}\n_CAFFE_BGR_MEAN = (103.939, 116.779, 123.68)\n_TORCH_RGB_MEAN = (0.485, 0.456, 0.406)\n_TORCH_RGB_STD = (0.229, 0.224, 0.225)\n\n\ndef _backbone_input_adapter(backbone_name: str, inputs: keras.KerasTensor) -> keras.KerasTensor:\n    \"\"\"Convert [0, 1] images to the numeric range each ImageNet backbone was trained on.\"\"\"\n    if backbone_name in (\"efficientnetb3\", \"efficientnetb0\"):\n        return layers.Rescaling(255.0, name=\"restore_efficientnet_input_range\")(inputs)\n    if backbone_name == \"mobilenetv2\":\n        return layers.Rescaling(2.0, offset=-1.0, name=\"mobilenetv2_scale_to_pm1\")(inputs)\n    if backbone_name in (\"resnet50\", \"vgg16\"):\n        scaled = layers.Rescaling(255.0, name=f\"{backbone_name}_scale_to_0_255\")(inputs)\n        return layers.Lambda(\n            lambda t: keras.ops.flip(t, axis=-1)\n            - keras.ops.convert_to_tensor(_CAFFE_BGR_MEAN, dtype=t.dtype),\n            name=f\"{backbone_name}_caffe_bgr_mean_subtract\",\n        )(scaled)\n    if backbone_name == \"densenet121\":\n        return layers.Normalization(\n            mean=list(_TORCH_RGB_MEAN),\n            variance=[std ** 2 for std in _TORCH_RGB_STD],\n            name=\"densenet121_torch_mean_std\",\n        )(inputs)\n    if backbone_name == \"scratch_cnn\":\n        return inputs  # trained from scratch on [0, 1] inputs; no pretrained convention\n    raise ValueError(f\"Unknown backbone: {backbone_name}\")\n\n\ndef build_scratch_cnn(image_size: int) -> keras.Model:\n    \"\"\"Small 4-block CNN (Conv-BN-ReLU-MaxPool, 32/64/128/256 filters) with random initialisation.\"\"\"\n    inputs = keras.Input(shape=(image_size, image_size, 3))\n    x = inputs\n    for block, filters in enumerate((32, 64, 128, 256), start=1):\n        x = layers.Conv2D(filters, 3, padding=\"same\", use_bias=False, name=f\"scratch_block{block}_conv\")(x)\n        x = layers.BatchNormalization(name=f\"scratch_block{block}_bn\")(x)\n        x = layers.Activation(\"relu\", name=f\"scratch_block{block}_relu\")(x)\n        x = layers.MaxPooling2D(2, name=f\"scratch_block{block}_pool\")(x)\n    return keras.Model(inputs, x, name=\"scratch_cnn_backbone\")\n\n\ndef _build_tuning_classifier(backbone_name: str = \"efficientnetb3\") -> tuple:\n    \"\"\"Build the shared classification head on top of the requested backbone.\"\"\"\n    if backbone_name == \"scratch_cnn\":\n        backbone = build_scratch_cnn(Config.IMG_SIZE)\n        backbone.trainable = True  # no pretrained weights to protect\n    elif backbone_name in _BACKBONE_CONSTRUCTORS:\n        backbone = _BACKBONE_CONSTRUCTORS[backbone_name](\n            include_top=False,\n            weights=\"imagenet\",\n            input_shape=(Config.IMG_SIZE, Config.IMG_SIZE, 3),\n        )\n        backbone.trainable = False\n    else:\n        raise ValueError(f\"Unknown backbone: {backbone_name}\")\n\n    inputs = keras.Input(shape=(Config.IMG_SIZE, Config.IMG_SIZE, 3))\n    features = _backbone_input_adapter(backbone_name, inputs)\n    if backbone_name == \"scratch_cnn\":\n        features = backbone(features)  # BatchNorm follows the fit/predict mode\n    else:\n        features = backbone(features, training=False)  # keep pretrained BatchNorm in inference mode\n    features = layers.GlobalAveragePooling2D(name=\"gap\")(features)\n    features = layers.BatchNormalization(name=\"head_bn\")(features)\n    features = layers.Dense(Config.DENSE_UNITS, activation=\"relu\", name=\"head_dense\")(features)\n    features = layers.Dropout(Config.DROPOUT_RATE, name=\"head_dropout\")(features)\n    outputs = layers.Dense(\n        Config.NUM_CLASSES,\n        activation=\"softmax\",\n        dtype=\"float32\",\n        name=\"predictions\",\n    )(features)\n    model_name = \"RetinaGuard_EfficientNetB3\" if backbone_name == \"efficientnetb3\" else f\"Pilot_{backbone_name}\"\n    classifier = keras.Model(inputs=inputs, outputs=outputs, name=model_name)\n    return classifier, backbone\n\n\ndef _apply_trial_to_config(trial: Dict[str, Any]) -> None:\n    check_balancing_choice(bool(trial[\"class_weights\"]), bool(trial[\"oversample\"]))\n    Config.IMG_SIZE = int(trial[\"image_size\"])\n    Config.DROPOUT_RATE = float(trial[\"dropout_rate\"])\n    Config.PHASE1_LR = float(trial[\"phase1_lr\"])\n    Config.PHASE2_LR = float(trial[\"phase2_lr\"])\n    Config.UNFREEZE_TOP_N = int(trial[\"unfreeze_top_n\"])\n    Config.USE_CLASS_WEIGHTS = bool(trial[\"class_weights\"])\n    Config.CLASS_WEIGHT_MODE = str(trial[\"class_weight_mode\"])\n    Config.USE_OVERSAMPLING = bool(trial[\"oversample\"])\n    Config.BATCH_SIZE = int(trial[\"batch_size\"])\n    Config.WEIGHT_DECAY = float(trial[\"weight_decay\"])\n    Config.PHASE2_LR_SCHEDULE = str(trial[\"phase2_lr_schedule\"])\n\n\ndef _count_trainable_parameters(m: keras.Model) -> int:\n    return int(sum(int(np.prod(v.shape)) for v in m.trainable_variables))\n\n\ndef _combine_histories(histories: List[Tuple[str, keras.callbacks.History]]) -> pd.DataFrame:\n    \"\"\"Concatenate per-phase Keras histories into one epoch-indexed table.\"\"\"\n    frames = []\n    for phase_name, history in histories:\n        frame = pd.DataFrame(history.history)\n        frame.insert(0, \"phase\", phase_name)\n        frame.insert(1, \"phase_epoch\", np.arange(1, len(frame) + 1))\n        frames.append(frame)\n    combined = pd.concat(frames, ignore_index=True)\n    combined.insert(0, \"epoch\", np.arange(1, len(combined) + 1))\n    return combined\n\n\ndef _save_pilot_evidence(trial_id: str, history_df: pd.DataFrame, y_true: np.ndarray, y_pred: np.ndarray) -> None:\n    \"\"\"Save history.csv, training-curve PNG and validation confusion-matrix PNG for one trial.\"\"\"\n    trial_dir = PILOT_DIR / trial_id\n    trial_dir.mkdir(parents=True, exist_ok=True)\n    history_df.to_csv(trial_dir / \"history.csv\", index=False)\n\n    fig, axes = plt.subplots(1, 3, figsize=(17, 5))\n    epochs = history_df[\"epoch\"]\n    panels = (\n        (\"accuracy\", \"clean_train_accuracy\", \"val_accuracy\", \"Accuracy\"),\n        (\"loss\", \"clean_train_loss\", \"val_loss\", \"Loss (label-smoothed CE)\"),\n        (None, None, \"val_qwk\", \"Validation QWK\"),\n    )\n    boundaries = history_df.index[history_df[\"phase\"] != history_df[\"phase\"].shift()].tolist()[1:]\n    for axis, (train_key, clean_key, val_key, title) in zip(axes, panels):\n        if train_key and train_key in history_df:\n            axis.plot(epochs, history_df[train_key], \"b-o\", markersize=3, label=\"Train (augmented)\")\n        if clean_key and clean_key in history_df:\n            axis.plot(epochs, history_df[clean_key], \"c--s\", markersize=3, label=\"Train (clean subset)\")\n        if val_key in history_df:\n            axis.plot(epochs, history_df[val_key], \"r-o\", markersize=3, label=\"Validation\")\n        for boundary in boundaries:\n            axis.axvline(history_df.loc[boundary, \"epoch\"] - 0.5, color=\"grey\", linestyle=\"--\", linewidth=1)\n        axis.set_title(title)\n        axis.set_xlabel(\"Epoch (Phase 1 then Phase 2)\")\n        axis.xaxis.set_major_locator(plt.MaxNLocator(integer=True))\n        axis.grid(alpha=0.3)\n        axis.legend(fontsize=8)\n    fig.suptitle(f\"Pilot trial {trial_id} — train vs validation (dashed line = phase change)\", fontsize=12)\n    fig.text(0.5, -0.02, CURVE_LEGEND_NOTE, ha=\"center\", fontsize=9, color=\"#475569\")\n    fig.tight_layout()\n    fig.savefig(trial_dir / \"training_curves.png\", dpi=130, bbox_inches=\"tight\")\n    plt.close(fig)\n\n    matrix = confusion_matrix(y_true, y_pred, labels=range(Config.NUM_CLASSES))\n    fig, axis = plt.subplots(figsize=(6.5, 5.5))\n    sns.heatmap(matrix, annot=True, fmt=\"d\", cmap=\"Blues\", ax=axis,\n                xticklabels=Config.CLASS_NAMES, yticklabels=Config.CLASS_NAMES)\n    axis.set_title(f\"{trial_id}\\nValidation confusion matrix\")\n    axis.set_xlabel(\"Predicted stage\")\n    axis.set_ylabel(\"True stage\")\n    axis.tick_params(axis=\"x\", rotation=30)\n    fig.tight_layout()\n    fig.savefig(trial_dir / \"validation_confusion_matrix.png\", dpi=130, bbox_inches=\"tight\")\n    plt.close(fig)\n\n\ndef _process_memory_gb() -> Tuple[float, float]:\n    \"\"\"Process RAM (resident set size) and GPU:0 memory in use, in GB (GPU is NaN when unavailable).\"\"\"\n    ram_gb = psutil.Process().memory_info().rss / 1024 ** 3\n    try:\n        gpu_gb = tf.config.experimental.get_memory_info(\"GPU:0\")[\"current\"] / 1024 ** 3\n    except Exception:  # no GPU, or memory info not supported on this device\n        gpu_gb = float(\"nan\")\n    return float(ram_gb), float(gpu_gb)\n\n\ndef _release_keras_memory() -> None:\n    \"\"\"Clear Keras' global state (freeing memory where supported), then run the garbage collector.\"\"\"\n    # clear_session() also resets the global dtype policy to float32, so the\n    # policy set in Section 5.1 (mixed_float16 on GPU) is restored afterwards.\n    precision_policy = keras.mixed_precision.global_policy().name\n    if \"free_memory\" in inspect.signature(keras.backend.clear_session).parameters:\n        keras.backend.clear_session(free_memory=True)\n    else:\n        keras.backend.clear_session()\n    keras.mixed_precision.set_global_policy(precision_policy)\n    gc.collect()\n\n\ndef _pilot_selection_key(validation_qwk: float, validation_loss: float) -> Tuple[float, float]:\n    \"\"\"Selection rule: highest validation QWK, tie-break lower validation loss.\"\"\"\n    return (validation_qwk if np.isfinite(validation_qwk) else -1.0, -validation_loss)\n\n\ndef _select_best_trial(\n    trial_records: List[Dict[str, Any]],\n    trials_by_id: Dict[str, Dict[str, Any]],\n) -> Optional[Dict[str, Any]]:\n    \"\"\"Best completed selectable trial (new or resumed rows alike), in trial order.\"\"\"\n    best_trial, best_key = None, (-np.inf, -np.inf)\n    for record in trial_records:\n        trial = trials_by_id[str(record[\"trial_id\"])]\n        if record[\"status\"] != \"completed\" or trial[\"group\"] not in SELECTABLE_GROUPS:\n            continue\n        key = _pilot_selection_key(float(record[\"pilot_validation_qwk\"]), float(record[\"pilot_validation_loss\"]))\n        if key > best_key:\n            best_key, best_trial = key, trial.copy()\n    return best_trial\n\n\ndef _not_run_record(trial: Dict[str, Any], status: str) -> Dict[str, Any]:\n    \"\"\"Tuning-CSV row for a trial that was not trained in this run (metrics left empty).\"\"\"\n    return {\n        \"run_id\": _RUN_ID,\n        \"trial_id\": str(trial[\"trial_id\"]),\n        \"group\": trial[\"group\"],\n        \"status\": status,\n        \"selectable_for_final_model\": trial[\"group\"] in SELECTABLE_GROUPS,\n        \"change_vs_baseline\": trial[\"change_vs_baseline\"],\n        \"trial_rationale\": trial[\"trial_rationale\"],\n        \"backbone\": trial[\"backbone\"],\n        \"split_fingerprint\": _split_fingerprint,\n    }\n\n\ndef _load_resumable_records(\n    tuning_results_path: Path,\n    trials_by_id: Dict[str, Dict[str, Any]],\n) -> Dict[str, Dict[str, Any]]:\n    \"\"\"Completed rows from an earlier session of this run folder, reusable only with the same split.\"\"\"\n    if not tuning_results_path.is_file():\n        return {}\n    previous = pd.read_csv(tuning_results_path, dtype={\"trial_id\": str, \"split_fingerprint\": str})\n    if \"status\" not in previous.columns:\n        # Files written before the status column existed only ever held completed trials.\n        previous[\"status\"] = \"completed\"\n    resumable = {}\n    for row in previous.to_dict(\"records\"):\n        trial_id = str(row[\"trial_id\"])\n        if row[\"status\"] != \"completed\" or trial_id not in trials_by_id:\n            continue\n        if str(row.get(\"split_fingerprint\")) != _split_fingerprint:\n            print(f\"[Resume] {trial_id}: saved with a different split fingerprint; it will be retrained.\")\n            continue\n        resumable[trial_id] = row\n    return resumable\n\n\ndef _write_tuning_results(trial_records: List[Dict[str, Any]], pending_resumed: List[Dict[str, Any]],\n                          tuning_results_path: Path) -> pd.DataFrame:\n    \"\"\"Write every row so far (plus resumed rows not yet reached, so a crash never loses them).\"\"\"\n    results = pd.DataFrame(trial_records + pending_resumed)\n    front = [\"run_id\", \"trial_id\", \"group\", \"status\"]\n    results = results[front + [column for column in results.columns if column not in front]]\n    results.to_csv(tuning_results_path, index=False)\n    return results\n\n\ndef _run_pilot_trial(\n    trial: Dict[str, Any],\n    pilot_train_df: pd.DataFrame,\n    pilot_val_df: pd.DataFrame,\n    pilot_true: np.ndarray,\n) -> Dict[str, Any]:\n    \"\"\"Train and score one pilot trial, save its evidence, and free everything it built.\"\"\"\n    global augmentation_layer\n\n    set_all_seeds(Config.SEED)\n    trial_id = str(trial[\"trial_id\"])\n    _apply_trial_to_config(trial)\n\n    print(\"\\n\" + \"=\" * 76)\n    print(f\"[Pilot] {trial_id}  (group: {trial['group']})\")\n    print(f\"        change: {trial['change_vs_baseline']}\")\n    print(\"=\" * 76)\n\n    augmentation_layer = build_augmentation_layer()\n    pilot_train_dataset = _build_tuning_dataset(\n        pilot_train_df, Config.IMG_SIZE, shuffle=True,\n        augment=bool(trial[\"augment\"]), preprocessing=trial[\"preprocessing\"],\n        oversample=Config.USE_OVERSAMPLING,\n    )\n    pilot_val_dataset = _build_tuning_dataset(\n        pilot_val_df, Config.IMG_SIZE, shuffle=False,\n        augment=False, preprocessing=trial[\"preprocessing\"],\n    )\n    pilot_steps = int(np.ceil(len(pilot_train_df) / Config.BATCH_SIZE)) if Config.USE_OVERSAMPLING else None\n    pilot_class_weights = active_class_weights()\n    pilot_clean_callback = clean_train_callback(Config.IMG_SIZE, trial[\"preprocessing\"])\n\n    def _pilot_callbacks(schedule: str, total_epochs: int, peak_lr: float) -> list:\n        return [\n            ValidationQWK(pilot_val_dataset, pilot_true),   # monitoring only\n            pilot_clean_callback,                           # monitoring only\n            LearningRateLogger(),\n            EarlyStopping(monitor=\"val_loss\", patience=PILOT_ES_PATIENCE,\n                          restore_best_weights=True, mode=\"min\", verbose=1),\n            *lr_schedule_callbacks(schedule, total_epochs, peak_lr),\n        ]\n\n    pilot_model, pilot_base_model = _build_tuning_classifier(trial[\"backbone\"])\n    # Different backbones have different depths; never unfreeze more layers than exist.\n    Config.UNFREEZE_TOP_N = min(Config.UNFREEZE_TOP_N, len(pilot_base_model.layers))\n    # XLA off for pilots: each trial would otherwise keep its own compiled XLA cluster in host RAM.\n    compile_model_phase1(pilot_model, jit_compile=False)\n\n    fit_start = time.perf_counter()\n    if trial[\"backbone\"] == \"scratch_cnn\":\n        # No pretrained features to freeze or fine-tune: train end to end for the same total budget.\n        scratch_callbacks = _pilot_callbacks(\"plateau\", TUNING_PHASE1_EPOCHS + TUNING_PHASE2_EPOCHS, Config.PHASE1_LR)\n        histories = [(\"Scratch (end to end)\", pilot_model.fit(\n            pilot_train_dataset,\n            epochs=TUNING_PHASE1_EPOCHS + TUNING_PHASE2_EPOCHS,\n            steps_per_epoch=pilot_steps,\n            validation_data=pilot_val_dataset,\n            callbacks=scratch_callbacks,\n            class_weight=pilot_class_weights,\n            verbose=1,\n        ))]\n        del scratch_callbacks\n    else:\n        phase1_callbacks = _pilot_callbacks(\"plateau\", TUNING_PHASE1_EPOCHS, Config.PHASE1_LR)\n        history_phase1_pilot = pilot_model.fit(\n            pilot_train_dataset,\n            epochs=TUNING_PHASE1_EPOCHS,\n            steps_per_epoch=pilot_steps,\n            validation_data=pilot_val_dataset,\n            callbacks=phase1_callbacks,\n            class_weight=pilot_class_weights,\n            verbose=1,\n        )\n        configure_phase2(pilot_model, pilot_base_model, jit_compile=False)\n        phase2_callbacks = _pilot_callbacks(Config.PHASE2_LR_SCHEDULE, TUNING_PHASE2_EPOCHS, Config.PHASE2_LR)\n        history_phase2_pilot = pilot_model.fit(\n            pilot_train_dataset,\n            epochs=TUNING_PHASE2_EPOCHS,\n            steps_per_epoch=pilot_steps,\n            validation_data=pilot_val_dataset,\n            callbacks=phase2_callbacks,\n            class_weight=pilot_class_weights,\n            verbose=1,\n        )\n        histories = [(\"Phase 1\", history_phase1_pilot), (\"Phase 2\", history_phase2_pilot)]\n        del phase1_callbacks, phase2_callbacks, history_phase1_pilot, history_phase2_pilot\n    training_seconds = time.perf_counter() - fit_start\n    trainable_parameters = _count_trainable_parameters(pilot_model)\n\n    pilot_probabilities = pilot_model.predict(pilot_val_dataset, verbose=0)\n    pilot_pred = np.argmax(pilot_probabilities, axis=1)\n    if len(pilot_true) != len(pilot_pred):\n        raise ValueError(f\"Pilot validation prediction count mismatch for {trial_id}.\")\n\n    history_df = _combine_histories(histories)\n    best_index = int(np.nanargmin(history_df[\"val_loss\"]))\n    final_phase_val_loss = float(np.nanmin(histories[-1][1].history[\"val_loss\"]))  # restored weights\n    _save_pilot_evidence(trial_id, history_df, pilot_true, pilot_pred)\n\n    pilot_accuracy = float(np.mean(pilot_true == pilot_pred))\n    pilot_qwk = float(cohen_kappa_score(pilot_true, pilot_pred, weights=\"quadratic\"))\n    pilot_report = classification_report(\n        pilot_true, pilot_pred,\n        labels=list(range(Config.NUM_CLASSES)),\n        target_names=Config.CLASS_NAMES,\n        output_dict=True, zero_division=0,\n    )\n    record = {\n        \"run_id\": _RUN_ID,\n        \"trial_id\": trial_id,\n        \"group\": trial[\"group\"],\n        \"status\": \"completed\",\n        \"selectable_for_final_model\": trial[\"group\"] in SELECTABLE_GROUPS,\n        \"change_vs_baseline\": trial[\"change_vs_baseline\"],\n        \"trial_rationale\": trial[\"trial_rationale\"],\n        \"result_stage\": \"pilot_screen\",\n        \"backbone\": trial[\"backbone\"],\n        \"backbone_parameters\": int(pilot_base_model.count_params()),\n        \"trainable_parameters\": trainable_parameters,\n        \"preprocessing_variant\": _preprocessing_tag(trial[\"preprocessing\"]).lstrip(\"_\") or \"default_ben_graham\",\n        \"augmentation\": bool(trial[\"augment\"]),\n        \"image_size\": Config.IMG_SIZE,\n        \"dropout_rate\": Config.DROPOUT_RATE,\n        \"phase1_learning_rate\": Config.PHASE1_LR,\n        \"phase2_learning_rate\": Config.PHASE2_LR,\n        \"phase2_unfreeze_top_n\": Config.UNFREEZE_TOP_N,\n        \"class_weights_enabled\": bool(Config.USE_CLASS_WEIGHTS),\n        \"class_weight_mode\": Config.CLASS_WEIGHT_MODE,\n        \"oversampling_enabled\": bool(Config.USE_OVERSAMPLING),\n        \"batch_size\": Config.BATCH_SIZE,\n        \"weight_decay\": Config.WEIGHT_DECAY,\n        \"phase2_lr_schedule\": Config.PHASE2_LR_SCHEDULE,\n        \"seed\": Config.SEED,\n        \"split_fingerprint\": _split_fingerprint,\n        \"pilot_train_images\": len(pilot_train_df),\n        \"pilot_validation_images\": len(pilot_val_df),\n        \"pilot_train_fraction\": TUNING_TRAIN_FRACTION,\n        \"pilot_validation_fraction\": TUNING_VALIDATION_FRACTION,\n        \"epochs_run\": len(history_df),\n        \"pilot_phase1_epochs_run\": len(histories[0][1].history[\"val_loss\"]),\n        \"pilot_phase2_epochs_run\": len(histories[1][1].history[\"val_loss\"]) if len(histories) > 1 else 0,\n        \"best_val_loss_epoch\": best_index + 1,\n        \"train_val_accuracy_gap_at_best_epoch\": float(\n            history_df.loc[best_index, \"accuracy\"] - history_df.loc[best_index, \"val_accuracy\"]\n        ),\n        \"clean_train_minus_validation_accuracy_gap_at_best_epoch\": float(\n            history_df.loc[best_index, \"clean_train_accuracy\"] - history_df.loc[best_index, \"val_accuracy\"]\n        ),\n        \"val_loss_rise_after_min\": float(history_df[\"val_loss\"].iloc[-1] - history_df[\"val_loss\"].iloc[best_index]),\n        \"training_time_seconds\": round(training_seconds, 1),\n        \"pilot_validation_accuracy\": pilot_accuracy,\n        \"pilot_validation_qwk\": pilot_qwk,\n        \"pilot_validation_macro_f1\": float(pilot_report[\"macro avg\"][\"f1-score\"]),\n        \"pilot_validation_loss\": final_phase_val_loss,\n        **{\n            f\"pilot_recall_class_{class_id}\": float(pilot_report[class_name][\"recall\"])\n            for class_id, class_name in enumerate(Config.CLASS_NAMES)\n        },\n    }\n    print(\n        f\"[Pilot] {trial_id}: validation accuracy={pilot_accuracy:.4f}, \"\n        f\"QWK={pilot_qwk:.4f}, macro-F1={record['pilot_validation_macro_f1']:.4f}, \"\n        f\"val loss={final_phase_val_loss:.4f}, {record['epochs_run']} epochs in {training_seconds:.0f}s\"\n    )\n\n    # Free everything this trial built before the next one starts.\n    del pilot_model, pilot_base_model\n    del pilot_train_dataset, pilot_val_dataset, pilot_clean_callback, _pilot_callbacks\n    del augmentation_layer\n    del histories, history_df, pilot_probabilities, pilot_pred\n    return record\n\n\ndef run_hyperparameter_tuning() -> pd.DataFrame:\n    \"\"\"Run the pilot trials (skip list, resume, limits), then prepare the winning selectable config.\"\"\"\n    global model, base_model, augmentation_layer\n    global train_dataset, val_dataset, test_dataset\n    global steps_per_epoch, validation_steps, test_steps\n    global selected_tuning_trial, tuning_results_df, EXECUTED_TRIALS\n\n    if not RUN_HYPERPARAMETER_TUNING:\n        print(\"[Tuning] Disabled; the ordinary single-configuration path will be used.\")\n        return pd.DataFrame()\n\n    EXECUTED_TRIALS = TUNING_TRIALS + (COMPARISON_TRIALS if RUN_COMPARISON_TRIALS else [])\n    trials_by_id = {str(trial[\"trial_id\"]): trial for trial in EXECUTED_TRIALS}\n    defined_ids = {str(trial[\"trial_id\"]) for trial in TUNING_TRIALS + COMPARISON_TRIALS}\n    skip_ids = set(Config.PILOT_SKIP_TRIALS)\n    unknown_skips = sorted(skip_ids - defined_ids)\n    if unknown_skips:\n        raise ValueError(f\"Config.PILOT_SKIP_TRIALS lists unknown trial IDs: {unknown_skips}\")\n    # The first N non-skipped trials (all of them when PILOT_TRIAL_LIMIT is None).\n    allowed_ids = [trial_id for trial_id in trials_by_id if trial_id not in skip_ids]\n    if Config.PILOT_TRIAL_LIMIT is not None:\n        allowed_ids = allowed_ids[: int(Config.PILOT_TRIAL_LIMIT)]\n    if not any(trials_by_id[trial_id][\"group\"] in SELECTABLE_GROUPS for trial_id in allowed_ids):\n        raise ValueError(\"At least one selectable (hyperparameter/balancing) trial must be allowed to run.\")\n\n    pilot_train_df = _stratified_pilot_sample(train_df, TUNING_TRAIN_FRACTION, Config.SEED)\n    pilot_val_df = _stratified_pilot_sample(val_df, TUNING_VALIDATION_FRACTION, Config.SEED + 1000)\n    assert set(pilot_train_df[\"split\"]) == {\"train\"} and set(pilot_val_df[\"split\"]) == {\"validation\"}, \\\n        \"Pilot trials may only use the training and validation splits.\"\n    pilot_true = pilot_val_df[\"label\"].to_numpy(dtype=int)\n\n    tuning_results_path = Path(Config.REPORTS_DIR) / \"hyperparameter_tuning_results.csv\"\n    PILOT_DIR.mkdir(parents=True, exist_ok=True)\n    resumable_records = _load_resumable_records(tuning_results_path, trials_by_id)\n    resumed_ids = [trial_id for trial_id in allowed_ids if trial_id in resumable_records]\n    if resumed_ids:\n        print(f\"[Resume] Reusing {len(resumed_ids)} completed trial(s) from {tuning_results_path}:\")\n        for trial_id in resumed_ids:\n            row = resumable_records[trial_id]\n            print(f\"         {trial_id}: validation QWK={float(row['pilot_validation_qwk']):.4f}, \"\n                  f\"val loss={float(row['pilot_validation_loss']):.4f}\")\n    elif Config.RESUME_RUN_ID:\n        print(\"[Resume] No completed trials with the current split fingerprint were found.\")\n\n    print(\n        f\"[Tuning] Pilot data: {len(pilot_train_df):,}/{len(train_df):,} training and \"\n        f\"{len(pilot_val_df):,}/{len(val_df):,} validation images. \"\n        f\"{len(allowed_ids) - len(resumed_ids)} trials to train, {len(resumed_ids)} resumed, \"\n        f\"{len(skip_ids & set(trials_by_id))} skipped by config, \"\n        f\"{len(trials_by_id) - len(skip_ids & set(trials_by_id)) - len(allowed_ids)} beyond PILOT_TRIAL_LIMIT \"\n        f\"(up to {TUNING_PHASE1_EPOCHS} Phase 1 + {TUNING_PHASE2_EPOCHS} Phase 2 epochs each; \"\n        f\"time budget {Config.PILOT_TIME_BUDGET_HOURS} h).\"\n    )\n\n    trial_records = []\n    time_budget_exceeded = False\n    pilot_loop_start = time.perf_counter()\n    for trial in EXECUTED_TRIALS:\n        trial_id = str(trial[\"trial_id\"])\n        if trial_id in skip_ids:\n            trial_records.append(_not_run_record(trial, \"skipped_by_config\"))\n            continue\n        if trial_id not in allowed_ids:\n            trial_records.append(_not_run_record(trial, \"skipped_trial_limit\"))\n            continue\n        if trial_id in resumable_records:\n            trial_records.append(resumable_records[trial_id])\n            continue\n\n        elapsed_hours = (time.perf_counter() - pilot_loop_start) / 3600\n        if (not time_budget_exceeded and Config.PILOT_TIME_BUDGET_HOURS is not None\n                and elapsed_hours > Config.PILOT_TIME_BUDGET_HOURS):\n            time_budget_exceeded = True\n            print(f\"\\n[Tuning] Time budget reached ({elapsed_hours:.2f} h > {Config.PILOT_TIME_BUDGET_HOURS} h); \"\n                  \"remaining trials are not started.\")\n        if time_budget_exceeded:\n            trial_records.append(_not_run_record(trial, \"skipped_time_budget\"))\n            continue\n\n        record = _run_pilot_trial(trial, pilot_train_df, pilot_val_df, pilot_true)\n        _release_keras_memory()\n        ram_gb, gpu_gb = _process_memory_gb()\n        record[\"ram_gb_after_trial\"] = round(ram_gb, 3)\n        record[\"gpu_mem_gb_after_trial\"] = round(gpu_gb, 3)\n        print(f\"[Memory] After {trial_id}: process RAM {ram_gb:.2f} GB, GPU memory in use {gpu_gb:.2f} GB \"\n              f\"({(time.perf_counter() - pilot_loop_start) / 3600:.2f} h since the pilot loop began)\")\n        trial_records.append(record)\n        reached_ids = {str(row[\"trial_id\"]) for row in trial_records}\n        _write_tuning_results(\n            trial_records,\n            [resumable_records[i] for i in resumed_ids if i not in reached_ids],\n            tuning_results_path,\n        )\n\n    tuning_results_df = _write_tuning_results(trial_records, [], tuning_results_path)\n    print(\"\\n[Pilot] Trial status counts:\", tuning_results_df[\"status\"].value_counts().to_dict())\n\n    best_trial = _select_best_trial(trial_records, trials_by_id)\n    if best_trial is None:\n        raise RuntimeError(\"No valid selectable pilot trial completed.\")\n\n    # Rebuild the final EfficientNetB3 with the selected settings and default preprocessing.\n    selected_tuning_trial = best_trial\n    _apply_trial_to_config(best_trial)\n\n    augmentation_layer = build_augmentation_layer()\n    train_dataset = _build_tuning_dataset(\n        train_df, Config.IMG_SIZE, shuffle=True, augment=True, oversample=Config.USE_OVERSAMPLING,\n    )\n    val_dataset = _build_tuning_dataset(val_df, Config.IMG_SIZE, shuffle=False, augment=False)\n    test_dataset = _build_tuning_dataset(test_df, Config.IMG_SIZE, shuffle=False, augment=False)\n    steps_per_epoch = int(np.ceil(len(train_df) / Config.BATCH_SIZE))\n    validation_steps = len(val_df) // Config.BATCH_SIZE\n    test_steps = len(test_df) // Config.BATCH_SIZE\n\n    model, base_model = _build_tuning_classifier(\"efficientnetb3\")\n    Config.UNFREEZE_TOP_N = min(Config.UNFREEZE_TOP_N, len(base_model.layers))\n    compile_model_phase1(model)  # final model: default jit_compile=\"auto\"\n\n    print(\"\\n[Pilot] Selected settings for the final full-data fit:\",\n          {key: best_trial[key] for key in (\"trial_id\", \"image_size\", \"dropout_rate\", \"phase1_lr\",\n                                            \"phase2_lr\", \"unfreeze_top_n\", \"class_weight_mode\", \"oversample\",\n                                            \"batch_size\", \"weight_decay\", \"phase2_lr_schedule\")})\n    print(\"[Pilot] Selection: highest validation QWK (tie-break: lower validation loss) among completed \"\n          f\"trials in groups {sorted(SELECTABLE_GROUPS)}.\")\n    # Trade-off evidence for the report — printed only, NOT used for selection.\n    selectable_rows = tuning_results_df[\n        (tuning_results_df[\"status\"] == \"completed\") & tuning_results_df[\"group\"].isin(SELECTABLE_GROUPS)\n    ]\n    best_accuracy_row = selectable_rows.loc[selectable_rows[\"pilot_validation_accuracy\"].idxmax()]\n    smallest_gap_row = selectable_rows.loc[\n        selectable_rows[\"clean_train_minus_validation_accuracy_gap_at_best_epoch\"].abs().idxmin()\n    ]\n    print(f\"[Pilot] (Not used for selection) Highest validation accuracy: {best_accuracy_row['trial_id']} \"\n          f\"({best_accuracy_row['pilot_validation_accuracy']:.4f})\")\n    print(f\"[Pilot] (Not used for selection) Smallest clean-train vs validation accuracy gap: \"\n          f\"{smallest_gap_row['trial_id']} \"\n          f\"({smallest_gap_row['clean_train_minus_validation_accuracy_gap_at_best_epoch']:+.4f})\")\n    print(\"[Pilot] All pilot rows saved to:\", tuning_results_path)\n    print(\"[Pilot] Per-trial history, curves and confusion matrices saved under:\", PILOT_DIR)\n    return tuning_results_df\n\n\nif RUN_HYPERPARAMETER_TUNING:\n    hyperparameter_tuning_results = run_hyperparameter_tuning()\nelse:\n    print(\"[Tuning] Disabled; use the ordinary single-configuration training path.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:01:07.138062Z","iopub.execute_input":"2026-09-30T20:01:07.138472Z","iopub.status.idle":"2026-09-30T20:01:09.467255Z","shell.execute_reply.started":"2026-09-30T20:01:07.138445Z","shell.execute_reply":"2026-09-30T20:01:09.466397Z"}},"outputs":[],"execution_count":null},{"id":"3999fc1e","cell_type":"code","source":"# ============================================================\n# SECTION 6.2c — MODEL COMPARISON TABLE AND DIFFERENCE FROM BASELINE\n# One row per trial with validation metrics and overfitting evidence,\n# plus each non-hyperparameter trial's difference from the baseline.\n# All values are computed from the pilot runs above; validation split only.\n# ============================================================\n\nif RUN_HYPERPARAMETER_TUNING:\n    # Skipped trials (config, trial limit, time budget) stay in the tuning CSV but not in these tables.\n    completed_pilot_results = hyperparameter_tuning_results[\n        hyperparameter_tuning_results[\"status\"] == \"completed\"\n    ].reset_index(drop=True)\n    comparison_columns = [\n        \"trial_id\", \"group\", \"change_vs_baseline\", \"backbone\",\n        \"trainable_parameters\", \"epochs_run\", \"training_time_seconds\",\n        \"pilot_validation_qwk\", \"pilot_validation_macro_f1\", \"pilot_validation_accuracy\",\n        \"pilot_validation_loss\", \"train_val_accuracy_gap_at_best_epoch\",\n        \"clean_train_minus_validation_accuracy_gap_at_best_epoch\", \"val_loss_rise_after_min\",\n        *[f\"pilot_recall_class_{class_id}\" for class_id in range(Config.NUM_CLASSES)],\n    ]\n    model_comparison_table = completed_pilot_results[comparison_columns].copy()\n    model_comparison_path = Path(Config.REPORTS_DIR) / \"model_comparison_table.csv\"\n    model_comparison_table.to_csv(model_comparison_path, index=False)\n    print(\"[Pilot] Final model comparison table (full validation split):\")\n    display(model_comparison_table.round(4))\n    print(f\"[Pilot] Saved -> {model_comparison_path}\")\n    print(f\"[Pilot] Selected for the final fit: {selected_tuning_trial['trial_id']}\")\n\n    baseline_rows = completed_pilot_results.loc[\n        completed_pilot_results[\"trial_id\"] == TUNING_TRIALS[0][\"trial_id\"]\n    ]\n    if baseline_rows.empty:\n        raise RuntimeError(\"The baseline pilot trial did not complete, so differences from it cannot be computed.\")\n    baseline_row = baseline_rows.iloc[0]\n    delta_metrics = [\n        \"pilot_validation_accuracy\",\n        \"pilot_validation_qwk\",\n        \"pilot_validation_macro_f1\",\n        *[f\"pilot_recall_class_{class_id}\" for class_id in range(Config.NUM_CLASSES)],\n    ]\n    comparison_rows = completed_pilot_results[\n        completed_pilot_results[\"group\"] != \"hyperparameter\"\n    ].copy()\n    for metric in delta_metrics:\n        comparison_rows[f\"delta_{metric.replace('pilot_', '')}\"] = (\n            comparison_rows[metric] - float(baseline_row[metric])\n        )\n    comparison_path = Path(Config.REPORTS_DIR) / \"comparison_trials_results.csv\"\n    comparison_rows.to_csv(comparison_path, index=False)\n\n    if not comparison_rows.empty:\n        print(\"\\n[Comparison] Difference from the EfficientNetB3 + Ben Graham baseline \"\n              \"(positive = better than baseline):\")\n        display(\n            comparison_rows[\n                [\"trial_id\", \"group\", \"change_vs_baseline\",\n                 \"delta_validation_qwk\", \"delta_validation_macro_f1\", \"delta_validation_accuracy\",\n                 \"delta_recall_class_1\", \"delta_recall_class_3\"]\n            ].round(4)\n        )\n        print(f\"The validation split has only {len(val_df):,} images \"\n              f\"({int((val_df['label'] == 3).sum())} Severe), so differences of a few hundredths \"\n              \"can be noise; discuss them as indicative, not conclusive.\")\n    print(f\"[Comparison] Saved -> {comparison_path}\")\nelse:\n    print(\"[Pilot] Skipped because hyperparameter tuning is disabled.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:01:43.176932Z","iopub.execute_input":"2026-09-30T20:01:43.177341Z","iopub.status.idle":"2026-09-30T20:01:43.228389Z","shell.execute_reply.started":"2026-09-30T20:01:43.177314Z","shell.execute_reply":"2026-09-30T20:01:43.227795Z"}},"outputs":[],"execution_count":null},{"id":"21266198","cell_type":"code","source":"# ============================================================\n# SECTION 6.2d — PILOT COMPARISON GRAPHS\n# (a) Validation QWK and macro-F1 per trial, grouped by trial group.\n# (b) Per-stage validation recall heatmap.\n# These are pilot runs (up to 4 + 6 epochs), not final model results.\n# ============================================================\n\nif RUN_HYPERPARAMETER_TUNING:\n    group_order = [\"hyperparameter\", \"balancing\", \"preprocessing_ablation\",\n                   \"augmentation_ablation\", \"backbone\", \"scratch_baseline\"]\n    group_colors = {\n        \"hyperparameter\": \"#367C8A\",\n        \"balancing\": \"#7A8F55\",\n        \"preprocessing_ablation\": \"#D8784B\",\n        \"augmentation_ablation\": \"#B5655E\",\n        \"backbone\": \"#6E6A9E\",\n        \"scratch_baseline\": \"#8C8C8C\",\n    }\n    # Only completed trials are plotted; skipped rows have no metrics.\n    pilot_plot_df = hyperparameter_tuning_results[hyperparameter_tuning_results[\"status\"] == \"completed\"].copy()\n    pilot_plot_df[\"group\"] = pd.Categorical(pilot_plot_df[\"group\"], categories=group_order, ordered=True)\n    pilot_plot_df = pilot_plot_df.sort_values([\"group\", \"pilot_validation_qwk\"], ascending=[True, False])\n    trial_labels = pilot_plot_df[\"trial_id\"].astype(str).tolist()\n    positions = np.arange(len(trial_labels))\n    bar_colors = [group_colors.get(str(group), \"#888888\") for group in pilot_plot_df[\"group\"]]\n    baseline_values = pilot_plot_df.set_index(\"trial_id\").loc[TUNING_TRIALS[0][\"trial_id\"]]\n\n    fig, axis = plt.subplots(figsize=(max(12, 0.75 * len(trial_labels) + 4), 6))\n    bar_width = 0.4\n    axis.bar(positions - bar_width / 2, pilot_plot_df[\"pilot_validation_qwk\"], bar_width,\n             color=bar_colors, edgecolor=\"black\", linewidth=0.5, label=\"Validation QWK\")\n    axis.bar(positions + bar_width / 2, pilot_plot_df[\"pilot_validation_macro_f1\"], bar_width,\n             color=bar_colors, edgecolor=\"black\", linewidth=0.5, hatch=\"//\", alpha=0.75, label=\"Validation macro-F1\")\n    axis.axhline(float(baseline_values[\"pilot_validation_qwk\"]), color=\"black\", linestyle=\"--\",\n                 linewidth=1, label=\"Baseline QWK\")\n    group_edges = pilot_plot_df[\"group\"].astype(str).ne(pilot_plot_df[\"group\"].astype(str).shift()).to_numpy()\n    for edge in np.flatnonzero(group_edges)[1:]:\n        axis.axvline(edge - 0.5, color=\"grey\", linewidth=0.8, alpha=0.6)\n    axis.set_xticks(positions, trial_labels, rotation=60, ha=\"right\", fontsize=8)\n    axis.set_ylabel(\"Score (validation split)\")\n    axis.set_title(\"Pilot trials — validation QWK (solid) and macro-F1 (hatched), grouped by trial group\")\n    handles = [plt.Rectangle((0, 0), 1, 1, color=group_colors[group]) for group in group_order\n               if group in set(pilot_plot_df[\"group\"].astype(str))]\n    labels = [group for group in group_order if group in set(pilot_plot_df[\"group\"].astype(str))]\n    legend_groups = axis.legend(handles, labels, title=\"Trial group\", loc=\"upper right\", fontsize=8)\n    axis.add_artist(legend_groups)\n    axis.legend(loc=\"upper center\", fontsize=8)\n    axis.grid(axis=\"y\", alpha=0.25)\n    fig.tight_layout()\n    pilot_comparison_plot_path = Path(Config.REPORTS_DIR) / \"pilot_comparison.png\"\n    fig.savefig(pilot_comparison_plot_path, dpi=160, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"[Pilot Comparison] Saved QWK / macro-F1 chart -> {pilot_comparison_plot_path}\")\n\n    recall_columns = [f\"pilot_recall_class_{class_id}\" for class_id in range(Config.NUM_CLASSES)]\n    recall_matrix = pilot_plot_df.set_index(\"trial_id\")[recall_columns].copy()\n    recall_matrix.columns = Config.CLASS_NAMES\n    fig, axis = plt.subplots(figsize=(9, max(5, 0.45 * len(recall_matrix) + 1.5)))\n    sns.heatmap(recall_matrix, annot=True, fmt=\".2f\", vmin=0, vmax=1, cmap=\"YlGnBu\",\n                ax=axis, cbar_kws={\"label\": \"Recall\"})\n    axis.set_title(\"Per-stage validation recall by pilot trial\")\n    axis.set_xlabel(\"Dataset grade\")\n    axis.set_ylabel(\"\")\n    fig.tight_layout()\n    recall_plot_path = Path(Config.REPORTS_DIR) / \"pilot_recall_heatmap.png\"\n    fig.savefig(recall_plot_path, dpi=160, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"[Pilot Comparison] Saved recall heatmap -> {recall_plot_path}\")\nelse:\n    print(\"[Pilot Comparison] Skipped because hyperparameter tuning is disabled.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:01:58.551431Z","iopub.execute_input":"2026-09-30T20:01:58.551975Z","iopub.status.idle":"2026-09-30T20:02:00.148655Z","shell.execute_reply.started":"2026-09-30T20:01:58.551948Z","shell.execute_reply":"2026-09-30T20:02:00.147974Z"}},"outputs":[],"execution_count":null},{"id":"ec921d08632e4c80a8736c88c1db5aaf","cell_type":"code","source":"# ============================================================\n# SECTION 6.3 — FINAL PHASE 1 TRAINING (FROZEN BACKBONE)\n# After the optional pilot selects settings, train the head on the\n# complete training split using the selected configuration.\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  FINAL PHASE 1 — Frozen Backbone (Head Only)\")\nprint(f\"  Image size   : {Config.IMG_SIZE}\")\nprint(f\"  Dropout      : {Config.DROPOUT_RATE}\")\nprint(f\"  Max epochs   : {Config.PHASE1_EPOCHS}\")\nprint(f\"  Learning rate: {Config.PHASE1_LR}\")\nprint(f\"  Batch size   : {Config.BATCH_SIZE}\")\nprint(\"=\" * 60)\n\ncallbacks_phase1 = build_callbacks(\"phase1\", lr_schedule=\"plateau\",\n                                   total_epochs=Config.PHASE1_EPOCHS, peak_lr=Config.PHASE1_LR)\nhistory_phase1 = model.fit(\n    train_dataset,\n    epochs=Config.PHASE1_EPOCHS,\n    validation_data=val_dataset,\n    callbacks=callbacks_phase1,\n    steps_per_epoch=steps_per_epoch if Config.USE_OVERSAMPLING else None,\n    class_weight=active_class_weights(),\n    verbose=1,\n)\n\nlog_experiment(\n    phase=\"Phase 1 - Final Full-Data Fit\",\n    lr=Config.PHASE1_LR,\n    batch_size=Config.BATCH_SIZE,\n    epochs_planned=Config.PHASE1_EPOCHS,\n    history=history_phase1,\n    notes=(\n        f\"Selected pilot trial={selected_tuning_trial['trial_id'] if RUN_HYPERPARAMETER_TUNING else 'single_config'}; \"\n        f\"image_size={Config.IMG_SIZE}; dropout={Config.DROPOUT_RATE}; \"\n        f\"class weights enabled={Config.USE_CLASS_WEIGHTS}; class weight mode={Config.CLASS_WEIGHT_MODE}; \"\n        f\"oversampling={Config.USE_OVERSAMPLING}; batch={Config.BATCH_SIZE}\"\n    ),\n)\n\npd.DataFrame(history_phase1.history).to_csv(\n    os.path.join(Config.REPORTS_DIR, \"phase1_training_history.csv\"), index=False\n)\nprint(\"[Phase 1] Final full-data training complete.\")","metadata":{"id":"ec921d08632e4c80a8736c88c1db5aaf","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:02:08.550635Z","iopub.execute_input":"2026-09-30T20:02:08.551103Z","iopub.status.idle":"2026-09-30T20:10:26.4391Z","shell.execute_reply.started":"2026-09-30T20:02:08.551075Z","shell.execute_reply":"2026-09-30T20:10:26.438379Z"}},"outputs":[],"execution_count":null},{"id":"c14b7e770a26476589dd3b0ca8ee3139","cell_type":"code","source":"# ============================================================\n# SECTION 6.4 — PHASE 1 LEARNING CURVES\n# Plot and save training vs validation accuracy and loss curves, plus\n# validation QWK. Each accuracy/loss panel shows three lines:\n#   Train (augmented)    -> Keras training metric on augmented batches,\n#                           dropout on, class-weighted loss\n#   Train (clean subset) -> 500 fixed training images, no augmentation,\n#                           inference mode, unweighted loss (like-for-like)\n#   Validation           -> validation split, inference mode, unweighted\n# These give an early visual indicator of overfitting (validation\n# diverging from the clean training curve) or underfitting (all\n# curves plateauing at a poor value) before fine-tuning begins.\n# ============================================================\n\ndef plot_training_curves(\n    history: keras.callbacks.History,\n    phase_name: str,\n    save_dir: str = Config.REPORTS_DIR,\n) -> None:\n    \"\"\"\n    Plot and save training vs validation accuracy, loss and validation QWK.\n\n    Three subplots side by side:\n    - Left: accuracy — train (augmented), train (clean subset), validation\n    - Middle: loss — train (augmented), train (clean subset), validation\n    - Right: val_qwk (logged by the ValidationQWK callback)\n\n    The clean-subset line is the like-for-like comparison with validation:\n    a widening gap between it and validation indicates overfitting; all\n    curves plateauing at a poor value indicates underfitting. A vertical\n    line marks the lowest-validation-loss epoch, whose weights early\n    stopping restores. Curves are plotted exactly as logged (no smoothing).\n\n    Args:\n        history:    Keras History object from model.fit().\n        phase_name: Label for the figure title and filename.\n        save_dir:   Directory to save the PNG.\n    \"\"\"\n    acc     = history.history.get(\"accuracy\",     history.history.get(\"categorical_accuracy\", []))\n    val_acc = history.history.get(\"val_accuracy\", history.history.get(\"val_categorical_accuracy\", []))\n    loss     = history.history[\"loss\"]\n    val_loss = history.history[\"val_loss\"]\n    clean_acc = history.history.get(\"clean_train_accuracy\", [])\n    clean_loss = history.history.get(\"clean_train_loss\", [])\n    epochs   = range(1, len(loss) + 1)\n\n    val_qwk = history.history.get(\"val_qwk\", [])\n    fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(19, 5.5))\n    fig.suptitle(f\"Training Curves — {phase_name}\", fontsize=13, fontweight=\"bold\")\n\n    # ---- Accuracy subplot ----\n    ax1.plot(epochs, acc,     \"b-o\", markersize=4, label=\"Train (augmented)\")\n    if clean_acc:\n        ax1.plot(epochs, clean_acc, \"c--s\", markersize=4, label=\"Train (clean subset)\")\n    ax1.plot(epochs, val_acc, \"r-o\", markersize=4, label=\"Validation\")\n    ax1.set_title(\"Accuracy\")\n    ax1.set_xlabel(\"Epoch\")\n    ax1.set_ylabel(\"Accuracy\")\n    ax1.legend(fontsize=8)\n    ax1.grid(alpha=0.3)\n    ax1.set_ylim(0, 1)\n\n    # ---- Loss subplot ----\n    ax2.plot(epochs, loss,     \"b-o\", markersize=4, label=\"Train (augmented)\")\n    if clean_loss:\n        ax2.plot(epochs, clean_loss, \"c--s\", markersize=4, label=\"Train (clean subset)\")\n    ax2.plot(epochs, val_loss, \"r-o\", markersize=4, label=\"Validation\")\n    ax2.set_title(\"Loss (label-smoothed cross-entropy)\")\n    ax2.set_xlabel(\"Epoch\")\n    ax2.set_ylabel(\"Loss\")\n    ax2.legend(fontsize=8)\n    ax2.grid(alpha=0.3)\n\n    # ---- Validation QWK subplot (monitoring only; not used for stopping) ----\n    if val_qwk:\n        ax3.plot(epochs, val_qwk, \"g-o\", markersize=4, label=\"Validation QWK\")\n    ax3.set_title(\"Validation Quadratic Weighted Kappa\")\n    ax3.set_xlabel(\"Epoch\")\n    ax3.set_ylabel(\"QWK\")\n    ax3.legend(fontsize=8)\n    ax3.grid(alpha=0.3)\n\n    # Mark the lowest-validation-loss epoch (the weights that early stopping restores)\n    best_epoch = int(np.argmin(val_loss)) + 1\n    for ax in (ax1, ax2, ax3):\n        ax.axvline(\n            x=best_epoch,\n            color=\"green\",\n            linestyle=\"--\",\n            linewidth=1.2,\n            alpha=0.7,\n            label=f\"Best epoch ({best_epoch})\",\n        )\n        ax.xaxis.set_major_locator(plt.MaxNLocator(integer=True))\n\n    fig.text(0.5, -0.02, CURVE_LEGEND_NOTE, ha=\"center\", fontsize=9, color=\"#475569\")\n    plt.tight_layout()\n    save_path = os.path.join(save_dir, f\"curves_{phase_name.replace(' ', '_')}.png\")\n    os.makedirs(save_dir, exist_ok=True)\n    plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    print(f\"[Curves] Saved -> {save_path}\")\n    plt.show()\n\n\nplot_training_curves(history_phase1, \"Phase_1_Frozen_Base\")\n","metadata":{"id":"c14b7e770a26476589dd3b0ca8ee3139","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:11:53.579532Z","iopub.execute_input":"2026-09-30T20:11:53.580135Z","iopub.status.idle":"2026-09-30T20:11:54.703689Z","shell.execute_reply.started":"2026-09-30T20:11:53.580105Z","shell.execute_reply":"2026-09-30T20:11:54.703031Z"}},"outputs":[],"execution_count":null},{"id":"64ea143081634da796644510d780dea0","cell_type":"code","source":"# ============================================================\n# SECTION 6.5 — FINAL PHASE 2 TRAINING (FINE-TUNING)\n# Fine-tune the selected configuration on the complete\n# training split; these are the weights used by final evaluation.\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  FINAL PHASE 2 — Fine-Tuning (BatchNorm Frozen)\")\nprint(f\"  Selected pilot trial: {selected_tuning_trial['trial_id'] if RUN_HYPERPARAMETER_TUNING else 'single_config'}\")\nprint(f\"  Image size   : {Config.IMG_SIZE}\")\nprint(f\"  Dropout      : {Config.DROPOUT_RATE}\")\nprint(f\"  Unfreeze top : {Config.UNFREEZE_TOP_N} layers\")\nprint(f\"  Max epochs   : {Config.PHASE2_EPOCHS}\")\nprint(f\"  Learning rate: {Config.PHASE2_LR} ({Config.PHASE2_LR_SCHEDULE} schedule)\")\nprint(f\"  Weight decay : {Config.WEIGHT_DECAY}\")\nprint(f\"  Batch size   : {Config.BATCH_SIZE}\")\nprint(\"=\" * 60)\n\nmodel.load_weights(os.path.join(Config.CHECKPOINT_DIR, \"best_phase1.weights.h5\"))\nconfigure_phase2(model, base_model)\ncallbacks_phase2 = build_callbacks(\"phase2\", lr_schedule=Config.PHASE2_LR_SCHEDULE,\n                                   total_epochs=Config.PHASE2_EPOCHS, peak_lr=Config.PHASE2_LR)\nhistory_phase2 = model.fit(\n    train_dataset,\n    epochs=Config.PHASE2_EPOCHS,\n    validation_data=val_dataset,\n    callbacks=callbacks_phase2,\n    steps_per_epoch=steps_per_epoch if Config.USE_OVERSAMPLING else None,\n    class_weight=active_class_weights(),\n    verbose=1,\n)\n\nlog_experiment(\n    phase=\"Phase 2 - Final Full-Data Fine-Tuning\",\n    lr=Config.PHASE2_LR,\n    batch_size=Config.BATCH_SIZE,\n    epochs_planned=Config.PHASE2_EPOCHS,\n    history=history_phase2,\n    notes=(\n        f\"Top {Config.UNFREEZE_TOP_N} layers trainable except BatchNorm; \"\n        f\"dropout={Config.DROPOUT_RATE}; class weights={Config.USE_CLASS_WEIGHTS}; \"\n        f\"oversampling={Config.USE_OVERSAMPLING}; class weight mode={Config.CLASS_WEIGHT_MODE}; \"\n        f\"lr schedule={Config.PHASE2_LR_SCHEDULE}; weight decay={Config.WEIGHT_DECAY}; batch={Config.BATCH_SIZE}; \"\n        f\"image_size={Config.IMG_SIZE}; \"\n        f\"selected pilot trial={selected_tuning_trial['trial_id'] if RUN_HYPERPARAMETER_TUNING else 'single_config'}\"\n    ),\n)\n\npd.DataFrame(history_phase2.history).to_csv(\n    os.path.join(Config.REPORTS_DIR, \"phase2_training_history.csv\"), index=False\n)\nplot_training_curves(history_phase2, \"Phase_2_Final_Full_Data_Fine_Tuning\")\nexperiment_log.to_csv(os.path.join(Config.REPORTS_DIR, \"experiment_log.csv\"), index=False)\nprint(\"[Phase 2] Final full-data training complete.\")\nmodel.load_weights(os.path.join(Config.CHECKPOINT_DIR, \"best_phase2.weights.h5\"))","metadata":{"id":"64ea143081634da796644510d780dea0","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:12:03.989846Z","iopub.execute_input":"2026-09-30T20:12:03.990109Z","iopub.status.idle":"2026-09-30T20:27:13.415302Z","shell.execute_reply.started":"2026-09-30T20:12:03.99009Z","shell.execute_reply":"2026-09-30T20:27:13.414226Z"}},"outputs":[],"execution_count":null},{"id":"42c74d5b","cell_type":"code","source":"# ============================================================\n# SECTION 6.5b — POST-TRAINING OVERFITTING DIAGNOSTICS\n# Summarize observed train/validation gaps and validation-loss movement.\n# These are diagnostic signals, not a clinical or statistical proof.\n# ============================================================\n\n\ndef _summarize_overfitting_signals(history, phase_name: str) -> Dict[str, Any]:\n    accuracy = history.history.get(\"accuracy\", history.history.get(\"categorical_accuracy\"))\n    validation_accuracy = history.history.get(\n        \"val_accuracy\", history.history.get(\"val_categorical_accuracy\")\n    )\n    training_loss = history.history.get(\"loss\")\n    validation_loss = history.history.get(\"val_loss\")\n    if any(values is None for values in (accuracy, validation_accuracy, training_loss, validation_loss)):\n        raise ValueError(f\"Training history is missing required accuracy/loss metrics for {phase_name}.\")\n\n    best_accuracy_index = int(np.nanargmax(validation_accuracy))\n    best_loss_index = int(np.nanargmin(validation_loss))  # restored by early stopping\n    accuracy_gap = float(accuracy[best_loss_index] - validation_accuracy[best_loss_index])\n    clean_accuracy = history.history.get(\"clean_train_accuracy\")\n    clean_gap = (\n        float(clean_accuracy[best_loss_index] - validation_accuracy[best_loss_index])\n        if clean_accuracy else float(\"nan\")\n    )\n    late_validation_loss_rise = float(validation_loss[-1] - validation_loss[best_loss_index])\n    if late_validation_loss_rise > 1e-6:\n        loss_observation = (\n            f\"Validation loss rose by {late_validation_loss_rise:.4f} after its minimum; \"\n            \"inspect the curves and accuracy gap for a possible overfitting signal.\"\n        )\n    else:\n        loss_observation = (\n            \"No late validation-loss rise was observed; this alone does not rule out overfitting.\"\n        )\n\n    return {\n        \"phase\": phase_name,\n        \"epochs_run\": len(validation_loss),\n        \"best_val_accuracy_epoch\": best_accuracy_index + 1,\n        \"train_accuracy_at_best_val_accuracy\": float(accuracy[best_accuracy_index]),\n        \"best_validation_accuracy\": float(validation_accuracy[best_accuracy_index]),\n        \"train_minus_validation_accuracy_gap_at_best_val_loss\": accuracy_gap,\n        \"clean_train_minus_validation_accuracy_gap\": clean_gap,\n        \"best_val_qwk\": float(np.nanmax(history.history[\"val_qwk\"])) if history.history.get(\"val_qwk\") else float(\"nan\"),\n        \"best_val_loss_epoch\": best_loss_index + 1,\n        \"minimum_validation_loss\": float(validation_loss[best_loss_index]),\n        \"final_validation_loss\": float(validation_loss[-1]),\n        \"final_minus_min_validation_loss\": late_validation_loss_rise,\n        \"curve_observation\": loss_observation,\n        \"overfitting_controls\": (\n            f\"dropout={Config.DROPOUT_RATE}; training-only augmentation; \"\n            f\"val_loss early stopping (patience {Config.ES_PATIENCE}); \"\n            f\"top {Config.UNFREEZE_TOP_N} layers fine-tuned; \"\n            + (f\"AdamW weight_decay={Config.WEIGHT_DECAY}\" if phase_name == \"Phase 2\" else \"Phase 1 backbone frozen\")\n        ),\n    }\n\n\noverfitting_diagnostics = pd.DataFrame(\n    [\n        _summarize_overfitting_signals(history_phase1, \"Phase 1\"),\n        _summarize_overfitting_signals(history_phase2, \"Phase 2\"),\n    ]\n)\noverfitting_diagnostics_path = Path(Config.REPORTS_DIR) / \"overfitting_diagnostics.csv\"\noverfitting_diagnostics.to_csv(overfitting_diagnostics_path, index=False)\nprint(\"Observed train/validation diagnostics (not proof of overfitting):\")\ndisplay(overfitting_diagnostics)\nprint(\"Saved curve interpretation inputs to:\", overfitting_diagnostics_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:29:10.697493Z","iopub.execute_input":"2026-09-30T20:29:10.698117Z","iopub.status.idle":"2026-09-30T20:29:10.723122Z","shell.execute_reply.started":"2026-09-30T20:29:10.698089Z","shell.execute_reply":"2026-09-30T20:29:10.72209Z"}},"outputs":[],"execution_count":null},{"id":"ec805c1a","cell_type":"code","source":"# ============================================================\n# SECTION 6.5c - LEARNING-RATE CURVE\n# Plots the learning rate used at each epoch of the final Phase 1 and\n# Phase 2 runs. Phase 1 uses ReduceLROnPlateau; Phase 2 uses either\n# ReduceLROnPlateau or warm-up + cosine (Config.PHASE2_LR_SCHEDULE).\n# Read from the saved history CSVs so it also works after a restart.\n# ============================================================\n\ndef _load_lr_history(csv_name: str, history=None) -> pd.Series:\n    frame = pd.read_csv(os.path.join(Config.REPORTS_DIR, csv_name))\n    for key in (\"learning_rate\", \"lr\"):\n        if key in frame.columns:\n            return frame[key].astype(float)\n    if history is not None:\n        for key in (\"learning_rate\", \"lr\"):\n            if key in history.history:\n                return pd.Series(history.history[key], dtype=float)\n    raise KeyError(f\"No learning-rate column found in {csv_name}; rerun training with build_callbacks().\")\n\n\n_lr_phase1 = _load_lr_history(\"phase1_training_history.csv\", globals().get(\"history_phase1\"))\n_lr_phase2 = _load_lr_history(\"phase2_training_history.csv\", globals().get(\"history_phase2\"))\nlearning_rate_history = pd.DataFrame({\n    \"phase\": [\"Phase 1\"] * len(_lr_phase1) + [\"Phase 2\"] * len(_lr_phase2),\n    \"phase_epoch\": list(range(1, len(_lr_phase1) + 1)) + list(range(1, len(_lr_phase2) + 1)),\n    \"overall_epoch\": list(range(1, len(_lr_phase1) + len(_lr_phase2) + 1)),\n    \"learning_rate\": pd.concat([_lr_phase1, _lr_phase2], ignore_index=True),\n})\nlearning_rate_history.to_csv(os.path.join(Config.REPORTS_DIR, \"learning_rate_history.csv\"), index=False)\n\nfig, ax = plt.subplots(figsize=(10, 4.5))\nfor _phase, _color in ((\"Phase 1\", \"#367C8A\"), (\"Phase 2\", \"#D8784B\")):\n    _rows = learning_rate_history[learning_rate_history[\"phase\"] == _phase]\n    ax.step(_rows[\"overall_epoch\"], _rows[\"learning_rate\"], where=\"mid\",\n            color=_color, marker=\"o\", markersize=4, label=_phase)\nax.axvline(len(_lr_phase1) + 0.5, color=\"grey\", linestyle=\"--\", linewidth=1, alpha=0.7)\nax.set_yscale(\"log\")\nax.set_title(f\"Learning Rate per Epoch - Final Training Run (Phase 2: {Config.PHASE2_LR_SCHEDULE})\")\nax.set_xlabel(\"Epoch (Phase 1 then Phase 2)\")\nax.set_ylabel(\"Learning rate (log scale)\")\nax.legend()\nax.grid(alpha=0.3, which=\"both\")\nfig.tight_layout()\nlr_curve_path = os.path.join(Config.REPORTS_DIR, \"lr_curve_final.png\")\nfig.savefig(lr_curve_path, dpi=150, bbox_inches=\"tight\")\nplt.show()\n\nfor _phase, _rows in learning_rate_history.groupby(\"phase\"):\n    _reductions = int((_rows[\"learning_rate\"].diff() < 0).sum())\n    _schedule = \"plateau\" if _phase == \"Phase 1\" else Config.PHASE2_LR_SCHEDULE\n    print(\n        f\"[LR Curve] {_phase} ({_schedule}): start {_rows['learning_rate'].iloc[0]:.2e} -> \"\n        f\"peak {_rows['learning_rate'].max():.2e} -> end {_rows['learning_rate'].iloc[-1]:.2e} \"\n        f\"({_reductions} epoch-to-epoch decreases)\"\n    )\nprint(f\"[LR Curve] Saved -> {lr_curve_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:29:52.003367Z","iopub.execute_input":"2026-09-30T20:29:52.00383Z","iopub.status.idle":"2026-09-30T20:29:52.665238Z","shell.execute_reply.started":"2026-09-30T20:29:52.003802Z","shell.execute_reply":"2026-09-30T20:29:52.664549Z"}},"outputs":[],"execution_count":null},{"id":"f8ee2eb8","cell_type":"code","source":"# ============================================================\n# SECTION 6.6 — FINAL RUN CONFIGURATION AUDIT\n# Confirms the full-data final fit used the selected pilot settings.\n# ============================================================\n\nphase2_rows = experiment_log[\n    experiment_log[\"phase\"].str.contains(\"Phase 2\", case=False, na=False)\n].copy()\nif phase2_rows.empty:\n    raise RuntimeError(\"No completed Phase 2 record is available.\")\n\nfinal_phase2_record = phase2_rows.iloc[-1]\nrecorded_notes = str(final_phase2_record[\"notes\"])\nphase2_config_matches = all([\n    np.isclose(float(final_phase2_record[\"learning_rate\"]), Config.PHASE2_LR),\n    int(final_phase2_record[\"batch_size\"]) == Config.BATCH_SIZE,\n    f\"Top {Config.UNFREEZE_TOP_N} layers trainable\" in recorded_notes,\n    f\"dropout={Config.DROPOUT_RATE}\" in recorded_notes,\n    f\"class weights={Config.USE_CLASS_WEIGHTS}\" in recorded_notes,\n    f\"image_size={Config.IMG_SIZE}\" in recorded_notes,\n])\nif RUN_HYPERPARAMETER_TUNING:\n    phase2_config_matches = phase2_config_matches and (\n        f\"selected pilot trial={selected_tuning_trial['trial_id']}\" in recorded_notes\n    )\n\nbest_val_accuracy_epoch = int(np.argmax(history_phase2.history[\"val_accuracy\"])) + 1\nbest_val_loss_epoch = int(np.argmin(history_phase2.history[\"val_loss\"])) + 1\nfinal_run_config = {\n    \"image_size\": Config.IMG_SIZE,\n    \"batch_size\": Config.BATCH_SIZE,\n    \"phase1_learning_rate\": Config.PHASE1_LR,\n    \"phase2_learning_rate\": Config.PHASE2_LR,\n    \"phase2_unfreeze_top_n\": Config.UNFREEZE_TOP_N,\n    \"dropout_rate\": Config.DROPOUT_RATE,\n    \"class_weights\": Config.USE_CLASS_WEIGHTS,\n    \"oversampling\": Config.USE_OVERSAMPLING,\n    \"class_weight_mode\": Config.CLASS_WEIGHT_MODE,\n    \"phase2_lr_schedule\": Config.PHASE2_LR_SCHEDULE,\n    \"weight_decay\": Config.WEIGHT_DECAY,\n    \"early_stopping_monitor\": \"val_loss\",\n    \"selected_pilot_trial\": selected_tuning_trial[\"trial_id\"] if RUN_HYPERPARAMETER_TUNING else None,\n}\n\nprint(\"Final full-data model configuration:\")\nfor key, value in final_run_config.items():\n    print(f\"  {key}: {value}\")\nprint(f\"Best final Phase 2 validation-accuracy epoch: {best_val_accuracy_epoch}\")\nprint(f\"Best final Phase 2 validation-loss epoch    : {best_val_loss_epoch}\")\nprint(f\"Final Phase 2 record matches selected settings: {phase2_config_matches}\")\nprint(\"Final weights loaded from:\", os.path.join(Config.CHECKPOINT_DIR, \"best_phase2.weights.h5\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:30:03.308409Z","iopub.execute_input":"2026-09-30T20:30:03.309169Z","iopub.status.idle":"2026-09-30T20:30:03.319772Z","shell.execute_reply.started":"2026-09-30T20:30:03.30914Z","shell.execute_reply":"2026-09-30T20:30:03.318856Z"}},"outputs":[],"execution_count":null},{"id":"70a6273bd77b46fe9f7609251eda07be","cell_type":"markdown","source":"## Section 7: Evaluation\n\n**Goal:** Conduct a rigorous clinical and statistical evaluation on the held-out test set\n(15% of the total dataset, never seen during training or validation).\n\n> **Headline numbers.** Section 7.2 reports both decision rules on the same test predictions: **\"Argmax (plain)\"** is the headline **accuracy**, and **\"Validation-QWK thresholds\"** (cut points chosen on the validation split) is the headline **QWK**. Both rows, with macro-F1, are saved to `test_decision_rule_summary.csv`; thresholding can raise QWK while lowering accuracy, so both are always reported.\n\n> **Small test set.** The APTOS 2019 test split holds 550 images, with only 29 Severe cases (per-class counts are printed in Section 4.3). Treat per-class test metrics — especially Severe recall — as having wide uncertainty; a handful of images changes them by several percentage points.\n\n### Clinical Evaluation Metric: Why Quadratic Weighted Kappa (QWK)?\n\nStandard **classification accuracy** treats all errors identically:\n- Misclassifying **No DR (Class 0)** as **Mild NPDR (Class 1)** is penalized by 1 error.\n- Misclassifying **No DR (Class 0)** as **Proliferative DR (Class 4)** is also penalized by 1 error.\n\nClinically, this is fundamentally flawed. Diabetic Retinopathy staging is an **ordinal scale**\n(ordered severity levels from 0 to 4):\n1. Confusing adjacent stages (e.g., Mild vs. Moderate) has minor clinical consequence because both\n   require clinical follow-up and monitoring.\n2. Confusing distant stages (e.g., predicting \"No DR\" when the patient has sight-threatening\n   \"Severe\" or \"Proliferative DR\") is catastrophic, delaying emergency laser photocoagulation or\n   anti-VEGF injections and risking irreversible vision loss.\n\n**Quadratic Weighted Kappa (QWK)** measures inter-rater agreement between the ground truth and\nmodel predictions on an ordinal scale, penalizing misclassifications quadratically according to\nthe distance $|i - j|^2$:\n$$\\kappa = 1 - \\frac{\\sum_{i,j} w_{ij} O_{ij}}{\\sum_{i,j} w_{ij} E_{ij}}, \\quad \\text{where } w_{ij} = \\frac{(i - j)^2}{(N - 1)^2}$$\n- An error between Class 0 and Class 1 carries a penalty weight of $(0-1)^2 = 1$.\n- An error between Class 0 and Class 4 carries a penalty weight of $(0-4)^2 = 16$.\n\nThis mirrors the official metric used in the landmark Kaggle Diabetic Retinopathy and APTOS 2019\ncompetitions, providing the gold-standard benchmark for automated retinal image grading.","metadata":{"id":"70a6273bd77b46fe9f7609251eda07be","language":"markdown"}},{"id":"6136dbed","cell_type":"code","source":"def predict_probs_tta(eval_model, images):\n    probabilities = eval_model(images, training=False)\n    probabilities += eval_model(tf.image.flip_left_right(images), training=False)\n    probabilities += eval_model(tf.image.flip_up_down(images), training=False)\n    probabilities += eval_model(\n        tf.image.flip_left_right(tf.image.flip_up_down(images)), training=False\n    )\n    return (probabilities / 4.0).numpy()\n\n\ndef _val_metrics(use_tta):  # CHANGED: compare both accuracy and ordinal QWK\n    y_true_batches = []  # CHANGED: collect validation truth alongside predictions\n    y_pred_batches = []  # CHANGED: collect validation predictions\n    for images, labels in val_dataset:\n        if use_tta:\n            probabilities = predict_probs_tta(model, images)\n        else:\n            probabilities = model(images, training=False).numpy()\n        y_true_batches.append(np.argmax(labels.numpy(), axis=1))  # CHANGED: use this batch's labels\n        y_pred_batches.append(np.argmax(probabilities, axis=1))  # CHANGED: collect predictions\n    y_true = np.concatenate(y_true_batches)  # CHANGED: aggregate validation labels\n    y_pred = np.concatenate(y_pred_batches)  # CHANGED: aggregate validation predictions\n    accuracy = float(np.mean(y_true == y_pred))  # CHANGED: validation accuracy\n    qwk = float(cohen_kappa_score(y_true, y_pred, weights=\"quadratic\"))  # CHANGED: validation QWK\n    return accuracy, qwk  # CHANGED: return both selection metrics\n\n\nplain_val_accuracy, plain_val_qwk = _val_metrics(False)  # CHANGED: score plain validation inference\ntta_val_accuracy, tta_val_qwk = _val_metrics(True)  # CHANGED: score four-view TTA on validation only\nUSE_TTA_FOR_EVALUATION = tta_val_qwk > plain_val_qwk  # CHANGED: select TTA only when validation QWK improves\nprint(\n    f\"Validation plain — accuracy: {plain_val_accuracy:.4f}, QWK: {plain_val_qwk:.4f} | \"\n    f\"TTA — accuracy: {tta_val_accuracy:.4f}, QWK: {tta_val_qwk:.4f} | \"\n    f\"Selected: {'TTA' if USE_TTA_FOR_EVALUATION else 'plain'}\"\n)  # CHANGED: report both metrics and validation-only selection","metadata":{"language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:30:30.22248Z","iopub.execute_input":"2026-09-30T20:30:30.223057Z","iopub.status.idle":"2026-09-30T20:32:04.100353Z","shell.execute_reply.started":"2026-09-30T20:30:30.22303Z","shell.execute_reply":"2026-09-30T20:32:04.09967Z"}},"outputs":[],"execution_count":null},{"id":"69fc0190","cell_type":"code","source":"# ============================================================\n# SECTION 7.0b — VALIDATION-ONLY QWK THRESHOLD CALIBRATION\n# Thresholds are selected from validation predictions, never test labels.\n# ============================================================\n\nfrom itertools import combinations\n\n\ndef apply_qwk_thresholds(probabilities: np.ndarray, thresholds: np.ndarray) -> np.ndarray:\n    \"\"\"Map expected ordinal grades to classes using increasing cut points.\"\"\"\n    probabilities = np.asarray(probabilities, dtype=np.float64)\n    thresholds = np.asarray(thresholds, dtype=np.float64)\n    if probabilities.ndim != 2 or probabilities.shape[1] != Config.NUM_CLASSES:\n        raise ValueError(\"probabilities must have shape (n_samples, NUM_CLASSES).\")\n    if thresholds.size != Config.NUM_CLASSES - 1 or np.any(np.diff(thresholds) <= 0):\n        raise ValueError(\"Provide NUM_CLASSES - 1 strictly increasing thresholds.\")\n    expected_grade = probabilities @ np.arange(Config.NUM_CLASSES, dtype=np.float64)\n    return np.digitize(expected_grade, thresholds).astype(int)\n\n\ndef tune_qwk_thresholds(\n    y_true: np.ndarray,\n    probabilities: np.ndarray,\n) -> tuple[np.ndarray, float]:\n    \"\"\"Choose ordinal cut points by maximizing QWK over validation predictions.\"\"\"\n    probabilities = np.asarray(probabilities, dtype=np.float64)\n    y_true = np.asarray(y_true, dtype=int)\n    expected_grade = probabilities @ np.arange(Config.NUM_CLASSES, dtype=np.float64)\n    quantile_candidates = np.quantile(expected_grade, np.linspace(0.05, 0.95, 10))\n    default_candidates = np.arange(0.5, Config.NUM_CLASSES - 0.5, 1.0)\n    candidates = np.unique(np.concatenate([quantile_candidates, default_candidates]))\n\n    best_thresholds = None\n    best_qwk = -np.inf\n    for trial_thresholds in combinations(candidates, Config.NUM_CLASSES - 1):\n        predictions = apply_qwk_thresholds(probabilities, np.asarray(trial_thresholds))\n        score = cohen_kappa_score(y_true, predictions, weights=\"quadratic\")\n        if np.isfinite(score) and score > best_qwk:\n            best_thresholds = np.asarray(trial_thresholds, dtype=np.float64)\n            best_qwk = float(score)\n\n    if best_thresholds is None:\n        raise RuntimeError(\"Could not find valid QWK thresholds from the validation predictions.\")\n    return best_thresholds, best_qwk\n\n\nval_probability_batches = []\nfor images, _ in val_dataset:\n    if USE_TTA_FOR_EVALUATION:\n        val_probabilities = predict_probs_tta(model, images)\n    else:\n        val_probabilities = model(images, training=False).numpy()\n    val_probability_batches.append(val_probabilities)\n\nval_probabilities = np.concatenate(val_probability_batches, axis=0)\nval_true = val_df[\"label\"].to_numpy(dtype=int)\nif len(val_true) != len(val_probabilities):\n    raise ValueError(\"Validation labels and predictions have different lengths.\")\n\nQWK_THRESHOLDS, calibrated_val_qwk = tune_qwk_thresholds(val_true, val_probabilities)\nargmax_val_predictions = np.argmax(val_probabilities, axis=1)\ncalibrated_val_predictions = apply_qwk_thresholds(val_probabilities, QWK_THRESHOLDS)\nprint(f\"Validation argmax QWK      : {cohen_kappa_score(val_true, argmax_val_predictions, weights='quadratic'):.4f}\")\nprint(f\"Validation calibrated QWK  : {calibrated_val_qwk:.4f}\")\nprint(f\"Validation argmax accuracy : {np.mean(val_true == argmax_val_predictions):.4f}\")\nprint(f\"Validation calibrated acc. : {np.mean(val_true == calibrated_val_predictions):.4f}\")\nprint(\"Validation-selected QWK thresholds:\", np.round(QWK_THRESHOLDS, 4).tolist())\nprint(\"Freeze these thresholds before evaluating the held-out test set.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:32:12.864666Z","iopub.execute_input":"2026-09-30T20:32:12.865123Z","iopub.status.idle":"2026-09-30T20:32:32.84789Z","shell.execute_reply.started":"2026-09-30T20:32:12.865094Z","shell.execute_reply":"2026-09-30T20:32:32.846953Z"}},"outputs":[],"execution_count":null},{"id":"2aa1ad03","cell_type":"code","source":"# ============================================================\n# SECTION 7.0c — CLASS-WEIGHT VALIDATION EVIDENCE\n# For each arm, restart the kernel and rerun training through this cell:\n# once with USE_CLASS_WEIGHTS=False, then once with it True.\n# Keep the seed, manifests, and other settings fixed. Stop before\n# Section 7.1 and evaluate the held-out test set only once after selection.\n# ============================================================\n\nimport hashlib\n\n_expected_weight_note = f\"class weights enabled={Config.USE_CLASS_WEIGHTS}\"\n_phase1_records = experiment_log[\n    experiment_log[\"phase\"].str.contains(\"Phase 1\", case=False, na=False)\n]\n_phase2_records = experiment_log[\n    experiment_log[\"phase\"].str.contains(\"Phase 2\", case=False, na=False)\n]\nif _phase1_records.empty or _phase2_records.empty:\n    raise RuntimeError(\"Complete both classifier training phases before recording evidence.\")\nif _expected_weight_note not in str(_phase1_records.iloc[-1][\"notes\"]):\n    raise RuntimeError(\"Phase 1 was not trained with the currently configured class-weight setting.\")\nif f\"class weights={Config.USE_CLASS_WEIGHTS}\" not in str(_phase2_records.iloc[-1][\"notes\"]):\n    raise RuntimeError(\"Phase 2 was not trained with the currently configured class-weight setting.\")\n\n_validation_true = val_df[\"label\"].to_numpy(dtype=int)\n_validation_probability_batches = []\nfor _images, _ in val_dataset:\n    _validation_probability_batches.append(model(_images, training=False).numpy())\n_validation_probabilities = np.concatenate(_validation_probability_batches, axis=0)\n_validation_pred = np.argmax(_validation_probabilities, axis=1)\nif len(_validation_true) != len(_validation_pred):\n    raise ValueError(\"Validation labels and predictions have different lengths.\")\n\n_validation_report = classification_report(\n    _validation_true,\n    _validation_pred,\n    labels=list(range(Config.NUM_CLASSES)),\n    target_names=Config.CLASS_NAMES,\n    output_dict=True,\n    zero_division=0,\n)\n_validation_accuracy = float(np.mean(_validation_true == _validation_pred))\n_validation_qwk = float(\n    cohen_kappa_score(_validation_true, _validation_pred, weights=\"quadratic\")\n)\n\n_manifest_rows = pd.concat(\n    [\n        frame[[\"patient_id\", \"duplicate_group\", \"label\"]].assign(split=split_name)\n        for split_name, frame in (\n            (\"train\", train_df),\n            (\"validation\", val_df),\n            (\"test\", test_df),\n        )\n    ],\n    ignore_index=True,\n)\n_manifest_rows = _manifest_rows.astype(\"string\").fillna(\"\").sort_values(\n    [\"split\", \"patient_id\", \"duplicate_group\", \"label\"]\n)\n_split_fingerprint = hashlib.sha256(\n    _manifest_rows.to_csv(index=False, lineterminator=\"\\n\").encode(\"utf-8\")\n).hexdigest()\n\n_validation_records = []\nfor _class_id, _class_name in enumerate(Config.CLASS_NAMES):\n    _class_metrics = _validation_report[_class_name]\n    _validation_records.append(\n        {\n            \"run_id\": _RUN_ID,\n            \"split_fingerprint\": _split_fingerprint,\n            \"class_weights_enabled\": bool(Config.USE_CLASS_WEIGHTS),\n            \"seed\": Config.SEED,\n            \"image_size\": Config.IMG_SIZE,\n            \"batch_size\": Config.BATCH_SIZE,\n            \"phase1_learning_rate\": Config.PHASE1_LR,\n            \"phase2_learning_rate\": Config.PHASE2_LR,\n            \"phase1_epochs\": Config.PHASE1_EPOCHS,\n            \"phase2_epochs\": Config.PHASE2_EPOCHS,\n            \"dropout_rate\": Config.DROPOUT_RATE,\n            \"phase2_unfreeze_top_n\": Config.UNFREEZE_TOP_N,\n            \"class_id\": _class_id,\n            \"class_name\": _class_name,\n            \"validation_accuracy\": _validation_accuracy,\n            \"validation_qwk\": _validation_qwk,\n            \"precision\": float(_class_metrics[\"precision\"]),\n            \"recall\": float(_class_metrics[\"recall\"]),\n            \"f1\": float(_class_metrics[\"f1-score\"]),\n            \"support\": int(_class_metrics[\"support\"]),\n        }\n    )\n\n_validation_evidence = pd.DataFrame(_validation_records)\n_run_evidence_path = Path(Config.REPORTS_DIR) / \"class_weight_validation_evidence.csv\"\n_validation_evidence.to_csv(_run_evidence_path, index=False)\n\n_comparison_path = Path(Config.REPORTS_DIR).parent.parent / \"class_weight_validation_runs.csv\"\nif _comparison_path.is_file():\n    _all_validation_evidence = pd.read_csv(_comparison_path)\n    _all_validation_evidence = _all_validation_evidence[\n        _all_validation_evidence[\"run_id\"] != _RUN_ID\n    ]\n    _all_validation_evidence = pd.concat(\n        [_all_validation_evidence, _validation_evidence], ignore_index=True\n    )\nelse:\n    _all_validation_evidence = _validation_evidence.copy()\n_all_validation_evidence.to_csv(_comparison_path, index=False)\n_all_validation_evidence.to_csv(\n    Path(Config.REPORTS_DIR) / \"class_weight_validation_comparison.csv\", index=False\n)\n\nprint(\n    f\"Validation only | class weights={Config.USE_CLASS_WEIGHTS} | \"\n    f\"accuracy={_validation_accuracy:.4f} | QWK={_validation_qwk:.4f}\"\n)\nprint(\"Per-class recall and support:\")\nprint(\n    _validation_evidence[[\"class_id\", \"class_name\", \"recall\", \"support\"]]\n    .to_string(index=False)\n)\nprint(f\"Saved this run's metrics -> {_run_evidence_path}\")\nprint(f\"Saved accumulated comparison -> {_comparison_path}\")\n\n_latest_run_ids = (\n    _all_validation_evidence.sort_values(\"run_id\")\n    .groupby(\"class_weights_enabled\")[\"run_id\"]\n    .last()\n)\nif {False, True}.issubset(set(_latest_run_ids.index)):\n    _off_run_id = _latest_run_ids.loc[False]\n    _on_run_id = _latest_run_ids.loc[True]\n    _run_summaries = _all_validation_evidence.drop_duplicates(\"run_id\").set_index(\"run_id\")\n    _off_summary = _run_summaries.loc[_off_run_id]\n    _on_summary = _run_summaries.loc[_on_run_id]\n    _matched_fields = [\n        \"split_fingerprint\",\n        \"seed\",\n        \"image_size\",\n        \"batch_size\",\n        \"phase1_learning_rate\",\n        \"phase2_learning_rate\",\n        \"phase1_epochs\",\n        \"phase2_epochs\",\n        \"dropout_rate\",\n        \"phase2_unfreeze_top_n\",\n    ]\n    _matched = all(_off_summary[field] == _on_summary[field] for field in _matched_fields)\n    if _matched:\n        print(\"Matched validation comparison (class weights ON minus OFF):\")\n        print(\n            f\"  Accuracy delta: {_on_summary['validation_accuracy'] - _off_summary['validation_accuracy']:+.4f}\"\n        )\n        print(\n            f\"  QWK delta     : {_on_summary['validation_qwk'] - _off_summary['validation_qwk']:+.4f}\"\n        )\n        _recall_comparison = _all_validation_evidence[\n            _all_validation_evidence[\"run_id\"].isin([_off_run_id, _on_run_id])\n        ].pivot(index=[\"class_id\", \"class_name\"], columns=\"class_weights_enabled\", values=\"recall\")\n        _recall_comparison[\"recall_delta_on_minus_off\"] = (\n            _recall_comparison[True] - _recall_comparison[False]\n        )\n        print(_recall_comparison.to_string())\n    else:\n        print(\n            \"Both weight settings are logged, but their split fingerprint, seed, or \"\n            \"training settings differ; this is not a matched comparison.\"\n        )\nelse:\n    print(\n        \"Paired comparison pending: run the other class-weight setting with the same \"\n        \"saved splits, seed, and training settings.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:32:45.650413Z","iopub.execute_input":"2026-09-30T20:32:45.65087Z","iopub.status.idle":"2026-09-30T20:33:04.724505Z","shell.execute_reply.started":"2026-09-30T20:32:45.65084Z","shell.execute_reply":"2026-09-30T20:33:04.723896Z"}},"outputs":[],"execution_count":null},{"id":"38fa4ec8deaf4630b570adb2aa63a4ec","cell_type":"code","source":"# ============================================================\n# SECTION 7.1 — TEST INFERENCE, SAVED PREDICTIONS & BASELINE\n# Run one model inference pass on test; all test metrics derive from this NPZ.\n# The majority-class baseline predicts the most common TRAIN label.\n# ============================================================\n\ndef generate_test_predictions(\n    eval_model: keras.Model,\n    test_ds: tf.data.Dataset,\n) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:\n    \"\"\"Return test labels, raw argmax predictions, and class probabilities.\"\"\"\n    print(\"[Evaluation] Running inference on the held-out test set once...\")\n    y_true_batches = []\n    y_prob_batches = []\n\n    for images, labels in tqdm(test_ds, desc=\"Evaluating test set\"):\n        if USE_TTA_FOR_EVALUATION:\n            probabilities = predict_probs_tta(eval_model, images)\n        else:\n            probabilities = eval_model(images, training=False).numpy()\n        y_true_batches.append(np.argmax(labels.numpy(), axis=-1))\n        y_prob_batches.append(probabilities)\n\n    probabilities = np.concatenate(y_prob_batches, axis=0)\n    y_true = np.concatenate(y_true_batches, axis=0)\n    y_pred_argmax = np.argmax(probabilities, axis=1)\n    print(f\"[Evaluation] Inference complete for {len(y_true):,} test samples.\")\n    return y_true, y_pred_argmax, probabilities\n\n\n# Define the naive baseline from training-label frequencies, never test-label frequencies.\n_train_class_counts = train_df[\"label\"].value_counts().sort_index()\nMAJORITY_CLASS = int(_train_class_counts.idxmax())\nprint(\n    f\"[Baseline] Majority class from training split: {MAJORITY_CLASS} \"\n    f\"({Config.CLASS_NAMES[MAJORITY_CLASS]}), \"\n    f\"{int(_train_class_counts.loc[MAJORITY_CLASS]):,}/{len(train_df):,} training images.\"\n)\n\n# Preserve raw argmax metrics, then apply thresholds fixed using validation only.\ny_test_true, y_test_pred_argmax, y_test_prob = generate_test_predictions(model, test_dataset)\ny_test_pred = apply_qwk_thresholds(y_test_prob, QWK_THRESHOLDS)\ny_test_pred_majority = np.full(y_test_true.shape, MAJORITY_CLASS, dtype=int)\ntest_predictions_path = os.path.join(Config.REPORTS_DIR, \"test_predictions.npz\")\nnp.savez_compressed(\n    test_predictions_path,\n    y_true=y_test_true,\n    y_pred=y_test_pred,\n    y_pred_argmax=y_test_pred_argmax,\n    y_pred_majority=y_test_pred_majority,\n    majority_class=np.asarray(MAJORITY_CLASS),\n    softmax_probabilities=y_test_prob,\n    qwk_thresholds=QWK_THRESHOLDS,\n    use_tta=np.asarray(USE_TTA_FOR_EVALUATION),\n    filepaths=np.asarray(test_df[\"filepath\"].tolist()),\n)\n\nwith np.load(test_predictions_path) as saved_test_predictions:\n    y_test_true = saved_test_predictions[\"y_true\"]\n    y_test_pred = saved_test_predictions[\"y_pred\"]\n    y_test_pred_argmax = saved_test_predictions[\"y_pred_argmax\"]\n    y_test_pred_majority = saved_test_predictions[\"y_pred_majority\"]\n    y_test_prob = saved_test_predictions[\"softmax_probabilities\"]\n    MAJORITY_CLASS = int(saved_test_predictions[\"majority_class\"].item())\n    test_filepaths = saved_test_predictions[\"filepaths\"].astype(str)\n\nif not np.array_equal(y_test_pred_argmax, np.argmax(y_test_prob, axis=1)):\n    raise AssertionError(\"Saved argmax predictions do not match saved softmax probabilities.\")\nif not np.array_equal(y_test_pred_majority, np.full_like(y_test_true, MAJORITY_CLASS)):\n    raise AssertionError(\"Saved majority-baseline predictions do not match the saved baseline class.\")\nif len(test_filepaths) != len(y_test_true):\n    raise ValueError(\"Saved test paths and prediction arrays have different lengths.\")\n\n# Compare all decision rules from the same persisted truth/prediction arrays.\n_test_prediction_sets = {\n    \"majority_class_baseline\": y_test_pred_majority,\n    \"efficientnet_argmax\": y_test_pred_argmax,\n    \"efficientnet_validation_qwk_thresholds\": y_test_pred,\n}\n_test_comparison_records = []\nfor _method, _predictions in _test_prediction_sets.items():\n    _method_report = classification_report(\n        y_test_true,\n        _predictions,\n        labels=list(range(Config.NUM_CLASSES)),\n        target_names=Config.CLASS_NAMES,\n        output_dict=True,\n        zero_division=0,\n    )\n    _test_comparison_records.append({\n        \"method\": _method,\n        \"majority_class\": MAJORITY_CLASS,\n        \"accuracy\": float(np.mean(y_test_true == _predictions)),\n        \"quadratic_weighted_kappa\": float(\n            cohen_kappa_score(y_test_true, _predictions, weights=\"quadratic\")\n        ),\n        \"macro_f1\": float(_method_report[\"macro avg\"][\"f1-score\"]),\n        \"weighted_f1\": float(_method_report[\"weighted avg\"][\"f1-score\"]),\n        **{\n            f\"recall_{class_name.replace(' ', '_')}\": float(_method_report[class_name][\"recall\"])\n            for class_name in Config.CLASS_NAMES\n        },\n    })\n\n_test_comparison_df = pd.DataFrame(_test_comparison_records)\n_test_comparison_path = Path(Config.REPORTS_DIR) / \"model_vs_majority_baseline.csv\"\n_test_comparison_df.to_csv(_test_comparison_path, index=False)\nprint(\"[Evaluation] Metrics derived from the saved one-pass test prediction artifact:\")\ndisplay(_test_comparison_df)\nprint(f\"[Evaluation] Saved model/baseline comparison -> {_test_comparison_path}\")\nprint(f\"[Evaluation] Saved one-pass test predictions -> {test_predictions_path}\")","metadata":{"id":"38fa4ec8deaf4630b570adb2aa63a4ec","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:33:11.737577Z","iopub.execute_input":"2026-09-30T20:33:11.738003Z","iopub.status.idle":"2026-09-30T20:33:40.46091Z","shell.execute_reply.started":"2026-09-30T20:33:11.737975Z","shell.execute_reply":"2026-09-30T20:33:40.460258Z"}},"outputs":[],"execution_count":null},{"id":"012bbef2afc54d5b874ad21bba16d47a","cell_type":"code","source":"# ============================================================\n# SECTION 7.2 — STANDARD ACCURACY & QUADRATIC WEIGHTED KAPPA\n# ============================================================\n\ndef calculate_clinical_metrics(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n) -> Dict[str, float]:\n    \"\"\"Calculate accuracy and quadratic weighted kappa for ordinal grades.\"\"\"\n    accuracy = float(np.mean(y_true == y_pred))\n    qwk = float(cohen_kappa_score(y_true, y_pred, weights=\"quadratic\"))\n\n    print(\"\\n\" + \"=\" * 65)\n    print(\"  CLINICAL & STATISTICAL METRICS SUMMARY\")\n    print(\"=\" * 65)\n    print(f\"  Total Test Samples          : {len(y_true):,}\")\n    print(f\"  Standard Accuracy           : {accuracy:.4f}  ({accuracy * 100:.2f}%)\")\n    print(f\"  Quadratic Weighted Kappa    : {qwk:.4f}\")\n    print(\"-\" * 65)\n\n    if qwk >= 0.80:\n        agreement_desc = \"Very strong ordinal agreement; not evidence of specialist equivalence\"\n    elif qwk >= 0.60:\n        agreement_desc = \"Substantial ordinal agreement\"\n    elif qwk >= 0.40:\n        agreement_desc = \"Moderate ordinal agreement\"\n    else:\n        agreement_desc = \"Fair or slight ordinal agreement\"\n\n    print(f\"  QWK interpretation          : {agreement_desc}\")\n    print(\"=\" * 65 + \"\\n\")\n    return {\"accuracy\": accuracy, \"quadratic_weighted_kappa\": qwk}\n\n\nargmax_eval_metrics = calculate_clinical_metrics(y_test_true, y_test_pred_argmax)  # CHANGED: report plain argmax metrics\nprint(\"Thresholded predictions:\")  # CHANGED: separate calibrated decision rule\neval_metrics = calculate_clinical_metrics(y_test_true, y_test_pred)  # CHANGED: keep calibrated metrics for existing downstream cells\nprint(\"Thresholding may raise QWK while lowering plain accuracy.\")  # CHANGED: disclose calibration trade-off\n\n# Headline table: both decision rules on the same saved test predictions.\nfrom sklearn.metrics import f1_score\n\n_labels = list(range(Config.NUM_CLASSES))\nheadline_metrics = pd.DataFrame([\n    {\n        \"decision_rule\": \"Argmax (plain)\",\n        \"accuracy\": argmax_eval_metrics[\"accuracy\"],\n        \"qwk\": argmax_eval_metrics[\"quadratic_weighted_kappa\"],\n        \"macro_f1\": float(f1_score(y_test_true, y_test_pred_argmax, labels=_labels, average=\"macro\", zero_division=0)),\n        \"headline_for\": \"accuracy\",\n    },\n    {\n        \"decision_rule\": \"Validation-QWK thresholds\",\n        \"accuracy\": eval_metrics[\"accuracy\"],\n        \"qwk\": eval_metrics[\"quadratic_weighted_kappa\"],\n        \"macro_f1\": float(f1_score(y_test_true, y_test_pred, labels=_labels, average=\"macro\", zero_division=0)),\n        \"headline_for\": \"QWK\",\n    },\n])\nheadline_metrics_path = os.path.join(Config.REPORTS_DIR, \"test_decision_rule_summary.csv\")\nheadline_metrics.to_csv(headline_metrics_path, index=False)\nprint(\"\\nTEST SET — both decision rules (thresholds chosen on validation only):\")\nfor _row in headline_metrics.itertuples():\n    print(f\"  {_row.decision_rule:<27} — accuracy {_row.accuracy:.4f} / QWK {_row.qwk:.4f} / \"\n          f\"macro-F1 {_row.macro_f1:.4f}   (headline {_row.headline_for})\")\nprint(f\"[Evaluation] Saved -> {headline_metrics_path}\")\n\n\n# ------------------------------------------------------------\n# TEST LOSS - one number, computed from the saved one-pass test\n# probabilities with the same loss the model was trained on\n# (categorical cross-entropy, label_smoothing=0.1). This equals what\n# model.evaluate() reports for plain inference, without a second test pass.\n# ------------------------------------------------------------\n_y_test_one_hot = tf.one_hot(y_test_true, Config.NUM_CLASSES)\nTEST_LOSS = float(\n    keras.losses.CategoricalCrossentropy(label_smoothing=0.1)(_y_test_one_hot, y_test_prob)\n)\nTEST_LOSS_UNSMOOTHED = float(\n    keras.losses.CategoricalCrossentropy()(_y_test_one_hot, y_test_prob)\n)\nprint(f\"Test Loss: {TEST_LOSS:.4f}  (training loss definition, label smoothing 0.1)\")\nprint(f\"Test cross-entropy without label smoothing: {TEST_LOSS_UNSMOOTHED:.4f}\")\nif USE_TTA_FOR_EVALUATION:\n    print(\"Note: computed on the validation-selected 4-view TTA probabilities.\")\npd.DataFrame([{\n    \"test_loss\": TEST_LOSS,\n    \"test_cross_entropy_unsmoothed\": TEST_LOSS_UNSMOOTHED,\n    \"label_smoothing\": 0.1,\n    \"use_tta\": bool(USE_TTA_FOR_EVALUATION),\n    \"test_samples\": int(len(y_test_true)),\n    \"argmax_accuracy\": argmax_eval_metrics[\"accuracy\"],\n    \"argmax_qwk\": argmax_eval_metrics[\"quadratic_weighted_kappa\"],\n    \"thresholded_accuracy\": eval_metrics[\"accuracy\"],\n    \"thresholded_qwk\": eval_metrics[\"quadratic_weighted_kappa\"],\n}]).to_csv(os.path.join(Config.REPORTS_DIR, \"test_loss_and_metrics.csv\"), index=False)\nprint(f\"[Evaluation] Saved -> {os.path.join(Config.REPORTS_DIR, 'test_loss_and_metrics.csv')}\")\n","metadata":{"id":"012bbef2afc54d5b874ad21bba16d47a","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:33:57.825434Z","iopub.execute_input":"2026-09-30T20:33:57.825868Z","iopub.status.idle":"2026-09-30T20:33:57.860449Z","shell.execute_reply.started":"2026-09-30T20:33:57.82584Z","shell.execute_reply":"2026-09-30T20:33:57.859794Z"}},"outputs":[],"execution_count":null},{"id":"a9bff4f061694a008eb45b21f2c113a0","cell_type":"code","source":"# ============================================================\n# SECTION 7.3 — DETAILED CLASSIFICATION REPORT\n#\n# Generates per-class Precision, Recall, and F1-Score along with\n# Macro and Weighted Averages.\n#\n# Note on Imbalanced Datasets:\n#   - Macro Average: Unweighted mean across all 5 classes — treats\n#     rare Proliferative DR equally with frequent No DR.\n#   - Weighted Average: Weighted by class support — dominated by No DR.\n# In clinical disease staging, the Macro Average and minority class\n# recalls (classes 3 & 4) are the most critical safety indicators.\n# ============================================================\n\ndef generate_classification_summary(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    class_names: List[str] = Config.CLASS_NAMES,\n    save_path: Optional[str] = None,\n) -> pd.DataFrame:\n    \"\"\"\n    Generate and display a structured classification report DataFrame.\n\n    Args:\n        y_true: True integer labels.\n        y_pred: Predicted integer labels.\n        class_names: List of class names ordered by class index 0-4.\n        save_path: Optional path to save CSV report.\n\n    Returns:\n        Formatted pandas DataFrame containing per-class and aggregated metrics.\n    \"\"\"\n    report_dict = classification_report(\n        y_true,\n        y_pred,\n        target_names=class_names,\n        output_dict=True,\n        zero_division=0,\n    )\n    report_df = pd.DataFrame(report_dict).transpose()\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\"  PER-STAGE CLASSIFICATION REPORT (PRECISION, RECALL, F1-SCORE)\")\n    print(\"=\" * 70)\n    print(report_df.round(4).to_string())\n    print(\"=\" * 70)\n\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        report_df.to_csv(save_path)\n        print(f\"[Evaluation] Saved classification report -> {save_path}\")\n\n    return report_df\n\n\nreport_df = generate_classification_summary(\n    y_test_true,\n    y_test_pred,\n    save_path=os.path.join(Config.REPORTS_DIR, \"classification_report.csv\"),\n)","metadata":{"id":"a9bff4f061694a008eb45b21f2c113a0","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:34:52.176188Z","iopub.execute_input":"2026-09-30T20:34:52.176615Z","iopub.status.idle":"2026-09-30T20:34:52.199194Z","shell.execute_reply.started":"2026-09-30T20:34:52.176588Z","shell.execute_reply":"2026-09-30T20:34:52.198233Z"}},"outputs":[],"execution_count":null},{"id":"7b44587afd6046f68edcfc9957272a81","cell_type":"code","source":"# ============================================================\n# SECTION 7.4 — CONFUSION MATRIX & CLINICAL ERROR ANALYSIS\n#\n# Generates normalized and raw confusion matrices as a heatmap\n# and produces an automated clinical interpretation of misclassifications:\n# - Identifies adjacent stage confusions (e.g. Mild vs Moderate)\n# - Flags high-risk distant confusions (e.g. No DR vs Severe/PDR)\n# - Provides clinical hypotheses for boundary classification ambiguities.\n# ============================================================\n\ndef plot_confusion_matrix_heatmap(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    class_names: List[str] = Config.CLASS_NAMES,\n    save_path: Optional[str] = None,\n) -> np.ndarray:\n    \"\"\"\n    Plot and save dual confusion matrices (raw counts and normalized percentages).\n\n    Args:\n        y_true: Ground truth integer class labels.\n        y_pred: Predicted integer class labels.\n        class_names: Human-readable class names.\n        save_path: Optional path to save PNG.\n\n    Returns:\n        Raw confusion matrix 2D NumPy array.\n    \"\"\"\n    cm = confusion_matrix(y_true, y_pred, labels=range(len(class_names)))\n    # Normalized by true row support (recall per class)\n    cm_norm = cm.astype(\"float\") / (cm.sum(axis=1)[:, np.newaxis] + 1e-10)\n\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\n    # Subplot 1: Absolute Counts\n    sns.heatmap(\n        cm,\n        annot=True,\n        fmt=\"d\",\n        cmap=\"Blues\",\n        xticklabels=class_names,\n        yticklabels=class_names,\n        ax=axes[0],\n        cbar=True,\n    )\n    axes[0].set_title(\"Confusion Matrix (Counts)\", fontsize=13, fontweight=\"bold\")\n    axes[0].set_xlabel(\"Predicted Stage\", fontsize=11)\n    axes[0].set_ylabel(\"True Stage\", fontsize=11)\n    axes[0].tick_params(axis=\"x\", rotation=30)\n\n    # Subplot 2: Normalized (Recall) Proportions\n    sns.heatmap(\n        cm_norm,\n        annot=True,\n        fmt=\".2%\",\n        cmap=\"Blues\",\n        xticklabels=class_names,\n        yticklabels=class_names,\n        ax=axes[1],\n        cbar=True,\n    )\n    axes[1].set_title(\"Confusion Matrix (Normalized by True Class)\", fontsize=13, fontweight=\"bold\")\n    axes[1].set_xlabel(\"Predicted Stage\", fontsize=11)\n    axes[1].set_ylabel(\"True Stage\", fontsize=11)\n    axes[1].tick_params(axis=\"x\", rotation=30)\n\n    plt.tight_layout()\n\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n        print(f\"[Confusion Matrix] Saved figure -> {save_path}\")\n\n    plt.show()\n    return cm\n\n\ndef interpret_confusion_matrix(\n    cm: np.ndarray,\n    class_names: List[str] = Config.CLASS_NAMES,\n    save_path: Optional[str] = None,\n) -> Optional[pd.DataFrame]:\n    \"\"\"\n    Analyze the confusion matrix to identify primary error patterns and print clinical hypotheses.\n\n    Args:\n        cm: 2D integer confusion matrix array (rows=true, cols=pred).\n        class_names: Class labels.\n        save_path: Optional CSV path for every misclassification pair, ranked by count.\n    \"\"\"\n    print(\"\\n\" + \"=\" * 70)\n    print(\"  CLINICAL CONFUSION INTERPRETATION & HYPOTHESIS ANALYSIS\")\n    print(\"=\" * 70)\n\n    # Collect all off-diagonal confusion counts\n    off_diagonals = []\n    n_classes = len(class_names)\n    for i in range(n_classes):\n        for j in range(n_classes):\n            if i != j and cm[i, j] > 0:\n                off_diagonals.append((cm[i, j], i, j, abs(i - j)))\n\n    # Sort by error count descending\n    off_diagonals.sort(key=lambda x: x[0], reverse=True)\n\n    total_errors = sum(count for count, *_ in off_diagonals)\n    pairs_df = pd.DataFrame([\n        {\n            \"rank\": rank,\n            \"true_label\": true_cls,\n            \"true_class\": class_names[true_cls],\n            \"predicted_label\": pred_cls,\n            \"predicted_class\": class_names[pred_cls],\n            \"count\": int(count),\n            \"percent_of_true_class\": float(count / cm[true_cls].sum() * 100),\n            \"percent_of_all_errors\": float(count / total_errors * 100) if total_errors else 0.0,\n            \"stage_distance\": distance,\n            \"error_type\": \"adjacent\" if distance == 1 else \"distant (high clinical risk)\",\n        }\n        for rank, (count, true_cls, pred_cls, distance) in enumerate(off_diagonals, start=1)\n    ])\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        pairs_df.to_csv(save_path, index=False)\n        print(f\"[Confusion Pairs] Saved all {len(pairs_df)} misclassification pairs -> {save_path}\")\n\n    print(\"Top Misclassification Pairs (Observed in Test Data):\")\n    for count, true_cls, pred_cls, distance in off_diagonals[:5]:\n        dist_type = \"Adjacent (Distance 1)\" if distance == 1 else f\"Distant (Distance {distance}) - High Clinical Risk!\"\n        print(\n            f\"  - True [{class_names[true_cls]}] misclassified as [{class_names[pred_cls]}]: \"\n            f\"{count:,} cases | {dist_type}\"\n        )\n\n    print(\"\\nGeneral clinical background (not derived from this run's results):\")\n    print(\"Clinical Diagnosis & Boundary Ambiguity Hypotheses:\")\n    print(\"1. Mild NPDR (Stage 1) vs. Moderate NPDR (Stage 2):\")\n    print(\"   - Clinical criterion for Stage 1 is microaneurysms only. Stage 2 involves early dot\")\n    print(\"     blot hemorrhages, venous beading, or hard exudates.\")\n    print(f\"   - In downsampled {Config.IMG_SIZE}x{Config.IMG_SIZE} images, faint microaneurysms and small dot hemorrhages\")\n    print(\"     can look identical at single-pixel scales, causing the model to blur this boundary.\")\n\n    print(\"2. Moderate NPDR (Stage 2) vs. Severe NPDR (Stage 3):\")\n    print(\"   - Severe NPDR requires fulfilling the '4-2-1 rule' (>20 intraretinal hemorrhages in\")\n    print(\"     all 4 quadrants, venous beading in 2+ quadrants, or IRMA in 1+ quadrant).\")\n    print(\"   - Without explicit lesion segmentation or full-retina multi-quadrant counting,\")\n    print(\"     global average pooled features capture lesion density but cannot formally count\")\n    print(\"     quadrant coverage, resulting in adjacent class confusion.\")\n\n    print(\"3. Distant Misclassification Safeguards:\")\n    print(\"   - Severe/Proliferative cases misdiagnosed as No DR represent fatal false negatives.\")\n    print(\"   - The GovernanceAgent (Section 10) explicitly gates cases where model confidence is\")\n    print(\"     uncertain (< 70%), diverting them to human ophthalmologist triage.\")\n\n    # Computed from this run's confusion matrix (rows = true stage, columns = predicted).\n    def _rate(count: int, total: int) -> str:\n        return f\"{count / total * 100:.1f}%\" if total else \"n/a\"\n\n    row_totals = cm.sum(axis=1)\n    print(\"\\nWhat this run's confusion matrix shows for each hypothesis:\")\n    if n_classes >= 5:\n        print(\n            f\"1. Mild -> Moderate: {cm[1, 2]:,} of {row_totals[1]:,} Mild ({_rate(cm[1, 2], row_totals[1])}); \"\n            f\"Moderate -> Mild: {cm[2, 1]:,} of {row_totals[2]:,} Moderate ({_rate(cm[2, 1], row_totals[2])}).\"\n        )\n        print(\n            f\"2. Moderate -> Severe: {cm[2, 3]:,} of {row_totals[2]:,} Moderate ({_rate(cm[2, 3], row_totals[2])}); \"\n            f\"Severe -> Moderate: {cm[3, 2]:,} of {row_totals[3]:,} Severe ({_rate(cm[3, 2], row_totals[3])}).\"\n        )\n        severe_as_no_dr = int(cm[3, 0] + cm[4, 0])\n        severe_total = int(row_totals[3] + row_totals[4])\n        print(\n            f\"3. Severe/Proliferative predicted as No DR: {severe_as_no_dr:,} of {severe_total:,} \"\n            f\"({_rate(severe_as_no_dr, severe_total)}) — Severe {cm[3, 0]:,}, Proliferative {cm[4, 0]:,}.\"\n        )\n    print(\"=\" * 70 + \"\\n\")\n    return pairs_df\n\n\ntest_cm = plot_confusion_matrix_heatmap(\n    y_test_true,\n    y_test_pred,\n    save_path=os.path.join(Config.REPORTS_DIR, \"confusion_matrix.png\"),\n)\ntop_confusion_pairs = interpret_confusion_matrix(\n    test_cm,\n    save_path=os.path.join(Config.REPORTS_DIR, \"top_confusion_pairs.csv\"),\n)\ndisplay(top_confusion_pairs.head(10))","metadata":{"id":"7b44587afd6046f68edcfc9957272a81","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:35:05.315458Z","iopub.execute_input":"2026-09-30T20:35:05.315908Z","iopub.status.idle":"2026-09-30T20:35:06.259437Z","shell.execute_reply.started":"2026-09-30T20:35:05.315879Z","shell.execute_reply":"2026-09-30T20:35:06.258799Z"}},"outputs":[],"execution_count":null},{"id":"24155d6d","cell_type":"code","source":"# ============================================================\n# SECTION 7.4b — ARGMAX REPORTS AND ROC FROM SAVED PREDICTIONS\n# No model inference is performed here; all results come from the NPZ.\n# ============================================================\n\nfrom sklearn.metrics import auc, roc_curve  # CHANGED: compute ROC from persisted softmax scores\n\nwith np.load(test_predictions_path) as saved_test_predictions:  # CHANGED: use the one-pass test artifact\n    artifact_y_true = saved_test_predictions[\"y_true\"]  # CHANGED: load labels from artifact\n    artifact_y_pred = saved_test_predictions[\"y_pred\"]  # CHANGED: load thresholded predictions\n    artifact_y_pred_argmax = saved_test_predictions[\"y_pred_argmax\"]  # CHANGED: load plain predictions\n    artifact_probabilities = saved_test_predictions[\"softmax_probabilities\"]  # CHANGED: load scores for ROC\n\nif not np.array_equal(artifact_y_pred_argmax, np.argmax(artifact_probabilities, axis=1)):  # CHANGED: verify stored argmax consistency\n    raise AssertionError(\"Saved argmax predictions do not match the saved softmax probabilities.\")  # CHANGED: stop on inconsistent evidence\n\nargmax_report_df = generate_classification_summary(  # CHANGED: report raw argmax metrics from saved arrays\n    artifact_y_true,\n    artifact_y_pred_argmax,\n    save_path=os.path.join(Config.REPORTS_DIR, \"classification_report_argmax.csv\"),\n)\nargmax_cm = plot_confusion_matrix_heatmap(  # CHANGED: create the complementary raw argmax confusion matrix\n    artifact_y_true,\n    artifact_y_pred_argmax,\n    save_path=os.path.join(Config.REPORTS_DIR, \"confusion_matrix_argmax.png\"),\n)\ninterpret_confusion_matrix(argmax_cm)  # CHANGED: interpret only predictions stored in the artifact\n\nroc_records = []  # CHANGED: collect per-class ROC summary from saved probabilities\nfig, ax = plt.subplots(figsize=(8, 7))  # CHANGED: create one-vs-rest ROC figure\nfor class_id, class_name in enumerate(Config.CLASS_NAMES):  # CHANGED: calculate one ROC curve per ICDR stage\n    binary_truth = (artifact_y_true == class_id).astype(int)  # CHANGED: form one-vs-rest labels\n    if np.unique(binary_truth).size != 2:  # CHANGED: require positives and negatives for AUC\n        raise ValueError(f\"Test artifact is missing one-vs-rest examples for class {class_id}.\")  # CHANGED: fail instead of reporting invalid AUC\n    false_positive_rate, true_positive_rate, _ = roc_curve(binary_truth, artifact_probabilities[:, class_id])  # CHANGED: use saved softmax column\n    class_auc = float(auc(false_positive_rate, true_positive_rate))  # CHANGED: integrate ROC curve\n    roc_records.append({\"class_id\": class_id, \"class_name\": class_name, \"roc_auc\": class_auc})  # CHANGED: persist class AUC\n    ax.plot(false_positive_rate, true_positive_rate, label=f\"{class_name} (AUC={class_auc:.3f})\")  # CHANGED: plot per-stage ROC\n\nmacro_auc = float(np.mean([record[\"roc_auc\"] for record in roc_records]))  # CHANGED: macro-average one-vs-rest AUC\nroc_records.append({\"class_id\": \"macro\", \"class_name\": \"Macro average\", \"roc_auc\": macro_auc})  # CHANGED: include aggregate AUC in CSV\nax.plot([0, 1], [0, 1], linestyle=\"--\", color=\"gray\", label=\"Chance\")  # CHANGED: show chance baseline\nax.set(title=f\"One-vs-Rest ROC Curves (Macro AUC={macro_auc:.3f})\", xlabel=\"False Positive Rate\", ylabel=\"True Positive Rate\")  # CHANGED: label plot\nax.legend(loc=\"lower right\", fontsize=8)  # CHANGED: identify each class curve\nax.grid(alpha=0.25)  # CHANGED: improve plot readability\nfig.tight_layout()  # CHANGED: fit labels and legend\nroc_plot_path = os.path.join(Config.REPORTS_DIR, \"roc_auc_curves.png\")  # CHANGED: save report figure in run directory\nfig.savefig(roc_plot_path, dpi=150, bbox_inches=\"tight\")  # CHANGED: persist ROC figure\nplt.show()  # CHANGED: display ROC figure\npd.DataFrame(roc_records).to_csv(os.path.join(Config.REPORTS_DIR, \"roc_auc_summary.csv\"), index=False)  # CHANGED: persist per-class and macro AUC values\nprint(f\"[Evaluation] Macro one-vs-rest ROC-AUC: {macro_auc:.4f}\")  # CHANGED: report artifact-derived AUC\nprint(f\"[Evaluation] ROC figure saved -> {roc_plot_path}\")  # CHANGED: disclose ROC artifact path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:35:47.110787Z","iopub.execute_input":"2026-09-30T20:35:47.11108Z","iopub.status.idle":"2026-09-30T20:35:48.448031Z","shell.execute_reply.started":"2026-09-30T20:35:47.111056Z","shell.execute_reply":"2026-09-30T20:35:48.447192Z"}},"outputs":[],"execution_count":null},{"id":"db006431","cell_type":"code","source":"# ============================================================\n# SECTION 7.4c — PER-STAGE RECALL COMPARISON\n# Reuses the saved one-pass test predictions; performs no new inference.\n# ============================================================\n\nwith np.load(test_predictions_path) as saved_predictions:\n    recall_true = saved_predictions[\"y_true\"].astype(int)\n    recall_prediction_sets = {\n        \"Majority baseline\": saved_predictions[\"y_pred_majority\"].astype(int),\n        \"EfficientNet argmax\": saved_predictions[\"y_pred_argmax\"].astype(int),\n        \"Validation-calibrated\": saved_predictions[\"y_pred\"].astype(int),\n    }\n\nrecall_records = []\nfor method_name, method_predictions in recall_prediction_sets.items():\n    method_report = classification_report(\n        recall_true,\n        method_predictions,\n        labels=list(range(Config.NUM_CLASSES)),\n        target_names=Config.CLASS_NAMES,\n        output_dict=True,\n        zero_division=0,\n    )\n    for class_id, class_name in enumerate(Config.CLASS_NAMES):\n        recall_records.append({\n            \"class_id\": class_id,\n            \"class_name\": class_name,\n            \"method\": method_name,\n            \"recall\": float(method_report[class_name][\"recall\"]),\n            \"support\": int(method_report[class_name][\"support\"]),\n        })\n\nper_class_recall_comparison = pd.DataFrame(recall_records)\nrecall_table = per_class_recall_comparison.pivot(\n    index=[\"class_id\", \"class_name\", \"support\"],\n    columns=\"method\",\n    values=\"recall\",\n).reset_index()\nrecall_csv_path = Path(Config.REPORTS_DIR) / \"per_class_recall_comparison.csv\"\nper_class_recall_comparison.to_csv(recall_csv_path, index=False)\ndisplay(recall_table.round(4))\n\nrecall_plot_data = per_class_recall_comparison.pivot(\n    index=\"class_name\", columns=\"method\", values=\"recall\"\n).reindex(Config.CLASS_NAMES)\nfig, ax = plt.subplots(figsize=(12, 6))\nrecall_plot_data.plot(\n    kind=\"bar\",\n    ax=ax,\n    color=[\"#7A8F55\", \"#367C8A\", \"#D8784B\"],\n    width=0.78,\n)\nax.set_title(\"Per-Stage Recall from the Saved Held-Out Test Predictions\")\nax.set_xlabel(\"Dataset grade\")\nax.set_ylabel(\"Recall\")\nax.set_ylim(0, 1)\nax.legend(title=\"Decision rule\")\nax.grid(axis=\"y\", alpha=0.25)\nax.tick_params(axis=\"x\", rotation=18)\nfig.tight_layout()\nrecall_plot_path = Path(Config.REPORTS_DIR) / \"per_class_recall_comparison.png\"\nfig.savefig(recall_plot_path, dpi=160, bbox_inches=\"tight\")\nplt.show()\nprint(f\"[Recall Comparison] Reused saved test predictions; no model inference was run.\")\nprint(f\"[Recall Comparison] Saved table -> {recall_csv_path}\")\nprint(f\"[Recall Comparison] Saved plot -> {recall_plot_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:04.955217Z","iopub.execute_input":"2026-09-30T20:37:04.955518Z","iopub.status.idle":"2026-09-30T20:37:05.416463Z","shell.execute_reply.started":"2026-09-30T20:37:04.955494Z","shell.execute_reply":"2026-09-30T20:37:05.415631Z"}},"outputs":[],"execution_count":null},{"id":"1a4f57c6","cell_type":"markdown","source":"### 7.4d DR detection and referable-DR screening\n\nThe brief asks the model to classify diabetic retinopathy **as well as** its stage. The next cell collapses the five-stage predictions into *any DR* (stage ≥ 1) and *referable DR* (stage ≥ 2), and reports sensitivity, specificity, PPV, NPV and AUC for each.","metadata":{}},{"id":"3a74c2a0","cell_type":"code","source":"# ============================================================\n# SECTION 7.4d — DR DETECTION (BINARY) AND REFERABLE-DR SCREENING\n#\n# The brief asks the model to detect diabetic retinopathy AS WELL AS\n# its stage. The 5-stage predictions are collapsed into two screening\n# questions, using the saved one-pass test predictions (no new inference):\n#   - Any DR       : stage >= 1  (DR present vs No DR)\n#   - Referable DR : stage >= 2  (the usual screening-programme cut-off\n#                    for referral to an ophthalmologist)\n# Sensitivity = share of diseased eyes that are flagged (missed cases hurt);\n# specificity = share of healthy eyes correctly cleared (false alarms cost).\n# AUC uses the summed softmax probability of the positive stages.\n# ============================================================\n\nfrom sklearn.metrics import roc_auc_score\n\nSCREENING_CUTOFFS = {\"Any DR (stage >= 1)\": 1, \"Referable DR (stage >= 2)\": 2}\n\n\ndef screening_metrics(y_true: np.ndarray, y_pred: np.ndarray, cutoff: int) -> Dict[str, float]:\n    \"\"\"Binary confusion counts and rates after collapsing stages at `cutoff`.\"\"\"\n    true_positive_mask = np.asarray(y_true) >= cutoff\n    pred_positive_mask = np.asarray(y_pred) >= cutoff\n    tp = int(np.sum(true_positive_mask & pred_positive_mask))\n    tn = int(np.sum(~true_positive_mask & ~pred_positive_mask))\n    fp = int(np.sum(~true_positive_mask & pred_positive_mask))\n    fn = int(np.sum(true_positive_mask & ~pred_positive_mask))\n\n    def _ratio(numerator: int, denominator: int) -> float:\n        return float(numerator / denominator) if denominator else float(\"nan\")\n\n    return {\n        \"tp\": tp, \"fn\": fn, \"fp\": fp, \"tn\": tn,\n        \"sensitivity\": _ratio(tp, tp + fn),\n        \"specificity\": _ratio(tn, tn + fp),\n        \"ppv\": _ratio(tp, tp + fp),\n        \"npv\": _ratio(tn, tn + fn),\n        \"accuracy\": _ratio(tp + tn, tp + tn + fp + fn),\n    }\n\n\nwith np.load(test_predictions_path) as saved_predictions:\n    screening_true = saved_predictions[\"y_true\"].astype(int)\n    screening_probabilities = saved_predictions[\"softmax_probabilities\"]\n    screening_rules = {\n        \"EfficientNet argmax\": saved_predictions[\"y_pred_argmax\"].astype(int),\n        \"Validation-calibrated\": saved_predictions[\"y_pred\"].astype(int),\n    }\n\nscreening_records = []\nfor task_name, cutoff in SCREENING_CUTOFFS.items():\n    positive_probability = screening_probabilities[:, cutoff:].sum(axis=1)\n    task_auc = float(roc_auc_score(screening_true >= cutoff, positive_probability))\n    for rule_name, rule_predictions in screening_rules.items():\n        screening_records.append({\n            \"task\": task_name,\n            \"decision_rule\": rule_name,\n            \"positives_in_test\": int(np.sum(screening_true >= cutoff)),\n            \"negatives_in_test\": int(np.sum(screening_true < cutoff)),\n            **screening_metrics(screening_true, rule_predictions, cutoff),\n            \"roc_auc\": task_auc,\n        })\n\nscreening_df = pd.DataFrame(screening_records)\nscreening_csv_path = Path(Config.REPORTS_DIR) / \"screening_metrics.csv\"\nscreening_df.to_csv(screening_csv_path, index=False)\ndisplay(screening_df.round(4))\n\n# Clinically critical errors under the calibrated rule used for the headline results.\ncalibrated_predictions = screening_rules[\"Validation-calibrated\"]\nmissed_referable = int(np.sum((screening_true >= 2) & (calibrated_predictions < 2)))\nsight_threatening_as_no_dr = int(np.sum((screening_true >= 3) & (calibrated_predictions == 0)))\nprint(f\"Referable-DR eyes missed (true >= 2, predicted < 2): {missed_referable:,} \"\n      f\"of {int(np.sum(screening_true >= 2)):,}\")\nprint(f\"Severe/Proliferative eyes predicted as No DR       : {sight_threatening_as_no_dr:,}\")\n\nfig, axes = plt.subplots(1, len(SCREENING_CUTOFFS), figsize=(12, 5))\nfor axis, (task_name, cutoff) in zip(axes, SCREENING_CUTOFFS.items()):\n    row = screening_df[(screening_df[\"task\"] == task_name)\n                       & (screening_df[\"decision_rule\"] == \"Validation-calibrated\")].iloc[0]\n    matrix = np.array([[row[\"tn\"], row[\"fp\"]], [row[\"fn\"], row[\"tp\"]]])\n    negative_label = \"No DR\" if cutoff == 1 else \"Stage 0-1\"\n    positive_label = \"DR\" if cutoff == 1 else \"Stage 2-4\"\n    sns.heatmap(matrix, annot=True, fmt=\"d\", cmap=\"Blues\", cbar=False, ax=axis,\n                xticklabels=[negative_label, positive_label],\n                yticklabels=[negative_label, positive_label])\n    axis.set_title(f\"{task_name}\\nSens {row['sensitivity']:.3f} | Spec {row['specificity']:.3f} \"\n                   f\"| AUC {row['roc_auc']:.3f}\", fontsize=11)\n    axis.set_xlabel(\"Predicted\")\n    axis.set_ylabel(\"True\")\nfig.suptitle(\"Binary screening performance on the held-out test set (validation-calibrated rule)\",\n             fontsize=12, fontweight=\"bold\")\nfig.tight_layout()\nscreening_plot_path = Path(Config.REPORTS_DIR) / \"screening_confusion_matrices.png\"\nfig.savefig(screening_plot_path, dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(f\"[Screening] Saved table -> {screening_csv_path}\")\nprint(f\"[Screening] Saved figure -> {screening_plot_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:12.569613Z","iopub.execute_input":"2026-09-30T20:37:12.570333Z","iopub.status.idle":"2026-09-30T20:37:12.96144Z","shell.execute_reply.started":"2026-09-30T20:37:12.570306Z","shell.execute_reply":"2026-09-30T20:37:12.960557Z"}},"outputs":[],"execution_count":null},{"id":"9e9e05e40e7c406eb28d4589e597fe7d","cell_type":"code","source":"# ============================================================\n# SECTION 7.5 — SECTION SUMMARY\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  SECTION 7 COMPLETE — Evaluation\")\nprint(\"=\" * 60)\nprint(f\"  Accuracy                : {eval_metrics['accuracy']:.4f}\")\nprint(f\"  Quadratic Weighted Kappa: {eval_metrics['quadratic_weighted_kappa']:.4f}\")\nprint(f\"  Artifacts exported to {Config.REPORTS_DIR}:\")\nprint(\"    - classification_report.csv\")\nprint(\"    - confusion_matrix.png, confusion_matrix_argmax.png, roc_auc_curves.png\")\nprint(\"    - per_class_recall_comparison.png\")\nprint(\"    - screening_metrics.csv, screening_confusion_matrices.png\")\nprint(\"  Next -> Section 8: Explainability — Grad-CAM\")\nprint(\"=\" * 60)","metadata":{"id":"9e9e05e40e7c406eb28d4589e597fe7d","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:20.572782Z","iopub.execute_input":"2026-09-30T20:37:20.573391Z","iopub.status.idle":"2026-09-30T20:37:20.579144Z","shell.execute_reply.started":"2026-09-30T20:37:20.573348Z","shell.execute_reply":"2026-09-30T20:37:20.578497Z"}},"outputs":[],"execution_count":null},{"id":"3ae6942b7f254a7b87f3ff9f443ef2be","cell_type":"markdown","source":"## Section 8: Explainability — Grad-CAM\n\n**Goal:** Provide visual and linguistic interpretability for the model's predictions using\n**Gradient-weighted Class Activation Mapping (Grad-CAM)** and an automated retinal quadrant describer.\n\n### Clinical & Regulatory Rationale for Explainability in Medical AI\n\nDeep neural networks are frequently criticized as \"black boxes\" in healthcare. The European Union\nAI Act and FDA SaMD (Software as a Medical Device) guidelines require high-risk clinical AI\nsystems to provide explainability to clinical end-users.\n\n**Grad-CAM** works by:\n1. Identifying the last convolutional layer (`top_activation` in EfficientNetB3), which preserves\n   both high-level semantic features (lesions, exudates, neovascular tufts) and spatial coordinates.\n2. Computing the gradient of the winning class score $y^c$ with respect to feature activation maps $A^k$:\n   $$\\alpha_k^c = \\frac{1}{Z} \\sum_i \\sum_j \\frac{\\partial y^c}{\\partial A_{i,j}^k}$$\n3. Computing a weighted combination followed by a rectified linear activation (ReLU):\n   $$L_{\\text{Grad-CAM}}^c = \\text{ReLU}\\left(\\sum_k \\alpha_k^c A^k\\right)$$\n   The $\\text{ReLU}$ ensures the heatmap displays features that **positively contribute** to the\n   predicted disease stage, rather than features that suppress it.\n\n### Automated Retinal Quadrant Description\nTo make explainability immediately actionable for busy clinicians who do not want to inspect raw\nheatmaps, we segment the retinal field into anatomical quadrants (Superior-Temporal, Superior-Nasal,\nInferior-Temporal, Inferior-Nasal, and Central Macular) and generate a natural language explanation.","metadata":{"id":"3ae6942b7f254a7b87f3ff9f443ef2be","language":"markdown"}},{"id":"a1a01e4eb3154aceaa0206d662071be1","cell_type":"code","source":"# ============================================================\n# SECTION 8.1 — GRAD-CAM MODEL CONSTRUCTION\n#\n# Builds a gradient-tappable sub-model using the Keras functional API.\n# It returns both:\n#   1. The feature activation map from the final conv layer ('top_activation')\n#   2. The final softmax prediction probabilities\n# ============================================================\n\ndef build_gradcam_model(\n    full_model: keras.Model,\n    base: keras.Model,\n    last_conv_layer_name: str = \"top_activation\",\n) -> keras.Model:\n    \"\"\"\n    Construct a dual-output Grad-CAM inspection model.\n\n    Args:\n        full_model: The complete trained DR_EfficientNetB3 model.\n        base: The EfficientNetB3 base model sub-layer.\n        last_conv_layer_name: Name of target convolutional layer in base model.\n                              Defaults to 'top_activation' (last activation of EfficientNetB3).\n\n    Returns:\n        Keras Model with inputs=raw_image, outputs=[conv_features, class_predictions].\n    \"\"\"\n    # 1. Access the target conv layer from the base feature extractor\n    target_conv_layer = base.get_layer(last_conv_layer_name)\n\n    # 2. Build intermediate sub-model for the base extractor\n    base_submodel = keras.Model(\n        inputs=base.inputs,\n        outputs=[target_conv_layer.output, base.output],\n        name=\"base_submodel\",\n    )\n\n    # 3. Wire the full computational graph from raw inputs through the classification head\n    inputs = keras.Input(shape=(Config.IMG_SIZE, Config.IMG_SIZE, 3), name=\"gradcam_input\")\n    scaled_inputs = layers.Rescaling(255.0, name=\"gradcam_restore_input_range\")(inputs)\n    conv_features, base_features = base_submodel(scaled_inputs)\n\n    # Pass through head layers (reusing trained weights from full_model)\n    x = full_model.get_layer(\"gap\")(base_features)\n    x = full_model.get_layer(\"head_bn\")(x)\n    x = full_model.get_layer(\"head_dense\")(x)\n    x = full_model.get_layer(\"head_dropout\")(x)\n    preds = full_model.get_layer(\"predictions\")(x)\n\n    gradcam_model = keras.Model(\n        inputs=inputs,\n        outputs=[conv_features, preds],\n        name=\"gradcam_extractor\",\n    )\n    return gradcam_model\n\n\n# Instantiate Grad-CAM inspection model\ngradcam_model = build_gradcam_model(model, base_model, last_conv_layer_name=\"top_activation\")\nprint(f\"[Grad-CAM] Model built successfully with target layer: 'top_activation'\")","metadata":{"id":"a1a01e4eb3154aceaa0206d662071be1","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:29.059318Z","iopub.execute_input":"2026-09-30T20:37:29.059705Z","iopub.status.idle":"2026-09-30T20:37:29.092803Z","shell.execute_reply.started":"2026-09-30T20:37:29.059679Z","shell.execute_reply":"2026-09-30T20:37:29.092189Z"}},"outputs":[],"execution_count":null},{"id":"02b27c5560ce486fa4556ab239b474a9","cell_type":"code","source":"# ============================================================\n# SECTION 8.2 — GRAD-CAM COMPUTATION VIA TF.GRADIENTTAPE\n#\n# Computes Grad-CAM heatmaps for a target class.\n# ============================================================\n\ndef compute_gradcam_heatmap(\n    img_array: np.ndarray,\n    cam_model: keras.Model = gradcam_model,\n    pred_index: Optional[int] = None,\n) -> Tuple[np.ndarray, int, float]:\n    \"\"\"Return a normalized float32 Grad-CAM map, class index, and confidence.\"\"\"\n    if img_array.ndim == 3:\n        img_tensor = tf.expand_dims(img_array, axis=0)\n    else:\n        img_tensor = tf.convert_to_tensor(img_array)\n\n    with tf.GradientTape() as tape:\n        tape.watch(img_tensor)\n        conv_outputs, predictions = cam_model(img_tensor, training=False)\n        if pred_index is None:\n            pred_index = tf.argmax(predictions[0]).numpy()\n        class_channel = predictions[:, pred_index]\n\n    grads = tape.gradient(class_channel, conv_outputs)\n    pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2))\n    conv_outputs_val = conv_outputs[0]\n    heatmap = conv_outputs_val @ pooled_grads[..., tf.newaxis]\n    heatmap = tf.squeeze(heatmap)\n    heatmap = tf.maximum(heatmap, 0)\n\n    max_val = tf.math.reduce_max(heatmap)\n    if max_val > 0:\n        heatmap = heatmap / max_val\n\n    confidence = float(predictions[0, pred_index].numpy())\n    heatmap = np.asarray(heatmap.numpy(), dtype=np.float32)\n    return heatmap, int(pred_index), confidence","metadata":{"id":"02b27c5560ce486fa4556ab239b474a9","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:34.346685Z","iopub.execute_input":"2026-09-30T20:37:34.34746Z","iopub.status.idle":"2026-09-30T20:37:34.354249Z","shell.execute_reply.started":"2026-09-30T20:37:34.347433Z","shell.execute_reply":"2026-09-30T20:37:34.353575Z"}},"outputs":[],"execution_count":null},{"id":"3d745079e13f4b168fd5adaefceba996","cell_type":"code","source":"# ============================================================\n# SECTION 8.3 — HEATMAP SUPERIMPOSITION & QUADRANT GENERATOR\n#\n# 1. superimpose_gradcam(): Blends Jet colormap with RGB fundus image.\n# 2. generate_quadrant_explanation(): Divides heatmap into retinal\n#    anatomical regions and generates a natural-language description.\n# ============================================================\n\ndef superimpose_gradcam(\n    original_rgb: np.ndarray,\n    heatmap: np.ndarray,\n    alpha: float = 0.4,\n    colormap: int = cv2.COLORMAP_JET,\n) -> np.ndarray:\n    \"\"\"\n    Overlay Grad-CAM heatmap onto the RGB fundus image.\n\n    Args:\n        original_rgb: RGB image (H, W, 3), float [0,1] or uint8 [0,255].\n        heatmap: 2D float array (h, w) with values in [0, 1].\n        alpha: Blending weight for heatmap (1-alpha weight for original image).\n        colormap: OpenCV colormap constant.\n\n    Returns:\n        Superimposed RGB image as uint8 NumPy array of shape (H, W, 3).\n    \"\"\"\n    # Standardize image to uint8 [0, 255]\n    if original_rgb.dtype != np.uint8:\n        base_img = (np.clip(original_rgb, 0.0, 1.0) * 255).astype(np.uint8)\n    else:\n        base_img = original_rgb.copy()\n\n    h, w = base_img.shape[:2]\n\n    # Resize heatmap to match original image dimensions\n    heatmap_resized = cv2.resize(heatmap, (w, h), interpolation=cv2.INTER_LINEAR)\n    heatmap_uint8 = (heatmap_resized * 255).astype(np.uint8)\n\n    # Apply color palette (Jet: blue=low, red=peak activation)\n    heatmap_color = cv2.applyColorMap(heatmap_uint8, colormap)\n    heatmap_color = cv2.cvtColor(heatmap_color, cv2.COLOR_BGR2RGB)\n\n    # Alpha blending: overlay = alpha * heatmap + (1 - alpha) * original\n    overlay = cv2.addWeighted(heatmap_color, alpha, base_img, 1.0 - alpha, 0)\n    return overlay\n\n\ndef generate_quadrant_explanation(\n    heatmap: np.ndarray,\n    predicted_stage: int,\n    confidence: float,\n) -> str:\n    \"\"\"\n    Analyze the spatial distribution of Grad-CAM activations and generate a clinical description.\n\n    Divides the retinal field into 5 key regions:\n      - Superior-Temporal (upper outer quadrant)\n      - Superior-Nasal (upper inner quadrant)\n      - Inferior-Temporal (lower outer quadrant)\n      - Inferior-Nasal (lower inner quadrant)\n      - Central Macular (central 50% zone)\n\n    Args:\n        heatmap: 2D float array normalized to [0, 1].\n        predicted_stage: Predicted DR stage (0 to 4).\n        confidence: Prediction confidence score.\n\n    Returns:\n        Structured natural-language clinical explanation sentence.\n    \"\"\"\n    h, w = heatmap.shape\n    mid_y, mid_x = h // 2, w // 2\n\n    # Define quadrant masks\n    quadrants = {\n        \"superior-temporal\": heatmap[:mid_y, :mid_x],\n        \"superior-nasal\":    heatmap[:mid_y, mid_x:],\n        \"inferior-temporal\": heatmap[mid_y:, :mid_x],\n        \"inferior-nasal\":    heatmap[mid_y:, mid_x:],\n    }\n\n    # Central macular region (central 50% of the image)\n    y1, y2 = int(0.25 * h), int(0.75 * h)\n    x1, x2 = int(0.25 * w), int(0.75 * w)\n    central_region = heatmap[y1:y2, x1:x2]\n\n    # Calculate mean activation density per region\n    scores = {name: float(np.mean(q)) for name, q in quadrants.items()}\n    scores[\"central-macular\"] = float(np.mean(central_region))\n\n    # Identify primary and secondary regions of interest\n    sorted_regions = sorted(scores.items(), key=lambda x: x[1], reverse=True)\n    peak_region, peak_score = sorted_regions[0]\n    sec_region, sec_score   = sorted_regions[1]\n\n    stage_name = Config.CLASS_NAMES[predicted_stage]\n\n    # Clinical pathology mapping associated with stage\n    pathology_notes = {\n        0: \"absence of diabetic microvascular abnormalities\",\n        1: \"isolated microaneurysms or focal vascular dilatation\",\n        2: \"microaneurysms, dot-and-blot hemorrhages, or hard exudates\",\n        3: \"widespread intraretinal microvascular abnormalities (IRMA) or multi-quadrant hemorrhages\",\n        4: \"neovascularization, vitreous preretinal hemorrhage, or fibrous proliferation\",\n    }\n\n    explanation = (\n        f\"The model predicted {stage_name} (Stage {predicted_stage}) with {confidence*100:.1f}% confidence. \"\n        f\"Grad-CAM salience indicates the model focused most strongly on the [{peak_region}] region \"\n        f\"(salience density: {peak_score:.2f}), with secondary attention on the [{sec_region}] region, \"\n        f\"consistent with {pathology_notes[predicted_stage]}.\"\n    )\n    return explanation","metadata":{"id":"3d745079e13f4b168fd5adaefceba996","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:39.111444Z","iopub.execute_input":"2026-09-30T20:37:39.112272Z","iopub.status.idle":"2026-09-30T20:37:39.123187Z","shell.execute_reply.started":"2026-09-30T20:37:39.112241Z","shell.execute_reply":"2026-09-30T20:37:39.122456Z"}},"outputs":[],"execution_count":null},{"id":"aaab0616b528444f84008365daf3f51f","cell_type":"code","source":"# ============================================================\n# SECTION 8.4 — GRAD-CAM EXAMPLES: CORRECT AND INCORRECT VALIDATION PREDICTIONS\n# Select examples from validation predictions; do not inspect the held-out test set.\n# The Grad-CAM target is the model's predicted class, not the ground-truth label.\n# ============================================================\n\ndef evaluate_and_visualize_gradcam(\n    validation_df: pd.DataFrame,\n    validation_true: np.ndarray,\n    validation_pred: np.ndarray,\n    save_dir: str = Config.REPORTS_DIR,\n) -> None:\n    \"\"\"Show one correct and one incorrect validation prediction per true stage when available.\"\"\"\n    validation_df = validation_df.reset_index(drop=True)\n    validation_true = np.asarray(validation_true, dtype=int)\n    validation_pred = np.asarray(validation_pred, dtype=int)\n    if len(validation_df) != len(validation_true) or len(validation_true) != len(validation_pred):\n        raise ValueError(\"Validation rows, labels, and predictions must have matching lengths.\")\n\n    selected_examples = []\n    for true_stage in range(Config.NUM_CLASSES):\n        for outcome, mask in (\n            (\"Correct\", validation_pred == validation_true),\n            (\"Incorrect\", validation_pred != validation_true),\n        ):\n            candidates = np.flatnonzero((validation_true == true_stage) & mask)\n            if candidates.size:\n                selected_examples.append((int(candidates[0]), outcome))\n\n    if not selected_examples:\n        raise ValueError(\"No validation examples were available for Grad-CAM visualization.\")\n\n    os.makedirs(save_dir, exist_ok=True)\n    fig, axes = plt.subplots(\n        len(selected_examples),\n        3,\n        figsize=(15, 4 * len(selected_examples)),\n        squeeze=False,\n    )\n    column_headers = [\"Validation fundus\", \"Grad-CAM for predicted class\", \"Predicted-class overlay\"]\n    for column, header in enumerate(column_headers):\n        axes[0, column].set_title(header, fontsize=12, fontweight=\"bold\")\n\n    for row_index, (sample_index, outcome) in enumerate(selected_examples):\n        sample_row = validation_df.iloc[sample_index]\n        true_stage = int(validation_true[sample_index])\n        predicted_stage = int(validation_pred[sample_index])\n        image = preprocess_image(sample_row[\"filepath\"])\n        heatmap, _, confidence = compute_gradcam_heatmap(\n            image,\n            cam_model=gradcam_model,\n            pred_index=predicted_stage,\n        )\n        overlay = superimpose_gradcam(image, heatmap, alpha=0.4)\n        axes[row_index, 0].imshow(image)\n        axes[row_index, 0].axis(\"off\")\n        axes[row_index, 0].set_ylabel(\n            f\"{outcome}\\nTrue: {Config.CLASS_NAMES[true_stage]}\\n\"\n            f\"Predicted: {Config.CLASS_NAMES[predicted_stage]} ({confidence:.1%})\",\n            fontsize=9,\n            rotation=0,\n            labelpad=65,\n            va=\"center\",\n        )\n        axes[row_index, 1].imshow(heatmap, cmap=\"jet\")\n        axes[row_index, 1].axis(\"off\")\n        axes[row_index, 2].imshow(overlay)\n        axes[row_index, 2].axis(\"off\")\n        print(\n            f\"[{outcome}] True={Config.CLASS_NAMES[true_stage]}, \"\n            f\"predicted={Config.CLASS_NAMES[predicted_stage]}, \"\n            f\"predicted-class confidence={confidence:.4f}. \"\n            \"Grad-CAM indicates influential regions, not verified lesions or causal evidence.\"\n        )\n\n    fig.suptitle(\n        \"Grad-CAM examples from validation predictions (correct and incorrect where available)\",\n        fontsize=13,\n        fontweight=\"bold\",\n    )\n    plt.tight_layout()\n    save_path = os.path.join(save_dir, \"gradcam_correct_and_incorrect_validation.png\")\n    plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n    print(f\"[Grad-CAM] Validation examples saved -> {save_path}\")\n    plt.show()\n\n\nevaluate_and_visualize_gradcam(\n    val_df,\n    val_true,\n    calibrated_val_predictions,\n)","metadata":{"id":"aaab0616b528444f84008365daf3f51f","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:37:45.701874Z","iopub.execute_input":"2026-09-30T20:37:45.702259Z","iopub.status.idle":"2026-09-30T20:38:42.051548Z","shell.execute_reply.started":"2026-09-30T20:37:45.702234Z","shell.execute_reply":"2026-09-30T20:38:42.050294Z"}},"outputs":[],"execution_count":null},{"id":"23a0eb724f884f72a7d13adc434be207","cell_type":"code","source":"# ============================================================\n# SECTION 8.5 — SECTION SUMMARY\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  SECTION 8 COMPLETE — Explainability (Grad-CAM)\")\nprint(\"=\" * 60)\nprint(\"  build_gradcam_model()          : dual-output model tapping 'top_activation'\")\nprint(\"  compute_gradcam_heatmap()      : tf.GradientTape computation + ReLU\")\nprint(\"  superimpose_gradcam()          : Jet colormap alpha blending\")\nprint(\"  generate_quadrant_explanation(): anatomical quadrant natural-language generator\")\nprint(f\"  Artifact saved                 : {Config.REPORTS_DIR}/gradcam_correct_and_incorrect_validation.png\")\nprint(\"  Next -> Innovation A: Embedding-Based Similar-Case Retrieval\")\nprint(\"=\" * 60)","metadata":{"id":"23a0eb724f884f72a7d13adc434be207","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:38:56.485598Z","iopub.execute_input":"2026-09-30T20:38:56.486112Z","iopub.status.idle":"2026-09-30T20:38:56.492297Z","shell.execute_reply.started":"2026-09-30T20:38:56.486082Z","shell.execute_reply":"2026-09-30T20:38:56.491359Z"}},"outputs":[],"execution_count":null},{"id":"d27969fafbc749f8a14b9e80f3d6d23b","cell_type":"markdown","source":"## Section 9: Innovation Feature A — Embedding-Based Similar-Case Retrieval\n\n**Purpose:** Show a query fundus image beside visually similar examples from the saved training split. The gallery uses real image files and labels from the supplied dataset; it does not generate synthetic photographs or independently verified clinical records.\n\n### Rationale and Mechanism\n\nThe classifier's 256-dimensional `head_dense` representation is extracted before the softmax layer, L2-normalized, and compared with cached training embeddings using cosine similarity. For normalized vectors, cosine similarity is their dot product. The notebook's demonstration selects a validation image as the query and retrieves from the training split, keeping the query out of the reference set.\n\nThis was added as a case-comparison aid alongside the five-stage prediction. It is **not** a Siamese or metric-learning model: the embedding was learned through the classification objective. Similarity does not establish pathological equivalence, and dataset labels are not independent clinical confirmation. No retrieval accuracy, clinician-utility study, or controlled evidence that retrieval improves grading has been produced. Treat the gallery as an exploratory interface, not a clinical explanation or diagnosis.","metadata":{"id":"d27969fafbc749f8a14b9e80f3d6d23b","language":"markdown"}},{"id":"412df5acc7d94f91b8bd75005b078e54","cell_type":"code","source":"# ============================================================\n# SECTION 9.1 — PENULTIMATE EMBEDDING EXTRACTOR MODEL\n#\n# Creates a dedicated feature extractor from the trained network.\n# It taps the 256-dimensional Dense layer ('head_dense') immediately\n# before the final 5-class softmax output layer.\n# ============================================================\n\ndef build_embedding_extractor(trained_model: keras.Model) -> keras.Model:\n    \"\"\"\n    Construct a feature extractor model outputting 256-D penultimate embeddings.\n\n    Args:\n        trained_model: Full trained DR_EfficientNetB3 model.\n\n    Returns:\n        Keras Model with inputs=raw_images, outputs=256-dimensional penultimate features.\n    \"\"\"\n    # Tap the intermediate dense layer before dropout and softmax\n    dense_output = trained_model.get_layer(\"head_dense\").output\n\n    embedding_extractor = keras.Model(\n        inputs=trained_model.input,\n        outputs=dense_output,\n        name=\"penultimate_embedding_extractor\",\n    )\n    return embedding_extractor\n\n\n# Build embedding extractor\nembedding_extractor = build_embedding_extractor(model)\nprint(f\"[Embedding] Extractor built: {embedding_extractor.name}\")\nprint(f\"  Input shape : {embedding_extractor.input_shape}\")\nprint(f\"  Output shape: {embedding_extractor.output_shape} (256-dimensional latent representation)\")","metadata":{"id":"412df5acc7d94f91b8bd75005b078e54","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:39:01.865354Z","iopub.execute_input":"2026-09-30T20:39:01.866038Z","iopub.status.idle":"2026-09-30T20:39:01.873998Z","shell.execute_reply.started":"2026-09-30T20:39:01.866008Z","shell.execute_reply":"2026-09-30T20:39:01.87302Z"}},"outputs":[],"execution_count":null},{"id":"9647e214","cell_type":"code","source":"# ============================================================\n# SECTION 9.2 — EXTRACT & CACHE TRAINING SET EMBEDDINGS\n#\n# Generates 256-D penultimate embeddings for every training image,\n# L2-normalises them (so cosine similarity becomes a dot product),\n# and saves them to Config.EMBEDDINGS_PATH for the retrieval demo and app.\n#\n# Images are read from the preprocessed PNG cache built for training,\n# so extraction does not repeat the OpenCV preprocessing for the training images.\n# ============================================================\n\ndef extract_and_cache_training_embeddings(\n    extractor: keras.Model,\n    train_dataframe: pd.DataFrame,\n    cache_path: str = Config.EMBEDDINGS_PATH,\n    force_recompute: bool = False,\n) -> Tuple[np.ndarray, np.ndarray, List[str]]:\n    \"\"\"\n    Extract L2-normalised embeddings for all training images and persist them as NPZ.\n\n    For unit vectors u and v, cos_sim(u, v) = u . v, and\n    ||u - v||^2 = 2 - 2 (u . v), so ranking by dot product is equivalent\n    to nearest-neighbour search in Euclidean distance.\n\n    Returns:\n        embeddings (N, 256) float32, labels (N,) int, filepaths list[str]\n    \"\"\"\n    if os.path.exists(cache_path) and not force_recompute:\n        print(f\"[Embedding] Loading cached embeddings from {cache_path}...\")\n        with np.load(cache_path, allow_pickle=True) as data:\n            return data[\"embeddings\"], data[\"labels\"].astype(int), list(data[\"filepaths\"])\n\n    filepaths = train_dataframe[\"filepath\"].tolist()\n    labels = train_dataframe[\"label\"].to_numpy(dtype=int)\n    print(f\"[Embedding] Computing embeddings for {len(filepaths):,} training images...\")\n\n    # shuffle=False keeps the dataset order identical to train_dataframe row order.\n    embedding_dataset = _build_tuning_dataset(\n        train_dataframe, Config.IMG_SIZE, shuffle=False, augment=False\n    )\n    batches = []\n    for images, _ in tqdm(embedding_dataset, desc=\"Extracting embeddings\"):\n        features = extractor(images, training=False).numpy().astype(np.float32)\n        features /= np.linalg.norm(features, axis=1, keepdims=True) + 1e-10\n        batches.append(features)\n    embeddings = np.concatenate(batches, axis=0)\n    if len(embeddings) != len(filepaths):\n        raise ValueError(\"Embedding count does not match the number of training images.\")\n\n    os.makedirs(os.path.dirname(cache_path), exist_ok=True)\n    np.savez_compressed(cache_path, embeddings=embeddings, labels=labels, filepaths=np.asarray(filepaths))\n    print(f\"[Embedding] Saved {len(embeddings):,} embeddings -> {cache_path}\")\n    return embeddings, labels, filepaths\n\n\ntrain_embeddings, train_labels_arr, train_fps = extract_and_cache_training_embeddings(\n    embedding_extractor,\n    train_df,\n    cache_path=Config.EMBEDDINGS_PATH,\n)\nprint(f\"[Embedding] Matrix shape: {train_embeddings.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:39:06.577953Z","iopub.execute_input":"2026-09-30T20:39:06.578618Z","iopub.status.idle":"2026-09-30T20:40:34.567312Z","shell.execute_reply.started":"2026-09-30T20:39:06.578588Z","shell.execute_reply":"2026-09-30T20:40:34.566559Z"}},"outputs":[],"execution_count":null},{"id":"b772b6c2c6fb4e03852881fc6d2c639f","cell_type":"code","source":"# ============================================================\n# SECTION 9.3 — REUSABLE SIMILAR-CASE RETRIEVAL FUNCTION\n#\n# find_similar_cases(query_image, k=3):\n#   - Accepts an image array or file path\n#   - Extracts and normalizes the 256-D query vector\n#   - Computes cosine similarity to cached training embeddings\n#   - Returns real training image paths with their dataset labels\n# Similarity is a feature-space ranking, not evidence of clinical equivalence.\n# ============================================================\n\n\ndef find_similar_cases(\n    query_image: Union[str, np.ndarray],\n    k: int = 3,\n    extractor: keras.Model = embedding_extractor,\n    stored_embeddings: np.ndarray = train_embeddings,\n    stored_labels: np.ndarray = train_labels_arr,\n    stored_filepaths: List[str] = train_fps,\n) -> List[Dict[str, Any]]:\n    \"\"\"Return the top-k nearest dataset-labeled training examples.\"\"\"\n    if isinstance(query_image, str):\n        img_arr = preprocess_image(query_image)\n    else:\n        img_arr = query_image.copy()\n\n    query_tensor = np.expand_dims(img_arr, axis=0) if img_arr.ndim == 3 else img_arr\n    query_embedding = extractor(query_tensor, training=False).numpy()\n    query_embedding /= np.linalg.norm(query_embedding, axis=1, keepdims=True) + 1e-10\n    similarities = np.dot(stored_embeddings, query_embedding.T).squeeze()\n    top_indices = np.argsort(similarities)[::-1][:k]\n\n    results = []\n    for rank, index in enumerate(top_indices, start=1):\n        stage = int(stored_labels[index])\n        results.append({\n            \"rank\": rank,\n            \"filepath\": stored_filepaths[index],\n            \"label\": stage,\n            \"stage_name\": Config.CLASS_NAMES[stage],\n            \"cosine_similarity\": float(similarities[index]),\n        })\n    return results\n\n\nprint(\"[Section 9] find_similar_cases(query_image, k=3) is defined.\")","metadata":{"id":"b772b6c2c6fb4e03852881fc6d2c639f","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:41:05.476003Z","iopub.execute_input":"2026-09-30T20:41:05.476412Z","iopub.status.idle":"2026-09-30T20:41:05.484769Z","shell.execute_reply.started":"2026-09-30T20:41:05.476386Z","shell.execute_reply":"2026-09-30T20:41:05.483898Z"}},"outputs":[],"execution_count":null},{"id":"d07ac056386a4f48be8c37e8f1cae973","cell_type":"code","source":"# ============================================================\n# SECTION 9.4 — VISUAL DEMONSTRATION OF SIMILAR-CASE RETRIEVAL\n# The query is validation-only; reference images come from train_df.\n# All gallery images are real files. Captions show supplied dataset labels.\n# ============================================================\n\n\ndef visualize_similar_cases_demo(\n    query_filepath: str,\n    k: int = 3,\n    save_path: Optional[str] = None,\n) -> List[Dict[str, Any]]:\n    \"\"\"Plot a validation fundus image beside its nearest training references.\"\"\"\n    similar_cases = find_similar_cases(query_filepath, k=k)\n    query_img = preprocess_image(query_filepath)\n    fig, axes = plt.subplots(1, k + 1, figsize=(4 * (k + 1), 4.5))\n\n    axes[0].imshow(query_img)\n    axes[0].set_title(\"QUERY: VALIDATION IMAGE\", fontsize=11, fontweight=\"bold\")\n    axes[0].axis(\"off\")\n\n    for index, case in enumerate(similar_cases, start=1):\n        case_img = preprocess_image(case[\"filepath\"])\n        axes[index].imshow(case_img)\n        axes[index].set_title(\n            f\"Training reference #{case['rank']}\\n\"\n            f\"Dataset label: {case['stage_name']} ({case['label']})\\n\"\n            f\"Cosine similarity: {case['cosine_similarity']:.3f}\",\n            fontsize=9,\n        )\n        axes[index].axis(\"off\")\n\n    plt.suptitle(\n        \"Similar-Case Retrieval Demo — Similarity Is Not Clinical Confirmation\",\n        fontsize=12,\n        fontweight=\"bold\",\n    )\n    plt.tight_layout()\n\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n        print(f\"[Retrieval] Saved demonstration figure -> {save_path}\")\n\n    plt.show()\n    print(\"\\nNearest training references for validation query:\")\n    for case in similar_cases:\n        print(\n            f\"  Rank {case['rank']}: dataset label [{case['stage_name']}] | \"\n            f\"cosine similarity={case['cosine_similarity']:.4f}\"\n        )\n        print(f\"         Image file: {case['filepath']}\")\n    return similar_cases\n\n\nsample_query_row = val_df[val_df[\"label\"] >= 2].iloc[0]\nsample_retrievals = visualize_similar_cases_demo(\n    sample_query_row[\"filepath\"],\n    k=3,\n    save_path=os.path.join(Config.REPORTS_DIR, \"similar_cases_demo.png\"),\n)","metadata":{"id":"d07ac056386a4f48be8c37e8f1cae973","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:41:17.616412Z","iopub.execute_input":"2026-09-30T20:41:17.616774Z","iopub.status.idle":"2026-09-30T20:41:21.139969Z","shell.execute_reply.started":"2026-09-30T20:41:17.616707Z","shell.execute_reply":"2026-09-30T20:41:21.139027Z"}},"outputs":[],"execution_count":null},{"id":"a71869f6","cell_type":"code","source":"# ============================================================\n# SECTION 9.4b — VALIDATION RETRIEVAL COMPARISON\n# Compare learned-embedding neighbors with random training references.\n# Metrics measure dataset-label agreement, not clinical usefulness.\n# ============================================================\n\nRETRIEVAL_EVAL_PER_CLASS = 100\nRETRIEVAL_RANDOM_REPEATS = 20\nRETRIEVAL_K = 3\n\nif not set(val_df[\"patient_id\"].dropna()).isdisjoint(set(train_df[\"patient_id\"].dropna())):\n    raise AssertionError(\"Retrieval validation queries overlap training patients.\")\nif not set(val_df[\"duplicate_group\"].dropna()).isdisjoint(set(train_df[\"duplicate_group\"].dropna())):\n    raise AssertionError(\"Retrieval validation queries overlap training duplicate groups.\")\n\nretrieval_eval_parts = []\nfor retrieval_class_id, class_rows in val_df.groupby(\"label\", sort=True):\n    retrieval_eval_parts.append(\n        class_rows.sample(\n            min(RETRIEVAL_EVAL_PER_CLASS, len(class_rows)),\n            random_state=Config.SEED + int(retrieval_class_id),\n        )\n    )\nretrieval_eval_df = pd.concat(retrieval_eval_parts).reset_index(drop=True)\nretrieval_eval_dataset = _build_tuning_dataset(\n    retrieval_eval_df,\n    Config.IMG_SIZE,\n    shuffle=False,\n    augment=False,\n)\n\nquery_embedding_batches = []\nquery_label_batches = []\nfor query_images, query_labels in retrieval_eval_dataset:\n    query_features = embedding_extractor(query_images, training=False).numpy()\n    query_features /= np.linalg.norm(query_features, axis=1, keepdims=True) + 1e-10\n    query_embedding_batches.append(query_features)\n    query_label_batches.append(np.argmax(query_labels.numpy(), axis=1))\nquery_embeddings = np.concatenate(query_embedding_batches, axis=0)\nquery_true_labels = np.concatenate(query_label_batches, axis=0).astype(int)\nif len(query_true_labels) != len(retrieval_eval_df):\n    raise ValueError(\"Retrieval query labels and validation rows are misaligned.\")\n\nretrieval_similarities = query_embeddings @ train_embeddings.T\nnearest_indices = np.argpartition(\n    retrieval_similarities,\n    -RETRIEVAL_K,\n    axis=1,\n)[:, -RETRIEVAL_K:]\nnearest_scores = np.take_along_axis(retrieval_similarities, nearest_indices, axis=1)\nnearest_order = np.argsort(nearest_scores, axis=1)[:, ::-1]\nnearest_indices = np.take_along_axis(nearest_indices, nearest_order, axis=1)\nnearest_labels = train_labels_arr[nearest_indices]\nembedding_top1_agreement = nearest_labels[:, 0] == query_true_labels\nembedding_top3_hit = np.any(nearest_labels == query_true_labels[:, np.newaxis], axis=1)\nembedding_top3_grade_error = np.min(\n    np.abs(nearest_labels - query_true_labels[:, np.newaxis]), axis=1\n)\n\nrandom_generator = np.random.default_rng(Config.SEED)\nrandom_top1_agreement_repeats = []\nrandom_top3_hit_repeats = []\nrandom_top3_grade_error_repeats = []\nrandom_top3_hit_by_stage_repeats = []\nfor _ in range(RETRIEVAL_RANDOM_REPEATS):\n    random_indices = np.stack([\n        random_generator.choice(len(train_labels_arr), size=RETRIEVAL_K, replace=False)\n        for _ in range(len(query_true_labels))\n    ])\n    random_labels = train_labels_arr[random_indices]\n    random_top1_agreement_repeats.append(random_labels[:, 0] == query_true_labels)\n    random_top3_hit_repeats.append(\n        np.any(random_labels == query_true_labels[:, np.newaxis], axis=1)\n    )\n    random_top3_grade_error_repeats.append(\n        np.min(np.abs(random_labels - query_true_labels[:, np.newaxis]), axis=1)\n    )\n    random_top3_hit_by_stage_repeats.append([\n        np.mean(\n            np.any(random_labels[query_true_labels == stage] == stage, axis=1)\n        )\n        for stage in range(Config.NUM_CLASSES)\n    ])\n\nrandom_top1_agreement = np.asarray(random_top1_agreement_repeats)\nrandom_top3_hit = np.asarray(random_top3_hit_repeats)\nrandom_top3_grade_error = np.asarray(random_top3_grade_error_repeats)\nretrieval_overall_comparison = pd.DataFrame([\n    {\n        \"method\": \"Embedding nearest neighbors\",\n        \"queries\": len(query_true_labels),\n        \"top1_dataset_label_agreement\": float(np.mean(embedding_top1_agreement)),\n        \"top3_dataset_label_hit_rate\": float(np.mean(embedding_top3_hit)),\n        \"top3_best_absolute_grade_error\": float(np.mean(embedding_top3_grade_error)),\n        \"random_repeat_sd\": np.nan,\n    },\n    {\n        \"method\": f\"Random references (mean of {RETRIEVAL_RANDOM_REPEATS} repeats)\",\n        \"queries\": len(query_true_labels),\n        \"top1_dataset_label_agreement\": float(np.mean(random_top1_agreement)),\n        \"top3_dataset_label_hit_rate\": float(np.mean(random_top3_hit)),\n        \"top3_best_absolute_grade_error\": float(np.mean(random_top3_grade_error)),\n        \"random_repeat_sd\": float(np.std(np.mean(random_top3_hit, axis=1), ddof=1)),\n    },\n])\nretrieval_stage_records = []\nrandom_stage_hits = np.asarray(random_top3_hit_by_stage_repeats)\nfor stage_id, stage_name in enumerate(Config.CLASS_NAMES):\n    stage_queries = query_true_labels == stage_id\n    retrieval_stage_records.extend([\n        {\n            \"class_id\": stage_id,\n            \"class_name\": stage_name,\n            \"method\": \"Embedding nearest neighbors\",\n            \"queries\": int(stage_queries.sum()),\n            \"top3_dataset_label_hit_rate\": float(np.mean(embedding_top3_hit[stage_queries])),\n        },\n        {\n            \"class_id\": stage_id,\n            \"class_name\": stage_name,\n            \"method\": \"Random references\",\n            \"queries\": int(stage_queries.sum()),\n            \"top3_dataset_label_hit_rate\": float(np.mean(random_stage_hits[:, stage_id])),\n        },\n    ])\nretrieval_by_stage_comparison = pd.DataFrame(retrieval_stage_records)\nretrieval_overall_comparison.to_csv(\n    Path(Config.REPORTS_DIR) / \"retrieval_validation_metrics.csv\", index=False\n)\nretrieval_by_stage_comparison.to_csv(\n    Path(Config.REPORTS_DIR) / \"retrieval_validation_by_stage.csv\", index=False\n)\ndisplay(retrieval_overall_comparison.round(4))\ndisplay(retrieval_by_stage_comparison.pivot(\n    index=[\"class_id\", \"class_name\", \"queries\"],\n    columns=\"method\",\n    values=\"top3_dataset_label_hit_rate\",\n).round(4))\n\nretrieval_plot_data = retrieval_by_stage_comparison.pivot(\n    index=\"class_name\", columns=\"method\", values=\"top3_dataset_label_hit_rate\"\n).reindex(Config.CLASS_NAMES)\nfig, ax = plt.subplots(figsize=(11, 6))\nretrieval_plot_data.plot(\n    kind=\"bar\",\n    ax=ax,\n    color=[\"#367C8A\", \"#D8784B\"],\n    width=0.78,\n)\nax.set_title(\"Validation Top-3 Dataset-Label Match Rate: Embeddings vs Random References\")\nax.set_xlabel(\"Validation query grade\")\nax.set_ylabel(\"Fraction with at least one matching dataset label\")\nax.set_ylim(0, 1)\nax.legend(title=\"Reference selection\")\nax.grid(axis=\"y\", alpha=0.25)\nax.tick_params(axis=\"x\", rotation=18)\nfig.tight_layout()\nretrieval_comparison_plot_path = Path(Config.REPORTS_DIR) / \"retrieval_validation_comparison.png\"\nfig.savefig(retrieval_comparison_plot_path, dpi=160, bbox_inches=\"tight\")\nplt.show()\nprint(\"[Retrieval Evaluation] Dataset-label agreement is not evidence of clinical equivalence or improved diagnosis.\")\nprint(f\"[Retrieval Evaluation] Metrics saved -> {Config.REPORTS_DIR}/retrieval_validation_metrics.csv\")\nprint(f\"[Retrieval Evaluation] Plot saved -> {retrieval_comparison_plot_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:41:28.161808Z","iopub.execute_input":"2026-09-30T20:41:28.162194Z","iopub.status.idle":"2026-09-30T20:41:58.834147Z","shell.execute_reply.started":"2026-09-30T20:41:28.162157Z","shell.execute_reply":"2026-09-30T20:41:58.833357Z"}},"outputs":[],"execution_count":null},{"id":"72c4481d39ec4107b228ff87f3d1c7cc","cell_type":"code","source":"# ============================================================\n# SECTION 9.5 — SECTION SUMMARY\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  SECTION 9 COMPLETE — Innovation Feature A\")\nprint(\"=\" * 60)\nprint(\"  build_embedding_extractor()           : 256-D latent projection\")\nprint(\"  extract_and_cache_training_embeddings(): cached to embeddings.npz\")\nprint(\"  find_similar_cases(query_img, k=3)    : cosine-metric CBR engine\")\nprint(f\"  Visual demonstration saved            : {Config.REPORTS_DIR}/similar_cases_demo.png\")\nprint(\"  Next -> Innovation B: Multi-Agent Clinical Decision Pipeline\")\nprint(\"=\" * 60)","metadata":{"id":"72c4481d39ec4107b228ff87f3d1c7cc","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:06.819662Z","iopub.execute_input":"2026-09-30T20:42:06.82012Z","iopub.status.idle":"2026-09-30T20:42:06.825847Z","shell.execute_reply.started":"2026-09-30T20:42:06.820092Z","shell.execute_reply":"2026-09-30T20:42:06.825016Z"}},"outputs":[],"execution_count":null},{"id":"0bf012827b564ebb9d6f785afa7eaf1a","cell_type":"markdown","source":"## Section 10: Innovation Feature B — Multi-Agent Clinical Decision Pipeline\n*(Grounded in Coursework Material: `6. Multi-Agent_AI_Blueprint.pdf` & `7/8. Defense_Multi_Agent_LLM.ipynb`)*\n\n**Goal:** Build a robust, safety-critical multi-agent clinical decision support (CDS) system\ncomposed of four specialized, decoupled agents coordinated by a strict governance layer.\n\n```\n                  ┌─────────────────────────────────────┐\n                  │          Query Fundus Image         │\n                  └──────────────────┬──────────────────┘\n                                     │\n                                     ▼\n                  ┌─────────────────────────────────────┐\n                  │           DiagnosisAgent            │\n                  │   - EfficientNetB3 CNN Inference    │\n                  │   - Stage & Confidence Prediction   │\n                  └─────────┬─────────────────┬─────────┘\n                            │                 │\n              ┌─────────────┘                 └─────────────┐\n              ▼                                             ▼\n┌─────────────────────────────┐               ┌─────────────────────────────┐\n│     ExplainabilityAgent     │               │        AdvisoryAgent        │\n│   - Grad-CAM Heatmap        │               │   - Protocol Guidance       │\n│   - Quadrant Description    │               │   - Clinical Next Steps     │\n│   - Similar-Case Retrieval  │               │   - Medical Disclaimer      │\n└─────────────┬───────────────┘               └─────────────┬───────────────┘\n              │                                             │\n              └─────────────────────┬───────────────────────┘\n                                    │\n                                    ▼\n                  ┌─────────────────────────────────────┐\n                  │           GovernanceAgent           │\n                  │   - Safety Gating Layer             │\n                  │   - Is Confidence >= Threshold?     │\n                  └─────────┬─────────────────┬─────────┘\n                            │                 │\n             NO (Conf < 70%)│                 │YES (Conf >= 70%)\n                            ▼                 ▼\n          ┌───────────────────────────┐     ┌───────────────────────────┐\n          │ OVERRIDE OUTPUT:          │     │ RELEASE OUTPUT:           │\n          │ \"FLAGGED FOR HUMAN REVIEW\"│     │ Validated Diagnosis,      │\n          │ Guidance withheld         │     │ Guidance & Explanations   │\n          └───────────────────────────┘     └───────────────────────────┘\n```\n\n### Architectural & Safety Rationale in a Clinical Context\n\n1. **Why not a monolithic chatbot or single model wrapper?**\n   - Single-model wrappers conflate perceptual classification with clinical protocol and risk management.\n   - If an unconstrained model outputs both a probability and unstructured text advice, hallucination\n     or low-confidence guesswork can be presented as medical guidance with disastrous patient consequences.\n\n2. **Specialized Decoupled Responsibilities:**\n   - **`DiagnosisAgent`** has one isolated role: perceptual classification.\n   - **`ExplainabilityAgent`** provides verifiable evidentiary backing (Grad-CAM spatial localization +\n     historical case retrieval).\n   - **`AdvisoryAgent`** references validated clinical protocols (NHS/AAO staging guidelines), ensuring\n     recommendations strictly adhere to evidence-based medical consensus.\n   - **`GovernanceAgent`** functions as an **autonomous safety officer**: it intercepts predictions and has\n     the authority to veto automated recommendations when the model's epistemic uncertainty is too high.","metadata":{"id":"0bf012827b564ebb9d6f785afa7eaf1a","language":"markdown"}},{"id":"5f16b3a7ca48420f9621530910cf76ed","cell_type":"code","source":"# ============================================================\n# SECTION 10.1 — AGENT 1: DiagnosisAgent\n#\n# Runs the trained EfficientNetB3 CNN model on an input image,\n# extracts softmax probabilities, and returns the discrete predicted\n# stage, stage name, and confidence score.\n# ============================================================\n\nclass DiagnosisAgent:\n    \"\"\"\n    Agent responsible for perceptual classification and disease staging.\n\n    Evaluates fundus photographs through the trained EfficientNetB3 network\n    and extracts class probabilities across all 5 ICDR stages.\n    \"\"\"\n\n    def __init__(self, trained_model: keras.Model):\n        \"\"\"\n        Initialize with the trained DR detection model.\n\n        Args:\n            trained_model: Trained Keras model outputting 5-class softmax distribution.\n        \"\"\"\n        self.model = trained_model\n        self.class_names = Config.CLASS_NAMES\n\n    def process(self, image: Union[str, np.ndarray]) -> Dict[str, Any]:\n        \"\"\"\n        Run inference and return structured diagnosis outputs.\n\n        Args:\n            image: Image filepath string or preprocessed RGB array of shape (Config.IMG_SIZE, Config.IMG_SIZE, 3).\n\n        Returns:\n            Dictionary containing:\n                - 'predicted_stage': Integer ICDR stage (0 to 4).\n                - 'stage_name': Human-readable stage string.\n                - 'confidence': Float confidence score in [0.0, 1.0].\n                - 'probabilities': Dict mapping each stage name to its softmax probability.\n                - 'preprocessed_image': Normalized float32 image array (Config.IMG_SIZE, Config.IMG_SIZE, 3).\n        \"\"\"\n        # Preprocess if path string is provided\n        if isinstance(image, str):\n            preproc_img = preprocess_image(image)\n        else:\n            preproc_img = image.copy()\n\n        # Add batch dimension for inference\n        if preproc_img.ndim == 3:\n            tensor_in = np.expand_dims(preproc_img, axis=0)\n        else:\n            tensor_in = preproc_img\n\n        # Execute model forward pass\n        raw_probs = self.model(tensor_in, training=False).numpy()[0]\n\n        pred_stage = int(np.argmax(raw_probs))\n        confidence = float(raw_probs[pred_stage])\n\n        stage_probs = {\n            self.class_names[i]: float(raw_probs[i])\n            for i in range(len(self.class_names))\n        }\n\n        return {\n            \"predicted_stage\": pred_stage,\n            \"stage_name\": self.class_names[pred_stage],\n            \"confidence\": confidence,\n            \"probabilities\": stage_probs,\n            \"preprocessed_image\": preproc_img,\n        }\n\n\n# Instantiate DiagnosisAgent\ndiagnosis_agent = DiagnosisAgent(model)\nprint(\"[DiagnosisAgent] Initialized and verified.\")","metadata":{"id":"5f16b3a7ca48420f9621530910cf76ed","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:13.21239Z","iopub.execute_input":"2026-09-30T20:42:13.212825Z","iopub.status.idle":"2026-09-30T20:42:13.221317Z","shell.execute_reply.started":"2026-09-30T20:42:13.212799Z","shell.execute_reply":"2026-09-30T20:42:13.220407Z"}},"outputs":[],"execution_count":null},{"id":"e16e88968a27491c84c197106c7542a8","cell_type":"code","source":"# ============================================================\n# SECTION 10.2 — AGENT 2: ExplainabilityAgent\n#\n# Generates dual-channel explainability:\n#   1. Intra-image Grad-CAM heatmap & anatomical quadrant description\n#   2. Inter-image Similar-Case Retrieval (Innovation A)\n# ============================================================\n\nclass ExplainabilityAgent:\n    \"\"\"\n    Agent responsible for clinical interpretability and evidentiary backing.\n\n    Combines spatial Grad-CAM saliency with historical case-based retrieval\n    to substantiate the automated diagnosis.\n    \"\"\"\n\n    def __init__(\n        self,\n        cam_model: keras.Model = gradcam_model,\n        extractor: keras.Model = embedding_extractor,\n        stored_embeddings: np.ndarray = train_embeddings,\n        stored_labels: np.ndarray = train_labels_arr,\n        stored_filepaths: List[str] = train_fps,\n    ):\n        \"\"\"\n        Initialize with Grad-CAM model and retrieval components.\n        \"\"\"\n        self.cam_model = cam_model\n        self.extractor = extractor\n        self.stored_embeddings = stored_embeddings\n        self.stored_labels = stored_labels\n        self.stored_filepaths = stored_filepaths\n\n    def process(\n        self,\n        image_input: Union[str, np.ndarray],\n        diagnosis_result: Dict[str, Any],\n        k_similar: int = 3,\n    ) -> Dict[str, Any]:\n        \"\"\"\n        Generate comprehensive spatial and comparative explanations.\n\n        Args:\n            image_input: Image filepath or preprocessed array.\n            diagnosis_result: Output dict from DiagnosisAgent.process().\n            k_similar: Number of nearest cases to retrieve.\n\n        Returns:\n            Dictionary containing:\n                - 'gradcam_heatmap': 2D float array normalized to [0, 1].\n                - 'overlay_image': RGB uint8 array with superimposed Jet heatmap.\n                - 'quadrant_explanation': Natural language description of peak lesions.\n                - 'similar_cases': List of k retrieval metadata dicts.\n        \"\"\"\n        preproc_img = diagnosis_result[\"preprocessed_image\"]\n        stage = diagnosis_result[\"predicted_stage\"]\n        conf = diagnosis_result[\"confidence\"]\n\n        # 1. Compute Grad-CAM heatmap\n        heatmap, _, _ = compute_gradcam_heatmap(\n            preproc_img,\n            cam_model=self.cam_model,\n            pred_index=stage,\n        )\n\n        # 2. Superimpose heatmap onto preprocessed image\n        overlay = superimpose_gradcam(preproc_img, heatmap, alpha=0.4)\n\n        # 3. Generate natural-language quadrant description\n        quadrant_text = generate_quadrant_explanation(heatmap, stage, conf)\n\n        # 4. Perform case-based retrieval\n        similar_cases = find_similar_cases(\n            preproc_img,\n            k=k_similar,\n            extractor=self.extractor,\n            stored_embeddings=self.stored_embeddings,\n            stored_labels=self.stored_labels,\n            stored_filepaths=self.stored_filepaths,\n        )\n\n        return {\n            \"gradcam_heatmap\": heatmap,\n            \"overlay_image\": overlay,\n            \"quadrant_explanation\": quadrant_text,\n            \"similar_cases\": similar_cases,\n        }\n\n\n# Instantiate ExplainabilityAgent\nexplainability_agent = ExplainabilityAgent()\nprint(\"[ExplainabilityAgent] Initialized and verified.\")","metadata":{"id":"e16e88968a27491c84c197106c7542a8","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:23.907318Z","iopub.execute_input":"2026-09-30T20:42:23.908169Z","iopub.status.idle":"2026-09-30T20:42:23.916803Z","shell.execute_reply.started":"2026-09-30T20:42:23.908137Z","shell.execute_reply":"2026-09-30T20:42:23.915973Z"}},"outputs":[],"execution_count":null},{"id":"ca875f5bbf934db38e5547ce8d681340","cell_type":"code","source":"# ============================================================\n# SECTION 10.3 — AGENT 3: AdvisoryAgent\n#\n# Maps predicted disease stages to standard clinical protocols\n# (follow-up intervals, specialist referrals, diagnostic testing)\n# based on international ophthalmology guidelines (AAO / NHS).\n# Appends a mandatory medical disclaimer.\n# ============================================================\n\nclass AdvisoryAgent:\n    \"\"\"\n    Agent responsible for clinical protocol alignment and next-step recommendations.\n\n    Translates discrete DR stages into evidence-based action plans and attaches\n    mandatory regulatory disclaimers.\n    \"\"\"\n\n    # Clinical protocol lookup adhering to American Academy of Ophthalmology (AAO) guidelines\n    CLINICAL_PROTOCOLS: Dict[int, Dict[str, str]] = {\n        0: {\n            \"clinical_urgency\": \"Routine / Annual\",\n            \"action_plan\": (\n                \"No diabetic microvascular lesions observed. Continue routine annual diabetic retinopathy \"\n                \"surveillance. Reinforce systemic glycemic control (HbA1c target < 7.0%), blood pressure \"\n                \"management, and lipid monitoring in primary care.\"\n            ),\n            \"follow_up\": \"Annual fundus photography examination (12 months).\",\n        },\n        1: {\n            \"clinical_urgency\": \"Non-Urgent / Close Monitoring\",\n            \"action_plan\": (\n                \"Mild non-proliferative changes detected (microaneurysms only). Primary care optimization \"\n                \"of blood glucose, blood pressure, and renal profile. Assess for potential subclinical diabetic \"\n                \"macular edema if visual acuity is reduced.\"\n            ),\n            \"follow_up\": \"Repeat dilated fundus examination within 6 to 9 months.\",\n        },\n        2: {\n            \"clinical_urgency\": \"Semi-Urgent / Retinal Specialist Referral\",\n            \"action_plan\": (\n                \"Moderate non-proliferative changes present (microaneurysms, dot-and-blot hemorrhages, hard \"\n                \"exudates, or cotton wool spots). High clinical risk of disease progression. Schedule formal \"\n                \"optical coherence tomography (OCT) and comprehensive biomicroscopic dilated evaluation.\"\n            ),\n            \"follow_up\": \"Refer to an ophthalmologist for comprehensive evaluation within 3 to 6 months.\",\n        },\n        3: {\n            \"clinical_urgency\": \"Urgent Specialist Intervention\",\n            \"action_plan\": (\n                \"Severe non-proliferative retinopathy detected (meeting the '4-2-1 rule': severe intraretinal \"\n                \"hemorrhages in all 4 quadrants, venous beading in 2+ quadrants, or IRMA in 1+ quadrant). \"\n                \"Imminent risk of progression to proliferative DR (50% risk within 1 year). Urgent vitreoretinal \"\n                \"assessment for panretinal photocoagulation (PRP) or anti-VEGF injection readiness.\"\n            ),\n            \"follow_up\": \"Urgent ophthalmology referral within 2 to 4 weeks.\",\n        },\n        4: {\n            \"clinical_urgency\": \"Emergent / Sight-Threatening\",\n            \"action_plan\": (\n                \"Proliferative Diabetic Retinopathy (PDR) identified. Presence of neovascularization, \"\n                \"vitreous/preretinal hemorrhage, or fibrovascular proliferation posing immediate risk of tractional \"\n                \"retinal detachment and irreversible vision loss. Requires immediate specialist intervention \"\n                \"(panretinal photocoagulation, anti-VEGF therapy, or vitrectomy).\"\n            ),\n            \"follow_up\": \"Emergent referral to a vitreoretinal subspecialist within 24 to 72 hours.\",\n        },\n    }\n\n    DISCLAIMER: str = (\n        \"IMPORTANT REGULATORY NOTICE: This automated advisory summary is generated by an artificial \"\n        \"intelligence clinical decision-support research system. It does NOT constitute a confirmed medical \"\n        \"diagnosis, prescription, or therapeutic mandate. All findings must be clinically corroborated by a \"\n        \"licensed medical professional (ophthalmologist, optometrist, or retinal specialist).\"\n    )\n\n    def process(self, stage: int) -> Dict[str, str]:\n        \"\"\"\n        Retrieve evidence-based guidance for the diagnosed DR stage.\n\n        Args:\n            stage: Integer ICDR DR stage (0 to 4).\n\n        Returns:\n            Dictionary containing 'clinical_urgency', 'action_plan', 'follow_up', and 'disclaimer'.\n        \"\"\"\n        protocol = self.CLINICAL_PROTOCOLS.get(\n            stage,\n            {\n                \"clinical_urgency\": \"Unknown Stage\",\n                \"action_plan\": \"Unrecognized disease stage. Manual clinical review required.\",\n                \"follow_up\": \"Immediate manual examination.\",\n            },\n        )\n\n        return {\n            \"clinical_urgency\": protocol[\"clinical_urgency\"],\n            \"action_plan\": protocol[\"action_plan\"],\n            \"follow_up\": protocol[\"follow_up\"],\n            \"disclaimer\": self.DISCLAIMER,\n        }\n\n\n# Instantiate AdvisoryAgent\nadvisory_agent = AdvisoryAgent()\nprint(\"[AdvisoryAgent] Initialized and verified.\")","metadata":{"id":"ca875f5bbf934db38e5547ce8d681340","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:29.063359Z","iopub.execute_input":"2026-09-30T20:42:29.063705Z","iopub.status.idle":"2026-09-30T20:42:29.072947Z","shell.execute_reply.started":"2026-09-30T20:42:29.063678Z","shell.execute_reply":"2026-09-30T20:42:29.07215Z"}},"outputs":[],"execution_count":null},{"id":"d23225522aa14ec082ca986e18ca05f8","cell_type":"code","source":"# ============================================================\n# SECTION 10.4 — AGENT 4: GovernanceAgent & RUN_PIPELINE\n#\n# Safety Gating Layer:\n# If DiagnosisAgent confidence is below CONFIDENCE_THRESHOLD (default 70%):\n#   - Actively OVERRIDES the output\n#   - Sets 'flagged_for_review': True\n#   - Withholds automated action plans\n#   - Returns an urgent human review triage notice\n#\n# run_pipeline(image) orchestrates all 4 agents in sequence.\n# ============================================================\n\nclass GovernanceAgent:\n    \"\"\"\n    Safety and regulatory oversight agent.\n\n    Monitors the outputs of upstream agents and enforces clinical risk thresholds.\n    Actively gates and overrides recommendations when model uncertainty is elevated.\n    \"\"\"\n\n    def __init__(self, confidence_threshold: float = Config.CONFIDENCE_THRESHOLD):\n        \"\"\"\n        Initialize with configurable safety threshold.\n\n        Args:\n            confidence_threshold: Minimum confidence [0.0, 1.0] required to release\n                                  automated guidance. Default 0.70 (70%).\n        \"\"\"\n        self.confidence_threshold = confidence_threshold\n\n    def evaluate(\n        self,\n        diagnosis: Dict[str, Any],\n        explanation: Dict[str, Any],\n        advisory: Dict[str, str],\n    ) -> Dict[str, Any]:\n        \"\"\"\n        Evaluate upstream agent outputs and enforce safety gating.\n\n        Args:\n            diagnosis: Output from DiagnosisAgent.\n            explanation: Output from ExplainabilityAgent.\n            advisory: Output from AdvisoryAgent.\n\n        Returns:\n            Unified structured clinical result dictionary.\n        \"\"\"\n        conf = diagnosis[\"confidence\"]\n        stage = diagnosis[\"predicted_stage\"]\n        stage_name = diagnosis[\"stage_name\"]\n\n        # Check safety threshold\n        if conf < self.confidence_threshold:\n            flagged = True\n            governance_status = \"FLAGGED_FOR_HUMAN_REVIEW\"\n            gating_message = (\n                f\"SAFETY OVERRIDE ACTIVATED: Model confidence ({conf*100:.1f}%) is below the required \"\n                f\"clinical safety threshold ({self.confidence_threshold*100:.0f}%). Automated recommendations \"\n                f\"have been withheld to protect patient safety. This case has been routed to human ophthalmologist \"\n                f\"triage for manual visual inspection.\"\n            )\n            # Override advisory guidance with safe triage protocol\n            controlled_advisory = {\n                \"clinical_urgency\": \"Triage Required (Low AI Confidence)\",\n                \"action_plan\": gating_message,\n                \"follow_up\": \"Do NOT initiate treatment based on AI output. Perform manual dilated examination.\",\n                \"disclaimer\": advisory[\"disclaimer\"],\n            }\n        else:\n            flagged = False\n            governance_status = \"AUTOMATED_RECOMMENDATION_APPROVED\"\n            gating_message = (\n                f\"Safety check passed: Model confidence ({conf*100:.1f}%) meets or exceeds the required \"\n                f\"safety threshold ({self.confidence_threshold*100:.0f}%).\"\n            )\n            controlled_advisory = advisory\n\n        return {\n            \"flagged_for_review\": flagged,\n            \"governance_status\": governance_status,\n            \"governance_message\": gating_message,\n            \"diagnosis\": {\n                \"predicted_stage\": stage,\n                \"stage_name\": stage_name,\n                \"confidence\": conf,\n                \"probabilities\": diagnosis[\"probabilities\"],\n            },\n            \"explanation\": {\n                \"gradcam_heatmap\": explanation[\"gradcam_heatmap\"],\n                \"overlay_image\": explanation[\"overlay_image\"],\n                \"quadrant_explanation\": explanation[\"quadrant_explanation\"],\n                \"similar_cases\": explanation[\"similar_cases\"],\n            },\n            \"advisory\": controlled_advisory,\n        }\n\n\n# Instantiate GovernanceAgent\ngovernance_agent = GovernanceAgent(confidence_threshold=Config.CONFIDENCE_THRESHOLD)\n\n\ndef run_pipeline(\n    image_input: Union[str, np.ndarray],\n    diag_agent: DiagnosisAgent = diagnosis_agent,\n    expl_agent: ExplainabilityAgent = explainability_agent,\n    adv_agent: AdvisoryAgent = advisory_agent,\n    gov_agent: GovernanceAgent = governance_agent,\n) -> Dict[str, Any]:\n    \"\"\"\n    Execute the multi-agent clinical decision pipeline on a patient fundus image.\n\n    Chains all 4 specialist agents in strict sequence:\n        DiagnosisAgent -> ExplainabilityAgent -> AdvisoryAgent -> GovernanceAgent (Gate)\n\n    Args:\n        image_input: Absolute file path or preprocessed RGB image array (Config.IMG_SIZE, Config.IMG_SIZE, 3).\n        diag_agent: Initialized DiagnosisAgent instance.\n        expl_agent: Initialized ExplainabilityAgent instance.\n        adv_agent: Initialized AdvisoryAgent instance.\n        gov_agent: Initialized GovernanceAgent instance.\n\n    Returns:\n        Structured result dictionary with diagnosis, explainability, advisory,\n        and safety gating status.\n    \"\"\"\n    # 1. Run perception / diagnosis\n    diag_result = diag_agent.process(image_input)\n\n    # 2. Run explainability & case retrieval\n    expl_result = expl_agent.process(image_input, diag_result, k_similar=3)\n\n    # 3. Formulate standard clinical advisory\n    advis_result = adv_agent.process(diag_result[\"predicted_stage\"])\n\n    # 4. Enforce safety gating & governance\n    final_result = gov_agent.evaluate(diag_result, expl_result, advis_result)\n\n    return final_result\n\n\nprint(\"[Multi-Agent Pipeline] run_pipeline(image) compiled and ready.\")","metadata":{"id":"d23225522aa14ec082ca986e18ca05f8","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:34.546633Z","iopub.execute_input":"2026-09-30T20:42:34.547245Z","iopub.status.idle":"2026-09-30T20:42:34.558108Z","shell.execute_reply.started":"2026-09-30T20:42:34.547219Z","shell.execute_reply":"2026-09-30T20:42:34.557134Z"}},"outputs":[],"execution_count":null},{"id":"86e1cb4938ec47f0b6ab24103f787498","cell_type":"code","source":"# ============================================================\n# SECTION 10.5 — VALIDATION EXAMPLES FOR MULTI-AGENT PIPELINE  # CHANGED: avoid post-evaluation test reuse\n#\n# Case 1: Standard pipeline flow\n# Case 2: Safety-gate threshold demonstration\n# ============================================================\n\ndef display_pipeline_result(result: Dict[str, Any], query_title: str) -> None:\n    \"\"\"\n    Format and print a structured report of the multi-agent pipeline outcome.\n    \"\"\"\n    print(\"\\n\" + \"=\" * 80)\n    print(f\"  MULTI-AGENT DECISION PIPELINE REPORT: {query_title.upper()}\")\n    print(\"=\" * 80)\n\n    diag = result[\"diagnosis\"]\n    gov  = result[\"governance_status\"]\n    flag = result[\"flagged_for_review\"]\n    adv  = result[\"advisory\"]\n    expl = result[\"explanation\"]\n\n    print(f\"  Diagnosis Result        : {diag['stage_name']} (Stage {diag['predicted_stage']})\")\n    print(f\"  Model Confidence        : {diag['confidence']*100:.2f}%\")\n    print(f\"  Governance Status       : {gov}\")\n    print(f\"  Flagged for Review?     : {'YES - HUMAN REVIEW MANDATORY' if flag else 'NO - AUTOMATION APPROVED'}\")\n    print(f\"  Safety Gate Details     : {result['governance_message']}\")\n    print(\"-\" * 80)\n    print(f\"  Clinical Urgency        : {adv['clinical_urgency']}\")\n    print(f\"  Recommended Action Plan : {adv['action_plan']}\")\n    print(f\"  Follow-up Protocol      : {adv['follow_up']}\")\n    print(\"-\" * 80)\n    print(f\"  Spatial Explanation     : {expl['quadrant_explanation']}\")\n    print(\"  Case-Based Evidence     : Retrieved top-3 matching historical training cases:\")\n    for c in expl[\"similar_cases\"]:\n        print(f\"    - Rank {c['rank']}: Confirmed {c['stage_name']} (Similarity: {c['cosine_similarity']:.3f})\")\n    print(\"=\" * 80 + \"\\n\")\n\n\nsample_fp = val_df.iloc[0][\"filepath\"]  # CHANGED: use validation image for qualitative demo\nprint(\"--> Running validation example: standard pipeline\")  # CHANGED: identify validation use\nres_standard = run_pipeline(sample_fp)\ndisplay_pipeline_result(res_standard, \"Validation Example (Standard Threshold)\")  # CHANGED: label source correctly\n\nprint(\"--> Running validation example: safety-gate demonstration\")  # CHANGED: identify validation use\nstrict_gov = GovernanceAgent(confidence_threshold=0.999)  # intentionally high to trigger gate\nres_gated = run_pipeline(sample_fp, gov_agent=strict_gov)\ndisplay_pipeline_result(res_gated, \"Validation Example (High Safety Threshold)\")  # CHANGED: label source correctly","metadata":{"id":"86e1cb4938ec47f0b6ab24103f787498","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:39.857344Z","iopub.execute_input":"2026-09-30T20:42:39.857624Z","iopub.status.idle":"2026-09-30T20:42:44.515819Z","shell.execute_reply.started":"2026-09-30T20:42:39.857602Z","shell.execute_reply":"2026-09-30T20:42:44.515278Z"}},"outputs":[],"execution_count":null},{"id":"fa537e1e6e6143e4aae5ed27ad9a3e2e","cell_type":"code","source":"# ============================================================\n# SECTION 10.6 — SECTION SUMMARY\n# ============================================================\n\nprint(\"=\" * 60)\nprint(\"  SECTION 10 COMPLETE — Innovation Feature B\")\nprint(\"=\" * 60)\nprint(\"  DiagnosisAgent     : CNN inference & 5-class stage/confidence outputs\")\nprint(\"  ExplainabilityAgent: Grad-CAM heatmap + quadrant text + case retrieval\")\nprint(\"  AdvisoryAgent      : AAO clinical protocol mapping + disclaimer\")\nprint(\"  GovernanceAgent    : active safety gate (<70% confidence overrides output)\")\nprint(\"  run_pipeline()     : unified multi-agent decision support function\")\nprint(\"  Next -> Section 11: UI with Gradio & Bonus Features C, D, E\")\nprint(\"=\" * 60)","metadata":{"id":"fa537e1e6e6143e4aae5ed27ad9a3e2e","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:42:51.165246Z","iopub.execute_input":"2026-09-30T20:42:51.16606Z","iopub.status.idle":"2026-09-30T20:42:51.171559Z","shell.execute_reply.started":"2026-09-30T20:42:51.166033Z","shell.execute_reply":"2026-09-30T20:42:51.170608Z"}},"outputs":[],"execution_count":null},{"id":"eb2735033e444aba85324d5bdf475b3e","cell_type":"markdown","source":"## Section 11: Optional Gradio Research Interface\n\nThe notebook builds an optional Gradio interface for image upload, five-stage probabilities, Grad-CAM visualization, advisory text, and a similar-case gallery. The gallery displays real training reference files with dataset labels when the paths are available. This is a coursework prototype, not a deployed or clinically validated service; launching with a public share URL is optional and is not needed for the assessment.\n\nThe interface demonstrates how the analysis components can be presented together. Its confidence threshold is a configurable softmax cutoff, not calibrated clinical uncertainty, and neither the interface nor its outputs should be used for diagnosis or treatment decisions.","metadata":{"id":"eb2735033e444aba85324d5bdf475b3e","language":"markdown"}},{"id":"4ea27daa76cf4b058e6b7f7eede7d1d6","cell_type":"code","source":"# ============================================================\n# SECTION 11.1 — GRADIO INFERENCE WRAPPER FUNCTION\n# Formats outputs from the optional research pipeline for Gradio.\n# ============================================================\n\n\ndef gradio_predict(\n    input_img: np.ndarray,\n    confidence_threshold: float = 0.70,\n) -> Tuple[str, str, np.ndarray, str, List[Tuple[np.ndarray, str]]]:\n    \"\"\"Format prototype outputs for the Gradio interface.\"\"\"\n    if input_img is None:\n        empty_img = np.zeros((Config.IMG_SIZE, Config.IMG_SIZE, 3), dtype=np.uint8)\n        return (\n            \"### Please upload a retinal fundus photograph.\",\n            \"\",\n            empty_img,\n            \"\",\n            [],\n        )\n\n    gov_agent = GovernanceAgent(confidence_threshold=confidence_threshold)\n    preprocessed_input = preprocess_image(input_img, img_size=Config.IMG_SIZE)\n    result = run_pipeline(preprocessed_input, gov_agent=gov_agent)\n    diagnosis = result[\"diagnosis\"]\n    advisory = result[\"advisory\"]\n    explanation = result[\"explanation\"]\n\n    if result[\"flagged_for_review\"]:\n        status_md = (\n            \"### FLAGGED FOR HUMAN REVIEW\\n\"\n            f\"Model softmax confidence ({diagnosis['confidence']:.1%}) is below the configured \"\n            f\"display threshold ({confidence_threshold:.0%}). This threshold is not calibrated clinical uncertainty.\"\n        )\n    else:\n        status_md = (\n            \"### Research output only\\n\"\n            f\"Model softmax confidence: {diagnosis['confidence']:.1%}. A threshold pass does not establish safety.\"\n        )\n\n    probability_lines = [\n        f\"- **{class_name}:** {probability:.1%}\"\n        for class_name, probability in diagnosis[\"probabilities\"].items()\n    ]\n    diagnosis_md = (\n        f\"## Predicted stage: {diagnosis['stage_name']} (ICDR {diagnosis['predicted_stage']})\\n\\n\"\n        \"**Five-stage probability distribution:**\\n\"\n        + \"\\n\".join(probability_lines)\n    )\n    similar_gallery = []\n    for case in explanation[\"similar_cases\"]:\n        if not os.path.isfile(case[\"filepath\"]):\n            print(f\"[Gradio] Skipping unavailable reference image: {case['filepath']}\")\n            continue\n        case_img = preprocess_image(case[\"filepath\"], img_size=Config.IMG_SIZE)\n        case_img_uint8 = (np.clip(case_img, 0.0, 1.0) * 255).astype(np.uint8)\n        caption = (\n            f\"Rank #{case['rank']} | Dataset label: {case['stage_name']} \"\n            f\"(Stage {case['label']})\\nCosine similarity: {case['cosine_similarity']:.3f}\"\n        )\n        similar_gallery.append((case_img_uint8, caption))\n\n    advisory_md = (\n        f\"**Prototype advisory text:** {advisory['action_plan']}\\n\\n\"\n        f\"**Follow-up text:** {advisory['follow_up']}\\n\\n\"\n        f\"*{advisory['disclaimer']}*\"\n    )\n    quadrant_text = (\n        \"**Grad-CAM saliency (not verified lesion localization):**\\n\"\n        f\"{explanation['quadrant_explanation']}\"\n    )\n    return (\n        status_md,\n        diagnosis_md,\n        explanation[\"overlay_image\"],\n        f\"{quadrant_text}\\n\\n{advisory_md}\",\n        similar_gallery,\n    )\n\n\nprint(\"[Gradio] Optional prediction wrapper defined; gallery entries use existing reference image files.\")","metadata":{"id":"4ea27daa76cf4b058e6b7f7eede7d1d6","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:43:23.501387Z","iopub.execute_input":"2026-09-30T20:43:23.502225Z","iopub.status.idle":"2026-09-30T20:43:23.513508Z","shell.execute_reply.started":"2026-09-30T20:43:23.502193Z","shell.execute_reply":"2026-09-30T20:43:23.5128Z"}},"outputs":[],"execution_count":null},{"id":"a7900dd7c230453e9132445e18dc233d","cell_type":"code","source":"# ============================================================\n# SECTION 11.2 — BUILD & LAUNCH THE GRADIO APP\n#\n# Constructs an ergonomic clinical decision support dashboard\n# using Gradio Blocks and launches with share=True.\n# ============================================================\n\nimport gradio as gr\n\n\ndef build_gradio_app() -> gr.Blocks:\n    \"\"\"\n    Construct the Gradio clinical decision-support interface.\n    \"\"\"\n    custom_theme = gr.themes.Soft(\n        primary_hue=\"teal\",\n        secondary_hue=\"blue\",\n    )\n\n    with gr.Blocks(theme=custom_theme, title=\"Diabetic Retinopathy Clinical AI\") as demo:\n        gr.Markdown(\n            \"# 👁️ Diabetic Retinopathy Multi-Agent Decision Support System\\n\"\n            \"### *Deep Learning (EfficientNetB3) + Grad-CAM + Case-Based Reasoning + Safety Governance*\\n\"\n            \"Upload a fundus photograph to initiate automated multi-agent staging and explainability analysis.\"\n        )\n\n        with gr.Row():\n            # Left Column: Inputs & Controls\n            with gr.Column(scale=4):\n                input_image = gr.Image(\n                    label=\"Upload Fundus Photograph\",\n                    type=\"numpy\",\n                    sources=[\"upload\", \"clipboard\"],\n                )\n                threshold_slider = gr.Slider(\n                    minimum=0.50,\n                    maximum=0.95,\n                    value=0.70,\n                    step=0.05,\n                    label=\"Governance Safety Threshold (Default: 70%)\",\n                    info=\"Predictions below this confidence are gated and flagged for manual human review.\",\n                )\n                analyze_btn = gr.Button(\"🔍 Run Multi-Agent Analysis\", variant=\"primary\", size=\"lg\")\n\n                # Validation examples keep the held-out test images out of the demo.  # CHANGED: preserve test holdout\n                sample_examples = []\n                for cls_id in range(Config.NUM_CLASSES):\n                    sample_subset = val_df[val_df[\"label\"] == cls_id]  # CHANGED: use validation split\n                    if not sample_subset.empty:\n                        sample_examples.append(sample_subset.iloc[0][\"filepath\"])\n\n                if sample_examples:\n                    gr.Examples(\n                        examples=sample_examples[:5],\n                        inputs=input_image,\n                        label=\"Validation Sample Cases (Stages 0 to 4)\",  # CHANGED: label source\n                    )\n\n            # Right Column: Diagnostic & Governance Outputs\n            with gr.Column(scale=6):\n                status_box = gr.Markdown(\"### Upload an image and click 'Run Multi-Agent Analysis'\")\n                diagnosis_box = gr.Markdown()\n\n        gr.Markdown(\"---\")\n        gr.Markdown(\"## 🔬 Multi-Agent Clinical Explainability & Evidence Dossier\")\n\n        with gr.Row():\n            with gr.Column(scale=5):\n                gr.Markdown(\"### **1. Spatial Explainability: Grad-CAM Lesion Heatmap**\")\n                overlay_output = gr.Image(label=\"Superimposed Retinal Lesion Heatmap\", type=\"numpy\")\n\n            with gr.Column(scale=5):\n                gr.Markdown(\"### **2. Comparative Evidence: Similar Verified Historical Cases**\")\n                gallery_output = gr.Gallery(\n                    label=\"Top-3 Nearest Verified Training Cases (CBR)\",\n                    columns=3,\n                    rows=1,\n                    height=280,\n                    object_fit=\"contain\",\n                )\n\n        gr.Markdown(\"---\")\n        with gr.Row():\n            advisory_box = gr.Markdown()\n\n        analyze_btn.click(\n            fn=gradio_predict,\n            inputs=[input_image, threshold_slider],\n            outputs=[status_box, diagnosis_box, overlay_output, advisory_box, gallery_output],\n        )\n\n    return demo\n\n\ngradio_app = build_gradio_app()\nprint(\"[Gradio] App constructed successfully.\")\nprint(\"[Launch Note] Call 'gradio_app.launch(share=True)' to launch the interface with a public video-demo link.\")","metadata":{"id":"a7900dd7c230453e9132445e18dc233d","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:43:29.342796Z","iopub.execute_input":"2026-09-30T20:43:29.343076Z","iopub.status.idle":"2026-09-30T20:43:29.720496Z","shell.execute_reply.started":"2026-09-30T20:43:29.343052Z","shell.execute_reply":"2026-09-30T20:43:29.719595Z"}},"outputs":[],"execution_count":null},{"id":"c8e184b1c84c4b72b96558130f2718a1","cell_type":"code","source":"# ============================================================\n# SECTION 11.3 — OPTIONAL GRADIO APP LAUNCH\n# This cell does not launch the server during Run All.\n# Run one of the launch commands below manually when needed.\n# ============================================================\n\nprint(\"[Gradio] App is ready. Launch it manually if you want to use the UI.\")\n\n# Local-only access from this notebook environment:\n# gradio_app.launch(share=False, debug=False, show_error=True)\n\n# Optional: creates a publicly accessible Gradio URL. Use only when needed.\n# gradio_app.launch(share=True, debug=False, show_error=True)","metadata":{"id":"c8e184b1c84c4b72b96558130f2718a1","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:43:34.980643Z","iopub.execute_input":"2026-09-30T20:43:34.980977Z","iopub.status.idle":"2026-09-30T20:43:34.985852Z","shell.execute_reply.started":"2026-09-30T20:43:34.980951Z","shell.execute_reply":"2026-09-30T20:43:34.985081Z"}},"outputs":[],"execution_count":null},{"id":"6734c5aba7964a99a91844d23c7d006c","cell_type":"code","source":"# ============================================================\n# SECTION 11.4 — BONUS FEATURE E: MULTI-STAGE VERIFICATION\n#\n# Code audit cell confirming that the model operates as a true\n# 5-stage classifier (ICDR Stages 0-4) rather than a collapsed\n# binary model (DR present vs absent).\n# ============================================================\n\ndef verify_multistage_classification_compliance(\n    m: keras.Model,\n    class_names: List[str] = Config.CLASS_NAMES,\n) -> bool:\n    \"\"\"\n    Formally verify and print multi-stage classification compliance.\n\n    Coursework Criterion E:\n    Explicitly confirm the final model outputs all 5 stages (not collapsed to DR/no-DR)\n    and note that this satisfies stage-level (not just presence/absence) detection.\n    \"\"\"\n    output_shape = m.output_shape\n    n_outputs = output_shape[-1]\n\n    print(\"\\n\" + \"=\" * 75)\n    print(\"  BONUS FEATURE E COMPLIANCE: MULTI-STAGE CLASSIFICATION AUDIT\")\n    print(\"=\" * 75)\n    print(f\"  Model Name                : {m.name}\")\n    print(f\"  Final Softmax Output Shape: {output_shape} (Dimension: {n_outputs})\")\n    print(f\"  Configured Class Names    : {class_names}\")\n    print(f\"  Number of Discrete Stages : {len(class_names)}\")\n    print(\"-\" * 75)\n\n    is_compliant = (n_outputs == 5) and (len(class_names) == 5)\n\n    if is_compliant:\n        print(\"  [CONFIRMED] Multi-Stage Classification Requirement Fulfilled:\")\n        print(\"  - The model outputs distinct probabilities for all 5 international ICDR stages:\")\n        print(\"      Stage 0: No DR\")\n        print(\"      Stage 1: Mild NPDR\")\n        print(\"      Stage 2: Moderate NPDR\")\n        print(\"      Stage 3: Severe NPDR\")\n        print(\"      Stage 4: Proliferative DR\")\n        print(\"  - The architecture does NOT collapse outputs to binary (DR present / absent).\")\n        print(\"  - It satisfies granular, disease-severity clinical grading standards.\")\n    else:\n        print(\"  [ERROR] Model output dimension does not match 5 classes!\")\n\n    print(\"=\" * 75 + \"\\n\")\n    return is_compliant\n\n\nmultistage_compliant = verify_multistage_classification_compliance(model)","metadata":{"id":"6734c5aba7964a99a91844d23c7d006c","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:43:46.987477Z","iopub.execute_input":"2026-09-30T20:43:46.988252Z","iopub.status.idle":"2026-09-30T20:43:46.995874Z","shell.execute_reply.started":"2026-09-30T20:43:46.988225Z","shell.execute_reply":"2026-09-30T20:43:46.994832Z"}},"outputs":[],"execution_count":null},{"id":"7afa68525fda438690e106eb9f037070","cell_type":"code","source":"# ============================================================\n# SECTION 11.5 — END-TO-END NOTEBOOK COMPLETION SUMMARY\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"  END-TO-END PIPELINE COMPLETE: DIABETIC RETINOPATHY STAGE DETECTION\")\nprint(\"=\" * 70)\nprint(\"  1.  Setup & Data Acquisition      : APTOS 2019 (3,662 labelled images)\")\nprint(\"  2.  Preprocessing                 : Crop + Ben Graham (CLAHE/denoise/edges in gallery + ablation)\")\nprint(\"  3.  Augmentation & Balancing      : Train-only augmentation; class weights tested in a pilot\")\nprint(\"  4.  Train/Validation/Test         : Duplicate-aware stratified 70% / 15% / 15% split\")\nprint(f\"  5.  Model Architecture            : EfficientNetB3 + Dropout({Config.DROPOUT_RATE})\")\nprint(f\"  6.  Two-Phase Training            : top {Config.UNFREEZE_TOP_N} layers, Phase 2 LR={Config.PHASE2_LR}\")\nprint(\"      Pilot screen                  : hyperparameters, ablations, ResNet50/MobileNetV2 comparison\")\nprint(\"  7.  Evaluation                    : Accuracy, QWK, P/R/F1, confusion, ROC, DR/referable screening\")\nprint(\"  8.  Explainability                : Grad-CAM heatmaps + anatomical quadrant text\")\nprint(\"  9.  Innovation Feature A          : Penultimate embedding similar-case retrieval\")\nprint(\"  10. Innovation Feature B          : 4-Agent Clinical Decision Pipeline + Safety Gate\")\nprint(\"  11. Bonus Features C, D, E        : Gradio Web UI + HF Spaces + 5-Stage Multi-class\")\nprint(\"  12. Exploratory U-Net             : Pseudo-mask lesion localisation demo\")\nprint(\"  13. Research evidence             : Error analysis + reproducibility manifest\")\nprint(\"=\" * 70)","metadata":{"id":"7afa68525fda438690e106eb9f037070","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:44:06.062235Z","iopub.execute_input":"2026-09-30T20:44:06.062542Z","iopub.status.idle":"2026-09-30T20:44:06.068956Z","shell.execute_reply.started":"2026-09-30T20:44:06.062519Z","shell.execute_reply":"2026-09-30T20:44:06.068114Z"}},"outputs":[],"execution_count":null},{"id":"6dc71dc9ce154458b646d0d5adaa866f","cell_type":"markdown","source":"## Section 12: Exploratory Innovation — U-Net Pseudo-Mask Demonstration\n\n**Purpose:** Explore whether an encoder-decoder U-Net can produce spatially detailed candidate regions alongside five-stage grading and Grad-CAM. The input fundus photographs are real dataset images, but the segmentation targets are synthetic heuristic pseudo-masks, not expert annotations.\n\n### Rationale and Method\n\nSmall lesions can be difficult to inspect after image downsampling, while a U-Net's encoder-decoder structure and skip connections preserve spatial detail. This motivates the experiment, but does not establish that its output is a lesion mask. The notebook synthesizes targets from green-channel top-hat/black-hat morphology gated by a Grad-CAM threshold, then trains an auxiliary U-Net with a hybrid binary-cross-entropy and soft-Dice loss.\n\n$$\\mathcal{L}_{\\text{hybrid}} = 0.5\\,\\mathcal{L}_{\\text{BCE}} + 0.5\\left(1 - \\frac{2 \\sum_i y_i \\hat{y}_i + \\epsilon}{\\sum_i y_i + \\sum_i \\hat{y}_i + \\epsilon}\\right)$$\n\nThe generated target masks and U-Net overlays are **synthetic pseudo-mask demos**. Any Dice value measures agreement with those heuristic targets only. There are no expert pixel labels or independent segmentation results here, so the overlay is not validated lesion localization and must not be interpreted as clinical evidence. The visualization uses real validation fundus images and should be captioned as a pseudo-mask demonstration.","metadata":{"id":"6dc71dc9ce154458b646d0d5adaa866f","language":"markdown"}},{"id":"0eed7197fce34712af3a08b784ded04d","cell_type":"code","source":"# ============================================================\n# SECTION 12.1 — SOFT DICE LOSS & METRICS\n#\n# Implements differentiable Soft Dice Loss and Dice Coefficient\n# for extreme foreground-background imbalance in lesion masks.\n# ============================================================\n\ndef dice_coefficient(y_true: tf.Tensor, y_pred: tf.Tensor, smooth: float = 1e-6) -> tf.Tensor:\n    \"\"\"\n    Compute the Sorensen-Dice coefficient for binary segmentation masks.\n\n    Dice = (2 * |X ∩ Y| + smooth) / (|X| + |Y| + smooth)\n\n    Args:\n        y_true: Ground truth binary tensor of shape (B, H, W, 1).\n        y_pred: Predicted probability tensor of shape (B, H, W, 1) in [0, 1].\n        smooth: Smoothing epsilon to prevent division by zero.\n\n    Returns:\n        Scalar Dice coefficient tensor in [0, 1].\n    \"\"\"\n    y_true_f = tf.cast(tf.reshape(y_true, [-1]), tf.float32)\n    y_pred_f = tf.cast(tf.reshape(y_pred, [-1]), tf.float32)\n    intersection = tf.reduce_sum(y_true_f * y_pred_f)\n    cardinality  = tf.reduce_sum(y_true_f) + tf.reduce_sum(y_pred_f)\n    return (2.0 * intersection + smooth) / (cardinality + smooth)\n\n\ndef dice_loss(y_true: tf.Tensor, y_pred: tf.Tensor) -> tf.Tensor:\n    \"\"\"\n    Differentiable Soft Dice Loss: 1.0 - Dice Coefficient.\n    \"\"\"\n    return 1.0 - dice_coefficient(y_true, y_pred)\n\n\ndef hybrid_dice_bce_loss(y_true: tf.Tensor, y_pred: tf.Tensor) -> tf.Tensor:\n    \"\"\"\n    Hybrid loss combining Binary Cross-Entropy (pixel-level fidelity)\n    with Soft Dice Loss (macro spatial overlap).\n\n    Loss = 0.5 * BCE + 0.5 * Dice_Loss\n    \"\"\"\n    bce = keras.losses.binary_crossentropy(y_true, y_pred)\n    dice = dice_loss(y_true, y_pred)\n    return 0.5 * tf.reduce_mean(bce) + 0.5 * dice\n\n\nprint(\"[U-Net] Custom Soft Dice Loss and Hybrid Loss compiled.\")","metadata":{"id":"0eed7197fce34712af3a08b784ded04d","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:44:12.272034Z","iopub.execute_input":"2026-09-30T20:44:12.272791Z","iopub.status.idle":"2026-09-30T20:44:12.280165Z","shell.execute_reply.started":"2026-09-30T20:44:12.272754Z","shell.execute_reply":"2026-09-30T20:44:12.279247Z"}},"outputs":[],"execution_count":null},{"id":"edb42fe2-4f1d-432f-a9b7-b2b7d7457e10","cell_type":"code","source":"   # U-Net demo runs in float32: the classifier's mixed_float16 policy breaks XLA in the U-Net gradients.\n   keras.mixed_precision.set_global_policy(\"float32\")\n   print(\"[U-Net] Precision policy:\", keras.mixed_precision.global_policy().name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:49:49.570157Z","iopub.execute_input":"2026-09-30T20:49:49.570469Z","iopub.status.idle":"2026-09-30T20:49:49.575475Z","shell.execute_reply.started":"2026-09-30T20:49:49.570444Z","shell.execute_reply":"2026-09-30T20:49:49.574574Z"}},"outputs":[],"execution_count":null},{"id":"22a4d519577c473aa5dde342f968cd23","cell_type":"code","source":"# ============================================================\n# SECTION 12.2 — AUXILIARY U-NET ARCHITECTURE\n# Encoder-decoder network with skip connections, sized from Config.\n# ============================================================\n\ndef build_unet(\n    input_shape: Tuple[int, int, int] = (Config.IMG_SIZE, Config.IMG_SIZE, 3),\n) -> keras.Model:\n    \"\"\"Construct an encoder-decoder network for auxiliary lesion masks.\"\"\"\n    inputs = keras.Input(shape=input_shape, name=\"unet_input\")\n\n    def conv_block(x, filters, name_prefix):\n        x = layers.Conv2D(filters, (3, 3), padding=\"same\", name=f\"{name_prefix}_conv1\")(x)\n        x = layers.BatchNormalization(name=f\"{name_prefix}_bn1\")(x)\n        x = layers.Activation(\"relu\", name=f\"{name_prefix}_relu1\")(x)\n        x = layers.Conv2D(filters, (3, 3), padding=\"same\", name=f\"{name_prefix}_conv2\")(x)\n        x = layers.BatchNormalization(name=f\"{name_prefix}_bn2\")(x)\n        x = layers.Activation(\"relu\", name=f\"{name_prefix}_relu2\")(x)\n        return x\n\n    c1 = conv_block(inputs, 32, \"enc1\")\n    p1 = layers.MaxPooling2D((2, 2), name=\"pool1\")(c1)\n    c2 = conv_block(p1, 64, \"enc2\")\n    p2 = layers.MaxPooling2D((2, 2), name=\"pool2\")(c2)\n    c3 = conv_block(p2, 128, \"enc3\")\n    p3 = layers.MaxPooling2D((2, 2), name=\"pool3\")(c3)\n    bottleneck = conv_block(p3, 256, \"bottleneck\")\n\n    up3 = layers.Conv2DTranspose(128, (2, 2), strides=(2, 2), padding=\"same\", name=\"up3\")(bottleneck)\n    up3 = layers.Resizing(c3.shape[1], c3.shape[2], name=\"align_up3\")(up3)\n    dec3 = conv_block(layers.concatenate([up3, c3], name=\"concat3\"), 128, \"dec3\")\n\n    up2 = layers.Conv2DTranspose(64, (2, 2), strides=(2, 2), padding=\"same\", name=\"up2\")(dec3)\n    up2 = layers.Resizing(c2.shape[1], c2.shape[2], name=\"align_up2\")(up2)\n    dec2 = conv_block(layers.concatenate([up2, c2], name=\"concat2\"), 64, \"dec2\")\n\n    up1 = layers.Conv2DTranspose(32, (2, 2), strides=(2, 2), padding=\"same\", name=\"up1\")(dec2)\n    up1 = layers.Resizing(c1.shape[1], c1.shape[2], name=\"align_up1\")(up1)\n    dec1 = conv_block(layers.concatenate([up1, c1], name=\"concat1\"), 32, \"dec1\")\n\n    outputs = layers.Conv2D(1, (1, 1), activation=\"sigmoid\", name=\"lesion_mask_output\")(dec1)\n    return keras.Model(inputs=inputs, outputs=outputs, name=\"Retinal_Lesion_UNet\")\n\n\nunet_model = build_unet()\nunet_model.compile(\n    optimizer=keras.optimizers.Adam(learning_rate=1e-4),\n    loss=hybrid_dice_bce_loss,\n    metrics=[dice_coefficient, \"binary_accuracy\"],\n)\nprint(\"[U-Net] Retinal Lesion U-Net constructed successfully.\")\nprint(f\"  Input resolution: {Config.IMG_SIZE}x{Config.IMG_SIZE}\")\nprint(f\"  Total parameters: {unet_model.count_params():,}\")","metadata":{"id":"22a4d519577c473aa5dde342f968cd23","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:50:01.339943Z","iopub.execute_input":"2026-09-30T20:50:01.34022Z","iopub.status.idle":"2026-09-30T20:50:01.531868Z","shell.execute_reply.started":"2026-09-30T20:50:01.3402Z","shell.execute_reply":"2026-09-30T20:50:01.530931Z"}},"outputs":[],"execution_count":null},{"id":"c59c68e25bf840c4945709cf6966b4fb","cell_type":"code","source":"# ============================================================\n# SECTION 12.3 — SEMI-SUPERVISED LESION MASK SYNTHESIZER\n#\n# Generates heuristic pseudo-masks from Grad-CAM and morphology.\n# These are synthesized masks, not manually annotated ground truth.\n# ============================================================\n\ndef synthesize_retinal_lesion_mask(\n    preprocessed_img: np.ndarray,\n    stage: int,\n    cam_model: keras.Model = gradcam_model,\n) -> np.ndarray:\n    \"\"\"Create a heuristic pseudo-mask with shape (IMG_SIZE, IMG_SIZE, 1).\"\"\"\n    if stage == 0:\n        return np.zeros((Config.IMG_SIZE, Config.IMG_SIZE, 1), dtype=np.float32)\n\n    heatmap, _, _ = compute_gradcam_heatmap(\n        preprocessed_img, cam_model=cam_model, pred_index=stage\n    )\n    green = (preprocessed_img[:, :, 1] * 255).astype(np.uint8)\n\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5))\n    tophat = cv2.morphologyEx(green, cv2.MORPH_TOPHAT, kernel)\n    blackhat = cv2.morphologyEx(green, cv2.MORPH_BLACKHAT, kernel)\n    lesion_signals = cv2.add(tophat, blackhat)\n\n    _, binary_lesions = cv2.threshold(lesion_signals, 25, 255, cv2.THRESH_BINARY)\n    heatmap_full = cv2.resize(\n        heatmap.astype(np.float32),\n        (Config.IMG_SIZE, Config.IMG_SIZE),\n        interpolation=cv2.INTER_LINEAR,\n    )\n    saliency_gate = (heatmap_full > 0.35).astype(np.uint8)\n    gated_mask = cv2.bitwise_and(binary_lesions, binary_lesions, mask=saliency_gate)\n\n    cleaned_mask = cv2.morphologyEx(gated_mask, cv2.MORPH_OPEN, kernel)\n    return (cleaned_mask > 0).astype(np.float32)[..., np.newaxis]\n\n\nprint(\"[U-Net] Lesion pseudo-mask generator ready; masks are not clinical annotations.\")","metadata":{"id":"c59c68e25bf840c4945709cf6966b4fb","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:50:07.250615Z","iopub.execute_input":"2026-09-30T20:50:07.251028Z","iopub.status.idle":"2026-09-30T20:50:07.260767Z","shell.execute_reply.started":"2026-09-30T20:50:07.251Z","shell.execute_reply":"2026-09-30T20:50:07.259884Z"}},"outputs":[],"execution_count":null},{"id":"8c32d74f07e64ca58e7f248849348de3","cell_type":"code","source":"# ============================================================\n# SECTION 12.4 — TRAIN AUXILIARY U-NET ON SYNTHETIC PSEUDO-MASKS\n# Targets are generated from image morphology and Grad-CAM, not annotations.\n# ============================================================\n\n\ndef train_auxiliary_unet(\n    unet: keras.Model,\n    dataframe: pd.DataFrame,\n    n_samples: int = 400,\n    epochs: int = 5,\n    batch_size: int = 16,\n) -> keras.callbacks.History:\n    \"\"\"Train a demonstration U-Net against heuristic pseudo-mask targets.\"\"\"\n    sample_df = dataframe.sample(min(n_samples, len(dataframe)), random_state=Config.SEED)\n    images = []\n    masks = []\n\n    for _, row in tqdm(sample_df.iterrows(), total=len(sample_df), desc=\"Creating synthetic pseudo-mask demo targets\"):\n        image = preprocess_image(row[\"filepath\"])\n        mask = synthesize_retinal_lesion_mask(image, int(row[\"label\"]))\n        images.append(image)\n        masks.append(mask)\n\n    x_train = np.asarray(images, dtype=np.float32)\n    y_pseudo = np.asarray(masks, dtype=np.float32)\n    print(\n        f\"[U-Net DEMO] Training on {len(x_train)} real images with synthetic pseudo-masks; \"\n        f\"image shape={x_train.shape}, target shape={y_pseudo.shape}.\"\n    )\n\n    history = unet.fit(\n        x_train,\n        y_pseudo,\n        batch_size=batch_size,\n        epochs=epochs,\n        validation_split=0.2,\n        verbose=1,\n    )\n    print(\n        \"[U-Net DEMO] Pseudo-mask validation Dice is target self-consistency only, \"\n        f\"not clinical lesion accuracy. Final Dice: {history.history['dice_coefficient'][-1]:.4f}\"\n    )\n    return history\n\n\nunet_history = train_auxiliary_unet(unet_model, train_df, n_samples=300, epochs=5)\nunet_model.save_weights(os.path.join(Config.CHECKPOINT_DIR, \"unet_pseudomask.weights.h5\"))\nprint(\"[U-Net DEMO] Saved pseudo-mask demonstration checkpoint; no expert lesion ground truth was used.\")","metadata":{"id":"8c32d74f07e64ca58e7f248849348de3","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:50:12.367996Z","iopub.execute_input":"2026-09-30T20:50:12.36825Z","iopub.status.idle":"2026-09-30T20:55:34.527205Z","shell.execute_reply.started":"2026-09-30T20:50:12.368231Z","shell.execute_reply":"2026-09-30T20:55:34.526238Z"}},"outputs":[],"execution_count":null},{"id":"15425e647b8744a18bcadad596cad8af","cell_type":"code","source":"# ============================================================\n# SECTION 12.5 — PSEUDO-MASK TARGET VS U-NET OUTPUT DEMO\n# Real validation photos; synthetic heuristic targets and unverified predictions.\n# ============================================================\n\n\ndef visualize_three_layer_explainability_stack(\n    validation_dataframe: pd.DataFrame,\n    cls_model: keras.Model = model,\n    seg_model: keras.Model = unet_model,\n    save_path: Optional[str] = None,\n) -> None:\n    \"\"\"Compare real validation images with synthetic targets and U-Net demo outputs.\"\"\"\n    selected_samples = []\n    for stage in [1, 2, 3, 4]:\n        subset = validation_dataframe[validation_dataframe[\"label\"] == stage]\n        if not subset.empty:\n            selected_samples.append(subset.iloc[0])\n\n    if not selected_samples:\n        raise ValueError(\"No validation examples are available for the pseudo-mask demonstration.\")\n\n    n_rows = len(selected_samples)\n    fig, axes = plt.subplots(n_rows, 6, figsize=(25, 4.5 * n_rows), squeeze=False)\n    col_headers = [\n        \"Real validation fundus\",\n        \"Synthetic heuristic target\",\n        \"Five-stage classifier output\",\n        \"Grad-CAM saliency\",\n        \"U-Net probability (unverified)\",\n        \"U-Net overlay demo\",\n    ]\n    for column, title in enumerate(col_headers):\n        axes[0, column].set_title(title, fontsize=10, fontweight=\"bold\", pad=12)\n\n    for row_index, row in enumerate(selected_samples):\n        image = preprocess_image(row[\"filepath\"])\n        true_stage = int(row[\"label\"])\n        probabilities = cls_model(image[np.newaxis, ...], training=False).numpy()[0]\n        predicted_stage = int(np.argmax(probabilities))\n        confidence = float(probabilities[predicted_stage])\n        heatmap, _, _ = compute_gradcam_heatmap(\n            image, cam_model=gradcam_model, pred_index=predicted_stage\n        )\n        cam_overlay = superimpose_gradcam(image, heatmap, alpha=0.45)\n        pseudo_target = synthesize_retinal_lesion_mask(image, true_stage)[..., 0]\n        predicted_mask = seg_model(image[np.newaxis, ...], training=False).numpy()[0, :, :, 0]\n        binary_mask = (predicted_mask > 0.4).astype(np.uint8)\n\n        base_image = (image * 255).astype(np.uint8)\n        pseudo_overlay = base_image.copy()\n        pseudo_overlay[binary_mask == 1] = [0, 255, 64]\n        pseudo_overlay = cv2.addWeighted(pseudo_overlay, 0.65, base_image, 0.35, 0)\n\n        axes[row_index, 0].imshow(image)\n        axes[row_index, 0].axis(\"off\")\n        axes[row_index, 0].set_ylabel(\n            f\"Dataset label: {Config.CLASS_NAMES[true_stage]}\\n(Stage {true_stage})\",\n            fontsize=9, rotation=0, labelpad=75, va=\"center\",\n        )\n        axes[row_index, 1].imshow(pseudo_target, cmap=\"gray\", vmin=0, vmax=1)\n        axes[row_index, 1].axis(\"off\")\n        bar_colors = [\"#4A90E2\" if i != predicted_stage else \"#D0021B\" for i in range(5)]\n        axes[row_index, 2].barh(\n            Config.CLASS_NAMES, probabilities, color=bar_colors, edgecolor=\"black\", linewidth=0.6\n        )\n        axes[row_index, 2].set_xlim(0, 1.0)\n        axes[row_index, 2].set_xlabel(\"Probability\")\n        axes[row_index, 2].text(\n            0.5, 0.85,\n            f\"Predicted: {Config.CLASS_NAMES[predicted_stage]}\\nConfidence: {confidence:.1%}\",\n            transform=axes[row_index, 2].transAxes,\n            bbox=dict(boxstyle=\"round,pad=0.3\", fc=\"#FFF9C4\", ec=\"#FBC02D\"),\n            fontsize=9, fontweight=\"bold\",\n        )\n        axes[row_index, 3].imshow(cam_overlay)\n        axes[row_index, 3].axis(\"off\")\n        axes[row_index, 4].imshow(predicted_mask, cmap=\"gray\", vmin=0, vmax=1)\n        axes[row_index, 4].axis(\"off\")\n        axes[row_index, 5].imshow(pseudo_overlay)\n        axes[row_index, 5].axis(\"off\")\n\n    fig.suptitle(\n        \"Exploratory Demo — Synthetic Pseudo-Targets vs U-Net Output (Not Clinical Segmentation)\",\n        fontsize=12,\n        fontweight=\"bold\",\n    )\n    plt.tight_layout()\n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        plt.savefig(save_path, dpi=150, bbox_inches=\"tight\")\n        print(f\"[Pseudo-mask DEMO] Saved target/output comparison -> {save_path}\")\n    plt.show()\n\n\nvisualize_three_layer_explainability_stack(\n    val_df,\n    save_path=os.path.join(Config.REPORTS_DIR, \"three_layer_explainability_stack.png\"),\n)","metadata":{"id":"15425e647b8744a18bcadad596cad8af","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:55:41.710506Z","iopub.execute_input":"2026-09-30T20:55:41.710901Z","iopub.status.idle":"2026-09-30T20:55:59.085156Z","shell.execute_reply.started":"2026-09-30T20:55:41.710872Z","shell.execute_reply":"2026-09-30T20:55:59.084205Z"}},"outputs":[],"execution_count":null},{"id":"162e81ccd0cc4dc5a077a541043541c6","cell_type":"code","source":"print(\"=\" * 70)\nprint(\"  SECTION 12 COMPLETE — Exploratory U-Net Pseudo-Mask Demonstration\")\nprint(\"=\" * 70)\nprint(\"  build_unet()                     : encoder-decoder with skip connections\")\nprint(\"  hybrid_dice_bce_loss()           : BCE + soft Dice against pseudo-mask targets\")\nprint(\"  synthesize_retinal_lesion_mask() : heuristic synthetic demo-target generator\")\nprint(\"  train_auxiliary_unet()           : exploratory self-consistency training only\")\nprint(f\"  Artifact generated               : {Config.REPORTS_DIR}/three_layer_explainability_stack.png\")\nprint(\"  The green overlay is not expert-annotated or clinically validated lesion segmentation.\")\nprint(\"=\" * 70)","metadata":{"id":"162e81ccd0cc4dc5a077a541043541c6","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:56:05.49187Z","iopub.execute_input":"2026-09-30T20:56:05.492134Z","iopub.status.idle":"2026-09-30T20:56:05.498043Z","shell.execute_reply.started":"2026-09-30T20:56:05.492114Z","shell.execute_reply":"2026-09-30T20:56:05.497172Z"}},"outputs":[],"execution_count":null},{"id":"8a0e8d1c","cell_type":"markdown","source":"## Section 13 — Research Evidence: Error Analysis and Reproducibility\n\nThis section is research-only and does not modify the deployed EfficientNetB3 checkpoint or the Gradio application. The model-level ablation and backbone comparison results are produced by the pilot runner in Section 6 (`comparison_trials_results.csv`, `pilot_comparison.png`). This section records the test-set error analysis (adjacent vs distant ordinal errors, high-confidence mistakes, low-confidence Governance cases) and a reproducibility manifest (library versions, hardware, seed, dataset, split files, hyperparameters and checkpoints).","metadata":{}},{"id":"c8f951552ce74fe4bbe11f446b80bd33","cell_type":"code","source":"import os\nimport sys\nfrom pathlib import Path\n\nPROJECT_ROOT = Path.cwd()\nif not (PROJECT_ROOT / \"core\").is_dir():\n    PROJECT_ROOT = Path(_PROJECT_ROOT)\nif str(PROJECT_ROOT) not in sys.path:\n    sys.path.insert(0, str(PROJECT_ROOT))\n\nfrom core.research_evidence import (\n    ABLATION_VARIANTS,\n    build_error_analysis,\n    build_reproducibility_manifest,\n    save_error_analysis,\n    set_reproducible_seed,\n)\n\nset_reproducible_seed(getattr(Config, \"SEED\", 42))\nREPORT_DIR = Path(getattr(Config, \"REPORTS_DIR\", \"report_images\"))\nREPORT_DIR.mkdir(parents=True, exist_ok=True)\n\nfor split_name, split_frame in ((\"train\", train_df), (\"validation\", val_df), (\"test\", test_df)):\n    split_frame.to_csv(REPORT_DIR / f\"{split_name}_split.csv\", index=False)\n\nerror_analysis = build_error_analysis(\n    y_test_true,\n    y_test_pred,\n    probabilities=y_test_prob,\n)\nsave_error_analysis(error_analysis, REPORT_DIR / \"error_analysis\")\n\nreproducibility_manifest = build_reproducibility_manifest(\n    config={\n        \"random_seed\": getattr(Config, \"SEED\", 42),\n        \"image_size\": [Config.IMG_SIZE, Config.IMG_SIZE],\n        \"batch_size\": Config.BATCH_SIZE,\n        \"phase_1_learning_rate\": Config.PHASE1_LR,\n        \"phase_2_learning_rate\": Config.PHASE2_LR,\n        \"phase_1_max_epochs\": Config.PHASE1_EPOCHS,\n        \"phase_2_max_epochs\": Config.PHASE2_EPOCHS,\n        \"phase_2_unfreeze_top_n\": Config.UNFREEZE_TOP_N,\n        \"dropout_rate\": Config.DROPOUT_RATE,\n        \"class_weights_enabled\": Config.USE_CLASS_WEIGHTS,\n        \"early_stopping_monitor\": \"val_loss\",\n        \"oversampling_enabled\": Config.USE_OVERSAMPLING,\n        \"tta_selected_on_validation\": USE_TTA_FOR_EVALUATION,\n        \"qwk_thresholds_selected_on_validation\": np.round(QWK_THRESHOLDS, 6).tolist(),\n        \"test_predictions\": \"ordinal_expected_grade_thresholds\",\n        \"confidence_threshold\": Config.CONFIDENCE_THRESHOLD,\n        \"model\": \"EfficientNetB3\",\n    },\n    output_path=REPORT_DIR / \"reproducibility_manifest.json\",\n    dataset_version=\"APTOS 2019 Blindness Detection - labelled train_images (Kaggle competition)\",\n    split_files=[\n        str(REPORT_DIR / \"train_split.csv\"),\n        str(REPORT_DIR / \"validation_split.csv\"),\n        str(REPORT_DIR / \"test_split.csv\"),\n    ],\n    checkpoint_files=[\n        str(Path(Config.CHECKPOINT_DIR) / \"best_phase1.weights.h5\"),\n        str(Path(Config.CHECKPOINT_DIR) / \"best_phase2.weights.h5\"),\n        str(Path(Config.CHECKPOINT_DIR) / \"unet_pseudomask.weights.h5\"),\n    ],\n    training_duration_seconds=None,\n)\n\nprint(\"[Research Evidence] Error analysis and reproducibility artifacts saved.\")\nprint(\"[Research Evidence] Ablation variants:\")\nfor variant in ABLATION_VARIANTS:\n    print(\" -\", variant.name)\nprint({key: value for key, value in error_analysis.items() if key != \"error_rows\"})","metadata":{"id":"c8f951552ce74fe4bbe11f446b80bd33","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:56:11.318853Z","iopub.execute_input":"2026-09-30T20:56:11.319255Z","iopub.status.idle":"2026-09-30T20:56:11.546544Z","shell.execute_reply.started":"2026-09-30T20:56:11.319229Z","shell.execute_reply":"2026-09-30T20:56:11.545937Z"}},"outputs":[],"execution_count":null},{"id":"d2d968bb","cell_type":"code","source":"run_archive = shutil.make_archive(str(Path(Config.REPORTS_DIR).parent), \"zip\", root_dir=Path(Config.REPORTS_DIR).parent)\nprint(\"Download the complete run archive from Kaggle Output:\", run_archive)\nprint(\"Run folder:\", Path(Config.REPORTS_DIR).parent)","metadata":{"id":"d2d968bb","language":"python","trusted":true,"execution":{"iopub.status.busy":"2026-09-30T20:56:21.818812Z","iopub.execute_input":"2026-09-30T20:56:21.819215Z","iopub.status.idle":"2026-09-30T20:56:37.933966Z","shell.execute_reply.started":"2026-09-30T20:56:21.81919Z","shell.execute_reply":"2026-09-30T20:56:37.932936Z"}},"outputs":[],"execution_count":null}]}