{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"479d8876-48cf-4036-9955-76bfb10bed0d","cell_type":"markdown","source":"# Diabetic Retinopathy Stage Detection - Training Notebook\n\nThis notebook builds a diabetic retinopathy (DR) screening model from retinal fundus photographs. It grades each image on the five-stage scale (No DR, Mild, Moderate, Severe, Proliferative), detects whether DR is present, checks whether the photo is good enough to grade, and reports how certain it is.\n\n","metadata":{}},{"id":"4f3553ad-b613-4eda-928b-11b2354eabce","cell_type":"markdown","source":"## 0. Setup\n\nInstall the extra packages, import libraries, create output folders, check the GPU, fix the random seed and define every setting in one place.\n\n### 0.1 Install packages\n`timm` provides the pretrained CNN; `onnx` and `onnxruntime` are used for the deployment export.","metadata":{}},{"id":"9d263eaa-c741-42ed-b72e-d0b3f903d40d","cell_type":"code","source":"!pip install -q timm onnx onnxruntime","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:04:48.145426Z","iopub.execute_input":"2026-09-21T08:04:48.146008Z","iopub.status.idle":"2026-09-21T08:04:54.81137Z","shell.execute_reply.started":"2026-09-21T08:04:48.145977Z","shell.execute_reply":"2026-09-21T08:04:54.810623Z"}},"outputs":[],"execution_count":null},{"id":"771ff462-97cd-4ea5-82ae-abbc4d8c1f8f","cell_type":"markdown","source":"### 0.2 Imports and output folders\nAll figures, tables and logs are written under `/kaggle/working/results` so they can be downloaded for the report.","metadata":{}},{"id":"3f5cb522-1483-4288-abd1-48518f5b2870","cell_type":"code","source":"import os, json, time, random, copy, zipfile, inspect, warnings\nfrom pathlib import Path\nfrom dataclasses import dataclass, asdict\nfrom concurrent.futures import ThreadPoolExecutor\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport timm\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings('ignore', category=UserWarning)\nwarnings.filterwarnings('ignore', category=FutureWarning)\n\nWORKING_DIR = Path('/kaggle/working')\nWORKING_DIR.mkdir(parents=True, exist_ok=True)\nos.chdir(WORKING_DIR)\nfor folder in ['results/figures', 'results/tables', 'results/logs', 'models/export']:\n    Path(folder).mkdir(parents=True, exist_ok=True)\n\n\ndef show_table(title, table):\n    # print a table as plain text\n    print(f'\\n=== {title} ===')\n    print(table.to_string())\n\n\ndef save_figure(path, dpi=200):\n\n    plt.tight_layout()\n    plt.savefig(path, dpi=dpi)\n    print(f'figure saved: {path}')\n    plt.show()\n\n\nprint('torch version:', torch.__version__)\nprint('timm version :', timm.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:04:54.813052Z","iopub.execute_input":"2026-09-21T08:04:54.813783Z","iopub.status.idle":"2026-09-21T08:05:08.085871Z","shell.execute_reply.started":"2026-09-21T08:04:54.813754Z","shell.execute_reply":"2026-09-21T08:05:08.084965Z"}},"outputs":[],"execution_count":null},{"id":"dee75187-d97a-46c3-8d6d-7e2585f1488a","cell_type":"markdown","source":"### 0.3 Check the GPU\nTraining needs a GPU accelerator (P100 or T4).","metadata":{}},{"id":"d516fe6c-03ce-4264-9af7-fbbf790ba79c","cell_type":"code","source":"if torch.cuda.is_available():\n    gpu = torch.cuda.get_device_properties(0)\n    print(f'GPU: {gpu.name} | memory: {gpu.total_memory / 1e9:.1f} GB')\nelse:\n    print('GPU: NONE - turn on the accelerator in the notebook settings')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:08.086943Z","iopub.execute_input":"2026-09-21T08:05:08.087409Z","iopub.status.idle":"2026-09-21T08:05:08.405416Z","shell.execute_reply.started":"2026-09-21T08:05:08.087384Z","shell.execute_reply":"2026-09-21T08:05:08.404568Z"}},"outputs":[],"execution_count":null},{"id":"ece4a15e-1689-4530-bc0f-678efe2adb02","cell_type":"markdown","source":"### 0.4 Fix the random seed\nThe same seed gives the same splits, sampling and initial weights on every run.","metadata":{}},{"id":"0c25e493-8fee-4aa6-9771-4ee49aa98634","cell_type":"code","source":"RANDOM_SEED = 42\n\ndef seed_everything(seed=RANDOM_SEED):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\nseed_everything()\nprint('random seed:', RANDOM_SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:08.407169Z","iopub.execute_input":"2026-09-21T08:05:08.407471Z","iopub.status.idle":"2026-09-21T08:05:08.417384Z","shell.execute_reply.started":"2026-09-21T08:05:08.407448Z","shell.execute_reply":"2026-09-21T08:05:08.416576Z"}},"outputs":[],"execution_count":null},{"id":"32d6d836-2ced-43d9-a886-95c73b12fcee","cell_type":"markdown","source":"### 0.5 Configuration\nEvery hyperparameter is stored in one `config` object. It is saved inside each checkpoint and with the final results, so every number can be traced back to the settings that produced it.","metadata":{}},{"id":"1ba0837d-33ef-4e19-ae19-9471f42c81c8","cell_type":"code","source":"@dataclass\nclass Config:\n    # data\n    image_size: int = 384              # input resolution in pixels\n    num_classes: int = 5               # DR stages 0-4\n    num_quality_classes: int = 2       # gradable / ungradable\n    val_frac: float = 0.15\n    test_frac: float = 0.15\n    seed: int = 42\n\n    # preprocessing switches\n    use_graham: bool = True            # illumination correction / edge enhancement\n    use_clahe: bool = True             # contrast enhancement\n\n    # model\n    backbone: str = 'tf_efficientnetv2_s.in21k_ft_in1k'\n    pretrained: bool = True            # start from ImageNet weights\n    dropout: float = 0.4\n\n    # optimisation\n    epochs: int = 25\n    batch_size: int = 32\n    head_lr: float = 1e-3              # learning rate for the new output layers\n    backbone_lr: float = 1e-4          # learning rate for the pretrained layers\n    weight_decay: float = 1e-5\n    warmup_steps: int = 300\n    grad_clip: float = 1.0\n    amp: bool = True                   # mixed precision\n    use_ema: bool = True               # moving average of weights\n    ema_decay: float = 0.999\n    freeze_first_epoch: bool = True    # train heads only in epoch 0\n    quality_loss_weight: float = 0.3\n    label_smoothing: float = 0.05\n\n    # class balancing\n    sampler_beta: float = 0.999\n\n    # early stopping\n    patience: int = 4\n    min_delta: float = 0.002\n\n    # uncertainty\n    mc_dropout_samples: int = 20\n\n    # triage thresholds\n    ungradable_threshold: float = 0.50\n    uncertainty_threshold: float = 0.45\n    confidence_threshold: float = 0.55\n\n    # paths and hardware\n    checkpoint_dir: str = '/kaggle/working/checkpoints'\n    cache_dir: str = '/kaggle/working/cache'\n    device: str = 'cuda'\n    num_workers: int = 2\n\n    def to_dict(self):\n        return asdict(self)\n\n\nconfig = Config()\nconfig.device = 'cuda' if torch.cuda.is_available() else 'cpu'\nconfig.amp = config.amp and config.device == 'cuda'\nconfig.num_workers = min(4, os.cpu_count() or 2)\n\nshow_table('Configuration', pd.Series(config.to_dict()).to_frame('value'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:08.418458Z","iopub.execute_input":"2026-09-21T08:05:08.418777Z","iopub.status.idle":"2026-09-21T08:05:08.444788Z","shell.execute_reply.started":"2026-09-21T08:05:08.418755Z","shell.execute_reply":"2026-09-21T08:05:08.444107Z"}},"outputs":[],"execution_count":null},{"id":"195e49d9-ecc1-4234-b5c6-b813c86cd7fa","cell_type":"markdown","source":"## 1. Dataset\n\n**DDR** (13,673 fundus images from 147 hospitals, graded 0-4 plus an \"ungradable\" class 5) is used for training, validation and internal testing. **APTOS 2019** (3,662 images from India, different cameras and population) is never trained on and is used only as an external test set.\n\nTasks: locate both datasets, build one table of all images, and analyse the class distribution, the ungradable images, the split composition and the image sizes.\n\n### 1.1 Find the attached datasets","metadata":{}},{"id":"1cbace07-90f6-4e60-9921-436809d34132","cell_type":"code","source":"!ls /kaggle/input/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:08.445714Z","iopub.execute_input":"2026-09-21T08:05:08.446044Z","iopub.status.idle":"2026-09-21T08:05:08.581612Z","shell.execute_reply.started":"2026-09-21T08:05:08.446016Z","shell.execute_reply":"2026-09-21T08:05:08.580883Z"}},"outputs":[],"execution_count":null},{"id":"a48aad08-3669-49af-8cc5-22651f2c5ec7","cell_type":"code","source":"def find_dataset_folder(is_match, root='/kaggle/input', max_depth=4):\n    root = Path(root)\n    if not root.exists():\n        return None\n    for folder_path, subfolders, files in os.walk(root):\n        folder = Path(folder_path)\n        depth = len(folder.relative_to(root).parts)\n        if depth >= max_depth:\n            subfolders[:] = []\n        if depth > 0 and is_match(folder, subfolders, files):\n            return str(folder)\n    return None\n\n\nDDR_FOLDER = find_dataset_folder(lambda folder, subfolders, files: 'ddr' in folder.name.lower())\nAPTOS_FOLDER = find_dataset_folder(lambda folder, subfolders, files:\n                                   'train.csv' in files and 'train_images' in subfolders)\n\n# set by hand if the wrong folder is picked\n# DDR_FOLDER = '/kaggle/input/ddrdataset'\n# APTOS_FOLDER = '/kaggle/input/aptos2019-blindness-detection'\n\nprint('DDR folder  :', DDR_FOLDER)\nprint('APTOS folder:', APTOS_FOLDER)\nassert DDR_FOLDER and APTOS_FOLDER, 'dataset not found - attach both datasets with Add Data'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:08.582768Z","iopub.execute_input":"2026-09-21T08:05:08.582985Z","iopub.status.idle":"2026-09-21T08:05:08.639687Z","shell.execute_reply.started":"2026-09-21T08:05:08.582962Z","shell.execute_reply":"2026-09-21T08:05:08.639033Z"}},"outputs":[],"execution_count":null},{"id":"12b06d23-47a7-4e40-a4d3-89c1c7fbbe29","cell_type":"markdown","source":"### 1.2 Inspect the DDR folder layout\nDifferent Kaggle copies of DDR store their labels differently, so the folder tree is printed before loading.","metadata":{}},{"id":"60a779d3-2a45-46d4-af6d-d696af540de0","cell_type":"code","source":"def print_folder_tree(root, max_depth=3):\n    root = Path(root)\n    if not root.exists():\n        print('PATH DOES NOT EXIST:', root)\n        return\n    for path in sorted(root.rglob('*'))[:400]:\n        depth = len(path.relative_to(root).parts)\n        if depth > max_depth:\n            continue\n        indent = '  ' * (depth - 1)\n        if path.is_dir():\n            print(f'{indent}{path.name}/  ({sum(1 for _ in path.iterdir())} entries)')\n        elif depth <= 2 or path.suffix in {'.csv', '.txt'}:\n            print(f'{indent}{path.name}')\n\n\nprint_folder_tree(DDR_FOLDER)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:08.640616Z","iopub.execute_input":"2026-09-21T08:05:08.640918Z","iopub.status.idle":"2026-09-21T08:05:32.961054Z","shell.execute_reply.started":"2026-09-21T08:05:08.640884Z","shell.execute_reply":"2026-09-21T08:05:32.960407Z"}},"outputs":[],"execution_count":null},{"id":"f0856942-6224-4dad-ba48-da76eabc4103","cell_type":"markdown","source":"### 1.3 DDR loader\nReads the grading labels (official `train/valid/test` lists, a single CSV, or one folder per class), matches every label to its image file and assigns train / validation / test splits. If the copy has no official split, a stratified 70 / 15 / 15 split is created.\n\nUngradable images (label 5) are **kept**: they get grade `-1` (left out of the grading loss) and quality `1`, so they can teach the quality head to reject bad photos.","metadata":{}},{"id":"a25791a4-f614-44a9-be97-150bd36847c6","cell_type":"code","source":"UNGRADABLE_LABEL = 5\nGRADE_NAMES = {0: 'No DR', 1: 'Mild NPDR', 2: 'Moderate NPDR',\n               3: 'Severe NPDR', 4: 'Proliferative DR'}\nLABEL_NAMES = {**GRADE_NAMES, 5: 'Ungradable'}\n\n\ndef read_label_file(path):\n    # txt: \"<file> <label>\" per line | csv: name and label columns\n    path = Path(path)\n    if path.suffix.lower() == '.csv':\n        table = pd.read_csv(path)\n        table.columns = [c.strip().lower() for c in table.columns]\n        name_column = next((c for c in table.columns if c in\n                            {'image', 'id_code', 'filename', 'name', 'image_name', 'img'}),\n                           table.columns[0])\n        label_column = next((c for c in table.columns if c in\n                             {'label', 'level', 'grade', 'diagnosis', 'dr_grade', 'class'}),\n                            table.columns[-1])\n        return pd.DataFrame({'filename': table[name_column].astype(str),\n                             'raw_label': table[label_column].astype(int)})\n    rows = []\n    for line in path.read_text().strip().splitlines():\n        parts = line.strip().split()\n        if len(parts) >= 2:\n            rows.append({'filename': parts[0], 'raw_label': int(parts[1])})\n    return pd.DataFrame(rows)\n\n\ndef index_image_files(root):\n    # file name (and name without extension) -> path relative to root\n    index = {}\n    for extension in ('*.jpg', '*.jpeg', '*.png', '*.JPG', '*.JPEG', '*.PNG'):\n        for path in Path(root).rglob(extension):\n            relative_path = str(path.relative_to(root))\n            index.setdefault(path.name, relative_path)\n            index.setdefault(path.stem, relative_path)\n    return index\n\n\ndef make_stratified_split(table, val_frac=0.15, test_frac=0.15, seed=42):\n    from sklearn.model_selection import train_test_split\n\n    def safe_strata(labels):\n        # merge classes too small to stratify\n        counts = labels.value_counts()\n        rare = set(counts[counts < 2].index)\n        merged = labels.where(~labels.isin(rare), -99)\n        return None if merged.value_counts().min() < 2 else merged\n\n    train_index, holdout_index = train_test_split(\n        table.index, test_size=val_frac + test_frac, random_state=seed,\n        stratify=safe_strata(table['raw_label']))\n    val_index, test_index = train_test_split(\n        holdout_index, test_size=test_frac / (val_frac + test_frac), random_state=seed,\n        stratify=safe_strata(table.loc[holdout_index, 'raw_label']))\n\n    table = table.copy()\n    table['split'] = 'train'\n    table.loc[val_index, 'split'] = 'val'\n    table.loc[test_index, 'split'] = 'test'\n    return table\n\n\ndef load_ddr(root):\n    root = Path(root)\n    image_index = index_image_files(root)\n    if not image_index:\n        raise FileNotFoundError(f'no images found under {root}')\n\n    # layout A: official train / valid / test label files\n    split_files = {'train': ['train.txt', 'train.csv'],\n                   'val': ['valid.txt', 'val.txt', 'valid.csv', 'val.csv'],\n                   'test': ['test.txt', 'test.csv']}\n    found_files = {}\n    for split, names in split_files.items():\n        for name in names:\n            matches = list(root.rglob(name))\n            if matches:\n                found_files[split] = matches[0]\n                break\n\n    split_source = 'official'\n    if len(found_files) >= 2:\n        parts = []\n        for split, path in found_files.items():\n            part = read_label_file(path)\n            part['split'] = split\n            parts.append(part)\n        table = pd.concat(parts, ignore_index=True)\n    else:\n        split_source = 'generated'\n        # layout B: one grading csv\n        label_csvs = [p for p in root.rglob('*.csv')\n                      if 'grading' in p.name.lower() or 'label' in p.name.lower()]\n        if label_csvs:\n            table = read_label_file(label_csvs[0])\n        else:\n            # layout C: one folder per class\n            rows = []\n            for class_folder in sorted(root.rglob('*')):\n                if class_folder.is_dir() and class_folder.name.isdigit():\n                    for image_file in class_folder.iterdir():\n                        if image_file.suffix.lower() in {'.jpg', '.jpeg', '.png'}:\n                            rows.append({'filename': image_file.name,\n                                         'raw_label': int(class_folder.name)})\n            if not rows:\n                raise RuntimeError('unknown DDR layout - check the folder tree above')\n            table = pd.DataFrame(rows)\n        table = make_stratified_split(table, config.val_frac, config.test_frac, config.seed)\n\n    # match labels to real files\n    table['image'] = table['filename'].map(\n        lambda name: image_index.get(name) or image_index.get(Path(name).name)\n        or image_index.get(Path(name).stem))\n    missing_count = int(table['image'].isna().sum())\n    if missing_count:\n        print(f'[DDR] {missing_count} listed images not found on disk - dropped')\n    table = table.dropna(subset=['image']).reset_index(drop=True)\n\n    # keep ungradable images: grade -1, quality 1\n    table['ungradable'] = table['raw_label'] == UNGRADABLE_LABEL\n    table['quality'] = table['ungradable'].astype(int)\n    table['level'] = table['raw_label'].where(~table['ungradable'], -1).astype(int)\n    table['source'] = 'DDR'\n    table['split_source'] = split_source\n    return table[['image', 'level', 'quality', 'split', 'ungradable', 'source',\n                  'split_source', 'raw_label']]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:32.962096Z","iopub.execute_input":"2026-09-21T08:05:32.962426Z","iopub.status.idle":"2026-09-21T08:05:32.980736Z","shell.execute_reply.started":"2026-09-21T08:05:32.962377Z","shell.execute_reply":"2026-09-21T08:05:32.979869Z"}},"outputs":[],"execution_count":null},{"id":"3dea7554-049d-40ab-ab35-c22b9d5ed719","cell_type":"code","source":"ddr_images = load_ddr(DDR_FOLDER)\n\nprint(f'DDR images loaded : {len(ddr_images)}')\nprint(f'split source      : {ddr_images[\"split_source\"].iloc[0]}')\nprint(f'ungradable images : {int(ddr_images[\"ungradable\"].sum())}')\nshow_table('First 5 rows', ddr_images.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:32.983399Z","iopub.execute_input":"2026-09-21T08:05:32.983735Z","iopub.status.idle":"2026-09-21T08:05:53.879555Z","shell.execute_reply.started":"2026-09-21T08:05:32.983704Z","shell.execute_reply":"2026-09-21T08:05:53.878605Z"}},"outputs":[],"execution_count":null},{"id":"ce42e3f4-4f11-4974-a59f-4f60bc51434e","cell_type":"markdown","source":"### 1.4 APTOS 2019 loader (external test set)\nEvery APTOS image is marked `split = \"external\"`. It is never used for training or for choosing thresholds, so its results measure how well the model generalises to new hospitals and cameras.","metadata":{}},{"id":"e65860fc-a339-4afd-98a8-4b2a6921b456","cell_type":"code","source":"def load_aptos(root, labels_csv='train.csv', image_folder='train_images'):\n    root = Path(root)\n    table = pd.read_csv(root / labels_csv)\n    table.columns = [c.strip().lower() for c in table.columns]\n    id_column = 'id_code' if 'id_code' in table.columns else table.columns[0]\n    label_column = 'diagnosis' if 'diagnosis' in table.columns else table.columns[1]\n    return pd.DataFrame({\n        'image': table[id_column].astype(str).map(lambda image_id: f'{image_folder}/{image_id}.png'),\n        'level': table[label_column].astype(int),\n        'quality': -1,\n        'split': 'external',\n        'ungradable': False,\n        'source': 'APTOS',\n        'raw_label': table[label_column].astype(int),\n    })\n\n\naptos_images = load_aptos(APTOS_FOLDER)\nprint(f'APTOS images loaded: {len(aptos_images)}')\nshow_table('First 3 rows', aptos_images.head(3))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:53.880729Z","iopub.execute_input":"2026-09-21T08:05:53.881475Z","iopub.status.idle":"2026-09-21T08:05:53.963823Z","shell.execute_reply.started":"2026-09-21T08:05:53.881447Z","shell.execute_reply":"2026-09-21T08:05:53.963164Z"}},"outputs":[],"execution_count":null},{"id":"c4d35985-0719-44ff-b760-d350851f190f","cell_type":"markdown","source":"### 1.5 Combined image table\nOne table listing every image from both datasets, saved as `manifest.csv`.","metadata":{}},{"id":"b5a9ae91-680f-4ebf-aebb-8e6b5366241c","cell_type":"code","source":"all_images = pd.concat([ddr_images, aptos_images], ignore_index=True)\nall_images.to_csv('results/tables/manifest.csv', index=False)\nshow_table('Images per dataset and split', all_images.groupby(['source', 'split']).size().to_frame('images'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:53.9646Z","iopub.execute_input":"2026-09-21T08:05:53.964817Z","iopub.status.idle":"2026-09-21T08:05:54.03891Z","shell.execute_reply.started":"2026-09-21T08:05:53.964783Z","shell.execute_reply":"2026-09-21T08:05:54.038081Z"}},"outputs":[],"execution_count":null},{"id":"b848bae4-da14-4d22-8124-8bf16d3f80ad","cell_type":"markdown","source":"### 1.6 Class distribution\nCounts images per DR stage. The data is heavily imbalanced (most images are \"No DR\"), which is why class balancing is used in Section 3 and why accuracy alone is not a sufficient metric.","metadata":{}},{"id":"35746cdb-f3d0-4a04-99fe-ca30107cf058","cell_type":"code","source":"def grade_distribution_table(table):\n    graded = table[~table['ungradable']]\n    counts = graded['level'].value_counts().sort_index()\n    result = pd.DataFrame({'grade': counts.index,\n                           'name': [GRADE_NAMES[g] for g in counts.index],\n                           'count': counts.values})\n    result['percent'] = (result['count'] / result['count'].sum() * 100).round(2)\n    result['ratio_to_majority'] = (result['count'].max() / result['count']).round(1)\n    return result\n\n\ndef split_summary_table(table):\n    graded = table[~table['ungradable']]\n    summary = pd.crosstab(graded['split'], graded['level'])\n    summary.columns = [f'{g} ({GRADE_NAMES[g]})' for g in summary.columns]\n    summary['Gradable total'] = summary.sum(axis=1)\n    summary['Ungradable'] = table[table['ungradable']].groupby('split').size()\n    summary['Ungradable'] = summary['Ungradable'].fillna(0).astype(int)\n    summary['All images'] = summary['Gradable total'] + summary['Ungradable']\n    referable_share = (graded['level'] >= 2).groupby(graded['split']).mean()\n    summary['Referable %'] = (referable_share * 100).round(1)\n    return summary\n\n\ngrade_distribution = grade_distribution_table(ddr_images)\ngrade_distribution.to_csv('results/tables/table1_class_distribution.csv', index=False)\nshow_table('DDR grade distribution', grade_distribution)\nprint(f\"\\nimbalance ratio (largest / smallest class): {grade_distribution['ratio_to_majority'].max():.1f} : 1\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:54.040012Z","iopub.execute_input":"2026-09-21T08:05:54.04025Z","iopub.status.idle":"2026-09-21T08:05:54.058485Z","shell.execute_reply.started":"2026-09-21T08:05:54.040228Z","shell.execute_reply":"2026-09-21T08:05:54.057843Z"}},"outputs":[],"execution_count":null},{"id":"342351a0-264b-46ed-971c-f568327e0137","cell_type":"code","source":"fig, ax = plt.subplots(figsize=(7.5, 3.6))\nbars = ax.bar(grade_distribution['name'], grade_distribution['count'], color='#7F77DD', edgecolor='none')\nax.set_yscale('log')\nax.set_ylabel('images (log scale)')\nax.set_title('DDR grade distribution')\nfor bar, count, percent in zip(bars, grade_distribution['count'], grade_distribution['percent']):\n    ax.text(bar.get_x() + bar.get_width() / 2, count, f'{count}\\n{percent}%',\n            ha='center', va='bottom', fontsize=8)\nplt.xticks(rotation=15)\nsave_figure('results/figures/fig2_class_distribution.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:54.059388Z","iopub.execute_input":"2026-09-21T08:05:54.0598Z","iopub.status.idle":"2026-09-21T08:05:54.567556Z","shell.execute_reply.started":"2026-09-21T08:05:54.059758Z","shell.execute_reply":"2026-09-21T08:05:54.566865Z"}},"outputs":[],"execution_count":null},{"id":"85eb033d-8375-4a41-8dad-0166e1fd223b","cell_type":"markdown","source":"### 1.7 Ungradable images\nCounts the photos marked ungradable in each split. These train the quality gate that asks for a re-capture instead of giving an unreliable grade.","metadata":{}},{"id":"7d205a08-9236-4bc2-8982-dad7c6fc51d4","cell_type":"code","source":"ungradable_count = int(ddr_images['ungradable'].sum())\nprint(f'ungradable images: {ungradable_count} / {len(ddr_images)} = {ungradable_count / len(ddr_images):.1%}')\nif ungradable_count == 0:\n    print('WARNING: no ungradable images in this copy of DDR - the quality head has no reject examples')\nshow_table('Ungradable images per split',\n           ddr_images.groupby('split')['ungradable'].agg(['sum', 'mean'])\n           .rename(columns={'sum': 'ungradable', 'mean': 'fraction'}))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:54.568496Z","iopub.execute_input":"2026-09-21T08:05:54.56888Z","iopub.status.idle":"2026-09-21T08:05:54.583169Z","shell.execute_reply.started":"2026-09-21T08:05:54.568853Z","shell.execute_reply":"2026-09-21T08:05:54.582509Z"}},"outputs":[],"execution_count":null},{"id":"f061436f-5539-4cd8-bc85-d71a428eca1b","cell_type":"markdown","source":"### 1.8 Split summary\nChecks that every grade appears in every split, and shows the share of referable DR (grade 2 or higher) per split.","metadata":{}},{"id":"757b1e3d-2a97-41a8-b894-93d32229b364","cell_type":"code","source":"split_summary = split_summary_table(all_images)\nsplit_summary.to_csv('results/tables/table2_split_summary.csv')\nshow_table('Images per split and grade', split_summary)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:54.584109Z","iopub.execute_input":"2026-09-21T08:05:54.584418Z","iopub.status.idle":"2026-09-21T08:05:54.623944Z","shell.execute_reply.started":"2026-09-21T08:05:54.584394Z","shell.execute_reply":"2026-09-21T08:05:54.623161Z"}},"outputs":[],"execution_count":null},{"id":"734c20ac-3608-468a-92c6-662438832df1","cell_type":"markdown","source":"### 1.9 Image sizes\nMeasures resolution and aspect ratio on a sample. Images come from many cameras, so they must be cropped and resized to one size.","metadata":{}},{"id":"e3f54dfa-ec23-4315-9a0d-5850842bce41","cell_type":"code","source":"size_rows = []\nfor relative_path in ddr_images.sample(min(60, len(ddr_images)), random_state=RANDOM_SEED)['image']:\n    image = cv2.imread(f'{DDR_FOLDER}/{relative_path}')\n    if image is not None:\n        height, width = image.shape[:2]\n        size_rows.append({'height': height, 'width': width,\n                          'aspect_ratio': round(width / height, 3),\n                          'megapixels': round(height * width / 1e6, 2)})\n\nimage_sizes = pd.DataFrame(size_rows)\nprint(f'{len(image_sizes)} images sampled')\nshow_table('Image size statistics', image_sizes.describe().round(2))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:54.624791Z","iopub.execute_input":"2026-09-21T08:05:54.624997Z","iopub.status.idle":"2026-09-21T08:05:57.510696Z","shell.execute_reply.started":"2026-09-21T08:05:54.624978Z","shell.execute_reply":"2026-09-21T08:05:57.509812Z"}},"outputs":[],"execution_count":null},{"id":"1e3381a2-bf03-4b50-9a06-95379073368d","cell_type":"markdown","source":"### 1.10 Example images per grade\nThree random images for each DR stage.","metadata":{}},{"id":"9bfff778-f8ea-4434-ab1d-b40550bd6125","cell_type":"code","source":"fig, axes = plt.subplots(3, 5, figsize=(15, 9))\nfor column, grade in enumerate(range(5)):\n    examples = ddr_images[ddr_images['level'] == grade]\n    examples = examples.sample(min(3, len(examples)), random_state=RANDOM_SEED)\n    for row in range(3):\n        axes[row, column].axis('off')\n    for row, (_, record) in enumerate(examples.iterrows()):\n        image = cv2.imread(f\"{DDR_FOLDER}/{record['image']}\")\n        axes[row, column].imshow(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))\n    axes[0, column].set_title(f'Grade {grade}: {GRADE_NAMES[grade]}', fontsize=10)\nplt.suptitle('DDR examples by DR grade', y=0.98)\nsave_figure('results/figures/fig3_samples_by_grade.png', dpi=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:05:57.511818Z","iopub.execute_input":"2026-09-21T08:05:57.512178Z","iopub.status.idle":"2026-09-21T08:06:02.669459Z","shell.execute_reply.started":"2026-09-21T08:05:57.512114Z","shell.execute_reply":"2026-09-21T08:06:02.668603Z"}},"outputs":[],"execution_count":null},{"id":"8e56c094-2782-4782-8761-fcaf3cb6809b","cell_type":"markdown","source":"### 1.11 Ungradable examples\nExamples of photos too poor to grade (blur, bad exposure, poor field of view).","metadata":{}},{"id":"fd93a153-4c3e-44d3-be40-6b8eb31fec30","cell_type":"code","source":"ungradable_images = ddr_images[ddr_images['ungradable']]\nif len(ungradable_images) == 0:\n    print('no ungradable images to show')\nelse:\n    fig, axes = plt.subplots(1, 5, figsize=(15, 3.2))\n    for ax in axes:\n        ax.axis('off')\n    for ax, (_, record) in zip(axes, ungradable_images.sample(min(5, len(ungradable_images)),\n                                                              random_state=RANDOM_SEED).iterrows()):\n        image = cv2.imread(f\"{DDR_FOLDER}/{record['image']}\")\n        ax.imshow(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))\n    plt.suptitle('DDR ungradable images')\n    save_figure('results/figures/fig3b_ungradable_samples.png', dpi=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:02.670895Z","iopub.execute_input":"2026-09-21T08:06:02.671279Z","iopub.status.idle":"2026-09-21T08:06:02.680913Z","shell.execute_reply.started":"2026-09-21T08:06:02.67124Z","shell.execute_reply":"2026-09-21T08:06:02.680124Z"}},"outputs":[],"execution_count":null},{"id":"2edc12d6-9c01-4cc0-9aa3-18bfe1504135","cell_type":"markdown","source":"## 2. Preprocessing\n\nRaw fundus photos have large black borders, uneven lighting from the camera flash, and low contrast around small lesions. The pipeline fixes each problem in turn:\n\n1. **Crop** to the retina and **pad** to a square (no stretching of lesion shapes)\n2. **Resize** to 384 x 384\n3. **CLAHE** on the green channel (contrast enhancement where vessels and lesions are most visible)\n4. **Illumination correction** by subtracting a blurred copy of the image - this removes slow lighting changes and sharpens edges and small lesions (edge enhancement)\n5. **Circular mask** to remove the bright rim at the edge of the field of view\n\n### 2.1 Preprocessing functions","metadata":{}},{"id":"7ea93c5e-bca9-43c0-b844-f7a1a2b6a3ff","cell_type":"code","source":"def retina_mask(image, threshold=10):\n    # bright pixels = retina, dark pixels = background\n    mask = image.max(axis=2) > threshold\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (7, 7))\n    mask = cv2.morphologyEx(mask.astype(np.uint8), cv2.MORPH_OPEN, kernel)\n    return mask.astype(bool)\n\n\ndef crop_to_retina(image, threshold=10):\n    mask = retina_mask(image, threshold)\n    if not mask.any():\n        return image\n    rows = np.where(mask.any(axis=1))[0]\n    columns = np.where(mask.any(axis=0))[0]\n    return image[rows[0]:rows[-1] + 1, columns[0]:columns[-1] + 1]\n\n\ndef pad_to_square(image):\n    height, width = image.shape[:2]\n    side = max(height, width)\n    top, left = (side - height) // 2, (side - width) // 2\n    canvas = np.zeros((side, side, image.shape[2]), dtype=image.dtype)\n    canvas[top:top + height, left:left + width] = image\n    return canvas\n\n\ndef circular_crop(image, shrink=0.97):\n    # black out everything outside a centred circle\n    height, width = image.shape[:2]\n    mask = np.zeros((height, width), dtype=np.uint8)\n    cv2.circle(mask, (width // 2, height // 2), int(min(height, width) / 2 * shrink), 255, thickness=-1)\n    return cv2.bitwise_and(image, image, mask=mask)\n\n\ndef subtract_local_average(image, sigma_ratio=30.0, weight=4.0, bias=128):\n    # illumination correction + edge enhancement: image minus its blurred copy\n    sigma = max(image.shape[1] / sigma_ratio, 1.0)\n    blurred = cv2.GaussianBlur(image, (0, 0), sigmaX=sigma)\n    return cv2.addWeighted(image, weight, blurred, -weight, bias)\n\n\ndef clahe_green(image, clip_limit=2.0, tile=8):\n    # contrast enhancement on the green channel only\n    enhanced = image.copy()\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=(tile, tile))\n    enhanced[:, :, 1] = clahe.apply(enhanced[:, :, 1])\n    return enhanced\n\n\ndef preprocess_fundus(image, size=384, use_graham=True, use_clahe=True):\n    # crop -> square -> resize -> CLAHE -> illumination correction -> circle mask\n    processed = crop_to_retina(image)\n    processed = pad_to_square(processed)\n    processed = cv2.resize(processed, (size, size), interpolation=cv2.INTER_AREA)\n    if use_clahe:\n        processed = clahe_green(processed)\n    if use_graham:\n        processed = subtract_local_average(processed)\n    return circular_crop(processed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:02.681814Z","iopub.execute_input":"2026-09-21T08:06:02.682439Z","iopub.status.idle":"2026-09-21T08:06:02.700984Z","shell.execute_reply.started":"2026-09-21T08:06:02.682399Z","shell.execute_reply":"2026-09-21T08:06:02.700267Z"}},"outputs":[],"execution_count":null},{"id":"738faf0b-b2d5-4c6a-a105-2acd3b22ba4d","cell_type":"markdown","source":"### 2.2 Pick a demonstration image\nOne moderate-DR image is followed through every step below.","metadata":{}},{"id":"01750aac-2637-4f26-a417-c4dadc6dc307","cell_type":"code","source":"demo_candidates = ddr_images[ddr_images['level'] == 2]\nif len(demo_candidates) == 0:\n    demo_candidates = ddr_images[~ddr_images['ungradable']]\ndemo_record = demo_candidates.sample(1, random_state=RANDOM_SEED).iloc[0]\ndemo_raw_image = cv2.imread(f\"{DDR_FOLDER}/{demo_record['image']}\")\nprint('demo image:', demo_record['image'])\nprint('size      :', demo_raw_image.shape, '| grade:', demo_record['level'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:02.70207Z","iopub.execute_input":"2026-09-21T08:06:02.702334Z","iopub.status.idle":"2026-09-21T08:06:02.723387Z","shell.execute_reply.started":"2026-09-21T08:06:02.702279Z","shell.execute_reply":"2026-09-21T08:06:02.722609Z"}},"outputs":[],"execution_count":null},{"id":"c68daa69-86b0-410b-9652-df7aa9620ec3","cell_type":"markdown","source":"### 2.3 Retina mask\nMeasures how much of the raw photo is black background that carries no information.","metadata":{}},{"id":"cda8319b-b9bc-4497-9641-e9bfcce41975","cell_type":"code","source":"retina_region = retina_mask(demo_raw_image)\nbackground_fraction = 1 - retina_region.mean()\nprint(f'background: {background_fraction:.1%} of pixels | retina: {1 - background_fraction:.1%}')\n\nfig, axes = plt.subplots(1, 2, figsize=(9, 4))\naxes[0].imshow(cv2.cvtColor(demo_raw_image, cv2.COLOR_BGR2RGB)); axes[0].set_title('raw')\naxes[1].imshow(retina_region, cmap='gray'); axes[1].set_title('retina mask')\nfor ax in axes:\n    ax.axis('off')\nsave_figure('results/figures/fig4_retina_mask.png', dpi=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:02.724218Z","iopub.execute_input":"2026-09-21T08:06:02.724995Z","iopub.status.idle":"2026-09-21T08:06:03.316666Z","shell.execute_reply.started":"2026-09-21T08:06:02.724973Z","shell.execute_reply":"2026-09-21T08:06:03.31586Z"}},"outputs":[],"execution_count":null},{"id":"2b74563a-9ed2-44ea-8eb7-35ebe51485d9","cell_type":"markdown","source":"### 2.4 Crop, pad and resize","metadata":{}},{"id":"aa1fcdf4-a28f-4429-be58-5863232738f1","cell_type":"code","source":"cropped_image = crop_to_retina(demo_raw_image)\nsquare_image = pad_to_square(cropped_image)\nresized_image = cv2.resize(square_image, (config.image_size, config.image_size), interpolation=cv2.INTER_AREA)\n\nprint(f'raw     : {demo_raw_image.shape[:2]}')\nprint(f'cropped : {cropped_image.shape[:2]}')\nprint(f'square  : {square_image.shape[:2]}')\nprint(f'resized : {resized_image.shape[:2]}')\nprint(f'pixels reduced by {1 - np.prod(resized_image.shape[:2]) / np.prod(demo_raw_image.shape[:2]):.1%}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:03.317883Z","iopub.execute_input":"2026-09-21T08:06:03.318289Z","iopub.status.idle":"2026-09-21T08:06:03.342002Z","shell.execute_reply.started":"2026-09-21T08:06:03.31826Z","shell.execute_reply":"2026-09-21T08:06:03.341391Z"}},"outputs":[],"execution_count":null},{"id":"055cc476-c53b-4f0c-b83d-8456a5ab9317","cell_type":"markdown","source":"### 2.5 Contrast enhancement (CLAHE on the green channel)\nThe histogram shows the green-channel intensities spreading out after CLAHE, which means higher contrast.","metadata":{}},{"id":"76864a9b-2e0f-4f47-b7ce-a987ae4d772e","cell_type":"code","source":"contrast_enhanced_image = clahe_green(resized_image)\n\ngreen_std_before = resized_image[:, :, 1].std()\ngreen_std_after = contrast_enhanced_image[:, :, 1].std()\nprint(f'green channel contrast (std): {green_std_before:.1f} -> {green_std_after:.1f}')\n\nfig, axes = plt.subplots(1, 3, figsize=(13, 4.2))\naxes[0].imshow(cv2.cvtColor(resized_image, cv2.COLOR_BGR2RGB)); axes[0].set_title('before')\naxes[1].imshow(cv2.cvtColor(contrast_enhanced_image, cv2.COLOR_BGR2RGB)); axes[1].set_title('after CLAHE')\naxes[2].hist(resized_image[:, :, 1].ravel(), bins=60, alpha=0.55, label='before', color='#7F77DD')\naxes[2].hist(contrast_enhanced_image[:, :, 1].ravel(), bins=60, alpha=0.55, label='after', color='#E8734A')\naxes[2].set_title('green channel histogram'); axes[2].legend()\nfor ax in axes[:2]:\n    ax.axis('off')\nsave_figure('results/figures/fig6_clahe_effect.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:03.342818Z","iopub.execute_input":"2026-09-21T08:06:03.343089Z","iopub.status.idle":"2026-09-21T08:06:04.640825Z","shell.execute_reply.started":"2026-09-21T08:06:03.343042Z","shell.execute_reply":"2026-09-21T08:06:04.639996Z"}},"outputs":[],"execution_count":null},{"id":"5a9443ab-d2af-4c1f-9370-414c1c99fcc7","cell_type":"markdown","source":"### 2.6 Illumination correction and edge enhancement\nMeasures uneven lighting as the brightness difference between the left and right thirds of the image, before and after correction.","metadata":{}},{"id":"a5e597c7-e983-42c4-9014-91029a515d5b","cell_type":"code","source":"def illumination_gradient(image):\n    # brightness difference between left and right thirds\n    third = image.shape[1] // 3\n    return abs(float(image[:, :third].mean()) - float(image[:, -third:].mean()))\n\n\nillumination_corrected_image = subtract_local_average(contrast_enhanced_image)\ngradient_before = illumination_gradient(contrast_enhanced_image)\ngradient_after = illumination_gradient(illumination_corrected_image)\nprint(f'illumination gradient: {gradient_before:.2f} -> {gradient_after:.2f} '\n      f'({1 - gradient_after / max(gradient_before, 1e-6):.0%} reduction)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:04.641938Z","iopub.execute_input":"2026-09-21T08:06:04.642438Z","iopub.status.idle":"2026-09-21T08:06:04.675996Z","shell.execute_reply.started":"2026-09-21T08:06:04.642414Z","shell.execute_reply":"2026-09-21T08:06:04.67519Z"}},"outputs":[],"execution_count":null},{"id":"e674a14b-b9ce-4a3b-a8a0-69eea7866d09","cell_type":"markdown","source":"### 2.7 Illumination check on 40 images\nRepeats the measurement on 40 random images to show the correction works across the dataset, not just on one image.","metadata":{}},{"id":"9ee42538-f305-46b2-a7a9-f2623e4bf496","cell_type":"code","source":"illumination_rows = []\nfor relative_path in ddr_images.sample(min(40, len(ddr_images)), random_state=RANDOM_SEED)['image']:\n    image = cv2.imread(f'{DDR_FOLDER}/{relative_path}')\n    if image is None:\n        continue\n    base = cv2.resize(pad_to_square(crop_to_retina(image)),\n                      (config.image_size, config.image_size), interpolation=cv2.INTER_AREA)\n    illumination_rows.append({'before': illumination_gradient(base),\n                              'after': illumination_gradient(subtract_local_average(clahe_green(base)))})\n\nillumination_results = pd.DataFrame(illumination_rows)\nillumination_results['reduction_%'] = (1 - illumination_results['after']\n                                       / illumination_results['before'].clip(lower=1e-6)) * 100\nillumination_results.to_csv('results/tables/table3_illumination_reduction.csv', index=False)\nshow_table('Illumination gradient on 40 images', illumination_results.describe().round(2))\nprint(f\"\\nmean reduction: {illumination_results['reduction_%'].mean():.1f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:04.677033Z","iopub.execute_input":"2026-09-21T08:06:04.678119Z","iopub.status.idle":"2026-09-21T08:06:07.612183Z","shell.execute_reply.started":"2026-09-21T08:06:04.678094Z","shell.execute_reply":"2026-09-21T08:06:07.611546Z"}},"outputs":[],"execution_count":null},{"id":"fdd91439-c8dd-481b-acc2-7b26b9337fe8","cell_type":"markdown","source":"### 2.8 All preprocessing stages","metadata":{}},{"id":"085346be-1184-42db-a181-a1822d970d35","cell_type":"code","source":"pipeline_stages = {\n    'raw': demo_raw_image,\n    '1. cropped + padded': square_image,\n    '2. resized': resized_image,\n    '3. CLAHE': contrast_enhanced_image,\n    '4. illumination corrected': illumination_corrected_image,\n    '5. circular mask': circular_crop(illumination_corrected_image),\n}\n\nfig, axes = plt.subplots(1, len(pipeline_stages), figsize=(3 * len(pipeline_stages), 3.4))\nfor ax, (stage_name, stage_image) in zip(axes, pipeline_stages.items()):\n    ax.imshow(cv2.cvtColor(stage_image, cv2.COLOR_BGR2RGB))\n    ax.set_title(stage_name, fontsize=9); ax.axis('off')\nsave_figure('results/figures/fig5_preprocessing_stages.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:07.613395Z","iopub.execute_input":"2026-09-21T08:06:07.614031Z","iopub.status.idle":"2026-09-21T08:06:09.304313Z","shell.execute_reply.started":"2026-09-21T08:06:07.614003Z","shell.execute_reply":"2026-09-21T08:06:09.302438Z"}},"outputs":[],"execution_count":null},{"id":"e28fd5bd-c17a-491f-95dd-d3d9f90db9ee","cell_type":"markdown","source":"### 2.9 Preprocessing time per image","metadata":{}},{"id":"fd0a67da-d9ea-49af-84b9-67fe5e0e713c","cell_type":"code","source":"timing_paths = ddr_images.sample(min(20, len(ddr_images)), random_state=1)['image']\nstart_time = time.perf_counter()\nfor relative_path in timing_paths:\n    image = cv2.imread(f'{DDR_FOLDER}/{relative_path}')\n    if image is not None:\n        preprocess_fundus(image, size=config.image_size)\nms_per_image = (time.perf_counter() - start_time) / len(timing_paths) * 1000\nprint(f'preprocessing time: {ms_per_image:.1f} ms per full-resolution image')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:09.30873Z","iopub.execute_input":"2026-09-21T08:06:09.308958Z","iopub.status.idle":"2026-09-21T08:06:10.585629Z","shell.execute_reply.started":"2026-09-21T08:06:09.308938Z","shell.execute_reply":"2026-09-21T08:06:10.584825Z"}},"outputs":[],"execution_count":null},{"id":"ed8443ca-540c-49df-b00d-22477d60c0ce","cell_type":"markdown","source":"### 2.10 Preprocess every image once and cache it\nEvery image (DDR and APTOS) is preprocessed once and saved as a 384 x 384 JPEG. Training then reads the small cached files, which makes each epoch much faster and guarantees training, evaluation and Grad-CAM all see exactly the same preprocessing. Already-cached files are skipped on later runs.","metadata":{}},{"id":"6cf52e7c-2441-4946-b91f-42979aa93c35","cell_type":"code","source":"CACHE_FOLDER = Path(config.cache_dir) / str(config.image_size)\nCACHE_FOLDER.mkdir(parents=True, exist_ok=True)\nDATASET_FOLDERS = {'DDR': DDR_FOLDER, 'APTOS': APTOS_FOLDER}\n\n\ndef cached_file_path(source, relative_path):\n    flat_name = Path(relative_path).with_suffix('.jpg').as_posix().replace('/', '__')\n    return str(CACHE_FOLDER / f'{source}__{flat_name}')\n\n\ndef cache_one_image(source_path, cached_path, size=config.image_size):\n    if os.path.exists(cached_path):\n        return True\n    image = cv2.imread(source_path, cv2.IMREAD_COLOR)\n    if image is None:\n        return False\n    # shrink very large images first to save time\n    height, width = image.shape[:2]\n    scale = 3 * size / max(height, width)\n    if scale < 1:\n        image = cv2.resize(image, (round(width * scale), round(height * scale)), interpolation=cv2.INTER_AREA)\n    processed = preprocess_fundus(image, size, config.use_graham, config.use_clahe)\n    return bool(cv2.imwrite(cached_path, processed, [cv2.IMWRITE_JPEG_QUALITY, 95]))\n\n\ndef build_cache(table):\n    table = table.copy()\n    table['cached'] = [cached_file_path(source, path) for source, path in zip(table['source'], table['image'])]\n    jobs = [(f'{DATASET_FOLDERS[source]}/{path}', cached)\n            for source, path, cached in zip(table['source'], table['image'], table['cached'])]\n    with ThreadPoolExecutor(max_workers=os.cpu_count() or 4) as pool:\n        succeeded = list(tqdm(pool.map(lambda job: cache_one_image(*job), jobs), total=len(jobs)))\n    succeeded = np.array(succeeded, dtype=bool)\n    if (~succeeded).sum():\n        print(f'{int((~succeeded).sum())} unreadable images dropped')\n    return table[succeeded].reset_index(drop=True)\n\n\nstart_time = time.time()\nddr_images = build_cache(ddr_images)\naptos_images = build_cache(aptos_images)\nprint(f'cache ready in {(time.time() - start_time) / 60:.1f} min')\nprint(f'cached images: {len(ddr_images)} DDR + {len(aptos_images)} APTOS -> {CACHE_FOLDER}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:06:10.586737Z","iopub.execute_input":"2026-09-21T08:06:10.587184Z","iopub.status.idle":"2026-09-21T08:18:33.852862Z","shell.execute_reply.started":"2026-09-21T08:06:10.587155Z","shell.execute_reply":"2026-09-21T08:18:33.852122Z"}},"outputs":[],"execution_count":null},{"id":"eec67b7f-fef5-4f72-a4a9-8626a44d2751","cell_type":"markdown","source":"## 3. Data augmentation and class balancing\n\n**Augmentation** creates a new random variation of each training image every time it is used, so the network does not memorise the rare classes. Only changes that keep the diagnosis the same are used:\n\n| Augmentation | Setting | Reason |\n|---|---|---|\n| Horizontal / vertical flip | 50% each | left and right eyes are mirror images |\n| Rotation | any angle, 70% | the camera has no fixed \"up\" |\n| Zoom (scale) | 90-110% | framing varies between photos |\n| Shift | up to 5% | the retina is not always centred |\n| Brightness / contrast | +/-12%, 50% | flash and exposure vary between cameras |\n| Light blur or noise | 20% | small focus and sensor differences |\n\nStrong colour changes, large cut-outs and heavy blur are **not** used: colour carries diagnostic information, cut-outs can erase the few lesions that define a grade, and heavy blur is exactly what the quality head must learn to detect.\n\n**Class balancing**: a weighted sampler draws rare grades (and ungradable images) more often, using effective-number weights that correct the imbalance without over-correcting.\n\n### 3.1 Augmentation pipeline","metadata":{}},{"id":"186d1261-bb75-4f12-a71f-8f980d17ed3a","cell_type":"code","source":"IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)\nIMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)\n\n\nclass FundusTransform:\n    # train: flips, rotation / zoom / shift, brightness / contrast, light blur or noise\n    # eval: resize + normalise only\n\n    def __init__(self, size, train):\n        self.size = size\n        self.train = train\n\n    def augment(self, image):\n        rng = np.random\n        if rng.rand() < 0.5:\n            image = image[:, ::-1]                      # horizontal flip\n        if rng.rand() < 0.5:\n            image = image[::-1]                         # vertical flip\n        image = np.ascontiguousarray(image)\n\n        if rng.rand() < 0.7:                            # rotation + zoom + shift\n            height, width = image.shape[:2]\n            matrix = cv2.getRotationMatrix2D((width / 2, height / 2),\n                                             rng.uniform(-180, 180), rng.uniform(0.9, 1.1))\n            matrix[:, 2] += rng.uniform(-0.05, 0.05, size=2) * np.array([width, height])\n            image = cv2.warpAffine(image, matrix, (width, height), flags=cv2.INTER_LINEAR,\n                                   borderMode=cv2.BORDER_CONSTANT, borderValue=0)\n\n        if rng.rand() < 0.5:                            # brightness / contrast\n            contrast = 1 + rng.uniform(-0.12, 0.12)\n            brightness = rng.uniform(-0.12, 0.12) * 255\n            image = np.clip(image.astype(np.float32) * contrast + brightness, 0, 255).astype(np.uint8)\n\n        if rng.rand() < 0.2:                            # light blur or noise\n            if rng.rand() < 0.5:\n                image = cv2.GaussianBlur(image, (3, 3), 0)\n            else:\n                noise = rng.normal(0, np.sqrt(rng.uniform(5, 20)), image.shape)\n                image = np.clip(image.astype(np.float32) + noise, 0, 255).astype(np.uint8)\n        return image\n\n    def __call__(self, image_rgb):\n        if image_rgb.shape[:2] != (self.size, self.size):\n            image_rgb = cv2.resize(image_rgb, (self.size, self.size), interpolation=cv2.INTER_AREA)\n        if self.train:\n            image_rgb = self.augment(image_rgb)\n        normalised = (image_rgb.astype(np.float32) / 255.0 - IMAGENET_MEAN) / IMAGENET_STD\n        return torch.from_numpy(np.ascontiguousarray(normalised.transpose(2, 0, 1)))\n\n\ndef denormalise(image_tensor):\n    # tensor -> RGB image in [0, 1] for plotting\n    image = image_tensor.permute(1, 2, 0).cpu().numpy() * IMAGENET_STD + IMAGENET_MEAN\n    return np.clip(image, 0, 1)\n\n\ntrain_transform = FundusTransform(config.image_size, train=True)\neval_transform = FundusTransform(config.image_size, train=False)\nprint('training transform  : flips, rotation, zoom, shift, brightness/contrast, blur or noise, normalise')\nprint('evaluation transform: resize, normalise')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:33.853862Z","iopub.execute_input":"2026-09-21T08:18:33.854212Z","iopub.status.idle":"2026-09-21T08:18:33.866455Z","shell.execute_reply.started":"2026-09-21T08:18:33.854185Z","shell.execute_reply":"2026-09-21T08:18:33.865592Z"}},"outputs":[],"execution_count":null},{"id":"734aab89-3813-4392-a7e2-e299e9486aea","cell_type":"markdown","source":"### 3.2 Eight augmented views of one image\nEach panel is a different random version of the same preprocessed image, as the network would see it during training.","metadata":{}},{"id":"300590c8-82c4-4490-a051-390427564c5b","cell_type":"code","source":"demo_preprocessed_rgb = cv2.cvtColor(preprocess_fundus(demo_raw_image, size=config.image_size), cv2.COLOR_BGR2RGB)\n\nfig, axes = plt.subplots(2, 4, figsize=(14, 7))\nfor view_number, ax in enumerate(axes.ravel(), start=1):\n    ax.imshow(denormalise(train_transform(demo_preprocessed_rgb)))\n    ax.axis('off'); ax.set_title(f'view {view_number}', fontsize=9)\nplt.suptitle('Eight augmented views of one image')\nsave_figure('results/figures/fig7_augmentation_grid.png', dpi=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:33.867431Z","iopub.execute_input":"2026-09-21T08:18:33.868098Z","iopub.status.idle":"2026-09-21T08:18:36.374049Z","shell.execute_reply.started":"2026-09-21T08:18:33.86807Z","shell.execute_reply":"2026-09-21T08:18:36.372771Z"}},"outputs":[],"execution_count":null},{"id":"323fa412-41b0-4887-ba09-83a25b3abd74","cell_type":"markdown","source":"### 3.3 Class weights\nCompares simple inverse-frequency weights with the effective-number weights used here. Ungradable images count as their own class.","metadata":{}},{"id":"f85cd982-8944-40e5-b727-0f9a98ea3f0f","cell_type":"code","source":"def effective_number_weights(counts, beta=0.999):\n    # class-balanced weights, mean weight = 1, zero for empty classes\n    counts = np.asarray(counts, dtype=np.float64)\n    effective_counts = (1.0 - np.power(beta, counts)) / (1.0 - beta)\n    weights = np.where(counts > 0, 1.0 / np.maximum(effective_counts, 1e-8), 0.0)\n    present = counts > 0\n    weights = weights / weights[present].sum() * present.sum()\n    return torch.tensor(weights, dtype=torch.float32)\n\n\ntrain_set = ddr_images[ddr_images['split'] == 'train'].reset_index(drop=True)\nvalidation_set = ddr_images[(ddr_images['split'] == 'val') & (~ddr_images['ungradable'])].reset_index(drop=True)\n\ntrain_label_counts = train_set['raw_label'].value_counts().sort_index()\nclass_weights = effective_number_weights(train_label_counts.values, beta=config.sampler_beta)\n\nclass_weight_table = pd.DataFrame({\n    'label': [LABEL_NAMES[label] for label in train_label_counts.index],\n    'count': train_label_counts.values,\n    'inverse_frequency_weight': (train_label_counts.max() / train_label_counts.values).round(2),\n    'effective_number_weight': class_weights.numpy().round(3)})\nshow_table('Class weights (training split)', class_weight_table)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:36.375239Z","iopub.execute_input":"2026-09-21T08:18:36.375645Z","iopub.status.idle":"2026-09-21T08:18:36.421413Z","shell.execute_reply.started":"2026-09-21T08:18:36.375618Z","shell.execute_reply":"2026-09-21T08:18:36.420502Z"}},"outputs":[],"execution_count":null},{"id":"ca8ba26e-604a-4eeb-8018-81c22cfb9166","cell_type":"markdown","source":"### 3.4 Check that the sampler rebalances the classes\nDraws one epoch from the sampler and compares the class shares with the original data.","metadata":{}},{"id":"a4478a75-6fb1-43de-825e-472c91253121","cell_type":"code","source":"def build_balanced_sampler(labels, beta=0.999):\n    labels = np.asarray(labels)\n    weight_per_class = effective_number_weights(np.bincount(labels), beta).numpy()\n    return WeightedRandomSampler(weights=torch.as_tensor(weight_per_class[labels], dtype=torch.double),\n                                 num_samples=len(labels), replacement=True)\n\n\n# ungradable images are their own sampling class (label 5)\nbalanced_sampler = build_balanced_sampler(train_set['raw_label'].values, beta=config.sampler_beta)\nsampled_labels = train_set['raw_label'].values[list(balanced_sampler)]\n\nsampler_effect = pd.DataFrame({\n    'original_%': (train_set['raw_label'].value_counts(normalize=True).sort_index() * 100).round(1),\n    'after_sampling_%': (pd.Series(sampled_labels).value_counts(normalize=True).sort_index() * 100).round(1),\n})\nsampler_effect.index = [LABEL_NAMES[label] for label in sampler_effect.index]\nsampler_effect.to_csv('results/tables/table4_sampler_effect.csv')\nshow_table('Class share before and after balanced sampling', sampler_effect)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:36.422624Z","iopub.execute_input":"2026-09-21T08:18:36.422996Z","iopub.status.idle":"2026-09-21T08:18:36.444675Z","shell.execute_reply.started":"2026-09-21T08:18:36.422953Z","shell.execute_reply":"2026-09-21T08:18:36.44408Z"}},"outputs":[],"execution_count":null},{"id":"a43878ac-555c-4928-95a7-e39f7998d2fe","cell_type":"markdown","source":"### 3.5 Datasets and data loaders\nThe training loader uses the balanced sampler and augmentation; the validation loader uses plain images in a fixed order. Validation contains gradable images only, because it measures grading performance.","metadata":{}},{"id":"5ff7a8aa-d9fd-4d8e-b4ae-812af81df9cf","cell_type":"code","source":"class FundusDataset(Dataset):\n    # reads cached preprocessed images, returns (image, grade, quality)\n\n    def __init__(self, table, image_size, train):\n        self.table = table.reset_index(drop=True)\n        self.transform = FundusTransform(image_size, train)\n\n    def __len__(self):\n        return len(self.table)\n\n    def __getitem__(self, index):\n        record = self.table.iloc[index]\n        image = cv2.imread(record['cached'], cv2.IMREAD_COLOR)\n        if image is None:\n            raise FileNotFoundError(record['cached'])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        return (self.transform(image),\n                torch.tensor(int(record['level'])),\n                torch.tensor(int(record['quality'])))\n\n\ndef seed_worker(worker_id):\n    # different random augmentations in each loader worker\n    worker_seed = torch.initial_seed() % 2 ** 32\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n    cv2.setNumThreads(0)\n\n\ndef make_loader(table, train, sampler=None, batch_size=None):\n    return DataLoader(FundusDataset(table, config.image_size, train),\n                      batch_size=batch_size or config.batch_size,\n                      sampler=sampler, shuffle=False,\n                      num_workers=config.num_workers,\n                      pin_memory=config.device == 'cuda',\n                      drop_last=train,\n                      worker_init_fn=seed_worker)\n\n\ntrain_loader = make_loader(train_set, train=True, sampler=balanced_sampler)\nvalidation_loader = make_loader(validation_set, train=False, batch_size=config.batch_size * 2)\n\nprint(f'training images   : {len(train_set)} ({int(train_set.ungradable.sum())} ungradable)')\nprint(f'validation images : {len(validation_set)}')\nprint(f'batches per epoch : {len(train_loader)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:36.445512Z","iopub.execute_input":"2026-09-21T08:18:36.445874Z","iopub.status.idle":"2026-09-21T08:18:36.455793Z","shell.execute_reply.started":"2026-09-21T08:18:36.445852Z","shell.execute_reply":"2026-09-21T08:18:36.455186Z"}},"outputs":[],"execution_count":null},{"id":"d73381f9-3aa8-4b18-89a8-8087b39618aa","cell_type":"markdown","source":"### 3.6 Check one batch\nConfirms image shape, value range and labels before training.","metadata":{}},{"id":"bf50e9e1-9a87-4e6c-906a-9089f77c2c94","cell_type":"code","source":"batch_images, batch_grades, batch_quality = next(iter(train_loader))\nprint('images  :', tuple(batch_images.shape), batch_images.dtype,\n      f'range [{batch_images.min():.2f}, {batch_images.max():.2f}]')\nprint('grades  :', tuple(batch_grades.shape), 'values', sorted(batch_grades.unique().tolist()))\nprint('quality :', tuple(batch_quality.shape), 'values', sorted(batch_quality.unique().tolist()))\nshow_table('Grades in this batch (-1 = ungradable)',\n           pd.Series(batch_grades.numpy()).value_counts().sort_index().to_frame('images'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:36.456776Z","iopub.execute_input":"2026-09-21T08:18:36.457094Z","iopub.status.idle":"2026-09-21T08:18:38.146402Z","shell.execute_reply.started":"2026-09-21T08:18:36.457061Z","shell.execute_reply":"2026-09-21T08:18:38.145617Z"}},"outputs":[],"execution_count":null},{"id":"8f5c88fb-b2b0-4d01-bdad-cf47b5a57b95","cell_type":"markdown","source":"## 4. CNN architecture and transfer learning\n\n- **Backbone**: EfficientNetV2-S pretrained on ImageNet-21k (about 20M parameters). It gives high accuracy for its size and trains fast, which fits the Kaggle GPU time limit.\n- **Transfer learning**: the pretrained layers are frozen for the first epoch while the new output layers learn, then the whole network is fine-tuned with a 10x smaller learning rate for the pretrained layers.\n- **Two heads** on shared features:\n  - **Grade head**: 4 outputs for ordinal (CORN) classification of the 5 DR stages\n  - **Quality head**: 2 outputs, gradable / ungradable\n- **Ordinal loss (CORN)**: DR stages are ordered, so predicting \"No DR\" for a proliferative eye must cost more than predicting \"Severe\". CORN splits the problem into four yes/no questions (\"is the grade above 0?\", \"above 1?\", ...).\n\n### 4.1 Ordinal (CORN) loss functions","metadata":{}},{"id":"088f462a-4281-4e43-8b90-03a8209edd67","cell_type":"code","source":"def corn_loss(logits, targets, num_classes=5):\n    # 4 yes/no tasks \"is grade > k\"; task k only trained on images with grade > k-1\n    total_loss = torch.zeros((), device=logits.device, dtype=torch.float32)\n    term_count = 0\n    for k in range(num_classes - 1):\n        in_task = torch.ones_like(targets, dtype=torch.bool) if k == 0 else targets > (k - 1)\n        if in_task.sum() == 0:\n            continue\n        task_targets = (targets[in_task] > k).float()\n        total_loss = total_loss + F.binary_cross_entropy_with_logits(\n            logits[in_task, k].float(), task_targets, reduction='sum')\n        term_count += int(in_task.sum().item())\n    return total_loss / max(term_count, 1)\n\n\ndef corn_cumulative_probs(logits):\n    # probability that grade > k, for k = 0..3\n    return torch.cumprod(torch.sigmoid(logits), dim=1)\n\n\ndef corn_predict(logits, threshold=0.5):\n    # predicted grade = number of \"yes\" answers\n    return (corn_cumulative_probs(logits) > threshold).sum(dim=1)\n\n\ndef corn_expected_grade(logits):\n    # continuous severity score between 0 and 4\n    return corn_cumulative_probs(logits).sum(dim=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:38.147819Z","iopub.execute_input":"2026-09-21T08:18:38.14819Z","iopub.status.idle":"2026-09-21T08:18:38.155418Z","shell.execute_reply.started":"2026-09-21T08:18:38.148157Z","shell.execute_reply":"2026-09-21T08:18:38.154802Z"}},"outputs":[],"execution_count":null},{"id":"d541f039-be10-48e1-b6a3-a7987c5d1e6f","cell_type":"markdown","source":"### 4.2 Multi-task loss (grade + image quality)\nTotal loss = grade loss (gradable images only) + 0.3 x quality loss (class-weighted cross-entropy).","metadata":{}},{"id":"f47dca83-3df4-4d6f-bc56-b6ba60ac1fe7","cell_type":"code","source":"class MultiTaskLoss(nn.Module):\n\n    def __init__(self, num_classes=5, quality_weight=0.3, quality_class_weights=None,\n                 label_smoothing=0.05):\n        super().__init__()\n        self.num_classes = num_classes\n        self.quality_weight = quality_weight\n        self.label_smoothing = label_smoothing\n        if quality_class_weights is None:\n            quality_class_weights = torch.ones(2)\n        self.register_buffer('quality_class_weights', quality_class_weights.float())\n\n    def forward(self, grade_logits, quality_logits, grade_targets, quality_targets):\n        # ungradable images (grade -1) are left out of the grade loss\n        gradable = grade_targets >= 0\n        if gradable.any():\n            grade_loss = corn_loss(grade_logits[gradable], grade_targets[gradable], self.num_classes)\n        else:\n            grade_loss = torch.zeros((), device=grade_logits.device)\n\n        has_quality_label = quality_targets >= 0\n        if has_quality_label.any():\n            quality_loss = F.cross_entropy(quality_logits[has_quality_label].float(),\n                                           quality_targets[has_quality_label],\n                                           weight=self.quality_class_weights,\n                                           label_smoothing=self.label_smoothing)\n        else:\n            quality_loss = torch.zeros((), device=grade_logits.device)\n\n        total_loss = grade_loss + self.quality_weight * quality_loss\n        return total_loss, {'loss': float(total_loss.detach()),\n                            'loss_grade': float(grade_loss.detach()),\n                            'loss_quality': float(quality_loss.detach())}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:38.156302Z","iopub.execute_input":"2026-09-21T08:18:38.156492Z","iopub.status.idle":"2026-09-21T08:18:38.166476Z","shell.execute_reply.started":"2026-09-21T08:18:38.156472Z","shell.execute_reply":"2026-09-21T08:18:38.165906Z"}},"outputs":[],"execution_count":null},{"id":"5049b025-51d9-46a2-8378-7d96ada1fd7e","cell_type":"markdown","source":"### 4.3 Network definition","metadata":{}},{"id":"49d2148d-b88a-4208-b16d-cfeba50551c5","cell_type":"code","source":"class DRTriageNet(nn.Module):\n\n    def __init__(self, backbone_name, num_classes=5, num_quality_classes=2,\n                 dropout=0.4, pretrained=True):\n        super().__init__()\n        self.num_classes = num_classes\n        # pretrained CNN without its ImageNet classifier -> one feature vector per image\n        self.backbone = timm.create_model(backbone_name, pretrained=pretrained, num_classes=0)\n        feature_size = self.backbone.num_features\n        self.dropout = nn.Dropout(p=dropout)\n        self.grade_head = nn.Linear(feature_size, num_classes - 1)\n        self.quality_head = nn.Linear(feature_size, num_quality_classes)\n        for head in (self.grade_head, self.quality_head):\n            nn.init.trunc_normal_(head.weight, std=0.01)\n            nn.init.zeros_(head.bias)\n\n    def forward_features(self, images):\n        return self.backbone(images)\n\n    def heads_from_features(self, features):\n        dropped = self.dropout(features)\n        return self.grade_head(dropped), self.quality_head(dropped)\n\n    def forward(self, images):\n        return self.heads_from_features(self.backbone(images))\n\n    def freeze_backbone(self):\n        for parameter in self.backbone.parameters():\n            parameter.requires_grad = False\n\n    def unfreeze_backbone(self):\n        for parameter in self.backbone.parameters():\n            parameter.requires_grad = True\n\n    def param_groups(self, head_lr=1e-3, backbone_lr=1e-4):\n        # smaller learning rate for the pretrained layers\n        head_parameters = list(self.grade_head.parameters()) + list(self.quality_head.parameters())\n        return [{'params': list(self.backbone.parameters()), 'lr': backbone_lr, 'name': 'backbone'},\n                {'params': head_parameters, 'lr': head_lr, 'name': 'heads'}]\n\n\ndef build_model(settings, pretrained=None):\n    return DRTriageNet(settings.backbone, num_classes=settings.num_classes,\n                       num_quality_classes=settings.num_quality_classes,\n                       dropout=settings.dropout,\n                       pretrained=settings.pretrained if pretrained is None else pretrained)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:38.16732Z","iopub.execute_input":"2026-09-21T08:18:38.167619Z","iopub.status.idle":"2026-09-21T08:18:38.185614Z","shell.execute_reply.started":"2026-09-21T08:18:38.167588Z","shell.execute_reply":"2026-09-21T08:18:38.18495Z"}},"outputs":[],"execution_count":null},{"id":"8c35152e-e656-441b-8a9d-b1e54785b02c","cell_type":"markdown","source":"### 4.4 Build the model with pretrained weights","metadata":{}},{"id":"b9e35fa5-eb69-4382-9321-ea9a611f6877","cell_type":"code","source":"model = build_model(config).to(config.device)\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint('backbone          :', config.backbone)\nprint(f'total parameters  : {total_params / 1e6:.1f}M')\nprint(f'trainable         : {trainable_params / 1e6:.1f}M')\nprint(f'feature size      : {model.backbone.num_features}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:38.186596Z","iopub.execute_input":"2026-09-21T08:18:38.186884Z","iopub.status.idle":"2026-09-21T08:18:41.372272Z","shell.execute_reply.started":"2026-09-21T08:18:38.18685Z","shell.execute_reply":"2026-09-21T08:18:41.371675Z"}},"outputs":[],"execution_count":null},{"id":"7381e3fc-953d-4c87-b60b-1c767ae4a438","cell_type":"markdown","source":"### 4.5 Output shape check","metadata":{}},{"id":"7483f6fb-be66-4cb3-a864-351b9450f927","cell_type":"code","source":"model.eval()\nwith torch.no_grad():\n    grade_logits, quality_logits = model(batch_images[:4].to(config.device))\nprint('grade head output  :', tuple(grade_logits.shape), '-> 4 ordinal yes/no tasks')\nprint('quality head output:', tuple(quality_logits.shape), '-> gradable / ungradable')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:41.373157Z","iopub.execute_input":"2026-09-21T08:18:41.373743Z","iopub.status.idle":"2026-09-21T08:18:42.722651Z","shell.execute_reply.started":"2026-09-21T08:18:41.373719Z","shell.execute_reply":"2026-09-21T08:18:42.721944Z"}},"outputs":[],"execution_count":null},{"id":"4bca661b-51c0-45c7-8d0d-6da2a73f5abf","cell_type":"markdown","source":"### 4.6 Parameters per part of the network","metadata":{}},{"id":"eba357a4-4a27-40ee-b9ae-5c5c0ad2ae03","cell_type":"code","source":"print(f'shared backbone : {sum(p.numel() for p in model.backbone.parameters()) / 1e6:.1f}M parameters')\nprint(f'grade head      : {sum(p.numel() for p in model.grade_head.parameters()):,} parameters')\nprint(f'quality head    : {sum(p.numel() for p in model.quality_head.parameters()):,} parameters')\nprint(f'quality loss weight: {config.quality_loss_weight}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.723458Z","iopub.execute_input":"2026-09-21T08:18:42.723829Z","iopub.status.idle":"2026-09-21T08:18:42.730502Z","shell.execute_reply.started":"2026-09-21T08:18:42.723804Z","shell.execute_reply":"2026-09-21T08:18:42.729969Z"}},"outputs":[],"execution_count":null},{"id":"a3fedcf4-735c-4cd5-abdc-a7e88341d64a","cell_type":"markdown","source":"### 4.7 The ordinal loss grows with distance from the true grade\nPlain cross-entropy would give every wrong answer the same loss; CORN penalises bigger mistakes more.","metadata":{}},{"id":"527a332a-a78d-46d4-a925-67b836b708d5","cell_type":"code","source":"true_grades = torch.tensor([4, 4, 4])\n\ndef corn_loss_for_prediction(predicted_grade):\n    logits = torch.full((3, 4), -4.0)\n    logits[:, :predicted_grade] = 4.0\n    return float(corn_loss(logits, true_grades, config.num_classes))\n\nfor predicted_grade in [4, 3, 2, 1, 0]:\n    print(f'true grade 4, predicted {predicted_grade}: CORN loss {corn_loss_for_prediction(predicted_grade):.3f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.731593Z","iopub.execute_input":"2026-09-21T08:18:42.731973Z","iopub.status.idle":"2026-09-21T08:18:42.7626Z","shell.execute_reply.started":"2026-09-21T08:18:42.731948Z","shell.execute_reply":"2026-09-21T08:18:42.761829Z"}},"outputs":[],"execution_count":null},{"id":"2e8697cb-9c25-48fb-95ce-c63ef199bb88","cell_type":"markdown","source":"### 4.8 Probabilities are always in grade order\nP(grade > k) never increases with k, so the predicted stages are always consistent.","metadata":{}},{"id":"5acc863b-ad2a-4342-8c37-eba4eb3211f0","cell_type":"code","source":"with torch.no_grad():\n    cumulative_probs = corn_cumulative_probs(grade_logits.float())\norder_is_consistent = bool(((cumulative_probs[:, 1:] - cumulative_probs[:, :-1]) <= 1e-6).all())\nprint('P(grade > k) for 4 images (untrained heads):')\nprint(cumulative_probs.cpu().numpy().round(3))\nprint('probabilities non-increasing:', order_is_consistent)\nprint('expected grades:', corn_expected_grade(grade_logits.float()).cpu().numpy().round(2))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.763484Z","iopub.execute_input":"2026-09-21T08:18:42.763824Z","iopub.status.idle":"2026-09-21T08:18:42.868311Z","shell.execute_reply.started":"2026-09-21T08:18:42.763802Z","shell.execute_reply":"2026-09-21T08:18:42.867584Z"}},"outputs":[],"execution_count":null},{"id":"cb998247-cbb1-4aa3-b4b4-636216f6c86b","cell_type":"markdown","source":"### 4.9 Learning rate per parameter group","metadata":{}},{"id":"b4672b7a-bea0-4103-871d-b1513113d804","cell_type":"code","source":"for group in model.param_groups(config.head_lr, config.backbone_lr):\n    print(f\"{group['name']:10s} learning rate {group['lr']:.0e} | \"\n          f\"{sum(p.numel() for p in group['params']) / 1e6:.1f}M parameters\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.869193Z","iopub.execute_input":"2026-09-21T08:18:42.869515Z","iopub.status.idle":"2026-09-21T08:18:42.875741Z","shell.execute_reply.started":"2026-09-21T08:18:42.869474Z","shell.execute_reply":"2026-09-21T08:18:42.875041Z"}},"outputs":[],"execution_count":null},{"id":"c82e1350-27bd-4987-b149-b0c94a0836aa","cell_type":"markdown","source":"## 5. Training\n\n| Technique | Purpose |\n|---|---|\n| Frozen backbone in epoch 0 | protects pretrained features while the new heads start learning |\n| Different learning rates (1e-4 backbone, 1e-3 heads) | gentle fine-tuning of pretrained layers |\n| Warmup + cosine learning-rate schedule | stable start, smooth decay |\n| Mixed precision (AMP) | about half the memory, faster training |\n| Weight moving average (EMA) | smoother, more stable final weights |\n| Gradient clipping | prevents unstable updates |\n| Validation every epoch on QWK | measures ordinal grading quality |\n| Early stopping (patience 4) | stops when validation QWK stops improving, limits overfitting |\n| Checkpoints every epoch | training resumes automatically after a disconnect |\n\n**Quadratic weighted kappa (QWK)** measures agreement between predicted and true grades, penalising large grade errors more than small ones. It is the model-selection metric.\n\n### 5.1 Quadratic weighted kappa","metadata":{}},{"id":"7babf999-c9d8-46e3-ab03-7765a9f27bfc","cell_type":"code","source":"from sklearn.metrics import cohen_kappa_score\n\ndef quadratic_weighted_kappa(true_grades, predicted_grades, num_classes=5):\n    return float(cohen_kappa_score(true_grades, predicted_grades, weights='quadratic',\n                                   labels=list(range(num_classes))))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.876679Z","iopub.execute_input":"2026-09-21T08:18:42.87704Z","iopub.status.idle":"2026-09-21T08:18:42.891463Z","shell.execute_reply.started":"2026-09-21T08:18:42.877017Z","shell.execute_reply":"2026-09-21T08:18:42.890643Z"}},"outputs":[],"execution_count":null},{"id":"97d97617-6153-454a-8ca3-5adc51e39384","cell_type":"markdown","source":"### 5.2 Trainer\nRuns the training and validation loops, records loss, accuracy and QWK per epoch, saves `last.pt` and `best.pt`, and applies early stopping.","metadata":{}},{"id":"e558ec92-7d9f-45c6-bc3d-4d3442b0836a","cell_type":"code","source":"def make_grad_scaler(enabled):\n    if hasattr(torch.amp, 'GradScaler'):\n        return torch.amp.GradScaler('cuda', enabled=enabled)\n    return torch.cuda.amp.GradScaler(enabled=enabled)\n\n\nclass WeightAverage:\n    # exponential moving average (EMA) of the model weights\n\n    def __init__(self, model, decay=0.999):\n        self.decay = decay\n        self.shadow = copy.deepcopy(model).eval()\n        for parameter in self.shadow.parameters():\n            parameter.requires_grad_(False)\n\n    @torch.no_grad()\n    def update(self, model):\n        for average, current in zip(self.shadow.parameters(), model.parameters()):\n            average.mul_(self.decay).add_(current.detach(), alpha=1.0 - self.decay)\n        for average_buffer, current_buffer in zip(self.shadow.buffers(), model.buffers()):\n            average_buffer.copy_(current_buffer)\n\n\ndef warmup_cosine_schedule(warmup_steps, total_steps):\n    def learning_rate_factor(step):\n        if step < warmup_steps:\n            return (step + 1) / max(warmup_steps, 1)\n        progress = (step - warmup_steps) / max(total_steps - warmup_steps, 1)\n        return 0.5 * (1.0 + np.cos(np.pi * min(progress, 1.0)))\n    return learning_rate_factor\n\n\nclass Trainer:\n\n    def __init__(self, model, settings, train_loader, validation_loader, quality_class_weights=None):\n        self.model = model.to(settings.device)\n        self.settings = settings\n        self.train_loader = train_loader\n        self.validation_loader = validation_loader\n        self.device = settings.device\n        self.device_type = 'cuda' if settings.device.startswith('cuda') else 'cpu'\n\n        self.loss_fn = MultiTaskLoss(settings.num_classes, settings.quality_loss_weight,\n                                     quality_class_weights, settings.label_smoothing).to(settings.device)\n        self.optimizer = torch.optim.AdamW(model.param_groups(settings.head_lr, settings.backbone_lr),\n                                           weight_decay=settings.weight_decay)\n        total_steps = settings.epochs * max(len(train_loader), 1)\n        self.scheduler = torch.optim.lr_scheduler.LambdaLR(\n            self.optimizer, warmup_cosine_schedule(settings.warmup_steps, total_steps))\n        self.grad_scaler = make_grad_scaler(settings.amp)\n        self.weight_average = WeightAverage(model, settings.ema_decay) if settings.use_ema else None\n\n        self.start_epoch = 0\n        self.best_qwk = -1.0\n        self.epochs_without_improvement = 0\n        self.history = []\n\n        self.checkpoint_folder = Path(settings.checkpoint_dir)\n        self.checkpoint_folder.mkdir(parents=True, exist_ok=True)\n        self.last_checkpoint = self.checkpoint_folder / 'last.pt'\n        self.best_checkpoint = self.checkpoint_folder / 'best.pt'\n\n    def save_checkpoint(self, epoch, is_best=False):\n        state = {\n            'epoch': epoch,\n            'model': self.model.state_dict(),\n            'optimizer': self.optimizer.state_dict(),\n            'scheduler': self.scheduler.state_dict(),\n            'grad_scaler': self.grad_scaler.state_dict(),\n            'ema': self.weight_average.shadow.state_dict() if self.weight_average else None,\n            'best_qwk': self.best_qwk,\n            'epochs_without_improvement': self.epochs_without_improvement,\n            'history': self.history,\n            'torch_rng': torch.get_rng_state(),\n            'numpy_rng': np.random.get_state(),\n            'config': self.settings.to_dict(),\n        }\n        # write to a temp file then rename, so a crash cannot corrupt last.pt\n        temp_file = self.last_checkpoint.with_suffix('.tmp')\n        torch.save(state, temp_file)\n        os.replace(temp_file, self.last_checkpoint)\n        if is_best:\n            torch.save(state, self.best_checkpoint)\n        (self.checkpoint_folder / 'history.json').write_text(json.dumps(self.history, indent=2))\n\n    def resume_if_possible(self):\n        if not self.last_checkpoint.exists():\n            print('[resume] no checkpoint found - starting from scratch')\n            return False\n        state = torch.load(self.last_checkpoint, map_location=self.device, weights_only=False)\n        self.model.load_state_dict(state['model'])\n        self.optimizer.load_state_dict(state['optimizer'])\n        self.scheduler.load_state_dict(state['scheduler'])\n        self.grad_scaler.load_state_dict(state['grad_scaler'])\n        if self.weight_average and state.get('ema'):\n            self.weight_average.shadow.load_state_dict(state['ema'])\n        self.best_qwk = state['best_qwk']\n        self.epochs_without_improvement = state['epochs_without_improvement']\n        self.history = state['history']\n        self.start_epoch = state['epoch'] + 1\n        torch.set_rng_state(state['torch_rng'].cpu())\n        np.random.set_state(state['numpy_rng'])\n        print(f\"[resume] restored from epoch {state['epoch']}, best QWK so far {self.best_qwk:.4f} \"\n              f\"- continuing at epoch {self.start_epoch}\")\n        return True\n\n    def train_one_epoch(self):\n        self.model.train()\n        loss_sums, batch_count = {}, 0\n        correct, graded_total = 0, 0\n        start_time = time.time()\n        for images, grade_targets, quality_targets in tqdm(self.train_loader, leave=False):\n            images = images.to(self.device, non_blocking=True)\n            grade_targets = grade_targets.to(self.device, non_blocking=True)\n            quality_targets = quality_targets.to(self.device, non_blocking=True)\n\n            self.optimizer.zero_grad(set_to_none=True)\n            with torch.autocast(device_type=self.device_type, enabled=self.settings.amp):\n                grade_logits, quality_logits = self.model(images)\n            grade_logits, quality_logits = grade_logits.float(), quality_logits.float()\n            loss, loss_parts = self.loss_fn(grade_logits, quality_logits, grade_targets, quality_targets)\n\n            self.grad_scaler.scale(loss).backward()\n            self.grad_scaler.unscale_(self.optimizer)\n            torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.settings.grad_clip)\n            self.grad_scaler.step(self.optimizer)\n            self.grad_scaler.update()\n            self.scheduler.step()\n            if self.weight_average:\n                self.weight_average.update(self.model)\n\n            # training accuracy on gradable images\n            gradable = grade_targets >= 0\n            with torch.no_grad():\n                predicted = corn_predict(grade_logits.detach())\n            correct += int((predicted[gradable] == grade_targets[gradable]).sum())\n            graded_total += int(gradable.sum())\n\n            for name, value in loss_parts.items():\n                loss_sums[name] = loss_sums.get(name, 0.0) + value\n            batch_count += 1\n\n        result = {f'train_{name}': total / max(batch_count, 1) for name, total in loss_sums.items()}\n        result['train_accuracy'] = correct / max(graded_total, 1)\n        result['epoch_seconds'] = time.time() - start_time\n        result['lr'] = self.optimizer.param_groups[0]['lr']\n        return result\n\n    @torch.no_grad()\n    def validate(self, use_average=True):\n        model = self.weight_average.shadow if (self.weight_average and use_average) else self.model\n        model.eval()\n        predictions, targets, losses = [], [], []\n        for images, grade_targets, quality_targets in self.validation_loader:\n            images = images.to(self.device, non_blocking=True)\n            grade_targets = grade_targets.to(self.device, non_blocking=True)\n            quality_targets = quality_targets.to(self.device, non_blocking=True)\n            with torch.autocast(device_type=self.device_type, enabled=self.settings.amp):\n                grade_logits, quality_logits = model(images)\n            grade_logits, quality_logits = grade_logits.float(), quality_logits.float()\n            loss, _ = self.loss_fn(grade_logits, quality_logits, grade_targets, quality_targets)\n            losses.append(float(loss))\n            predictions.append(corn_predict(grade_logits).cpu().numpy())\n            targets.append(grade_targets.cpu().numpy())\n\n        predicted_grades, true_grades = np.concatenate(predictions), np.concatenate(targets)\n        return {'val_loss': float(np.mean(losses)),\n                'val_accuracy': float((predicted_grades == true_grades).mean()),\n                'val_qwk': quadratic_weighted_kappa(true_grades, predicted_grades)}\n\n    def fit(self):\n        # safe to re-run: continues from last.pt\n        self.resume_if_possible()\n        if self.epochs_without_improvement >= self.settings.patience:\n            print('[resume] training already finished (early stopped)')\n            return self.history\n\n        for epoch in range(self.start_epoch, self.settings.epochs):\n            if epoch == 0 and self.settings.freeze_first_epoch:\n                self.model.freeze_backbone()\n                print('[stage 1] backbone frozen - training the new heads only')\n            elif epoch == 1 or (epoch == 0 and not self.settings.freeze_first_epoch):\n                self.model.unfreeze_backbone()\n                print('[stage 2] backbone unfrozen - fine-tuning the whole network')\n\n            record = {'epoch': epoch, **self.train_one_epoch(), **self.validate()}\n            self.history.append(record)\n            print(f\"epoch {epoch:02d} | train loss {record['train_loss']:.4f} | \"\n                  f\"train acc {record['train_accuracy']:.3f} | val loss {record['val_loss']:.4f} | \"\n                  f\"val acc {record['val_accuracy']:.3f} | val QWK {record['val_qwk']:.4f} | \"\n                  f\"{record['epoch_seconds']:.0f}s\")\n\n            improved = record['val_qwk'] > self.best_qwk + self.settings.min_delta\n            if improved:\n                self.best_qwk = record['val_qwk']\n                self.epochs_without_improvement = 0\n                print(f'           new best validation QWK: {self.best_qwk:.4f} (saved best.pt)')\n            else:\n                self.epochs_without_improvement += 1\n            self.save_checkpoint(epoch, is_best=improved)\n\n            if self.epochs_without_improvement >= self.settings.patience:\n                print(f'[early stop] no QWK improvement for {self.settings.patience} epochs')\n                break\n        return self.history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.892511Z","iopub.execute_input":"2026-09-21T08:18:42.892831Z","iopub.status.idle":"2026-09-21T08:18:42.922707Z","shell.execute_reply.started":"2026-09-21T08:18:42.8928Z","shell.execute_reply":"2026-09-21T08:18:42.92197Z"}},"outputs":[],"execution_count":null},{"id":"d6bf889b-7759-406c-a9b2-33bd41aef3a2","cell_type":"markdown","source":"### 5.3 Create the trainer\nThe quality loss is class-weighted because ungradable images are much rarer than gradable ones.","metadata":{}},{"id":"8dd391dc-ace1-41d7-a2c9-74e706f51201","cell_type":"code","source":"quality_counts = np.bincount(train_set['quality'].values, minlength=2)\ntrainer = Trainer(model, config, train_loader, validation_loader,\n                  quality_class_weights=effective_number_weights(quality_counts))\nprint(f'quality labels in training: {quality_counts[0]} gradable, {quality_counts[1]} ungradable')\nprint(f'checkpoints saved to      : {config.checkpoint_dir}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:42.923505Z","iopub.execute_input":"2026-09-21T08:18:42.924188Z","iopub.status.idle":"2026-09-21T08:18:43.057088Z","shell.execute_reply.started":"2026-09-21T08:18:42.924158Z","shell.execute_reply":"2026-09-21T08:18:43.056255Z"}},"outputs":[],"execution_count":null},{"id":"810aa4bc-d44c-49a8-a570-46f0cc5bd724","cell_type":"markdown","source":"### 5.4 One-batch overfit test\nA quick check before the long run: a fresh model trained on only 8 images should drive its loss close to zero. If it cannot, something in the data, model or loss is broken.","metadata":{}},{"id":"b4fffa4e-3203-4bd1-8444-b2074a552792","cell_type":"code","source":"probe_model = build_model(config).to(config.device)\nprobe_model.train()\nprobe_loss_fn = MultiTaskLoss(config.num_classes, config.quality_loss_weight).to(config.device)\nprobe_optimizer = torch.optim.AdamW(probe_model.parameters(), lr=3e-4)\n\nprobe_images = batch_images[:8].to(config.device)\nprobe_grades = batch_grades[:8].to(config.device)\nprobe_quality = batch_quality[:8].to(config.device)\n\nfor step in range(30):\n    probe_optimizer.zero_grad()\n    loss, loss_parts = probe_loss_fn(*probe_model(probe_images), probe_grades, probe_quality)\n    loss.backward()\n    probe_optimizer.step()\n    if step % 10 == 0 or step == 29:\n        print(f\"step {step:2d} | total loss {loss_parts['loss']:.4f} | grade loss {loss_parts['loss_grade']:.4f}\")\n\ndel probe_model, probe_optimizer\nif config.device == 'cuda':\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:43.058133Z","iopub.execute_input":"2026-09-21T08:18:43.05848Z","iopub.status.idle":"2026-09-21T08:18:53.529924Z","shell.execute_reply.started":"2026-09-21T08:18:43.058437Z","shell.execute_reply":"2026-09-21T08:18:53.52909Z"}},"outputs":[],"execution_count":null},{"id":"ca6bc689-6fdf-43f5-aedd-088bc9e9f9c4","cell_type":"markdown","source":"### 5.5 Train the model\nPrints training and validation loss, accuracy and QWK for every epoch. If the session disconnects, running this cell again continues from the last saved epoch.","metadata":{}},{"id":"2b32eb11-62d0-45eb-8668-5f7f4bebf2fa","cell_type":"code","source":"seed_everything()\ntraining_history = trainer.fit()\nprint(f'\\nbest validation QWK: {trainer.best_qwk:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T08:18:53.530984Z","iopub.execute_input":"2026-09-21T08:18:53.53168Z","iopub.status.idle":"2026-09-21T09:02:52.647851Z","shell.execute_reply.started":"2026-09-21T08:18:53.531649Z","shell.execute_reply":"2026-09-21T09:02:52.646983Z"}},"outputs":[],"execution_count":null},{"id":"f6b777ac-8d45-49da-8c3e-8eaa276253ce","cell_type":"markdown","source":"### 5.6 Accuracy and loss curves\nTraining vs validation **loss** and **accuracy**, validation QWK, and the learning-rate schedule. Training accuracy is measured on augmented, class-balanced batches, so it is not directly comparable with validation accuracy early in training; the trend over epochs is what matters.","metadata":{}},{"id":"ed55c1d8-9a02-4a61-94ee-5f6a8250b227","cell_type":"code","source":"history_table = pd.DataFrame(training_history)\nhistory_table.to_csv('results/logs/history.csv', index=False)\nshow_table('Training history',\n           history_table[['epoch', 'train_loss', 'val_loss', 'train_accuracy', 'val_accuracy', 'val_qwk', 'lr']].round(4))\n\nfig, axes = plt.subplots(2, 2, figsize=(12, 8))\naxes[0, 0].plot(history_table.epoch, history_table.train_loss, 'o-', ms=3, label='train', color='#7F77DD')\naxes[0, 0].plot(history_table.epoch, history_table.val_loss, 'o-', ms=3, label='validation', color='#E8734A')\naxes[0, 0].set_title('Loss'); axes[0, 0].legend()\n\naxes[0, 1].plot(history_table.epoch, history_table.train_accuracy, 'o-', ms=3, label='train', color='#7F77DD')\naxes[0, 1].plot(history_table.epoch, history_table.val_accuracy, 'o-', ms=3, label='validation', color='#E8734A')\naxes[0, 1].set_title('Accuracy'); axes[0, 1].legend()\n\naxes[1, 0].plot(history_table.epoch, history_table.val_qwk, 'o-', ms=3, color='#1D9E75')\naxes[1, 0].axhline(history_table.val_qwk.max(), ls='--', lw=0.8, color='grey')\naxes[1, 0].set_title(f'Validation QWK (best {history_table.val_qwk.max():.3f})')\n\naxes[1, 1].plot(history_table.epoch, history_table.lr, color='#333')\naxes[1, 1].set_title('Backbone learning rate')\nfor ax in axes.ravel():\n    ax.set_xlabel('epoch')\nsave_figure('results/figures/fig9_training_curves.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:52.649697Z","iopub.execute_input":"2026-09-21T09:02:52.650032Z","iopub.status.idle":"2026-09-21T09:02:53.875021Z","shell.execute_reply.started":"2026-09-21T09:02:52.649992Z","shell.execute_reply":"2026-09-21T09:02:53.874343Z"}},"outputs":[],"execution_count":null},{"id":"c3ae3c37-4d06-441f-a024-d53adea0f7d9","cell_type":"markdown","source":"### 5.7 Overfitting check\nA growing gap between training and validation loss, while validation QWK stops improving, would indicate overfitting.","metadata":{}},{"id":"b4cd71ef-f4c3-46e9-9483-c2553da7a6bc","cell_type":"code","source":"loss_gap = history_table.val_loss - history_table.train_loss\nbest_epoch = int(history_table.loc[history_table.val_qwk.idxmax(), 'epoch'])\nprint(f'validation - training loss gap: first epoch {loss_gap.iloc[0]:.3f} -> last epoch {loss_gap.iloc[-1]:.3f}')\nprint(f'best epoch (highest validation QWK): {best_epoch}')\nif len(history_table) >= 3:\n    print(f'QWK still improving in final 3 epochs: {bool(history_table.val_qwk.iloc[-1] >= history_table.val_qwk.iloc[-3])}')\nstopped_early = len(history_table) < config.epochs\nprint(f'epochs run: {len(history_table)} / {config.epochs} ({\"early stopped\" if stopped_early else \"completed\"})')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:53.876102Z","iopub.execute_input":"2026-09-21T09:02:53.876493Z","iopub.status.idle":"2026-09-21T09:02:53.884075Z","shell.execute_reply.started":"2026-09-21T09:02:53.876466Z","shell.execute_reply":"2026-09-21T09:02:53.883421Z"}},"outputs":[],"execution_count":null},{"id":"64f45512-da1d-4f4e-a640-711b3abc2dce","cell_type":"markdown","source":"## 6. Evaluation\n\nThe best checkpoint is tested on the held-out DDR test split and on the external APTOS set.\n\n- **Stage classification (5 classes)**: accuracy, precision, recall and F1-score per stage, macro and weighted averages, confusion matrix, QWK\n- **DR detection** (any DR vs No DR) and **referable DR** (grade 2 or higher vs lower): accuracy, precision, recall (sensitivity), specificity, F1, ROC AUC\n- **Uncertainty**: each image is predicted 20 times with dropout switched on (MC dropout); the spread of these predictions is the model's uncertainty. Calibration and risk-coverage show whether that uncertainty is useful.\n- **Triage policy**: ungradable -> re-capture; uncertain -> human grader; grade 2 or higher -> refer; otherwise routine rescreen.\n\n### 6.1 Metric functions","metadata":{}},{"id":"b32c3cdc-e25d-4652-bedd-2c1d4e7954ba","cell_type":"code","source":"from sklearn.metrics import (accuracy_score, classification_report, confusion_matrix, f1_score,\n                             precision_score, recall_score, roc_auc_score, roc_curve)\n\n_trapezoid = getattr(np, 'trapezoid', None) or np.trapz\n\n\ndef per_class_report(true_grades, predicted_grades, num_classes=5):\n    labels = list(range(num_classes))\n    return {\n        'accuracy': float(accuracy_score(true_grades, predicted_grades)),\n        'qwk': quadratic_weighted_kappa(true_grades, predicted_grades, num_classes),\n        'macro_precision': float(precision_score(true_grades, predicted_grades, average='macro', labels=labels, zero_division=0)),\n        'macro_recall': float(recall_score(true_grades, predicted_grades, average='macro', labels=labels, zero_division=0)),\n        'macro_f1': float(f1_score(true_grades, predicted_grades, average='macro', labels=labels, zero_division=0)),\n        'weighted_f1': float(f1_score(true_grades, predicted_grades, average='weighted', labels=labels, zero_division=0)),\n        'precision_per_class': precision_score(true_grades, predicted_grades, average=None, labels=labels, zero_division=0).tolist(),\n        'recall_per_class': recall_score(true_grades, predicted_grades, average=None, labels=labels, zero_division=0).tolist(),\n        'f1_per_class': f1_score(true_grades, predicted_grades, average=None, labels=labels, zero_division=0).tolist(),\n        'confusion_matrix': confusion_matrix(true_grades, predicted_grades, labels=labels).tolist(),\n    }\n\n\ndef binary_report(true_positive, predicted_positive):\n    # accuracy, precision, recall (sensitivity), specificity, F1 for a yes/no decision\n    true_positive = np.asarray(true_positive).astype(int)\n    predicted_positive = np.asarray(predicted_positive).astype(int)\n    tn, fp, fn, tp = confusion_matrix(true_positive, predicted_positive, labels=[0, 1]).ravel()\n    return {'accuracy': float((tp + tn) / max(tp + tn + fp + fn, 1)),\n            'precision': float(tp / max(tp + fp, 1)),\n            'recall_sensitivity': float(tp / max(tp + fn, 1)),\n            'specificity': float(tn / max(tn + fp, 1)),\n            'f1': float(2 * tp / max(2 * tp + fp + fn, 1)),\n            'tp': int(tp), 'fp': int(fp), 'fn': int(fn), 'tn': int(tn)}\n\n\ndef sensitivity_at_specificity(true_positive, scores, target_specificity=0.95):\n    false_positive_rate, true_positive_rate, thresholds = roc_curve(true_positive, scores)\n    specificity = 1.0 - false_positive_rate\n    feasible = np.where(specificity >= target_specificity)[0]\n    if feasible.size == 0:\n        return {'specificity': float(specificity.max()), 'sensitivity': 0.0, 'threshold': float(thresholds[0])}\n    best = feasible[np.argmax(true_positive_rate[feasible])]\n    return {'specificity': float(specificity[best]), 'sensitivity': float(true_positive_rate[best]),\n            'threshold': float(thresholds[best])}\n\n\ndef referable_dr_metrics(true_grades, expected_grades, threshold=2):\n    # referable DR = grade 2 or higher\n    is_referable = (np.asarray(true_grades) >= threshold).astype(int)\n    scores = np.asarray(expected_grades, dtype=float)\n    result = {'auc': float(roc_auc_score(is_referable, scores)), 'prevalence': float(is_referable.mean())}\n    for specificity in (0.90, 0.95):\n        result[f'sens_at_spec_{int(specificity * 100)}'] = sensitivity_at_specificity(is_referable, scores, specificity)\n    return result\n\n\ndef expected_calibration_error(confidences, is_correct, n_bins=15):\n    confidences = np.asarray(confidences, dtype=float)\n    is_correct = np.asarray(is_correct, dtype=float)\n    edges = np.linspace(0.0, 1.0, n_bins + 1)\n    ece, mce, bins = 0.0, 0.0, []\n    for low, high in zip(edges[:-1], edges[1:]):\n        in_bin = (confidences > low) & (confidences <= high)\n        if not in_bin.any():\n            continue\n        bin_accuracy, bin_confidence = is_correct[in_bin].mean(), confidences[in_bin].mean()\n        gap = abs(bin_accuracy - bin_confidence)\n        ece += in_bin.mean() * gap\n        mce = max(mce, gap)\n        bins.append({'low': float(low), 'high': float(high), 'images': int(in_bin.sum()),\n                     'accuracy': float(bin_accuracy), 'confidence': float(bin_confidence)})\n    return {'ece': float(ece), 'mce': float(mce), 'bins': bins}\n\n\ndef risk_coverage_curve(true_grades, predicted_grades, uncertainty, n_points=21):\n    # keep the most certain images, defer the rest to a human\n    true_grades, predicted_grades = np.asarray(true_grades), np.asarray(predicted_grades)\n    most_certain_first = np.argsort(np.asarray(uncertainty, dtype=float))\n    total = len(true_grades)\n    points = []\n    for coverage in np.linspace(1.0, 0.1, n_points):\n        kept_count = max(int(round(coverage * total)), 2)\n        kept = most_certain_first[:kept_count]\n        accuracy = float(accuracy_score(true_grades[kept], predicted_grades[kept]))\n        points.append({'coverage': float(kept_count / total), 'accuracy': accuracy, 'risk': 1.0 - accuracy,\n                       'qwk': quadratic_weighted_kappa(true_grades[kept], predicted_grades[kept]),\n                       'images_kept': kept_count})\n    coverages = np.array([p['coverage'] for p in points])\n    risks = np.array([p['risk'] for p in points])\n    return {'points': points, 'aurc': float(_trapezoid(risks[::-1], coverages[::-1]))}\n\n\ndef full_evaluation(true_grades, predicted_grades, expected_grades, uncertainty=None, confidences=None):\n    results = per_class_report(true_grades, predicted_grades)\n    results['dr_detection'] = binary_report(np.asarray(true_grades) >= 1, np.asarray(predicted_grades) >= 1)\n    results['referable_decision'] = binary_report(np.asarray(true_grades) >= 2, np.asarray(predicted_grades) >= 2)\n    results['referable'] = referable_dr_metrics(true_grades, expected_grades)\n    if confidences is not None:\n        is_correct = (np.asarray(true_grades) == np.asarray(predicted_grades)).astype(float)\n        results['calibration'] = expected_calibration_error(confidences, is_correct)\n    if uncertainty is not None:\n        results['selective'] = risk_coverage_curve(true_grades, predicted_grades, uncertainty)\n    return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:53.885182Z","iopub.execute_input":"2026-09-21T09:02:53.886028Z","iopub.status.idle":"2026-09-21T09:02:53.910732Z","shell.execute_reply.started":"2026-09-21T09:02:53.885992Z","shell.execute_reply":"2026-09-21T09:02:53.90984Z"}},"outputs":[],"execution_count":null},{"id":"a9e8fc69-0e95-446d-a935-83c2fe83d429","cell_type":"markdown","source":"### 6.2 MC dropout uncertainty and triage policy\nThe expensive backbone runs once per image; only the dropout and the small heads are repeated 20 times, so uncertainty costs almost nothing extra.","metadata":{}},{"id":"1525b306-4f60-436a-8702-b3e06f20618a","cell_type":"code","source":"def enable_mc_dropout(model):\n    # evaluation mode, but dropout stays active\n    model.eval()\n    for module in model.modules():\n        if isinstance(module, nn.Dropout):\n            module.train()\n\n\ndef cumulative_to_grade_probs(cumulative):\n    # P(grade > k) -> probability of each grade 0-4\n    ones = torch.ones_like(cumulative[:, :1])\n    padded = torch.cat([ones, cumulative, torch.zeros_like(cumulative[:, :1])], dim=1)\n    return padded[:, :-1] - padded[:, 1:]\n\n\n@torch.no_grad()\ndef mc_dropout_predict(model, images, n_samples=20):\n    enable_mc_dropout(model)\n    features = model.forward_features(images).float()\n    cumulative_samples, severity_samples, quality_samples = [], [], []\n    for _ in range(n_samples):\n        grade_logits, quality_logits = model.heads_from_features(features)\n        cumulative_samples.append(corn_cumulative_probs(grade_logits))\n        severity_samples.append(corn_expected_grade(grade_logits))\n        quality_samples.append(F.softmax(quality_logits, dim=1))\n    mean_cumulative = torch.stack(cumulative_samples).mean(dim=0)\n    severity = torch.stack(severity_samples)\n    grade_probs = cumulative_to_grade_probs(mean_cumulative).clamp(min=0)\n    return {\n        'grade': (mean_cumulative > 0.5).sum(dim=1).cpu().numpy(),\n        'expected_grade': severity.mean(dim=0).cpu().numpy(),\n        'uncertainty': severity.std(dim=0).cpu().numpy(),\n        'confidence': grade_probs.max(dim=1).values.cpu().numpy(),\n        'entropy': (-(grade_probs * torch.log(grade_probs + 1e-12)).sum(dim=1)).cpu().numpy(),\n        'quality_probs': torch.stack(quality_samples).mean(dim=0).cpu().numpy(),\n    }\n\n\ndef make_triage_decision(predictions, index, settings):\n    # 1 ungradable -> re-capture | 2 uncertain -> human | 3 grade >= 2 -> refer | 4 routine\n    p_ungradable = float(predictions['quality_probs'][index][1])\n    grade = int(predictions['grade'][index])\n    confidence = float(predictions['confidence'][index])\n    uncertainty = float(predictions['uncertainty'][index])\n    borderline_quality = p_ungradable > 0.6 * settings.ungradable_threshold\n\n    if p_ungradable > settings.ungradable_threshold:\n        action = 'Re-capture image'\n    elif (uncertainty > settings.uncertainty_threshold or confidence < settings.confidence_threshold\n          or (borderline_quality and confidence < settings.confidence_threshold + 0.10)):\n        action = 'Refer for human grading'\n    elif grade >= 2:\n        action = 'Refer to ophthalmology'\n    else:\n        action = 'Routine rescreen'\n    return {'grade': grade, 'grade_name': GRADE_NAMES[grade], 'confidence': confidence,\n            'uncertainty': uncertainty, 'p_ungradable': p_ungradable, 'action': action}\n\n\ndef sweep_referral_thresholds(uncertainty, true_grades, predicted_grades,\n                              deferral_rates=(0.0, 0.05, 0.10, 0.20, 0.30)):\n    uncertainty = np.asarray(uncertainty, dtype=float)\n    true_grades, predicted_grades = np.asarray(true_grades), np.asarray(predicted_grades)\n    rows = []\n    for deferral_rate in deferral_rates:\n        kept_count = max(int(round((1.0 - deferral_rate) * len(true_grades))), 1)\n        cutoff = np.sort(uncertainty)[kept_count - 1]\n        kept = uncertainty <= cutoff\n        truly_referable, predicted_referable = true_grades[kept] >= 2, predicted_grades[kept] >= 2\n        rows.append({'deferred_%': int(deferral_rate * 100),\n                     'uncertainty_cutoff': float(cutoff),\n                     'images_kept': int(kept.sum()),\n                     'accuracy_on_kept': float((true_grades[kept] == predicted_grades[kept]).mean()),\n                     'referable_sensitivity_on_kept': float((predicted_referable & truly_referable).sum()\n                                                            / max(truly_referable.sum(), 1))})\n    return rows","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:53.912252Z","iopub.execute_input":"2026-09-21T09:02:53.912619Z","iopub.status.idle":"2026-09-21T09:02:53.930072Z","shell.execute_reply.started":"2026-09-21T09:02:53.912582Z","shell.execute_reply":"2026-09-21T09:02:53.92943Z"}},"outputs":[],"execution_count":null},{"id":"95e394d2-1835-4728-9ec0-c5a264545004","cell_type":"markdown","source":"### 6.3 Load the best checkpoint\nThe weight-averaged (EMA) weights from the epoch with the highest validation QWK.","metadata":{}},{"id":"682eaf76-1182-4b6e-9dd0-2126208d7f2d","cell_type":"code","source":"best_checkpoint = torch.load(Path(config.checkpoint_dir) / 'best.pt', map_location=config.device, weights_only=False)\nbest_weights = best_checkpoint['ema'] if (config.use_ema and best_checkpoint.get('ema')) else best_checkpoint['model']\n\nevaluation_model = build_model(config, pretrained=False).to(config.device)\nevaluation_model.load_state_dict(best_weights)\nevaluation_model.eval()\nprint(f\"loaded best checkpoint: epoch {best_checkpoint['epoch']} | validation QWK {best_checkpoint['best_qwk']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:53.931151Z","iopub.execute_input":"2026-09-21T09:02:53.931514Z","iopub.status.idle":"2026-09-21T09:02:54.991499Z","shell.execute_reply.started":"2026-09-21T09:02:53.931474Z","shell.execute_reply":"2026-09-21T09:02:54.990699Z"}},"outputs":[],"execution_count":null},{"id":"893aae45-38fe-47ae-a5e9-58587d73b98f","cell_type":"markdown","source":"### 6.4 Prediction helper","metadata":{}},{"id":"962d3289-d216-4a65-a15a-49213bb53bdb","cell_type":"code","source":"def predict_images(table, name):\n    loader = make_loader(table, train=False)\n    keys = ['grade', 'expected_grade', 'uncertainty', 'confidence', 'entropy', 'quality_probs']\n    collected, true_grades = {key: [] for key in keys}, []\n    start_time = time.time()\n    for images, grades, _ in tqdm(loader, leave=False):\n        outputs = mc_dropout_predict(evaluation_model, images.to(config.device), config.mc_dropout_samples)\n        for key in keys:\n            collected[key].append(outputs[key])\n        true_grades.append(grades.numpy())\n    print(f'[{name}] {len(table)} images predicted in {time.time() - start_time:.0f}s')\n    return {key: np.concatenate(values) for key, values in collected.items()}, np.concatenate(true_grades)\n\n\ndef select_rows(predictions, mask):\n    return {key: values[mask] for key, values in predictions.items()}\n\n\ndef print_binary_report(title, report):\n    print(f'\\n=== {title} ===')\n    print(f\"accuracy             : {report['accuracy']:.4f}\")\n    print(f\"precision            : {report['precision']:.4f}\")\n    print(f\"recall (sensitivity) : {report['recall_sensitivity']:.4f}\")\n    print(f\"specificity          : {report['specificity']:.4f}\")\n    print(f\"F1-score             : {report['f1']:.4f}\")\n    print(f\"TP {report['tp']} | FP {report['fp']} | FN {report['fn']} | TN {report['tn']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:54.99256Z","iopub.execute_input":"2026-09-21T09:02:54.99332Z","iopub.status.idle":"2026-09-21T09:02:55.001107Z","shell.execute_reply.started":"2026-09-21T09:02:54.993285Z","shell.execute_reply":"2026-09-21T09:02:55.000306Z"}},"outputs":[],"execution_count":null},{"id":"9dbe507f-d1fa-4e12-98c2-35a8910bf5b7","cell_type":"markdown","source":"### 6.5 DDR test set - overall results\nGrading metrics use gradable test images only; the ungradable ones are used for the triage check in 6.13.","metadata":{}},{"id":"831c953d-bad3-4171-a4c5-1d073c1f4dff","cell_type":"code","source":"ddr_test_all = ddr_images[ddr_images['split'] == 'test'].reset_index(drop=True)\ntest_all_predictions, test_all_true_grades = predict_images(ddr_test_all, 'DDR test')\n\ngradable_mask = ~ddr_test_all['ungradable'].values\nddr_test_set = ddr_test_all[gradable_mask].reset_index(drop=True)\nddr_test_predictions = select_rows(test_all_predictions, gradable_mask)\nddr_test_true_grades = test_all_true_grades[gradable_mask]\n\nddr_test_results = full_evaluation(ddr_test_true_grades, ddr_test_predictions['grade'],\n                                   ddr_test_predictions['expected_grade'],\n                                   uncertainty=ddr_test_predictions['uncertainty'],\n                                   confidences=ddr_test_predictions['confidence'])\n\nmajority_baseline = (ddr_test_true_grades == 0).mean()\nprint(f'\\n=== DDR test set: 5-stage classification ({len(ddr_test_true_grades)} images) ===')\nprint(f\"accuracy          : {ddr_test_results['accuracy']:.4f}  (always predicting 'No DR' would give {majority_baseline:.4f})\")\nprint(f\"macro precision   : {ddr_test_results['macro_precision']:.4f}\")\nprint(f\"macro recall      : {ddr_test_results['macro_recall']:.4f}\")\nprint(f\"macro F1-score    : {ddr_test_results['macro_f1']:.4f}\")\nprint(f\"weighted F1-score : {ddr_test_results['weighted_f1']:.4f}\")\nprint(f\"quadratic kappa   : {ddr_test_results['qwk']:.4f}\")\nprint(f\"referable DR AUC  : {ddr_test_results['referable']['auc']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:02:55.00252Z","iopub.execute_input":"2026-09-21T09:02:55.002866Z","iopub.status.idle":"2026-09-21T09:03:16.908411Z","shell.execute_reply.started":"2026-09-21T09:02:55.00283Z","shell.execute_reply":"2026-09-21T09:03:16.907508Z"}},"outputs":[],"execution_count":null},{"id":"c04516e3-fcd4-4dd3-a17c-1e722a4b2031","cell_type":"markdown","source":"### 6.6 Precision, recall and F1-score per stage","metadata":{}},{"id":"2ff944c9-6d5f-4afb-b7c0-817272562f26","cell_type":"code","source":"print(classification_report(ddr_test_true_grades, ddr_test_predictions['grade'], labels=list(range(5)),\n                            target_names=[GRADE_NAMES[g] for g in range(5)], digits=3, zero_division=0))\n\nper_grade_table = pd.DataFrame({\n    'grade': range(5),\n    'name': [GRADE_NAMES[g] for g in range(5)],\n    'images': [int((ddr_test_true_grades == g).sum()) for g in range(5)],\n    'precision': np.round(ddr_test_results['precision_per_class'], 3),\n    'recall': np.round(ddr_test_results['recall_per_class'], 3),\n    'f1_score': np.round(ddr_test_results['f1_per_class'], 3),\n})\nper_grade_table.to_csv('results/tables/table8_per_class.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:16.911178Z","iopub.execute_input":"2026-09-21T09:03:16.91145Z","iopub.status.idle":"2026-09-21T09:03:16.932846Z","shell.execute_reply.started":"2026-09-21T09:03:16.911414Z","shell.execute_reply":"2026-09-21T09:03:16.932076Z"}},"outputs":[],"execution_count":null},{"id":"b30583c4-9093-4ad3-af3d-0c043504feb3","cell_type":"markdown","source":"### 6.7 DR detection and referable DR decisions\nTwo yes/no questions derived from the predicted stage: **is any DR present** (grade 1 or higher), and **should the patient be referred** (grade 2 or higher).","metadata":{}},{"id":"41821db0-2a27-4790-bf7a-b029a6c006cc","cell_type":"code","source":"print_binary_report('DR detection: any DR (grade >= 1) vs No DR', ddr_test_results['dr_detection'])\nprint_binary_report('Referable DR: grade >= 2 vs grade 0-1', ddr_test_results['referable_decision'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:16.933907Z","iopub.execute_input":"2026-09-21T09:03:16.934681Z","iopub.status.idle":"2026-09-21T09:03:16.93951Z","shell.execute_reply.started":"2026-09-21T09:03:16.934646Z","shell.execute_reply":"2026-09-21T09:03:16.938753Z"}},"outputs":[],"execution_count":null},{"id":"526c9990-d2db-4ba0-9d46-0857b78a2e05","cell_type":"markdown","source":"### 6.8 Confusion matrix\nShows where the errors fall. Most errors should be between neighbouring stages.","metadata":{}},{"id":"584d4705-17dc-4bc3-87bb-2239919a43f3","cell_type":"code","source":"confusion_counts = np.array(ddr_test_results['confusion_matrix'])\nconfusion_normalised = confusion_counts / confusion_counts.sum(axis=1, keepdims=True).clip(min=1)\nshort_names = ['No DR', 'Mild', 'Moderate', 'Severe', 'PDR']\n\nshow_table('Confusion matrix (rows = true, columns = predicted)',\n           pd.DataFrame(confusion_counts, index=[f'true {n}' for n in short_names],\n                        columns=[f'pred {n}' for n in short_names]))\n\nfig, ax = plt.subplots(figsize=(6, 5))\nimage = ax.imshow(confusion_normalised, cmap='Purples', vmin=0, vmax=1)\nax.set_xticks(range(5)); ax.set_xticklabels(short_names, rotation=45, ha='right')\nax.set_yticks(range(5)); ax.set_yticklabels(short_names)\nax.set_xlabel('predicted'); ax.set_ylabel('true')\nfor i in range(5):\n    for j in range(5):\n        ax.text(j, i, f'{confusion_normalised[i, j]:.2f}\\n({confusion_counts[i, j]})', ha='center', va='center',\n                fontsize=8, color='white' if confusion_normalised[i, j] > 0.5 else 'black')\nax.set_title('Row-normalised confusion matrix')\nplt.colorbar(image)\nsave_figure('results/figures/fig10_confusion_matrix.png')\n\nadjacent_errors = sum(confusion_counts[i, j] for i in range(5) for j in range(5) if abs(i - j) == 1)\nfar_errors = sum(confusion_counts[i, j] for i in range(5) for j in range(5) if abs(i - j) >= 2)\ntotal_errors = confusion_counts.sum() - np.trace(confusion_counts)\nprint(f'total errors: {total_errors} | off by one stage: {adjacent_errors} | off by two or more: {far_errors}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:16.940587Z","iopub.execute_input":"2026-09-21T09:03:16.940939Z","iopub.status.idle":"2026-09-21T09:03:37.502607Z","shell.execute_reply.started":"2026-09-21T09:03:16.940914Z","shell.execute_reply":"2026-09-21T09:03:37.501917Z"}},"outputs":[],"execution_count":null},{"id":"c0a3f15e-1c8c-4672-ae6a-fccad11354bd","cell_type":"markdown","source":"### 6.9 ROC curve for referable DR\nUses the continuous severity score, with operating points at 90% and 95% specificity. The red line marks the common screening target of 80% sensitivity at 95% specificity.","metadata":{}},{"id":"96198f84-9bcb-4338-95ad-d3e04c6f482e","cell_type":"code","source":"is_referable = (ddr_test_true_grades >= 2).astype(int)\nfalse_positive_rate, true_positive_rate, _ = roc_curve(is_referable, ddr_test_predictions['expected_grade'])\nreferable_results = ddr_test_results['referable']\n\nfig, ax = plt.subplots(figsize=(5.5, 5))\nax.plot(false_positive_rate, true_positive_rate, color='#7F77DD', lw=2, label=f\"AUC = {referable_results['auc']:.3f}\")\nax.plot([0, 1], [0, 1], ls='--', lw=0.8, color='grey')\nfor specificity, colour in [(90, '#E8734A'), (95, '#1D9E75')]:\n    point = referable_results[f'sens_at_spec_{specificity}']\n    ax.plot(1 - point['specificity'], point['sensitivity'], 'o', color=colour,\n            label=f\"sens {point['sensitivity']:.3f} @ spec {point['specificity']:.2f}\")\nax.axhline(0.80, ls=':', lw=0.8, color='red')\nax.set_xlabel('1 - specificity'); ax.set_ylabel('sensitivity')\nax.set_title('Referable DR (grade >= 2)'); ax.legend(loc='lower right', fontsize=8)\nsave_figure('results/figures/fig11_referable_roc.png')\n\nprint(f\"referable DR AUC: {referable_results['auc']:.4f} | prevalence {referable_results['prevalence']:.1%}\")\nfor specificity in (90, 95):\n    point = referable_results[f'sens_at_spec_{specificity}']\n    print(f\"at {point['specificity']:.1%} specificity -> sensitivity {point['sensitivity']:.1%}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:37.503617Z","iopub.execute_input":"2026-09-21T09:03:37.504465Z","iopub.status.idle":"2026-09-21T09:03:37.854064Z","shell.execute_reply.started":"2026-09-21T09:03:37.504427Z","shell.execute_reply":"2026-09-21T09:03:37.85313Z"}},"outputs":[],"execution_count":null},{"id":"7c447a88-93c7-4b66-a3cc-27a4c8e8d1c8","cell_type":"markdown","source":"### 6.10 Calibration\nChecks whether the model's confidence matches its real accuracy. A well-calibrated model lies on the diagonal; ECE is the average gap.","metadata":{}},{"id":"65486e6b-6ba1-4bcc-bd14-5f46dc7198ae","cell_type":"code","source":"calibration = ddr_test_results['calibration']\ncalibration_bins = pd.DataFrame(calibration['bins'])\n\nfig, ax = plt.subplots(figsize=(5, 5))\nax.plot([0, 1], [0, 1], ls='--', color='grey', lw=0.8, label='perfect calibration')\nax.plot(calibration_bins['confidence'], calibration_bins['accuracy'], 'o-', color='#7F77DD', label='model')\nax.set_xlabel('mean confidence'); ax.set_ylabel('observed accuracy')\nax.set_title(f\"Reliability diagram (ECE = {calibration['ece']:.4f})\")\nax.legend(fontsize=8)\nsave_figure('results/figures/fig12_calibration.png')\n\nprint(f\"expected calibration error (ECE): {calibration['ece']:.4f}\")\nprint(f\"maximum calibration error (MCE) : {calibration['mce']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:37.855028Z","iopub.execute_input":"2026-09-21T09:03:37.855238Z","iopub.status.idle":"2026-09-21T09:03:38.2172Z","shell.execute_reply.started":"2026-09-21T09:03:37.855217Z","shell.execute_reply":"2026-09-21T09:03:38.216514Z"}},"outputs":[],"execution_count":null},{"id":"b1cd8e7b-42b3-4f5b-a21f-afd0642e541f","cell_type":"markdown","source":"### 6.11 Risk-coverage: does uncertainty find the mistakes?\nThe model grades only its most certain cases and passes the rest to a human. If uncertainty is useful, accuracy rises as coverage falls.","metadata":{}},{"id":"705c5da3-f55c-4cc6-b376-d95188a7ba43","cell_type":"code","source":"selective = ddr_test_results['selective']\ncoverage_points = pd.DataFrame(selective['points'])\n\nfig, ax = plt.subplots(figsize=(6, 4))\nax.plot(coverage_points['coverage'] * 100, coverage_points['accuracy'], 'o-', color='#7F77DD', label='accuracy')\nax.plot(coverage_points['coverage'] * 100, coverage_points['qwk'], 's-', color='#1D9E75', label='QWK')\nax.invert_xaxis()\nax.set_xlabel('coverage % (images graded by the model)')\nax.set_title(f\"Risk-coverage (AURC = {selective['aurc']:.4f})\")\nax.legend(fontsize=8)\nsave_figure('results/figures/fig13_risk_coverage.png')\n\nprint(f\"accuracy grading 100% of images: {coverage_points['accuracy'].iloc[0]:.4f}\")\nprint(f\"accuracy grading  80% of images: {coverage_points[coverage_points.coverage <= 0.8]['accuracy'].iloc[0]:.4f}\")\nprint(f\"accuracy grading  50% of images: {coverage_points[coverage_points.coverage <= 0.5]['accuracy'].iloc[0]:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:38.218179Z","iopub.execute_input":"2026-09-21T09:03:38.218476Z","iopub.status.idle":"2026-09-21T09:03:38.578989Z","shell.execute_reply.started":"2026-09-21T09:03:38.218452Z","shell.execute_reply":"2026-09-21T09:03:38.578211Z"}},"outputs":[],"execution_count":null},{"id":"b6a5e74e-0184-4d71-acf9-ddff05a0f94c","cell_type":"markdown","source":"### 6.12 Referral threshold sweep\nFor each share of images passed to a human, shows the accuracy and the referable-DR sensitivity on the images the model keeps. Deferring uncertain cases must not remove the sick patients.","metadata":{}},{"id":"495c0359-bee8-444b-b51a-b0cb0356614a","cell_type":"code","source":"referral_sweep = pd.DataFrame(sweep_referral_thresholds(ddr_test_predictions['uncertainty'],\n                                                        ddr_test_true_grades, ddr_test_predictions['grade']))\nreferral_sweep.to_csv('results/tables/table9_referral_sweep.csv', index=False)\nshow_table('Deferral rate vs accuracy and sensitivity on kept images', referral_sweep.round(4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:38.579936Z","iopub.execute_input":"2026-09-21T09:03:38.580362Z","iopub.status.idle":"2026-09-21T09:03:38.590349Z","shell.execute_reply.started":"2026-09-21T09:03:38.580336Z","shell.execute_reply":"2026-09-21T09:03:38.589752Z"}},"outputs":[],"execution_count":null},{"id":"c9a998ee-31ac-428f-816f-930ead735a13","cell_type":"markdown","source":"### 6.13 Triage decisions on the full test split (including ungradable images)\nApplies the four-step triage policy to every test image. Ungradable photos should mostly be sent for re-capture.","metadata":{}},{"id":"abc48a82-10ba-4b9b-a034-85639b4c1874","cell_type":"code","source":"triage_decisions = pd.DataFrame([make_triage_decision(test_all_predictions, i, config)\n                                 for i in range(len(ddr_test_all))])\ntriage_decisions['true_label'] = ddr_test_all['raw_label'].map(LABEL_NAMES).values\n\ntriage_table = pd.crosstab(triage_decisions['true_label'], triage_decisions['action'])\ntriage_table = triage_table.reindex([name for name in LABEL_NAMES.values() if name in triage_table.index])\ntriage_table.to_csv('results/tables/table9b_triage_actions.csv')\nshow_table('Share of each triage action', triage_decisions['action'].value_counts(normalize=True).round(3).to_frame('share'))\nshow_table('True label vs triage action', triage_table)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:38.591792Z","iopub.execute_input":"2026-09-21T09:03:38.592156Z","iopub.status.idle":"2026-09-21T09:03:38.626013Z","shell.execute_reply.started":"2026-09-21T09:03:38.592118Z","shell.execute_reply":"2026-09-21T09:03:38.625195Z"}},"outputs":[],"execution_count":null},{"id":"9d67fc50-48b7-452f-bba9-b2fc92e2f988","cell_type":"markdown","source":"### 6.14 External validation on APTOS 2019\nThe same model on a dataset from a different country, population and cameras. A drop in performance here shows the effect of this distribution shift.","metadata":{}},{"id":"9334aebe-ea4c-4b9f-8947-0b9f57705a14","cell_type":"code","source":"aptos_predictions, aptos_true_grades = predict_images(aptos_images, 'APTOS external')\naptos_results = full_evaluation(aptos_true_grades, aptos_predictions['grade'], aptos_predictions['expected_grade'],\n                                uncertainty=aptos_predictions['uncertainty'], confidences=aptos_predictions['confidence'])\n\nmetric_rows = {\n    'accuracy': 'accuracy', 'macro precision': 'macro_precision', 'macro recall': 'macro_recall',\n    'macro F1-score': 'macro_f1', 'weighted F1-score': 'weighted_f1', 'quadratic kappa': 'qwk'}\nexternal_comparison = pd.DataFrame({\n    'DDR test (internal)': [ddr_test_results[key] for key in metric_rows.values()] + [ddr_test_results['referable']['auc']],\n    'APTOS (external)': [aptos_results[key] for key in metric_rows.values()] + [aptos_results['referable']['auc']],\n}, index=list(metric_rows.keys()) + ['referable DR AUC']).round(4)\nexternal_comparison['difference'] = (external_comparison.iloc[:, 0] - external_comparison.iloc[:, 1]).round(4)\nexternal_comparison.to_csv('results/tables/table10_external_validation.csv')\nshow_table('Internal vs external results', external_comparison)\n\nprint('\\nAPTOS per-stage report:')\nprint(classification_report(aptos_true_grades, aptos_predictions['grade'], labels=list(range(5)),\n                            target_names=[GRADE_NAMES[g] for g in range(5)], digits=3, zero_division=0))\nprint_binary_report('APTOS DR detection: any DR vs No DR', aptos_results['dr_detection'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:03:38.62712Z","iopub.execute_input":"2026-09-21T09:03:38.627445Z","iopub.status.idle":"2026-09-21T09:04:15.598199Z","shell.execute_reply.started":"2026-09-21T09:03:38.627423Z","shell.execute_reply":"2026-09-21T09:04:15.597274Z"}},"outputs":[],"execution_count":null},{"id":"01c4a914-9d1e-41ac-98aa-35ab97c2f7d7","cell_type":"markdown","source":"### 6.15 Save all metrics","metadata":{}},{"id":"ebad3640-927d-4aba-9d62-27d64a9c0fea","cell_type":"code","source":"with open('results/logs/evaluation.json', 'w') as file:\n    json.dump({'ddr_test': ddr_test_results, 'aptos_external': aptos_results, 'config': config.to_dict()},\n              file, indent=2, default=float)\nprint('all metrics saved to results/logs/evaluation.json')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:04:15.600229Z","iopub.execute_input":"2026-09-21T09:04:15.600589Z","iopub.status.idle":"2026-09-21T09:04:15.608839Z","shell.execute_reply.started":"2026-09-21T09:04:15.600501Z","shell.execute_reply":"2026-09-21T09:04:15.607991Z"}},"outputs":[],"execution_count":null},{"id":"cf2c7ede-febd-4d0e-ae54-3413e2757123","cell_type":"markdown","source":"### 6.16 Largest errors\nThe six test images with the biggest grade error, with the model's confidence and uncertainty. Useful for error analysis (e.g. poor image quality, subtle lesions, label noise).","metadata":{}},{"id":"d989f216-eacf-4c36-92da-21deeec702fc","cell_type":"code","source":"grade_errors = np.abs(ddr_test_true_grades - ddr_test_predictions['grade'])\nworst_indices = np.argsort(-grade_errors)[:6]\n\nerror_rows = []\nfig, axes = plt.subplots(2, 3, figsize=(12, 8))\nfor ax in axes.ravel():\n    ax.axis('off')\nfor ax, index in zip(axes.ravel(), worst_indices):\n    image = cv2.imread(ddr_test_set.loc[index, 'cached'])\n    ax.imshow(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))\n    ax.set_title(f\"true {ddr_test_true_grades[index]} -> predicted {ddr_test_predictions['grade'][index]}\\n\"\n                 f\"confidence {ddr_test_predictions['confidence'][index]:.2f}, \"\n                 f\"uncertainty {ddr_test_predictions['uncertainty'][index]:.2f}\", fontsize=9)\n    error_rows.append({'image': ddr_test_set.loc[index, 'image'],\n                       'true_grade': int(ddr_test_true_grades[index]),\n                       'predicted_grade': int(ddr_test_predictions['grade'][index]),\n                       'confidence': round(float(ddr_test_predictions['confidence'][index]), 3),\n                       'uncertainty': round(float(ddr_test_predictions['uncertainty'][index]), 3)})\nplt.suptitle('Largest grading errors')\nsave_figure('results/figures/fig14_error_gallery.png', dpi=150)\nshow_table('Largest errors', pd.DataFrame(error_rows))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:04:15.609989Z","iopub.execute_input":"2026-09-21T09:04:15.610672Z","iopub.status.idle":"2026-09-21T09:05:37.937691Z","shell.execute_reply.started":"2026-09-21T09:04:15.610636Z","shell.execute_reply":"2026-09-21T09:05:37.935623Z"}},"outputs":[],"execution_count":null},{"id":"efd45890-8ace-42f3-8541-6f4d700dcb8e","cell_type":"markdown","source":"## 7. Explainability (Grad-CAM)\n\nGrad-CAM shows which parts of the image pushed the predicted severity up. For DR, the heatmap should highlight lesions (haemorrhages, exudates, microaneurysms) rather than the optic disc or the image border.\n\nA sanity check compares the trained model's heatmap with one from a randomly initialised network. If they looked the same, the heatmap would only reflect image edges, not what the model learned.\n\n### 7.1 Grad-CAM functions","metadata":{}},{"id":"4812432b-e57d-400e-9a79-a93531713839","cell_type":"code","source":"class GradCAM:\n    # heatmap of the regions that increase the predicted severity\n\n    def __init__(self, model, target_layer=None):\n        self.model = model\n        self.activations = None\n        self.gradients = None\n        layer = target_layer or self._last_conv_layer()\n        self.hook_handle = layer.register_forward_hook(self._save_activations)\n\n    def _last_conv_layer(self):\n        last = None\n        for module in self.model.backbone.modules():\n            if isinstance(module, nn.Conv2d):\n                last = module\n        return last\n\n    def _save_activations(self, module, inputs, output):\n        self.activations = output.detach()\n        if output.requires_grad:\n            output.register_hook(self._save_gradients)\n\n    def _save_gradients(self, gradient):\n        self.gradients = gradient.detach()\n\n    def remove(self):\n        self.hook_handle.remove()\n\n    def __call__(self, image_tensor):\n        self.model.eval()\n        self.model.zero_grad(set_to_none=True)\n        with torch.enable_grad():\n            grade_logits, _ = self.model(image_tensor)\n            corn_expected_grade(grade_logits.float()).sum().backward()\n        channel_weights = self.gradients.mean(dim=(2, 3), keepdim=True)\n        heatmap = F.relu((channel_weights * self.activations).sum(dim=1, keepdim=True))\n        heatmap = F.interpolate(heatmap, size=image_tensor.shape[-2:], mode='bilinear', align_corners=False)\n        heatmap = heatmap[0, 0].cpu().numpy()\n        if heatmap.max() > heatmap.min():\n            heatmap = (heatmap - heatmap.min()) / (heatmap.max() - heatmap.min())\n        return heatmap\n\n\ndef overlay_heatmap(image_bgr, heatmap, alpha=0.4):\n    heatmap = cv2.resize(heatmap.astype(np.float32), (image_bgr.shape[1], image_bgr.shape[0]))\n    coloured = cv2.applyColorMap(np.uint8(255 * heatmap), cv2.COLORMAP_JET)\n    return cv2.addWeighted(coloured, alpha, image_bgr, 1 - alpha, 0)\n\n\ndef image_to_tensor(image_bgr):\n    return eval_transform(cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB))[None].to(config.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:05:37.939152Z","iopub.execute_input":"2026-09-21T09:05:37.939566Z","iopub.status.idle":"2026-09-21T09:05:37.952952Z","shell.execute_reply.started":"2026-09-21T09:05:37.939508Z","shell.execute_reply":"2026-09-21T09:05:37.952284Z"}},"outputs":[],"execution_count":null},{"id":"100bc026-37cc-46f4-a33f-e43159ccd782","cell_type":"markdown","source":"### 7.2 Grad-CAM for one test image per grade","metadata":{}},{"id":"b72be7e2-3b9f-46f9-bf7d-17de8314e70e","cell_type":"code","source":"gradcam = GradCAM(evaluation_model)\n\nfig, axes = plt.subplots(2, 5, figsize=(16, 6.6))\nfor ax in axes.ravel():\n    ax.axis('off')\nfor column, grade in enumerate(range(5)):\n    candidates = ddr_test_set[ddr_test_set['level'] == grade]\n    if len(candidates) == 0:\n        continue\n    preprocessed_image = cv2.imread(candidates.sample(1, random_state=RANDOM_SEED).iloc[0]['cached'])\n    heatmap = gradcam(image_to_tensor(preprocessed_image))\n    axes[0, column].imshow(cv2.cvtColor(preprocessed_image, cv2.COLOR_BGR2RGB))\n    axes[0, column].set_title(f'grade {grade}: {GRADE_NAMES[grade]}', fontsize=9)\n    axes[1, column].imshow(cv2.cvtColor(overlay_heatmap(preprocessed_image, heatmap), cv2.COLOR_BGR2RGB))\nplt.suptitle('Grad-CAM by DR grade (top: input, bottom: heatmap)')\nsave_figure('results/figures/fig15_gradcam_by_grade.png', dpi=150)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:05:37.95379Z","iopub.execute_input":"2026-09-21T09:05:37.954014Z","iopub.status.idle":"2026-09-21T09:05:41.723255Z","shell.execute_reply.started":"2026-09-21T09:05:37.953978Z","shell.execute_reply":"2026-09-21T09:05:41.72212Z"}},"outputs":[],"execution_count":null},{"id":"66a35066-aab6-47f5-a388-8d514c772e28","cell_type":"markdown","source":"### 7.3 Sanity check: trained model vs random weights\nA low correlation between the two heatmaps means the trained heatmap reflects learned features.","metadata":{}},{"id":"ef03a560-dc86-4aae-bb0b-b4c61569b11c","cell_type":"code","source":"random_model = build_model(config, pretrained=False).to(config.device)\nrandom_gradcam = GradCAM(random_model)\n\nreferable_candidates = ddr_test_set[ddr_test_set['level'] >= 2]\npreprocessed_image = cv2.imread(referable_candidates.sample(1, random_state=RANDOM_SEED).iloc[0]['cached'])\nimage_tensor = image_to_tensor(preprocessed_image)\n\ntrained_heatmap, random_heatmap = gradcam(image_tensor), random_gradcam(image_tensor)\nmap_correlation = float(np.corrcoef(trained_heatmap.ravel(), random_heatmap.ravel())[0, 1])\n\nfig, axes = plt.subplots(1, 3, figsize=(13, 4.2))\npanels = [('input', cv2.cvtColor(preprocessed_image, cv2.COLOR_BGR2RGB)),\n          ('trained model', cv2.cvtColor(overlay_heatmap(preprocessed_image, trained_heatmap), cv2.COLOR_BGR2RGB)),\n          ('random weights', cv2.cvtColor(overlay_heatmap(preprocessed_image, random_heatmap), cv2.COLOR_BGR2RGB))]\nfor ax, (title, panel) in zip(axes, panels):\n    ax.imshow(panel); ax.set_title(title); ax.axis('off')\nplt.suptitle(f'Saliency sanity check - correlation {map_correlation:.3f} (low is good)')\nsave_figure('results/figures/fig15b_saliency_sanity_check.png', dpi=150)\nprint(f'correlation between trained and random heatmaps: {map_correlation:.3f}')\n\ngradcam.remove(); random_gradcam.remove()\ndel random_model\nif config.device == 'cuda':\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:05:41.724744Z","iopub.execute_input":"2026-09-21T09:05:41.7251Z","iopub.status.idle":"2026-09-21T09:05:43.200884Z","shell.execute_reply.started":"2026-09-21T09:05:41.725067Z","shell.execute_reply":"2026-09-21T09:05:43.200189Z"}},"outputs":[],"execution_count":null},{"id":"a7f38627-cde9-4b83-9ccf-5942558fca21","cell_type":"markdown","source":"## 8. Export for CPU deployment\n\nThe prototype app runs on a laptop without a GPU. The backbone is exported to **ONNX** (plus a smaller **int8** version), and the two small heads are saved as NumPy arrays so the app can still run MC dropout for uncertainty. The export is checked numerically against PyTorch and timed on CPU.\n\n### 8.1 Export functions","metadata":{}},{"id":"110da8e2-eeee-47b7-9a0d-c17cd5f226fe","cell_type":"code","source":"class BackboneOnly(nn.Module):\n    # ONNX graph = backbone features only\n\n    def __init__(self, network):\n        super().__init__()\n        self.network = network\n\n    def forward(self, images):\n        return self.network.forward_features(images)\n\n\ndef export_backbone_onnx(model, path, image_size=384, opset=17):\n    model.eval()\n    extra_options = {}\n    if 'dynamo' in inspect.signature(torch.onnx.export).parameters:\n        extra_options['dynamo'] = False\n    torch.onnx.export(BackboneOnly(model), torch.randn(1, 3, image_size, image_size), path,\n                      input_names=['image'], output_names=['features'],\n                      dynamic_axes={'image': {0: 'batch'}, 'features': {0: 'batch'}},\n                      opset_version=opset, do_constant_folding=True, **extra_options)\n    print(f'[onnx] {path} ({Path(path).stat().st_size / 1e6:.1f} MB)')\n    return path\n\n\ndef export_head_weights(model, path):\n    np.savez(path,\n             grade_w=model.grade_head.weight.detach().cpu().numpy(),\n             grade_b=model.grade_head.bias.detach().cpu().numpy(),\n             quality_w=model.quality_head.weight.detach().cpu().numpy(),\n             quality_b=model.quality_head.bias.detach().cpu().numpy(),\n             dropout_p=np.array(model.dropout.p, dtype=np.float32))\n    print(f'[heads] {path}')\n    return path\n\n\ndef quantize_int8(source_path, target_path):\n    try:\n        from onnxruntime.quantization import QuantType, quantize_dynamic\n        quantize_dynamic(source_path, target_path, weight_type=QuantType.QUInt8)\n    except Exception as error:\n        print('[int8] skipped:', error)\n        return None\n    print(f'[int8] {Path(source_path).stat().st_size / 1e6:.1f} MB -> {Path(target_path).stat().st_size / 1e6:.1f} MB')\n    return target_path\n\n\ndef benchmark_cpu(path, image_size=384, warmup_runs=3, timed_runs=20):\n    import onnxruntime as ort\n    try:\n        session = ort.InferenceSession(path, providers=['CPUExecutionProvider'])\n    except Exception as error:\n        print(f'[benchmark] {Path(path).name} skipped:', error)\n        return None\n    input_name = session.get_inputs()[0].name\n    dummy_image = np.random.randn(1, 3, image_size, image_size).astype(np.float32)\n    for _ in range(warmup_runs):\n        session.run(None, {input_name: dummy_image})\n    times_ms = []\n    for _ in range(timed_runs):\n        start_time = time.perf_counter()\n        session.run(None, {input_name: dummy_image})\n        times_ms.append((time.perf_counter() - start_time) * 1000)\n    times_ms = np.array(times_ms)\n    return {'model': Path(path).name, 'median_ms': round(float(np.median(times_ms)), 1),\n            'mean_ms': round(float(times_ms.mean()), 1), 'p95_ms': round(float(np.percentile(times_ms, 95)), 1),\n            'size_mb': round(Path(path).stat().st_size / 1e6, 1)}\n\n\ndef verify_export(model, path, image_size=384, sample_count=4, tolerance=1e-3):\n    # ONNX features must match PyTorch features\n    import onnxruntime as ort\n    model.eval()\n    test_batch = torch.randn(sample_count, 3, image_size, image_size)\n    with torch.no_grad():\n        pytorch_features = model.forward_features(test_batch).numpy()\n    session = ort.InferenceSession(path, providers=['CPUExecutionProvider'])\n    onnx_features = session.run(None, {session.get_inputs()[0].name: test_batch.numpy()})[0]\n    difference = np.abs(pytorch_features - onnx_features)\n    return {'max_abs_diff': float(difference.max()), 'mean_abs_diff': float(difference.mean()),\n            'verified': bool(difference.max() < tolerance), 'tolerance': tolerance}\n\n\ndef export_all(model, output_folder, image_size=384):\n    output_folder = Path(output_folder)\n    output_folder.mkdir(parents=True, exist_ok=True)\n    fp32_path = export_backbone_onnx(model, str(output_folder / 'backbone_fp32.onnx'), image_size)\n    export_head_weights(model, str(output_folder / 'heads.npz'))\n    int8_path = quantize_int8(fp32_path, str(output_folder / 'backbone_int8.onnx'))\n    benchmarks = [benchmark_cpu(path, image_size) for path in (fp32_path, int8_path) if path]\n    manifest = {'image_size': image_size,\n                'backbone': config.backbone,\n                'verification': verify_export(model, fp32_path, image_size),\n                'benchmarks': [b for b in benchmarks if b]}\n    (output_folder / 'deployment_manifest.json').write_text(json.dumps(manifest, indent=2))\n    return manifest","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:05:43.201814Z","iopub.execute_input":"2026-09-21T09:05:43.202798Z","iopub.status.idle":"2026-09-21T09:05:43.219448Z","shell.execute_reply.started":"2026-09-21T09:05:43.202762Z","shell.execute_reply":"2026-09-21T09:05:43.218787Z"}},"outputs":[],"execution_count":null},{"id":"c31f1691-afa8-40f8-9ffa-6bd8f41d7549","cell_type":"markdown","source":"### 8.2 Export the best model","metadata":{}},{"id":"98eaedd5-ea03-4933-b073-2bdf596cd408","cell_type":"code","source":"export_model = build_model(config, pretrained=False)\nexport_model.load_state_dict({name: weights.cpu() for name, weights in best_weights.items()})\nexport_model.eval()\n\ndeployment_manifest = export_all(export_model, 'models/export', config.image_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:05:43.22039Z","iopub.execute_input":"2026-09-21T09:05:43.220759Z","iopub.status.idle":"2026-09-21T09:06:04.832502Z","shell.execute_reply.started":"2026-09-21T09:05:43.220736Z","shell.execute_reply":"2026-09-21T09:06:04.831446Z"}},"outputs":[],"execution_count":null},{"id":"be8e01ef-5db3-4785-ac7e-2ab36ebdd12b","cell_type":"markdown","source":"### 8.3 Verification and CPU speed","metadata":{}},{"id":"3c8ca1d4-ae10-49d2-a2c4-f84c3ce9a385","cell_type":"code","source":"verification = deployment_manifest['verification']\nprint(f\"export verified: {verification['verified']} | max difference {verification['max_abs_diff']:.2e} \"\n      f\"(tolerance {verification['tolerance']})\")\nshow_table('CPU latency per image', pd.DataFrame(deployment_manifest['benchmarks']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:06:04.833976Z","iopub.execute_input":"2026-09-21T09:06:04.834352Z","iopub.status.idle":"2026-09-21T09:06:04.84262Z","shell.execute_reply.started":"2026-09-21T09:06:04.834315Z","shell.execute_reply":"2026-09-21T09:06:04.841717Z"}},"outputs":[],"execution_count":null},{"id":"b188a20b-ed2a-4c96-ba6a-d045f14bbd8e","cell_type":"markdown","source":"### 8.4 Package the results for download\n`dr_triage_artifacts.zip` contains the exported model, all figures, tables and logs.","metadata":{}},{"id":"33676859-9b04-411d-b9c6-e2201f22daff","cell_type":"code","source":"artifacts_zip = WORKING_DIR / 'dr_triage_artifacts.zip'\nwith zipfile.ZipFile(artifacts_zip, 'w', zipfile.ZIP_DEFLATED) as archive:\n    for folder in ['models/export', 'results']:\n        for path in sorted(Path(folder).rglob('*')):\n            if path.is_file():\n                archive.write(path, path.as_posix())\nprint(f'{artifacts_zip} ({artifacts_zip.stat().st_size / 1e6:.1f} MB)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:06:04.843679Z","iopub.execute_input":"2026-09-21T09:06:04.843975Z","iopub.status.idle":"2026-09-21T09:06:10.658448Z","shell.execute_reply.started":"2026-09-21T09:06:04.843932Z","shell.execute_reply":"2026-09-21T09:06:10.657611Z"}},"outputs":[],"execution_count":null},{"id":"740fe537-c227-4aa6-82b2-21cd55b9baa8","cell_type":"markdown","source":"## 9. Summary of results","metadata":{}},{"id":"0a053755-aadd-4b16-b08e-74dbf3af525c","cell_type":"code","source":"ddr_detection = ddr_test_results['dr_detection']\nprint('=' * 70)\nprint('DATA')\nprint(f\"  DDR         {len(ddr_images)} images ({int(ddr_images.ungradable.sum())} ungradable) | \"\n      f\"split: {ddr_images['split_source'].iloc[0]}\")\nprint(f\"  APTOS 2019  {len(aptos_images)} images (external test only)\")\nprint('MODEL')\nprint(f'  {config.backbone} | {len(history_table)} / {config.epochs} epochs | '\n      f'best validation QWK {history_table.val_qwk.max():.4f}')\nprint('DDR TEST - 5-stage classification')\nprint(f\"  accuracy {ddr_test_results['accuracy']:.4f} | macro precision {ddr_test_results['macro_precision']:.4f} | \"\n      f\"macro recall {ddr_test_results['macro_recall']:.4f} | macro F1 {ddr_test_results['macro_f1']:.4f} | \"\n      f\"QWK {ddr_test_results['qwk']:.4f}\")\nprint('DDR TEST - DR detection (any DR vs No DR)')\nprint(f\"  accuracy {ddr_detection['accuracy']:.4f} | precision {ddr_detection['precision']:.4f} | \"\n      f\"recall {ddr_detection['recall_sensitivity']:.4f} | F1 {ddr_detection['f1']:.4f}\")\nprint('DDR TEST - referable DR')\nprint(f\"  AUC {ddr_test_results['referable']['auc']:.4f} | \"\n      f\"sensitivity at 95% specificity {ddr_test_results['referable']['sens_at_spec_95']['sensitivity']:.4f}\")\nprint('APTOS EXTERNAL')\nprint(f\"  accuracy {aptos_results['accuracy']:.4f} | macro F1 {aptos_results['macro_f1']:.4f} | \"\n      f\"QWK {aptos_results['qwk']:.4f} | referable AUC {aptos_results['referable']['auc']:.4f}\")\nprint('UNCERTAINTY')\nprint(f\"  calibration ECE {ddr_test_results['calibration']['ece']:.4f} | \"\n      f\"risk-coverage AURC {ddr_test_results['selective']['aurc']:.4f}\")\nprint('DEPLOYMENT')\nprint(f\"  ONNX export verified: {deployment_manifest['verification']['verified']}\")\nprint('=' * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-21T09:06:10.659628Z","iopub.execute_input":"2026-09-21T09:06:10.660056Z","iopub.status.idle":"2026-09-21T09:06:10.669352Z","shell.execute_reply.started":"2026-09-21T09:06:10.660023Z","shell.execute_reply":"2026-09-21T09:06:10.668484Z"}},"outputs":[],"execution_count":null}]}