| """ |
| Optic disc / cup segmentation demo (research use only). |
| |
| Serves the Notebook 10 public-data model: U-Net++ / resnet18, trained from |
| scratch on ORIGA + G1020 + REFUGE + PAPILA, and reports the vertical CDR. |
| |
| Preprocessing here MUST stay identical to |
| glaucoma_segmentation/data/segmentation_dataset.py::load_rgb_image |
| - resize to (256, 256), PIL BILINEAR |
| - scale to [0, 1] by dividing by 255 |
| - NO ImageNet mean/std normalisation (the encoder was trained from scratch) |
| Any drift between the two silently degrades predictions. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import os |
| from pathlib import Path |
|
|
| import gradio as gr |
| import numpy as np |
| import segmentation_models_pytorch as smp |
| import torch |
| import torch.nn.functional as F |
| from PIL import Image |
|
|
| |
| ENCODER_NAME = "resnet18" |
| ENCODER_WEIGHTS = None |
| NUM_CLASSES = 3 |
| IMAGE_SIZE = (256, 256) |
|
|
| CKPT_FILENAME = os.environ.get("CKPT_FILENAME", "unetplusplus_resnet18_public_best.pt") |
| HF_MODEL_REPO = os.environ.get("MODEL_REPO", "") |
|
|
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
|
|
| |
| |
| |
| def resolve_checkpoint() -> Path: |
| """Prefer a checkpoint committed next to app.py; else pull from a model repo.""" |
| local = Path(CKPT_FILENAME) |
| if local.exists(): |
| return local |
|
|
| if HF_MODEL_REPO: |
| from huggingface_hub import hf_hub_download |
|
|
| return Path(hf_hub_download(repo_id=HF_MODEL_REPO, filename=CKPT_FILENAME)) |
|
|
| raise FileNotFoundError( |
| f"No checkpoint found. Commit {CKPT_FILENAME} to this Space, or set the " |
| "MODEL_REPO environment variable to a Hugging Face model repo containing it." |
| ) |
|
|
|
|
| def extract_state_dict(obj) -> dict: |
| """Pull a state_dict out of whatever torch.save wrote.""" |
| if not isinstance(obj, dict): |
| raise TypeError(f"Unsupported checkpoint object: {type(obj)!r}") |
|
|
| for key in ("model", "state_dict", "model_state_dict", "model_state", "best_state"): |
| value = obj.get(key) |
| if isinstance(value, dict) and value: |
| return value |
|
|
| if obj and all(torch.is_tensor(v) for v in obj.values()): |
| return obj |
|
|
| raise KeyError( |
| "Could not locate a state_dict in the checkpoint. " |
| f"Top-level keys: {sorted(obj.keys())[:12]}" |
| ) |
|
|
|
|
| def load_model() -> torch.nn.Module: |
| model = smp.UnetPlusPlus( |
| encoder_name=ENCODER_NAME, |
| encoder_weights=ENCODER_WEIGHTS, |
| in_channels=3, |
| classes=NUM_CLASSES, |
| ) |
|
|
| path = resolve_checkpoint() |
| try: |
| raw = torch.load(path, map_location="cpu", weights_only=True) |
| except Exception: |
| raw = torch.load(path, map_location="cpu", weights_only=False) |
|
|
| state = extract_state_dict(raw) |
| state = {k.removeprefix("module."): v for k, v in state.items()} |
|
|
| missing, unexpected = model.load_state_dict(state, strict=False) |
| if missing or unexpected: |
| print(f"[warn] missing={len(missing)} unexpected={len(unexpected)} keys") |
| if len(missing) > 10: |
| raise RuntimeError( |
| "Checkpoint does not match the U-Net++/resnet18 architecture. " |
| "Confirm this is the Notebook 10 checkpoint." |
| ) |
|
|
| model.to(DEVICE).eval() |
| return model |
|
|
|
|
| MODEL = load_model() |
|
|
|
|
| |
| |
| |
| def preprocess(pil_image: Image.Image) -> torch.Tensor: |
| image = pil_image.convert("RGB").resize(IMAGE_SIZE, resample=Image.BILINEAR) |
| array = np.asarray(image, dtype=np.float32) / 255.0 |
| return torch.from_numpy(array).permute(2, 0, 1).unsqueeze(0) |
|
|
|
|
| def vertical_extent(binary: np.ndarray) -> int: |
| rows = np.where(np.asarray(binary).any(axis=1))[0] |
| if len(rows) == 0: |
| return 0 |
| return int(rows[-1] - rows[0] + 1) |
|
|
|
|
| def vertical_cdr(label_mask: np.ndarray) -> float: |
| """Vertical CDR. Disc = labels > 0 (inclusive), cup = label == 2.""" |
| disc_height = vertical_extent(label_mask > 0) |
| cup_height = vertical_extent(label_mask == 2) |
| if disc_height <= 0: |
| return float("nan") |
| return cup_height / disc_height |
|
|
|
|
| |
| |
| |
| def inner_boundary(binary: np.ndarray) -> np.ndarray: |
| b = np.asarray(binary, dtype=bool) |
| edge = np.zeros_like(b) |
| edge[:-1, :] |= b[:-1, :] != b[1:, :] |
| edge[1:, :] |= b[1:, :] != b[:-1, :] |
| edge[:, :-1] |= b[:, :-1] != b[:, 1:] |
| edge[:, 1:] |= b[:, 1:] != b[:, :-1] |
| return edge & b |
|
|
|
|
| def crop_to_disc(image: np.ndarray, label_mask: np.ndarray, pad_frac: float = 0.35): |
| """Crop to the predicted disc bounding box, padded. Returns None if no disc.""" |
| rows = np.where((label_mask > 0).any(axis=1))[0] |
| cols = np.where((label_mask > 0).any(axis=0))[0] |
| if len(rows) == 0 or len(cols) == 0: |
| return None |
|
|
| pad = int(round(max(rows[-1] - rows[0], cols[-1] - cols[0]) * pad_frac)) + 1 |
| r0 = max(int(rows[0]) - pad, 0) |
| r1 = min(int(rows[-1]) + pad + 1, image.shape[0]) |
| c0 = max(int(cols[0]) - pad, 0) |
| c1 = min(int(cols[-1]) + pad + 1, image.shape[1]) |
| return image[r0:r1, c0:c1] |
|
|
|
|
| def make_overlay(rgb: np.ndarray, label_mask: np.ndarray) -> np.ndarray: |
| """Disc red, cup blue -- same convention as the project's QA overlays.""" |
| out = rgb.astype(np.float32).copy() |
|
|
| disc = label_mask == 1 |
| cup = label_mask == 2 |
| out[disc] = 0.60 * out[disc] + 0.40 * np.array([255.0, 0.0, 0.0]) |
| out[cup] = 0.55 * out[cup] + 0.45 * np.array([0.0, 0.0, 255.0]) |
|
|
| out = out.astype(np.uint8) |
| out[inner_boundary(label_mask > 0)] = [255, 255, 0] |
| out[inner_boundary(label_mask == 2)] = [0, 255, 255] |
| return out |
|
|
|
|
| |
| |
| |
| @torch.no_grad() |
| def predict(image: Image.Image): |
| if image is None: |
| return None, None, "" |
|
|
| original = np.asarray(image.convert("RGB")) |
| logits = MODEL(preprocess(image).to(DEVICE)) |
|
|
| |
| mask_eval = logits.argmax(1)[0].cpu().numpy() |
| cdr = vertical_cdr(mask_eval) |
|
|
| |
| logits_full = F.interpolate( |
| logits, size=original.shape[:2], mode="bilinear", align_corners=False |
| ) |
| mask_full = logits_full.argmax(1)[0].cpu().numpy() |
| overlay = make_overlay(original, mask_full) |
| zoom = crop_to_disc(overlay, mask_full) |
|
|
| disc_px = int((mask_eval > 0).sum()) |
| cup_px = int((mask_eval == 2).sum()) |
|
|
| if disc_px == 0: |
| cdr_text = "No disc detected" |
| else: |
| cdr_text = f"{cdr:.3f}" |
|
|
| return overlay, zoom, cdr_text |
|
|
|
|
| |
| |
| |
| TITLE = "Optic Disc & Cup Segmentation in Retinal Fundus Images" |
|
|
| DESCRIPTION = """ |
| Segments the **optic disc** and **optic cup** in a retinal fundus photograph and |
| reports the **vertical cup-to-disc ratio (CDR)** — a measurement clinicians use |
| when assessing glaucoma risk. |
| |
| Model: U-Net++ with a resnet18 encoder, trained from scratch on four public fundus datasets (ORIGA, G1020, REFUGE, PAPILA) with a leakage-aware, group-wise split by patient/image group. |
| """ |
|
|
| DISCLAIMER = """ |
| > ### ⚠️ Research demonstration — not a medical device |
| > This model does **not** diagnose glaucoma and must not be used for clinical |
| > decisions. It reports a CDR measurement, not a diagnosis. A CDR from an |
| > automated segmentation is not a substitute for examination by an |
| > ophthalmologist. |
| """ |
|
|
| with gr.Blocks(title=TITLE) as demo: |
| gr.Markdown(f"# {TITLE}") |
| gr.Markdown(DESCRIPTION) |
| gr.Markdown(DISCLAIMER) |
|
|
| with gr.Row(): |
| with gr.Column(): |
| |
| |
| |
| image_input = gr.Image( |
| label="Retinal fundus image", |
| type="pil", |
| sources=["upload", "clipboard"], |
| ) |
| run_button = gr.Button("Segment", variant="primary") |
|
|
| example_dir = Path("examples") |
| IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff", ".webp"} |
| example_files = ( |
| sorted( |
| str(p) for p in example_dir.glob("*") |
| if p.is_file() |
| and not p.name.startswith(".") |
| and p.suffix.lower() in IMAGE_EXTS |
| ) |
| if example_dir.exists() |
| else [] |
| ) |
| if example_files: |
| gr.Examples(examples=example_files, inputs=image_input, examples_per_page=12) |
|
|
| with gr.Column(): |
| cdr_output = gr.Textbox(label="Vertical cup-to-disc ratio") |
| with gr.Row(): |
| overlay_output = gr.Image(label="Segmentation (disc red, cup blue)") |
| zoom_output = gr.Image(label="Optic nerve head (zoomed)") |
|
|
| outputs = [overlay_output, zoom_output, cdr_output] |
| run_button.click(fn=predict, inputs=image_input, outputs=outputs) |
| image_input.change(fn=predict, inputs=image_input, outputs=outputs) |
|
|
| if __name__ == "__main__": |
| demo.launch() |