From eb9eafe715b12ee4589e9ddaa05f51fb2c215670 Mon Sep 17 00:00:00 2001 From: rpotter6298 Date: Thu, 19 Mar 2026 16:58:29 +0100 Subject: [PATCH] began work on v3 --- .gitignore | 3 +- portable_versions/image_loader.py | 351 ++++ portable_versions/refuge_mask_adapter.py | 270 +++ v3/__init__.py | 0 v3/classes/SE_attention.py | 123 ++ v3/classes/__init__.py | 107 + v3/classes/backbones.py | 178 ++ v3/classes/bridges.py | 93 + v3/classes/config_builder.py | 276 +++ v3/classes/croppers.py | 418 ++++ v3/classes/data_bundle.py | 241 +++ v3/classes/dataset.py | 115 + v3/classes/filters.py | 119 ++ v3/classes/geometry_features.py | 87 + v3/classes/hypertower_logger.py | 128 ++ v3/classes/image_loader.py | 149 ++ v3/classes/loader_factory.py | 244 +++ v3/classes/metrics.py | 243 +++ v3/classes/model_builder.py | 148 ++ v3/classes/models.py | 858 ++++++++ v3/classes/network_manager.py | 173 ++ v3/classes/papila_builders.py | 240 +++ v3/classes/papila_data.py | 61 + v3/classes/predictions.py | 194 ++ v3/classes/profiles/__init__.py | 10 + v3/classes/profiles/base.py | 53 + v3/classes/profiles/papila.py | 155 ++ v3/classes/results.py | 138 ++ v3/classes/slot_dataset.py | 114 + v3/classes/split_manager.py | 172 ++ v3/classes/towers.py | 279 +++ v3/classes/transforms.py | 320 +++ v3/classes/unet_segmenter.py | 894 ++++++++ v3/classes/utils.py | 47 + v3/classes/v2_hypertower.py | 1853 +++++++++++++++++ v3/classes/v3_hypertower.py | 1250 +++++++++++ v3/scripts/__init__.py | 0 v3/scripts/main/__init__.py | 0 v3/scripts/main/phase_1_papila_reproduce.py | 847 ++++++++ v3/scripts/main/run_cv.py | 87 + v3/scripts/output_analysis/__init__.py | 0 .../output_analysis/plot_phase1_boxplots.py | 177 ++ 42 files changed, 11214 insertions(+), 1 deletion(-) create mode 100644 portable_versions/image_loader.py create mode 100644 portable_versions/refuge_mask_adapter.py create mode 100644 v3/__init__.py create mode 100755 v3/classes/SE_attention.py create mode 100644 v3/classes/__init__.py create mode 100755 v3/classes/backbones.py create mode 100644 v3/classes/bridges.py create mode 100644 v3/classes/config_builder.py create mode 100644 v3/classes/croppers.py create mode 100644 v3/classes/data_bundle.py create mode 100644 v3/classes/dataset.py create mode 100644 v3/classes/filters.py create mode 100755 v3/classes/geometry_features.py create mode 100644 v3/classes/hypertower_logger.py create mode 100644 v3/classes/image_loader.py create mode 100644 v3/classes/loader_factory.py create mode 100644 v3/classes/metrics.py create mode 100644 v3/classes/model_builder.py create mode 100644 v3/classes/models.py create mode 100644 v3/classes/network_manager.py create mode 100644 v3/classes/papila_builders.py create mode 100644 v3/classes/papila_data.py create mode 100644 v3/classes/predictions.py create mode 100644 v3/classes/profiles/__init__.py create mode 100644 v3/classes/profiles/base.py create mode 100644 v3/classes/profiles/papila.py create mode 100644 v3/classes/results.py create mode 100644 v3/classes/slot_dataset.py create mode 100644 v3/classes/split_manager.py create mode 100644 v3/classes/towers.py create mode 100644 v3/classes/transforms.py create mode 100755 v3/classes/unet_segmenter.py create mode 100644 v3/classes/utils.py create mode 100644 v3/classes/v2_hypertower.py create mode 100644 v3/classes/v3_hypertower.py create mode 100644 v3/scripts/__init__.py create mode 100644 v3/scripts/main/__init__.py create mode 100644 v3/scripts/main/phase_1_papila_reproduce.py create mode 100644 v3/scripts/main/run_cv.py create mode 100644 v3/scripts/output_analysis/__init__.py create mode 100644 v3/scripts/output_analysis/plot_phase1_boxplots.py diff --git a/.gitignore b/.gitignore index ad54bdb..98d783a 100644 --- a/.gitignore +++ b/.gitignore @@ -11,4 +11,5 @@ models/refuge/ models/v2/refuge/ **/.archive/ .archive/ -scripts/deprecated/ \ No newline at end of file +scripts/deprecated/ +v3/results/* diff --git a/portable_versions/image_loader.py b/portable_versions/image_loader.py new file mode 100644 index 0000000..612bf7e --- /dev/null +++ b/portable_versions/image_loader.py @@ -0,0 +1,351 @@ +""" +portable_versions/image_loader.py +================================== +A self-contained image loader with in-memory caching, an optional +preprocessing pipeline (e.g. disc cropping), and composable augmentations. + +Returns plain NumPy arrays — works with PyTorch, TensorFlow, JAX, or +anything else that can consume an ndarray. + +Dependencies: Pillow, numpy (nothing else) + +Quickstart +---------- + from portable_versions.image_loader import ImageLoader, RandomHorizontalFlip, RandomRotation, ColorJitter + + # 1. Build the loader (once per run) + loader = ImageLoader( + target_size=(200, 200), + normalize=True, # float32 in [0, 1] with ImageNet mean/std + cache=True, # each image decoded from disk only once + workers=4, # parallel cache warm-up threads + preprocessor=my_crop_fn, # optional callable(PIL.Image) -> PIL.Image + ) + + # 2. Attach augmentations (applied randomly and independently per call) + loader.augmentation = [ + RandomHorizontalFlip(p=0.5), + RandomRotation(degrees=15), + ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05), + ] + + # 3. Warm the cache up front (optional but fast) + loader.warm(all_paths) + + # 4. Fetch images by path list — call as many times as you like + # Returns ndarray of shape (N, H, W, 3), dtype float32 + imgs = loader.get_img(train_paths) + + # For TensorFlow: + import tensorflow as tf + tensor = tf.constant(imgs) # (N, H, W, 3) + + # For PyTorch: + import torch + tensor = torch.from_numpy(imgs).permute(0, 3, 1, 2) # (N, C, H, W) + + +Augmentations reference +----------------------- +All augmentation classes live in this file and depend only on PIL + numpy. + + RandomHorizontalFlip(p=0.5) + RandomVerticalFlip(p=0.5) + RandomRotation(degrees=15) + ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05) + RandomGrayscale(p=0.1) + +You can also pass any callable(PIL.Image.Image) -> PIL.Image.Image as an +augmentation step. +""" +from __future__ import annotations + +import random +import threading +from concurrent.futures import ThreadPoolExecutor, as_completed +from pathlib import Path +from typing import Callable, Iterable, List, Optional, Tuple, Union + +import numpy as np +from PIL import Image, ImageEnhance, ImageOps + +# ImageNet channel statistics (RGB) +_IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) +_IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) + +PathLike = Union[str, Path] + + +def _call_preprocessor( + fn: Callable[..., Image.Image], + img: Image.Image, + path: Path, +) -> Image.Image: + """Call preprocessor as fn(img, path) if it accepts two args, else fn(img).""" + try: + return fn(img, path) + except TypeError: + return fn(img) + + +# --------------------------------------------------------------------------- +# Core loader +# --------------------------------------------------------------------------- + +class ImageLoader: + """ + Preprocessing pipeline + in-memory cache + augmentation, returning NumPy. + + Parameters + ---------- + target_size : (height, width) + Output spatial dimensions. Applied after ``preprocessor`` (if any). + Ignored when a ``preprocessor`` already resizes to the right size. + normalize : bool + When True, output is float32 with ImageNet mean/std subtraction. + When False, output is uint8 in [0, 255]. + cache : bool + Store decoded+preprocessed images in RAM so each file is read from + disk at most once. The cache persists across ``get_img`` calls. + workers : int + Thread count for ``warm()``. 0 or 1 = single-threaded. + preprocessor : callable, optional + Called as ``preprocessor(img: PIL.Image) -> PIL.Image`` before + resizing and caching. Use this for disc cropping, padding, etc. + augmentation : list of callables + Each element is called as ``fn(img: PIL.Image) -> PIL.Image``. + Applied **after** cache retrieval, so augmentations are NOT cached — + they are re-sampled independently on every ``get_img`` call. + """ + + def __init__( + self, + target_size: Tuple[int, int] = (224, 224), + *, + normalize: bool = True, + cache: bool = True, + workers: int = 4, + preprocessor: Optional[Callable[[Image.Image], Image.Image]] = None, + ) -> None: + self.target_size = target_size + self.normalize = normalize + self.workers = workers + self.preprocessor = preprocessor + self.augmentation: List[Callable[[Image.Image], Image.Image]] = [] + + self._cache: Optional[dict[str, np.ndarray]] = {} if cache else None + self._lock = threading.Lock() + + # ------------------------------------------------------------------ + # Public + # ------------------------------------------------------------------ + + def warm(self, paths: Iterable[PathLike]) -> None: + """ + Pre-load all *paths* into the cache in parallel. + + Already-cached paths are skipped, so calling ``warm`` multiple + times (e.g. once per fold) is safe and only loads new images. + """ + if self._cache is None: + return + + paths = [str(p) for p in paths] + to_warm = [p for p in paths if p not in self._cache] + if not to_warm: + return + + already = len(paths) - len(to_warm) + print( + f"[ImageLoader] warming {len(to_warm)} images" + + (f" ({already} already cached)" if already else ""), + flush=True, + ) + + def _load_one(path_str: str) -> None: + arr = self._decode(path_str) + with self._lock: + self._cache.setdefault(path_str, arr) + + if self.workers <= 1: + for p in to_warm: + _load_one(p) + else: + with ThreadPoolExecutor(max_workers=self.workers) as ex: + futures = {ex.submit(_load_one, p): p for p in to_warm} + for fut in as_completed(futures): + fut.result() + + def get_img( + self, + paths: Iterable[PathLike], + augment: bool = True, + ) -> np.ndarray: + """ + Return images for the given paths as a single NumPy array. + + Parameters + ---------- + paths : iterable of path-like + File paths to load. If the cache is enabled and a path has + been warmed (or loaded before), it is served from RAM. + augment : bool + Apply ``self.augmentation`` pipeline. Set to False at eval time. + + Returns + ------- + np.ndarray, shape (N, H, W, 3) + float32 in [0, 1] (or normalised) if ``self.normalize`` is True, + otherwise uint8 in [0, 255]. + """ + imgs = [] + for p in paths: + img = self._get_one(str(p), augment=augment) + imgs.append(img) + return np.stack(imgs, axis=0) + + # ------------------------------------------------------------------ + # Internal + # ------------------------------------------------------------------ + + def _decode(self, path_str: str) -> np.ndarray: + """Open, preprocess, and resize → uint8 HWC ndarray (for the cache).""" + img = Image.open(path_str).convert("RGB") + if self.preprocessor is not None: + img = _call_preprocessor(self.preprocessor, img, Path(path_str)) + # Only resize here if preprocessor didn't already produce target_size + if img.size != (self.target_size[1], self.target_size[0]): + img = img.resize((self.target_size[1], self.target_size[0]), Image.BILINEAR) + return np.asarray(img, dtype=np.uint8) + + def _get_one(self, path_str: str, augment: bool) -> np.ndarray: + if self._cache is not None: + arr = self._cache.get(path_str) + if arr is None: + arr = self._decode(path_str) + with self._lock: + self._cache.setdefault(path_str, arr) + img = Image.fromarray(arr, mode="RGB") + else: + img = Image.open(path_str).convert("RGB") + if self.preprocessor is not None: + img = self.preprocessor(img) + if img.size != (self.target_size[1], self.target_size[0]): + img = img.resize((self.target_size[1], self.target_size[0]), Image.BILINEAR) + + if augment and self.augmentation: + for fn in self.augmentation: + img = fn(img) + + arr = np.asarray(img, dtype=np.float32) / 255.0 + if self.normalize: + arr = (arr - _IMAGENET_MEAN) / _IMAGENET_STD + else: + arr = (arr * 255).clip(0, 255).astype(np.uint8) + return arr + + def __len__(self) -> int: + """Number of images currently in the cache.""" + return len(self._cache) if self._cache is not None else 0 + + def __repr__(self) -> str: + return ( + f"ImageLoader(target_size={self.target_size}, " + f"normalize={self.normalize}, " + f"cached={len(self)}, " + f"augmentations={len(self.augmentation)})" + ) + + +# --------------------------------------------------------------------------- +# Augmentation primitives (PIL-only, no torch/tf dependencies) +# --------------------------------------------------------------------------- + +class RandomHorizontalFlip: + """Flip image left-right with probability *p*.""" + def __init__(self, p: float = 0.5): + self.p = p + + def __call__(self, img: Image.Image) -> Image.Image: + return ImageOps.mirror(img) if random.random() < self.p else img + + +class RandomVerticalFlip: + """Flip image top-bottom with probability *p*.""" + def __init__(self, p: float = 0.5): + self.p = p + + def __call__(self, img: Image.Image) -> Image.Image: + return ImageOps.flip(img) if random.random() < self.p else img + + +class RandomRotation: + """Rotate by a uniformly-sampled angle in [-degrees, +degrees].""" + def __init__(self, degrees: float = 15): + self.degrees = degrees + + def __call__(self, img: Image.Image) -> Image.Image: + angle = random.uniform(-self.degrees, self.degrees) + return img.rotate(angle, resample=Image.BILINEAR, expand=False) + + +class ColorJitter: + """ + Randomly jitter brightness, contrast, saturation, and hue. + + Each factor is sampled uniformly from [1 - amount, 1 + amount]. + Hue shift is sampled from [-hue, +hue] (range 0–0.5). + Pass 0 for any channel to leave it unchanged. + """ + def __init__( + self, + brightness: float = 0.2, + contrast: float = 0.2, + saturation: float = 0.1, + hue: float = 0.05, + ): + self.brightness = brightness + self.contrast = contrast + self.saturation = saturation + self.hue = hue + + def __call__(self, img: Image.Image) -> Image.Image: + ops = [] + if self.brightness: + ops.append(("brightness", self.brightness)) + if self.contrast: + ops.append(("contrast", self.contrast)) + if self.saturation: + ops.append(("saturation", self.saturation)) + if self.hue: + ops.append(("hue", self.hue)) + random.shuffle(ops) + + for kind, amount in ops: + factor = random.uniform(1 - amount, 1 + amount) + if kind == "brightness": + img = ImageEnhance.Brightness(img).enhance(factor) + elif kind == "contrast": + img = ImageEnhance.Contrast(img).enhance(factor) + elif kind == "saturation": + img = ImageEnhance.Color(img).enhance(factor) + elif kind == "hue": + # PIL has no direct hue enhancer — shift via HSV in numpy + arr = np.asarray(img.convert("HSV"), dtype=np.int16) + shift = int(random.uniform(-self.hue, self.hue) * 255) + arr[:, :, 0] = (arr[:, :, 0] + shift) % 256 + img = Image.fromarray(arr.astype(np.uint8), mode="HSV").convert("RGB") + return img + + +class RandomGrayscale: + """Convert to grayscale (keeping 3 channels) with probability *p*.""" + def __init__(self, p: float = 0.1): + self.p = p + + def __call__(self, img: Image.Image) -> Image.Image: + if random.random() < self.p: + img = ImageOps.grayscale(img).convert("RGB") + return img + + diff --git a/portable_versions/refuge_mask_adapter.py b/portable_versions/refuge_mask_adapter.py new file mode 100644 index 0000000..fb21f85 --- /dev/null +++ b/portable_versions/refuge_mask_adapter.py @@ -0,0 +1,270 @@ +""" +portable_versions/refuge_mask_adapter.py +========================================= +Optic-disc cropper for REFUGE (and REFUGE2) fundus images, designed as a +drop-in ``preprocessor`` for ``ImageLoader``. + +Given an image and its corresponding segmentation mask, it: + 1. Extracts the optic disc region from the mask + 2. Computes a padded bounding box around it + 3. Crops and resizes the original image + +Dependencies: Pillow, numpy (nothing else) + +Quickstart +---------- + from portable_versions.image_loader import ImageLoader, RandomHorizontalFlip, RandomRotation, ColorJitter + from portable_versions.refuge_mask_adapter import RefugeMaskCropper + + cropper = RefugeMaskCropper( + mask_dir="REFUGE/Annotations/Training400/Disc_Cup_Masks", + scale=1.5, # context around disc (1.0 = tight, 2.0 = lots of context) + target_size=(200, 200), # output size — should match ImageLoader target_size + mask_suffix=".bmp", # REFUGE1 uses .bmp; REFUGE2 uses .png + ) + + loader = ImageLoader( + target_size=(200, 200), + normalize=True, + preprocessor=cropper, + ) + loader.augmentation = [ + RandomHorizontalFlip(), + RandomRotation(15), + ColorJitter(0.2, 0.2, 0.1, 0.05), + ] + + imgs = loader.get_img(image_paths, augment=True) # (N, 200, 200, 3) + +REFUGE mask formats +------------------- +REFUGE1 Grayscale BMP: background=128, disc=255, cup=0 +REFUGE2 RGB PNG: background detected from image borders, disc/cup by colour + +Both are handled automatically. + +Directory structure assumption +------------------------------ +The cropper looks for the mask with the same stem as the image file, inside +``mask_dir``. If your layout differs, pass a custom ``mask_path_fn``: + + cropper = RefugeMaskCropper( + mask_path_fn=lambda img_path: img_path.with_suffix(".bmp"), + scale=1.5, + target_size=(200, 200), + ) +""" +from __future__ import annotations + +from collections import Counter +from pathlib import Path +from typing import Optional, Tuple + +import numpy as np +from PIL import Image + + +_MASK_DIR_NAMES = {"Disc_Cup_Masks", "Disc_Masks", "Disc_Mask"} +_MASK_SUFFIXES = {".bmp", ".png"} + + +class RefugeMaskCropper: + """ + Crop a fundus image to the optic disc region using its segmentation mask. + + Pass the REFUGE root directory and the cropper will automatically index + all masks underneath it — no need to specify which subdirectory or + file extension. + + cropper = RefugeMaskCropper("REFUGE/", scale=1.5) + loader = ImageLoader(target_size=(200, 200), preprocessor=cropper) + imgs = loader.get_img(test_set) # test_set = any list of image paths + + Parameters + ---------- + refuge_root : str or Path + Top-level REFUGE directory. All mask files under directories named + ``Disc_Cup_Masks``, ``Disc_Masks``, or ``Disc_Mask`` are indexed + automatically (supports both .bmp and .png). + scale : float + Padding multiplier applied to the disc radius. + 1.0 = tight crop, 1.5 = moderate context, 2.5 = lots of context. + target_size : (height, width) + Output size after cropping. Should match ``ImageLoader.target_size``. + """ + + def __init__( + self, + refuge_root: str | Path, + *, + scale: float = 1.5, + target_size: Tuple[int, int] = (200, 200), + ) -> None: + self.refuge_root = Path(refuge_root) + self.scale = scale + self.target_size = target_size + self._index: dict[str, list[Path]] = {} + self._build_index() + + def _build_index(self) -> None: + """Walk refuge_root and index all mask files by stem (stem → [paths]).""" + for mask_dir in self.refuge_root.rglob("*"): + if mask_dir.is_dir() and mask_dir.name in _MASK_DIR_NAMES: + for f in mask_dir.rglob("*"): + if f.is_file() and f.suffix.lower() in _MASK_SUFFIXES: + self._index.setdefault(f.stem, []).append(f) + if not self._index: + raise FileNotFoundError( + f"No mask files found under {self.refuge_root!r}. " + f"Expected directories named: {_MASK_DIR_NAMES}" + ) + n_masks = sum(len(v) for v in self._index.values()) + print(f"[RefugeMaskCropper] indexed {n_masks} masks ({len(self._index)} unique stems)", flush=True) + + # ------------------------------------------------------------------ + # Callable interface — drop-in preprocessor for ImageLoader + # ------------------------------------------------------------------ + + def __call__( + self, + img: Image.Image, + img_path: Optional[str | Path] = None, + ) -> Image.Image: + stem = Path(img_path).stem if img_path else None + mask_path = self._lookup(stem, img_path) + disc_mask = _load_disc_mask(mask_path, img.size) + box = _mask_to_crop_box(disc_mask, scale=self.scale, img_size=img.size) + cropped = img.crop(box) + return cropped.resize( + (self.target_size[1], self.target_size[0]), Image.Resampling.BILINEAR + ) + + def _lookup(self, stem: Optional[str], img_path: Optional[str | Path] = None) -> Path: + if stem is None: + raise ValueError("img_path is required to match the mask.") + candidates = self._index.get(stem) + if not candidates: + raise KeyError( + f"No mask found for image stem {stem!r}. " + f"Available stems (sample): {list(self._index)[:5]}" + ) + if len(candidates) == 1: + return candidates[0] + # Pick the mask whose directory components best overlap with img_path + # (ignores the filename itself to handle extension differences) + img_parts = set(Path(img_path).parent.parts) if img_path else set() + return max(candidates, key=lambda m: len(set(m.parent.parts) & img_parts)) + + def __repr__(self) -> str: + return ( + f"RefugeMaskCropper(refuge_root={str(self.refuge_root)!r}, " + f"scale={self.scale}, target_size={self.target_size}, " + f"masks_indexed={len(self._index)})" + ) + + +# --------------------------------------------------------------------------- +# Mask parsing +# --------------------------------------------------------------------------- + +def _load_disc_mask(mask_path: Path, img_size: Tuple[int, int]) -> np.ndarray: + """ + Return a binary disc mask (uint8, 1=disc) from a REFUGE mask file. + + Handles: + - Grayscale BMP (REFUGE1): background≈128, disc=255, cup=0 + - RGB PNG (REFUGE2): background detected from image borders + """ + mask_img = Image.open(mask_path) + + if mask_img.mode == "L" or mask_img.mode == "P": + arr = np.asarray(mask_img.convert("L"), dtype=np.uint8) + bg = _border_mode(arr) + disc_mask = (arr != bg).astype(np.uint8) + else: + arr = np.asarray(mask_img.convert("RGB"), dtype=np.uint8) + bg = _border_mode_rgb(arr) + # disc = any non-background pixel + bg_mask = np.all(arr == bg, axis=2) + disc_mask = (~bg_mask).astype(np.uint8) + + # Ensure mask matches image spatial size + mh, mw = disc_mask.shape + iw, ih = img_size + if (mw, mh) != (iw, ih): + disc_img = Image.fromarray(disc_mask * 255).resize((iw, ih), Image.NEAREST) + disc_mask = (np.asarray(disc_img) > 0).astype(np.uint8) + + return disc_mask + + +def _border_mode(arr: np.ndarray, border: int = 5) -> int: + """Most common pixel value along the image border (grayscale).""" + h, w = arr.shape + border_pixels = np.concatenate([ + arr[:border, :].ravel(), + arr[-border:, :].ravel(), + arr[:, :border].ravel(), + arr[:, -border:].ravel(), + ]) + return int(Counter(border_pixels.tolist()).most_common(1)[0][0]) + + +def _border_mode_rgb(arr: np.ndarray, border: int = 5) -> np.ndarray: + """Most common RGB colour along the image border.""" + h, w, _ = arr.shape + border_pixels = np.concatenate([ + arr[:border, :].reshape(-1, 3), + arr[-border:, :].reshape(-1, 3), + arr[:, :border].reshape(-1, 3), + arr[:, -border:].reshape(-1, 3), + ], axis=0) + tuples = [tuple(row) for row in border_pixels.tolist()] + most_common = Counter(tuples).most_common(1)[0][0] + return np.array(most_common, dtype=np.uint8) + + +# --------------------------------------------------------------------------- +# Bounding box from mask +# --------------------------------------------------------------------------- + +def _mask_to_crop_box( + disc_mask: np.ndarray, + scale: float, + img_size: Tuple[int, int], +) -> Tuple[int, int, int, int]: + """ + Compute a square crop box centred on the disc with padding = scale * radius. + + Returns (left, upper, right, lower) — ready for PIL Image.crop(). + Falls back to the full image if no disc pixels are found. + """ + coords = np.argwhere(disc_mask > 0) # (N, 2) in (row, col) order + if coords.size == 0: + w, h = img_size + return (0, 0, w, h) + + ys, xs = coords[:, 0], coords[:, 1] + centre_x = float(xs.mean()) + centre_y = float(ys.mean()) + radius = max(float(xs.max() - xs.min()), float(ys.max() - ys.min())) / 2.0 + crop_radius = radius * scale + + iw, ih = img_size + left = int(max(0, centre_x - crop_radius)) + upper = int(max(0, centre_y - crop_radius)) + right = int(min(iw, centre_x + crop_radius)) + lower = int(min(ih, centre_y + crop_radius)) + + # Make square by expanding the shorter side + cw, ch = right - left, lower - upper + if cw < ch: + diff = ch - cw + left = max(0, left - diff // 2) + right = min(iw, right + diff // 2) + elif ch < cw: + diff = cw - ch + upper = max(0, upper - diff // 2) + lower = min(ih, lower + diff // 2) + + return (left, upper, right, lower) diff --git a/v3/__init__.py b/v3/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/v3/classes/SE_attention.py b/v3/classes/SE_attention.py new file mode 100755 index 0000000..178dc30 --- /dev/null +++ b/v3/classes/SE_attention.py @@ -0,0 +1,123 @@ +# se_block.py +import torch +import torch.nn as nn + +class SEGateLogger: + """ + Lightweight stats over SE gates. + Use: logger.accumulate(gates) each batch; logger.get() at epoch end. + """ + def __init__(self, enabled: bool = True, track_channels: bool = False, dim: int | None = None): + self.enabled = enabled + self.track_channels = track_channels + self.dim = dim + self.reset() + + def reset(self): + self._n = 0 + self._sum = 0.0 + self._sum2 = 0.0 + self._lt02 = 0 + self._gt08 = 0 + # optional per-channel + self._ch_sum = None + self._ch_count = 0 + if self.track_channels and self.dim is not None: + self._ch_sum = torch.zeros(self.dim, dtype=torch.float32) + + @torch.no_grad() + def accumulate(self, gates: torch.Tensor): + if not self.enabled: + return + # gates expected shape [N, C]; if a map/sequence gate is passed, reduce to [N, C] + if gates.dim() == 4: # [N,C,H,W] gates (uncommon) + g = gates.mean(dim=(2,3)) + elif gates.dim() == 3: # [N,T,C] gates (sequence) + g = gates.mean(dim=1) + elif gates.dim() == 2: # [N,C] + g = gates + else: + g = gates.view(gates.size(0), -1) + + g = g.detach() + self._n += g.numel() + self._sum += g.sum().item() + self._sum2 += (g*g).sum().item() + self._lt02 += (g < 0.2).sum().item() + self._gt08 += (g > 0.8).sum().item() + + if self._ch_sum is not None: + self._ch_sum += g.sum(dim=0).cpu() + self._ch_count += g.size(0) + + def get(self, reset: bool = True): + if self._n == 0: + return None + mean = self._sum / self._n + var = max(0.0, self._sum2 / self._n - mean * mean) + out = { + "mean": mean, + "std": var ** 0.5, + "pct_lt_0.2": self._lt02 / self._n, + "pct_gt_0.8": self._gt08 / self._n, + } + if self._ch_sum is not None and self._ch_count > 0: + out["channel_mean"] = (self._ch_sum / float(self._ch_count)).tolist() + if reset: + self.reset() + return out + +class SEBlock(nn.Module): + """ + SE-style channel gating that works for vectors and maps. + + Input: + - [N, C] (vector) -> squeeze = identity + - [N, C, H, W] (image map) -> squeeze over H,W + - [N, T, C] (sequence) -> squeeze over T + + Gate modes: + - residual (default): gate = 1 + tanh(MLP(s)) in (0, 2) [identity at init] + - plain: gate = sigmoid(MLP(s)) in (0, 1) + """ + def __init__(self, dim: int, reduction: int = 16, residual: bool = True, identity_init: bool = True): + super().__init__() + hid = max(1, dim // max(1, reduction)) + self.fc1 = nn.Linear(dim, hid, bias=True) + self.act = nn.ReLU(inplace=True) + self.fc2 = nn.Linear(hid, dim, bias=True) + self.residual = residual + + if residual and identity_init: + # make MLP output ~0 at start → gate ≈ 1.0 + nn.init.zeros_(self.fc2.weight) + nn.init.zeros_(self.fc2.bias) + + def _squeeze(self, x: torch.Tensor) -> torch.Tensor: + if x.dim() == 2: # [N,C] + return x + if x.dim() == 4: # [N,C,H,W] + return x.mean(dim=(2,3)) + if x.dim() == 3: # [N,T,C] + return x.mean(dim=1) + # fallback: flatten non-batch dims into channels + return x.view(x.size(0), -1) + + def _broadcast(self, gate: torch.Tensor, like: torch.Tensor) -> torch.Tensor: + if like.dim() == 2: + return gate + if like.dim() == 3: + return gate.unsqueeze(1) # [N,1,C] + if like.dim() == 4: + return gate.unsqueeze(-1).unsqueeze(-1) # [N,C,1,1] + return gate.view_as(like) + + def forward(self, x: torch.Tensor): + s = self._squeeze(x) # [N,C] + u = self.fc2(self.act(self.fc1(s))) # [N,C] + if self.residual: + gate = 1.0 + torch.tanh(u) # (0, 2) with identity at 1.0 + else: + gate = torch.sigmoid(u) # (0, 1) + y = x * self._broadcast(gate, x) + return y, gate # return both the reweighted tensor and the gate for logging diff --git a/v3/classes/__init__.py b/v3/classes/__init__.py new file mode 100644 index 0000000..01a94bd --- /dev/null +++ b/v3/classes/__init__.py @@ -0,0 +1,107 @@ +from .network_manager import ( + FoldResult, + LoaderBundle, + NetworkManager, + PatientSplit, +) +from .split_manager import ( + PatientFirstSplitManager, + SplitPlan, + build_patient_split_plans, +) +from .profiles import ( + DatasetProfile, + SimpleDatasetProfile, + SlotDescriptor, + PapilaProfile, + build_papila_profile, +) +from .loader_factory import SlotLoaderFactory +from .slot_dataset import SlotDataset, slot_collate +from .papila_data import PapilaData +from .papila_builders import build_papila_data +from .data_bundle import DataBundle +from .dataset import ClinicalDataset +from .config_builder import ( + ConfigAssembly, + assemble_config, + load_config, + resolve_imports, +) +from .filters import RegexFilter, ColumnFilter, apply_regex_filters, apply_column_filters +from .transforms import ( + ImageTransformConfig, + backbone_transform_config, + build_backbone_transform, + build_eval_transform, + build_imagenet_transform, + ResizeTransform, + CenterCropTransform, + ROICropTransform, + JitterBundleTransform, + UnetMaskProvider, + TRANSFORM_REGISTRY, + build_transform_chain, +) +from .model_builder import V2ModelBundle, build_model_bundle +from .towers import ImageTower, ClinicalTower, SiameseImageTower, build_backbone +from .bridges import Bridge, VoteBridge +from .models import SingleEyeHT, BilateralHT +from .v2_hypertower import V2HyperTower, V2ModeComparisonOps, V2ModeComparator +from .hypertower_logger import HypertowerLogger + +__all__ = [ + "NetworkManager", + "PatientSplit", + "LoaderBundle", + "FoldResult", + "PatientFirstSplitManager", + "SplitPlan", + "build_patient_split_plans", + "DatasetProfile", + "SimpleDatasetProfile", + "SlotDescriptor", + "PapilaProfile", + "build_papila_profile", + "PapilaData", + "build_papila_data", + "DataBundle", + "ClinicalDataset", + "SlotLoaderFactory", + "SlotDataset", + "slot_collate", + "ConfigAssembly", + "assemble_config", + "load_config", + "resolve_imports", + "RegexFilter", + "ColumnFilter", + "apply_regex_filters", + "apply_column_filters", + "ImageTransformConfig", + "backbone_transform_config", + "build_backbone_transform", + "build_eval_transform", + "build_imagenet_transform", + "ResizeTransform", + "CenterCropTransform", + "ROICropTransform", + "JitterBundleTransform", + "UnetMaskProvider", + "TRANSFORM_REGISTRY", + "build_transform_chain", + "V2ModelBundle", + "build_model_bundle", + "ImageTower", + "ClinicalTower", + "SiameseImageTower", + "build_backbone", + "Bridge", + "VoteBridge", + "SingleEyeHT", + "BilateralHT", + "V2HyperTower", + "V2ModeComparisonOps", + "V2ModeComparator", + "HypertowerLogger", +] diff --git a/v3/classes/backbones.py b/v3/classes/backbones.py new file mode 100755 index 0000000..2484e62 --- /dev/null +++ b/v3/classes/backbones.py @@ -0,0 +1,178 @@ +# classes/backbones.py +from __future__ import annotations +from dataclasses import dataclass +from pathlib import Path +from typing import Callable, Dict, List + +import torch +from torch import nn +from torchvision import models + +@dataclass(frozen=True) +class BackboneSpec: + ctor: Callable # torchvision constructor + weights_default: object # torchvision Weights enum DEFAULT member + strip: Callable[[nn.Module], tuple] # fn(model)->(out_dim, model_no_head) + blocks: Callable[[nn.Module], List[nn.Module]] # fn(model)->ordered blocks for freezing + +REFUGELIKE_BACKBONE_PATH = Path("models/v2/refuge/refugelike_backbone.pt") +REFUGE_DENSENET_PATH = Path("models/refuge/classifier/refuge_densenet_backbone.pt") +REFUGE_EFFICIENT_B0_PATH = Path("models/refuge/classifier/refuge_efficient_b0_backbone.pt") +REFUGE_EFFICIENT_B7_PATH = Path("models/refuge/classifier/refuge_efficient_b7_backbone.pt") + +# --- strip fns --- +def _strip_efficientnet_b0(m: models.EfficientNet): + from torch import nn as _nn + out_dim = m.classifier[1].in_features + m.classifier = _nn.Identity() + return out_dim, m + +def _strip_resnet(m: models.ResNet): + out_dim = m.fc.in_features + m.fc = nn.Identity() + return out_dim, m + +def _strip_densenet(m: models.DenseNet): + out_dim = m.classifier.in_features + m.classifier = nn.Identity() + return out_dim, m + +def _strip_vgg(m: models.VGG): + out_dim = m.classifier[0].in_features # 25088 for VGG16 at 224×224 + m.classifier = nn.Identity() + return out_dim, m + +def _strip_mobilenet_v2(m: models.MobileNetV2): + out_dim = m.classifier[1].in_features + m.classifier = nn.Identity() + return out_dim, m + +def _strip_inception_v3(m: models.Inception3): + out_dim = m.fc.in_features + m.fc = nn.Identity() + m.aux_logits = False + m.AuxLogits = None # torchvision checks `AuxLogits is not None`, not the flag + return out_dim, m + +# --- block splitters for ratio-based freezing --- +def _blocks_efficientnet_b0(m: models.EfficientNet): + return list(m.features) + +def _blocks_resnet(m: models.ResNet): + stem = nn.Sequential(m.conv1, m.bn1, m.relu, m.maxpool) + return [stem, m.layer1, m.layer2, m.layer3, m.layer4] + +def _blocks_densenet(m: models.DenseNet): + f = m.features + stem = nn.Sequential(f.conv0, f.norm0, f.relu0, f.pool0) + return [stem, f.denseblock1, f.transition1, f.denseblock2, f.transition2, + f.denseblock3, f.transition3, f.denseblock4, f.norm5] + +def _blocks_vgg(m: models.VGG): + stages, cur = [], [] + for mod in m.features: + cur.append(mod) + if isinstance(mod, nn.MaxPool2d): + stages.append(nn.Sequential(*cur)); cur = [] + if cur: stages.append(nn.Sequential(*cur)) + return stages + +def _blocks_mobilenet_v2(m: models.MobileNetV2): + return list(m.features) + +def _blocks_inception_v3(m: models.Inception3): + blocks = [] + for name, child in m.named_children(): + if name in ("fc", "AuxLogits"): + continue + blocks.append(child) + return blocks + +# --- registry (covers paper models available in torchvision) --- +BACKBONES: Dict[str, BackboneSpec] = { + "efficientnet_b0": BackboneSpec( + ctor=models.efficientnet_b0, + weights_default=models.EfficientNet_B0_Weights.DEFAULT, + strip=_strip_efficientnet_b0, + blocks=_blocks_efficientnet_b0, + ), + "resnet50": BackboneSpec( + ctor=models.resnet50, + weights_default=models.ResNet50_Weights.DEFAULT, + strip=_strip_resnet, + blocks=_blocks_resnet, + ), + "densenet121": BackboneSpec( + ctor=models.densenet121, + weights_default=models.DenseNet121_Weights.DEFAULT, + strip=_strip_densenet, + blocks=_blocks_densenet, + ), + "vgg16": BackboneSpec( + ctor=models.vgg16, + weights_default=models.VGG16_Weights.DEFAULT, + strip=_strip_vgg, + blocks=_blocks_vgg, + ), + "mobilenet_v2": BackboneSpec( + ctor=models.mobilenet_v2, + weights_default=models.MobileNet_V2_Weights.DEFAULT, + strip=_strip_mobilenet_v2, + blocks=_blocks_mobilenet_v2, + ), + "inception_v3": BackboneSpec( + ctor=models.inception_v3, + weights_default=models.Inception_V3_Weights.DEFAULT, + strip=_strip_inception_v3, + blocks=_blocks_inception_v3, + ), + "refugelike": BackboneSpec( + ctor=models.resnet50, + weights_default=None, + strip=_strip_resnet, + blocks=_blocks_resnet, + ), + "refuge_densenet": BackboneSpec( + ctor=models.densenet121, + weights_default=None, + strip=_strip_densenet, + blocks=_blocks_densenet, + ), + "refuge_efficient_b0": BackboneSpec( + ctor=models.efficientnet_b0, + weights_default=None, + strip=_strip_efficientnet_b0, + blocks=_blocks_efficientnet_b0, + ), + "refuge_efficient_b7": BackboneSpec( + ctor=models.efficientnet_b7, + weights_default=None, + strip=_strip_efficientnet_b0, + blocks=_blocks_efficientnet_b0, + ), + # Xception isn’t in torchvision +} + +def list_names() -> List[str]: + return list(BACKBONES.keys()) + + +def load_backbone_weights(key: str, model: nn.Module) -> None: + if key == "refugelike": + path = REFUGELIKE_BACKBONE_PATH + elif key == "refuge_densenet": + path = REFUGE_DENSENET_PATH + elif key == "refuge_efficient_b0": + path = REFUGE_EFFICIENT_B0_PATH + elif key == "refuge_efficient_b7": + path = REFUGE_EFFICIENT_B7_PATH + else: + return + + if not path.exists(): + raise FileNotFoundError( + "Custom REFUGE backbone not found at " + f"{path}. Export it via refuge_build.py --export-backbone first." + ) + state = torch.load(path, map_location="cpu") + model.load_state_dict(state, strict=False) diff --git a/v3/classes/bridges.py b/v3/classes/bridges.py new file mode 100644 index 0000000..b2783bb --- /dev/null +++ b/v3/classes/bridges.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import torch +import torch.nn as nn + +from v3.classes.SE_attention import SEBlock, SEGateLogger + + +class Bridge(nn.Module): + def __init__( + self, + img_dim, + meta_dim, + num_classes, + fusion_dim=256, + mode="fused", + use_se: bool = True, + se_reduction: int = 16, + se_pre_norm: bool = True, + ): + super().__init__() + self.mode = mode + self.use_se = use_se + + # project towers to equal width + self.W_img = nn.Linear(img_dim, fusion_dim) + self.W_md = nn.Linear(meta_dim, fusion_dim) + + # optional: layernorm before SE + self.ln_img = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity() + self.ln_md = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity() + + # SE gate on the fused vector + self.se = SEBlock(fusion_dim, reduction=se_reduction, residual=True) if use_se else None + self.se_log = SEGateLogger(enabled=use_se, track_channels=False, dim=fusion_dim) + + # heads + self.classifier_fused = nn.Sequential( + nn.ReLU(), + nn.Dropout(0.5), + nn.Linear(fusion_dim, num_classes), + ) + self.classifier_img = nn.Linear(img_dim, num_classes) + self.classifier_cd = nn.Linear(meta_dim, num_classes) + + def reset_se_stats(self): + """Call at epoch start.""" + if getattr(self, "se_log", None): + self.se_log.reset() + + def get_se_stats(self, reset: bool = True): + """Call after eval. Returns dict or None.""" + if getattr(self, "se_log", None) and self.se_log.enabled: + return self.se_log.get(reset=reset) + return None + + def forward(self, img_feats, md_feats): + out_img = None if self.mode == "clinical_only" else self.classifier_img(img_feats) + out_md = None if self.mode == "image_only" else self.classifier_cd(md_feats) + + if self.mode == "fused": + hi = self.ln_img(self.W_img(img_feats)) # image features + hm = self.ln_md(self.W_md(md_feats)) # clinical data features + fused = hi * hm # elementwise product + # apply SE gates + if self.se is not None: + fused, gates = self.se(fused) + if self.se_log.enabled: + self.se_log.accumulate(gates) + + if self.se is not None and self.training and self.se_log.enabled: + if not hasattr(self, "_dbg_seen"): + self._dbg_seen = 0 + if self._dbg_seen < 3: # print only a few times + print("[SE] gate mean this batch:", gates.mean().item()) + self._dbg_seen += 1 + out_f = self.classifier_fused(fused) + return out_f, out_img, out_md + # if ablation modes: + if self.mode == "image_only": + return out_img, out_img, None + if self.mode == "clinical_only": + return out_md, None, out_md + + +class VoteBridge(nn.Module): + def __init__(self, num_classes): + super().__init__() + self.vote_combiner = nn.Linear(num_classes * 2, num_classes) # two sets of logits + + def forward(self, out_img, out_md): + votes = torch.cat([out_img, out_md], dim=1) + return self.vote_combiner(votes) diff --git a/v3/classes/config_builder.py b/v3/classes/config_builder.py new file mode 100644 index 0000000..9ec4df6 --- /dev/null +++ b/v3/classes/config_builder.py @@ -0,0 +1,276 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Dict, Iterable, List, Optional + +import json + +from v3.classes.papila_data import PapilaData + + +@dataclass +class ImportSpec: + id: str + class_name: str + params: Dict[str, Any] + + +@dataclass +class DataSourceSpec: + node_id: str + label: str + output_type: str + source: Optional[Dict[str, Any]] + source_ref: Optional[Dict[str, Any]] + + +@dataclass +class TransformSpec: + node_id: str + label: str + transform_type: str + params: Dict[str, Any] + + +@dataclass +class LoaderSpec: + node_id: str + label: str + input_type: str + input_index: str + input_key: str + output_key: str + transforms: List[TransformSpec] + data_source: Optional[DataSourceSpec] + + +@dataclass +class TowerSpec: + node_id: str + label: str + tower_type: str + params: Dict[str, Any] + + +@dataclass +class BridgeSpec: + node_id: str + label: str + method: str + params: Dict[str, Any] + + +@dataclass +class ClassifierSpec: + node_id: str + label: str + + +@dataclass +class ConfigAssembly: + raw: Dict[str, Any] + imports: Dict[str, ImportSpec] + data_sources: Dict[str, DataSourceSpec] + transforms: Dict[str, TransformSpec] + loaders: Dict[str, LoaderSpec] + towers: Dict[str, TowerSpec] + bridges: Dict[str, BridgeSpec] + classifiers: Dict[str, ClassifierSpec] + + +def load_config(path: Path) -> Dict[str, Any]: + payload = json.loads(Path(path).read_text()) + if not isinstance(payload, dict): + raise ValueError("Config JSON must be an object.") + return payload + + +def assemble_config(path: Path) -> ConfigAssembly: + config = load_config(path) + meta = config.get("meta", {}) + imports = _build_imports(meta.get("imports", [])) + nodes = {node["id"]: node for node in config.get("nodes", [])} + edges = config.get("edges", []) + + data_sources: Dict[str, DataSourceSpec] = {} + transforms: Dict[str, TransformSpec] = {} + loaders: Dict[str, LoaderSpec] = {} + towers: Dict[str, TowerSpec] = {} + bridges: Dict[str, BridgeSpec] = {} + classifiers: Dict[str, ClassifierSpec] = {} + + for node in nodes.values(): + ntype = node.get("type") + if ntype == "data": + data_sources[node["id"]] = DataSourceSpec( + node_id=node["id"], + label=node.get("label", ""), + output_type=node.get("outputType", ""), + source=node.get("source"), + source_ref=node.get("sourceRef"), + ) + elif ntype == "transform": + transforms[node["id"]] = TransformSpec( + node_id=node["id"], + label=node.get("label", ""), + transform_type=node.get("transformType", ""), + params=_extract_transform_params(node), + ) + elif ntype == "loader": + loaders[node["id"]] = LoaderSpec( + node_id=node["id"], + label=node.get("label", ""), + input_type=node.get("inputType", ""), + input_index=node.get("inputIndex", ""), + input_key=node.get("inputKey", ""), + output_key=node.get("outputKey", ""), + transforms=[], + data_source=None, + ) + elif ntype in ("image_tower", "metadata_tower"): + towers[node["id"]] = TowerSpec( + node_id=node["id"], + label=node.get("label", ""), + tower_type=node.get("towerType", "image" if ntype == "image_tower" else "clinical data"), + params=_extract_tower_params(node), + ) + elif ntype == "bridge": + bridges[node["id"]] = BridgeSpec( + node_id=node["id"], + label=node.get("label", ""), + method=node.get("bridgeMethod", "fusion"), + params=_extract_bridge_params(node), + ) + elif ntype == "classifier": + classifiers[node["id"]] = ClassifierSpec( + node_id=node["id"], + label=node.get("label", ""), + ) + + # attach transforms + data sources to loaders by walking upstream + for loader_id, loader in loaders.items(): + chain = _upstream_chain(loader_id, nodes, edges) + for node_id in reversed(chain): + if node_id in transforms: + loader.transforms.append(transforms[node_id]) + if node_id in data_sources: + loader.data_source = data_sources[node_id] + + return ConfigAssembly( + raw=config, + imports=imports, + data_sources=data_sources, + transforms=transforms, + loaders=loaders, + towers=towers, + bridges=bridges, + classifiers=classifiers, + ) + + +def resolve_imports(assembly: ConfigAssembly) -> Dict[str, Any]: + resolved: Dict[str, Any] = {} + for import_id, spec in assembly.imports.items(): + if spec.class_name == "PapilaData": + params = spec.params + resolved[import_id] = PapilaData.from_dirs( + image_dir=params.get("image_dir", "Papila/FundusImages"), + clinical_dir=params.get("clinical_dir", "Papila/ClinicalData"), + label_col=params.get("label_col", "Diagnosis"), + cat_cols=params.get("cat_cols", ["Gender", "Phakic/Pseudophakic"]), + ) + else: + raise ValueError(f"Unsupported import class {spec.class_name!r}") + return resolved + + +def _build_imports(entries: Iterable[Dict[str, Any]]) -> Dict[str, ImportSpec]: + specs: Dict[str, ImportSpec] = {} + for entry in entries or []: + import_id = entry.get("id") + if not import_id: + continue + specs[import_id] = ImportSpec( + id=import_id, + class_name=entry.get("className", ""), + params=entry.get("params", {}) or {}, + ) + return specs + + +def _extract_transform_params(node: Dict[str, Any]) -> Dict[str, Any]: + return { + "transformType": node.get("transformType"), + "roiMaskSource": node.get("roiMaskSource"), + "roiScale": node.get("roiScale"), + "roiTargetSize": node.get("roiTargetSize"), + "roiFallback": node.get("roiFallback"), + "centerCropSize": node.get("centerCropSize"), + "jitterHFlip": node.get("jitterHFlip"), + "jitterVFlip": node.get("jitterVFlip"), + "jitterRotation": node.get("jitterRotation"), + "jitterColorEnabled": node.get("jitterColorEnabled"), + "jitterColor": node.get("jitterColor"), + "resizeSize": node.get("resizeSize"), + } + + +def _extract_tower_params(node: Dict[str, Any]) -> Dict[str, Any]: + if node.get("towerType") == "clinical data": + return { + "hidden_dim": node.get("mdHiddenDim"), + "dropout": node.get("mdDropout"), + "use_se": node.get("mdUseSe"), + "se_reduction": node.get("mdSeReduction"), + "se_pre_norm": node.get("mdSePreNorm"), + "freeze_ratio": node.get("mdFreezeRatio"), + } + return { + "backbone": node.get("imageBackbone"), + "freeze_ratio": node.get("imageFreezeRatio"), + "augment": node.get("imageAugment"), + "geometry_dim": node.get("imageGeometryDim"), + "use_se": node.get("imageUseSe"), + "se_reduction": node.get("imageSeReduction"), + "se_pre_norm": node.get("imageSePreNorm"), + } + + +def _extract_bridge_params(node: Dict[str, Any]) -> Dict[str, Any]: + return { + "fusion_dim": node.get("bridgeFusionDim"), + "use_se": node.get("bridgeUseSe"), + "se_reduction": node.get("bridgeSeReduction"), + "se_pre_norm": node.get("bridgeSePreNorm"), + } + + +def _edge_from(edge: Dict[str, Any]) -> Optional[str]: + return edge.get("from") or edge.get("source") + + +def _edge_to(edge: Dict[str, Any]) -> Optional[str]: + return edge.get("to") or edge.get("target") + + +def _upstream_chain(start_id: str, nodes: Dict[str, Dict[str, Any]], edges: List[Dict[str, Any]]) -> List[str]: + chain: List[str] = [] + visited = set() + current = start_id + while True: + if current in visited: + break + visited.add(current) + incoming = [edge for edge in edges if _edge_to(edge) == current] + if not incoming: + break + # prefer first incoming edge for now + current = _edge_from(incoming[0]) + if not current: + break + chain.append(current) + node = nodes.get(current) + if node and node.get("type") == "data": + break + return chain diff --git a/v3/classes/croppers.py b/v3/classes/croppers.py new file mode 100644 index 0000000..fe91273 --- /dev/null +++ b/v3/classes/croppers.py @@ -0,0 +1,418 @@ +"""Optic-disc image croppers and preprocessor factory for V2.""" +from __future__ import annotations + +from pathlib import Path +from typing import Dict, Optional, Tuple + +import numpy as np +import pandas as pd +import torch +from PIL import Image, ImageDraw +from torchvision import transforms + +from v3.classes.geometry_features import compute_geometry_features, disc_cup_from_mask_image +from v3.classes.unet_segmenter import UNetSegmenter + + +def _geometry_from_mask(mask: np.ndarray, scale: float) -> Dict: + mask = np.asarray(mask) > 0 + coords = np.argwhere(mask) + if coords.size == 0: + raise RuntimeError("Empty mask; cannot derive geometry") + ys, xs = coords[:, 0], coords[:, 1] + centre_x = float(xs.mean()) + centre_y = float(ys.mean()) + width = float(xs.max() - xs.min()) + height = float(ys.max() - ys.min()) + diameter = max(width, height) + radius = diameter / 2.0 + crop_radius = radius * scale + return { + "centre_x": centre_x, + "centre_y": centre_y, + "radius": radius, + "crop_radius": crop_radius, + "crop_size": crop_radius * 2.0, + } + + +class UNetImageCropper: + def __init__( + self, + manifest_path: Path, + weights_path: Path, + normalize: str = "per_image", + threshold: float = 0.5, + tta: bool = False, + scale: float = 2.5, + target_size: int = 224, + cache_dir: Optional[Path] = None, + ) -> None: + self.segmenter = UNetSegmenter( + manifest_path=manifest_path, + normalize=normalize, + ) + state = torch.load(weights_path, map_location=self.segmenter.device) + state_dict = state.get("model", state) + self.segmenter.model.load_state_dict(state_dict) + self.segmenter.model.to(self.segmenter.device) + self.segmenter.model.eval() + + self.threshold = threshold + self.tta = tta + self.scale = scale + self.target_size = target_size + self.cache_dir = Path(cache_dir) if cache_dir is not None else None + if self.cache_dir is not None: + self.cache_dir.mkdir(parents=True, exist_ok=True) + + self.to_tensor = transforms.ToTensor() + + def _cache_path(self, image_path: Path) -> Optional[Path]: + if self.cache_dir is None: + return None + stem = image_path.stem + return self.cache_dir / f"{stem}_s{int(self.scale * 100)}.npz" + + def clear_cache(self) -> None: + if self.cache_dir is None or not self.cache_dir.exists(): + return + removed = sum(1 for f in self.cache_dir.glob("*.npz") if f.unlink() or True) + print(f"[UNetImageCropper] Cleared {removed} cached crop files from {self.cache_dir}") + + def _infer_masks(self, image: Image.Image) -> Optional[Tuple[np.ndarray, np.ndarray]]: + resized = self.segmenter.preprocess_image(image) + tensor = self.segmenter._normalize_tensor( + self.to_tensor(resized).to(self.segmenter.device) + ).unsqueeze(0) + + with torch.no_grad(): + logits = self.segmenter.model(tensor) + if self.tta: + t_h = torch.flip(tensor, dims=[3]) + log_h = self.segmenter.model(t_h) + log_h = torch.flip(log_h, dims=[3]) + t_v = torch.flip(tensor, dims=[2]) + log_v = self.segmenter.model(t_v) + log_v = torch.flip(log_v, dims=[2]) + logits = (logits + log_h + log_v) / 3.0 + probs = torch.sigmoid(logits)[0].cpu().numpy() + + disc_pred = (probs[0] > self.threshold).astype(np.uint8) * 255 + cup_pred = (probs[1] > self.threshold).astype(np.uint8) * 255 + disc_img = Image.fromarray(disc_pred, mode="L").resize(image.size, Image.NEAREST) + disc_mask = np.array(disc_img, dtype=np.uint8) + cup_img = Image.fromarray(cup_pred, mode="L").resize(image.size, Image.NEAREST) + cup_mask = (np.array(cup_img, dtype=np.uint8) > 0).astype(np.uint8) + cup_mask = (cup_mask > 0) & (disc_mask > 0) + cup_mask = cup_mask.astype(np.uint8) + disc_mask = (disc_mask > 0).astype(np.uint8) + return disc_mask, cup_mask + + def _compute_crop_info(self, image: Image.Image, image_path: Path) -> Optional[dict]: + image_path = Path(image_path).resolve() + cache_path = self._cache_path(image_path) + cached_bounds = None + if cache_path is not None and cache_path.exists(): + data = np.load(cache_path, allow_pickle=False) + try: + cached_bounds = { + "left": float(data["left"]), + "upper": float(data["upper"]), + "right": float(data["right"]), + "lower": float(data["lower"]), + } + if "features" in data.files: + cached_bounds["features"] = data["features"].astype(np.float32) + return cached_bounds + except KeyError: + cached_bounds = None + + masks = self._infer_masks(image) + if masks is None: + return cached_bounds + disc_mask, cup_mask = masks + try: + geom = _geometry_from_mask(disc_mask, self.scale) + except Exception: + return cached_bounds + cx = geom["centre_x"] + cy = geom["centre_y"] + r = geom["crop_radius"] + left = max(0.0, cx - r) + upper = max(0.0, cy - r) + right = min(float(image.width), cx + r) + lower = min(float(image.height), cy + r) + features = compute_geometry_features(disc_mask, cup_mask) + + info = { + "left": left, + "upper": upper, + "right": right, + "lower": lower, + "features": features, + } + if cache_path is not None: + np.savez( + cache_path, + left=left, + upper=upper, + right=right, + lower=lower, + width=float(image.width), + height=float(image.height), + scale=self.scale, + target_size=self.target_size, + features=features, + ) + return info + + def __call__(self, image: Image.Image, image_path: Path) -> Image.Image: + info = self._compute_crop_info(image, image_path) + if info is None: + return image + left = info["left"] + upper = info["upper"] + right = info["right"] + lower = info["lower"] + if right <= left or lower <= upper: + return image + crop = image.crop((left, upper, right, lower)) + return crop.resize((self.target_size, self.target_size), Image.BILINEAR) + + def geometry_features(self, image: Image.Image, image_path: Path) -> Optional[np.ndarray]: + info = self._compute_crop_info(image, image_path) + if info is None: + return None + features = info.get("features") + if features is None: + return None + return np.asarray(features, dtype=np.float32) + + +class ManifestImageCropper: + def __init__( + self, + manifest_path: Path, + scale: float = 2.5, + target_size: int = 224, + cache_dir: Optional[Path] = None, + ) -> None: + self.scale = scale + self.target_size = target_size + self.cache_dir = Path(cache_dir) if cache_dir is not None else None + if self.cache_dir is not None: + self.cache_dir.mkdir(parents=True, exist_ok=True) + + df = pd.read_csv(manifest_path) + self.entries: Dict[str, dict] = {} + for _, row in df.iterrows(): + img_path = Path(row["image_path"]).resolve() + self.entries[str(img_path)] = { + "annotation_disc": row.get("annotation_disc"), + "annotation_cup": row.get("annotation_cup"), + "annotation_type_disc": row.get("annotation_type_disc"), + "annotation_type_cup": row.get("annotation_type_cup"), + } + + def _cache_path(self, image_path: Path) -> Optional[Path]: + if self.cache_dir is None: + return None + return self.cache_dir / f"{image_path.stem}_s{int(self.scale * 100)}.npz" + + def clear_cache(self) -> None: + if self.cache_dir is None or not self.cache_dir.exists(): + return + removed = sum(1 for f in self.cache_dir.glob("*.npz") if f.unlink() or True) + print(f"[ManifestImageCropper] Cleared {removed} cached crop files from {self.cache_dir}") + + @staticmethod + def _load_contour(path: Path) -> np.ndarray: + coords = np.loadtxt(path) + if coords.ndim == 1: + coords = coords.reshape(-1, 2) + return coords + + @staticmethod + def _contour_to_mask(coords: np.ndarray, size: tuple[int, int]) -> np.ndarray: + if coords is None or coords.size == 0: + return np.zeros((size[1], size[0]), dtype=np.uint8) + img = Image.new("L", size, 0) + draw = ImageDraw.Draw(img) + points = [tuple(map(float, pt)) for pt in coords] + draw.polygon(points, outline=1, fill=1) + return np.array(img, dtype=np.uint8) + + def _load_masks(self, entry: dict, image: Image.Image) -> Optional[Tuple[np.ndarray, np.ndarray]]: + disc_path = entry.get("annotation_disc") + cup_path = entry.get("annotation_cup") + disc_type = (entry.get("annotation_type_disc") or "").lower() + cup_type = (entry.get("annotation_type_cup") or "").lower() + + disc_mask: Optional[np.ndarray] = None + cup_mask: Optional[np.ndarray] = None + + if disc_path and not pd.isna(disc_path): + disc_path = Path(disc_path) + try: + if disc_type == "mask": + mask_img = Image.open(disc_path) + mask_img = mask_img.resize(image.size, Image.NEAREST) + disc_mask, cup_from_mask = disc_cup_from_mask_image(mask_img) + if cup_from_mask.sum() > 0: + cup_mask = cup_from_mask + elif disc_type == "contour": + coords = self._load_contour(disc_path) + disc_mask = self._contour_to_mask(coords, image.size) + except Exception: + disc_mask = None + + if cup_mask is None and cup_path and not pd.isna(cup_path): + cup_path = Path(cup_path) + try: + if cup_type == "mask": + mask_img = Image.open(cup_path) + mask_img = mask_img.resize(image.size, Image.NEAREST) + _, cup_mask = disc_cup_from_mask_image(mask_img) + elif cup_type == "contour": + coords = self._load_contour(cup_path) + cup_mask = self._contour_to_mask(coords, image.size) + except Exception: + cup_mask = None + + if disc_mask is None: + return None + disc_mask = (disc_mask > 0).astype(np.uint8) + if cup_mask is None: + cup_mask = np.zeros_like(disc_mask, dtype=np.uint8) + cup_mask = ((cup_mask > 0) & (disc_mask > 0)).astype(np.uint8) + return disc_mask, cup_mask + + def _compute_crop_info(self, image: Image.Image, image_path: Path) -> Optional[dict]: + image_path = Path(image_path).resolve() + entry = self.entries.get(str(image_path)) + if entry is None: + return None + cache_path = self._cache_path(image_path) + cached_bounds = None + if cache_path is not None and cache_path.exists(): + data = np.load(cache_path, allow_pickle=False) + try: + cached_bounds = { + "left": float(data["left"]), + "upper": float(data["upper"]), + "right": float(data["right"]), + "lower": float(data["lower"]), + } + if "features" in data.files: + cached_bounds["features"] = data["features"].astype(np.float32) + return cached_bounds + except KeyError: + cached_bounds = None + + masks = self._load_masks(entry, image) + if masks is None: + return cached_bounds + disc_mask, cup_mask = masks + try: + geom = _geometry_from_mask(disc_mask, self.scale) + except Exception: + return cached_bounds + cx = geom["centre_x"] + cy = geom["centre_y"] + r = geom["crop_radius"] + left = max(0.0, cx - r) + upper = max(0.0, cy - r) + right = min(float(image.width), cx + r) + lower = min(float(image.height), cy + r) + features = compute_geometry_features(disc_mask, cup_mask) + + info = { + "left": left, + "upper": upper, + "right": right, + "lower": lower, + "features": features, + } + if cache_path is not None: + np.savez( + cache_path, + left=left, + upper=upper, + right=right, + lower=lower, + width=float(image.width), + height=float(image.height), + scale=self.scale, + target_size=self.target_size, + features=features, + ) + return info + + def __call__(self, image: Image.Image, image_path: Path) -> Image.Image: + info = self._compute_crop_info(image, image_path) + if info is None: + return image + left = info["left"] + upper = info["upper"] + right = info["right"] + lower = info["lower"] + if right <= left or lower <= upper: + return image + crop = image.crop((left, upper, right, lower)) + return crop.resize((self.target_size, self.target_size), Image.BILINEAR) + + def geometry_features(self, image: Image.Image, image_path: Path) -> Optional[np.ndarray]: + info = self._compute_crop_info(image, image_path) + if info is None: + return None + features = info.get("features") + if features is None: + return None + return np.asarray(features, dtype=np.float32) + + +# --------------------------------------------------------------------------- +# Factory +# --------------------------------------------------------------------------- + +def build_image_preprocessor_from_args(args): + """Construct the correct image cropper from CLI args, or return None.""" + crop_manifest = getattr(args, "img_crop_manifest", None) + crop_weights = getattr(args, "img_crop_weights", None) + use_gt = bool(getattr(args, "img_crop_gt", False)) + if not crop_manifest: + return None + crop_cache = Path(getattr(args, "img_crop_cache", Path("cache_data/hypertower_crops"))) + persist_cache = bool(getattr(args, "persist_img_crop_cache", False)) + if use_gt: + pre = ManifestImageCropper( + manifest_path=Path(crop_manifest), + scale=getattr(args, "img_crop_scale", 2.5), + target_size=getattr(args, "img_crop_size", 224), + cache_dir=crop_cache, + ) + if not persist_cache: + pre.clear_cache() + print(f"[V2 modes] GT disc cropper enabled -> cache at {crop_cache}", flush=True) + return pre + if crop_weights: + pre = UNetImageCropper( + manifest_path=Path(crop_manifest), + weights_path=Path(crop_weights), + normalize=getattr(args, "img_crop_normalize", "per_image"), + threshold=getattr(args, "img_crop_threshold", 0.5), + tta=getattr(args, "img_crop_tta", False), + scale=getattr(args, "img_crop_scale", 2.5), + target_size=getattr(args, "img_crop_size", 224), + cache_dir=crop_cache, + ) + if not persist_cache: + pre.clear_cache() + print(f"[V2 modes] UNet disc cropper enabled -> cache at {crop_cache}", flush=True) + return pre + print( + "[V2 modes] img_crop_manifest provided but no --img-crop-gt or --img-crop-weights; cropping disabled.", + flush=True, + ) + return None diff --git a/v3/classes/data_bundle.py b/v3/classes/data_bundle.py new file mode 100644 index 0000000..f9c9767 --- /dev/null +++ b/v3/classes/data_bundle.py @@ -0,0 +1,241 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Callable, Dict, Iterable, List, Optional, Tuple + +import numpy as np +import pandas as pd + + +class DataBundle: + """ + Generic, torch-free container for metadata and file/label bookkeeping. + + Keeps feature typing, vectorization, and patient-level splits generic. + Dataset-specific preprocessing (e.g., eye canonicalization) should live + in the dataset builder (e.g., papila_builders in v2). + """ + + def __init__( + self, + *, + image_dir: str, + clinical_dir: Optional[str] = None, + label_col: str, + patient_col: str = "Patient ID", + cat_cols: Optional[Iterable[str]] = None, + max_unique_for_cat: int = 4, + n_splits: int = 5, + random_seed: int = 42, + filename_template: str = "RET{pid:03d}{eye}.jpg", + image_path_fn: Optional[Callable[[pd.Series], Path]] = None, + ) -> None: + self.image_dir = Path(image_dir) + self.label_col = label_col + self.patient_col = patient_col + self.max_unique_for_cat = max_unique_for_cat + self.n_splits = n_splits + self.filename_template = filename_template + self.image_path_fn = image_path_fn + self.clinical_dir = Path(clinical_dir) if clinical_dir else None + + # Internal state + self.frames: List[pd.DataFrame] = [] + self.df: pd.DataFrame = pd.DataFrame() + self.scalar_cols: List[str] = [] + self.cat_cols: List[str] = list(cat_cols) if cat_cols is not None else [] + self.scalar_stats: Dict[str, Dict[str, float]] = {} + self.cat_maps: Dict[str, Dict[object, int]] = {} + self.feature_dim: int = 0 + self.folds: Dict[int, Dict[str, List[object]]] = {} + self.random_seed = int(random_seed) + + # ------------------- Public API ------------------- + def add_df( + self, + df: pd.DataFrame, + *, + id_column: Optional[str] = None, + exclude_cols: Optional[Iterable[str]] = None, + ) -> None: + """ + Add a dataframe and re-run typing, stats, and K-fold indices. + QC rules: + - Must have patient ID column; if not provided under that name, specify id_column. + """ + df = df.copy() + self._ensure_patient_id(df, id_column) + if self.label_col not in df.columns: + raise ValueError(f"label_col '{self.label_col}' not found in added dataframe") + + self.frames.append(df) + self._refresh_master_df(exclude_cols=exclude_cols) + self._infer_or_validate_feature_types(exclude_cols=exclude_cols) + self._compute_numeric_stats() + self._build_cat_maps() + self._compute_feature_dim() + self._build_kfold_indices() + + def get_split_ids(self, fold: int) -> Tuple[List[object], List[object]]: + rec = self.folds.get(fold) + if not rec: + raise KeyError(f"Fold {fold} not available. Built folds: {sorted(self.folds.keys())}") + return rec["train_ids"], rec["test_ids"] + + def get_split_dfs(self, fold: int) -> Tuple[pd.DataFrame, pd.DataFrame]: + train_ids, test_ids = self.get_split_ids(fold) + train_df = self.df[self.df[self.patient_col].isin(train_ids)].reset_index(drop=True) + test_df = self.df[self.df[self.patient_col].isin(test_ids)].reset_index(drop=True) + return train_df, test_df + + def vectorize_row(self, row: pd.Series) -> np.ndarray: + """Return a numpy feature vector (torch-free).""" + feats: List[float] = [] + miss: List[float] = [] + # numeric + for col in self.scalar_cols: + v = pd.to_numeric(row.get(col), errors="coerce") + if pd.isna(v): + miss.append(1.0) + v = self.scalar_stats[col]["median"] + else: + miss.append(0.0) + lo = self.scalar_stats[col]["min"] + hi = self.scalar_stats[col]["max"] + feats.append((float(v) - lo) / (hi - lo) if hi > lo else 0.0) + # categorical + for col in self.cat_cols: + mapping = self.cat_maps[col] + one = [0.0] * len(mapping) + key = row.get(col) + one[mapping.get(key, 0)] = 1.0 # 0 is + feats.extend(one) + # numeric missing flags + feats.extend(miss) + return np.asarray(feats, dtype=np.float32) + + def get_image_path(self, row: pd.Series) -> Path: + if self.image_path_fn is not None: + return Path(self.image_path_fn(row)) + pid = int(row[self.patient_col]) + eye = row.get("eyeID", "") + if eye in ("OS", "OD"): + eye_str = eye + else: + eye_str = str(eye) + return self.image_dir / self.filename_template.format(pid=pid, eye=eye_str) + + def encode_metadata(self, row: pd.Series) -> np.ndarray: + return self.vectorize_row(row) + + def get_label(self, row: pd.Series) -> int: + return int(row[self.label_col]) + + # ------------------- Internal helpers ------------------- + def _ensure_patient_id(self, df: pd.DataFrame, id_column: Optional[str]) -> None: + if self.patient_col in df.columns: + return + if id_column and id_column in df.columns: + df.rename(columns={id_column: self.patient_col}, inplace=True) + return + candidates = [ + c + for c in df.columns + if c.lower().replace(" ", "") in {"patientid", "patient", "pid"} + ] + if len(candidates) == 1: + df.rename(columns={candidates[0]: self.patient_col}, inplace=True) + return + raise ValueError( + f"A '{self.patient_col}' column is required; provide id_column=... if it has a different name." + ) + + def _refresh_master_df(self, exclude_cols: Optional[Iterable[str]] = None) -> None: + self.df = pd.concat(self.frames, axis=0, ignore_index=True) + if exclude_cols: + self.df = self.df.drop(columns=[c for c in exclude_cols if c in self.df.columns]) + + def _infer_or_validate_feature_types(self, exclude_cols: Optional[Iterable[str]] = None) -> None: + excluded = set(exclude_cols or []) | {self.label_col, self.patient_col} + feature_candidates = [c for c in self.df.columns if c not in excluded] + cats = set(self.cat_cols) if self.cat_cols else set() + scalars = set() + for c in feature_candidates: + if c in cats: + continue + s = self.df[c] + as_num = pd.to_numeric(s, errors="coerce") + num_missing = as_num.isna().mean() + num_unique = s.dropna().nunique() + if as_num.notna().any() and num_missing < 1.0 and num_unique > self.max_unique_for_cat: + scalars.add(c) + else: + if num_unique <= self.max_unique_for_cat or as_num.isna().mean() > 0.0: + cats.add(c) + else: + scalars.add(c) + self.cat_cols = sorted(cats) + self.scalar_cols = sorted(scalars) + + def _compute_numeric_stats(self) -> None: + self.scalar_stats.clear() + for col in self.scalar_cols: + s = pd.to_numeric(self.df[col], errors="coerce") + vals = s.dropna().astype(float).values + if vals.size == 0: + lo, hi, med = 0.0, 1.0, 0.0 + else: + lo, hi = float(np.min(vals)), float(np.max(vals)) + med = float(np.median(vals)) + if hi <= lo: + hi = lo + 1.0 + self.scalar_stats[col] = {"min": lo, "max": hi, "median": med} + + def _build_cat_maps(self) -> None: + self.cat_maps.clear() + for col in self.cat_cols: + cats = [v for v in self.df[col].dropna().unique().tolist()] + try: + cats = sorted(cats) + except Exception: + pass + mapping = {"": 0} + for i, v in enumerate(cats, start=1): + mapping[v] = i + self.cat_maps[col] = mapping + + def _compute_feature_dim(self) -> None: + self.feature_dim = len(self.scalar_cols) + sum(len(m) for m in self.cat_maps.values()) + len(self.scalar_cols) + + # ------------------- K-fold on unique patients ------------------- + def _build_kfold_indices(self) -> None: + pats = self.df[self.patient_col].unique().tolist() + labels_by_pat: Dict[object, object] = {} + for pid, grp in self.df.groupby(self.patient_col): + lab = grp[self.label_col].dropna() + if len(lab) == 0: + labels_by_pat[pid] = 0 + else: + labels_by_pat[pid] = lab.mode().iloc[0] + y_pat = np.array([labels_by_pat[p] for p in pats]) + + try: + from sklearn.model_selection import StratifiedGroupKFold + + sgkf = StratifiedGroupKFold( + n_splits=self.n_splits, shuffle=True, random_state=self.random_seed + ) + split_iter = sgkf.split(X=pats, y=y_pat, groups=pats) + except Exception: + from sklearn.model_selection import StratifiedKFold + + skf = StratifiedKFold( + n_splits=self.n_splits, shuffle=True, random_state=self.random_seed + ) + split_iter = skf.split(X=np.zeros(len(pats)), y=y_pat) + + self.folds.clear() + for i, (train_idx, test_idx) in enumerate(split_iter): + train_ids = [pats[j] for j in train_idx] + test_ids = [pats[j] for j in test_idx] + self.folds[i] = {"train_ids": train_ids, "test_ids": test_ids} diff --git a/v3/classes/dataset.py b/v3/classes/dataset.py new file mode 100644 index 0000000..c8c2f1e --- /dev/null +++ b/v3/classes/dataset.py @@ -0,0 +1,115 @@ +from torch.utils.data import Dataset +from PIL import Image +import numpy as np +import torch + + +class ClinicalDataset(Dataset): + """Generic dataset wrapping a DataBundle-like instance. + Returns (img_tensor, meta_tensor, label).""" + + def __init__( + self, + clinical_data, + img_transform, + meta_transform=None, + image_preprocessor=None, + geometry_provider=None, + geometry_dim: int = 0, + image_cache: "dict | None" = None, + ): + self.clinical = clinical_data + self.transform_image = img_transform + self.meta_transform = meta_transform or (lambda x: x) + self.image_preprocessor = image_preprocessor + self.geometry_provider = geometry_provider + self.geometry_dim = geometry_dim if geometry_provider is not None else 0 + self.image_cache = image_cache + + def __len__(self): + return len(self.clinical.df) + + def __getitem__(self, idx: int): + row = self.clinical.df.iloc[idx] + # load & transform image + img_path = self.clinical.get_image_path(row) + cache_key = str(img_path) + if self.image_cache is not None and cache_key in self.image_cache: + orig_img = Image.fromarray(self.image_cache[cache_key]) + else: + orig_img = Image.open(img_path).convert("RGB") + if self.image_cache is not None: + self.image_cache[cache_key] = np.asarray(orig_img, dtype=np.uint8) + img = orig_img + if self.image_preprocessor is not None: + img = self.image_preprocessor(img, img_path) + img_t = self.transform_image(img) + # encode & transform metadata + meta = self.clinical.encode_metadata(row) + meta_t = self.meta_transform(meta) + # label + label = self.clinical.get_label(row) + if self.geometry_dim > 0: + features = None + if self.geometry_provider is not None and hasattr(self.geometry_provider, "geometry_features"): + features = self.geometry_provider.geometry_features(orig_img, img_path) + if features is None: + geom_vec = torch.zeros(self.geometry_dim, dtype=torch.float32) + else: + features = np.asarray(features, dtype=np.float32) + if features.shape[0] != self.geometry_dim: + geom_vec = torch.zeros(self.geometry_dim, dtype=torch.float32) + else: + geom_vec = torch.from_numpy(features) + return img_t, meta_t, geom_vec, label + return img_t, meta_t, label + + +# --------------------------------------------------------------------------- +# _ClinicalView — shim used by V2HyperTower._run_fold +# --------------------------------------------------------------------------- + +from .data_bundle import DataBundle # noqa: E402 + + +class _ClinicalView: + """Minimal shim so ClinicalDataset can iterate an epoch-specific DataFrame + while still delegating encoding/paths/labels to the DataBundle object.""" + + def __init__(self, base: DataBundle, df): + self.base = base + self.df = df + + @property + def image_dir(self): + return self.base.image_dir + + @property + def clinical_dir(self): + return self.base.clinical_dir + + @property + def id_cols(self): + return ("Patient ID", "eyeID") + + @property + def label_col(self): + return self.base.label_col + + @property + def filename_template(self): + return getattr(self.base, "filename_template", "RET{pid:03d}{eye}.jpg") + + @property + def dim(self): + return self.base.feature_dim + + def encode_metadata(self, row): + vec = self.base.vectorize_row(row) + return torch.as_tensor(vec, dtype=torch.float32) + + def get_image_path(self, row): + return self.base.get_image_path(row) + + def get_label(self, row): + return int(row[self.base.label_col]) diff --git a/v3/classes/filters.py b/v3/classes/filters.py new file mode 100644 index 0000000..96b6922 --- /dev/null +++ b/v3/classes/filters.py @@ -0,0 +1,119 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Iterable, List, Sequence, Tuple, Union +import re + +import pandas as pd + + +@dataclass +class RegexFilter: + pattern: str + flags: int = 0 + + def apply_paths(self, paths: Sequence[str]) -> Tuple[List[str], List[str]]: + if not self.pattern: + return list(paths), [] + try: + regex = re.compile(self.pattern, self.flags) + except re.error as err: + return list(paths), [f'Invalid regex "{self.pattern}": {err}'] + filtered = [p for p in paths if regex.search(p)] + return filtered, [] + + +@dataclass +class ColumnFilter: + column: str + operator: str + value: str + case_insensitive: bool = True + + def apply_df(self, df: pd.DataFrame) -> Tuple[pd.DataFrame, List[str]]: + warnings: List[str] = [] + if not self.column: + return df, ["Column filter missing column name."] + columns = list(df.columns) + col_index = _resolve_column_index(columns, self.column, warnings) + if col_index is None: + return df, warnings + col_name = columns[col_index] + if self.value is None or self.value == "": + return df, [f'Column filter "{self.column}" missing value.'] + series = df[col_name] + mask = series.apply( + lambda cell: compare_cell( + cell, self.value, self.operator, case_insensitive=self.case_insensitive + ) + ) + return df[mask], warnings + + +FilterSpec = Union[RegexFilter, ColumnFilter] + + +def apply_regex_filters(paths: Sequence[str], filters: Iterable[RegexFilter]) -> Tuple[List[str], List[str]]: + filtered = list(paths) + warnings: List[str] = [] + for filt in filters: + filtered, warn = filt.apply_paths(filtered) + warnings.extend(warn) + return filtered, warnings + + +def apply_column_filters(df: pd.DataFrame, filters: Iterable[ColumnFilter]) -> Tuple[pd.DataFrame, List[str]]: + filtered = df + warnings: List[str] = [] + for filt in filters: + filtered, warn = filt.apply_df(filtered) + warnings.extend(warn) + return filtered, warnings + + +def compare_cell(cell, raw_value: str, operator: str, case_insensitive: bool = True) -> bool: + cell_str = "" if cell is None else str(cell).strip() + value_str = "" if raw_value is None else str(raw_value).strip() + if case_insensitive: + cell_str = cell_str.lower() + value_str = value_str.lower() + if operator == "=": + return cell_str == value_str + if operator == "!=": + return cell_str != value_str + cell_num = _to_float(cell_str) + value_num = _to_float(value_str) + if cell_num is None or value_num is None: + return False + if operator == ">": + return cell_num > value_num + if operator == ">=": + return cell_num >= value_num + if operator == "<": + return cell_num < value_num + if operator == "<=": + return cell_num <= value_num + return False + + +def _resolve_column_index(columns: Sequence[str], column: str, warnings: List[str]) -> int | None: + try: + return columns.index(column) + except ValueError: + lower = column.lower() + matches = [idx for idx, col in enumerate(columns) if str(col).lower() == lower] + if matches: + if len(matches) > 1: + warnings.append( + f'Column "{column}" matched multiple headers; using "{columns[matches[0]]}".' + ) + return matches[0] + warnings.append(f'Column "{column}" not found.') + return None + + +def _to_float(value: str) -> float | None: + try: + return float(value) + except (TypeError, ValueError): + return None diff --git a/v3/classes/geometry_features.py b/v3/classes/geometry_features.py new file mode 100755 index 0000000..6539cae --- /dev/null +++ b/v3/classes/geometry_features.py @@ -0,0 +1,87 @@ +"""Shared helpers for deriving disc/cup geometry features.""" + +from __future__ import annotations + +from collections import Counter +from typing import Tuple + +import numpy as np +from PIL import Image + +EPS = 1e-6 +FEATURE_DIM = 5 + + +def disc_cup_from_mask_image(mask_img: Image.Image) -> Tuple[np.ndarray, np.ndarray]: + """Return binary disc/cup masks from a REFUGE-style annotation image.""" + arr = np.asarray(mask_img) + if arr.ndim == 3: + h, w, c = arr.shape + border = np.concatenate( + [arr[0, :, :], arr[-1, :, :], arr[:, 0, :], arr[:, -1, :]], + axis=0, + ) + border_counts = Counter(map(tuple, border)) + bg_color = border_counts.most_common(1)[0][0] + flat = arr.reshape(-1, c) + colors = Counter(map(tuple, flat)) + colors.pop(bg_color, None) + disc = (~np.all(arr == bg_color, axis=-1)).astype(np.uint8) + if colors: + cup_color = min(colors.keys(), key=lambda col: sum(col)) + cup = np.all(arr == cup_color, axis=-1).astype(np.uint8) + else: + cup = np.zeros((h, w), dtype=np.uint8) + else: + border = np.concatenate([arr[0, :], arr[-1, :], arr[:, 0], arr[:, -1]]) + counts = Counter(border.tolist()) + bg_value = counts.most_common(1)[0][0] + disc = (arr != bg_value).astype(np.uint8) + fg = arr[arr != bg_value] + if fg.size > 0: + cup_value = int(np.min(fg)) + cup = (arr == cup_value).astype(np.uint8) + else: + cup = np.zeros_like(arr, dtype=np.uint8) + cup = (cup > 0) & (disc > 0) + return disc.astype(np.uint8), cup.astype(np.uint8) + + +def compute_geometry_features(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray: + """Compute cup/disc geometry descriptors (area, rim, diameter ratios, centre shift).""" + disc = (disc_mask > 0).astype(np.float32) + cup = (cup_mask > 0).astype(np.float32) + + disc_area = disc.sum() + cup_area = cup.sum() + area_ratio = cup_area / (disc_area + EPS) + rim_ratio = (disc_area - cup_area) / (disc_area + EPS) + + disc_rows = np.any(disc > 0, axis=1) + cup_rows = np.any(cup > 0, axis=1) + disc_cols = np.any(disc > 0, axis=0) + cup_cols = np.any(cup > 0, axis=0) + + disc_height = float(disc_rows.sum()) + cup_height = float(cup_rows.sum()) + disc_width = float(disc_cols.sum()) + cup_width = float(cup_cols.sum()) + + vertical_ratio = cup_height / (disc_height + EPS) + horizontal_ratio = cup_width / (disc_width + EPS) + + def _centre(mask: np.ndarray) -> Tuple[float, float]: + coords = np.argwhere(mask > 0) + if coords.size == 0: + return 0.5, 0.5 + ys, xs = coords[:, 0], coords[:, 1] + return float(xs.mean()) / mask.shape[1], float(ys.mean()) / mask.shape[0] + + disc_cx, disc_cy = _centre(disc) + cup_cx, cup_cy = _centre(cup) + centre_shift = float(np.hypot(cup_cx - disc_cx, cup_cy - disc_cy)) + + return np.array( + [area_ratio, rim_ratio, vertical_ratio, horizontal_ratio, centre_shift], + dtype=np.float32, + ) diff --git a/v3/classes/hypertower_logger.py b/v3/classes/hypertower_logger.py new file mode 100644 index 0000000..c46d6f7 --- /dev/null +++ b/v3/classes/hypertower_logger.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +import csv +import json +import logging +from pathlib import Path +from typing import Optional + + +DEFAULT_OPTIONAL_EPOCH_COLS = [ + "pct_fused", + "pct_img", + "pct_md", + "phase", + "se_mean", + "se_std", + "se_pct_lt_0.2", + "se_pct_gt_0.8", + "holdout_loss", + "holdout_acc_fused", + "holdout_acc_img", + "holdout_acc_cd", + "holdout_auc_fused", + "holdout_auc_img", + "holdout_auc_cd", + "best_monitor", + "best_so_far", + "best_epoch", + "early_best_so_far", + "early_bad_epochs", + "early_improved", + "early_monitor", + "holdout_best_monitor", + "holdout_best_so_far", + "holdout_best_epoch", +] + + +class HypertowerLogger: + """ + Shared logging utility for V2 tower workflows. + - train.log line logging + - epoch_log.csv row logging with stable header + - lightweight JSON/array artifact helpers + """ + + def __init__( + self, + *, + run_dir: Path, + train_log_path: Optional[Path] = None, + epoch_log_path: Optional[Path] = None, + logger_name: Optional[str] = None, + ) -> None: + self.run_dir = Path(run_dir).resolve() + self.run_dir.mkdir(parents=True, exist_ok=True) + self.train_log_path = Path(train_log_path) if train_log_path else (self.run_dir / "train.log") + self.epoch_log_path = Path(epoch_log_path) if epoch_log_path else (self.run_dir / "epoch_log.csv") + + self._logger_name = logger_name or f"hypertower.{id(self)}" + self.logger = logging.getLogger(self._logger_name) + self.logger.setLevel(logging.INFO) + self.logger.handlers = [] + fh = logging.FileHandler(str(self.train_log_path)) + fh.setFormatter(logging.Formatter("%(asctime)s - %(message)s")) + self.logger.addHandler(fh) + self.logger.propagate = False + + self._epoch_log_fp = None + self._epoch_log_writer = None + self._epoch_log_fields: list[str] | None = None + + def info(self, msg: str) -> None: + self.logger.info(msg) + + def warning(self, msg: str) -> None: + self.logger.warning(msg) + + def error(self, msg: str) -> None: + self.logger.error(msg) + + def write_epoch_row( + self, + row: dict, + *, + path: str | Path | None = None, + optional_cols: Optional[list[str]] = None, + ) -> None: + optional = optional_cols if optional_cols is not None else DEFAULT_OPTIONAL_EPOCH_COLS + if self._epoch_log_writer is None: + fieldnames = list(dict.fromkeys([*row.keys(), *optional])) + target_path = Path(path) if path is not None else self.epoch_log_path + target_path.parent.mkdir(parents=True, exist_ok=True) + self._epoch_log_fp = open(target_path, "w", newline="", encoding="utf-8") + self._epoch_log_writer = csv.DictWriter(self._epoch_log_fp, fieldnames=fieldnames) + self._epoch_log_writer.writeheader() + self._epoch_log_fields = fieldnames + + assert self._epoch_log_fields is not None + assert self._epoch_log_writer is not None + assert self._epoch_log_fp is not None + for key in self._epoch_log_fields: + row.setdefault(key, None) + self._epoch_log_writer.writerow({k: row.get(k) for k in self._epoch_log_fields}) + self._epoch_log_fp.flush() + + def write_json(self, path: str | Path, payload: dict) -> None: + target = Path(path) + if not target.is_absolute(): + target = self.run_dir / target + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(json.dumps(payload, indent=2), encoding="utf-8") + + def close(self) -> None: + if self._epoch_log_fp is not None: + try: + self._epoch_log_fp.close() + except Exception: + pass + self._epoch_log_fp = None + self._epoch_log_writer = None + self._epoch_log_fields = None + for handler in list(self.logger.handlers): + try: + handler.close() + except Exception: + pass + self.logger.removeHandler(handler) diff --git a/v3/classes/image_loader.py b/v3/classes/image_loader.py new file mode 100644 index 0000000..0cf531e --- /dev/null +++ b/v3/classes/image_loader.py @@ -0,0 +1,149 @@ +""" +CachedImageLoader — shared image loading with optional in-memory cache. + +A single instance can be passed to any dataset (SlotDataset, EyeDataset, etc.) +so that images are decoded from disk at most once per run, regardless of how +many folds or loaders reference the same file. + +Usage: + loader = CachedImageLoader(enabled=True, workers=4) + loader.warm(paths, preprocessor=my_crop_fn) # optional: parallel pre-fill + img = loader.load(path, preprocessor=my_crop_fn) # returns PIL Image +""" +from __future__ import annotations + +from concurrent.futures import ThreadPoolExecutor, as_completed +from pathlib import Path +from typing import Callable, Iterable, Optional + +import numpy as np +from PIL import Image + + +class CachedImageLoader: + """ + Loads PIL Images from disk with an optional shared in-memory cache. + + The cache stores decoded, pre-preprocessed images as uint8 numpy arrays + (RGB, HWC). Storing after preprocessing means the preprocessor (e.g. + cropper) runs only once per image across all folds and epochs. + + Parameters + ---------- + enabled : bool + When False the cache is disabled and every call hits disk. + workers : int + Number of threads used by warm(). 0 or 1 = single-threaded. + """ + + def __init__(self, *, enabled: bool = True, workers: int = 4) -> None: + self._cache: dict[str, np.ndarray] | None = {} if enabled else None + self._workers = workers + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + @property + def enabled(self) -> bool: + return self._cache is not None + + def __len__(self) -> int: + return len(self._cache) if self._cache is not None else 0 + + def load( + self, + path: str | Path, + preprocessor: Optional[Callable[..., Image.Image]] = None, + ) -> Image.Image: + """ + Return a PIL Image for *path*. + + If the cache is enabled the image is stored after preprocessing so + the preprocessor only runs once. Subsequent calls return a copy + reconstructed from the cached array. + """ + key = str(path) + + if self._cache is not None: + cached = self._cache.get(key) + if cached is not None: + return Image.fromarray(cached, mode="RGB") + + img = Image.open(path).convert("RGB") + if preprocessor is not None: + img = _call_preprocessor(preprocessor, img, path) + + if self._cache is not None: + self._cache[key] = np.asarray(img, dtype=np.uint8) + + return img + + def warm( + self, + paths: Iterable[str | Path], + preprocessor: Optional[Callable[..., Image.Image]] = None, + ) -> None: + """ + Pre-populate the cache for all *paths* (no-op when disabled). + + Already-cached paths are skipped so calling warm() multiple times + (e.g. once per fold) is safe and only loads new images. + """ + if self._cache is None: + return + + to_warm = [str(p) for p in paths if str(p) not in self._cache] + if not to_warm: + return + + already = len(paths if isinstance(paths, (list, tuple)) else list(paths)) - len(to_warm) + print( + f"[image_cache] warming {len(to_warm)} images" + + (f" ({already} already cached)" if already else ""), + flush=True, + ) + + def _warm_one(path_str: str) -> None: + if path_str in self._cache: # guard against races + return + img = Image.open(path_str).convert("RGB") + if preprocessor is not None: + img = _call_preprocessor(preprocessor, img, Path(path_str)) + self._cache[path_str] = np.asarray(img, dtype=np.uint8) + + try: + from tqdm import tqdm + except ImportError: + tqdm = None + + if self._workers <= 1: + it = tqdm(to_warm, desc="Warm image cache", unit="img") if tqdm else to_warm + for p in it: + _warm_one(p) + else: + with ThreadPoolExecutor(max_workers=self._workers) as ex: + futures = {ex.submit(_warm_one, p): p for p in to_warm} + it = ( + tqdm(as_completed(futures), total=len(futures), + desc="Warm image cache", unit="img") + if tqdm else as_completed(futures) + ) + for fut in it: + fut.result() + + +# ------------------------------------------------------------------ +# Internal helpers +# ------------------------------------------------------------------ + +def _call_preprocessor( + fn: Callable[..., Image.Image], + img: Image.Image, + path: Path, +) -> Image.Image: + """Call preprocessor with (img, path) or just (img) depending on arity.""" + try: + return fn(img, path) + except TypeError: + return fn(img) diff --git a/v3/classes/loader_factory.py b/v3/classes/loader_factory.py new file mode 100644 index 0000000..f27d940 --- /dev/null +++ b/v3/classes/loader_factory.py @@ -0,0 +1,244 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable, Optional + +import torch +from torch.utils.data import DataLoader, WeightedRandomSampler + +from .network_manager import LoaderBundle, PatientSplit +from .slot_dataset import SlotDataset, slot_collate +from .profiles.base import SlotDescriptor, SimpleDatasetProfile + + +def _default_slot_descriptors(patient_col: str, label_col: str) -> dict[str, SlotDescriptor]: + return { + "id_1": SlotDescriptor( + key="id_1", + kind="id", + description=f"Patient identifier column ({patient_col})", + required=True, + shape_hint="scalar", + ), + "eye_id_1": SlotDescriptor( + key="eye_id_1", + kind="id", + description="Eye side identifier (OD/OS)", + required=False, + shape_hint="scalar", + ), + "label_1": SlotDescriptor( + key="label_1", + kind="label", + description=f"Label column ({label_col})", + required=True, + shape_hint="scalar", + ), + "image_1": SlotDescriptor( + key="image_1", + kind="image", + description="Primary image slot", + required=False, + shape_hint="HWC or CHW", + ), + "matrix_1": SlotDescriptor( + key="matrix_1", + kind="matrix", + description="Primary matrix slot", + required=False, + shape_hint="[feature_dim]", + ), + } + + +def _row_to_sample( + row: Any, + *, + clinical: Any, + patient_col: str, + label_col: str, +) -> dict[str, Any]: + return { + "id_1": row[patient_col], + "eye_id_1": str(row.get("eyeID", "")), + "label_1": row[label_col], + "image_1": clinical.get_image_path(row) if hasattr(clinical, "get_image_path") else None, + "matrix_1": clinical.vectorize_row(row) if hasattr(clinical, "vectorize_row") else None, + } + + +@dataclass +class SlotLoaderFactory: + """ + Generic loader factory that emits dict batches keyed by slot names. + """ + + image_transform: Optional[Callable] = None + matrix_transform: Optional[Callable] = None + num_workers: int = 0 + + def build( + self, + *, + clinical: Any, + split: PatientSplit, + args: Any, + fold: int, + profile: Optional[Any] = None, + ) -> LoaderBundle: + batch_size = int(getattr(args, "batch_size", 8)) + slot_desc = self._resolve_slot_descriptors(clinical=clinical, profile=profile) + + train_samples = self._build_samples(split.train, clinical, profile, slot_desc) + val_samples = self._build_samples(split.val, clinical, profile, slot_desc) + holdout_samples = ( + self._build_samples(split.holdout, clinical, profile, slot_desc) + if split.holdout is not None + else None + ) + + train_loader = DataLoader( + SlotDataset( + train_samples, + slot_desc, + image_transform=self.image_transform, + matrix_transform=self.matrix_transform, + ), + batch_size=batch_size, + shuffle=True, + num_workers=self.num_workers, + collate_fn=slot_collate, + ) + val_loader = DataLoader( + SlotDataset( + val_samples, + slot_desc, + image_transform=self.image_transform, + matrix_transform=self.matrix_transform, + ), + batch_size=batch_size, + shuffle=False, + num_workers=self.num_workers, + collate_fn=slot_collate, + ) + holdout_loader = None + if holdout_samples is not None: + holdout_loader = DataLoader( + SlotDataset( + holdout_samples, + slot_desc, + image_transform=self.image_transform, + matrix_transform=self.matrix_transform, + ), + batch_size=batch_size, + shuffle=False, + num_workers=self.num_workers, + collate_fn=slot_collate, + ) + return LoaderBundle(train=train_loader, val=val_loader, holdout=holdout_loader) + + @staticmethod + def _resolve_slot_descriptors( + *, + clinical: Any, + profile: Optional[Any], + ) -> dict[str, SlotDescriptor]: + if profile is not None and hasattr(profile, "slot_descriptors"): + return profile.slot_descriptors() + patient_col = getattr(clinical, "patient_col", "Patient ID") + label_col = getattr(clinical, "label_col", "Diagnosis") + return _default_slot_descriptors(patient_col, label_col) + + @staticmethod + def _build_samples( + df, + clinical: Any, + profile: Optional[Any], + slot_desc: dict[str, SlotDescriptor], + ) -> list[dict[str, Any]]: + if df is None or df.empty: + return [] + if profile is not None and hasattr(profile, "build_samples"): + return profile.build_samples(df=df, clinical=clinical) + + patient_col = getattr(profile, "patient_col", None) if profile is not None else None + label_col = getattr(profile, "label_col", None) if profile is not None else None + pcol = patient_col or "Patient ID" + lcol = label_col or getattr(clinical, "label_col", "Diagnosis") + samples = [] + for _, row in df.iterrows(): + sample = _row_to_sample(row, clinical=clinical, patient_col=pcol, label_col=lcol) + for key in slot_desc.keys(): + sample.setdefault(key, None) + samples.append(sample) + return samples + + +# --------------------------------------------------------------------------- +# V2 filter / loader helpers (used by V2HyperTower._run_fold) +# --------------------------------------------------------------------------- + +def filter_eye_samples(samples: list[dict]) -> list[dict]: + """Keep any single-eye sample with a valid image, matrix, and label.""" + return [ + s for s in samples + if s.get("image_1") is not None + and s.get("matrix_1") is not None + and s.get("label_1") is not None + ] + + +def filter_bilateral_samples(samples: list[dict]) -> list[dict]: + """Keep only patient-level samples where both eyes are fully present.""" + return [ + s for s in samples + if s.get("image_1") is not None + and s.get("matrix_1") is not None + and s.get("image_2") is not None + and s.get("matrix_2") is not None + and s.get("label_1") is not None + ] + + +def make_loader( + samples: list[dict], + slots: dict, + *, + image_transform, + image_preprocessor=None, + image_cache=None, + batch_size: int, + shuffle: bool, + num_workers: int, + sampler: Optional[WeightedRandomSampler] = None, +) -> DataLoader: + ds = SlotDataset( + samples, + slots, + image_transform=image_transform, + image_preprocessor=image_preprocessor, + image_cache=image_cache, + ) + return DataLoader( + ds, + batch_size=batch_size, + shuffle=(shuffle if sampler is None else False), + sampler=sampler, + num_workers=num_workers, + collate_fn=slot_collate, + ) + + +def build_balanced_sampler(samples: list[dict], label_key: str = "label_1") -> WeightedRandomSampler: + """Return a WeightedRandomSampler that equalises class frequency for training.""" + from collections import Counter + labels = [s[label_key] for s in samples] + counts = Counter(labels) + weights = [1.0 / counts[lbl] for lbl in labels] + return WeightedRandomSampler(weights, num_samples=len(weights), replacement=True) + + +def to_label_tensor(labels, device: torch.device) -> torch.Tensor: + if torch.is_tensor(labels): + return labels.to(device=device, dtype=torch.long) + return torch.as_tensor(labels, dtype=torch.long, device=device) diff --git a/v3/classes/metrics.py b/v3/classes/metrics.py new file mode 100644 index 0000000..f65e091 --- /dev/null +++ b/v3/classes/metrics.py @@ -0,0 +1,243 @@ +"""Metric computation, calibration, and threshold/bias tuning for V2.""" +from __future__ import annotations + +from typing import Optional + +import numpy as np +import torch +import torch.nn.functional as F +from sklearn.metrics import ( + cohen_kappa_score, + f1_score, + matthews_corrcoef, + recall_score, + roc_auc_score, + roc_curve, +) + + +# --------------------------------------------------------------------------- +# Loss +# --------------------------------------------------------------------------- + +def focal_loss( + logits: torch.Tensor, + targets: torch.Tensor, + gamma: float = 0.0, + weight: Optional[torch.Tensor] = None, + reduction: str = "mean", +) -> torch.Tensor: + """ + Standard focal loss wrapper. When gamma=0 it reduces to cross entropy. + weight should be per-class weights (same semantics as CrossEntropyLoss). + """ + if gamma <= 0: + return F.cross_entropy(logits, targets, weight=weight, reduction=reduction) + + log_probs = F.log_softmax(logits, dim=1) + probs = log_probs.exp() + + targets = targets.long().view(-1, 1) + logpt = log_probs.gather(1, targets) + pt = probs.gather(1, targets) + + focal_factor = (1.0 - pt).clamp_min(0.0) ** gamma + loss = -focal_factor * logpt + + if weight is not None: + class_weight = weight.gather(0, targets.view(-1)) + loss = loss * class_weight.view(-1, 1) + + loss = loss.view(-1) + if reduction == "sum": + return loss.sum() + if reduction == "mean": + return loss.mean() + return loss + + +# --------------------------------------------------------------------------- +# Basic array scoring +# --------------------------------------------------------------------------- + +def _score_arrays(y_true: np.ndarray, probs: np.ndarray, num_classes: int): + """Returns (acc, auc, n).""" + if y_true.size == 0: + return float("nan"), float("nan"), 0 + acc = float((probs.argmax(1) == y_true).mean()) + try: + auc = ( + float(roc_auc_score(y_true, probs[:, 1])) + if num_classes == 2 + else float(roc_auc_score(y_true, probs, multi_class="ovr", average="macro")) + ) + except Exception: + auc = float("nan") + return acc, auc, int(len(y_true)) + + +# --------------------------------------------------------------------------- +# Calibration +# --------------------------------------------------------------------------- + +def compute_ece(y_true: np.ndarray, probs: np.ndarray, n_bins: int = 10) -> float: + """Expected Calibration Error: weighted mean of |confidence - accuracy| per bin.""" + if y_true.size == 0: + return float("nan") + confidences = probs.max(axis=1) + predictions = probs.argmax(axis=1) + bin_edges = np.linspace(0.0, 1.0, n_bins + 1) + ece = 0.0 + n = len(y_true) + for i, (lo, hi) in enumerate(zip(bin_edges[:-1], bin_edges[1:])): + mask = (confidences >= lo) & ( + confidences <= hi if i == n_bins - 1 else confidences < hi + ) + if not mask.any(): + continue + bin_acc = float((predictions[mask] == y_true[mask]).mean()) + bin_conf = float(confidences[mask].mean()) + ece += float(mask.sum()) / n * abs(bin_conf - bin_acc) + return float(ece) + + +def compute_extended_metrics( + y_true: np.ndarray, + probs: np.ndarray, + num_classes: int, + n_bins: int = 10, + preds_override: Optional[np.ndarray] = None, +) -> dict: + nan = float("nan") + if y_true.size == 0: + return dict( + kappa=nan, mcc=nan, macro_f1=nan, + per_class_recall=np.full(num_classes, nan), ece=nan, + ) + preds = preds_override if preds_override is not None else probs.argmax(axis=1) + try: + kappa = float(cohen_kappa_score(y_true, preds)) + except Exception: + kappa = nan + try: + mcc = float(matthews_corrcoef(y_true, preds)) + except Exception: + mcc = nan + try: + macro_f1 = float(f1_score(y_true, preds, average="macro", zero_division=0)) + except Exception: + macro_f1 = nan + try: + pcr = recall_score( + y_true, preds, average=None, + labels=list(range(num_classes)), zero_division=0, + ).astype(float) + except Exception: + pcr = np.full(num_classes, nan) + ece = compute_ece(y_true, probs, n_bins=n_bins) + return dict(kappa=kappa, mcc=mcc, macro_f1=macro_f1, per_class_recall=pcr, ece=ece) + + +# --------------------------------------------------------------------------- +# Threshold / bias tuning +# --------------------------------------------------------------------------- + +def tune_binary_threshold(y_true: np.ndarray, p1: np.ndarray) -> float: + """Pick threshold via Youden's J (sensitivity + specificity − 1). + + This is class-distribution independent, unlike maximising raw accuracy, + which is biased toward the majority class on imbalanced validation sets. + Falls back to 0.5 if both classes are not present. + """ + if y_true.size == 0 or len(np.unique(y_true)) < 2: + return 0.5 + fpr, tpr, thresholds = roc_curve(y_true, p1) + j = tpr + (1.0 - fpr) - 1.0 + return float(thresholds[np.argmax(j)]) + + +def multiclass_acc_with_bias(y_true: np.ndarray, probs: np.ndarray, bias: np.ndarray) -> float: + """Balanced accuracy (mean per-class recall) after applying log-space bias.""" + if y_true.size == 0: + return float("nan") + logits = np.log(np.clip(probs, 1e-8, 1.0)) + bias.reshape(1, -1) + preds = np.argmax(logits, axis=1) + classes = np.unique(y_true) + per_class = [(preds[y_true == c] == c).mean() for c in classes] + return float(np.mean(per_class)) + + +def tune_multiclass_bias(y_true: np.ndarray, probs: np.ndarray, *, iters: int = 2) -> np.ndarray: + """Grid-search per-class log-space bias to maximise balanced accuracy. + + Balanced accuracy (mean per-class recall) is class-distribution independent, + unlike raw accuracy which is biased toward the majority class on imbalanced + validation sets. + """ + if y_true.size == 0 or probs.size == 0: + return np.zeros((0,), dtype=float) + c = probs.shape[1] + bias = np.zeros((c,), dtype=float) + grid = np.linspace(-1.0, 1.0, 41) + for _ in range(iters): + for k in range(c): + best_v = bias[k] + best_acc = multiclass_acc_with_bias(y_true, probs, bias) + old = bias[k] + for v in grid: + bias[k] = float(v) + acc = multiclass_acc_with_bias(y_true, probs, bias) + if acc > best_acc or (acc == best_acc and abs(v) < abs(best_v)): + best_acc, best_v = acc, float(v) + bias[k] = best_v + if np.isnan(best_acc): + bias[k] = old + return bias + + +def _svf(vec) -> Optional[str]: + """Serialise a float vector to pipe-separated string, or None if empty.""" + if vec is None: + return None + arr = np.asarray(vec, dtype=float) + if arr.size == 0: + return None + return "|".join(f"{float(v):.4f}" for v in arr.tolist()) + + +def _tune_and_snap( + y: np.ndarray, + p: np.ndarray, + acc: float, + num_classes: int, + args, + n_bins: int, +) -> tuple[dict, float, Optional[np.ndarray], Optional[np.ndarray]]: + """ + Apply threshold/bias tuning and compute extended metrics. + Returns (snap_dict, tuned_auc, threshold, bias). + """ + thr = 0.5 if num_classes == 2 else float("nan") + bias = None + ext_preds = None + + if args.tune_binary_threshold and num_classes == 2 and y.size > 0: + thr = tune_binary_threshold(y, p[:, 1]) + ext_preds = (p[:, 1] >= thr).astype(int) + acc = float((ext_preds == y).mean()) + elif args.tune_multiclass_bias and num_classes > 2 and y.size > 0: + bias = tune_multiclass_bias(y, p) + logits = np.log(np.clip(p, 1e-8, 1.0)) + bias.reshape(1, -1) + ext_preds = np.argmax(logits, axis=1) + acc = float((ext_preds == y).mean()) + + ext = compute_extended_metrics(y, p, num_classes, n_bins=n_bins, preds_override=ext_preds) + _, auc, n = _score_arrays(y, p, num_classes) + + snap = dict( + auc=auc, acc=acc, n=n, + kappa=ext["kappa"], mcc=ext["mcc"], macro_f1=ext["macro_f1"], + per_class_recall=ext["per_class_recall"], ece=ext["ece"], + threshold=thr, bias=bias, + ) + return snap, auc, thr, bias diff --git a/v3/classes/model_builder.py b/v3/classes/model_builder.py new file mode 100644 index 0000000..8bac831 --- /dev/null +++ b/v3/classes/model_builder.py @@ -0,0 +1,148 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable, Optional + +import torch +from torch import nn + +from v3.classes.bridges import Bridge, VoteBridge +from v3.classes.towers import ImageTower, ClinicalTower + +from .config_builder import ConfigAssembly +from .transforms import build_transform_chain + + +@dataclass +class V2ModelBundle: + image_tower: Optional[ImageTower] + metadata_tower: Optional[ClinicalTower] + bridge: Optional[nn.Module] + classifier: Optional[nn.Module] + image_transform: Optional[Callable] + matrix_transform: Optional[Callable] + + +def build_model_bundle( + assembly: ConfigAssembly, + clinical: Any, + *, + device: Optional[torch.device] = None, + strict: bool = True, +) -> V2ModelBundle: + """ + Build torch modules and input transforms from a V2 config assembly. + """ + image_tower_spec = _pick_tower(assembly, "image") + cd_tower_spec = _pick_tower(assembly, "clinical data") + bridge_spec = _pick_bridge(assembly) + image_loader = _pick_loader(assembly, input_type="image") + + clinical_core = getattr(clinical, "clinical", clinical) + num_classes = _infer_num_classes(clinical) + + img_tower = None + if image_tower_spec is not None: + img_tower = ImageTower( + backbone=image_tower_spec.params.get("backbone", "efficientnet_b0"), + freeze_ratio=float(image_tower_spec.params.get("freeze_ratio", 0.0) or 0.0), + use_se=bool(image_tower_spec.params.get("use_se", False)), + se_reduction=int(image_tower_spec.params.get("se_reduction", 16) or 16), + se_pre_norm=bool(image_tower_spec.params.get("se_pre_norm", True)), + augment=bool(image_tower_spec.params.get("augment", True)), + geometry_dim=int(image_tower_spec.params.get("geometry_dim", 0) or 0), + ) + if device is not None: + img_tower = img_tower.to(device) + + cd_tower = None + if cd_tower_spec is not None: + cd_tower = ClinicalTower( + clinical_core, + hidden_dim=int(cd_tower_spec.params.get("hidden_dim", 128) or 128), + dropout=float(cd_tower_spec.params.get("dropout", 0.1) or 0.1), + use_se=bool(cd_tower_spec.params.get("use_se", False)), + se_reduction=int(cd_tower_spec.params.get("se_reduction", 16) or 16), + se_pre_norm=bool(cd_tower_spec.params.get("se_pre_norm", True)), + ) + if device is not None: + cd_tower = cd_tower.to(device) + + bridge = None + if bridge_spec is not None and img_tower is not None and cd_tower is not None: + if bridge_spec.method == "consensus": + bridge = VoteBridge(num_classes=num_classes) + else: + bridge = Bridge( + img_dim=img_tower.out_dim, + meta_dim=cd_tower.out_dim, + num_classes=num_classes, + fusion_dim=int(bridge_spec.params.get("fusion_dim", 256) or 256), + mode="fused", + use_se=bool(bridge_spec.params.get("use_se", True)), + se_reduction=int(bridge_spec.params.get("se_reduction", 16) or 16), + se_pre_norm=bool(bridge_spec.params.get("se_pre_norm", True)), + ) + if device is not None: + bridge = bridge.to(device) + + classifier = None + if assembly.classifiers: + classifier = nn.Identity() + if device is not None: + classifier = classifier.to(device) + + image_transform = None + if image_loader is not None and image_tower_spec is not None: + image_transform = build_transform_chain( + image_loader.transforms, + backbone_name=image_tower_spec.params.get("backbone", "efficientnet_b0"), + augment=bool(image_tower_spec.params.get("augment", True)), + strict=strict, + ) + + return V2ModelBundle( + image_tower=img_tower, + metadata_tower=cd_tower, + bridge=bridge, + classifier=classifier, + image_transform=image_transform, + matrix_transform=None, + ) + + +def _pick_tower(assembly: ConfigAssembly, tower_type: str): + matches = [tower for tower in assembly.towers.values() if tower.tower_type == tower_type] + if not matches: + return None + if len(matches) > 1: + raise ValueError(f"Multiple {tower_type} towers found; only one is supported for now.") + return matches[0] + + +def _pick_bridge(assembly: ConfigAssembly): + if not assembly.bridges: + return None + if len(assembly.bridges) > 1: + raise ValueError("Multiple bridges found; only one is supported for now.") + return next(iter(assembly.bridges.values())) + + +def _pick_loader(assembly: ConfigAssembly, input_type: str): + matches = [loader for loader in assembly.loaders.values() if loader.input_type == input_type] + if not matches: + return None + if len(matches) > 1: + raise ValueError(f"Multiple loaders with input_type={input_type!r} found.") + return matches[0] + + +def _infer_num_classes(clinical: Any) -> int: + df = getattr(clinical, "df", None) + label_col = getattr(clinical, "label_col", None) + if df is None and hasattr(clinical, "clinical"): + df = clinical.clinical.df + label_col = clinical.clinical.label_col + if df is None or label_col is None or label_col not in df.columns: + return 2 + return int(df[label_col].dropna().nunique()) diff --git a/v3/classes/models.py b/v3/classes/models.py new file mode 100644 index 0000000..f581f45 --- /dev/null +++ b/v3/classes/models.py @@ -0,0 +1,858 @@ +"""V2 model classes and training/inference helpers.""" +from __future__ import annotations + +from random import random +from typing import Optional + +import numpy as np +import torch +import torch.nn.functional as F +from torch import nn +from torch.utils.data import DataLoader + +from v3.classes.bridges import Bridge +from v3.classes.towers import ImageTower, ClinicalTower + + +# --------------------------------------------------------------------------- +# Model classes +# --------------------------------------------------------------------------- + +class SingleEyeHT(nn.Module): + """ + ImageTower + ClinicalTower + Bridge, trained on eye-level samples. + Supports both Classic (eye-level) and Ensemble (patient-level averaging) eval. + """ + + def __init__( + self, + *, + backbone: str, + freeze_ratio: float, + augment: bool, + clinical_data, + num_classes: int, + cd_hidden_dim: int = 128, + fusion_dim: int = 256, + bridge_mode: str = "fused", + ): + super().__init__() + self.img_tower = ImageTower( + backbone=backbone, + freeze_ratio=freeze_ratio, + augment=augment, + use_se=False, + ) + self.cd_tower = ClinicalTower( + clinical_data=clinical_data, + hidden_dim=cd_hidden_dim, + use_se=False, + ) + self.bridge = Bridge( + img_dim=self.img_tower.out_dim, + meta_dim=self.cd_tower.out_dim, + num_classes=num_classes, + fusion_dim=fusion_dim, + mode=bridge_mode, + use_se=False, + ) + + @property + def transform(self): + return self.img_tower.transform + + def forward(self, x: torch.Tensor, meta: torch.Tensor) -> torch.Tensor: + img_feats = None if self.bridge.mode == "clinical_only" else self.img_tower(x) + md_feats = None if self.bridge.mode == "image_only" else self.cd_tower(meta) + out_f, _, _ = self.bridge(img_feats, md_feats) + return out_f + + +class BilateralHT(nn.Module): + """ + Bilateral mode with joint towers: + - shared eye-level towers encode OD/OS independently + - joint image and clinical data towers combine OD/OS embeddings + - standard Bridge fuses joint image + joint image + joint clinical data embeddings + """ + + def __init__( + self, + *, + backbone: str, + freeze_ratio: float, + augment: bool, + clinical_data, + num_classes: int, + cd_hidden_dim: int = 128, + fusion_dim: int = 256, + ): + super().__init__() + self.eye_img_tower = ImageTower( + backbone=backbone, + freeze_ratio=freeze_ratio, + augment=augment, + use_se=False, + ) + self.eye_cd_tower = ClinicalTower( + clinical_data=clinical_data, + hidden_dim=cd_hidden_dim, + use_se=False, + ) + img_dim = self.eye_img_tower.out_dim + md_dim = self.eye_cd_tower.out_dim + self.joint_img = nn.Sequential( + nn.Linear(2 * img_dim, fusion_dim), + nn.LayerNorm(fusion_dim), + nn.ReLU(), + nn.Dropout(0.3), + nn.Linear(fusion_dim, img_dim), + ) + self.joint_md = nn.Sequential( + nn.Linear(2 * md_dim, fusion_dim), + nn.LayerNorm(fusion_dim), + nn.ReLU(), + nn.Dropout(0.3), + nn.Linear(fusion_dim, md_dim), + ) + self.bridge = Bridge( + img_dim=img_dim, + meta_dim=md_dim, + num_classes=num_classes, + fusion_dim=fusion_dim, + mode="fused", + use_se=False, + ) + # Auxiliary heads for tower warmup / BCD tower steps. + self.aux_img = nn.Linear(img_dim, num_classes) + self.aux_md = nn.Linear(md_dim, num_classes) + + @property + def transform(self): + return self.eye_img_tower.transform + + def encode_joint( + self, + x_od: torch.Tensor, + meta_od: torch.Tensor, + x_os: torch.Tensor, + meta_os: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + img_od = self.eye_img_tower(x_od) + md_od = self.eye_cd_tower(meta_od) + img_os = self.eye_img_tower(x_os) + md_os = self.eye_cd_tower(meta_os) + joint_img = self.joint_img(torch.cat([img_od, img_os], dim=1)) + joint_md = self.joint_md(torch.cat([md_od, md_os], dim=1)) + return joint_img, joint_md + + def forward( + self, + x_od: torch.Tensor, + meta_od: torch.Tensor, + x_os: torch.Tensor, + meta_os: torch.Tensor, + ) -> torch.Tensor: + joint_img, joint_md = self.encode_joint(x_od, meta_od, x_os, meta_os) + out_f, _, _ = self.bridge(joint_img, joint_md) + return out_f + + +class FusedEnsembleHT(nn.Module): + """ + SingleEyeHT base with a per-eye attention scorer for bilateral fusion. + + The base model is trained eye-level (identical to ensemble mode). + After base training completes, the base is frozen and only the + eye_scorer is trained on bilateral (patient-level) samples. + + At inference, eye_scorer is applied independently to each eye's logit + vector to produce a scalar attention score. Softmax over the two scores + gives attention weights; the final logit is a weighted sum: + + score_od = eye_scorer(logit_od) # [B, 1] + score_os = eye_scorer(logit_os) # [B, 1] + alpha = softmax([score_od, score_os]) # [B, 2], sums to 1 + out = alpha[:,0:1]*logit_od + alpha[:,1:2]*logit_os + + Because eye_scorer is applied to each eye with the same weights, the + mechanism is permutation-equivariant — there is no left/right positional + bias. Through training on bilateral labels the scorer learns to give high + scores to logits that point strongly toward the GC class, creating the + desired asymmetry: a confidently GC eye dominates the patient prediction + more than a comparably confident healthy eye would. + """ + + def __init__(self, base: SingleEyeHT, num_classes: int): + super().__init__() + self.base = base + # Applied independently to each eye's logit → scalar attention score. + # Learns the GC-direction in logit space from bilateral labels. + self.eye_scorer = nn.Linear(num_classes, 1, bias=True) + + def forward( + self, + x_od: torch.Tensor, + meta_od: torch.Tensor, + x_os: torch.Tensor, + meta_os: torch.Tensor, + ) -> torch.Tensor: + logit_od = self.base(x_od, meta_od) # [B, C] + logit_os = self.base(x_os, meta_os) # [B, C] + scores = torch.cat([self.eye_scorer(logit_od), + self.eye_scorer(logit_os)], dim=1) # [B, 2] + alpha = torch.softmax(scores, dim=1) # [B, 2] + return alpha[:, 0:1] * logit_od + alpha[:, 1:2] * logit_os # [B, C] + + +# --------------------------------------------------------------------------- +# Phase control +# --------------------------------------------------------------------------- + +def _set_requires_grad(module: nn.Module, enabled: bool) -> None: + for p in module.parameters(): + p.requires_grad = enabled + + +def _set_single_phase(model: SingleEyeHT, phase: str) -> None: + bridge_mode = model.bridge.mode + # Ablation modes have no fusion bridge; fused_warmup is meaningless — treat as tower_warmup + if bridge_mode in ("image_only", "clinical_only") and phase == "fused_warmup": + phase = "tower_warmup" + if phase == "cd_warmup": + _set_requires_grad(model.img_tower, False) + _set_requires_grad(model.cd_tower, True) + _set_requires_grad(model.bridge.classifier_img, False) + _set_requires_grad(model.bridge.classifier_cd, True) + _set_requires_grad(model.bridge.W_img, False) + _set_requires_grad(model.bridge.W_md, False) + _set_requires_grad(model.bridge.classifier_fused, False) + return + if phase == "tower_warmup": + _set_requires_grad(model.img_tower, bridge_mode != "clinical_only") + _set_requires_grad(model.cd_tower, bridge_mode != "image_only") + _set_requires_grad(model.bridge.classifier_img, bridge_mode != "clinical_only") + _set_requires_grad(model.bridge.classifier_cd, bridge_mode != "image_only") + _set_requires_grad(model.bridge.W_img, False) + _set_requires_grad(model.bridge.W_md, False) + _set_requires_grad(model.bridge.classifier_fused, False) + return + if phase == "fused_warmup": + _set_requires_grad(model.img_tower, False) + _set_requires_grad(model.cd_tower, False) + _set_requires_grad(model.bridge.classifier_img, False) + _set_requires_grad(model.bridge.classifier_cd, False) + _set_requires_grad(model.bridge.W_img, True) + _set_requires_grad(model.bridge.W_md, True) + _set_requires_grad(model.bridge.classifier_fused, True) + return + _set_requires_grad(model, True) + + +def _set_bilateral_phase(model: BilateralHT, phase: str) -> None: + if phase == "tower_warmup": + _set_requires_grad(model.eye_img_tower, True) + _set_requires_grad(model.eye_cd_tower, True) + _set_requires_grad(model.joint_img, True) + _set_requires_grad(model.joint_md, True) + _set_requires_grad(model.aux_img, True) + _set_requires_grad(model.aux_md, True) + _set_requires_grad(model.bridge, False) + return + if phase == "fused_warmup": + _set_requires_grad(model.eye_img_tower, False) + _set_requires_grad(model.eye_cd_tower, False) + _set_requires_grad(model.joint_img, False) + _set_requires_grad(model.joint_md, False) + _set_requires_grad(model.aux_img, False) + _set_requires_grad(model.aux_md, False) + _set_requires_grad(model.bridge, True) + return + _set_requires_grad(model, True) + + +# --------------------------------------------------------------------------- +# Training helpers +# --------------------------------------------------------------------------- + +def train_single_epoch( + model: SingleEyeHT, + loader: DataLoader, + opt, + device: torch.device, + *, + phase: str, + bcd_prob: float = 0.5, + tower_loss_mode: str = "bcd", +) -> tuple[float, float]: + model.train() + _set_single_phase(model, phase) + total_loss = total_correct = total_n = 0 + for batch in loader: + x = batch.get("image_1") + m = batch.get("matrix_1") + y = batch.get("label_1") + if phase == "cd_warmup": + if not torch.is_tensor(m): + continue + m = m.to(device) + y = _to_label_tensor(y, device) + md_feats = model.cd_tower(m) + logits = model.bridge.classifier_cd(md_feats) + loss = F.cross_entropy(logits, y) + opt.zero_grad(); loss.backward(); opt.step() + bs = y.shape[0] + total_loss += float(loss.item()) * bs + total_correct += int((logits.argmax(1) == y).sum()) + total_n += bs + continue + if not torch.is_tensor(x) or not torch.is_tensor(m): + continue + x = x.to(device) + m = m.to(device) + y = _to_label_tensor(y, device) + bridge_mode = model.bridge.mode + + img_feats = None if bridge_mode == "clinical_only" else model.img_tower(x) + md_feats = None if bridge_mode == "image_only" else model.cd_tower(m) + + if phase == "tower_warmup": + if bridge_mode == "clinical_only": + logits = model.bridge.classifier_cd(md_feats) + loss = F.cross_entropy(logits, y) + elif bridge_mode == "image_only": + logits = model.bridge.classifier_img(img_feats) + loss = F.cross_entropy(logits, y) + else: + logits_i = model.bridge.classifier_img(img_feats) + logits_m = model.bridge.classifier_cd(md_feats) + loss = 0.5 * (F.cross_entropy(logits_i, y) + F.cross_entropy(logits_m, y)) + logits = 0.5 * (F.softmax(logits_i, dim=1) + F.softmax(logits_m, dim=1)) + elif phase == "fused_warmup": + logits, _, _ = model.bridge(img_feats, md_feats) + loss = F.cross_entropy(logits, y) + else: + if bridge_mode == "clinical_only": + logits = model.bridge.classifier_cd(md_feats) + loss = F.cross_entropy(logits, y) + elif bridge_mode == "image_only": + logits = model.bridge.classifier_img(img_feats) + loss = F.cross_entropy(logits, y) + elif tower_loss_mode == "all": + loss_i = F.cross_entropy(model.bridge.classifier_img(img_feats), y) + loss_m = F.cross_entropy(model.bridge.classifier_cd(md_feats), y) + logits, _, _ = model.bridge(img_feats, md_feats) + loss = F.cross_entropy(logits, y) + loss_i + loss_m + elif random() < bcd_prob: + if random() < 0.5: + logits = model.bridge.classifier_img(img_feats) + else: + logits = model.bridge.classifier_cd(md_feats) + loss = F.cross_entropy(logits, y) + else: + logits, _, _ = model.bridge(img_feats, md_feats) + loss = F.cross_entropy(logits, y) + + opt.zero_grad() + loss.backward() + opt.step() + bs = y.shape[0] + total_loss += float(loss.item()) * bs + total_correct += int((logits.argmax(1) == y).sum()) + total_n += bs + return ( + total_loss / total_n if total_n else float("nan"), + total_correct / total_n if total_n else float("nan"), + ) + + +def train_bilateral_epoch( + model: BilateralHT, + loader: DataLoader, + opt, + device: torch.device, + *, + phase: str, + bcd_prob: float = 0.5, + tower_loss_mode: str = "bcd", +) -> tuple[float, float]: + model.train() + _set_bilateral_phase(model, phase) + total_loss = total_correct = total_n = 0 + for batch in loader: + x1 = batch.get("image_1") + m1 = batch.get("matrix_1") + x2 = batch.get("image_2") + m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + x1 = x1.to(device); m1 = m1.to(device) + x2 = x2.to(device); m2 = m2.to(device) + y = _to_label_tensor(y, device) + joint_img, joint_md = model.encode_joint(x1, m1, x2, m2) + + if phase == "tower_warmup": + logits_i = model.aux_img(joint_img) + logits_m = model.aux_md(joint_md) + loss = 0.5 * (F.cross_entropy(logits_i, y) + F.cross_entropy(logits_m, y)) + logits = 0.5 * (F.softmax(logits_i, dim=1) + F.softmax(logits_m, dim=1)) + elif phase == "fused_warmup": + logits, _, _ = model.bridge(joint_img, joint_md) + loss = F.cross_entropy(logits, y) + else: + if tower_loss_mode == "all": + loss_i = F.cross_entropy(model.aux_img(joint_img), y) + loss_m = F.cross_entropy(model.aux_md(joint_md), y) + logits, _, _ = model.bridge(joint_img, joint_md) + loss = F.cross_entropy(logits, y) + loss_i + loss_m + elif random() < bcd_prob: + if random() < 0.5: + logits = model.aux_img(joint_img) + else: + logits = model.aux_md(joint_md) + loss = F.cross_entropy(logits, y) + else: + logits, _, _ = model.bridge(joint_img, joint_md) + loss = F.cross_entropy(logits, y) + + opt.zero_grad() + loss.backward() + opt.step() + bs = y.shape[0] + total_loss += float(loss.item()) * bs + total_correct += int((logits.argmax(1) == y).sum()) + total_n += bs + return ( + total_loss / total_n if total_n else float("nan"), + total_correct / total_n if total_n else float("nan"), + ) + + +def train_fusion_epoch( + model: FusedEnsembleHT, + loader: DataLoader, + opt, + device: torch.device, +) -> tuple[float, float]: + """Train only the fusion head; the base SingleEyeHT is frozen in eval mode.""" + model.base.eval() + model.eye_scorer.train() + total_loss = total_correct = total_n = 0 + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and + torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + out = model(x1.to(device), m1.to(device), x2.to(device), m2.to(device)) + loss = F.cross_entropy(out, y_t) + opt.zero_grad() + loss.backward() + opt.step() + bs = y_t.shape[0] + total_loss += float(loss.item()) * bs + total_correct += int((out.argmax(1) == y_t).sum()) + total_n += bs + return ( + total_loss / total_n if total_n else float("nan"), + total_correct / total_n if total_n else float("nan"), + ) + + +# --------------------------------------------------------------------------- +# Inference helpers +# --------------------------------------------------------------------------- + +def _to_label_tensor(labels, device: torch.device) -> torch.Tensor: + if torch.is_tensor(labels): + return labels.to(device=device, dtype=torch.long) + return torch.as_tensor(labels, dtype=torch.long, device=device) + + +def collect_probs_classic( + model: SingleEyeHT, + loader: DataLoader, + device: torch.device, +) -> tuple[np.ndarray, np.ndarray]: + """ + Classic eye-level eval using the bilateral val loader. + OD and OS are treated as independent samples (both contribute to the + arrays with the same patient label). Returns (y_true [2N], probs [2N, C]). + """ + model.eval() + y_chunks, p_chunks = [], [] + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + p_od = F.softmax(model(x1.to(device), m1.to(device)), dim=1) + p_os = F.softmax(model(x2.to(device), m2.to(device)), dim=1) + y_np = y_t.cpu().numpy() + y_chunks += [y_np, y_np] + p_chunks += [p_od.cpu().numpy(), p_os.cpu().numpy()] + if not y_chunks: + return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) + return np.concatenate(y_chunks), np.concatenate(p_chunks, axis=0) + + +def collect_probs_ensemble_pereye( + model: "SingleEyeHT", + loader: DataLoader, + device: torch.device, + *, + return_ids: bool = False, +): + """ + Per-patient, per-eye probs for all 3 heads from a bilateral loader (ensemble mode). + + OD corresponds to image_1/matrix_1; OS to image_2/matrix_2. + Arrays are in patient order (not interleaved at sample level). + + Returns: + (y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os) + or, when return_ids=True: + (y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os, patient_ids) + + Patient-level averaged ensemble probs can be recovered as: + p_en = 0.5 * (pf_od + pf_os) + """ + model.eval() + y_chunks: list = [] + pf_od_c, pi_od_c, pm_od_c = [], [], [] + pf_os_c, pi_os_c, pm_os_c = [], [], [] + id_chunks: list[str] = [] + + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and + torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + + def _fwd(x, m): + img_feats = None if model.bridge.mode == "clinical_only" else model.img_tower(x.to(device)) + md_feats = None if model.bridge.mode == "image_only" else model.cd_tower(m.to(device)) + out_f, out_i, out_m = model.bridge(img_feats, md_feats) + pf = F.softmax(out_f, dim=1) + pi = F.softmax(out_i, dim=1) if out_i is not None else pf + pm = F.softmax(out_m, dim=1) if out_m is not None else pf + return pf, pi, pm + + pf_od, pi_od, pm_od = _fwd(x1, m1) + pf_os, pi_os, pm_os = _fwd(x2, m2) + + y_chunks.append(y_t.cpu().numpy()) + pf_od_c.append(pf_od.cpu().numpy()); pi_od_c.append(pi_od.cpu().numpy()); pm_od_c.append(pm_od.cpu().numpy()) + pf_os_c.append(pf_os.cpu().numpy()); pi_os_c.append(pi_os.cpu().numpy()); pm_os_c.append(pm_os.cpu().numpy()) + + if return_ids: + ids = batch.get("id_1", [""] * len(y_t)) + if torch.is_tensor(ids): + ids = ids.tolist() + id_chunks.extend([str(i) for i in ids]) + + if not y_chunks: + z = np.zeros((0, 0), dtype=np.float32) + empty_i = np.array([], dtype=np.int64) + base = (empty_i, z, z, z, z, z, z) + return base + (np.array([], dtype=object),) if return_ids else base + + y = np.concatenate(y_chunks) + pf_od = np.concatenate(pf_od_c, axis=0); pi_od = np.concatenate(pi_od_c, axis=0); pm_od = np.concatenate(pm_od_c, axis=0) + pf_os = np.concatenate(pf_os_c, axis=0); pi_os = np.concatenate(pi_os_c, axis=0); pm_os = np.concatenate(pm_os_c, axis=0) + if return_ids: + return y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os, np.array(id_chunks, dtype=object) + return y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os + + +def collect_probs_ensemble( + model: SingleEyeHT, + loader: DataLoader, + device: torch.device, +) -> tuple[np.ndarray, np.ndarray]: + """ + Patient-level ensemble eval: average OD and OS softmax probabilities. + Returns (y_true [N], probs [N, C]). + """ + model.eval() + y_chunks, p_chunks = [], [] + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + p_od = F.softmax(model(x1.to(device), m1.to(device)), dim=1) + p_os = F.softmax(model(x2.to(device), m2.to(device)), dim=1) + p = 0.5 * (p_od + p_os) + y_chunks.append(y_t.cpu().numpy()) + p_chunks.append(p.cpu().numpy()) + if not y_chunks: + return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) + return np.concatenate(y_chunks), np.concatenate(p_chunks, axis=0) + + +def collect_probs_bilateral( + model: BilateralHT, + loader: DataLoader, + device: torch.device, +) -> tuple[np.ndarray, np.ndarray]: + """Patient-level bilateral eval. Returns (y_true [N], probs [N, C]).""" + model.eval() + y_chunks, p_chunks = [], [] + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + p = F.softmax(model(x1.to(device), m1.to(device), x2.to(device), m2.to(device)), dim=1) + y_chunks.append(y_t.cpu().numpy()) + p_chunks.append(p.cpu().numpy()) + if not y_chunks: + return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) + return np.concatenate(y_chunks), np.concatenate(p_chunks, axis=0) + + +def collect_probs_fused( + model: FusedEnsembleHT, + loader: DataLoader, + device: torch.device, +) -> tuple[np.ndarray, np.ndarray]: + """Patient-level fused-head eval. Returns (y_true [N], probs [N, C]).""" + model.eval() + y_chunks, p_chunks = [], [] + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and + torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + p = F.softmax(model(x1.to(device), m1.to(device), + x2.to(device), m2.to(device)), dim=1) + y_chunks.append(y_t.cpu().numpy()) + p_chunks.append(p.cpu().numpy()) + if not y_chunks: + return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) + return np.concatenate(y_chunks), np.concatenate(p_chunks, axis=0) + + +def collect_probs_single_components( + model: SingleEyeHT, + loader: DataLoader, + device: torch.device, + *, + aggregate_patient: bool, + return_logits: bool = False, +): + """ + Collect fused/img/md probabilities (and optionally raw logits) for SingleEyeHT. + - aggregate_patient=False: eye-level (OD/OS as independent samples) + - aggregate_patient=True : patient-level (average OD/OS per head) + - return_logits=False: returns (y, probs_f, probs_i, probs_m) + - return_logits=True: returns (y, probs_f, probs_i, probs_m, + logits_f, logits_i, logits_m) + Note: logits are averaged across eyes when aggregate_patient=True, + which is equivalent to averaging in logit space (before softmax). + """ + model.eval() + y_chunks = [] + pf_chunks, pi_chunks, pm_chunks = [], [], [] + lf_chunks, li_chunks, lm_chunks = [], [], [] + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + + def _per_eye(x, m): + img_feats = None if model.bridge.mode == "clinical_only" else model.img_tower(x.to(device)) + md_feats = None if model.bridge.mode == "image_only" else model.cd_tower(m.to(device)) + out_f, out_i, out_m = model.bridge(img_feats, md_feats) + pf = F.softmax(out_f, dim=1) + pi = F.softmax(out_i, dim=1) if out_i is not None else pf + pm = F.softmax(out_m, dim=1) if out_m is not None else pf + lf = out_f + li = out_i if out_i is not None else out_f + lm = out_m if out_m is not None else out_f + return pf, pi, pm, lf, li, lm + + pf_od, pi_od, pm_od, lf_od, li_od, lm_od = _per_eye(x1, m1) + pf_os, pi_os, pm_os, lf_os, li_os, lm_os = _per_eye(x2, m2) + + if aggregate_patient: + y_chunks.append(y_t.cpu().numpy()) + pf_chunks.append((0.5 * (pf_od + pf_os)).cpu().numpy()) + pi_chunks.append((0.5 * (pi_od + pi_os)).cpu().numpy()) + pm_chunks.append((0.5 * (pm_od + pm_os)).cpu().numpy()) + lf_chunks.append((0.5 * (lf_od + lf_os)).cpu().numpy()) + li_chunks.append((0.5 * (li_od + li_os)).cpu().numpy()) + lm_chunks.append((0.5 * (lm_od + lm_os)).cpu().numpy()) + else: + y_np = y_t.cpu().numpy() + y_chunks += [y_np, y_np] + pf_chunks += [pf_od.cpu().numpy(), pf_os.cpu().numpy()] + pi_chunks += [pi_od.cpu().numpy(), pi_os.cpu().numpy()] + pm_chunks += [pm_od.cpu().numpy(), pm_os.cpu().numpy()] + lf_chunks += [lf_od.cpu().numpy(), lf_os.cpu().numpy()] + li_chunks += [li_od.cpu().numpy(), li_os.cpu().numpy()] + lm_chunks += [lm_od.cpu().numpy(), lm_os.cpu().numpy()] + + if not y_chunks: + z = np.zeros((0, 0), dtype=np.float32) + if return_logits: + return np.array([], dtype=np.int64), z, z, z, z, z, z + return np.array([], dtype=np.int64), z, z, z + + y = np.concatenate(y_chunks) + pf = np.concatenate(pf_chunks, axis=0) + pi = np.concatenate(pi_chunks, axis=0) + pm = np.concatenate(pm_chunks, axis=0) + if return_logits: + lf = np.concatenate(lf_chunks, axis=0) + li = np.concatenate(li_chunks, axis=0) + lm = np.concatenate(lm_chunks, axis=0) + return y, pf, pi, pm, lf, li, lm + return y, pf, pi, pm + + +def collect_probs_eye_level( + model: "SingleEyeHT", + loader: DataLoader, + device: torch.device, + *, + return_ids: bool = False, +): + """ + Collect fused/img/md probabilities from a single-eye loader (image_1/matrix_1 only). + Used for eval-mode passes over the training set. + + Returns (y, probs_f, probs_i, probs_m) or, when return_ids=True, + (y, probs_f, probs_i, probs_m, sample_ids) where sample_ids is an + array of strings like "2OD", "4OS". + """ + model.eval() + y_chunks, pf_chunks, pi_chunks, pm_chunks, id_chunks = [], [], [], [], [] + with torch.no_grad(): + for batch in loader: + x = batch.get("image_1") + m = batch.get("matrix_1") + y = batch.get("label_1") + if not (torch.is_tensor(x) and torch.is_tensor(m)): + continue + y_t = _to_label_tensor(y, device) + img_feats = None if model.bridge.mode == "clinical_only" else model.img_tower(x.to(device)) + md_feats = None if model.bridge.mode == "image_only" else model.cd_tower(m.to(device)) + out_f, out_i, out_m = model.bridge(img_feats, md_feats) + pf = F.softmax(out_f, dim=1) + pi = F.softmax(out_i, dim=1) if out_i is not None else pf + pm = F.softmax(out_m, dim=1) if out_m is not None else pf + y_chunks.append(y_t.cpu().numpy()) + pf_chunks.append(pf.cpu().numpy()) + pi_chunks.append(pi.cpu().numpy()) + pm_chunks.append(pm.cpu().numpy()) + if return_ids: + ids = batch.get("id_1", [""] * len(y_t)) + eyes = batch.get("eye_id_1", [""] * len(y_t)) + # ids/eyes may be tensors (int) or lists of strings + if torch.is_tensor(ids): + ids = ids.tolist() + if torch.is_tensor(eyes): + eyes = eyes.tolist() + id_chunks.extend( + [f"{pid}{eye}" for pid, eye in zip(ids, eyes)] + ) + + if not y_chunks: + z = np.zeros((0, 0), dtype=np.float32) + empty_ids = np.array([], dtype=object) + if return_ids: + return np.array([], dtype=np.int64), z, z, z, empty_ids + return np.array([], dtype=np.int64), z, z, z + + y = np.concatenate(y_chunks) + pf = np.concatenate(pf_chunks, axis=0) + pi = np.concatenate(pi_chunks, axis=0) + pm = np.concatenate(pm_chunks, axis=0) + if return_ids: + return y, pf, pi, pm, np.array(id_chunks, dtype=object) + return y, pf, pi, pm + + +def collect_probs_bilateral_components( + model: BilateralHT, + loader: DataLoader, + device: torch.device, +) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + """Collect fused/img/md probabilities for bilateral joint-tower model.""" + model.eval() + y_chunks = [] + pf_chunks, pi_chunks, pm_chunks = [], [], [] + with torch.no_grad(): + for batch in loader: + x1 = batch.get("image_1"); m1 = batch.get("matrix_1") + x2 = batch.get("image_2"); m2 = batch.get("matrix_2") + y = batch.get("label_1") + if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): + continue + y_t = _to_label_tensor(y, device) + joint_img, joint_md = model.encode_joint( + x1.to(device), m1.to(device), x2.to(device), m2.to(device) + ) + out_f, _, _ = model.bridge(joint_img, joint_md) + out_i = model.aux_img(joint_img) + out_m = model.aux_md(joint_md) + y_chunks.append(y_t.cpu().numpy()) + pf_chunks.append(F.softmax(out_f, dim=1).cpu().numpy()) + pi_chunks.append(F.softmax(out_i, dim=1).cpu().numpy()) + pm_chunks.append(F.softmax(out_m, dim=1).cpu().numpy()) + if not y_chunks: + z = np.zeros((0, 0), dtype=np.float32) + return np.array([], dtype=np.int64), z, z, z + return ( + np.concatenate(y_chunks), + np.concatenate(pf_chunks, axis=0), + np.concatenate(pi_chunks, axis=0), + np.concatenate(pm_chunks, axis=0), + ) + + +# --------------------------------------------------------------------------- +# V2ModeComparisonOps — thin class wrapper kept for external import compat +# --------------------------------------------------------------------------- + +class V2ModeComparisonOps: + """Namespace wrapper kept for backward-compatibility imports.""" + + _set_requires_grad = staticmethod(_set_requires_grad) + _set_single_phase = staticmethod(_set_single_phase) + _set_bilateral_phase = staticmethod(_set_bilateral_phase) + train_single_epoch = staticmethod(train_single_epoch) + train_bilateral_epoch = staticmethod(train_bilateral_epoch) + collect_probs_classic = staticmethod(collect_probs_classic) + collect_probs_ensemble = staticmethod(collect_probs_ensemble) + collect_probs_bilateral = staticmethod(collect_probs_bilateral) + + @staticmethod + def _to_label_tensor(labels, device: torch.device) -> torch.Tensor: + return _to_label_tensor(labels, device) diff --git a/v3/classes/network_manager.py b/v3/classes/network_manager.py new file mode 100644 index 0000000..05e860e --- /dev/null +++ b/v3/classes/network_manager.py @@ -0,0 +1,173 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Optional, Protocol + +import pandas as pd + + +@dataclass +class PatientSplit: + """Patient-disjoint split definition for a fold (V3: train/val/test).""" + + train: pd.DataFrame + val: pd.DataFrame + test: Optional[pd.DataFrame] = None + + +@dataclass +class LoaderBundle: + """All loaders needed by a training run.""" + + train: Any + val: Any + test: Optional[Any] = None + + +@dataclass +class FoldResult: + """Normalized fold output from trainer implementations.""" + + fold: int + metrics: dict[str, Any] + artifacts: dict[str, Any] + + +class SplitManager(Protocol): + def build_plans( + self, + *, + clinical: Any, + args: Any, + profile: Optional[Any] = None, + ) -> list[PatientSplit]: + ... + + +class GraphFactory(Protocol): + def build( + self, + *, + clinical: Any, + args: Any, + fold: int, + profile: Optional[Any] = None, + ) -> Any: + ... + + +class LoaderFactory(Protocol): + def build( + self, + *, + clinical: Any, + split: PatientSplit, + args: Any, + fold: int, + profile: Optional[Any] = None, + ) -> LoaderBundle: + ... + + +class Trainer(Protocol): + def fit( + self, + *, + graph: Any, + loaders: LoaderBundle, + args: Any, + fold: int, + profile: Optional[Any] = None, + ) -> FoldResult: + ... + + +class NetworkManager: + """V3 orchestration entrypoint. No holdout — test = current fold.""" + + def __init__( + self, + *, + clinical: Any, + args: Any, + split_manager: SplitManager, + graph_factory: GraphFactory, + loader_factory: LoaderFactory, + trainer: Trainer, + profile: Optional[Any] = None, + ) -> None: + self.clinical = clinical + self.args = args + self.split_manager = split_manager + self.graph_factory = graph_factory + self.loader_factory = loader_factory + self.trainer = trainer + self.profile = profile + self._split_plans: Optional[list[PatientSplit]] = None + + def run_fold(self, fold: int) -> FoldResult: + plans = self._get_split_plans() + if fold < 0 or fold >= len(plans): + raise IndexError(f"Requested fold {fold} but only {len(plans)} fold plans are available") + split = plans[fold] + self._validate_patient_disjointness(split) + self._validate_labels(split) + + graph = self.graph_factory.build( + clinical=self.clinical, args=self.args, fold=fold, profile=self.profile, + ) + loaders = self.loader_factory.build( + clinical=self.clinical, split=split, args=self.args, fold=fold, profile=self.profile, + ) + return self.trainer.fit( + graph=graph, loaders=loaders, args=self.args, fold=fold, profile=self.profile, + ) + + def run_all_folds(self, n_splits: Optional[int] = None) -> list[FoldResult]: + plans = self._get_split_plans() + max_folds = len(plans) + n = max_folds if n_splits is None else int(n_splits) + if n < 1: + raise ValueError("n_splits must be >= 1") + if n > max_folds: + raise ValueError(f"Requested {n} folds but only {max_folds} available") + return [self.run_fold(fold) for fold in range(n)] + + def _get_split_plans(self) -> list[PatientSplit]: + if self._split_plans is None: + self._split_plans = self.split_manager.build_plans( + clinical=self.clinical, args=self.args, profile=self.profile, + ) + if not self._split_plans: + raise ValueError("SplitManager returned no fold plans") + return self._split_plans + + def _validate_patient_disjointness(self, split: PatientSplit) -> None: + train_ids = self._patient_ids(split.train) + val_ids = self._patient_ids(split.val) + test_ids = self._patient_ids(split.test) if split.test is not None else set() + + if train_ids & val_ids: + raise ValueError(f"Patient leakage train/val: {sorted(train_ids & val_ids)[:10]}") + if train_ids & test_ids: + raise ValueError(f"Patient leakage train/test: {sorted(train_ids & test_ids)[:10]}") + if val_ids & test_ids: + raise ValueError(f"Patient leakage val/test: {sorted(val_ids & test_ids)[:10]}") + + def _validate_labels(self, split: PatientSplit) -> None: + label_col = getattr(self.clinical, "label_col", None) + if not label_col: + return + for name, df in (("train", split.train), ("val", split.val), ("test", split.test)): + if df is None: + continue + if label_col not in df.columns: + raise ValueError(f"{name} split is missing label column {label_col!r}") + + @staticmethod + def _patient_ids(df: Optional[pd.DataFrame]) -> set[Any]: + if df is None or df.empty: + return set() + if "Patient ID" not in df.columns: + raise ValueError("Split dataframes must include 'Patient ID'") + return set(df["Patient ID"].tolist()) diff --git a/v3/classes/papila_builders.py b/v3/classes/papila_builders.py new file mode 100644 index 0000000..19b649d --- /dev/null +++ b/v3/classes/papila_builders.py @@ -0,0 +1,240 @@ +from __future__ import annotations + +from typing import Callable, Dict, List, Optional + +import numpy as np +import pandas as pd + +from v3.classes.data_bundle import DataBundle + +# ---- Pachymetry → IOP correction (per PAPILA Table 3) ---- +_PACHY_TABLE: Dict[int, int] = { + 475: +5, + 485: +4, + 495: +4, + 505: +3, + 515: +2, + 525: +1, + 535: +1, + 545: 0, + 555: -1, + 565: -1, + 575: -2, + 585: -3, + 595: -4, + 605: -4, + 615: -5, +} +_PACHY_KEYS = np.array(sorted(_PACHY_TABLE.keys())) + + +def _nearest_pachy_key(x: float) -> int: + idx = int(np.argmin(np.abs(_PACHY_KEYS - float(x)))) + return int(_PACHY_KEYS[idx]) + + +def _fit_perkins_converter( + frames: List[pd.DataFrame], method: str +) -> Callable[[float, Optional[float]], float]: + """ + Fit a Perkins→Pneumatic converter from pooled paired observations across all frames. + Returns a callable: converter(perkins_value, pachymetry_value) -> float. + Supported methods: "ratio", "ols", "lad", "multi". + """ + combined = pd.concat(frames, ignore_index=True) + paired = combined.dropna(subset=["Pneumatic", "Perkins"]) + pneumatic = paired["Pneumatic"].values.astype(float) + perkins = paired["Perkins"].values.astype(float) + + if len(paired) == 0: + raise ValueError("No paired Pneumatic+Perkins observations found; cannot fit converter.") + + if method == "ratio": + ratio = float((pneumatic / perkins).mean()) + def converter_ratio(p: float, pachy: Optional[float] = None) -> float: + return p * ratio + return converter_ratio + + elif method == "ols": + from scipy import stats as _stats + slope, intercept, *_ = _stats.linregress(perkins, pneumatic) + slope, intercept = float(slope), float(intercept) + def converter_ols(p: float, pachy: Optional[float] = None) -> float: + return p * slope + intercept + return converter_ols + + elif method == "lad": + from scipy import stats as _stats + from scipy.optimize import minimize as _minimize + slope0, intercept0, *_ = _stats.linregress(perkins, pneumatic) + def _lad_loss(params): + a, b = params + return np.abs(pneumatic - (a * perkins + b)).mean() + res = _minimize(_lad_loss, x0=[slope0, intercept0], method="Nelder-Mead") + slope, intercept = float(res.x[0]), float(res.x[1]) + def converter_lad(p: float, pachy: Optional[float] = None) -> float: + return p * slope + intercept + return converter_lad + + elif method == "multi": + from numpy.linalg import lstsq as _lstsq + paired_multi = combined.dropna(subset=["Pneumatic", "Perkins", "Pachymetry"]) + if len(paired_multi) == 0: + raise ValueError("No paired Pneumatic+Perkins+Pachymetry rows; cannot fit multi method.") + pneu = paired_multi["Pneumatic"].values.astype(float) + perk = paired_multi["Perkins"].values.astype(float) + pachy_vals = paired_multi["Pachymetry"].values.astype(float) + X = np.column_stack([perk, pachy_vals, np.ones(len(perk))]) + coeffs, *_ = _lstsq(X, pneu, rcond=None) + slope, pachy_coef, intercept = float(coeffs[0]), float(coeffs[1]), float(coeffs[2]) + pachy_fallback = float(pachy_vals.mean()) + def converter_multi(p: float, pachy: Optional[float] = None) -> float: + pv = pachy if (pachy is not None and not np.isnan(pachy)) else pachy_fallback + return p * slope + pachy_coef * pv + intercept + return converter_multi + + else: + raise ValueError(f"Unknown iop_corr_method: {method!r}. Choose ratio/ols/lad/multi.") + + +def _pick_iop(row: pd.Series, converter: Callable) -> float: + """Prefer Pneumatic; convert Perkins to Pneumatic scale if Pneumatic is absent.""" + pneumatic = row.get("Pneumatic", np.nan) + if not pd.isna(pneumatic): + return float(pneumatic) + perkins = row.get("Perkins", np.nan) + if pd.isna(perkins): + return np.nan + pachy = row.get("Pachymetry", np.nan) + return converter(float(perkins), None if pd.isna(pachy) else float(pachy)) + + +def _correct_iop(raw_iop: float, pachy: float) -> float: + """Return corrected IOP using nearest pachymetry bin; if pachy missing, return raw.""" + if pd.isna(raw_iop): + return np.nan + if pd.isna(pachy): + return float(raw_iop) + key = _nearest_pachy_key(float(pachy)) + return float(raw_iop) + float(_PACHY_TABLE[key]) + + +def _apply_iop_and_drop_md( + df: pd.DataFrame, + converter: Callable, + drop_raw: bool = False, +) -> pd.DataFrame: + """Add IOP_raw/IOP_corr and drop source IOP columns + VF_MD if present (in-place safe).""" + df["IOP_raw"] = df.apply(lambda row: _pick_iop(row, converter), axis=1) + pachy = df.get("Pachymetry", pd.Series(np.nan, index=df.index)) + df["IOP_corr"] = [ + _correct_iop(r, p) for r, p in zip(df["IOP_raw"].values, pachy.values) + ] + drop_cols = [c for c in ("Pneumatic", "Perkins", "VF_MD") if c in df.columns] + if drop_raw: + drop_cols.append("IOP_raw") + if drop_cols: + df.drop(columns=drop_cols, inplace=True) + return df + + +def _canonicalize_eye_column(df: pd.DataFrame) -> None: + if "eyeID" in df.columns: + src = "eyeID" + else: + src = None + for c in df.columns: + if "eye" in c.lower(): + src = c + break + if src is None: + df["eyeID"] = "OS" + return + + s = df[src] + + def norm(v): + if pd.isna(v): + return None + x = str(v).strip().upper() + if x in {"OS", "L", "LEFT", "0"}: + return "OS" + if x in {"OD", "R", "RIGHT", "1"}: + return "OD" + try: + num = int(float(x)) + return "OD" if num % 2 == 1 else "OS" + except Exception: + return None + + mapped = s.map(norm) + uniq = {u for u in mapped.dropna().unique().tolist()} + if not uniq.issubset({"OS", "OD"}): + raise ValueError(f"eyeID must be binary; found values {sorted(uniq)}") + df["eyeID"] = mapped.fillna("OS") + + +def build_papila_data( + *, + image_dir: str, + clinical_dir: str, + label_col: str, + cat_cols: List[str], + n_splits: int = 5, + random_seed: int = 42, + iop_corr_method: str = "ratio", + iop_drop_raw: bool = False, + exclude_cols: Optional[List[str]] = None, +) -> DataBundle: + """ + Build a DataBundle for PAPILA with dataset-specific preprocessing: + - load OD/OS Excel sheets + - normalize Patient ID + - canonicalize eyeID + - compute IOP_raw / IOP_corr, drop VF_MD + - build feature typing & folds + """ + _exclude = list(exclude_cols) if exclude_cols else [] + + # Remove excluded cols from cat_cols too so the bundle doesn't try to encode them + effective_cat_cols = [c for c in cat_cols if c not in _exclude] + + bundle = DataBundle( + image_dir=image_dir, + clinical_dir=clinical_dir, + label_col=label_col, + patient_col="Patient ID", + cat_cols=effective_cat_cols, + n_splits=n_splits, + random_seed=random_seed, + filename_template="RET{pid:03d}{eye}.jpg", + ) + + od = pd.read_excel(f"{clinical_dir}/patient_data_od.xlsx", header=1) + od["eyeID"] = "OD" + os = pd.read_excel(f"{clinical_dir}/patient_data_os.xlsx", header=1) + os["eyeID"] = "OS" + + for frame in (od, os): + if "Patient ID" not in frame.columns and "ID" in frame.columns: + frame.rename(columns={"ID": "Patient ID"}, inplace=True) + frame["Patient ID"] = frame["Patient ID"].astype(str).str.extract(r"(\d+)")[0].astype(int) + _canonicalize_eye_column(frame) + + bundle.add_df(od, id_column="ID", exclude_cols=_exclude or None) + bundle.add_df(os, id_column="ID", exclude_cols=_exclude or None) + + converter = _fit_perkins_converter(bundle.frames, method=iop_corr_method) + for i in range(len(bundle.frames)): + bundle.frames[i] = _apply_iop_and_drop_md( + bundle.frames[i], converter=converter, drop_raw=iop_drop_raw + ) + + bundle._refresh_master_df(exclude_cols=_exclude or None) + bundle._infer_or_validate_feature_types(exclude_cols=_exclude or None) + bundle._compute_numeric_stats() + bundle._build_cat_maps() + bundle._compute_feature_dim() + bundle._build_kfold_indices() + + return bundle diff --git a/v3/classes/papila_data.py b/v3/classes/papila_data.py new file mode 100644 index 0000000..b4e8ee3 --- /dev/null +++ b/v3/classes/papila_data.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Iterable, Optional + +import pandas as pd + +from v3.classes.data_bundle import DataBundle +from v3.classes.papila_builders import build_papila_data + + +@dataclass +class PapilaData: + """ + V2-friendly wrapper around the DataBundle pipeline. + + Keeps all formatting/normalization behavior from build_papila_clinical, + but exposes a minimal surface area for the V2 engine. + """ + + clinical: DataBundle + patient_col: str = "Patient ID" + + @property + def df(self) -> pd.DataFrame: + return self.clinical.df + + @property + def label_col(self) -> str: + return self.clinical.label_col + + @property + def feature_dim(self) -> int: + return self.clinical.feature_dim + + def get_image_path(self, row: pd.Series): + return self.clinical.get_image_path(row) + + def vectorize_row(self, row: pd.Series): + return self.clinical.vectorize_row(row) + + @classmethod + def from_dirs( + cls, + *, + image_dir: str, + clinical_dir: str, + label_col: str, + cat_cols: Iterable[str], + n_splits: int = 5, + random_seed: int = 42, + ) -> "PapilaData": + clinical = build_papila_data( + image_dir=image_dir, + clinical_dir=clinical_dir, + label_col=label_col, + cat_cols=list(cat_cols), + n_splits=n_splits, + random_seed=random_seed, + ) + return cls(clinical=clinical) diff --git a/v3/classes/predictions.py b/v3/classes/predictions.py new file mode 100644 index 0000000..009ec39 --- /dev/null +++ b/v3/classes/predictions.py @@ -0,0 +1,194 @@ +"""PredictionStore — unified per-epoch prediction tensor across all folds. + +Tensor shape: (n_folds, n_epochs, n_samples, n_heads, n_classes) + +The meaning of "sample" depends on tower_mode: + single — each eye is a sample; sample_ids like "5OD", "14OS" + ensemble — each patient is a sample; sample_ids like "5", "14" + fused — same as ensemble + bilateral— same as ensemble + +Head names by mode: + single : ["fused", "img", "md"] + ensemble : ["od_fused", "od_img", "od_md", "os_fused", "os_img", "os_md"] + fused : ["od_fused", "od_img", "od_md", "os_fused", "os_img", "os_md", "bilat_fused"] + bilateral : ["fused", "img_joint", "md_joint"] +""" +from __future__ import annotations + +from pathlib import Path +from typing import Sequence + +import numpy as np + + +def head_names_for_mode(tower_mode: str, *, fused_head: bool = False) -> list[str]: + """Return canonical head name list for a given tower_mode.""" + if tower_mode in ("single", "classic"): + return ["fused", "img", "md"] + if tower_mode == "ensemble": + names = ["od_fused", "od_img", "od_md", "os_fused", "os_img", "os_md"] + return names + ["bilat_fused"] if fused_head else names + if tower_mode == "bilateral": + return ["fused", "img_joint", "md_joint"] + raise ValueError(f"Unknown tower_mode: {tower_mode!r}") + + +class PredictionStore: + """ + Stores per-epoch predictions for every sample, head, and fold in one tensor. + + Usage + ----- + # Build once before the fold loop: + store = PredictionStore( + sample_ids=all_eye_or_patient_ids, + y_true=all_labels, + head_names=head_names_for_mode(tower_mode, fused_head=args.fused_head), + n_folds=n_folds, + n_epochs=total_epochs, + n_classes=num_classes, + ) + + # Inside each epoch, after collecting probs: + store.record(fold, epoch, patient_ids_batch, "od_fused", probs_od) + store.set_split(fold, train_ids, "train") + store.set_split(fold, val_ids, "val") + + # After all folds: + store.save(run_dir / "predictions.npz") + + # Load and query: + store = PredictionStore.load("predictions.npz") + store.query("5", "od_fused", fold=0) # → (n_epochs, n_classes) + store.query("5", "od_fused") # → (n_folds, n_epochs, n_classes) + store.get_split("5", fold=0) # → "train" + """ + + def __init__( + self, + sample_ids: Sequence[str], + y_true: Sequence[int], + head_names: Sequence[str], + n_folds: int, + n_epochs: int, + n_classes: int, + ): + self.sample_ids = np.array(sample_ids, dtype=object) + self.y_true = np.array(y_true, dtype=np.int64) + self.head_names = np.array(head_names, dtype=object) + self.n_folds = n_folds + self.n_epochs = n_epochs + self.n_classes = n_classes + + n_samples = len(self.sample_ids) + n_heads = len(self.head_names) + + self.probs = np.full( + (n_folds, n_epochs, n_samples, n_heads, n_classes), + fill_value=np.nan, + dtype=np.float32, + ) + self.split = np.full((n_folds, n_samples), fill_value="", dtype=object) + + self._sid_index: dict[str, int] = {str(s): i for i, s in enumerate(self.sample_ids)} + self._head_index: dict[str, int] = {str(h): i for i, h in enumerate(self.head_names)} + + # ------------------------------------------------------------------ + # Writing + # ------------------------------------------------------------------ + + def record( + self, + fold: int, + epoch: int, + sample_ids: Sequence[str], + head_name: str, + probs: np.ndarray, + ) -> None: + """Record a batch of predictions for one head. + + Args: + fold: 0-indexed fold number + epoch: 0-indexed epoch number + sample_ids: sequence of sample ID strings (length B) + head_name: which head — must be in self.head_names + probs: (B, n_classes) probability array + """ + head_idx = self._head_index.get(head_name) + if head_idx is None: + return # head not active in this mode — skip silently + for i, sid in enumerate(sample_ids): + s_idx = self._sid_index.get(str(sid)) + if s_idx is not None: + self.probs[fold, epoch, s_idx, head_idx, :] = probs[i] + + def set_split( + self, + fold: int, + sample_ids: Sequence[str], + label: str, + ) -> None: + """Label a group of samples as 'train', 'val', or 'holdout' for a fold.""" + for sid in sample_ids: + s_idx = self._sid_index.get(str(sid)) + if s_idx is not None: + self.split[fold, s_idx] = label + + # ------------------------------------------------------------------ + # Querying + # ------------------------------------------------------------------ + + def query( + self, + sample_id: str, + head_name: str, + fold: int | None = None, + ) -> np.ndarray: + """Return epoch-level predictions for one sample + head. + + Returns: + fold=None → (n_folds, n_epochs, n_classes) + fold=int → (n_epochs, n_classes) + """ + s_idx = self._sid_index[str(sample_id)] + head_idx = self._head_index[str(head_name)] + if fold is None: + return self.probs[:, :, s_idx, head_idx, :] + return self.probs[fold, :, s_idx, head_idx, :] + + def get_split(self, sample_id: str, fold: int) -> str: + """Return the split label ('train'/'val'/'holdout') for a sample in a fold.""" + s_idx = self._sid_index[str(sample_id)] + return str(self.split[fold, s_idx]) + + # ------------------------------------------------------------------ + # Persistence + # ------------------------------------------------------------------ + + def save(self, path: str | Path) -> None: + np.savez_compressed( + path, + probs=self.probs, + split=self.split, + sample_ids=self.sample_ids, + y_true=self.y_true, + head_names=self.head_names, + ) + + @classmethod + def load(cls, path: str | Path) -> "PredictionStore": + data = np.load(path, allow_pickle=True) + probs = data["probs"] + n_folds, n_epochs, _, _, n_classes = probs.shape + store = cls( + sample_ids=data["sample_ids"].tolist(), + y_true=data["y_true"], + head_names=data["head_names"].tolist(), + n_folds=n_folds, + n_epochs=n_epochs, + n_classes=n_classes, + ) + store.probs = probs + store.split = data["split"] + return store diff --git a/v3/classes/profiles/__init__.py b/v3/classes/profiles/__init__.py new file mode 100644 index 0000000..b903aa0 --- /dev/null +++ b/v3/classes/profiles/__init__.py @@ -0,0 +1,10 @@ +from .base import DatasetProfile, SimpleDatasetProfile, SlotDescriptor +from .papila import PapilaProfile, build_papila_profile + +__all__ = [ + "DatasetProfile", + "SimpleDatasetProfile", + "SlotDescriptor", + "PapilaProfile", + "build_papila_profile", +] diff --git a/v3/classes/profiles/base.py b/v3/classes/profiles/base.py new file mode 100644 index 0000000..2c692b4 --- /dev/null +++ b/v3/classes/profiles/base.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Protocol, Any + +import pandas as pd + + +@dataclass(frozen=True) +class SlotDescriptor: + """ + Metadata for a generic batch slot key (e.g., image_1, matrix_1). + """ + + key: str + kind: str + description: str + required: bool = True + shape_hint: str | None = None + + +class DatasetProfile(Protocol): + """ + Dataset-specific wiring that stays outside the generic V2 engine. + """ + + name: str + patient_col: str + label_col: str + + def slot_descriptors(self) -> dict[str, SlotDescriptor]: + ... + + def semantic_aliases(self) -> dict[str, str]: + ... + + def build_samples(self, *, df: pd.DataFrame, clinical: Any) -> list[dict[str, Any]]: + ... + + +@dataclass(frozen=True) +class SimpleDatasetProfile: + name: str + patient_col: str + label_col: str + slots: dict[str, SlotDescriptor] + aliases: dict[str, str] + + def slot_descriptors(self) -> dict[str, SlotDescriptor]: + return dict(self.slots) + + def semantic_aliases(self) -> dict[str, str]: + return dict(self.aliases) diff --git a/v3/classes/profiles/papila.py b/v3/classes/profiles/papila.py new file mode 100644 index 0000000..4a33062 --- /dev/null +++ b/v3/classes/profiles/papila.py @@ -0,0 +1,155 @@ +from __future__ import annotations + +import pandas as pd + +from dataclasses import dataclass + +from .base import SimpleDatasetProfile, SlotDescriptor + + +@dataclass(frozen=True) +class PapilaProfile(SimpleDatasetProfile): + sample_mode: str = "patient" # "patient" | "eye" + + def build_samples(self, *, df: pd.DataFrame, clinical) -> list[dict[str, object]]: + samples: list[dict[str, object]] = [] + patient_col = self.patient_col + label_col = self.label_col + + mode = (self.sample_mode or "patient").lower() + if mode not in {"patient", "eye"}: + raise ValueError(f"Unsupported sample_mode '{self.sample_mode}'. Expected 'patient' or 'eye'.") + + if mode == "eye": + for _, row in df.iterrows(): + pid = row[patient_col] + label = row[label_col] + image_1 = clinical.get_image_path(row) if hasattr(clinical, "get_image_path") else None + matrix_1 = clinical.vectorize_row(row) if hasattr(clinical, "vectorize_row") else None + samples.append( + { + "id_1": pid, + "label_1": label, + "image_1": image_1, + "matrix_1": matrix_1, + } + ) + return samples + + for pid, grp in df.groupby(patient_col): + label_series = grp[label_col] + if label_series.empty: + continue + mode_vals = label_series.mode() + label = mode_vals.iloc[0] if not mode_vals.empty else label_series.iloc[0] + + def _row_for_eye(eye: str): + if "eyeID" not in grp.columns: + return None + match = grp[grp["eyeID"].astype(str).str.upper() == eye] + if match.empty: + return None + return match.iloc[0] + + row_od = _row_for_eye("OD") + row_os = _row_for_eye("OS") + row_any = grp.iloc[0] + + image_1 = clinical.get_image_path(row_od) if row_od is not None else None + image_2 = clinical.get_image_path(row_os) if row_os is not None else None + matrix_1 = clinical.vectorize_row(row_od) if row_od is not None else None + matrix_2 = clinical.vectorize_row(row_os) if row_os is not None else None + + if image_1 is None and hasattr(clinical, "get_image_path"): + image_1 = clinical.get_image_path(row_any) + if matrix_1 is None and hasattr(clinical, "vectorize_row"): + matrix_1 = clinical.vectorize_row(row_any) + + samples.append( + { + "id_1": pid, + "label_1": label, + "image_1": image_1, + "image_2": image_2, + "matrix_1": matrix_1, + "matrix_2": matrix_2, + } + ) + return samples + + +def build_papila_profile( + *, + patient_col: str = "Patient ID", + label_col: str = "Diagnosis", + sample_mode: str = "patient", +) -> PapilaProfile: + """ + PAPILA-specific semantic map for generic V2 slot keys. + + The engine remains slot-based (image_1/image_2/matrix_1/...). + PAPILA meaning is captured here so run config stays dataset-local. + """ + + slots = { + "id_1": SlotDescriptor( + key="id_1", + kind="id", + description=f"Patient identifier column ({patient_col})", + required=True, + shape_hint="scalar", + ), + "label_1": SlotDescriptor( + key="label_1", + kind="label", + description=f"Diagnosis label column ({label_col})", + required=True, + shape_hint="scalar", + ), + "image_1": SlotDescriptor( + key="image_1", + kind="image", + description="Fundus image slot 1 (PAPILA: OD / right eye)", + required=False, + shape_hint="HWC or CHW", + ), + "image_2": SlotDescriptor( + key="image_2", + kind="image", + description="Fundus image slot 2 (PAPILA: OS / left eye)", + required=False, + shape_hint="HWC or CHW", + ), + "matrix_1": SlotDescriptor( + key="matrix_1", + kind="matrix", + description="Clinical metadata feature vector", + required=False, + shape_hint="[feature_dim]", + ), + "matrix_2": SlotDescriptor( + key="matrix_2", + kind="matrix", + description="Optional auxiliary tabular vector (reserved for experiments)", + required=False, + shape_hint="[feature_dim_2]", + ), + } + + aliases = { + "id_1": "patient_id", + "label_1": "diagnosis", + "image_1": "od_fundus", + "image_2": "os_fundus", + "matrix_1": "clinical_metadata", + "matrix_2": "aux_metadata", + } + + return PapilaProfile( + name="papila", + patient_col=patient_col, + label_col=label_col, + slots=slots, + aliases=aliases, + sample_mode=sample_mode, + ) diff --git a/v3/classes/results.py b/v3/classes/results.py new file mode 100644 index 0000000..3ee9846 --- /dev/null +++ b/v3/classes/results.py @@ -0,0 +1,138 @@ +"""Result dataclasses and serialisation helpers for V3 fold outputs.""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +import numpy as np + + +def _nan() -> float: + return float("nan") + + +def _f(v) -> Optional[float]: + """Round a scalar to 6 dp, return None for nan/None.""" + if v is None or (isinstance(v, float) and np.isnan(v)): + return None + return round(float(v), 6) + + +def _sv(vec) -> Optional[str]: + """Serialise a vector to a pipe-separated string, or None.""" + if vec is None: + return None + return "|".join(f"{float(v):.4f}" for v in vec) + + +@dataclass +class FoldResult: + mode: str + fold: int + # Epoch where each model hit its peak val AUC + best_epoch_single: int + best_epoch_bilat: int + # Classic (eye-level eval of SingleEyeHT) + classic_val_auc: float + classic_val_acc: float + classic_val_kappa: float + classic_val_mcc: float + classic_val_f1: float + classic_val_recall: Optional[str] + classic_val_ece: float + classic_val_threshold: float + classic_val_bias: Optional[str] + classic_val_n: int + # Ensemble (patient-level eval of SingleEyeHT) + ensemble_val_auc: float + ensemble_val_acc: float + ensemble_val_kappa: float + ensemble_val_mcc: float + ensemble_val_f1: float + ensemble_val_recall: Optional[str] + ensemble_val_ece: float + ensemble_val_threshold: float + ensemble_val_bias: Optional[str] + ensemble_val_n: int + # Bilateral (BilateralHT patient-level) + bilat_val_auc: float + bilat_val_acc: float + bilat_val_kappa: float + bilat_val_mcc: float + bilat_val_f1: float + bilat_val_recall: Optional[str] + bilat_val_ece: float + bilat_val_threshold: float + bilat_val_bias: Optional[str] + bilat_val_n: int + # Test metrics (evaluated once after training on final-epoch model) + ensemble_test_auc: float = float("nan") + ensemble_test_acc: float = float("nan") + ensemble_test_kappa: float = float("nan") + ensemble_test_f1: float = float("nan") + ensemble_test_ece: float = float("nan") + classic_test_auc: float = float("nan") + classic_test_acc: float = float("nan") + classic_test_kappa: float = float("nan") + classic_test_f1: float = float("nan") + classic_test_ece: float = float("nan") + bilat_test_auc: float = float("nan") + bilat_test_acc: float = float("nan") + bilat_test_kappa: float = float("nan") + bilat_test_f1: float = float("nan") + bilat_test_ece: float = float("nan") + test_n: int = 0 + # Training sample counts + single_train_n: int = 0 + bilat_train_n: int = 0 + # Fused head (optional) + fused_val_auc: float = float("nan") + fused_val_acc: float = float("nan") + fused_val_kappa: float = float("nan") + fused_val_mcc: float = float("nan") + fused_val_f1: float = float("nan") + fused_val_recall: Optional[str] = None + fused_val_ece: float = float("nan") + fused_val_threshold: float = float("nan") + fused_val_bias: Optional[str] = None + fused_val_n: int = 0 + fused_test_auc: float = float("nan") + fused_test_acc: float = float("nan") + + +@dataclass +class FoldArtifacts: + # Val split + y_true_classic: Optional[np.ndarray] + probs_classic: Optional[np.ndarray] + y_true_ensemble: Optional[np.ndarray] + probs_ensemble: Optional[np.ndarray] + y_true_bilat: Optional[np.ndarray] + probs_bilat: Optional[np.ndarray] + y_true_fused: Optional[np.ndarray] = None + probs_fused: Optional[np.ndarray] = None + probs_ensemble_img: Optional[np.ndarray] = None + probs_ensemble_md: Optional[np.ndarray] = None + probs_classic_img: Optional[np.ndarray] = None + probs_classic_md: Optional[np.ndarray] = None + # per-eye (pre-averaged) for ensemble mode + y_true_ensemble_pereye: Optional[np.ndarray] = None + probs_ensemble_pereye: Optional[np.ndarray] = None + probs_ensemble_img_pereye: Optional[np.ndarray] = None + probs_ensemble_md_pereye: Optional[np.ndarray] = None + # raw logits — patient-level + logits_ensemble: Optional[np.ndarray] = None + logits_ensemble_img: Optional[np.ndarray] = None + logits_ensemble_md: Optional[np.ndarray] = None + logits_classic: Optional[np.ndarray] = None + logits_classic_img: Optional[np.ndarray] = None + logits_classic_md: Optional[np.ndarray] = None + # raw logits — per-eye + logits_ensemble_pereye: Optional[np.ndarray] = None + logits_ensemble_img_pereye: Optional[np.ndarray] = None + logits_ensemble_md_pereye: Optional[np.ndarray] = None + # Test split equivalents + y_true_test: Optional[np.ndarray] = None + probs_test: Optional[np.ndarray] = None + probs_test_img: Optional[np.ndarray] = None + probs_test_md: Optional[np.ndarray] = None diff --git a/v3/classes/slot_dataset.py b/v3/classes/slot_dataset.py new file mode 100644 index 0000000..c3ccd4c --- /dev/null +++ b/v3/classes/slot_dataset.py @@ -0,0 +1,114 @@ +from __future__ import annotations + +from typing import Any, Callable, Optional + +from pathlib import Path +from PIL import Image +import numpy as np +import torch +from torch.utils.data import Dataset +from torchvision import transforms + +from .image_loader import CachedImageLoader +from .profiles.base import SlotDescriptor + + +def slot_collate(batch: list[dict[str, Any]]) -> dict[str, Any]: + if not batch: + return {} + keys = batch[0].keys() + out: dict[str, Any] = {} + for key in keys: + vals = [item.get(key) for item in batch] + if all(isinstance(v, torch.Tensor) for v in vals): + try: + out[key] = torch.stack(vals, dim=0) + except Exception: + out[key] = vals + else: + out[key] = vals + return out + + +class SlotDataset(Dataset): + """ + Dataset that yields dicts of slot-keyed values. + + Sample records are expected to be dicts with keys matching slot descriptors. + Image slots accept filesystem paths; matrix slots accept array-like values. + """ + + def __init__( + self, + samples: list[dict[str, Any]], + slot_descriptors: dict[str, SlotDescriptor], + *, + image_transform: Optional[Callable[[Image.Image], torch.Tensor]] = None, + matrix_transform: Optional[Callable[[Any], torch.Tensor]] = None, + image_preprocessor: Optional[Callable[..., Image.Image]] = None, + image_cache: Optional[CachedImageLoader] = None, + ) -> None: + self.samples = samples + self.slot_descriptors = slot_descriptors + self.image_transform = image_transform or transforms.ToTensor() + self.matrix_transform = matrix_transform or self._default_matrix_transform + self.image_preprocessor = image_preprocessor + self.image_cache = image_cache + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__(self, idx: int) -> dict[str, Any]: + record = self.samples[idx] + out: dict[str, Any] = {} + for key, desc in self.slot_descriptors.items(): + val = record.get(key) + if desc.kind == "image": + out[key] = self._load_image(val, required=desc.required) + elif desc.kind == "matrix": + out[key] = self._load_matrix(val, required=desc.required) + else: + out[key] = val + return out + + def _load_image(self, value: Any, *, required: bool) -> Optional[torch.Tensor]: + if value is None: + if required: + raise ValueError("Missing required image slot") + return None + if self.image_cache is not None: + img = self.image_cache.load(value, preprocessor=self.image_preprocessor) + else: + from .image_loader import _call_preprocessor + img = Image.open(value).convert("RGB") + if self.image_preprocessor is not None: + img = _call_preprocessor(self.image_preprocessor, img, Path(value)) + return self.image_transform(img) + + def prebuild_image_cache(self, cache_workers: int = 0) -> None: + """Pre-populate the CachedImageLoader for all image paths in this dataset.""" + if self.image_cache is None: + return + paths = list({ + str(record[key]) + for record in self.samples + for key, desc in self.slot_descriptors.items() + if desc.kind == "image" and record.get(key) is not None + }) + self.image_cache._workers = cache_workers or self.image_cache._workers + self.image_cache.warm(paths, preprocessor=self.image_preprocessor) + + def _load_matrix(self, value: Any, *, required: bool) -> Optional[torch.Tensor]: + if value is None: + if required: + raise ValueError("Missing required matrix slot") + return None + return self.matrix_transform(value) + + @staticmethod + def _default_matrix_transform(value: Any) -> torch.Tensor: + if isinstance(value, torch.Tensor): + return value.float() + if isinstance(value, np.ndarray): + return torch.from_numpy(value.astype(np.float32, copy=False)) + return torch.as_tensor(value, dtype=torch.float32) diff --git a/v3/classes/split_manager.py b/v3/classes/split_manager.py new file mode 100644 index 0000000..723de75 --- /dev/null +++ b/v3/classes/split_manager.py @@ -0,0 +1,172 @@ +"""V3 split manager — proper outer/inner k-fold CV. + +Outer fold k = test set. +Val = outer fold (k+1) % n_splits (rotated). +Train = remaining n_splits-2 folds. + +Every patient appears in test exactly once and in val exactly once. +No pre-carved holdout. +""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Iterable, Optional + +import numpy as np +import pandas as pd +from sklearn.model_selection import KFold, StratifiedKFold + +from .network_manager import PatientSplit + + +@dataclass(frozen=True) +class SplitPlan: + train_patient_ids: set[Any] + val_patient_ids: set[Any] + test_patient_ids: set[Any] + + +def build_patient_split_plans( + patient_ids: Iterable[Any], + patient_labels: Iterable[Any], + *, + n_splits: int, + seed: int, +) -> list[SplitPlan]: + """ + Outer/inner k-fold splitter. + + For each outer fold k: + - test = patients in fold k + - val = patients in fold (k+1) % n_splits + - train = patients in remaining n_splits-2 folds + """ + ids = np.asarray(list(patient_ids)) + labels = np.asarray(list(patient_labels)) + if ids.ndim != 1 or labels.ndim != 1: + raise ValueError("patient_ids and patient_labels must be 1D arrays") + if ids.size != labels.size: + raise ValueError(f"Length mismatch: ids={ids.size}, labels={labels.size}") + if ids.size == 0: + raise ValueError("No patients available for splitting") + if len(set(ids.tolist())) != ids.size: + raise ValueError("patient_ids must be unique") + if n_splits < 3: + raise ValueError("n_splits must be >= 3 for outer/inner k-fold") + + use_stratified = _can_stratify(labels, n_splits) + if use_stratified: + splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed) + outer_folds = list(splitter.split(ids, labels)) + else: + splitter = KFold(n_splits=n_splits, shuffle=True, random_state=seed) + outer_folds = list(splitter.split(ids)) + + # Build index sets for each outer fold + fold_index_sets: list[set] = [] + for _, test_idx in outer_folds: + fold_index_sets.append(set(ids[test_idx].tolist())) + + plans: list[SplitPlan] = [] + for k in range(n_splits): + test_ids = fold_index_sets[k] + val_ids = fold_index_sets[(k + 1) % n_splits] + train_ids: set = set() + for j in range(n_splits): + if j != k and j != (k + 1) % n_splits: + train_ids |= fold_index_sets[j] + plans.append(SplitPlan( + train_patient_ids=train_ids, + val_patient_ids=val_ids, + test_patient_ids=test_ids, + )) + return plans + + +class PatientFirstSplitManager: + """Patient-level splitter for V3. Outer/inner k-fold, no holdout.""" + + def __init__( + self, + *, + patient_col: str = "Patient ID", + label_col: Optional[str] = None, + ) -> None: + self.patient_col = patient_col + self.label_col = label_col + + def build_plans( + self, + *, + clinical: Any, + args: Any, + profile: Optional[Any] = None, + ) -> list[PatientSplit]: + profile_label_col = getattr(profile, "label_col", None) if profile is not None else None + profile_patient_col = getattr(profile, "patient_col", None) if profile is not None else None + patient_col = profile_patient_col or self.patient_col + label_col = self.label_col or profile_label_col or getattr(clinical, "label_col", None) + if label_col is None: + raise ValueError("Could not resolve label column") + + if not hasattr(clinical, "df"): + raise ValueError("Clinical object must expose a dataframe at .df") + df_full = clinical.df.copy() + self._validate_columns(df_full, label_col, patient_col=patient_col) + + eval_mode = str(getattr(args, "eval_mode", "multiclass")).lower() + if eval_mode == "binary": + df_full = df_full[df_full[label_col].isin([0, 1])].reset_index(drop=True) + + n_splits = int(getattr(args, "n_splits", 5)) + fold_seed = int(getattr(args, "fold_seed", 42)) + + patient_table = self._patient_label_table(df_full, label_col, patient_col=patient_col) + plans = build_patient_split_plans( + patient_ids=patient_table[patient_col].to_numpy(), + patient_labels=patient_table["_label"].to_numpy(), + n_splits=n_splits, + seed=fold_seed, + ) + + out: list[PatientSplit] = [] + for plan in plans: + train_df = df_full[df_full[patient_col].isin(plan.train_patient_ids)].reset_index(drop=True) + val_df = df_full[df_full[patient_col].isin(plan.val_patient_ids)].reset_index(drop=True) + test_df = df_full[df_full[patient_col].isin(plan.test_patient_ids)].reset_index(drop=True) + out.append(PatientSplit(train=train_df, val=val_df, test=test_df)) + return out + + def _validate_columns(self, df: pd.DataFrame, label_col: str, patient_col: Optional[str] = None) -> None: + pcol = patient_col or self.patient_col + if pcol not in df.columns: + raise ValueError(f"Missing required patient column: {pcol!r}") + if label_col not in df.columns: + raise ValueError(f"Missing required label column: {label_col!r}") + + def _patient_label_table( + self, + df: pd.DataFrame, + label_col: str, + patient_col: Optional[str] = None, + ) -> pd.DataFrame: + pcol = patient_col or self.patient_col + grouped = ( + df.groupby(pcol, as_index=False)[label_col] + .agg(lambda x: x.mode().iloc[0] if not x.mode().empty else x.iloc[0]) + .rename(columns={label_col: "_label"}) + .sort_values(pcol) + .reset_index(drop=True) + ) + if grouped.empty: + raise ValueError("No patients available for splitting") + return grouped + + +def _can_stratify(labels: np.ndarray, n_splits: int) -> bool: + if labels.size == 0: + return False + unique, counts = np.unique(labels, return_counts=True) + if len(unique) < 2: + return False + return bool(np.all(counts >= n_splits)) diff --git a/v3/classes/towers.py b/v3/classes/towers.py new file mode 100644 index 0000000..4b0868a --- /dev/null +++ b/v3/classes/towers.py @@ -0,0 +1,279 @@ +from __future__ import annotations + +import math +from typing import Optional + +import torch +from torch import nn +from torchvision import transforms + +from v3.classes.backbones import BACKBONES, list_names, load_backbone_weights +from v3.classes.SE_attention import SEBlock +from v3.classes.data_bundle import DataBundle + + +def build_backbone(name: str, freeze_ratio: float = 0.0, augment: bool = True): + """ + Operational builder: + - instantiate with DEFAULT weights + - strip classifier → features + - apply ratio-based freezing over coarse blocks + - return (model, out_dim, transform) + """ + key = (name or "").lower() + if key not in BACKBONES: + raise ValueError(f"Unsupported backbone '{name}'. Valid options: {list_names()}") + + spec = BACKBONES[key] + m = spec.ctor(weights=spec.weights_default) + out_dim, m = spec.strip(m) + load_backbone_weights(key, m) + + # transforms: use the weights’ mean/std, but keep your augmentation pipeline + mean = getattr(spec.weights_default, "meta", {}).get("mean", (0.485, 0.456, 0.406)) + std = getattr(spec.weights_default, "meta", {}).get("std", (0.229, 0.224, 0.225)) + crop = 299 if key == "inception_v3" else 224 + + if augment: + transform = transforms.Compose( + [ + transforms.Resize(256), + transforms.CenterCrop(crop), + transforms.RandomHorizontalFlip(), + transforms.RandomVerticalFlip(), + transforms.RandomRotation(15), + transforms.ColorJitter(0.1, 0.1, 0.1, 0.05), + transforms.ToTensor(), + transforms.Normalize(mean=mean, std=std), + ] + ) + else: + transform = transforms.Compose( + [ + transforms.Resize(256), + transforms.CenterCrop(crop), + transforms.ToTensor(), + transforms.Normalize(mean=mean, std=std), + ] + ) + + # ratio-based freezing: freeze earliest floor(N * freeze_ratio) blocks + fr = max(0.0, min(1.0, float(freeze_ratio))) + blocks = spec.blocks(m) + n = len(blocks) + freeze_n = int(math.floor(n * fr)) + for b in blocks[:freeze_n]: + for p in b.parameters(): + p.requires_grad = False + + return m, out_dim, transform + + +class ImageTower(nn.Module): + """ + Vision backbone → pooled features. + - backbone: one of list_names() (default 'efficientnet_b0') + - always DEFAULT torchvision weights + - freeze_ratio ∈ [0,1] freezes earliest floor(N*freeze_ratio) blocks + - returns [N, out_dim] features from backbone forward + """ + + def __init__( + self, + backbone: str = "efficientnet_b0", + freeze_ratio: float = 0.0, + use_se: bool = False, + se_reduction: int = 16, + se_pre_norm: bool = True, + augment: bool = True, + geometry_dim: int = 0, + ): + super().__init__() + self.backbone, base_dim, self.transform = build_backbone( + backbone, freeze_ratio, augment=augment + ) + self._name = backbone + # Keep ordered blocks for dynamic freezing/thawing + key = (self._name or "").lower() + self._spec = BACKBONES[key] + self._blocks = self._spec.blocks(self.backbone) + # Optional tower-level SE over the final feature vector + self.base_dim = base_dim + self.geometry_dim = max(0, int(geometry_dim)) + self.out_dim = self.base_dim + self.geometry_dim + self.tower_ln = nn.LayerNorm(self.base_dim) if se_pre_norm else nn.Identity() + self.tower_se = ( + SEBlock(self.base_dim, reduction=se_reduction, residual=True) + if use_se + else None + ) + + def forward( + self, x: torch.Tensor, geometry: Optional[torch.Tensor] = None + ) -> torch.Tensor: + y = self.backbone(x) + # sanity: pooled features, not logits + assert y.dim() == 2 and y.size(1) == self.base_dim, ( + f"Expected features [N,{self.base_dim}], got {tuple(y.shape)}" + ) + if self.tower_se is not None: + y, _ = self.tower_se(self.tower_ln(y)) + if self.geometry_dim > 0: + if geometry is None or geometry.numel() == 0: + geom = torch.zeros( + y.size(0), self.geometry_dim, device=y.device, dtype=y.dtype + ) + else: + if geometry.dim() == 1: + geom = geometry.unsqueeze(0) + else: + geom = geometry + geom = geom.to(device=y.device, dtype=y.dtype) + if geom.size(0) != y.size(0): + raise ValueError( + f"Geometry batch size mismatch: {geom.size(0)} vs {y.size(0)}" + ) + if geom.size(1) != self.geometry_dim: + raise ValueError( + f"Expected geometry dim {self.geometry_dim}, got {geom.size(1)}" + ) + y = torch.cat([y, geom], dim=1) + return y + + def set_freeze_ratio(self, ratio: float): + """Dynamically freeze earliest floor(N*ratio) backbone blocks.""" + r = max(0.0, min(1.0, float(ratio))) + n = len(self._blocks) + freeze_n = int(math.floor(n * r)) + # Unfreeze all first + for b in self._blocks: + for p in b.parameters(): + p.requires_grad = True + # Freeze earliest blocks + for b in self._blocks[:freeze_n]: + for p in b.parameters(): + p.requires_grad = False + + +class SiameseImageTower(nn.Module): + """ + Shared-weight bilateral image tower. + + Runs OD and OS images through a single shared backbone, then returns + cat([f_mean, f_delta]) where: + f_mean = (f_od + f_os) / 2 -- shared bilateral representation + f_delta = f_od - f_os -- asymmetry, signed OD-relative + + out_dim = 2 * backbone_out_dim + + When x_os is None (single-eye fallback): + f_mean = f_od + f_delta = zeros + so the module degrades gracefully when only one eye is available. + + The shared backbone means both eyes contribute to every gradient update, + effectively doubling the training signal for the visual pathway without + doubling parameters. + """ + + def __init__( + self, + backbone: str = "efficientnet_b0", + freeze_ratio: float = 0.0, + use_se: bool = False, + se_reduction: int = 16, + se_pre_norm: bool = True, + augment: bool = True, + ): + super().__init__() + self._tower = ImageTower( + backbone=backbone, + freeze_ratio=freeze_ratio, + use_se=use_se, + se_reduction=se_reduction, + se_pre_norm=se_pre_norm, + augment=augment, + geometry_dim=0, + ) + self.out_dim = self._tower.out_dim * 2 + self.transform = self._tower.transform + + def forward( + self, + x_od: torch.Tensor, + x_os: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + f_od = self._tower(x_od) + if x_os is None: + f_mean = f_od + f_delta = torch.zeros_like(f_od) + else: + f_os = self._tower(x_os) + f_mean = (f_od + f_os) * 0.5 + f_delta = f_od - f_os + return torch.cat([f_mean, f_delta], dim=1) + + def set_freeze_ratio(self, ratio: float) -> None: + """Delegates to the shared inner tower.""" + self._tower.set_freeze_ratio(ratio) + + +class ClinicalTower(nn.Module): + """MLP over DataBundle.vectorize_row outputs (convert to torch inside tower).""" + + def __init__( + self, + clinical_data: DataBundle, + hidden_dim: int = 128, + dropout: float = 0.1, + use_se: bool = False, + se_reduction: int = 16, + se_pre_norm: bool = True, + ): + super().__init__() + self.feature_dim = clinical_data.feature_dim + self.out_dim = hidden_dim + # two-block MLP so we can optionally freeze/thaw per block + self.block0 = nn.Sequential( + nn.Linear(self.feature_dim, hidden_dim), + nn.LayerNorm(hidden_dim), + nn.ReLU(inplace=True), + nn.Dropout(dropout), + ) + self.block1 = nn.Sequential( + nn.Linear(hidden_dim, hidden_dim), + nn.ReLU(inplace=True), + ) + self.net = nn.Sequential(self.block0, self.block1) + self.tower_ln = nn.LayerNorm(hidden_dim) if se_pre_norm else nn.Identity() + self.tower_se = ( + SEBlock(hidden_dim, reduction=se_reduction, residual=True) + if use_se + else None + ) + + def forward(self, meta_np_or_torch) -> torch.Tensor: + if isinstance(meta_np_or_torch, torch.Tensor): + x = meta_np_or_torch + else: + x = torch.as_tensor(meta_np_or_torch, dtype=torch.float32) + h = self.net(x) + if self.tower_se is not None: + h, _ = self.tower_se(self.tower_ln(h)) + return h + + def set_freeze_ratio(self, ratio: float): + """Optionally freeze earliest blocks of the MLP.""" + r = max(0.0, min(1.0, float(ratio))) + # Unfreeze all + for p in self.block0.parameters(): + p.requires_grad = True + for p in self.block1.parameters(): + p.requires_grad = True + # Freeze earliest blocks based on ratio threshold + if r >= 0.5: + for p in self.block0.parameters(): + p.requires_grad = False + if r >= 1.0: + for p in self.block1.parameters(): + p.requires_grad = False diff --git a/v3/classes/transforms.py b/v3/classes/transforms.py new file mode 100644 index 0000000..2e4fd8c --- /dev/null +++ b/v3/classes/transforms.py @@ -0,0 +1,320 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Callable, Iterable, Optional, Tuple, Union +import numpy as np +from PIL import Image + +from torchvision import transforms + +from v3.classes.backbones import BACKBONES + + +IMAGENET_MEAN: Tuple[float, float, float] = (0.485, 0.456, 0.406) +IMAGENET_STD: Tuple[float, float, float] = (0.229, 0.224, 0.225) + + +@dataclass +class ImageTransformConfig: + """ + Mirrors the hypertower v1 preprocessing: + - Resize(256) + - CenterCrop(crop) + - Optional augmentations (H/V flip, rotation, color jitter) + - ToTensor + Normalize(mean/std) + """ + + crop_size: int = 224 + resize_size: int = 256 + mean: Tuple[float, float, float] = IMAGENET_MEAN + std: Tuple[float, float, float] = IMAGENET_STD + augment: bool = True + rotation_deg: int = 15 + color_jitter: Tuple[float, float, float, float] = (0.1, 0.1, 0.1, 0.05) + hflip: bool = True + vflip: bool = True + + def build(self) -> transforms.Compose: + ops = [ + transforms.Resize(self.resize_size), + transforms.CenterCrop(self.crop_size), + ] + if self.augment: + if self.hflip: + ops.append(transforms.RandomHorizontalFlip()) + if self.vflip: + ops.append(transforms.RandomVerticalFlip()) + if self.rotation_deg: + ops.append(transforms.RandomRotation(self.rotation_deg)) + if self.color_jitter: + ops.append(transforms.ColorJitter(*self.color_jitter)) + ops.extend( + [ + transforms.ToTensor(), + transforms.Normalize(mean=self.mean, std=self.std), + ] + ) + return transforms.Compose(ops) + + +def backbone_transform_config(backbone_name: str, augment: bool = True) -> ImageTransformConfig: + """ + Build a transform config that matches v1 ImageTower/backbone preprocessing. + Uses DEFAULT weights mean/std and InceptionV3 crop size when relevant. + """ + key = (backbone_name or "").lower() + if key not in BACKBONES: + raise ValueError(f"Unsupported backbone '{backbone_name}'.") + spec = BACKBONES[key] + mean = getattr(spec.weights_default, "meta", {}).get("mean", IMAGENET_MEAN) + std = getattr(spec.weights_default, "meta", {}).get("std", IMAGENET_STD) + crop = 299 if key == "inception_v3" else 224 + return ImageTransformConfig(crop_size=crop, mean=mean, std=std, augment=augment) + + +def build_backbone_transform(backbone_name: str, augment: bool = True) -> transforms.Compose: + return backbone_transform_config(backbone_name, augment=augment).build() + + +def build_eval_transform(backbone: str) -> transforms.Compose: + """Deterministic eval transform matching backbone normalisation (no augmentation).""" + return build_backbone_transform(backbone, augment=False) + + +def build_imagenet_transform(augment: bool = True, crop_size: int = 224) -> transforms.Compose: + return ImageTransformConfig(crop_size=crop_size, augment=augment).build() + + +@dataclass +class ResizeTransform: + size: Union[int, Tuple[int, int]] = 256 + interpolation: int = Image.BILINEAR + + def __post_init__(self) -> None: + self._op = transforms.Resize(self.size, interpolation=self.interpolation) + + def __call__(self, image: Image.Image) -> Image.Image: + return self._op(image) + + +@dataclass +class CenterCropTransform: + size: Union[int, Tuple[int, int]] = 224 + + def __post_init__(self) -> None: + self._op = transforms.CenterCrop(self.size) + + def __call__(self, image: Image.Image) -> Image.Image: + return self._op(image) + + +class UnetMaskProvider: + """ + Placeholder for a UNet-powered mask provider. + This will be replaced once a UNet tower is wired in. + """ + + def __call__(self, image: Image.Image, image_path: Optional[str] = None): + raise NotImplementedError("UNet mask provider is not wired yet.") + + +@dataclass +class ROICropTransform: + """ + Crop an image using a binary mask (GT or UNet). + Expects a mask of the same spatial size as the image; nonzero pixels are ROI. + """ + + mask_source: str = "gt" # "gt" | "unet" + mask_provider: Optional[Callable[[Image.Image, Optional[str]], np.ndarray]] = None + scale: float = 2.5 + target_size: Optional[Tuple[int, int]] = (224, 224) + fallback_to_original: bool = True + + def __post_init__(self) -> None: + if self.mask_source not in {"gt", "unet"}: + raise ValueError(f"mask_source must be 'gt' or 'unet', got '{self.mask_source}'.") + + def __call__( + self, + image: Image.Image, + mask: Optional[Union[np.ndarray, Image.Image]] = None, + image_path: Optional[str] = None, + ) -> Image.Image: + resolved_mask = mask + if resolved_mask is None and self.mask_provider is not None: + resolved_mask = self.mask_provider(image, image_path) + if resolved_mask is None: + if self.fallback_to_original: + return image + raise ValueError("ROI crop requested but no mask provided.") + + mask_arr = ( + np.asarray(resolved_mask) + if not isinstance(resolved_mask, Image.Image) + else np.array(resolved_mask) + ) + if mask_arr.ndim == 3: + mask_arr = mask_arr[..., 0] + mask_arr = mask_arr > 0 + if not np.any(mask_arr): + return image if self.fallback_to_original else image + + ys, xs = np.where(mask_arr) + y_min, y_max = ys.min(), ys.max() + x_min, x_max = xs.min(), xs.max() + cx = (x_min + x_max) / 2.0 + cy = (y_min + y_max) / 2.0 + width = (x_max - x_min + 1) + height = (y_max - y_min + 1) + size = max(width, height) * float(self.scale) + + left = int(round(cx - size / 2)) + right = int(round(cx + size / 2)) + upper = int(round(cy - size / 2)) + lower = int(round(cy + size / 2)) + + left = max(0, left) + upper = max(0, upper) + right = min(image.width, right) + lower = min(image.height, lower) + crop = image.crop((left, upper, right, lower)) + if self.target_size is not None: + crop = crop.resize(self.target_size, Image.BILINEAR) + return crop + + +@dataclass +class JitterBundleTransform: + """ + Augmentations bundle: flips, rotation, color jitter. + """ + + hflip: bool = True + vflip: bool = True + rotation_deg: int = 15 + color_jitter: Optional[Tuple[float, float, float, float]] = (0.1, 0.1, 0.1, 0.05) + + def __post_init__(self) -> None: + ops = [] + if self.hflip: + ops.append(transforms.RandomHorizontalFlip()) + if self.vflip: + ops.append(transforms.RandomVerticalFlip()) + if self.rotation_deg: + ops.append(transforms.RandomRotation(self.rotation_deg)) + if self.color_jitter: + ops.append(transforms.ColorJitter(*self.color_jitter)) + self._op = transforms.Compose(ops) if ops else None + + def __call__(self, image: Image.Image) -> Image.Image: + if self._op is None: + return image + return self._op(image) + + +TRANSFORM_REGISTRY = { + "resize": ResizeTransform, + "roi_crop": ROICropTransform, + "center_crop": CenterCropTransform, + "jitter_bundle": JitterBundleTransform, +} + + +def _parse_color_jitter(value: Optional[Union[str, Iterable[float]]]) -> Optional[Tuple[float, float, float, float]]: + if value is None: + return None + if isinstance(value, str): + parts = [p.strip() for p in value.split(",") if p.strip()] + if not parts: + return None + try: + nums = [float(p) for p in parts] + except ValueError: + return None + if len(nums) == 1: + return (nums[0], nums[0], nums[0], nums[0]) + if len(nums) >= 4: + return (nums[0], nums[1], nums[2], nums[3]) + return tuple(nums + [nums[-1]] * (4 - len(nums))) # pad to length 4 + try: + vals = list(value) + except TypeError: + return None + if not vals: + return None + vals = [float(v) for v in vals] + if len(vals) == 1: + return (vals[0], vals[0], vals[0], vals[0]) + if len(vals) >= 4: + return (vals[0], vals[1], vals[2], vals[3]) + return tuple(vals + [vals[-1]] * (4 - len(vals))) + + +def build_transform_chain( + transform_specs: Iterable[object], + *, + backbone_name: str, + augment: bool = True, + mask_provider: Optional[Callable[[Image.Image, Optional[str]], np.ndarray]] = None, + strict: bool = True, +) -> transforms.Compose: + """ + Build an image transform pipeline from a list of transform specs plus the + standard ToTensor + Normalize steps. This mirrors the V1 preprocessing + but uses the explicit transform nodes from config. + """ + ops: list[Callable[[Image.Image], Image.Image]] = [] + for spec in transform_specs: + transform_type = getattr(spec, "transform_type", None) + params = getattr(spec, "params", None) + if transform_type is None and isinstance(spec, dict): + transform_type = spec.get("transformType") or spec.get("transform_type") + params = spec + params = params or {} + + if transform_type == "resize": + size = params.get("resizeSize", 256) + ops.append(ResizeTransform(size=size)) + elif transform_type == "center_crop": + size = params.get("centerCropSize", 224) + ops.append(CenterCropTransform(size=size)) + elif transform_type == "jitter_bundle": + if not augment: + continue + jitter = JitterBundleTransform( + hflip=bool(params.get("jitterHFlip", True)), + vflip=bool(params.get("jitterVFlip", True)), + rotation_deg=int(params.get("jitterRotation", 15) or 0), + color_jitter=_parse_color_jitter(params.get("jitterColor")) + if params.get("jitterColorEnabled", True) + else None, + ) + ops.append(jitter) + elif transform_type == "roi_crop": + roi = ROICropTransform( + mask_source=params.get("roiMaskSource", "gt"), + mask_provider=mask_provider, + scale=float(params.get("roiScale", 2.5)), + target_size=(int(params.get("roiTargetSize", 224)), int(params.get("roiTargetSize", 224))) + if params.get("roiTargetSize") is not None + else None, + fallback_to_original=bool(params.get("roiFallback", True)), + ) + if roi.mask_provider is None and roi.mask_source == "unet": + if strict: + raise ValueError("ROI crop requires a mask provider for 'unet' source.") + ops.append(roi) + else: + if strict: + raise ValueError(f"Unsupported transform type: {transform_type!r}") + + # Always end with tensor + normalize, using backbone defaults + cfg = backbone_transform_config(backbone_name, augment=augment) + ops.extend( + [ + transforms.ToTensor(), + transforms.Normalize(mean=cfg.mean, std=cfg.std), + ] + ) + return transforms.Compose(ops) diff --git a/v3/classes/unet_segmenter.py b/v3/classes/unet_segmenter.py new file mode 100755 index 0000000..407f8c7 --- /dev/null +++ b/v3/classes/unet_segmenter.py @@ -0,0 +1,894 @@ +"""U-Net based optic disc/cup segmenter for REFUGE + Papila.""" + +from __future__ import annotations + +import math +import os +from concurrent.futures import ThreadPoolExecutor, as_completed +from contextlib import suppress +from dataclasses import dataclass +from pathlib import Path +from typing import Iterable, List, Optional, Set, Tuple + +import numpy as np +import pandas as pd +from PIL import Image, ImageDraw, ImageOps +from PIL.Image import Resampling +from skimage import measure +import torch +from torch import nn +from torch.utils.data import DataLoader, Dataset +from torchvision import transforms +from tqdm import tqdm + + +@dataclass +class ManifestEntry: + sample_id: str + dataset: str + image_path: Path + annotation_disc: Path + annotation_cup: Path + annotation_type_disc: str + annotation_type_cup: str + split: str # train / holdout / etc. + + +class UNet(nn.Module): + def __init__( + self, in_channels: int = 3, base_channels: int = 32, out_channels: int = 2 + ): + super().__init__() + self.enc1 = self._block(in_channels, base_channels) + self.enc2 = self._block(base_channels, base_channels * 2) + self.enc3 = self._block(base_channels * 2, base_channels * 4) + self.enc4 = self._block(base_channels * 4, base_channels * 8) + + self.pool = nn.MaxPool2d(2) + self.bottleneck = self._block(base_channels * 8, base_channels * 16) + + self.up4 = nn.ConvTranspose2d( + base_channels * 16, base_channels * 8, 2, stride=2 + ) + self.dec4 = self._block(base_channels * 16, base_channels * 8) + self.up3 = nn.ConvTranspose2d(base_channels * 8, base_channels * 4, 2, stride=2) + self.dec3 = self._block(base_channels * 8, base_channels * 4) + self.up2 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, 2, stride=2) + self.dec2 = self._block(base_channels * 4, base_channels * 2) + self.up1 = nn.ConvTranspose2d(base_channels * 2, base_channels, 2, stride=2) + self.dec1 = self._block(base_channels * 2, base_channels) + + self.out_conv = nn.Conv2d(base_channels, out_channels, kernel_size=1) + + @staticmethod + def _block(in_ch: int, out_ch: int) -> nn.Module: + return nn.Sequential( + nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1, bias=False), + nn.BatchNorm2d(out_ch), + nn.ReLU(inplace=True), + nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1, bias=False), + nn.BatchNorm2d(out_ch), + nn.ReLU(inplace=True), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + e1 = self.enc1(x) + e2 = self.enc2(self.pool(e1)) + e3 = self.enc3(self.pool(e2)) + e4 = self.enc4(self.pool(e3)) + b = self.bottleneck(self.pool(e4)) + + d4 = self.up4(b) + d4 = torch.cat([d4, e4], dim=1) + d4 = self.dec4(d4) + d3 = self.up3(d4) + d3 = torch.cat([d3, e3], dim=1) + d3 = self.dec3(d3) + d2 = self.up2(d3) + d2 = torch.cat([d2, e2], dim=1) + d2 = self.dec2(d2) + d1 = self.up1(d2) + d1 = torch.cat([d1, e1], dim=1) + d1 = self.dec1(d1) + return self.out_conv(d1) + + +class SegmentationDataset(Dataset): + def __init__( + self, + entries: List[ManifestEntry], + segmenter: "UNetSegmenter", + augment: bool, + ) -> None: + self.entries = entries + self.segmenter = segmenter + self.augment = augment + + def __len__(self) -> int: + return len(self.entries) + + def __getitem__(self, idx: int): + entry = self.entries[idx] + image = self.segmenter.load_preprocessed_image(entry) + disc_mask, cup_mask = self.segmenter.load_masks(entry) + + if self.augment: + image = self.segmenter.jitter_image(image) + image, disc_mask, cup_mask = self.segmenter.augment_geometric( + image, disc_mask, cup_mask + ) + image_tensor = transforms.ToTensor()(image) + image_tensor = self.segmenter._normalize_tensor(image_tensor) + + mask = np.stack([disc_mask, cup_mask], axis=0).astype(np.float32) + mask_tensor = torch.from_numpy(mask) + return image_tensor, mask_tensor + + +class UNetSegmenter: + def __init__( + self, + manifest_path: Path, + device: Optional[str] = None, + cup_weight: float = 1.0, + disc_weight: float = 1.0, + target_size: int = 512, + val_ratio: float = 0.1, + train_datasets: Optional[Iterable[str]] = None, + val_datasets: Optional[Iterable[str]] = None, + holdout_datasets: Optional[Iterable[str]] = None, + normalize: str = "none", + use_stronger_aug: bool = False, + mask_cache_dir: Optional[Path] = None, + image_cache_dir: Optional[Path] = None, + in_memory_cache: bool = False, + loader_workers: int = 0, + ) -> None: + self.manifest_path = manifest_path + self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") + self.cup_weight = cup_weight + self.disc_weight = disc_weight + self.target_size = target_size + self.val_ratio = val_ratio + self.normalize = (normalize or "none").lower() + self.use_stronger_aug = bool(use_stronger_aug) + self.mask_cache_dir = Path(mask_cache_dir).resolve() if mask_cache_dir else None + if self.mask_cache_dir: + self.mask_cache_dir.mkdir(parents=True, exist_ok=True) + self.image_cache_dir = Path(image_cache_dir).resolve() if image_cache_dir else None + if self.image_cache_dir: + self.image_cache_dir.mkdir(parents=True, exist_ok=True) + self.in_memory_cache = bool(in_memory_cache) + self._mem_image_cache: dict[str, np.ndarray] = {} + self._mem_mask_cache: dict[str, Tuple[np.ndarray, np.ndarray]] = {} + self.loader_workers = max(0, int(loader_workers)) + + self.train_dataset_filter = self._normalize_filter(train_datasets) + self.val_dataset_filter = self._normalize_filter(val_datasets) + self.holdout_dataset_filter = self._normalize_filter(holdout_datasets) + + self.model = UNet().to(self.device) + self._manifest: List[ManifestEntry] = [] + self.train_entries: List[ManifestEntry] = [] + self.val_entries: List[ManifestEntry] = [] + self.holdout_entries: List[ManifestEntry] = [] + self.read_manifest() + + def prebuild_in_memory_cache( + self, + *, + cache_workers: int = 0, + include_train: bool = True, + include_val: bool = True, + include_holdout: bool = False, + ) -> None: + if not self.in_memory_cache: + return + selected: List[ManifestEntry] = [] + if include_train: + selected.extend(self.train_entries) + if include_val: + selected.extend(self.val_entries) + if include_holdout: + selected.extend(self.holdout_entries) + if not selected: + return + + # Deduplicate by cache key. + dedup = {} + for entry in selected: + dedup[self._entry_cache_key(entry)] = entry + entries = list(dedup.values()) + workers = max(0, int(cache_workers)) + print( + f"[UNetSegmenter] prebuilding in-memory cache for {len(entries)} samples " + f"(cache_workers={workers})", + flush=True, + ) + + def _warm_one(entry: ManifestEntry) -> None: + self.load_preprocessed_image(entry) + self.load_masks(entry) + + if workers <= 1: + for entry in tqdm(entries, desc="Warm cache", unit="sample"): + _warm_one(entry) + else: + with ThreadPoolExecutor(max_workers=workers) as ex: + futures = [ex.submit(_warm_one, entry) for entry in entries] + for fut in tqdm(as_completed(futures), total=len(futures), desc="Warm cache", unit="sample"): + fut.result() + + # ------------------------------------------------------------------ + def read_manifest(self) -> None: + df = pd.read_csv(self.manifest_path) + entries: List[ManifestEntry] = [] + for _, row in df.iterrows(): + entry = ManifestEntry( + sample_id=row["sample_id"], + dataset=row["dataset"], + image_path=Path(row["image_path"]), + annotation_disc=Path(row["annotation_disc"]), + annotation_cup=Path(row["annotation_cup"]), + annotation_type_disc=row["annotation_type_disc"], + annotation_type_cup=row["annotation_type_cup"], + split=row["split"], + ) + entries.append(entry) + self._manifest = entries + self.holdout_entries = [e for e in entries if e.split == "holdout"] + if self.holdout_dataset_filter is not None: + self.holdout_entries = [ + e for e in self.holdout_entries if e.dataset in self.holdout_dataset_filter + ] + + trainable = [e for e in entries if e.split != "holdout"] + if self.train_dataset_filter is not None: + trainable = [ + e for e in trainable if e.dataset in self.train_dataset_filter + ] + + if not trainable: + self.val_entries = [] + self.train_entries = [] + return + + val_pool = trainable + if self.val_dataset_filter is not None: + filtered = [e for e in trainable if e.dataset in self.val_dataset_filter] + if filtered: + val_pool = filtered + + if len(trainable) == 1: + val_count = 0 + else: + val_count = max(1, int(len(trainable) * self.val_ratio)) + val_count = min(val_count, len(val_pool), len(trainable) - 1) + + selected_val: List[ManifestEntry] = [] + if val_count > 0: + selected_val = list(val_pool[:val_count]) + self.val_entries = selected_val + selected_ids = {id(item) for item in selected_val} + self.train_entries = [e for e in trainable if id(e) not in selected_ids] + + if not self.train_entries and trainable: + # Fallback when filtering removed all train entries (e.g. val_count forced entire set) + self.train_entries = trainable + self.val_entries = [] + + # ------------------------------------------------------------------ + def preprocess_image(self, image: Image.Image) -> Image.Image: + return image.resize((self.target_size, self.target_size), Resampling.BILINEAR) + + def jitter_image(self, image: Image.Image) -> Image.Image: + # Photometric jitter only; geometric ops are applied jointly (image+mask) + return transforms.ColorJitter(0.1, 0.1, 0.1, 0.05)(image) + + def augment_geometric( + self, + image: Image.Image, + disc_mask: np.ndarray, + cup_mask: np.ndarray, + ) -> tuple[Image.Image, np.ndarray, np.ndarray]: + if not self.use_stronger_aug: + return image, disc_mask, cup_mask + + img = image + disc_pil = Image.fromarray((disc_mask > 0).astype(np.uint8) * 255) + cup_pil = Image.fromarray((cup_mask > 0).astype(np.uint8) * 255) + + # Random horizontal flip + if np.random.rand() < 0.5: + img = ImageOps.mirror(img) + disc_pil = ImageOps.mirror(disc_pil) + cup_pil = ImageOps.mirror(cup_pil) + # Random vertical flip + if np.random.rand() < 0.5: + img = ImageOps.flip(img) + disc_pil = ImageOps.flip(disc_pil) + cup_pil = ImageOps.flip(cup_pil) + # Random rotation (multiples of 90° to keep masks aligned) + rotations = np.random.choice([0, 90, 180, 270]) + if rotations: + img = img.rotate(rotations, expand=False) + disc_pil = disc_pil.rotate(rotations, expand=False) + cup_pil = cup_pil.rotate(rotations, expand=False) + + disc_mask = (np.array(disc_pil) > 0).astype(np.float32) + cup_mask = (np.array(cup_pil) > 0).astype(np.float32) + return img, disc_mask, cup_mask + + @staticmethod + def _slugify(text: str) -> str: + return "".join(ch if ch.isalnum() or ch in ("-", "_") else "_" for ch in text) + + def _entry_cache_key(self, entry: ManifestEntry) -> str: + return self._slugify(f"{entry.dataset}_{entry.sample_id}_sz{self.target_size}") + + def _mask_cache_path(self, entry: ManifestEntry) -> Optional[Path]: + if self.mask_cache_dir is None: + return None + slug = self._slugify(f"{entry.dataset}_{entry.sample_id}") + fname = f"{slug}_sz{self.target_size}.npz" + return self.mask_cache_dir / fname + + def _image_cache_path(self, entry: ManifestEntry) -> Optional[Path]: + if self.image_cache_dir is None: + return None + slug = self._slugify(f"{entry.dataset}_{entry.sample_id}") + fname = f"{slug}_img_sz{self.target_size}.npz" + return self.image_cache_dir / fname + + def _load_image_cache(self, cache_path: Path) -> Optional[Image.Image]: + try: + data = np.load(str(cache_path), allow_pickle=False) + arr = data["image"].astype(np.uint8, copy=False) + if arr.ndim != 3 or arr.shape[2] != 3: + return None + return Image.fromarray(arr, mode="RGB") + except Exception: + with suppress(OSError, FileNotFoundError): + cache_path.unlink() + return None + + def _save_image_cache(self, cache_path: Optional[Path], image: Image.Image) -> None: + if cache_path is None: + return + cache_path.parent.mkdir(parents=True, exist_ok=True) + tmp_path = cache_path.with_suffix(cache_path.suffix + ".tmp.npz") + try: + arr = np.asarray(image, dtype=np.uint8) + np.savez_compressed(tmp_path, image=arr) + os.replace(tmp_path, cache_path) + except Exception: + with suppress(OSError, FileNotFoundError): + tmp_path.unlink() + + def load_preprocessed_image(self, entry: ManifestEntry) -> Image.Image: + key = self._entry_cache_key(entry) + if self.in_memory_cache: + cached = self._mem_image_cache.get(key) + if cached is not None: + return Image.fromarray(cached, mode="RGB") + cache_path = self._image_cache_path(entry) + if cache_path and cache_path.exists(): + cached = self._load_image_cache(cache_path) + if cached is not None: + if self.in_memory_cache: + self._mem_image_cache[key] = np.asarray(cached, dtype=np.uint8) + return cached + image = Image.open(entry.image_path).convert("RGB") + image = self.preprocess_image(image) + if self.in_memory_cache: + self._mem_image_cache[key] = np.asarray(image, dtype=np.uint8) + self._save_image_cache(cache_path, image) + return image + + def _load_mask_cache(self, cache_path: Path) -> Optional[Tuple[np.ndarray, np.ndarray]]: + try: + data = np.load(str(cache_path), allow_pickle=False) + disc = data["disc"].astype(np.float32) + cup = data["cup"].astype(np.float32) + return disc, cup + except Exception: + with suppress(OSError, FileNotFoundError): + cache_path.unlink() + return None + + def _save_mask_cache( + self, + cache_path: Optional[Path], + disc_mask: np.ndarray, + cup_mask: np.ndarray, + ) -> None: + if cache_path is None: + return + cache_path.parent.mkdir(parents=True, exist_ok=True) + tmp_path = cache_path.with_suffix(cache_path.suffix + ".tmp.npz") + try: + np.savez_compressed( + tmp_path, + disc=disc_mask.astype(np.uint8), + cup=cup_mask.astype(np.uint8), + ) + os.replace(tmp_path, cache_path) + except Exception: + with suppress(OSError, FileNotFoundError): + tmp_path.unlink() + + def _normalize_tensor(self, tensor: torch.Tensor) -> torch.Tensor: + if self.normalize == "per_image": + mean = tensor.mean(dim=(1, 2), keepdim=True) + std = tensor.std(dim=(1, 2), keepdim=True).clamp(min=1e-6) + return (tensor - mean) / std + if self.normalize == "imagenet": + mean = torch.tensor([0.485, 0.456, 0.406]).view(-1, 1, 1) + std = torch.tensor([0.229, 0.224, 0.225]).view(-1, 1, 1) + return (tensor - mean) / std + return tensor + + # ------------------------------------------------------------------ + def extract_masks_from_image( + self, + mask_path: Path, + disc_color: Optional[tuple[int, int, int]] = None, + cup_color: Optional[tuple[int, int, int]] = None, + ) -> Tuple[np.ndarray, Optional[np.ndarray], Tuple[int, int]]: + raw = Image.open(mask_path) + arr = np.array(raw) + if arr.ndim == 2: + h, w = arr.shape + flat = arr.reshape(-1).astype(np.int64, copy=False) + edges = np.concatenate([arr[0, :], arr[-1, :], arr[:, 0], arr[:, -1]], axis=0).astype(np.int64, copy=False) + edge_counts = np.bincount(edges, minlength=256) + bg_val = int(np.argmax(edge_counts)) + counts = np.bincount(flat, minlength=256) + counts[bg_val] = 0 + vals = np.where(counts > 0)[0] + if vals.size < 1: + raise ValueError(f"Mask {mask_path} does not contain discernible labels") + # Disc = ALL non-background pixels (full optic disc: rim + cup combined). + # Previously this was rim-only, which caused the cup structural prior + # (cup & disc) to produce empty cup masks since cup and rim don't overlap. + disc_mask = (arr != bg_val).astype(np.uint8) + # Cup = the darkest non-background value (0 in REFUGE = inner cup region). + # Using min-value rather than frequency avoids swapping when cup area > rim area. + cup_val = int(np.min(vals)) if vals.size > 1 else None + cup_mask = (arr == cup_val).astype(np.uint8) if cup_val is not None else np.zeros_like(disc_mask, dtype=np.uint8) + return disc_mask, cup_mask if cup_mask.any() else None, (w, h) + + image = raw.convert("RGB") + arr = np.array(image) + h, w, c = arr.shape + + if disc_color is None or cup_color is None: + # Fast color discovery via NumPy (avoid Python-level per-pixel tuple counting). + edges = np.concatenate( + [arr[0, :, :], arr[-1, :, :], arr[:, 0, :], arr[:, -1, :]], axis=0 + ) + edge_colors, edge_counts = np.unique(edges.reshape(-1, c), axis=0, return_counts=True) + bg_color_np = edge_colors[int(np.argmax(edge_counts))] + + colors_np, counts_np = np.unique(arr.reshape(-1, c), axis=0, return_counts=True) + keep = np.any(colors_np != bg_color_np.reshape(1, -1), axis=1) + colors_np = colors_np[keep] + counts_np = counts_np[keep] + if colors_np.shape[0] < 1: + raise ValueError(f"Mask {mask_path} does not contain discernible labels") + order = np.argsort(-counts_np) + colors_np = colors_np[order] + disc_color = tuple(int(v) for v in colors_np[0].tolist()) + cup_color = ( + tuple(int(v) for v in colors_np[1].tolist()) + if colors_np.shape[0] > 1 + else None + ) + + disc_mask = np.zeros((h, w), dtype=np.uint8) + cup_mask = np.zeros((h, w), dtype=np.uint8) + + if disc_color is not None: + disc_mask[np.all(arr == disc_color, axis=-1)] = 1 + if cup_color is not None: + cup_mask[np.all(arr == cup_color, axis=-1)] = 1 + + return disc_mask, cup_mask if cup_mask.any() else None, (w, h) + + def load_contour_from_file(self, contour_path: Path) -> np.ndarray: + # Fast path: contour files are typically CSV or whitespace-delimited x,y pairs. + try: + arr = np.loadtxt(str(contour_path), delimiter=",", comments="#", dtype=np.float32) + except Exception: + try: + arr = np.loadtxt(str(contour_path), comments="#", dtype=np.float32) + except Exception: + return np.zeros((0, 2), dtype=np.float32) + if arr.size == 0: + return np.zeros((0, 2), dtype=np.float32) + if arr.ndim == 1: + if arr.shape[0] < 2: + return np.zeros((0, 2), dtype=np.float32) + arr = arr.reshape(1, -1) + if arr.shape[1] < 2: + return np.zeros((0, 2), dtype=np.float32) + return arr[:, :2].astype(np.float32, copy=False) + + def coords_to_mask( + self, + coords: Optional[np.ndarray], + size: Tuple[int, int], + ) -> np.ndarray: + if coords is None or len(coords) == 0: + return np.zeros((self.target_size, self.target_size), dtype=np.float32) + + width, height = map(int, size) + target_shape = (height, width) + arr = np.asarray(coords) + if arr.size == 0: + return np.zeros((self.target_size, self.target_size), dtype=np.float32) + + if arr.ndim == 2 and arr.shape[-1] != 2: + mask = (arr > 0).astype(np.uint8) + return self._resize_mask(mask) + + if arr.ndim > 2: + arr = arr.reshape(-1, arr.shape[-1]) + arr = arr.astype(float, copy=False) + if arr.shape[-1] != 2: + raise ValueError(f"Expected coordinate pairs, got shape {arr.shape}") + + points = [tuple(map(float, pt)) for pt in arr] + if len(points) < 3: + return np.zeros(target_shape, dtype=np.float32) + + img = Image.new("L", size, 0) + draw = ImageDraw.Draw(img) + draw.polygon(points, outline=1, fill=1) + mask = np.array(img, dtype=np.uint8) + return self._resize_mask(mask) + + def _resize_mask(self, mask: np.ndarray) -> np.ndarray: + img = Image.fromarray((mask > 0).astype(np.uint8) * 255) + img = img.resize((self.target_size, self.target_size), Resampling.NEAREST) + return (np.array(img, dtype=np.uint8) > 0).astype(np.float32) + + def load_masks(self, entry: ManifestEntry) -> Tuple[np.ndarray, np.ndarray]: + key = self._entry_cache_key(entry) + if self.in_memory_cache: + cached = self._mem_mask_cache.get(key) + if cached is not None: + disc_u8, cup_u8 = cached + return disc_u8.astype(np.float32), cup_u8.astype(np.float32) + cache_path = self._mask_cache_path(entry) + if cache_path and cache_path.exists(): + cached = self._load_mask_cache(cache_path) + if cached is not None: + if self.in_memory_cache: + disc, cup = cached + self._mem_mask_cache[key] = ( + disc.astype(np.uint8), + cup.astype(np.uint8), + ) + return cached + + image = Image.open(entry.image_path) + size = image.size + + disc_coords = cup_coords = None + if entry.annotation_type_disc == "mask": + disc_coords, cup_coords_from_disc, size = self.extract_masks_from_image( + entry.annotation_disc + ) + if cup_coords_from_disc is not None: + cup_coords = cup_coords_from_disc + else: + disc_coords = self.load_contour_from_file(entry.annotation_disc) + + if entry.annotation_type_cup == "mask": + _, cup_coords_from_cup, size_cup = self.extract_masks_from_image( + entry.annotation_cup + ) + if cup_coords_from_cup is not None: + cup_coords = cup_coords_from_cup + if disc_coords is None: + disc_coords, _, size = self.extract_masks_from_image( + entry.annotation_cup + ) + else: + size = size_cup + else: + cup_coords = self.load_contour_from_file(entry.annotation_cup) + + disc_mask = self.coords_to_mask(disc_coords, size).astype(np.float32) + cup_mask = self.coords_to_mask(cup_coords, size).astype(np.float32) + if self.in_memory_cache: + self._mem_mask_cache[key] = ( + disc_mask.astype(np.uint8), + cup_mask.astype(np.uint8), + ) + self._save_mask_cache(cache_path, disc_mask, cup_mask) + return disc_mask, cup_mask + + # ------------------------------------------------------------------ + def build_loaders(self, batch_size: int = 4, num_workers: int = 0) -> Tuple[DataLoader, DataLoader]: + train_ds = SegmentationDataset(self.train_entries, self, augment=True) + val_ds = SegmentationDataset(self.val_entries, self, augment=False) + train_loader = DataLoader( + train_ds, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True + ) + val_loader = DataLoader( + val_ds, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True + ) + return train_loader, val_loader + + # ------------------------------------------------------------------ + def dice_score(self, preds: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: + preds = (preds > 0.5).float() + intersection = (preds * targets).sum(dim=(2, 3)) + union = preds.sum(dim=(2, 3)) + targets.sum(dim=(2, 3)) + dice = (2 * intersection + 1e-6) / (union + 1e-6) + return dice.mean(dim=0) + + def train( + self, + epochs: int = 40, + batch_size: int = 4, + lr: float = 1e-3, + weight_decay: float = 1e-5, + checkpoint_dir: Path = Path("models/unet_segmenter"), + ) -> None: + print( + f"[UNetSegmenter] training on device={self.device} " + f"(epochs={epochs}, batch_size={batch_size}, workers={self.loader_workers})" + ) + train_loader, val_loader = self.build_loaders(batch_size=batch_size, num_workers=self.loader_workers) + optimizer = torch.optim.Adam( + self.model.parameters(), lr=lr, weight_decay=weight_decay + ) + criterion = nn.BCEWithLogitsLoss() + best_dice = -math.inf + checkpoint_dir.mkdir(parents=True, exist_ok=True) + best_path = checkpoint_dir / "best.pt" + + epoch_bar = tqdm(range(1, epochs + 1), desc="Epochs", unit="epoch") + + for epoch in epoch_bar: + self.model.train() + batch_bar = tqdm( + train_loader, + desc=f"Train {epoch}/{epochs}", + leave=False, + unit="batch", + total=len(train_loader), + ) + train_loss_total = 0.0 + train_samples = 0 + for images, masks in batch_bar: + images = images.to(self.device) + masks = masks.to(self.device) + optimizer.zero_grad() + logits = self.model(images) + loss_disc = criterion(logits[:, 0:1], masks[:, 0:1]) + loss_cup = criterion(logits[:, 1:2], masks[:, 1:2]) + loss = self.disc_weight * loss_disc + self.cup_weight * loss_cup + loss.backward() + optimizer.step() + batch_size = images.size(0) + train_loss_total += loss.item() * batch_size + train_samples += batch_size + + train_loss = ( + train_loss_total / train_samples if train_samples else float("nan") + ) + + self.model.eval() + dices = [] + val_bar = tqdm( + val_loader, + desc="Validate", + leave=False, + unit="batch", + total=len(val_loader), + ) + with torch.no_grad(): + for images, masks in val_bar: + images = images.to(self.device) + masks = masks.to(self.device) + logits = self.model(images) + probs = torch.sigmoid(logits) + dice = self.dice_score(probs, masks) + dices.append(dice.cpu()) + if dices: + mean_dice = torch.stack(dices).mean(dim=0) + disc_dice = mean_dice[0].item() + cup_dice = mean_dice[1].item() + weight_sum = self.disc_weight + self.cup_weight + score = ( + (self.disc_weight * disc_dice + self.cup_weight * cup_dice) + / weight_sum + if weight_sum + else 0.0 + ) + epoch_bar.set_postfix( + loss=f"{train_loss:.4f}", + dice_disc=f"{disc_dice:.3f}", + dice_cup=f"{cup_dice:.3f}", + dice_w=f"{score:.3f}", + ) + else: + disc_dice = cup_dice = 0.0 + score = 0.0 + epoch_bar.set_postfix(loss=f"{train_loss:.4f}") + + if score > best_dice: + best_dice = score + torch.save({"model": self.model.state_dict()}, best_path) + + if best_path.exists(): + state = torch.load(best_path, map_location=self.device) + self.model.load_state_dict(state["model"]) + + # ------------------------------------------------------------------ + def evaluate_holdout( + self, output_dir: Path = Path("analysis_data/segmenter_eval") + ) -> pd.DataFrame: + return self.evaluate_dataset(split_filter={"holdout"}, output_dir=output_dir) + + @staticmethod + def overlay_masks( + image: Image.Image, disc: np.ndarray, cup: np.ndarray + ) -> Image.Image: + overlay = image.copy() + disc_img = Image.fromarray((disc * 255).astype(np.uint8)) + cup_img = Image.fromarray((cup * 255).astype(np.uint8)) + disc_color = Image.new("RGBA", image.size, (255, 0, 0, 0)) + cup_color = Image.new("RGBA", image.size, (0, 255, 0, 0)) + disc_color.paste((255, 0, 0, 100), mask=disc_img) + cup_color.paste((0, 255, 0, 100), mask=cup_img) + overlay = overlay.convert("RGBA") + overlay = Image.alpha_composite(overlay, disc_color) + overlay = Image.alpha_composite(overlay, cup_color) + return overlay.convert("RGB") + + # ------------------------------------------------------------------ + @staticmethod + def _normalize_filter(values: Optional[Iterable[str]]) -> Optional[Set[str]]: + if values is None: + return None + if isinstance(values, str): + return {values} + return {str(item) for item in values} + + @staticmethod + def _dice_from_masks(pred: np.ndarray, target: np.ndarray) -> float: + pred = (pred > 0).astype(np.float32) + target = (target > 0).astype(np.float32) + intersection = float((pred * target).sum()) + denom = float(pred.sum() + target.sum()) + return (2.0 * intersection + 1e-6) / (denom + 1e-6) + + def get_entries( + self, + dataset_filter: Optional[Iterable[str]] = None, + split_filter: Optional[Iterable[str]] = None, + ) -> List[ManifestEntry]: + dataset_set = self._normalize_filter(dataset_filter) + split_set = self._normalize_filter(split_filter) + entries = self._manifest + if dataset_set is not None: + entries = [e for e in entries if e.dataset in dataset_set] + if split_set is not None: + entries = [e for e in entries if e.split in split_set] + return list(entries) + + def evaluate_dataset( + self, + dataset_filter: Optional[Iterable[str]] = None, + split_filter: Optional[Iterable[str]] = None, + output_dir: Path = Path("analysis_data/segmenter_eval"), + save_overlays: bool = True, + metrics_path: Optional[Path] = None, + threshold: float = 0.5, + tta: bool = False, + ) -> pd.DataFrame: + entries = self.get_entries( + dataset_filter=dataset_filter, split_filter=split_filter + ) + if not entries: + return pd.DataFrame( + columns=[ + "sample_id", + "dataset", + "split", + "dice_disc", + "dice_cup", + ] + ) + + output_dir.mkdir(parents=True, exist_ok=True) + if metrics_path is None: + suffix_parts = [] + if dataset_filter is not None: + suffix_parts.append("-".join(sorted(self._normalize_filter(dataset_filter)))) + if split_filter is not None: + suffix_parts.append("-".join(sorted(self._normalize_filter(split_filter)))) + suffix = "_".join(part for part in suffix_parts if part) + csv_name = f"metrics{'_' + suffix if suffix else ''}.csv" + metrics_path = output_dir / csv_name + + records = [] + self.model.eval() + progress = tqdm( + entries, + desc="Evaluate", + unit="sample", + leave=False, + ) + for entry in progress: + orig_image = Image.open(entry.image_path).convert("RGB") + image = self.preprocess_image(orig_image) + tensor = transforms.ToTensor()(image) + tensor = self._normalize_tensor(tensor) + tensor = tensor.unsqueeze(0).to(self.device) + with torch.no_grad(): + logits = self.model(tensor) + if tta: + t_h = torch.flip(tensor, dims=[3]) + log_h = self.model(t_h) + log_h = torch.flip(log_h, dims=[3]) + t_v = torch.flip(tensor, dims=[2]) + log_v = self.model(t_v) + log_v = torch.flip(log_v, dims=[2]) + logits = (logits + log_h + log_v) / 3.0 + probs = torch.sigmoid(logits)[0].cpu().numpy() + + disc_pred = (probs[0] > threshold).astype(np.uint8) + cup_pred = (probs[1] > threshold).astype(np.uint8) + # Structural prior: cup within disc + cup_pred = (cup_pred > 0) & (disc_pred > 0) + cup_pred = cup_pred.astype(np.uint8) + + disc_gt, cup_gt = self.load_masks(entry) + disc_gt = disc_gt.astype(np.uint8) + cup_gt = cup_gt.astype(np.uint8) + + dice_disc = self._dice_from_masks(disc_pred, disc_gt) + dice_cup = self._dice_from_masks(cup_pred, cup_gt) + + records.append( + { + "sample_id": entry.sample_id, + "dataset": entry.dataset, + "split": entry.split, + "dice_disc": dice_disc, + "dice_cup": dice_cup, + } + ) + + progress.set_postfix( + dice_disc=f"{dice_disc:.3f}", dice_cup=f"{dice_cup:.3f}" + ) + + if save_overlays: + overlay_gt = self.overlay_masks(image, disc_gt, cup_gt) + overlay_pred = self.overlay_masks(image, disc_pred, cup_pred) + combined = Image.new("RGB", (image.width * 2, image.height)) + combined.paste(overlay_gt, (0, 0)) + combined.paste(overlay_pred, (image.width, 0)) + combined.save(output_dir / f"{entry.sample_id}_eval.png") + + metrics_df = pd.DataFrame(records) + summary = metrics_df[["dice_disc", "dice_cup"]].mean() + summary_row = { + "sample_id": "__mean__", + "dataset": "summary", + "split": "summary", + "dice_disc": summary["dice_disc"], + "dice_cup": summary["dice_cup"], + } + metrics_with_summary = pd.concat( + [metrics_df, pd.DataFrame([summary_row])], ignore_index=True + ) + metrics_with_summary.to_csv(metrics_path, index=False) + return metrics_with_summary diff --git a/v3/classes/utils.py b/v3/classes/utils.py new file mode 100644 index 0000000..142be8b --- /dev/null +++ b/v3/classes/utils.py @@ -0,0 +1,47 @@ +"""General-purpose utilities for the V2 hypertower pipeline.""" +from __future__ import annotations + +import random as pyrandom +from pathlib import Path + +import numpy as np +import pandas as pd +import torch + + +def seed_everything(seed: int) -> None: + pyrandom.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def choose_device(device_arg: str | None) -> torch.device: + if device_arg and device_arg != "auto": + return torch.device(device_arg) + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + + +def _drop_mixed_label_patients(df: pd.DataFrame, *, patient_col: str, label_col: str): + """Remove patients whose rows carry conflicting labels. Returns (clean_df, mixed_pids).""" + per_patient = ( + df.groupby(patient_col)[label_col] + .agg(lambda s: set(pd.to_numeric(s, errors="coerce").dropna().astype(int).tolist())) + ) + mixed = [pid for pid, labels in per_patient.items() if len(labels) > 1] + if not mixed: + return df, [] + return df[~df[patient_col].isin(mixed)].reset_index(drop=True), mixed + + +def _relabel_mixed_patients_to_max(df: pd.DataFrame, *, patient_col: str, label_col: str): + """Set all rows for each patient to that patient's max observed label.""" + out = df.copy() + labels = pd.to_numeric(out[label_col], errors="coerce") + patient_max = labels.groupby(out[patient_col]).transform("max") + changed_rows = int((labels != patient_max).fillna(False).sum()) + out[label_col] = patient_max.astype(int) + per_patient_unique = out.groupby(patient_col)[label_col].nunique(dropna=True) + still_mixed = per_patient_unique[per_patient_unique > 1].index.tolist() + return out.reset_index(drop=True), changed_rows, still_mixed diff --git a/v3/classes/v2_hypertower.py b/v3/classes/v2_hypertower.py new file mode 100644 index 0000000..8733b11 --- /dev/null +++ b/v3/classes/v2_hypertower.py @@ -0,0 +1,1853 @@ +"""V2HyperTower — central orchestrator for the V2 mode-comparison pipeline. + +All model classes, metric helpers, croppers and data utilities live in their +respective category-specific modules. This file owns only: + - V2HyperTower (the big-picture orchestrator) + - V2ModeComparator (thin backward-compat shim) + - module-level ``build_parser`` alias (used by old scripts that import it + directly; will be removed once all callers are updated) +""" +from __future__ import annotations + +import argparse +import copy +import csv +import json +import time +from pathlib import Path +from types import SimpleNamespace +from typing import Optional + +import numpy as np +import torch + +from v3.classes.croppers import build_image_preprocessor_from_args +from v3.classes.dataset import _ClinicalView # noqa: F401 (re-exported for compat) +from v3.classes.loader_factory import ( + build_balanced_sampler, + filter_bilateral_samples, + filter_eye_samples, + make_loader, +) +from v3.classes.metrics import _score_arrays, _svf, _tune_and_snap +from v3.classes.models import ( + BilateralHT, + FusedEnsembleHT, + SingleEyeHT, + V2ModeComparisonOps, + collect_probs_bilateral, + collect_probs_bilateral_components, + collect_probs_classic, + collect_probs_ensemble, + collect_probs_ensemble_pereye, + collect_probs_eye_level, + collect_probs_fused, + collect_probs_single_components, + train_bilateral_epoch, + train_fusion_epoch, + train_single_epoch, +) +from v3.classes.papila_builders import build_papila_data +from v3.classes.predictions import PredictionStore, head_names_for_mode +from v3.classes.profiles import build_papila_profile +from v3.classes.results import FoldArtifacts, FoldResult, _f, _nan, _sv +from v3.classes.split_manager import PatientFirstSplitManager +from v3.classes.transforms import build_eval_transform +from v3.classes.utils import ( + _drop_mixed_label_patients, + _relabel_mixed_patients_to_max, + choose_device, + seed_everything, +) +from v3.classes.hypertower_logger import HypertowerLogger + + +# --------------------------------------------------------------------------- +# Fusion-event helper +# --------------------------------------------------------------------------- + +def _fusion_events( + y: np.ndarray, + pf: np.ndarray, + pi: np.ndarray, + pm: np.ndarray, +) -> tuple[int, int]: + """ + Count fusion corrections and errors. + - correction: fused correct, both img and md wrong + - error: fused wrong, both img and md correct + Returns (n_corrections, n_errors). + """ + pred_f = pf.argmax(1); pred_i = pi.argmax(1); pred_m = pm.argmax(1) + corr = int(((pred_f == y) & (pred_i != y) & (pred_m != y)).sum()) + err = int(((pred_f != y) & (pred_i == y) & (pred_m == y)).sum()) + return corr, err + + +def _cm_cells(y: np.ndarray, p: np.ndarray, num_classes: int) -> dict[str, int]: + """ + Return confusion matrix cells as a flat dict. + Binary: keys tn/fp/fn/tp + Multiclass: keys cm_{i}_{j} for true class i, predicted class j + Returns empty dict if arrays are empty or wrong shape. + """ + if not y.size or p.ndim < 2 or p.shape[1] != num_classes: + return {} + pred = p.argmax(1) + if num_classes == 2: + tn = int(((pred == 0) & (y == 0)).sum()) + fp = int(((pred == 1) & (y == 0)).sum()) + fn = int(((pred == 0) & (y == 1)).sum()) + tp = int(((pred == 1) & (y == 1)).sum()) + return {"tn": tn, "fp": fp, "fn": fn, "tp": tp} + # multiclass: full NxN matrix + out: dict[str, int] = {} + for i in range(num_classes): + for j in range(num_classes): + out[f"cm_{i}_{j}"] = int(((y == i) & (pred == j)).sum()) + return out + + +# --------------------------------------------------------------------------- +# Per-sample prediction logging +# --------------------------------------------------------------------------- + +def _save_predictions_csv( + fold_dir: Path, + eval_mode: str, + y_true: np.ndarray, + heads: dict, # {"fused": probs_array, "img": probs_array, "md": probs_array, ...} + suffix: str = "", # e.g. "_pereye" +) -> None: + """ + Save a per-sample CSV with predicted class, per-class probabilities, + and TP/FP/TN/FN (binary) or correct flag (multiclass) for every head. + """ + N = len(y_true) + num_classes = next(p.shape[1] for p in heads.values() if p is not None) + rows = [] + for i in range(N): + true = int(y_true[i]) + row: dict = {"idx": i, "y_true": true} + for head_name, probs in heads.items(): + if probs is None: + continue + pred = int(probs[i].argmax()) + row[f"pred_{head_name}"] = pred + for c in range(num_classes): + row[f"prob_{head_name}_c{c}"] = float(probs[i, c]) + if eval_mode == "binary": + row[f"tp_{head_name}"] = int(pred == 1 and true == 1) + row[f"fp_{head_name}"] = int(pred == 1 and true == 0) + row[f"tn_{head_name}"] = int(pred == 0 and true == 0) + row[f"fn_{head_name}"] = int(pred == 0 and true == 1) + else: + row[f"correct_{head_name}"] = int(pred == true) + rows.append(row) + + if not rows: + return + csv_path = fold_dir / f"predictions{suffix}.csv" + with csv_path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + +# --------------------------------------------------------------------------- +# V2HyperTower +# --------------------------------------------------------------------------- + +class V2HyperTower: + """Central orchestrator. Construct with ``V2HyperTower(args)``, call ``.run()``.""" + + # ------------------------------------------------------------------ + # CLI + # ------------------------------------------------------------------ + + @staticmethod + def build_parser() -> argparse.ArgumentParser: + ap = argparse.ArgumentParser( + description=( + "Three HyperTower modes: Classic (eye-level), Ensemble (patient-level avg), " + "Bilateral (BilateralBridge with shared towers). Pure k-fold CV." + ) + ) + ap.add_argument("--image-dir", default="Papila/FundusImages") + ap.add_argument("--clinical-dir", default="Papila/ClinicalData") + ap.add_argument("--label-col", default="Diagnosis") + ap.add_argument("--cat-cols", nargs="*", default=["Gender", "Phakic/Pseudophakic"]) + ap.add_argument("--exclude-cols", nargs="*", default=[], + help="Feature columns to exclude entirely from the clinical feature matrix.") + ap.add_argument("--eval-mode", choices=["binary", "multiclass"], default="multiclass") + ap.add_argument( + "--tower-mode", choices=["single", "ensemble", "bilateral", "classic"], + default="single", + help="Train/evaluate a single tower mode.", + ) + ap.add_argument("--n-splits", type=int, default=5) + ap.add_argument("--fold-seed", type=int, default=42) + ap.add_argument("--holdout-per-class", type=int, default=5, + help="Patients per class reserved for holdout before train/test split (0 disables)") + ap.add_argument("--holdout-seed", type=int, default=123, + help="Random seed for holdout sampling") + ap.add_argument( + "--folds", type=int, default=None, + help="Optional cap on how many folds to run (default: all --n-splits).", + ) + ap.add_argument("--epochs", type=int, default=40) + ap.add_argument( + "--warmup-tower-epochs", type=int, default=None, + help="Extra tower warmup epochs (added before main epochs). Default: auto by mode.", + ) + ap.add_argument( + "--warmup-fused-epochs", type=int, default=None, + help="Extra fused warmup epochs (added before main epochs). Default: auto by mode.", + ) + ap.add_argument("--single-warmup-tower-epochs", type=int, default=None, + help="Single-eye model tower warmup (overrides --warmup-tower-epochs).") + ap.add_argument("--single-warmup-fused-epochs", type=int, default=None, + help="Single-eye model fused warmup (overrides --warmup-fused-epochs).") + ap.add_argument("--warmup-cd-epochs", type=int, default=0, + help="MD-only warmup epochs before tower warmup. Trains only cd_tower + " + "classifier_cd (no CNN forward pass, so 50-100 epochs is cheap).") + ap.add_argument("--bilat-warmup-tower-epochs", type=int, default=None, + help="Bilateral model tower warmup (overrides --warmup-tower-epochs).") + ap.add_argument("--bilat-warmup-fused-epochs", type=int, default=None, + help="Bilateral model fused warmup (overrides --warmup-fused-epochs).") + ap.add_argument("--batch-size", type=int, default=8) + ap.add_argument("--lr", type=float, default=1e-4) + ap.add_argument("--bcd-prob", type=float, default=0.5, + help="Tower-only step probability during main phase (per model).") + ap.add_argument("--tower-loss-mode", choices=["bcd", "all"], default="bcd", + help="Main-phase tower loss strategy: " + "'bcd' (Block Coordinate Descent — randomly train one tower or fused per step) " + "or 'all' (sum all three losses — fused + img + md — every step).") + ap.add_argument("--backbone", default="refugelike") + ap.add_argument("--freeze-ratio", type=float, default=0.0) + ap.add_argument("--augment", action="store_true") + ap.add_argument("--balanced-sampling", action="store_true", + help="Use WeightedRandomSampler during training to equalise class frequency (default: off).") + ap.add_argument("--num-workers", type=int, default=4) + ap.add_argument("--in-memory-cache", action="store_true", default=True, + help="Cache preprocessed images in RAM (default: on).") + ap.add_argument("--no-in-memory-cache", action="store_false", dest="in_memory_cache", + help="Disable in-memory image cache.") + ap.add_argument("--cache-workers", type=int, default=4, + help="Threads for prebuilding in-memory image cache (default: 4).") + ap.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto") + ap.add_argument("--seed", type=int, default=1234) + ap.add_argument("--run-name", default=None) + ap.add_argument("--output-root", default="analysis_data") + # Optional ROI cropping + ap.add_argument("--img-crop-manifest", type=str, default=None, + help="Path to crop manifest CSV for ROI cropping.") + ap.add_argument("--img-crop-gt", action="store_true", + help="Use ground-truth masks/contours from manifest for ROI crop.") + ap.add_argument("--img-crop-weights", type=str, default=None, + help="UNet weights path for ROI cropping from predicted masks.") + ap.add_argument("--img-crop-normalize", type=str, default="per_image", + choices=["per_image", "imagenet"], + help="UNet input normalization mode.") + ap.add_argument("--img-crop-threshold", type=float, default=0.5, + help="UNet mask threshold for ROI extraction.") + ap.add_argument("--img-crop-tta", action="store_true", + help="Enable flip-TTA during UNet mask inference.") + ap.add_argument("--img-crop-scale", type=float, default=2.5, + help="Disc-radius multiplier for square crop.") + ap.add_argument("--img-crop-size", type=int, default=224, + help="Output ROI size before tower transforms.") + ap.add_argument("--img-crop-cache", type=str, default="cache_data/hypertower_crops", + help="Cache directory for cropped images and geometry sidecars.") + ap.add_argument("--persist-img-crop-cache", action="store_true", + help="Keep existing cached crop .npz files instead of clearing at run start.") + # Architecture + ap.add_argument("--cd-hidden-dim", type=int, default=128, + help="ClinicalTower hidden dimension.") + ap.add_argument("--fusion-dim", type=int, default=256, + help="Bridge/BilateralBridge fusion dimension.") + ap.add_argument("--bridge-mode", default="fused", + choices=["fused", "image_only", "clinical_only"], + help="Bridge fusion mode: fused (default), image_only, or clinical_only.") + # Mixed patients + ap.add_argument( + "--exclude-mixed-patients", + dest="exclude_mixed_patients", action="store_true", + help="Drop patients whose two eyes have different labels before splitting.", + ) + ap.add_argument( + "--include-mixed-patients", + dest="exclude_mixed_patients", action="store_false", + ) + ap.add_argument( + "--relabel-mixed-patients-to-max", + dest="relabel_mixed_patients_to_max", + action="store_true", + help="When mixed patients are included, relabel both eyes to patient max severity.", + ) + ap.add_argument( + "--keep-mixed-raw-labels", + dest="relabel_mixed_patients_to_max", + action="store_false", + help=argparse.SUPPRESS, + ) + ap.set_defaults(exclude_mixed_patients=False, relabel_mixed_patients_to_max=False) + # Tuning + ap.add_argument( + "--tune-binary-threshold", action="store_true", + help="Tune per-model binary threshold on validation each epoch.", + ) + ap.add_argument( + "--tune-multiclass-bias", action="store_true", + help="Tune per-model multiclass log-prob bias on validation each epoch.", + ) + ap.add_argument("--ece-bins", type=int, default=10) + ap.add_argument("--log-every", type=int, default=1) + ap.add_argument("--save-checkpoints", action=argparse.BooleanOptionalAction, default=True, + help="Save best_single.pt / best_holdout_single.pt per fold (use --no-save-checkpoints to disable)") + ap.add_argument("--use-last-epoch", action="store_true", default=False, + help="Score using the final epoch's model state rather than the best-AUC checkpoint.") + # IOP feature options + ap.add_argument( + "--iop-corr-method", + choices=["ratio", "ols", "lad", "multi"], + default="ratio", + help="Perkins→Pneumatic conversion method: ratio (default), ols, lad, or multi (+CCT).", + ) + ap.add_argument( + "--iop-drop-raw", + action="store_true", + default=False, + help="Exclude IOP_raw from the feature matrix (keep only IOP_corr).", + ) + ap.add_argument( + "--fused-head", action="store_true", + help="(ensemble mode only) After base SingleEyeHT training, freeze it and train a " + "small logit-level MLP fusion head on bilateral samples instead of averaging " + "OD/OS softmax probabilities.", + ) + ap.add_argument( + "--fusion-epochs", type=int, default=10, + help="Number of epochs to train the fusion head (--fused-head, ensemble mode only).", + ) + return ap + + # ------------------------------------------------------------------ + # Construction + # ------------------------------------------------------------------ + + def __init__(self, args) -> None: + self.args = args + self.device = choose_device(args.device) + seed_everything(args.seed) + + print(f"Device: {self.device}", flush=True) + print("Loading PAPILA data...", flush=True) + self.data = build_papila_data( + image_dir=args.image_dir, + clinical_dir=args.clinical_dir, + label_col=args.label_col, + cat_cols=list(args.cat_cols), + n_splits=args.n_splits, + random_seed=args.fold_seed, + iop_corr_method=getattr(args, "iop_corr_method", "ratio"), + iop_drop_raw=getattr(args, "iop_drop_raw", False), + exclude_cols=list(getattr(args, "exclude_cols", []) or []), + ) + print(f"Loaded: {len(self.data.df)} rows feature_dim={self.data.feature_dim}", flush=True) + self.image_preprocessor = build_image_preprocessor_from_args(args) + self.profile_eye = build_papila_profile( + patient_col="Patient ID", label_col=args.label_col, sample_mode="eye" + ) + self.profile_patient = build_papila_profile( + patient_col="Patient ID", label_col=args.label_col, sample_mode="patient" + ) + + # ------------------------------------------------------------------ + # Orchestration + # ------------------------------------------------------------------ + + def run(self) -> Path: + """Execute the full fold loop for one eval_mode × tower_mode combination.""" + args = self.args + ts = time.strftime("%Y%m%d_%H%M%S") + run_name = args.run_name or f"hypertower_modes_{ts}" + out_dir = Path(args.output_root) / run_name + out_dir.mkdir(parents=True, exist_ok=True) + + mode = args.eval_mode + tower_mode = "single" if args.tower_mode == "classic" else args.tower_mode + df_mode = self.data.df.copy() + + if args.exclude_mixed_patients: + before = df_mode["Patient ID"].nunique() + df_mode, mixed = _drop_mixed_label_patients( + df_mode, patient_col="Patient ID", label_col=args.label_col + ) + print( + f"[{mode}] dropped {len(mixed)} mixed-label patients " + f"({before} → {df_mode['Patient ID'].nunique()})", + flush=True, + ) + else: + if args.relabel_mixed_patients_to_max: + before_rows = len(df_mode) + df_mode, changed_rows, still_mixed = _relabel_mixed_patients_to_max( + df_mode, patient_col="Patient ID", label_col=args.label_col + ) + print( + f"[{mode}] relabeled mixed patients to max severity " + f"(changed={changed_rows}, rows={before_rows}→{len(df_mode)}, " + f"remaining_mixed={len(still_mixed)}).", + flush=True, + ) + else: + print(f"[{mode}] keeping mixed-label patients with raw per-eye labels.", flush=True) + + if mode == "binary": + df_mode = df_mode[df_mode[args.label_col].isin([0, 1])].reset_index(drop=True) + + num_classes = 2 if mode == "binary" else int(df_mode[args.label_col].nunique()) + print( + f"\n[{mode}] num_classes={num_classes} rows={len(df_mode)} " + f"patients={df_mode['Patient ID'].nunique()}", + flush=True, + ) + + split_manager = PatientFirstSplitManager( + patient_col="Patient ID", label_col=args.label_col + ) + split_args = SimpleNamespace( + eval_mode=mode, + holdout_per_class=args.holdout_per_class, + holdout_seed=args.holdout_seed, + n_splits=args.n_splits, + fold_seed=args.fold_seed, + ) + clinical_ns = SimpleNamespace(df=df_mode, label_col=args.label_col) + plans = split_manager.build_plans(clinical=clinical_ns, args=split_args, profile=None) + requested_folds = args.n_splits if args.folds is None else int(args.folds) + n_folds = min(requested_folds, len(plans)) + + tm_dir = out_dir / mode / tower_mode + tm_dir.mkdir(parents=True, exist_ok=True) + fold_results: list[FoldResult] = [] + + # Override profiles with df_mode slice. + profile_eye = build_papila_profile( + patient_col="Patient ID", label_col=args.label_col, sample_mode="eye" + ) + profile_patient = build_papila_profile( + patient_col="Patient ID", label_col=args.label_col, sample_mode="patient" + ) + + # ---- PredictionStore — build once before fold loop --------------- + fused_head = getattr(args, "fused_head", False) + _head_names = head_names_for_mode(tower_mode, fused_head=fused_head) + fusion_epochs = int(getattr(args, "fusion_epochs", 10)) if fused_head else 0 + _global_warmup_tower = getattr(args, "warmup_tower_epochs", None) + _global_warmup_fused = getattr(args, "warmup_fused_epochs", None) + _warmup_tower = ( + int(args.single_warmup_tower_epochs) + if getattr(args, "single_warmup_tower_epochs", None) is not None + else int(_global_warmup_tower) if _global_warmup_tower is not None else 2 + ) + _warmup_fused = ( + int(args.single_warmup_fused_epochs) + if getattr(args, "single_warmup_fused_epochs", None) is not None + else int(_global_warmup_fused) if _global_warmup_fused is not None else 2 + ) + _warmup_cd = int(getattr(args, "warmup_cd_epochs", 0)) + _total_epochs = _warmup_cd + _warmup_tower + _warmup_fused + int(args.epochs) + fusion_epochs + + # sample IDs depend on mode: single uses eye IDs, others use patient IDs + if tower_mode in ("single", "classic"): + _sample_ids = [ + f"{row['Patient ID']}{row['eyeID']}" + for _, row in df_mode.iterrows() + ] + _y_true = df_mode[args.label_col].tolist() + else: + # one row per patient (deduplicate — take first occurrence per patient) + _pat_df = df_mode.drop_duplicates(subset="Patient ID") + _sample_ids = _pat_df["Patient ID"].astype(str).tolist() + _y_true = _pat_df[args.label_col].tolist() + + pred_store = PredictionStore( + sample_ids=_sample_ids, + y_true=_y_true, + head_names=_head_names, + n_folds=n_folds, + n_epochs=_total_epochs, + n_classes=num_classes, + ) + + image_cache: dict | None = {} if getattr(args, "in_memory_cache", False) else None + + for fold in range(n_folds): + seed_everything(args.seed + fold * 100) + fold_dir = tm_dir / f"fold{fold}" + fold_dir.mkdir(exist_ok=True) + + print(f"\n[{mode}:{tower_mode}] fold {fold+1}/{n_folds}", flush=True) + result, artifacts = self._run_fold( + fold=fold, + split=plans[fold], + mode=mode, + data=self.data, + num_classes=num_classes, + profile_eye=profile_eye, + profile_patient=profile_patient, + fold_dir=fold_dir, + tower_mode=tower_mode, + pred_store=pred_store, + image_cache=image_cache, + ) + fold_results.append(result) + if artifacts.y_true_ensemble is not None: + np.save(fold_dir / "y_true.npy", artifacts.y_true_ensemble) + if artifacts.probs_ensemble is not None: + np.save(fold_dir / "probs_fused.npy", artifacts.probs_ensemble) + if artifacts.probs_ensemble_img is not None: + np.save(fold_dir / "probs_img.npy", artifacts.probs_ensemble_img) + if artifacts.probs_ensemble_md is not None: + np.save(fold_dir / "probs_cd.npy", artifacts.probs_ensemble_md) + if artifacts.y_true_classic is not None: + np.save(fold_dir / "y_true.npy", artifacts.y_true_classic) + if artifacts.probs_classic is not None: + np.save(fold_dir / "probs_classic.npy", artifacts.probs_classic) + if artifacts.probs_classic_img is not None: + np.save(fold_dir / "probs_classic_img.npy", artifacts.probs_classic_img) + if artifacts.probs_classic_md is not None: + np.save(fold_dir / "probs_classic_md.npy", artifacts.probs_classic_md) + if artifacts.y_true_ensemble_pereye is not None: + np.save(fold_dir / "y_true_pereye.npy", artifacts.y_true_ensemble_pereye) + if artifacts.probs_ensemble_pereye is not None: + np.save(fold_dir / "probs_fused_pereye.npy", artifacts.probs_ensemble_pereye) + if artifacts.probs_ensemble_img_pereye is not None: + np.save(fold_dir / "probs_img_pereye.npy", artifacts.probs_ensemble_img_pereye) + if artifacts.probs_ensemble_md_pereye is not None: + np.save(fold_dir / "probs_cd_pereye.npy", artifacts.probs_ensemble_md_pereye) + if artifacts.logits_ensemble is not None: + np.save(fold_dir / "logits_fused.npy", artifacts.logits_ensemble) + if artifacts.logits_ensemble_img is not None: + np.save(fold_dir / "logits_img.npy", artifacts.logits_ensemble_img) + if artifacts.logits_ensemble_md is not None: + np.save(fold_dir / "logits_cd.npy", artifacts.logits_ensemble_md) + if artifacts.logits_classic is not None: + np.save(fold_dir / "logits_classic.npy", artifacts.logits_classic) + if artifacts.logits_classic_img is not None: + np.save(fold_dir / "logits_classic_img.npy", artifacts.logits_classic_img) + if artifacts.logits_classic_md is not None: + np.save(fold_dir / "logits_classic_md.npy", artifacts.logits_classic_md) + if artifacts.logits_ensemble_pereye is not None: + np.save(fold_dir / "logits_fused_pereye.npy", artifacts.logits_ensemble_pereye) + if artifacts.logits_ensemble_img_pereye is not None: + np.save(fold_dir / "logits_img_pereye.npy", artifacts.logits_ensemble_img_pereye) + if artifacts.logits_ensemble_md_pereye is not None: + np.save(fold_dir / "logits_cd_pereye.npy", artifacts.logits_ensemble_md_pereye) + if artifacts.probs_bilat is not None: + np.save(fold_dir / "probs_bilat.npy", artifacts.probs_bilat) + if artifacts.probs_fused is not None: + np.save(fold_dir / "probs_fused_head.npy", artifacts.probs_fused) + # y_true is shared across all heads for the same fold + if artifacts.y_true_bilat is not None and artifacts.y_true_ensemble is None: + np.save(fold_dir / "y_true.npy", artifacts.y_true_bilat) + # per-sample prediction CSVs + if artifacts.y_true_ensemble is not None: + _save_predictions_csv( + fold_dir, mode, artifacts.y_true_ensemble, + {"fused": artifacts.probs_ensemble, + "img": artifacts.probs_ensemble_img, + "md": artifacts.probs_ensemble_md}, + ) + if artifacts.y_true_ensemble_pereye is not None: + _save_predictions_csv( + fold_dir, mode, artifacts.y_true_ensemble_pereye, + {"fused": artifacts.probs_ensemble_pereye, + "img": artifacts.probs_ensemble_img_pereye, + "md": artifacts.probs_ensemble_md_pereye}, + suffix="_pereye", + ) + if artifacts.y_true_classic is not None: + _save_predictions_csv( + fold_dir, mode, artifacts.y_true_classic, + {"fused": artifacts.probs_classic, + "img": artifacts.probs_classic_img, + "md": artifacts.probs_classic_md}, + suffix="_classic", + ) + + fold_csv = tm_dir / "fold_results.csv" + csv_fields = list(FoldResult.__dataclass_fields__.keys()) + with fold_csv.open("w", newline="", encoding="utf-8") as fh: + w = csv.DictWriter(fh, fieldnames=csv_fields) + w.writeheader() + for r in fold_results: + w.writerow({k: getattr(r, k) for k in csv_fields}) + + summary = self._summary(fold_results) + self._print_summary(f"{mode}:{tower_mode}", summary, tower_mode=tower_mode) + metric_key = { + "single": "classic_val_auc", + "ensemble": "ensemble_val_auc", + "bilateral": "bilat_val_auc", + }[tower_mode] + fold_metrics = [] + best_vals = [] + for r in fold_results: + best_val = getattr(r, metric_key) + fold_metrics.append({ + "fold": r.fold, + "best_metric_value": _f(best_val), + "best_epoch": (r.best_epoch_bilat if tower_mode == "bilateral" else r.best_epoch_single), + "monitor": metric_key, + }) + if not np.isnan(float(best_val)): + best_vals.append(float(best_val)) + + ts_now = time.strftime("%Y%m%d_%H%M%S") + mode_summary = { + "run_id": run_name, + "backbone": args.backbone, + "epochs": args.epochs, + "warmup_cd_epochs": getattr(args, "warmup_cd_epochs", 0), + "warmup_tower_epochs": args.warmup_tower_epochs, + "warmup_fused_epochs": args.warmup_fused_epochs, + "single_warmup_tower_epochs": args.single_warmup_tower_epochs, + "single_warmup_fused_epochs": args.single_warmup_fused_epochs, + "bilat_warmup_tower_epochs": args.bilat_warmup_tower_epochs, + "bilat_warmup_fused_epochs": args.bilat_warmup_fused_epochs, + "batch_size": args.batch_size, + "lr": args.lr, + "eval_mode": mode, + "tower_mode": tower_mode, + "n_splits": n_folds, + "best_metric": metric_key, + "best_metric_mode": "max", + "best_metric_mean": (float(np.mean(best_vals)) if best_vals else None), + "best_metric_std": (float(np.std(best_vals)) if best_vals else None), + "fold_metrics": fold_metrics, + "mode_summary": summary, + } + (tm_dir / "summary.json").write_text(json.dumps(mode_summary, indent=2), encoding="utf-8") + pred_store.save(tm_dir / "predictions.npz") + + root_summary_path = out_dir / "summary.json" + if root_summary_path.exists(): + try: + payload = json.loads(root_summary_path.read_text(encoding="utf-8")) + except Exception: + payload = {} + else: + payload = {} + payload.setdefault("run_name", run_name) + payload.setdefault("timestamp", ts_now) + payload["config"] = vars(args) + payload.setdefault("summaries", {}) + payload["summaries"][f"{mode}:{tower_mode}"] = summary + root_summary_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") + + print(f"\nOutputs written to: {out_dir}") + return out_dir + + # ------------------------------------------------------------------ + # Fold runner + # ------------------------------------------------------------------ + + def _run_fold( + self, + fold: int, + split, + mode: str, + data, + num_classes: int, + profile_eye, + profile_patient, + fold_dir: Path, + tower_mode: str, + pred_store: "PredictionStore | None" = None, + image_cache: "dict | None" = None, + ) -> tuple[FoldResult, FoldArtifacts]: + args = self.args + device = self.device + image_preprocessor = self.image_preprocessor + nan = _nan() + tower_mode = "single" if tower_mode == "classic" else tower_mode + run_single = tower_mode in ("single", "ensemble") + run_bilat = tower_mode == "bilateral" + run_fused = (tower_mode == "ensemble") and bool(getattr(args, "fused_head", False)) + + global_warmup_tower = getattr(args, "warmup_tower_epochs", None) + global_warmup_fused = getattr(args, "warmup_fused_epochs", None) + single_warmup_tower = ( + int(args.single_warmup_tower_epochs) + if getattr(args, "single_warmup_tower_epochs", None) is not None + else int(global_warmup_tower) if global_warmup_tower is not None else 2 + ) + single_warmup_fused = ( + int(args.single_warmup_fused_epochs) + if getattr(args, "single_warmup_fused_epochs", None) is not None + else int(global_warmup_fused) if global_warmup_fused is not None else 2 + ) + bilat_warmup_tower = ( + int(args.bilat_warmup_tower_epochs) + if getattr(args, "bilat_warmup_tower_epochs", None) is not None + else int(global_warmup_tower) if global_warmup_tower is not None else 4 + ) + bilat_warmup_fused = ( + int(args.bilat_warmup_fused_epochs) + if getattr(args, "bilat_warmup_fused_epochs", None) is not None + else int(global_warmup_fused) if global_warmup_fused is not None else 3 + ) + single_warmup_cd = int(getattr(args, "warmup_cd_epochs", 0)) if run_single else 0 + if not run_single: + single_warmup_tower = 0 + single_warmup_fused = 0 + if not run_bilat: + bilat_warmup_tower = 0 + bilat_warmup_fused = 0 + main_epochs = int(args.epochs) + total_single_epochs = (single_warmup_cd + single_warmup_tower + single_warmup_fused + main_epochs) if run_single else 0 + total_bilat_epochs = (bilat_warmup_tower + bilat_warmup_fused + main_epochs) if run_bilat else 0 + total_epochs = max(total_single_epochs, total_bilat_epochs) + + # ---- samples --------------------------------------------------- + eye_train = filter_eye_samples(profile_eye.build_samples(df=split.train, clinical=data)) + bilat_train = filter_bilateral_samples(profile_patient.build_samples(df=split.train, clinical=data)) + bilat_val = filter_bilateral_samples(profile_patient.build_samples(df=split.val, clinical=data)) + + # Register split labels in the prediction store + if pred_store is not None: + if tower_mode in ("single", "classic"): + # eye-level IDs: "{patient_id}{eyeID}" + train_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in eye_train] + val_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in bilat_val] + else: + train_sids = [str(s["id_1"]) for s in bilat_train] + val_sids = [str(s["id_1"]) for s in bilat_val] + pred_store.set_split(fold, train_sids, "train") + pred_store.set_split(fold, val_sids, "val") + + if len(bilat_val) == 0: + print(f" [fold {fold+1}] WARNING: no bilateral val samples; skipping fold.", flush=True) + empty = FoldResult( + mode=mode, fold=fold, + best_epoch_single=0, best_epoch_bilat=0, + classic_val_auc=nan, classic_val_acc=nan, classic_val_kappa=nan, + classic_val_mcc=nan, classic_val_f1=nan, classic_val_recall=None, + classic_val_ece=nan, classic_val_threshold=nan, classic_val_bias=None, + classic_val_n=0, + ensemble_val_auc=nan, ensemble_val_acc=nan, ensemble_val_kappa=nan, + ensemble_val_mcc=nan, ensemble_val_f1=nan, ensemble_val_recall=None, + ensemble_val_ece=nan, ensemble_val_threshold=nan, ensemble_val_bias=None, + ensemble_val_n=0, + bilat_val_auc=nan, bilat_val_acc=nan, bilat_val_kappa=nan, + bilat_val_mcc=nan, bilat_val_f1=nan, bilat_val_recall=None, + bilat_val_ece=nan, bilat_val_threshold=nan, bilat_val_bias=None, + bilat_val_n=0, + classic_holdout_auc=nan, classic_holdout_acc=nan, + ensemble_holdout_auc=nan, ensemble_holdout_acc=nan, + bilat_holdout_auc=nan, bilat_holdout_acc=nan, + holdout_n=0, + single_train_n=len(eye_train), bilat_train_n=len(bilat_train), + ) + return empty, FoldArtifacts( + y_true_classic=None, probs_classic=None, + y_true_ensemble=None, probs_ensemble=None, + y_true_bilat=None, probs_bilat=None, + ) + + # ---- models ---------------------------------------------------- + single = None + bilateral = None + if run_single: + single = SingleEyeHT( + backbone=args.backbone, freeze_ratio=args.freeze_ratio, + augment=args.augment, clinical_data=data, + num_classes=num_classes, + cd_hidden_dim=args.cd_hidden_dim, fusion_dim=args.fusion_dim, + bridge_mode=getattr(args, "bridge_mode", "fused"), + ).to(device) + if run_bilat: + bilateral = BilateralHT( + backbone=args.backbone, freeze_ratio=args.freeze_ratio, + augment=args.augment, clinical_data=data, + num_classes=num_classes, + cd_hidden_dim=args.cd_hidden_dim, fusion_dim=args.fusion_dim, + ).to(device) + + slots_eye = profile_eye.slot_descriptors() + slots_patient = profile_patient.slot_descriptors() + loader_kw = dict(batch_size=args.batch_size, num_workers=args.num_workers, + image_cache=image_cache) + + # ---- loaders --------------------------------------------------- + use_balanced = bool(getattr(args, "balanced_sampling", False)) + train_single_loader = None + train_eval_loader = None # non-shuffled, no sampler — for per-epoch train logging + train_bilat_loader = None + cd_only_loader = None # image-free loader for md_warmup phase + if run_single: + single_sampler = build_balanced_sampler(eye_train) if use_balanced else None + train_single_loader = make_loader( + eye_train, slots_eye, + image_transform=single.transform, + image_preprocessor=image_preprocessor, + shuffle=True, + sampler=single_sampler, + **loader_kw, + ) + train_eval_loader = make_loader( + eye_train, slots_eye, + image_transform=build_eval_transform(args.backbone), + image_preprocessor=image_preprocessor, + shuffle=False, + **loader_kw, + ) + if single_warmup_cd > 0: + # MD-only loader: drop image_1 so PIL never opens files during md_warmup. + # Always use balanced sampling for md_warmup — MD features alone are weaker + # than images and collapse to majority class without class balancing. + slots_cd_only = {k: v for k, v in slots_eye.items() if k != "image_1"} + md_warmup_sampler = single_sampler if single_sampler is not None else build_balanced_sampler(eye_train) + cd_only_loader = make_loader( + eye_train, slots_cd_only, + image_transform=None, + image_preprocessor=None, + shuffle=True, + sampler=md_warmup_sampler, + **loader_kw, + ) + if run_bilat: + bilat_sampler = build_balanced_sampler(bilat_train) if use_balanced else None + train_bilat_loader = make_loader( + bilat_train, slots_patient, + image_transform=bilateral.transform, + image_preprocessor=image_preprocessor, + shuffle=True, + sampler=bilat_sampler, + **loader_kw, + ) + elif run_fused: + # Fused head trains on bilateral samples using the single model's transform. + fused_sampler = build_balanced_sampler(bilat_train) if use_balanced else None + train_bilat_loader = make_loader( + bilat_train, slots_patient, + image_transform=single.transform, + image_preprocessor=image_preprocessor, + shuffle=True, + sampler=fused_sampler, + **loader_kw, + ) + eval_transform = build_eval_transform(args.backbone) + val_loader = make_loader( + bilat_val, slots_patient, + image_transform=eval_transform, + image_preprocessor=image_preprocessor, + shuffle=False, + **loader_kw, + ) + + # ---- holdout loader (optional) --------------------------------- + holdout_bilat: list = [] + holdout_loader = None + if split.holdout is not None and not split.holdout.empty: + holdout_bilat = filter_bilateral_samples( + profile_patient.build_samples(df=split.holdout, clinical=data) + ) + if holdout_bilat: + holdout_loader = make_loader( + holdout_bilat, slots_patient, + image_transform=eval_transform, + image_preprocessor=image_preprocessor, + shuffle=False, + **loader_kw, + ) + print(f" [fold {fold+1}] holdout_n={len(holdout_bilat)} (bilateral patients)", flush=True) + if pred_store is not None: + pred_store.set_split(fold, [str(s["id_1"]) for s in holdout_bilat], "holdout") + + # ---- prebuild in-memory image cache (fold 0 only; shared dict fills for later folds) ---- + if image_cache is not None: + cache_workers = int(getattr(args, "cache_workers", 4)) + _loaders_to_warm = [ + train_single_loader, train_bilat_loader, val_loader, holdout_loader, + ] + for _ldr in _loaders_to_warm: + if _ldr is not None: + _ldr.dataset.prebuild_image_cache(cache_workers=cache_workers) + + opt_single = torch.optim.Adam(single.parameters(), lr=args.lr) if run_single else None + opt_bilateral = torch.optim.Adam(bilateral.parameters(), lr=args.lr) if run_bilat else None + + # ---- epoch log ------------------------------------------------- + epoch_fields = [ + "fold", "epoch", + "phase_single", "phase_bilat", + "main_epoch_single", "main_epoch_bilat", + "single_active", "bilat_active", + "single_train_loss", "single_train_acc", + # val — fused head (existing) + "classic_val_auc", "classic_val_acc", "classic_val_n", + "ensemble_val_auc", "ensemble_val_acc", "ensemble_val_n", + "bilat_train_loss", "bilat_train_acc", + "bilat_val_auc", "bilat_val_acc", "bilat_val_n", + # val — img/md heads + fusion events + "classic_val_auc_img", "classic_val_acc_img", + "classic_val_auc_cd", "classic_val_acc_cd", + "classic_val_fe_corr", "classic_val_fe_err", + "ensemble_val_auc_img", "ensemble_val_acc_img", + "ensemble_val_auc_cd", "ensemble_val_acc_cd", + "ensemble_val_fe_corr", "ensemble_val_fe_err", + "bilat_val_auc_img", "bilat_val_acc_img", + "bilat_val_auc_cd", "bilat_val_acc_cd", + "bilat_val_fe_corr", "bilat_val_fe_err", + # holdout — fused head (existing) + "classic_holdout_auc", "classic_holdout_acc", + "ensemble_holdout_auc", "ensemble_holdout_acc", + "bilat_holdout_auc", "bilat_holdout_acc", + # holdout — img/md heads + fusion events + "classic_holdout_auc_img", "classic_holdout_acc_img", + "classic_holdout_auc_cd", "classic_holdout_acc_cd", + "classic_holdout_fe_corr", "classic_holdout_fe_err", + "ensemble_holdout_auc_img", "ensemble_holdout_acc_img", + "ensemble_holdout_auc_cd", "ensemble_holdout_acc_cd", + "ensemble_holdout_fe_corr", "ensemble_holdout_fe_err", + # train-set eval pass (eval mode, all 3 heads) + "train_auc_fused", "train_acc_fused", + "train_auc_img", "train_acc_img", + "train_auc_cd", "train_acc_cd", + "train_fe_corr", "train_fe_err", + "train_n", + "is_best_single", "is_best_bilat", + "is_best_holdout_single", "is_best_holdout_bilat", + ] + # CM columns — named by num_classes so binary and multiclass both work + if num_classes == 2: + _cm_keys = ["tn", "fp", "fn", "tp"] + else: + _cm_keys = [f"cm_{i}_{j}" for i in range(num_classes) for j in range(num_classes)] + for _split in ("classic_val", "ensemble_val", "classic_holdout", "ensemble_holdout", "train"): + for _head in ("fused", "img", "md"): + for _k in _cm_keys: + epoch_fields.append(f"{_split}_{_head}_{_k}") + fold_logger = HypertowerLogger(run_dir=fold_dir) + + # per-epoch accumulation for npy tensors + _epoch_train_pf: list[np.ndarray] = [] + _epoch_train_pi: list[np.ndarray] = [] + _epoch_train_pm: list[np.ndarray] = [] + _epoch_train_ids: list[np.ndarray] = [] + _epoch_train_y: list[np.ndarray] = [] + # per-eye val accumulators (ensemble mode: OD and OS separate) + _epoch_val_pf_od: list[np.ndarray] = [] + _epoch_val_pi_od: list[np.ndarray] = [] + _epoch_val_pm_od: list[np.ndarray] = [] + _epoch_val_pf_os: list[np.ndarray] = [] + _epoch_val_pi_os: list[np.ndarray] = [] + _epoch_val_pm_os: list[np.ndarray] = [] + _epoch_val_y: list[np.ndarray] = [] + _epoch_val_ids: list[np.ndarray] = [] + + # ---- best-epoch trackers --------------------------------------- + best_single_auc = -1.0 + best_bilat_auc = -1.0 + best_epoch_single = 0 + best_epoch_bilat = 0 + best_single_state: Optional[dict] = None + best_bilat_state: Optional[dict] = None + snap_classic: dict = {} + snap_ensemble: dict = {} + snap_bilat: dict = {} + snap_holdout_single: dict = {} + snap_holdout_bilat: dict = {} + snap_fused: dict = {} + snap_holdout_fused: dict = {} + best_holdout_single_auc = -1.0 + best_holdout_bilat_auc = -1.0 + best_epoch_holdout_single = 0 + best_epoch_holdout_bilat = 0 + best_holdout_single_state: Optional[dict] = None + best_holdout_bilat_state: Optional[dict] = None + + if run_single: + print( + f" [fold {fold+1}] single_train_n={len(eye_train)} (eye-level) " + f"val_n={len(bilat_val)} " + f"single_warmup=md{single_warmup_cd}+twr{single_warmup_tower}+fus{single_warmup_fused} total={total_single_epochs}", + flush=True, + ) + else: + print( + f" [fold {fold+1}] bilat_train_n={len(bilat_train)} (bilateral) " + f"val_n={len(bilat_val)} " + f"bilat_warmup={bilat_warmup_tower}+{bilat_warmup_fused} total={total_bilat_epochs}", + flush=True, + ) + + # ---- epoch loop ------------------------------------------------ + _prev_phase_single = "inactive" # used to detect md_warmup → next phase transition + for epoch in range(total_epochs): + _epoch_t0 = time.time() + if not run_single: + phase_single, main_epoch_single, single_active = "inactive", 0, False + elif epoch < single_warmup_cd: + phase_single, main_epoch_single, single_active = "cd_warmup", 0, True + elif epoch < (single_warmup_cd + single_warmup_tower): + phase_single, main_epoch_single, single_active = "tower_warmup", 0, True + elif epoch < (single_warmup_cd + single_warmup_tower + single_warmup_fused): + phase_single, main_epoch_single, single_active = "fused_warmup", 0, True + elif epoch < total_single_epochs: + phase_single, main_epoch_single, single_active = ( + "main", + epoch - single_warmup_cd - single_warmup_tower - single_warmup_fused + 1, + True, + ) + else: + phase_single, main_epoch_single, single_active = "done", main_epochs, False + + if not run_bilat: + phase_bilat, main_epoch_bilat, bilat_active = "inactive", 0, False + elif epoch < bilat_warmup_tower: + phase_bilat, main_epoch_bilat, bilat_active = "tower_warmup", 0, True + elif epoch < (bilat_warmup_tower + bilat_warmup_fused): + phase_bilat, main_epoch_bilat, bilat_active = "fused_warmup", 0, True + elif epoch < total_bilat_epochs: + phase_bilat, main_epoch_bilat, bilat_active = ( + "main", + epoch - bilat_warmup_tower - bilat_warmup_fused + 1, + True, + ) + else: + phase_bilat, main_epoch_bilat, bilat_active = "done", main_epochs, False + + if run_single and single_active: + _active_loader = cd_only_loader if phase_single == "cd_warmup" else train_single_loader + sl_loss, sl_acc = train_single_epoch( + single, _active_loader, opt_single, device, + phase=phase_single, bcd_prob=float(args.bcd_prob), + tower_loss_mode=args.tower_loss_mode, + ) + else: + sl_loss, sl_acc = nan, nan + + if run_bilat and bilat_active: + bl_loss, bl_acc = train_bilateral_epoch( + bilateral, train_bilat_loader, opt_bilateral, device, + phase=phase_bilat, bcd_prob=float(args.bcd_prob), + tower_loss_mode=args.tower_loss_mode, + ) + else: + bl_loss, bl_acc = nan, nan + + _skip_val_eval = (phase_single == "cd_warmup") + + if run_single and tower_mode == "single" and not _skip_val_eval: + y_cl, p_cl, p_cl_img, p_cl_cd = collect_probs_single_components( + single, val_loader, device, aggregate_patient=False + ) + cl_acc, cl_auc, cl_n = _score_arrays(y_cl, p_cl, num_classes) + cl_acc_img = float((p_cl_img.argmax(1) == y_cl).mean()) if y_cl.size else nan + cl_acc_cd = float((p_cl_cd.argmax(1) == y_cl).mean()) if y_cl.size else nan + _, cl_auc_img, _ = _score_arrays(y_cl, p_cl_img, num_classes) + _, cl_auc_cd, _ = _score_arrays(y_cl, p_cl_cd, num_classes) + y_en = np.array([], dtype=np.int64) + p_en = p_en_img = p_en_cd = np.zeros((0, num_classes), dtype=np.float32) + en_acc = en_auc = nan + en_n = 0 + en_acc_img = en_acc_cd = en_auc_img = en_auc_cd = nan + elif run_single and tower_mode == "ensemble" and not _skip_val_eval: + (y_en, + _p_en_f_od, _p_en_i_od, _p_en_m_od, + _p_en_f_os, _p_en_i_os, _p_en_m_os, + _en_pat_ids) = collect_probs_ensemble_pereye( + single, val_loader, device, return_ids=True + ) + # patient-level averages (used for metrics, same as before) + p_en = 0.5 * (_p_en_f_od + _p_en_f_os) + p_en_img = 0.5 * (_p_en_i_od + _p_en_i_os) + p_en_cd = 0.5 * (_p_en_m_od + _p_en_m_os) + en_acc, en_auc, en_n = _score_arrays(y_en, p_en, num_classes) + en_acc_img = float((p_en_img.argmax(1) == y_en).mean()) if y_en.size else nan + en_acc_cd = float((p_en_cd.argmax(1) == y_en).mean()) if y_en.size else nan + _, en_auc_img, _ = _score_arrays(y_en, p_en_img, num_classes) + _, en_auc_cd, _ = _score_arrays(y_en, p_en_cd, num_classes) + y_cl = np.array([], dtype=np.int64) + p_cl = p_cl_img = p_cl_cd = np.zeros((0, num_classes), dtype=np.float32) + cl_acc = cl_auc = nan + cl_n = 0 + cl_acc_img = cl_acc_cd = cl_auc_img = cl_auc_cd = nan + else: + y_cl = y_en = np.array([], dtype=np.int64) + p_cl = p_cl_img = p_cl_cd = np.zeros((0, num_classes), dtype=np.float32) + p_en = p_en_img = p_en_cd = np.zeros((0, num_classes), dtype=np.float32) + cl_acc = cl_auc = en_acc = en_auc = nan + cl_n = en_n = 0 + cl_acc_img = cl_acc_cd = en_acc_img = en_acc_cd = nan + cl_auc_img = cl_auc_cd = en_auc_img = en_auc_cd = nan + + if run_bilat and not _skip_val_eval: + y_bi, p_bi, p_bi_img, p_bi_cd = collect_probs_bilateral_components( + bilateral, val_loader, device + ) + bi_acc, bi_auc, bi_n = _score_arrays(y_bi, p_bi, num_classes) + bi_acc_img = float((p_bi_img.argmax(1) == y_bi).mean()) if y_bi.size else nan + bi_acc_cd = float((p_bi_cd.argmax(1) == y_bi).mean()) if y_bi.size else nan + _, bi_auc_img, _ = _score_arrays(y_bi, p_bi_img, num_classes) + _, bi_auc_cd, _ = _score_arrays(y_bi, p_bi_cd, num_classes) + else: + y_bi = np.array([], dtype=np.int64) + p_bi = np.zeros((0, 0), dtype=np.float32) + bi_acc = bi_auc = nan + bi_n = 0 + bi_acc_img = bi_acc_cd = bi_auc_img = bi_auc_cd = nan + + # --- holdout evaluation ------------------------------------ + # defaults (overwritten below when holdout_loader is not None) + _z2 = np.zeros((0, num_classes), dtype=np.float32) + _e2 = np.array([], dtype=np.int64) + y_cl_h = y_en_h = _e2 + p_cl_h = p_cl_h_img = p_cl_h_cd = _z2 + p_en_h = p_en_h_img = p_en_h_cd = _z2 + + if holdout_loader is not None and not _skip_val_eval: + if run_single and tower_mode == "single": + y_cl_h, p_cl_h, p_cl_h_img, p_cl_h_cd = collect_probs_single_components( + single, holdout_loader, device, aggregate_patient=False + ) + _, cl_auc_h, _ = _score_arrays(y_cl_h, p_cl_h, num_classes) + cl_acc_h = float((p_cl_h.argmax(1) == y_cl_h).mean()) if y_cl_h.size else nan + _, cl_auc_h_img, _ = _score_arrays(y_cl_h, p_cl_h_img, num_classes) + cl_acc_h_img = float((p_cl_h_img.argmax(1) == y_cl_h).mean()) if y_cl_h.size else nan + _, cl_auc_h_md, _ = _score_arrays(y_cl_h, p_cl_h_cd, num_classes) + cl_acc_h_md = float((p_cl_h_cd.argmax(1) == y_cl_h).mean()) if y_cl_h.size else nan + cl_fe_h_corr, cl_fe_h_err = _fusion_events(y_cl_h, p_cl_h, p_cl_h_img, p_cl_h_cd) + en_auc_h = en_acc_h = nan + en_auc_h_img = en_acc_h_img = en_auc_h_md = en_acc_h_md = nan + en_fe_h_corr = en_fe_h_err = 0 + elif run_single and tower_mode == "ensemble": + y_en_h, p_en_h, p_en_h_img, p_en_h_cd = collect_probs_single_components( + single, holdout_loader, device, aggregate_patient=True + ) + _, en_auc_h, _ = _score_arrays(y_en_h, p_en_h, num_classes) + en_acc_h = float((p_en_h.argmax(1) == y_en_h).mean()) if y_en_h.size else nan + _, en_auc_h_img, _ = _score_arrays(y_en_h, p_en_h_img, num_classes) + en_acc_h_img = float((p_en_h_img.argmax(1) == y_en_h).mean()) if y_en_h.size else nan + _, en_auc_h_md, _ = _score_arrays(y_en_h, p_en_h_cd, num_classes) + en_acc_h_md = float((p_en_h_cd.argmax(1) == y_en_h).mean()) if y_en_h.size else nan + en_fe_h_corr, en_fe_h_err = _fusion_events(y_en_h, p_en_h, p_en_h_img, p_en_h_cd) + cl_auc_h = cl_acc_h = nan + cl_auc_h_img = cl_acc_h_img = cl_auc_h_md = cl_acc_h_md = nan + cl_fe_h_corr = cl_fe_h_err = 0 + else: + cl_auc_h = cl_acc_h = en_auc_h = en_acc_h = nan + cl_auc_h_img = cl_acc_h_img = cl_auc_h_md = cl_acc_h_md = nan + en_auc_h_img = en_acc_h_img = en_auc_h_md = en_acc_h_md = nan + cl_fe_h_corr = cl_fe_h_err = en_fe_h_corr = en_fe_h_err = 0 + if run_bilat: + y_bi_h, p_bi_h, _, _ = collect_probs_bilateral_components( + bilateral, holdout_loader, device + ) + _, bi_auc_h, _ = _score_arrays(y_bi_h, p_bi_h, num_classes) + bi_acc_h = float((p_bi_h.argmax(1) == y_bi_h).mean()) if y_bi_h.size else nan + else: + bi_auc_h = bi_acc_h = nan + else: + cl_auc_h = cl_acc_h = en_auc_h = en_acc_h = bi_auc_h = bi_acc_h = nan + cl_auc_h_img = cl_acc_h_img = cl_auc_h_md = cl_acc_h_md = nan + en_auc_h_img = en_acc_h_img = en_auc_h_md = en_acc_h_md = nan + cl_fe_h_corr = cl_fe_h_err = en_fe_h_corr = en_fe_h_err = 0 + + # --- fusion-event helpers for val sets ---------------------- + cl_fe_corr, cl_fe_err = _fusion_events(y_cl, p_cl, p_cl_img, p_cl_cd) if y_cl.size else (0, 0) + en_fe_corr, en_fe_err = _fusion_events(y_en, p_en, p_en_img, p_en_cd) if y_en.size else (0, 0) + bi_fe_corr, bi_fe_err = (0, 0) # bilateral components not separated the same way + + # --- train eval pass (eval mode, all 3 heads) ---------------- + tr_auc_f = tr_acc_f = tr_auc_i = tr_acc_i = tr_auc_m = tr_acc_m = nan + tr_fe_corr = tr_fe_err = tr_n = 0 + y_tr = np.array([], dtype=np.int64) + p_tr_f = p_tr_i = p_tr_m = np.zeros((0, num_classes), dtype=np.float32) + if run_single and train_eval_loader is not None and not _skip_val_eval: + y_tr, p_tr_f, p_tr_i, p_tr_m, tr_ids = collect_probs_eye_level( + single, train_eval_loader, device, return_ids=True + ) + if y_tr.size: + _, tr_auc_f, _ = _score_arrays(y_tr, p_tr_f, num_classes) + tr_acc_f = float((p_tr_f.argmax(1) == y_tr).mean()) + _, tr_auc_i, _ = _score_arrays(y_tr, p_tr_i, num_classes) + tr_acc_i = float((p_tr_i.argmax(1) == y_tr).mean()) + _, tr_auc_m, _ = _score_arrays(y_tr, p_tr_m, num_classes) + tr_acc_m = float((p_tr_m.argmax(1) == y_tr).mean()) + tr_fe_corr, tr_fe_err = _fusion_events(y_tr, p_tr_f, p_tr_i, p_tr_m) + tr_n = int(y_tr.size) + # accumulate for npy tensors + _epoch_train_pf.append(p_tr_f) + _epoch_train_pi.append(p_tr_i) + _epoch_train_pm.append(p_tr_m) + _epoch_train_ids.append(tr_ids) + _epoch_train_y.append(y_tr) + # record into PredictionStore + if pred_store is not None: + if tower_mode in ("single", "classic"): + pred_store.record(fold, epoch, tr_ids, "fused", p_tr_f) + pred_store.record(fold, epoch, tr_ids, "img", p_tr_i) + pred_store.record(fold, epoch, tr_ids, "md", p_tr_m) + else: # ensemble: separate OD and OS by eye suffix + od_mask = np.array([str(i).endswith("OD") for i in tr_ids]) + os_mask = ~od_mask + od_pids = [str(i)[:-2] for i in tr_ids[od_mask]] + os_pids = [str(i)[:-2] for i in tr_ids[os_mask]] + pred_store.record(fold, epoch, od_pids, "od_fused", p_tr_f[od_mask]) + pred_store.record(fold, epoch, od_pids, "od_img", p_tr_i[od_mask]) + pred_store.record(fold, epoch, od_pids, "od_md", p_tr_m[od_mask]) + pred_store.record(fold, epoch, os_pids, "os_fused", p_tr_f[os_mask]) + pred_store.record(fold, epoch, os_pids, "os_img", p_tr_i[os_mask]) + pred_store.record(fold, epoch, os_pids, "os_md", p_tr_m[os_mask]) + + # accumulate val for npy tensors + if run_single and tower_mode == "ensemble" and y_en.size: + _epoch_val_pf_od.append(_p_en_f_od) + _epoch_val_pi_od.append(_p_en_i_od) + _epoch_val_pm_od.append(_p_en_m_od) + _epoch_val_pf_os.append(_p_en_f_os) + _epoch_val_pi_os.append(_p_en_i_os) + _epoch_val_pm_os.append(_p_en_m_os) + _epoch_val_y.append(y_en) + _epoch_val_ids.append(_en_pat_ids) + elif run_single and tower_mode == "single" and y_cl.size: + # single mode: no per-eye split, reuse same array for both slots + _epoch_val_pf_od.append(p_cl) + _epoch_val_pi_od.append(p_cl_img) + _epoch_val_pm_od.append(p_cl_cd) + _epoch_val_pf_os.append(p_cl) + _epoch_val_pi_os.append(p_cl_img) + _epoch_val_pm_os.append(p_cl_cd) + _epoch_val_y.append(y_cl) + + # record val into PredictionStore + if pred_store is not None: + if run_single and tower_mode == "ensemble" and y_en.size: + pred_store.record(fold, epoch, _en_pat_ids, "od_fused", _p_en_f_od) + pred_store.record(fold, epoch, _en_pat_ids, "od_img", _p_en_i_od) + pred_store.record(fold, epoch, _en_pat_ids, "od_md", _p_en_m_od) + pred_store.record(fold, epoch, _en_pat_ids, "os_fused", _p_en_f_os) + pred_store.record(fold, epoch, _en_pat_ids, "os_img", _p_en_i_os) + pred_store.record(fold, epoch, _en_pat_ids, "os_md", _p_en_m_os) + elif run_single and tower_mode == "single" and y_cl.size: + # val in single mode: collect_probs_single_components(aggregate_patient=False) + # returns interleaved [all_OD, all_OS] per batch — IDs not tracked here yet + pass # single-mode val IDs not currently available; train IDs are sufficient + + # Best-epoch checks (restricted to main phase). + target_single_auc = cl_auc if tower_mode == "single" else en_auc + target_holdout_single_auc = cl_auc_h if tower_mode == "single" else en_auc_h + single_ckpt_eligible = run_single and (phase_single == "main") + is_best_single = ( + single_ckpt_eligible + and (not np.isnan(target_single_auc)) + and (target_single_auc > best_single_auc) + ) + if is_best_single: + best_single_auc = target_single_auc + best_epoch_single = epoch + 1 + best_single_state = copy.deepcopy(single.state_dict()) + if tower_mode == "single": + snap_cl, _, _, _ = _tune_and_snap(y_cl, p_cl, cl_acc, num_classes, args, args.ece_bins) + snap_classic = snap_cl + else: + snap_en, _, _, _ = _tune_and_snap(y_en, p_en, en_acc, num_classes, args, args.ece_bins) + snap_ensemble = snap_en + snap_holdout_single = { + "auc": float(target_holdout_single_auc), + "acc": float(cl_acc_h if tower_mode == "single" else en_acc_h), + } + + is_best_holdout_single = ( + holdout_loader is not None + and single_ckpt_eligible + and (not np.isnan(target_holdout_single_auc)) + and (target_holdout_single_auc > best_holdout_single_auc) + ) + if is_best_holdout_single: + best_holdout_single_auc = target_holdout_single_auc + best_epoch_holdout_single = epoch + 1 + best_holdout_single_state = copy.deepcopy(single.state_dict()) + + bilat_ckpt_eligible = run_bilat and (phase_bilat == "main") + is_best_bilat = ( + bilat_ckpt_eligible + and (not np.isnan(bi_auc)) + and (bi_auc > best_bilat_auc) + ) + if is_best_bilat: + best_bilat_auc = bi_auc + best_epoch_bilat = epoch + 1 + best_bilat_state = copy.deepcopy(bilateral.state_dict()) + snap_bi, _, _, _ = _tune_and_snap(y_bi, p_bi, bi_acc, num_classes, args, args.ece_bins) + snap_bilat = snap_bi + snap_holdout_bilat = {"auc": float(bi_auc_h), "acc": float(bi_acc_h)} + + is_best_holdout_bilat = ( + holdout_loader is not None + and bilat_ckpt_eligible + and (not np.isnan(bi_auc_h)) + and (bi_auc_h > best_holdout_bilat_auc) + ) + if is_best_holdout_bilat: + best_holdout_bilat_auc = bi_auc_h + best_epoch_holdout_bilat = epoch + 1 + best_holdout_bilat_state = copy.deepcopy(bilateral.state_dict()) + + # --- confusion matrix cells per split × head ------------------- + def _prefixed_cm(prefix: str, y: np.ndarray, pf: np.ndarray, + pi: np.ndarray, pm: np.ndarray) -> dict: + out: dict = {} + for head, p in (("fused", pf), ("img", pi), ("md", pm)): + for k, v in _cm_cells(y, p, num_classes).items(): + out[f"{prefix}_{head}_{k}"] = v + return out + + cm_row: dict = {} + cm_row.update(_prefixed_cm("classic_val", y_cl, p_cl, p_cl_img, p_cl_cd)) + cm_row.update(_prefixed_cm("ensemble_val", y_en, p_en, p_en_img, p_en_cd)) + cm_row.update(_prefixed_cm("classic_holdout", y_cl_h, p_cl_h, p_cl_h_img, p_cl_h_cd)) + cm_row.update(_prefixed_cm("ensemble_holdout", y_en_h, p_en_h, p_en_h_img, p_en_h_cd)) + cm_row.update(_prefixed_cm("train", y_tr, p_tr_f, p_tr_i, p_tr_m)) + + fold_logger.write_epoch_row({ + "fold": fold, "epoch": epoch + 1, + "phase_single": phase_single, "phase_bilat": phase_bilat, + "main_epoch_single": main_epoch_single, "main_epoch_bilat": main_epoch_bilat, + "single_active": int(single_active), "bilat_active": int(bilat_active), + "single_train_loss": _f(sl_loss), "single_train_acc": _f(sl_acc), + # val — fused + "classic_val_auc": _f(cl_auc), "classic_val_acc": _f(cl_acc), "classic_val_n": cl_n, + "ensemble_val_auc": _f(en_auc), "ensemble_val_acc": _f(en_acc), "ensemble_val_n": en_n, + "bilat_train_loss": _f(bl_loss), "bilat_train_acc": _f(bl_acc), + "bilat_val_auc": _f(bi_auc), "bilat_val_acc": _f(bi_acc), "bilat_val_n": bi_n, + # val — img/md + fusion events + "classic_val_auc_img": _f(cl_auc_img), "classic_val_acc_img": _f(cl_acc_img), + "classic_val_auc_cd": _f(cl_auc_cd), "classic_val_acc_cd": _f(cl_acc_cd), + "classic_val_fe_corr": cl_fe_corr, "classic_val_fe_err": cl_fe_err, + "ensemble_val_auc_img": _f(en_auc_img), "ensemble_val_acc_img": _f(en_acc_img), + "ensemble_val_auc_cd": _f(en_auc_cd), "ensemble_val_acc_cd": _f(en_acc_cd), + "ensemble_val_fe_corr": en_fe_corr, "ensemble_val_fe_err": en_fe_err, + "bilat_val_auc_img": _f(bi_auc_img), "bilat_val_acc_img": _f(bi_acc_img), + "bilat_val_auc_cd": _f(bi_auc_cd), "bilat_val_acc_cd": _f(bi_acc_cd), + "bilat_val_fe_corr": bi_fe_corr, "bilat_val_fe_err": bi_fe_err, + # holdout — fused + "classic_holdout_auc": _f(cl_auc_h), "classic_holdout_acc": _f(cl_acc_h), + "ensemble_holdout_auc": _f(en_auc_h), "ensemble_holdout_acc": _f(en_acc_h), + "bilat_holdout_auc": _f(bi_auc_h), "bilat_holdout_acc": _f(bi_acc_h), + # holdout — img/md + fusion events + "classic_holdout_auc_img": _f(cl_auc_h_img), "classic_holdout_acc_img": _f(cl_acc_h_img), + "classic_holdout_auc_cd": _f(cl_auc_h_md), "classic_holdout_acc_cd": _f(cl_acc_h_md), + "classic_holdout_fe_corr": cl_fe_h_corr, "classic_holdout_fe_err": cl_fe_h_err, + "ensemble_holdout_auc_img": _f(en_auc_h_img), "ensemble_holdout_acc_img": _f(en_acc_h_img), + "ensemble_holdout_auc_cd": _f(en_auc_h_md), "ensemble_holdout_acc_cd": _f(en_acc_h_md), + "ensemble_holdout_fe_corr": en_fe_h_corr, "ensemble_holdout_fe_err": en_fe_h_err, + # train eval pass + "train_auc_fused": _f(tr_auc_f), "train_acc_fused": _f(tr_acc_f), + "train_auc_img": _f(tr_auc_i), "train_acc_img": _f(tr_acc_i), + "train_auc_cd": _f(tr_auc_m), "train_acc_cd": _f(tr_acc_m), + "train_fe_corr": tr_fe_corr, "train_fe_err": tr_fe_err, + "train_n": tr_n, + "is_best_single": int(is_best_single), + "is_best_bilat": int(is_best_bilat), + "is_best_holdout_single": int(is_best_holdout_single), + "is_best_holdout_bilat": int(is_best_holdout_bilat), + **cm_row, + }, optional_cols=epoch_fields) + + # ---- md_warmup progress bar (replaces per-epoch print) -------- + if phase_single == "cd_warmup": + _bar_w = 30 + _filled = int(_bar_w * (epoch + 1) / single_warmup_cd) + _bar = "#" * _filled + "-" * (_bar_w - _filled) + _bar_msg = ( + f" [fold {fold+1}] md_warmup [{_bar}] " + f"{epoch + 1}/{single_warmup_cd} loss={sl_loss:.4f}" + ) + print(f"\r{_bar_msg}", end="", flush=True) + fold_logger.info(_bar_msg) + _prev_phase_single = phase_single + continue # skip normal log block entirely + + if _prev_phase_single == "cd_warmup": + print() # seal the progress bar line + + if args.log_every > 0 and (epoch + 1) % args.log_every == 0: + _epoch_secs = time.time() - _epoch_t0 + hld_auc = target_holdout_single_auc if run_single else bi_auc_h + hld_suffix = f" hld_auc={hld_auc:.4f}" if holdout_loader is not None else "" + + # Human-readable phase progress for console logs. + if phase_single == "tower_warmup": + single_phase_epoch = epoch - single_warmup_cd + 1 + single_phase_total = single_warmup_tower + elif phase_single == "fused_warmup": + single_phase_epoch = epoch - single_warmup_cd - single_warmup_tower + 1 + single_phase_total = single_warmup_fused + else: + single_phase_epoch = main_epoch_single + single_phase_total = main_epochs + + if phase_bilat == "tower_warmup": + bilat_phase_epoch = epoch + 1 + bilat_phase_total = bilat_warmup_tower + elif phase_bilat == "fused_warmup": + bilat_phase_epoch = epoch - bilat_warmup_tower + 1 + bilat_phase_total = bilat_warmup_fused + else: + bilat_phase_epoch = main_epoch_bilat + bilat_phase_total = main_epochs + + if run_single: + if tower_mode == "single": + msg = ( + f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) " + f"[single:{phase_single} {single_phase_epoch}/{single_phase_total}] " + f"fused(acc={cl_acc:.4f},auc={cl_auc:.4f}) " + f"img(acc={cl_acc_img:.4f},auc={cl_auc_img:.4f}) " + f"md(acc={cl_acc_cd:.4f},auc={cl_auc_cd:.4f}) " + f"(best_fused={best_single_auc:.4f} @ep{best_epoch_single})" + f"{hld_suffix}" + ) + else: + msg = ( + f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) " + f"[single:{phase_single} {single_phase_epoch}/{single_phase_total}] " + f"fused(acc={en_acc:.4f},auc={en_auc:.4f}) " + f"img(acc={en_acc_img:.4f},auc={en_auc_img:.4f}) " + f"md(acc={en_acc_cd:.4f},auc={en_auc_cd:.4f}) " + f"(best_fused={best_single_auc:.4f} @ep{best_epoch_single})" + f"{hld_suffix}" + ) + else: + msg = ( + f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) " + f"[bilat:{phase_bilat} {bilat_phase_epoch}/{bilat_phase_total}] " + f"fused(acc={bi_acc:.4f},auc={bi_auc:.4f}) " + f"img(acc={bi_acc_img:.4f},auc={bi_auc_img:.4f}) " + f"md(acc={bi_acc_cd:.4f},auc={bi_auc_cd:.4f}) " + f"(best_bilat={best_bilat_auc:.4f} @ep{best_epoch_bilat})" + f"{hld_suffix}" + ) + print(msg, flush=True) + fold_logger.info(msg) + + _prev_phase_single = phase_single + + fold_logger.close() + + # Override: use final epoch state instead of best-AUC checkpoint + if getattr(args, "use_last_epoch", False): + best_single_state = copy.deepcopy(single.state_dict()) + if run_bilat: + best_bilat_state = copy.deepcopy(bilat.state_dict()) + + if args.save_checkpoints: + if best_single_state is not None: + torch.save(best_single_state, fold_dir / "best_single.pt") + if best_bilat_state is not None: + torch.save(best_bilat_state, fold_dir / "best_bilateral.pt") + if best_holdout_single_state is not None: + torch.save(best_holdout_single_state, fold_dir / "best_holdout_single.pt") + if best_holdout_bilat_state is not None: + torch.save(best_holdout_bilat_state, fold_dir / "best_holdout_bilateral.pt") + + # ---- Save per-epoch per-patient npy tensors ------------------------- + if _epoch_train_pf: + # Use the order from the first epoch (consistent since loader is non-shuffled) + ids_ref = _epoch_train_ids[0] + y_ref = _epoch_train_y[0] + np.save(fold_dir / "train_patient_ids.npy", ids_ref) + np.save(fold_dir / "train_y_true.npy", y_ref) + np.save(fold_dir / "train_probs_fused.npy", np.stack(_epoch_train_pf)) # (n_ep, n_pts, n_cls) + np.save(fold_dir / "train_probs_img.npy", np.stack(_epoch_train_pi)) + np.save(fold_dir / "train_probs_cd.npy", np.stack(_epoch_train_pm)) + if _epoch_val_pf_od: + np.save(fold_dir / "val_y_true_epochs.npy", np.stack(_epoch_val_y)) # (n_ep, N) + np.save(fold_dir / "val_probs_fused_od_epochs.npy", np.stack(_epoch_val_pf_od)) # (n_ep, N, C) + np.save(fold_dir / "val_probs_img_od_epochs.npy", np.stack(_epoch_val_pi_od)) + np.save(fold_dir / "val_probs_cd_od_epochs.npy", np.stack(_epoch_val_pm_od)) + np.save(fold_dir / "val_probs_fused_os_epochs.npy", np.stack(_epoch_val_pf_os)) + np.save(fold_dir / "val_probs_img_os_epochs.npy", np.stack(_epoch_val_pi_os)) + np.save(fold_dir / "val_probs_cd_os_epochs.npy", np.stack(_epoch_val_pm_os)) + if _epoch_val_ids: # only ensemble mode populates this + np.save(fold_dir / "val_patient_ids.npy", _epoch_val_ids[0]) + + # ---- Phase 2: fused head training (ensemble + --fused-head only) ---- + best_fused_auc = -1.0 + best_fused_state: Optional[dict] = None + best_holdout_fused_auc = -1.0 + + if run_fused and best_single_state is not None: + # Revert base to its best val checkpoint, then freeze it. + single.load_state_dict(best_single_state) + for p in single.parameters(): + p.requires_grad_(False) + + fused = FusedEnsembleHT(single, num_classes).to(device) + opt_fused = torch.optim.Adam(fused.eye_scorer.parameters(), lr=args.lr) + fusion_epochs = int(getattr(args, "fusion_epochs", 10)) + + print( + f" [fold {fold+1}] Phase 2: training fusion head " + f"bilat_train_n={len(bilat_train)} fusion_epochs={fusion_epochs}", + flush=True, + ) + + _val_pids_for_store = [str(s["id_1"]) for s in bilat_val] + for fep in range(fusion_epochs): + fu_loss, fu_acc = train_fusion_epoch(fused, train_bilat_loader, opt_fused, device) + y_fu, p_fu = collect_probs_fused(fused, val_loader, device) + fu_auc = _score_arrays(y_fu, p_fu, num_classes)[1] + if pred_store is not None and y_fu.size: + _store_ep = total_single_epochs + fep + pred_store.record(fold, _store_ep, _val_pids_for_store, "bilat_fused", p_fu) + + # Holdout eval (if available) + fu_hld_auc = nan + if holdout_loader is not None: + y_fu_h, p_fu_h = collect_probs_fused(fused, holdout_loader, device) + fu_hld_auc = _score_arrays(y_fu_h, p_fu_h, num_classes)[1] + + is_best_fused = not np.isnan(fu_auc) and fu_auc > best_fused_auc + if is_best_fused: + best_fused_auc = fu_auc + best_fused_state = copy.deepcopy(fused.state_dict()) + fu_acc_val = _score_arrays(y_fu, p_fu, num_classes)[0] + snap_fused, _, _, _ = _tune_and_snap(y_fu, p_fu, fu_acc_val, num_classes, args, args.ece_bins) + + is_best_hld_fused = not np.isnan(fu_hld_auc) and fu_hld_auc > best_holdout_fused_auc + if is_best_hld_fused: + best_holdout_fused_auc = fu_hld_auc + snap_holdout_fused = {"auc": fu_hld_auc, "acc": _score_arrays(y_fu_h, p_fu_h, num_classes)[0]} + + if (fep + 1) % max(1, getattr(args, "log_every", 1)) == 0: + print( + f" [fold {fold+1}] fusion ep{fep+1:>3} " + f"loss={fu_loss:.4f} train_acc={fu_acc:.4f} " + f"val_auc={fu_auc:.4f} hld_auc={fu_hld_auc:.4f}" + f"{' *' if is_best_fused else ''}", + flush=True, + ) + + if best_fused_state is not None: + if args.save_checkpoints: + torch.save(best_fused_state, fold_dir / "best_fused.pt") + print( + f" [fold {fold+1}] BEST " + f"fused_head(acc={snap_fused.get('acc', nan):.4f}," + f"auc={snap_fused.get('auc', nan):.4f}) " + f"kappa={snap_fused.get('kappa', nan):.4f} " + f"F1={snap_fused.get('macro_f1', nan):.4f} " + f"ECE={snap_fused.get('ece', nan):.4f} " + f"holdout_auc={snap_holdout_fused.get('auc', nan):.4f}", + flush=True, + ) + + if run_single: + if tower_mode == "single": + print( + f" [fold {fold+1}] BEST " + f"fused(acc={snap_classic.get('acc', nan):.4f},auc={snap_classic.get('auc', nan):.4f}) " + f"kappa={snap_classic.get('kappa', nan):.4f} " + f"F1={snap_classic.get('macro_f1', nan):.4f} " + f"ECE={snap_classic.get('ece', nan):.4f} @ep{best_epoch_single}", + flush=True, + ) + else: + print( + f" [fold {fold+1}] BEST " + f"ensemble(acc={snap_ensemble.get('acc', nan):.4f},auc={snap_ensemble.get('auc', nan):.4f}) " + f"kappa={snap_ensemble.get('kappa', nan):.4f} " + f"F1={snap_ensemble.get('macro_f1', nan):.4f} " + f"ECE={snap_ensemble.get('ece', nan):.4f} @ep{best_epoch_single}", + flush=True, + ) + else: + print( + f" [fold {fold+1}] BEST " + f"fused(acc={snap_bilat.get('acc', nan):.4f},auc={snap_bilat.get('auc', nan):.4f}) " + f"kappa={snap_bilat.get('kappa', nan):.4f} " + f"F1={snap_bilat.get('macro_f1', nan):.4f} " + f"ECE={snap_bilat.get('ece', nan):.4f} @ep{best_epoch_bilat}", + flush=True, + ) + + # Export best-epoch prediction artifacts. + if run_single and best_single_state is not None: + single.load_state_dict(best_single_state) + if run_bilat and best_bilat_state is not None: + bilateral.load_state_dict(best_bilat_state) + y_en_pe_best = p_en_pe_best = p_en_pe_best_img = p_en_pe_best_md = None + l_en_best = l_en_best_img = l_en_best_md = None + l_cl_best = l_cl_best_img = l_cl_best_md = None + l_en_pe_best = l_en_pe_best_img = l_en_pe_best_md = None + if run_single and tower_mode == "single": + y_cl_best, p_cl_best, p_cl_best_img, p_cl_best_md, \ + l_cl_best, l_cl_best_img, l_cl_best_md = collect_probs_single_components( + single, val_loader, device, aggregate_patient=False, return_logits=True + ) + y_en_best = p_en_best = p_en_best_img = p_en_best_md = None + elif run_single and tower_mode == "ensemble": + y_en_best, p_en_best, p_en_best_img, p_en_best_md, \ + l_en_best, l_en_best_img, l_en_best_md = collect_probs_single_components( + single, val_loader, device, aggregate_patient=True, return_logits=True + ) + y_en_pe_best, p_en_pe_best, p_en_pe_best_img, p_en_pe_best_md, \ + l_en_pe_best, l_en_pe_best_img, l_en_pe_best_md = collect_probs_single_components( + single, val_loader, device, aggregate_patient=False, return_logits=True + ) + y_cl_best = p_cl_best = p_cl_best_img = p_cl_best_md = None + else: + y_cl_best = y_en_best = None + p_cl_best = p_en_best = p_en_best_img = p_en_best_md = None + p_cl_best_img = p_cl_best_md = None + if run_bilat: + y_bi_best, p_bi_best = collect_probs_bilateral(bilateral, val_loader, device) + else: + y_bi_best = p_bi_best = None + + y_fu_best = p_fu_best = None + if run_fused and best_fused_state is not None: + fused.load_state_dict(best_fused_state) + y_fu_best, p_fu_best = collect_probs_fused(fused, val_loader, device) + + return FoldResult( + mode=mode, fold=fold, + best_epoch_single=best_epoch_single, best_epoch_bilat=best_epoch_bilat, + classic_val_auc=snap_classic.get("auc", nan), + classic_val_acc=snap_classic.get("acc", nan), + classic_val_kappa=snap_classic.get("kappa", nan), + classic_val_mcc=snap_classic.get("mcc", nan), + classic_val_f1=snap_classic.get("macro_f1", nan), + classic_val_recall=_sv(snap_classic.get("per_class_recall")), + classic_val_ece=snap_classic.get("ece", nan), + classic_val_threshold=snap_classic.get("threshold", nan), + classic_val_bias=_svf(snap_classic.get("bias")), + classic_val_n=snap_classic.get("n", 0), + ensemble_val_auc=snap_ensemble.get("auc", nan), + ensemble_val_acc=snap_ensemble.get("acc", nan), + ensemble_val_kappa=snap_ensemble.get("kappa", nan), + ensemble_val_mcc=snap_ensemble.get("mcc", nan), + ensemble_val_f1=snap_ensemble.get("macro_f1", nan), + ensemble_val_recall=_sv(snap_ensemble.get("per_class_recall")), + ensemble_val_ece=snap_ensemble.get("ece", nan), + ensemble_val_threshold=snap_ensemble.get("threshold", nan), + ensemble_val_bias=_svf(snap_ensemble.get("bias")), + ensemble_val_n=snap_ensemble.get("n", 0), + bilat_val_auc=snap_bilat.get("auc", nan), + bilat_val_acc=snap_bilat.get("acc", nan), + bilat_val_kappa=snap_bilat.get("kappa", nan), + bilat_val_mcc=snap_bilat.get("mcc", nan), + bilat_val_f1=snap_bilat.get("macro_f1", nan), + bilat_val_recall=_sv(snap_bilat.get("per_class_recall")), + bilat_val_ece=snap_bilat.get("ece", nan), + bilat_val_threshold=snap_bilat.get("threshold", nan), + bilat_val_bias=_svf(snap_bilat.get("bias")), + bilat_val_n=snap_bilat.get("n", 0), + classic_holdout_auc=snap_holdout_single.get("auc", nan) if tower_mode == "single" else nan, + classic_holdout_acc=snap_holdout_single.get("acc", nan) if tower_mode == "single" else nan, + ensemble_holdout_auc=snap_holdout_single.get("auc", nan) if tower_mode == "ensemble" else nan, + ensemble_holdout_acc=snap_holdout_single.get("acc", nan) if tower_mode == "ensemble" else nan, + bilat_holdout_auc=snap_holdout_bilat.get("auc", nan), + bilat_holdout_acc=snap_holdout_bilat.get("acc", nan), + holdout_n=len(holdout_bilat), + single_train_n=len(eye_train), + bilat_train_n=len(bilat_train), + fused_val_auc=snap_fused.get("auc", nan), + fused_val_acc=snap_fused.get("acc", nan), + fused_val_kappa=snap_fused.get("kappa", nan), + fused_val_mcc=snap_fused.get("mcc", nan), + fused_val_f1=snap_fused.get("macro_f1", nan), + fused_val_recall=_sv(snap_fused.get("per_class_recall")), + fused_val_ece=snap_fused.get("ece", nan), + fused_val_threshold=snap_fused.get("threshold", nan), + fused_val_bias=_svf(snap_fused.get("bias")), + fused_val_n=snap_fused.get("n", 0), + fused_holdout_auc=snap_holdout_fused.get("auc", nan), + fused_holdout_acc=snap_holdout_fused.get("acc", nan), + ), FoldArtifacts( + y_true_classic=y_cl_best, probs_classic=p_cl_best, + y_true_ensemble=y_en_best, probs_ensemble=p_en_best, + y_true_bilat=y_bi_best, probs_bilat=p_bi_best, + y_true_fused=y_fu_best, probs_fused=p_fu_best, + probs_ensemble_img=p_en_best_img, + probs_ensemble_md=p_en_best_md, + probs_classic_img=p_cl_best_img, + probs_classic_md=p_cl_best_md, + y_true_ensemble_pereye=y_en_pe_best, + probs_ensemble_pereye=p_en_pe_best, + probs_ensemble_img_pereye=p_en_pe_best_img, + probs_ensemble_md_pereye=p_en_pe_best_md, + logits_ensemble=l_en_best, + logits_ensemble_img=l_en_best_img, + logits_ensemble_md=l_en_best_md, + logits_classic=l_cl_best, + logits_classic_img=l_cl_best_img, + logits_classic_md=l_cl_best_md, + logits_ensemble_pereye=l_en_pe_best, + logits_ensemble_img_pereye=l_en_pe_best_img, + logits_ensemble_md_pereye=l_en_pe_best_md, + ) + + # ------------------------------------------------------------------ + # Summary helpers + # ------------------------------------------------------------------ + + @staticmethod + def _summary(results: list[FoldResult]) -> dict: + def _ms(vals): + v = np.array( + [x for x in vals if x is not None and not np.isnan(float(x))], dtype=float + ) + return (float(np.mean(v)) if v.size else None, float(np.std(v)) if v.size else None) + + out = {} + for label, prefix in [ + ("classic_best_val", "classic_val"), + ("ensemble_best_val", "ensemble_val"), + ("bilat_best_val", "bilat_val"), + ("fused_best_val", "fused_val"), + ]: + sub = {} + for m in ["auc", "acc", "kappa", "mcc", "f1", "ece", "threshold"]: + vals = [getattr(r, f"{prefix}_{m}") for r in results] + mean, std = _ms(vals) + sub[f"{m}_mean"] = mean + if m in ("auc", "f1", "kappa"): + sub[f"{m}_std"] = std + out[label] = sub + + for label, prefix in [ + ("classic_holdout", "classic_holdout"), + ("ensemble_holdout", "ensemble_holdout"), + ("bilat_holdout", "bilat_holdout"), + ("fused_holdout", "fused_holdout"), + ]: + sub = {} + for m in ["auc", "acc"]: + vals = [getattr(r, f"{prefix}_{m}") for r in results] + mean, std = _ms(vals) + sub[f"{m}_mean"] = mean + if m == "auc": + sub[f"{m}_std"] = std + out[label] = sub + + for delta_label, prefix_a, prefix_b in [ + ("delta_ensemble_vs_classic", "classic_val", "ensemble_val"), + ("delta_bilat_vs_ensemble", "ensemble_val", "bilat_val"), + ("delta_fused_vs_ensemble", "ensemble_val", "fused_val"), + ]: + delta = {} + for m in ["auc", "f1", "kappa"]: + pairs = [ + getattr(r, f"{prefix_b}_{m}") - getattr(r, f"{prefix_a}_{m}") + for r in results + if not np.isnan(float(getattr(r, f"{prefix_a}_{m}"))) + and not np.isnan(float(getattr(r, f"{prefix_b}_{m}"))) + ] + delta[f"{m}_mean"] = float(np.mean(pairs)) if pairs else None + delta[f"{m}_std"] = float(np.std(pairs)) if pairs else None + out[delta_label] = delta + + out["n_folds_completed"] = len(results) + out["single_train_mode"] = "eye-level (all OD+OS samples)" + out["bilat_train_mode"] = "patient-level (bilateral only)" + out["eval_note"] = ( + "classic=eye-level SingleEyeHT; " + "ensemble=patient-level SingleEyeHT (OD+OS averaged); " + "bilateral=patient-level BilateralHT; " + "fused=ensemble base + learned logit-level fusion head" + ) + return out + + @staticmethod + def _print_summary(mode: str, s: dict, tower_mode: str | None = None) -> None: + def f(v): + return "nan" if v is None else f"{v:.4f}" + + cv = s["classic_best_val"] + ev = s["ensemble_best_val"] + bv = s["bilat_best_val"] + fv = s["fused_best_val"] + d1 = s["delta_ensemble_vs_classic"] + d2 = s["delta_bilat_vs_ensemble"] + d3 = s["delta_fused_vs_ensemble"] + has_fused = fv["auc_mean"] is not None + + print(f"\n=== Summary [{mode}] — best-epoch val ===") + print(f" {'':26s} {'AUC':>8} {'ACC':>8} {'Kappa':>8} {'F1-mac':>8} {'ECE':>8}") + if tower_mode == "single": + rows = [("single (eye-lvl eval)", cv)] + elif tower_mode == "ensemble": + rows = [("ensemble (pat-lvl eval)", ev)] + if has_fused: + rows.append(("fused_head(pat-lvl eval)", fv)) + elif tower_mode == "bilateral": + rows = [("bilateral (bilat eval)", bv)] + else: + rows = [ + ("classic (eye-lvl eval)", cv), + ("ensemble (pat-lvl eval)", ev), + ("bilateral (bilat eval)", bv), + ] + if has_fused: + rows.append(("fused_head(pat-lvl eval)", fv)) + for label, d in rows: + print( + f" {label:26s} " + f"{f(d['auc_mean']):>8} {f(d['acc_mean']):>8} " + f"{f(d['kappa_mean']):>8} {f(d['f1_mean']):>8} {f(d['ece_mean']):>8}" + ) + if tower_mode is None: + print( + f" {'Δ ensemble−classic':26s} " + f"{f(d1['auc_mean']):>8} {'':>8} " + f"{f(d1['kappa_mean']):>8} {f(d1['f1_mean']):>8}" + ) + print( + f" {'Δ bilateral−ensemble':26s} " + f"{f(d2['auc_mean']):>8} {'':>8} " + f"{f(d2['kappa_mean']):>8} {f(d2['f1_mean']):>8}" + ) + if has_fused and tower_mode in ("ensemble", None): + print( + f" {'Δ fused−ensemble':26s} " + f"{f(d3['auc_mean']):>8} {'':>8} " + f"{f(d3['kappa_mean']):>8} {f(d3['f1_mean']):>8}" + ) + + +# --------------------------------------------------------------------------- +# V2ModeComparator — thin backward-compat shim +# --------------------------------------------------------------------------- + +class V2ModeComparator: + """Backward-compat shim used by run_multifold_v2_modes.py.""" + + @staticmethod + def build_parser() -> argparse.ArgumentParser: + return V2HyperTower.build_parser() + + @staticmethod + def run(cli_args=None) -> Path: + args = V2HyperTower.build_parser().parse_args(cli_args) + return V2HyperTower(args).run() + + +# --------------------------------------------------------------------------- +# Module-level aliases (kept for backward compat; use V2HyperTower directly) +# --------------------------------------------------------------------------- + +def build_parser() -> argparse.ArgumentParser: + return V2HyperTower.build_parser() + + +def run_mode(args) -> Path: + return V2HyperTower(args).run() diff --git a/v3/classes/v3_hypertower.py b/v3/classes/v3_hypertower.py new file mode 100644 index 0000000..a11f9cb --- /dev/null +++ b/v3/classes/v3_hypertower.py @@ -0,0 +1,1250 @@ +"""V3HyperTower — central orchestrator for the V3 pipeline. + +Key differences from V2: + - Proper outer/inner k-fold CV: test = current fold, val = next fold, train = rest. + No pre-carved holdout — every patient appears in test exactly once. + - No checkpoint saving (.pt files). Model states are kept in memory only. + - Test set evaluated each main-phase epoch (logged to epoch CSV only, never used for model selection). + Final test metrics reported in summary use the last-epoch model state. + - Binary-focused defaults (multiclass still supported via --eval-mode multiclass). + - tune_binary_threshold uses Youden's J by default (class-distribution independent). + - tune_multiclass_bias maximises balanced accuracy (class-distribution independent). +""" +from __future__ import annotations + +import argparse +import csv +import json +import time +from pathlib import Path +from types import SimpleNamespace +from typing import Optional + +import numpy as np +import torch + +from v3.classes.croppers import build_image_preprocessor_from_args +from v3.classes.image_loader import CachedImageLoader +from v3.classes.dataset import _ClinicalView # noqa: F401 +from v3.classes.loader_factory import ( + build_balanced_sampler, + filter_bilateral_samples, + filter_eye_samples, + make_loader, +) +from v3.classes.metrics import _score_arrays, _svf, _tune_and_snap +from v3.classes.models import ( + BilateralHT, + FusedEnsembleHT, + SingleEyeHT, + V2ModeComparisonOps, + collect_probs_bilateral, + collect_probs_bilateral_components, + collect_probs_classic, + collect_probs_ensemble, + collect_probs_ensemble_pereye, + collect_probs_eye_level, + collect_probs_fused, + collect_probs_single_components, + train_bilateral_epoch, + train_fusion_epoch, + train_single_epoch, +) +from v3.classes.papila_builders import build_papila_data +from v3.classes.predictions import PredictionStore, head_names_for_mode +from v3.classes.profiles import build_papila_profile +from v3.classes.results import FoldArtifacts, FoldResult, _f, _nan, _sv +from v3.classes.split_manager import PatientFirstSplitManager +from v3.classes.transforms import build_eval_transform +from v3.classes.utils import ( + _drop_mixed_label_patients, + _relabel_mixed_patients_to_max, + choose_device, + seed_everything, +) +from v3.classes.hypertower_logger import HypertowerLogger + + +# --------------------------------------------------------------------------- +# Helpers (unchanged from V2) +# --------------------------------------------------------------------------- + +def _fusion_events(y, pf, pi, pm): + pred_f = pf.argmax(1); pred_i = pi.argmax(1); pred_m = pm.argmax(1) + corr = int(((pred_f == y) & (pred_i != y) & (pred_m != y)).sum()) + err = int(((pred_f != y) & (pred_i == y) & (pred_m == y)).sum()) + return corr, err + + +def _cm_cells(y, p, num_classes): + if not y.size or p.ndim < 2 or p.shape[1] != num_classes: + return {} + pred = p.argmax(1) + if num_classes == 2: + return { + "tn": int(((pred==0)&(y==0)).sum()), "fp": int(((pred==1)&(y==0)).sum()), + "fn": int(((pred==0)&(y==1)).sum()), "tp": int(((pred==1)&(y==1)).sum()), + } + out = {} + for i in range(num_classes): + for j in range(num_classes): + out[f"cm_{i}_{j}"] = int(((y==i)&(pred==j)).sum()) + return out + + +def _save_predictions_csv(fold_dir, eval_mode, y_true, heads, suffix=""): + N = len(y_true) + num_classes = next(p.shape[1] for p in heads.values() if p is not None) + rows = [] + for i in range(N): + true = int(y_true[i]) + row = {"idx": i, "y_true": true} + for head_name, probs in heads.items(): + if probs is None: + continue + pred = int(probs[i].argmax()) + row[f"pred_{head_name}"] = pred + for c in range(num_classes): + row[f"prob_{head_name}_c{c}"] = float(probs[i, c]) + if eval_mode == "binary": + row[f"tp_{head_name}"] = int(pred==1 and true==1) + row[f"fp_{head_name}"] = int(pred==1 and true==0) + row[f"tn_{head_name}"] = int(pred==0 and true==0) + row[f"fn_{head_name}"] = int(pred==0 and true==1) + else: + row[f"correct_{head_name}"] = int(pred==true) + rows.append(row) + if not rows: + return + path = fold_dir / f"predictions{suffix}.csv" + with path.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + +# --------------------------------------------------------------------------- +# V3HyperTower +# --------------------------------------------------------------------------- + +class V3HyperTower: + """V3 orchestrator. Construct with ``V3HyperTower(args)``, call ``.run()``.""" + + @staticmethod + def build_parser() -> argparse.ArgumentParser: + ap = argparse.ArgumentParser( + description=( + "V3 HyperTower — outer/inner k-fold CV (test=current fold). " + "No pre-carved holdout. No checkpoint saving." + ) + ) + ap.add_argument("--image-dir", default="Papila/FundusImages") + ap.add_argument("--clinical-dir", default="Papila/ClinicalData") + ap.add_argument("--label-col", default="Diagnosis") + ap.add_argument("--cat-cols", nargs="*", default=["Gender", "Phakic/Pseudophakic"]) + ap.add_argument("--exclude-cols", nargs="*", default=[]) + ap.add_argument("--eval-mode", choices=["binary", "multiclass"], default="binary") + ap.add_argument( + "--tower-mode", choices=["single", "ensemble", "bilateral", "classic"], + default="ensemble", + ) + ap.add_argument("--n-splits", type=int, default=5) + ap.add_argument("--fold-seed", type=int, default=42) + ap.add_argument( + "--folds", type=int, default=None, + help="Optional cap on number of folds to run.", + ) + ap.add_argument("--epochs", type=int, default=40) + ap.add_argument("--warmup-tower-epochs", type=int, default=None) + ap.add_argument("--warmup-fused-epochs", type=int, default=None) + ap.add_argument("--single-warmup-tower-epochs", type=int, default=None) + ap.add_argument("--single-warmup-fused-epochs", type=int, default=None) + ap.add_argument("--warmup-cd-epochs", type=int, default=0) + ap.add_argument("--bilat-warmup-tower-epochs", type=int, default=None) + ap.add_argument("--bilat-warmup-fused-epochs", type=int, default=None) + ap.add_argument("--batch-size", type=int, default=8) + ap.add_argument("--lr", type=float, default=1e-4) + ap.add_argument("--bcd-prob", type=float, default=0.5) + ap.add_argument("--tower-loss-mode", choices=["bcd", "all"], default="bcd") + ap.add_argument("--backbone", default="refugelike") + ap.add_argument("--freeze-ratio", type=float, default=0.0) + ap.add_argument("--augment", action="store_true") + ap.add_argument("--balanced-sampling", action="store_true") + ap.add_argument("--num-workers", type=int, default=4) + ap.add_argument("--in-memory-cache", action="store_true", default=True) + ap.add_argument("--no-in-memory-cache", action="store_false", dest="in_memory_cache") + ap.add_argument("--cache-workers", type=int, default=4) + ap.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto") + ap.add_argument("--seed", type=int, default=1234) + ap.add_argument("--run-name", default=None) + ap.add_argument("--output-root", default="analysis_data") + # ROI cropping + ap.add_argument("--img-crop-manifest", type=str, default=None) + ap.add_argument("--img-crop-gt", action="store_true") + ap.add_argument("--img-crop-weights", type=str, default=None) + ap.add_argument("--img-crop-normalize", type=str, default="per_image", + choices=["per_image", "imagenet"]) + ap.add_argument("--img-crop-threshold", type=float, default=0.5) + ap.add_argument("--img-crop-tta", action="store_true") + ap.add_argument("--img-crop-scale", type=float, default=2.5) + ap.add_argument("--img-crop-size", type=int, default=224) + ap.add_argument("--img-crop-cache", type=str, default="cache_data/hypertower_crops") + ap.add_argument("--persist-img-crop-cache", action="store_true") + # Architecture + ap.add_argument("--cd-hidden-dim", type=int, default=128) + ap.add_argument("--fusion-dim", type=int, default=256) + ap.add_argument("--bridge-mode", default="fused", + choices=["fused", "image_only", "clinical_only"]) + # Mixed patients + ap.add_argument("--exclude-mixed-patients", dest="exclude_mixed_patients", + action="store_true") + ap.add_argument("--include-mixed-patients", dest="exclude_mixed_patients", + action="store_false") + ap.add_argument("--relabel-mixed-patients-to-max", dest="relabel_mixed_patients_to_max", + action="store_true") + ap.add_argument("--keep-mixed-raw-labels", dest="relabel_mixed_patients_to_max", + action="store_false", help=argparse.SUPPRESS) + ap.set_defaults(exclude_mixed_patients=False, relabel_mixed_patients_to_max=False) + # Tuning + ap.add_argument("--tune-binary-threshold", action="store_true") + ap.add_argument("--tune-multiclass-bias", action="store_true") + ap.add_argument("--ece-bins", type=int, default=10) + ap.add_argument("--log-every", type=int, default=1) + # IOP + ap.add_argument("--iop-corr-method", choices=["ratio", "ols", "lad", "multi"], + default="ratio") + ap.add_argument("--iop-drop-raw", action="store_true", default=False) + # Fused head + ap.add_argument("--fused-head", action="store_true") + ap.add_argument("--fusion-epochs", type=int, default=10) + return ap + + def __init__(self, args) -> None: + self.args = args + self.device = choose_device(args.device) + seed_everything(args.seed) + + print(f"Device: {self.device}", flush=True) + print("Loading PAPILA data...", flush=True) + self.data = build_papila_data( + image_dir=args.image_dir, + clinical_dir=args.clinical_dir, + label_col=args.label_col, + cat_cols=list(args.cat_cols), + n_splits=args.n_splits, + random_seed=args.fold_seed, + iop_corr_method=getattr(args, "iop_corr_method", "ratio"), + iop_drop_raw=getattr(args, "iop_drop_raw", False), + exclude_cols=list(getattr(args, "exclude_cols", []) or []), + ) + print(f"Loaded: {len(self.data.df)} rows feature_dim={self.data.feature_dim}", flush=True) + self.image_preprocessor = build_image_preprocessor_from_args(args) + self.profile_eye = build_papila_profile( + patient_col="Patient ID", label_col=args.label_col, sample_mode="eye" + ) + self.profile_patient = build_papila_profile( + patient_col="Patient ID", label_col=args.label_col, sample_mode="patient" + ) + + def run(self) -> Path: + """Execute the full fold loop.""" + args = self.args + ts = time.strftime("%Y%m%d_%H%M%S") + run_name = args.run_name or f"v3_hypertower_{ts}" + out_dir = Path(args.output_root) / run_name + out_dir.mkdir(parents=True, exist_ok=True) + + mode = args.eval_mode + tower_mode = "single" if args.tower_mode == "classic" else args.tower_mode + df_mode = self.data.df.copy() + + if args.exclude_mixed_patients: + before = df_mode["Patient ID"].nunique() + df_mode, mixed = _drop_mixed_label_patients( + df_mode, patient_col="Patient ID", label_col=args.label_col + ) + print(f"[{mode}] dropped {len(mixed)} mixed-label patients ({before}→{df_mode['Patient ID'].nunique()})", flush=True) + elif args.relabel_mixed_patients_to_max: + df_mode, changed, still_mixed = _relabel_mixed_patients_to_max( + df_mode, patient_col="Patient ID", label_col=args.label_col + ) + print(f"[{mode}] relabeled {changed} mixed-patient rows to max severity", flush=True) + + if mode == "binary": + df_mode = df_mode[df_mode[args.label_col].isin([0, 1])].reset_index(drop=True) + + num_classes = 2 if mode == "binary" else int(df_mode[args.label_col].nunique()) + print(f"\n[{mode}] num_classes={num_classes} rows={len(df_mode)} patients={df_mode['Patient ID'].nunique()}", flush=True) + + split_manager = PatientFirstSplitManager(patient_col="Patient ID", label_col=args.label_col) + split_args = SimpleNamespace( + eval_mode=mode, + n_splits=args.n_splits, + fold_seed=args.fold_seed, + ) + clinical_ns = SimpleNamespace(df=df_mode, label_col=args.label_col) + plans = split_manager.build_plans(clinical=clinical_ns, args=split_args, profile=None) + requested_folds = args.n_splits if args.folds is None else int(args.folds) + n_folds = min(requested_folds, len(plans)) + + tm_dir = out_dir / mode / tower_mode + tm_dir.mkdir(parents=True, exist_ok=True) + fold_results: list[FoldResult] = [] + + profile_eye = build_papila_profile(patient_col="Patient ID", label_col=args.label_col, sample_mode="eye") + profile_patient = build_papila_profile(patient_col="Patient ID", label_col=args.label_col, sample_mode="patient") + + fused_head = getattr(args, "fused_head", False) + _head_names = head_names_for_mode(tower_mode, fused_head=fused_head) + fusion_epochs = int(getattr(args, "fusion_epochs", 10)) if fused_head else 0 + _warmup_tower = int(getattr(args, "single_warmup_tower_epochs", None) or getattr(args, "warmup_tower_epochs", None) or 2) + _warmup_fused = int(getattr(args, "single_warmup_fused_epochs", None) or getattr(args, "warmup_fused_epochs", None) or 2) + _warmup_cd = int(getattr(args, "warmup_cd_epochs", 0)) + _total_epochs = _warmup_cd + _warmup_tower + _warmup_fused + int(args.epochs) + fusion_epochs + + if tower_mode in ("single", "classic"): + _sample_ids = [f"{row['Patient ID']}{row['eyeID']}" for _, row in df_mode.iterrows()] + _y_true = df_mode[args.label_col].tolist() + else: + _pat_df = df_mode.drop_duplicates(subset="Patient ID") + _sample_ids = _pat_df["Patient ID"].astype(str).tolist() + _y_true = _pat_df[args.label_col].tolist() + + pred_store = PredictionStore( + sample_ids=_sample_ids, y_true=_y_true, + head_names=_head_names, n_folds=n_folds, + n_epochs=_total_epochs, n_classes=num_classes, + ) + + image_cache = CachedImageLoader( + enabled=getattr(args, "in_memory_cache", False), + workers=int(getattr(args, "cache_workers", 4)), + ) + + for fold in range(n_folds): + seed_everything(args.seed + fold * 100) + fold_dir = tm_dir / f"fold{fold}" + fold_dir.mkdir(exist_ok=True) + print(f"\n[{mode}:{tower_mode}] fold {fold+1}/{n_folds}", flush=True) + result, artifacts = self._run_fold( + fold=fold, split=plans[fold], mode=mode, data=self.data, + num_classes=num_classes, profile_eye=profile_eye, + profile_patient=profile_patient, fold_dir=fold_dir, + tower_mode=tower_mode, pred_store=pred_store, image_cache=image_cache, + ) + fold_results.append(result) + + # Save val artifacts + if artifacts.y_true_ensemble is not None: + np.save(fold_dir / "y_true.npy", artifacts.y_true_ensemble) + if artifacts.probs_ensemble is not None: + np.save(fold_dir / "probs_fused.npy", artifacts.probs_ensemble) + if artifacts.probs_ensemble_img is not None: + np.save(fold_dir / "probs_img.npy", artifacts.probs_ensemble_img) + if artifacts.probs_ensemble_md is not None: + np.save(fold_dir / "probs_cd.npy", artifacts.probs_ensemble_md) + if artifacts.y_true_classic is not None: + np.save(fold_dir / "y_true.npy", artifacts.y_true_classic) + if artifacts.probs_classic is not None: + np.save(fold_dir / "probs_classic.npy", artifacts.probs_classic) + if artifacts.y_true_ensemble_pereye is not None: + np.save(fold_dir / "y_true_pereye.npy", artifacts.y_true_ensemble_pereye) + if artifacts.probs_ensemble_pereye is not None: + np.save(fold_dir / "probs_fused_pereye.npy", artifacts.probs_ensemble_pereye) + if artifacts.logits_ensemble is not None: + np.save(fold_dir / "logits_fused.npy", artifacts.logits_ensemble) + if artifacts.logits_ensemble_img is not None: + np.save(fold_dir / "logits_img.npy", artifacts.logits_ensemble_img) + if artifacts.logits_ensemble_md is not None: + np.save(fold_dir / "logits_cd.npy", artifacts.logits_ensemble_md) + + # Save test artifacts + if artifacts.y_true_test is not None: + np.save(fold_dir / "test_y_true.npy", artifacts.y_true_test) + if artifacts.probs_test is not None: + np.save(fold_dir / "test_probs_fused.npy", artifacts.probs_test) + if artifacts.probs_test_img is not None: + np.save(fold_dir / "test_probs_img.npy", artifacts.probs_test_img) + if artifacts.probs_test_md is not None: + np.save(fold_dir / "test_probs_cd.npy", artifacts.probs_test_md) + + if pred_store is not None: + pred_store.save(tm_dir / "predictions.npz") + + # Summary + summary = self._summary(fold_results) + self._print_summary(mode, summary, tower_mode=tower_mode) + + # Save fold results CSV + import dataclasses + fold_csv = tm_dir / "fold_results.csv" + rows = [dataclasses.asdict(r) for r in fold_results] + if rows: + with fold_csv.open("w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=list(rows[0].keys())) + writer.writeheader() + writer.writerows(rows) + + # Save summary JSON + if tower_mode in ("single", "classic"): + _test_key = "classic_test" + elif tower_mode == "ensemble": + _test_key = "ensemble_test" + elif tower_mode == "bilateral": + _test_key = "bilat_test" + else: + _test_key = "classic_test" + with (tm_dir / "summary.json").open("w") as f: + json.dump({"mode_summary": {_test_key: summary.get(_test_key, {})}}, + f, indent=2, default=str) + + return out_dir + + def _run_fold( + self, + *, + fold: int, + split, + mode: str, + data, + num_classes: int, + profile_eye, + profile_patient, + fold_dir: Path, + tower_mode: str, + pred_store, + image_cache, + ): + args = self.args + device = self.device + image_preprocessor = self.image_preprocessor + nan = float("nan") + + run_single = tower_mode in ("single", "ensemble") + run_bilat = tower_mode == "bilateral" + run_fused = tower_mode == "ensemble" and getattr(args, "fused_head", False) + + # ---- warmup schedule ------------------------------------------- + global_warmup_tower = getattr(args, "warmup_tower_epochs", None) + global_warmup_fused = getattr(args, "warmup_fused_epochs", None) + single_warmup_tower = int( + getattr(args, "single_warmup_tower_epochs", None) or global_warmup_tower or 2 + ) + single_warmup_fused = int( + getattr(args, "single_warmup_fused_epochs", None) or global_warmup_fused or 2 + ) + bilat_warmup_tower = int( + getattr(args, "bilat_warmup_tower_epochs", None) or global_warmup_tower or 3 + ) + bilat_warmup_fused = int( + getattr(args, "bilat_warmup_fused_epochs", None) or global_warmup_fused or 3 + ) + single_warmup_cd = int(getattr(args, "warmup_cd_epochs", 0)) if run_single else 0 + if not run_single: + single_warmup_tower = single_warmup_fused = 0 + if not run_bilat: + bilat_warmup_tower = bilat_warmup_fused = 0 + # Warmup is meaningless in single-pathway modes — skip it entirely + _bridge_mode = getattr(args, "bridge_mode", "fused") + if _bridge_mode in ("image_only", "clinical_only"): + single_warmup_cd = single_warmup_tower = single_warmup_fused = 0 + bilat_warmup_tower = bilat_warmup_fused = 0 + main_epochs = int(args.epochs) + total_single_epochs = (single_warmup_cd + single_warmup_tower + single_warmup_fused + main_epochs) if run_single else 0 + total_bilat_epochs = (bilat_warmup_tower + bilat_warmup_fused + main_epochs) if run_bilat else 0 + total_epochs = max(total_single_epochs, total_bilat_epochs) + + # ---- samples --------------------------------------------------- + eye_train = filter_eye_samples(profile_eye.build_samples(df=split.train, clinical=data)) + bilat_train = filter_bilateral_samples(profile_patient.build_samples(df=split.train, clinical=data)) + bilat_val = filter_bilateral_samples(profile_patient.build_samples(df=split.val, clinical=data)) + bilat_test = filter_bilateral_samples(profile_patient.build_samples(df=split.test, clinical=data)) if split.test is not None else [] + + if pred_store is not None: + if tower_mode in ("single", "classic"): + train_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in eye_train] + val_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in bilat_val] + else: + train_sids = [str(s["id_1"]) for s in bilat_train] + val_sids = [str(s["id_1"]) for s in bilat_val] + pred_store.set_split(fold, train_sids, "train") + pred_store.set_split(fold, val_sids, "val") + if bilat_test: + test_sids = [str(s["id_1"]) for s in bilat_test] + pred_store.set_split(fold, test_sids, "test") + + if len(bilat_val) == 0: + print(f" [fold {fold+1}] WARNING: no bilateral val samples; skipping fold.", flush=True) + empty = FoldResult( + mode=mode, fold=fold, + best_epoch_single=0, best_epoch_bilat=0, + classic_val_auc=nan, classic_val_acc=nan, classic_val_kappa=nan, + classic_val_mcc=nan, classic_val_f1=nan, classic_val_recall=None, + classic_val_ece=nan, classic_val_threshold=nan, classic_val_bias=None, + classic_val_n=0, + ensemble_val_auc=nan, ensemble_val_acc=nan, ensemble_val_kappa=nan, + ensemble_val_mcc=nan, ensemble_val_f1=nan, ensemble_val_recall=None, + ensemble_val_ece=nan, ensemble_val_threshold=nan, ensemble_val_bias=None, + ensemble_val_n=0, + bilat_val_auc=nan, bilat_val_acc=nan, bilat_val_kappa=nan, + bilat_val_mcc=nan, bilat_val_f1=nan, bilat_val_recall=None, + bilat_val_ece=nan, bilat_val_threshold=nan, bilat_val_bias=None, + bilat_val_n=0, + ensemble_test_auc=nan, ensemble_test_acc=nan, + classic_test_auc=nan, classic_test_acc=nan, + bilat_test_auc=nan, bilat_test_acc=nan, + test_n=0, + single_train_n=len(eye_train), bilat_train_n=len(bilat_train), + ) + return empty, FoldArtifacts( + y_true_classic=None, probs_classic=None, + y_true_ensemble=None, probs_ensemble=None, + y_true_bilat=None, probs_bilat=None, + ) + + # ---- models ---------------------------------------------------- + single = None + bilateral = None + if run_single: + single = SingleEyeHT( + backbone=args.backbone, freeze_ratio=args.freeze_ratio, + augment=args.augment, clinical_data=data, num_classes=num_classes, + cd_hidden_dim=args.cd_hidden_dim, fusion_dim=args.fusion_dim, + bridge_mode=getattr(args, "bridge_mode", "fused"), + ).to(device) + if run_bilat: + bilateral = BilateralHT( + backbone=args.backbone, freeze_ratio=args.freeze_ratio, + augment=args.augment, clinical_data=data, num_classes=num_classes, + cd_hidden_dim=args.cd_hidden_dim, fusion_dim=args.fusion_dim, + ).to(device) + + slots_eye = profile_eye.slot_descriptors() + slots_patient = profile_patient.slot_descriptors() + loader_kw = dict(batch_size=args.batch_size, num_workers=args.num_workers, + image_cache=image_cache) + + # ---- loaders --------------------------------------------------- + use_balanced = bool(getattr(args, "balanced_sampling", False)) + train_single_loader = train_eval_loader = train_bilat_loader = cd_only_loader = None + + if run_single: + single_sampler = build_balanced_sampler(eye_train) if use_balanced else None + train_single_loader = make_loader( + eye_train, slots_eye, image_transform=single.transform, + image_preprocessor=image_preprocessor, shuffle=True, + sampler=single_sampler, **loader_kw, + ) + train_eval_loader = make_loader( + eye_train, slots_eye, image_transform=build_eval_transform(args.backbone), + image_preprocessor=image_preprocessor, shuffle=False, **loader_kw, + ) + if single_warmup_cd > 0: + slots_cd_only = {k: v for k, v in slots_eye.items() if k != "image_1"} + md_sampler = single_sampler if single_sampler is not None else build_balanced_sampler(eye_train) + cd_only_loader = make_loader( + eye_train, slots_cd_only, image_transform=None, + image_preprocessor=None, shuffle=True, sampler=md_sampler, **loader_kw, + ) + if run_bilat: + bilat_sampler = build_balanced_sampler(bilat_train) if use_balanced else None + train_bilat_loader = make_loader( + bilat_train, slots_patient, image_transform=bilateral.transform, + image_preprocessor=image_preprocessor, shuffle=True, + sampler=bilat_sampler, **loader_kw, + ) + elif run_fused: + fused_sampler = build_balanced_sampler(bilat_train) if use_balanced else None + train_bilat_loader = make_loader( + bilat_train, slots_patient, image_transform=single.transform, + image_preprocessor=image_preprocessor, shuffle=True, + sampler=fused_sampler, **loader_kw, + ) + + eval_transform = build_eval_transform(args.backbone) + val_loader = make_loader( + bilat_val, slots_patient, image_transform=eval_transform, + image_preprocessor=image_preprocessor, shuffle=False, **loader_kw, + ) + + # ---- test loader (never touched during training) --------------- + test_loader = None + if bilat_test: + test_loader = make_loader( + bilat_test, slots_patient, image_transform=eval_transform, + image_preprocessor=image_preprocessor, shuffle=False, **loader_kw, + ) + print(f" [fold {fold+1}] test_n={len(bilat_test)} (bilateral patients)", flush=True) + + # ---- prebuild image cache -------------------------------------- + for _ldr in [train_single_loader, train_bilat_loader, val_loader, test_loader]: + if _ldr is not None: + _ldr.dataset.prebuild_image_cache() + + opt_single = torch.optim.Adam(single.parameters(), lr=args.lr) if run_single else None + opt_bilateral = torch.optim.Adam(bilateral.parameters(), lr=args.lr) if run_bilat else None + + # ---- epoch log ------------------------------------------------- + epoch_fields = [ + "fold", "epoch", "phase_single", "phase_bilat", + "main_epoch_single", "main_epoch_bilat", + "single_active", "bilat_active", + "single_train_loss", "single_train_acc", + "classic_val_auc", "classic_val_acc", "classic_val_n", + "ensemble_val_auc", "ensemble_val_acc", "ensemble_val_n", + "bilat_train_loss", "bilat_train_acc", + "bilat_val_auc", "bilat_val_acc", "bilat_val_n", + "classic_val_auc_img", "classic_val_acc_img", + "classic_val_auc_cd", "classic_val_acc_cd", + "classic_val_fe_corr", "classic_val_fe_err", + "ensemble_val_auc_img", "ensemble_val_acc_img", + "ensemble_val_auc_cd", "ensemble_val_acc_cd", + "ensemble_val_fe_corr", "ensemble_val_fe_err", + "bilat_val_auc_img", "bilat_val_acc_img", + "bilat_val_auc_cd", "bilat_val_acc_cd", + "bilat_val_fe_corr", "bilat_val_fe_err", + "train_auc_fused", "train_acc_fused", + "train_auc_img", "train_acc_img", + "train_auc_cd", "train_acc_cd", + "train_fe_corr", "train_fe_err", + "train_n", + "is_best_single", "is_best_bilat", + ] + if num_classes == 2: + _cm_keys = ["tn", "fp", "fn", "tp"] + else: + _cm_keys = [f"cm_{i}_{j}" for i in range(num_classes) for j in range(num_classes)] + for _split in ("classic_val", "ensemble_val", "train"): + for _head in ("fused", "img", "md"): + for _k in _cm_keys: + epoch_fields.append(f"{_split}_{_head}_{_k}") + fold_logger = HypertowerLogger(run_dir=fold_dir) + + # Per-epoch accumulators + _epoch_train_pf: list[np.ndarray] = [] + _epoch_train_pi: list[np.ndarray] = [] + _epoch_train_pm: list[np.ndarray] = [] + _epoch_train_ids: list[np.ndarray] = [] + _epoch_train_y: list[np.ndarray] = [] + _epoch_val_pf_od: list[np.ndarray] = [] + _epoch_val_pi_od: list[np.ndarray] = [] + _epoch_val_pm_od: list[np.ndarray] = [] + _epoch_val_pf_os: list[np.ndarray] = [] + _epoch_val_pi_os: list[np.ndarray] = [] + _epoch_val_pm_os: list[np.ndarray] = [] + _epoch_val_y: list[np.ndarray] = [] + _epoch_val_ids: list[np.ndarray] = [] + + snap_classic: dict = {} + snap_ensemble: dict = {} + snap_bilat: dict = {} + snap_fused: dict = {} + + if run_single: + print( + f" [fold {fold+1}] single_train_n={len(eye_train)} val_n={len(bilat_val)} " + f"test_n={len(bilat_test)} " + f"warmup=md{single_warmup_cd}+twr{single_warmup_tower}+fus{single_warmup_fused} total={total_single_epochs}", + flush=True, + ) + else: + print( + f" [fold {fold+1}] bilat_train_n={len(bilat_train)} val_n={len(bilat_val)} " + f"test_n={len(bilat_test)} " + f"bilat_warmup={bilat_warmup_tower}+{bilat_warmup_fused} total={total_bilat_epochs}", + flush=True, + ) + + _prev_phase_single = "inactive" + + # ================================================================ + # EPOCH LOOP — no test evaluation during training + # ================================================================ + for epoch in range(total_epochs): + _epoch_t0 = time.time() + + # Phase logic + if not run_single: + phase_single, main_epoch_single, single_active = "inactive", 0, False + elif epoch < single_warmup_cd: + phase_single, main_epoch_single, single_active = "cd_warmup", 0, True + elif epoch < single_warmup_cd + single_warmup_tower: + phase_single, main_epoch_single, single_active = "tower_warmup", 0, True + elif epoch < single_warmup_cd + single_warmup_tower + single_warmup_fused: + phase_single, main_epoch_single, single_active = "fused_warmup", 0, True + elif epoch < total_single_epochs: + phase_single = "main" + main_epoch_single = epoch - single_warmup_cd - single_warmup_tower - single_warmup_fused + 1 + single_active = True + else: + phase_single, main_epoch_single, single_active = "done", main_epochs, False + + if not run_bilat: + phase_bilat, main_epoch_bilat, bilat_active = "inactive", 0, False + elif epoch < bilat_warmup_tower: + phase_bilat, main_epoch_bilat, bilat_active = "tower_warmup", 0, True + elif epoch < bilat_warmup_tower + bilat_warmup_fused: + phase_bilat, main_epoch_bilat, bilat_active = "fused_warmup", 0, True + elif epoch < total_bilat_epochs: + phase_bilat = "main" + main_epoch_bilat = epoch - bilat_warmup_tower - bilat_warmup_fused + 1 + bilat_active = True + else: + phase_bilat, main_epoch_bilat, bilat_active = "done", main_epochs, False + + # Training steps + if run_single and single_active: + _active_loader = cd_only_loader if phase_single == "cd_warmup" else train_single_loader + sl_loss, sl_acc = train_single_epoch( + single, _active_loader, opt_single, device, + phase=phase_single, bcd_prob=float(args.bcd_prob), + tower_loss_mode=args.tower_loss_mode, + ) + else: + sl_loss, sl_acc = nan, nan + + if run_bilat and bilat_active: + bl_loss, bl_acc = train_bilateral_epoch( + bilateral, train_bilat_loader, opt_bilateral, device, + phase=phase_bilat, bcd_prob=float(args.bcd_prob), + tower_loss_mode=args.tower_loss_mode, + ) + else: + bl_loss, bl_acc = nan, nan + + _skip_val_eval = (phase_single == "cd_warmup") + + # Val evaluation + if run_single and tower_mode == "single" and not _skip_val_eval: + y_cl, p_cl, p_cl_img, p_cl_cd = collect_probs_single_components( + single, val_loader, device, aggregate_patient=False + ) + cl_acc, cl_auc, cl_n = _score_arrays(y_cl, p_cl, num_classes) + cl_acc_img = float((p_cl_img.argmax(1)==y_cl).mean()) if y_cl.size else nan + cl_acc_cd = float((p_cl_cd.argmax(1) ==y_cl).mean()) if y_cl.size else nan + _, cl_auc_img, _ = _score_arrays(y_cl, p_cl_img, num_classes) + _, cl_auc_cd, _ = _score_arrays(y_cl, p_cl_cd, num_classes) + y_en = np.array([], dtype=np.int64) + p_en = p_en_img = p_en_cd = np.zeros((0, num_classes), dtype=np.float32) + en_acc = en_auc = nan; en_n = 0 + en_acc_img = en_acc_cd = en_auc_img = en_auc_cd = nan + elif run_single and tower_mode == "ensemble" and not _skip_val_eval: + (y_en, _p_en_f_od, _p_en_i_od, _p_en_m_od, + _p_en_f_os, _p_en_i_os, _p_en_m_os, + _en_pat_ids) = collect_probs_ensemble_pereye( + single, val_loader, device, return_ids=True + ) + p_en = 0.5 * (_p_en_f_od + _p_en_f_os) + p_en_img = 0.5 * (_p_en_i_od + _p_en_i_os) + p_en_cd = 0.5 * (_p_en_m_od + _p_en_m_os) + en_acc, en_auc, en_n = _score_arrays(y_en, p_en, num_classes) + en_acc_img = float((p_en_img.argmax(1)==y_en).mean()) if y_en.size else nan + en_acc_cd = float((p_en_cd.argmax(1) ==y_en).mean()) if y_en.size else nan + _, en_auc_img, _ = _score_arrays(y_en, p_en_img, num_classes) + _, en_auc_cd, _ = _score_arrays(y_en, p_en_cd, num_classes) + y_cl = np.array([], dtype=np.int64) + p_cl = p_cl_img = p_cl_cd = np.zeros((0, num_classes), dtype=np.float32) + cl_acc = cl_auc = nan; cl_n = 0 + cl_acc_img = cl_acc_cd = cl_auc_img = cl_auc_cd = nan + else: + y_cl = y_en = np.array([], dtype=np.int64) + p_cl = p_cl_img = p_cl_cd = np.zeros((0, num_classes), dtype=np.float32) + p_en = p_en_img = p_en_cd = np.zeros((0, num_classes), dtype=np.float32) + cl_acc = cl_auc = en_acc = en_auc = nan; cl_n = en_n = 0 + cl_acc_img = cl_acc_cd = en_acc_img = en_acc_cd = nan + cl_auc_img = cl_auc_cd = en_auc_img = en_auc_cd = nan + + if run_bilat and not _skip_val_eval: + y_bi, p_bi, p_bi_img, p_bi_cd = collect_probs_bilateral_components( + bilateral, val_loader, device + ) + bi_acc, bi_auc, bi_n = _score_arrays(y_bi, p_bi, num_classes) + bi_acc_img = float((p_bi_img.argmax(1)==y_bi).mean()) if y_bi.size else nan + bi_acc_cd = float((p_bi_cd.argmax(1) ==y_bi).mean()) if y_bi.size else nan + _, bi_auc_img, _ = _score_arrays(y_bi, p_bi_img, num_classes) + _, bi_auc_cd, _ = _score_arrays(y_bi, p_bi_cd, num_classes) + else: + y_bi = np.array([], dtype=np.int64) + p_bi = np.zeros((0, 0), dtype=np.float32) + bi_acc = bi_auc = nan; bi_n = 0 + bi_acc_img = bi_acc_cd = bi_auc_img = bi_auc_cd = nan + + # Fusion events + cl_fe_corr, cl_fe_err = _fusion_events(y_cl, p_cl, p_cl_img, p_cl_cd) if y_cl.size else (0, 0) + en_fe_corr, en_fe_err = _fusion_events(y_en, p_en, p_en_img, p_en_cd) if y_en.size else (0, 0) + bi_fe_corr, bi_fe_err = (0, 0) + + # Train eval pass + tr_auc_f = tr_acc_f = tr_auc_i = tr_acc_i = tr_auc_m = tr_acc_m = nan + tr_fe_corr = tr_fe_err = tr_n = 0 + y_tr = np.array([], dtype=np.int64) + p_tr_f = p_tr_i = p_tr_m = np.zeros((0, num_classes), dtype=np.float32) + if run_single and train_eval_loader is not None and not _skip_val_eval: + y_tr, p_tr_f, p_tr_i, p_tr_m, tr_ids = collect_probs_eye_level( + single, train_eval_loader, device, return_ids=True + ) + if y_tr.size: + _, tr_auc_f, _ = _score_arrays(y_tr, p_tr_f, num_classes) + tr_acc_f = float((p_tr_f.argmax(1)==y_tr).mean()) + _, tr_auc_i, _ = _score_arrays(y_tr, p_tr_i, num_classes) + tr_acc_i = float((p_tr_i.argmax(1)==y_tr).mean()) + _, tr_auc_m, _ = _score_arrays(y_tr, p_tr_m, num_classes) + tr_acc_m = float((p_tr_m.argmax(1)==y_tr).mean()) + tr_fe_corr, tr_fe_err = _fusion_events(y_tr, p_tr_f, p_tr_i, p_tr_m) + tr_n = int(y_tr.size) + _epoch_train_pf.append(p_tr_f) + _epoch_train_pi.append(p_tr_i) + _epoch_train_pm.append(p_tr_m) + _epoch_train_ids.append(tr_ids) + _epoch_train_y.append(y_tr) + if pred_store is not None: + if tower_mode in ("single", "classic"): + pred_store.record(fold, epoch, tr_ids, "fused", p_tr_f) + pred_store.record(fold, epoch, tr_ids, "img", p_tr_i) + pred_store.record(fold, epoch, tr_ids, "md", p_tr_m) + else: + od_mask = np.array([str(i).endswith("OD") for i in tr_ids]) + os_mask = ~od_mask + od_pids = [str(i)[:-2] for i in tr_ids[od_mask]] + os_pids = [str(i)[:-2] for i in tr_ids[os_mask]] + pred_store.record(fold, epoch, od_pids, "od_fused", p_tr_f[od_mask]) + pred_store.record(fold, epoch, od_pids, "od_img", p_tr_i[od_mask]) + pred_store.record(fold, epoch, od_pids, "od_md", p_tr_m[od_mask]) + pred_store.record(fold, epoch, os_pids, "os_fused", p_tr_f[os_mask]) + pred_store.record(fold, epoch, os_pids, "os_img", p_tr_i[os_mask]) + pred_store.record(fold, epoch, os_pids, "os_md", p_tr_m[os_mask]) + + # Accumulate val per-epoch npy + if run_single and tower_mode == "ensemble" and y_en.size: + _epoch_val_pf_od.append(_p_en_f_od); _epoch_val_pi_od.append(_p_en_i_od) + _epoch_val_pm_od.append(_p_en_m_od); _epoch_val_pf_os.append(_p_en_f_os) + _epoch_val_pi_os.append(_p_en_i_os); _epoch_val_pm_os.append(_p_en_m_os) + _epoch_val_y.append(y_en); _epoch_val_ids.append(_en_pat_ids) + elif run_single and tower_mode == "single" and y_cl.size: + _epoch_val_pf_od.append(p_cl); _epoch_val_pi_od.append(p_cl_img) + _epoch_val_pm_od.append(p_cl_cd); _epoch_val_pf_os.append(p_cl) + _epoch_val_pi_os.append(p_cl_img); _epoch_val_pm_os.append(p_cl_cd) + _epoch_val_y.append(y_cl) + + # Val PredictionStore + if pred_store is not None: + if run_single and tower_mode == "ensemble" and y_en.size: + pred_store.record(fold, epoch, _en_pat_ids, "od_fused", _p_en_f_od) + pred_store.record(fold, epoch, _en_pat_ids, "od_img", _p_en_i_od) + pred_store.record(fold, epoch, _en_pat_ids, "od_md", _p_en_m_od) + pred_store.record(fold, epoch, _en_pat_ids, "os_fused", _p_en_f_os) + pred_store.record(fold, epoch, _en_pat_ids, "os_img", _p_en_i_os) + pred_store.record(fold, epoch, _en_pat_ids, "os_md", _p_en_m_os) + + is_best_single = False + is_best_bilat = False + + # Per-epoch test eval (logged only, never used for model selection) + te_auc = te_acc = float("nan") + if test_loader is not None and phase_single == "main": + if run_single and tower_mode == "single": + _yte, _pte, _, _ = collect_probs_single_components( + single, test_loader, device, aggregate_patient=False) + elif run_single and tower_mode == "ensemble": + _yte, _pte, _, _ = collect_probs_single_components( + single, test_loader, device, aggregate_patient=True) + elif run_bilat: + _yte, _pte, _, _ = collect_probs_bilateral_components( + bilateral, test_loader, device) + else: + _yte = _pte = None + if _yte is not None and _yte.size and len(np.unique(_yte)) > 1: + te_auc = float(_score_arrays(_yte, _pte, num_classes)[1]) + te_acc = float((_pte.argmax(1) == _yte).mean()) + + # Confusion matrix cells + def _prefixed_cm(prefix, y, pf, pi, pm): + out = {} + for head, p in (("fused", pf), ("img", pi), ("md", pm)): + for k, v in _cm_cells(y, p, num_classes).items(): + out[f"{prefix}_{head}_{k}"] = v + return out + + cm_row = {} + cm_row.update(_prefixed_cm("classic_val", y_cl, p_cl, p_cl_img, p_cl_cd)) + cm_row.update(_prefixed_cm("ensemble_val", y_en, p_en, p_en_img, p_en_cd)) + cm_row.update(_prefixed_cm("train", y_tr, p_tr_f, p_tr_i, p_tr_m)) + + fold_logger.write_epoch_row({ + "fold": fold, "epoch": epoch + 1, + "phase_single": phase_single, "phase_bilat": phase_bilat, + "main_epoch_single": main_epoch_single, "main_epoch_bilat": main_epoch_bilat, + "single_active": int(single_active), "bilat_active": int(bilat_active), + "single_train_loss": _f(sl_loss), "single_train_acc": _f(sl_acc), + "classic_val_auc": _f(cl_auc), "classic_val_acc": _f(cl_acc), "classic_val_n": cl_n, + "ensemble_val_auc": _f(en_auc), "ensemble_val_acc": _f(en_acc), "ensemble_val_n": en_n, + "bilat_train_loss": _f(bl_loss), "bilat_train_acc": _f(bl_acc), + "bilat_val_auc": _f(bi_auc), "bilat_val_acc": _f(bi_acc), "bilat_val_n": bi_n, + "classic_val_auc_img": _f(cl_auc_img), "classic_val_acc_img": _f(cl_acc_img), + "classic_val_auc_cd": _f(cl_auc_cd), "classic_val_acc_cd": _f(cl_acc_cd), + "classic_val_fe_corr": cl_fe_corr, "classic_val_fe_err": cl_fe_err, + "ensemble_val_auc_img": _f(en_auc_img), "ensemble_val_acc_img": _f(en_acc_img), + "ensemble_val_auc_cd": _f(en_auc_cd), "ensemble_val_acc_cd": _f(en_acc_cd), + "ensemble_val_fe_corr": en_fe_corr, "ensemble_val_fe_err": en_fe_err, + "bilat_val_auc_img": _f(bi_auc_img), "bilat_val_acc_img": _f(bi_acc_img), + "bilat_val_auc_cd": _f(bi_auc_cd), "bilat_val_acc_cd": _f(bi_acc_cd), + "bilat_val_fe_corr": bi_fe_corr, "bilat_val_fe_err": bi_fe_err, + "train_auc_fused": _f(tr_auc_f), "train_acc_fused": _f(tr_acc_f), + "train_auc_img": _f(tr_auc_i), "train_acc_img": _f(tr_acc_i), + "train_auc_cd": _f(tr_auc_m), "train_acc_cd": _f(tr_acc_m), + "train_fe_corr": tr_fe_corr, "train_fe_err": tr_fe_err, + "train_n": tr_n, + "is_best_single": int(is_best_single), + "is_best_bilat": int(is_best_bilat), + "test_auc": _f(te_auc), "test_acc": _f(te_acc), + **cm_row, + }, optional_cols=epoch_fields) + + # md_warmup progress bar + if phase_single == "cd_warmup": + _bar_w = 30 + _filled = int(_bar_w * (epoch + 1) / single_warmup_cd) + _bar = "#" * _filled + "-" * (_bar_w - _filled) + msg = f" [fold {fold+1}] md_warmup [{_bar}] {epoch+1}/{single_warmup_cd} loss={sl_loss:.4f}" + print(f"\r{msg}", end="", flush=True) + fold_logger.info(msg) + _prev_phase_single = phase_single + continue + + if _prev_phase_single == "cd_warmup": + print() + + if args.log_every > 0 and (epoch + 1) % args.log_every == 0: + _epoch_secs = time.time() - _epoch_t0 + if tower_mode == "ensemble": + msg = f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={en_auc:.4f} acc={en_acc:.4f}" + elif tower_mode == "bilateral": + msg = f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={bi_auc:.4f} acc={bi_acc:.4f}" + else: + msg = f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={cl_auc:.4f} acc={cl_acc:.4f}" + print(msg, flush=True) + fold_logger.info(msg) + + _prev_phase_single = phase_single + + fold_logger.close() + + # No checkpoint saving in V3. + + # Save per-epoch npy tensors + if _epoch_train_pf: + ids_ref = _epoch_train_ids[0]; y_ref = _epoch_train_y[0] + np.save(fold_dir / "train_patient_ids.npy", ids_ref) + np.save(fold_dir / "train_y_true.npy", y_ref) + np.save(fold_dir / "train_probs_fused.npy", np.stack(_epoch_train_pf)) + np.save(fold_dir / "train_probs_img.npy", np.stack(_epoch_train_pi)) + np.save(fold_dir / "train_probs_cd.npy", np.stack(_epoch_train_pm)) + if _epoch_val_pf_od: + np.save(fold_dir / "val_y_true_epochs.npy", np.stack(_epoch_val_y)) + np.save(fold_dir / "val_probs_fused_od_epochs.npy", np.stack(_epoch_val_pf_od)) + np.save(fold_dir / "val_probs_img_od_epochs.npy", np.stack(_epoch_val_pi_od)) + np.save(fold_dir / "val_probs_cd_od_epochs.npy", np.stack(_epoch_val_pm_od)) + np.save(fold_dir / "val_probs_fused_os_epochs.npy", np.stack(_epoch_val_pf_os)) + np.save(fold_dir / "val_probs_img_os_epochs.npy", np.stack(_epoch_val_pi_os)) + np.save(fold_dir / "val_probs_cd_os_epochs.npy", np.stack(_epoch_val_pm_os)) + if _epoch_val_ids: + np.save(fold_dir / "val_patient_ids.npy", _epoch_val_ids[0]) + + # ================================================================ + # Phase 2: fused head (ensemble only) + # ================================================================ + snap_holdout_fused: dict = {} + + if run_fused and single is not None: + for p in single.parameters(): + p.requires_grad_(False) + fused = FusedEnsembleHT(single, num_classes).to(device) + opt_fused = torch.optim.Adam(fused.eye_scorer.parameters(), lr=args.lr) + fusion_epochs = int(getattr(args, "fusion_epochs", 10)) + print( + f" [fold {fold+1}] Phase 2: fusion head bilat_train_n={len(bilat_train)} epochs={fusion_epochs}", + flush=True, + ) + _val_pids_for_store = [str(s["id_1"]) for s in bilat_val] + for fep in range(fusion_epochs): + fu_loss, fu_acc = train_fusion_epoch(fused, train_bilat_loader, opt_fused, device) + y_fu, p_fu = collect_probs_fused(fused, val_loader, device) + fu_auc = _score_arrays(y_fu, p_fu, num_classes)[1] + if pred_store is not None and y_fu.size: + pred_store.record(fold, total_single_epochs + fep, _val_pids_for_store, "bilat_fused", p_fu) + if y_fu.size and not np.isnan(fu_auc): + fu_acc_val = _score_arrays(y_fu, p_fu, num_classes)[0] + snap_fused, _, _, _ = _tune_and_snap(y_fu, p_fu, fu_acc_val, num_classes, args, args.ece_bins) + if (fep + 1) % max(1, args.log_every) == 0: + print( + f" [fold {fold+1}] fusion ep{fep+1:>3} loss={fu_loss:.4f} " + f"val_auc={fu_auc:.4f}", + flush=True, + ) + # No checkpoint saving for fused head either. + + # ================================================================ + # Post-training: evaluate ONCE on val (final-epoch state) then test set + # ================================================================ + + # Val artifacts + y_en_pe_best = p_en_pe_best = p_en_pe_best_img = p_en_pe_best_md = None + l_en_best = l_en_best_img = l_en_best_md = None + l_cl_best = l_cl_best_img = l_cl_best_md = None + l_en_pe_best = l_en_pe_best_img = l_en_pe_best_md = None + + if run_single and tower_mode == "single": + y_cl_best, p_cl_best, p_cl_best_img, p_cl_best_md, \ + l_cl_best, l_cl_best_img, l_cl_best_md = collect_probs_single_components( + single, val_loader, device, aggregate_patient=False, return_logits=True + ) + y_en_best = p_en_best = p_en_best_img = p_en_best_md = None + elif run_single and tower_mode == "ensemble": + y_en_best, p_en_best, p_en_best_img, p_en_best_md, \ + l_en_best, l_en_best_img, l_en_best_md = collect_probs_single_components( + single, val_loader, device, aggregate_patient=True, return_logits=True + ) + y_en_pe_best, p_en_pe_best, p_en_pe_best_img, p_en_pe_best_md, \ + l_en_pe_best, l_en_pe_best_img, l_en_pe_best_md = collect_probs_single_components( + single, val_loader, device, aggregate_patient=False, return_logits=True + ) + y_cl_best = p_cl_best = p_cl_best_img = p_cl_best_md = None + else: + y_cl_best = y_en_best = None + p_cl_best = p_en_best = p_en_best_img = p_en_best_md = None + p_cl_best_img = p_cl_best_md = None + if run_bilat: + y_bi_best, p_bi_best = collect_probs_bilateral(bilateral, val_loader, device) + else: + y_bi_best = p_bi_best = None + + # Compute val snaps from final-epoch model state + if run_single and tower_mode == "single" and y_cl_best is not None: + snap_cl, _, _, _ = _tune_and_snap(y_cl_best, p_cl_best, float((p_cl_best.argmax(1) == y_cl_best).mean()), num_classes, args, args.ece_bins) + snap_classic = snap_cl + elif run_single and tower_mode == "ensemble" and y_en_best is not None: + snap_en, _, _, _ = _tune_and_snap(y_en_best, p_en_best, float((p_en_best.argmax(1) == y_en_best).mean()), num_classes, args, args.ece_bins) + snap_ensemble = snap_en + if run_bilat and y_bi_best is not None: + snap_bi, _, _, _ = _tune_and_snap(y_bi_best, p_bi_best, float((p_bi_best.argmax(1) == y_bi_best).mean()), num_classes, args, args.ece_bins) + snap_bilat = snap_bi + + y_fu_best = p_fu_best = None + if run_fused and single is not None: + y_fu_best, p_fu_best = collect_probs_fused(fused, val_loader, device) + + # Test set evaluation (once, never seen during training) + snap_test: dict = {} + y_test_out = p_test_out = p_test_img_out = p_test_md_out = None + + if test_loader is not None: + if run_single and tower_mode == "ensemble": + y_test_out, p_test_out, p_test_img_out, p_test_md_out = collect_probs_single_components( + single, test_loader, device, aggregate_patient=True + ) + elif run_single and tower_mode == "single": + y_test_out, p_test_out, p_test_img_out, p_test_md_out = collect_probs_single_components( + single, test_loader, device, aggregate_patient=False + ) + elif run_bilat: + y_test_out, p_test_out, _, _ = collect_probs_bilateral_components( + bilateral, test_loader, device + ) + if y_test_out is not None and y_test_out.size: + test_acc_raw = float((p_test_out.argmax(1) == y_test_out).mean()) + snap_test, _, _, _ = _tune_and_snap( + y_test_out, p_test_out, test_acc_raw, num_classes, args, args.ece_bins + ) + print( + f" [fold {fold+1}] TEST " + f"auc={snap_test.get('auc', nan):.4f} " + f"acc={snap_test.get('acc', nan):.4f} " + f"kappa={snap_test.get('kappa', nan):.4f} " + f"f1={snap_test.get('macro_f1', nan):.4f} " + f"ece={snap_test.get('ece', nan):.4f} " + f"n={snap_test.get('n', 0)}", + flush=True, + ) + _save_predictions_csv( + fold_dir, mode, y_test_out, + {"fused": p_test_out, "img": p_test_img_out, "md": p_test_md_out}, + suffix="_test", + ) + else: + print(f" [fold {fold+1}] WARNING: no test samples for this fold.", flush=True) + + test_n = snap_test.get("n", 0) + + return FoldResult( + mode=mode, fold=fold, + best_epoch_single=0, best_epoch_bilat=0, + classic_val_auc=snap_classic.get("auc", nan), + classic_val_acc=snap_classic.get("acc", nan), + classic_val_kappa=snap_classic.get("kappa", nan), + classic_val_mcc=snap_classic.get("mcc", nan), + classic_val_f1=snap_classic.get("macro_f1", nan), + classic_val_recall=_sv(snap_classic.get("per_class_recall")), + classic_val_ece=snap_classic.get("ece", nan), + classic_val_threshold=snap_classic.get("threshold", nan), + classic_val_bias=_svf(snap_classic.get("bias")), + classic_val_n=snap_classic.get("n", 0), + ensemble_val_auc=snap_ensemble.get("auc", nan), + ensemble_val_acc=snap_ensemble.get("acc", nan), + ensemble_val_kappa=snap_ensemble.get("kappa", nan), + ensemble_val_mcc=snap_ensemble.get("mcc", nan), + ensemble_val_f1=snap_ensemble.get("macro_f1", nan), + ensemble_val_recall=_sv(snap_ensemble.get("per_class_recall")), + ensemble_val_ece=snap_ensemble.get("ece", nan), + ensemble_val_threshold=snap_ensemble.get("threshold", nan), + ensemble_val_bias=_svf(snap_ensemble.get("bias")), + ensemble_val_n=snap_ensemble.get("n", 0), + bilat_val_auc=snap_bilat.get("auc", nan), + bilat_val_acc=snap_bilat.get("acc", nan), + bilat_val_kappa=snap_bilat.get("kappa", nan), + bilat_val_mcc=snap_bilat.get("mcc", nan), + bilat_val_f1=snap_bilat.get("macro_f1", nan), + bilat_val_recall=_sv(snap_bilat.get("per_class_recall")), + bilat_val_ece=snap_bilat.get("ece", nan), + bilat_val_threshold=snap_bilat.get("threshold", nan), + bilat_val_bias=_svf(snap_bilat.get("bias")), + bilat_val_n=snap_bilat.get("n", 0), + ensemble_test_auc=snap_test.get("auc", nan) if tower_mode == "ensemble" else nan, + ensemble_test_acc=snap_test.get("acc", nan) if tower_mode == "ensemble" else nan, + ensemble_test_kappa=snap_test.get("kappa", nan) if tower_mode == "ensemble" else nan, + ensemble_test_f1=snap_test.get("macro_f1", nan) if tower_mode == "ensemble" else nan, + ensemble_test_ece=snap_test.get("ece", nan) if tower_mode == "ensemble" else nan, + classic_test_auc=snap_test.get("auc", nan) if tower_mode == "single" else nan, + classic_test_acc=snap_test.get("acc", nan) if tower_mode == "single" else nan, + classic_test_kappa=snap_test.get("kappa", nan) if tower_mode == "single" else nan, + classic_test_f1=snap_test.get("macro_f1", nan) if tower_mode == "single" else nan, + classic_test_ece=snap_test.get("ece", nan) if tower_mode == "single" else nan, + bilat_test_auc=snap_test.get("auc", nan) if tower_mode == "bilateral" else nan, + bilat_test_acc=snap_test.get("acc", nan) if tower_mode == "bilateral" else nan, + bilat_test_kappa=snap_test.get("kappa", nan) if tower_mode == "bilateral" else nan, + bilat_test_f1=snap_test.get("macro_f1", nan) if tower_mode == "bilateral" else nan, + bilat_test_ece=snap_test.get("ece", nan) if tower_mode == "bilateral" else nan, + test_n=test_n, + single_train_n=len(eye_train), + bilat_train_n=len(bilat_train), + fused_val_auc=snap_fused.get("auc", nan), + fused_val_acc=snap_fused.get("acc", nan), + fused_val_kappa=snap_fused.get("kappa", nan), + fused_val_mcc=snap_fused.get("mcc", nan), + fused_val_f1=snap_fused.get("macro_f1", nan), + fused_val_recall=_sv(snap_fused.get("per_class_recall")), + fused_val_ece=snap_fused.get("ece", nan), + fused_val_threshold=snap_fused.get("threshold", nan), + fused_val_bias=_svf(snap_fused.get("bias")), + fused_val_n=snap_fused.get("n", 0), + fused_test_auc=nan, fused_test_acc=nan, + ), FoldArtifacts( + y_true_classic=y_cl_best, probs_classic=p_cl_best, + y_true_ensemble=y_en_best, probs_ensemble=p_en_best, + y_true_bilat=y_bi_best, probs_bilat=p_bi_best, + y_true_fused=y_fu_best, probs_fused=p_fu_best, + probs_ensemble_img=p_en_best_img, + probs_ensemble_md=p_en_best_md, + probs_classic_img=p_cl_best_img, + probs_classic_md=p_cl_best_md, + y_true_ensemble_pereye=y_en_pe_best, + probs_ensemble_pereye=p_en_pe_best, + probs_ensemble_img_pereye=p_en_pe_best_img, + probs_ensemble_md_pereye=p_en_pe_best_md, + logits_ensemble=l_en_best, + logits_ensemble_img=l_en_best_img, + logits_ensemble_md=l_en_best_md, + logits_classic=l_cl_best, + logits_classic_img=l_cl_best_img, + logits_classic_md=l_cl_best_md, + logits_ensemble_pereye=l_en_pe_best, + logits_ensemble_img_pereye=l_en_pe_best_img, + logits_ensemble_md_pereye=l_en_pe_best_md, + y_true_test=y_test_out, + probs_test=p_test_out, + probs_test_img=p_test_img_out, + probs_test_md=p_test_md_out, + ) + + @staticmethod + def _summary(results: list[FoldResult]) -> dict: + def _ms(vals): + v = np.array([x for x in vals if x is not None and not np.isnan(float(x))], dtype=float) + return (float(np.mean(v)) if v.size else None, float(np.std(v)) if v.size else None) + + out = {} + for label, prefix in [ + ("classic_best_val", "classic_val"), + ("ensemble_best_val", "ensemble_val"), + ("bilat_best_val", "bilat_val"), + ("fused_best_val", "fused_val"), + ]: + sub = {} + for m in ["auc", "acc", "kappa", "mcc", "f1", "ece", "threshold"]: + vals = [getattr(r, f"{prefix}_{m}") for r in results] + mean, std = _ms(vals) + sub[f"{m}_mean"] = mean + if m in ("auc", "f1", "kappa"): + sub[f"{m}_std"] = std + out[label] = sub + + for label, prefix in [ + ("ensemble_test", "ensemble_test"), + ("classic_test", "classic_test"), + ("bilat_test", "bilat_test"), + ]: + sub = {} + for m in ["auc", "acc", "kappa", "f1", "ece"]: + vals = [getattr(r, f"{prefix}_{m}") for r in results] + mean, std = _ms(vals) + sub[f"{m}_mean"] = mean + sub[f"{m}_std"] = std + out[label] = sub + + for delta_label, prefix_a, prefix_b in [ + ("delta_ensemble_vs_classic", "classic_val", "ensemble_val"), + ("delta_bilat_vs_ensemble", "ensemble_val", "bilat_val"), + ]: + delta = {} + for m in ["auc", "f1", "kappa"]: + pairs = [ + getattr(r, f"{prefix_b}_{m}") - getattr(r, f"{prefix_a}_{m}") + for r in results + if not np.isnan(float(getattr(r, f"{prefix_a}_{m}"))) + and not np.isnan(float(getattr(r, f"{prefix_b}_{m}"))) + ] + delta[f"{m}_mean"] = float(np.mean(pairs)) if pairs else None + delta[f"{m}_std"] = float(np.std(pairs)) if pairs else None + out[delta_label] = delta + + out["n_folds_completed"] = len(results) + return out + + @staticmethod + def _print_summary(mode: str, s: dict, tower_mode: str | None = None) -> None: + def f(v): + return " nan " if v is None else f"{v:.4f}" + def fsd(mean, std): + if mean is None: return " nan " + if std is None: return f"{mean:.4f} " + return f"{mean:.4f}±{std:.4f}" + + # Resolve test key + if tower_mode in ("single", "classic"): + test_key = "classic_test" + elif tower_mode == "ensemble": + test_key = "ensemble_test" + elif tower_mode == "bilateral": + test_key = "bilat_test" + else: + test_key = "classic_test" + + td = s.get(test_key, {}) + print(f"\n=== Summary [{mode}] ===") + print(f" {'':26s} {'AUC':>16} {'ACC':>16} {'Kappa':>16} {'F1-mac':>16} {'ECE':>16}") + print(f" {'Test':26s} " + f"{fsd(td.get('auc_mean'), td.get('auc_std')):>16} " + f"{fsd(td.get('acc_mean'), td.get('acc_std')):>16} " + f"{fsd(td.get('kappa_mean'), td.get('kappa_std')):>16} " + f"{fsd(td.get('f1_mean'), td.get('f1_std')):>16} " + f"{fsd(td.get('ece_mean'), td.get('ece_std')):>16}") + print() diff --git a/v3/scripts/__init__.py b/v3/scripts/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/v3/scripts/main/__init__.py b/v3/scripts/main/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/v3/scripts/main/phase_1_papila_reproduce.py b/v3/scripts/main/phase_1_papila_reproduce.py new file mode 100644 index 0000000..cd3f319 --- /dev/null +++ b/v3/scripts/main/phase_1_papila_reproduce.py @@ -0,0 +1,847 @@ +#!/usr/bin/env python +""" +Phase 1: Reproduce PAPILA paper baseline results. + +Runs classical ML classifiers on clinical data and/or a CNN on fundus images, +using the same 5-fold stratified CV scheme as the original paper. + +Classifiers available (enable with flags): + --knn K-Nearest Neighbours + --rf Random Forest + --svm Support Vector Machine + --logreg Logistic Regression + --cnn CNN (specify backbone with --backbone) + +Clinical data loader is self-contained here — tweak the ClinicalLoader class +below without touching anything in the main v3 classes. This lets you match +the paper's preprocessing (or lack thereof) independently. + +Usage examples: + # All classical + our default clinical preprocessing + python -m v3.scripts.main.phase_1_papila_reproduce --knn --rf --svm --logreg + + # Match paper more closely (no IOP correction, no feature engineering) + python -m v3.scripts.main.phase_1_papila_reproduce --knn --rf --svm --logreg \ + --no-iop-corr --keep-raw-iop --no-cat-cols + + # CNN only, refugelike backbone + python -m v3.scripts.main.phase_1_papila_reproduce --cnn --backbone refugelike + + # CNN with paper backbones + python -m v3.scripts.main.phase_1_papila_reproduce --cnn \ + --backbone resnet50 --backbone-pretrained + + # Everything + python -m v3.scripts.main.phase_1_papila_reproduce --knn --rf --svm --logreg \ + --cnn --backbone refugelike --output-dir analysis_data/papila_reproduce +""" +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path +from typing import List, Optional, Tuple + +import numpy as np +import pandas as pd +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +from sklearn.base import clone +from sklearn.ensemble import RandomForestClassifier +from sklearn.linear_model import LogisticRegression +from sklearn.metrics import accuracy_score, roc_auc_score, roc_curve +from sklearn.model_selection import StratifiedKFold +from sklearn.model_selection import StratifiedGroupKFold +from sklearn.neighbors import KNeighborsClassifier +from sklearn.pipeline import Pipeline +from sklearn.preprocessing import StandardScaler +from sklearn.svm import SVC + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + + +# --------------------------------------------------------------------------- +# Standalone clinical data loader +# --------------------------------------------------------------------------- +# This loader is intentionally independent of the v3 clinical data pipeline +# so that we can tune preprocessing to match the original PAPILA paper without +# modifying the production classes. + +class ClinicalLoader: + """ + Standalone loader for PAPILA clinical data. + + Parameters + ---------- + clinical_dir : str + Path to Papila/ClinicalData directory. + label_col : str + Column containing ground-truth labels (default: "Diagnosis"). + cat_cols : list[str] | None + Categorical columns to one-hot encode. Pass [] to disable. + exclude_cols : list[str] | None + Extra columns to drop from the feature matrix. + iop_corr : bool + Apply Perkins→Pneumatic IOP correction (ratio method). Default True. + keep_raw_iop : bool + If True, keep Perkins IOP column alongside corrected IOP. Default False. + drop_suspects : bool + Drop Diagnosis==2 (Suspect) rows — binary task only. Default True. + """ + + # PAPILA column names + _PATIENT_COL = "Patient ID" + _EYE_COL = "eyeID" + + # These are always excluded from feature matrix + _ALWAYS_EXCLUDE = {"ID", "Patient ID", "eyeID", "Diagnosis", "VF_MD"} + + def __init__( + self, + clinical_dir: str = "Papila/ClinicalData", + label_col: str = "Diagnosis", + cat_cols: Optional[List[str]] = None, + exclude_cols: Optional[List[str]] = None, + iop_corr: bool = True, + keep_raw_iop: bool = False, + drop_suspects: bool = True, + ) -> None: + self.clinical_dir = Path(clinical_dir) + self.label_col = label_col + self.cat_cols = cat_cols if cat_cols is not None else ["Gender", "Phakic/Pseudophakic"] + self.exclude_cols = set(exclude_cols or []) + self.iop_corr = iop_corr + self.keep_raw_iop = keep_raw_iop + self.drop_suspects = drop_suspects + self._df: Optional[pd.DataFrame] = None + + @property + def df(self) -> pd.DataFrame: + if self._df is None: + self._df = self._load() + return self._df + + def _load(self) -> pd.DataFrame: + # Load OD and OS files (xlsx, header on row 1) + od_path = self.clinical_dir / "patient_data_od.xlsx" + os_path = self.clinical_dir / "patient_data_os.xlsx" + frames = [] + for path, eye in ((od_path, "OD"), (os_path, "OS")): + if not path.exists(): + raise FileNotFoundError(f"Clinical data file not found: {path}") + df = pd.read_excel(path, header=1) + df["eyeID"] = eye + # Normalise patient ID: '#002' → 2 + id_col = "Patient ID" if "Patient ID" in df.columns else "ID" + df["Patient ID"] = ( + df[id_col].astype(str).str.extract(r"(\d+)")[0].astype(int) + ) + frames.append(df) + df = pd.concat(frames, ignore_index=True) + + # IOP: average Perkins and Pneumatic when both present, else use whichever is available + has_perk = "Perkins" in df.columns + has_pneu = "Pneumatic" in df.columns + if has_perk and has_pneu: + both = df["Perkins"].notna() & df["Pneumatic"].notna() + df["IOP_raw"] = df["Perkins"].copy() + df.loc[both, "IOP_raw"] = (df.loc[both, "Perkins"] + df.loc[both, "Pneumatic"]) / 2 + df.loc[~both & df["Pneumatic"].notna(), "IOP_raw"] = df.loc[~both & df["Pneumatic"].notna(), "Pneumatic"] + df = df.drop(columns=["Perkins", "Pneumatic"]) + elif has_pneu: + df = df.rename(columns={"Pneumatic": "IOP_raw"}) + elif has_perk: + df = df.rename(columns={"Perkins": "IOP_raw"}) + + if self.drop_suspects: + df = df[df[self.label_col] != 2].reset_index(drop=True) + + return df + + def feature_matrix(self) -> Tuple[np.ndarray, np.ndarray, List[str]]: + """ + Returns (X, y, feature_names, patient_ids) at the eye level. + + Each eye is one row. Patient IDs are returned so that CV can split + at the patient level (preventing OD/OS leakage across folds). + + X shape: (n_eyes, n_features) + y: binary labels (0=Normal, 1=Glaucoma) + patient_ids: (n_eyes,) int array — group labels for GroupKFold + """ + df = self.df.copy() + exclude = self._ALWAYS_EXCLUDE | self.exclude_cols + + numeric_cols = [ + c for c in df.columns + if c not in exclude and c not in self.cat_cols + and c not in ("eyeID",) + and pd.to_numeric(df[c], errors="coerce").notna().any() + ] + for col in numeric_cols: + df[col] = pd.to_numeric(df[col], errors="coerce") + df[col] = df[col].fillna(df[col].median()) + + X_num = df[numeric_cols].values.astype(np.float32) + names = list(numeric_cols) + + parts = [X_num] + cat_present = [c for c in (self.cat_cols or []) if c in df.columns] + if cat_present: + dummies = pd.get_dummies(df[cat_present].astype("category"), + drop_first=False) + parts.append(dummies.values.astype(np.float32)) + names.extend(list(dummies.columns)) + + X = np.concatenate(parts, axis=1) + y = (df[self.label_col].values.astype(int) == 1).astype(int) + patient_ids = df["Patient ID"].values.astype(int) + return X, y, names, patient_ids + + +# --------------------------------------------------------------------------- +# Adapter: wraps v3 DataBundle to match ClinicalLoader.feature_matrix() API +# --------------------------------------------------------------------------- + +class _BundleLoaderAdapter: + """Thin wrapper around a v3 DataBundle for use in phase_1 classical CV.""" + + def __init__(self, bundle, label_col: str, drop_suspects: bool = True): + self._bundle = bundle + self.label_col = label_col + self.drop_suspects = drop_suspects + self._df_cache: Optional[pd.DataFrame] = None + + @property + def df(self) -> pd.DataFrame: + if self._df_cache is None: + df = self._bundle.df.copy() + if self.drop_suspects: + df = df[df[self.label_col] != 2].reset_index(drop=True) + self._df_cache = df + return self._df_cache + + def feature_matrix(self) -> Tuple[np.ndarray, np.ndarray, List[str], np.ndarray]: + df = self.df.copy() + patient_col = self._bundle.patient_col + scalar_cols = [c for c in self._bundle.scalar_cols if c in df.columns] + for col in scalar_cols: + df[col] = pd.to_numeric(df[col], errors="coerce") + df[col] = df[col].fillna(df[col].median()) + X_num = df[scalar_cols].values.astype(np.float32) + names = list(scalar_cols) + + parts = [X_num] + cat_present = [c for c in self._bundle.cat_cols if c in df.columns] + if cat_present: + dummies = pd.get_dummies(df[cat_present].astype("category"), drop_first=False) + parts.append(dummies.values.astype(np.float32)) + names.extend(list(dummies.columns)) + + X = np.concatenate(parts, axis=1) + y = (df[self.label_col].values.astype(int) == 1).astype(int) + groups = df[patient_col].values.astype(int) + return X, y, names, groups + + +# --------------------------------------------------------------------------- +# Shared CV utilities +# --------------------------------------------------------------------------- + +def _oof_scores(model, X, y, groups, n_splits, seed, patient_level_cv=True): + if patient_level_cv: + splitter = StratifiedGroupKFold(n_splits=n_splits) + split_iter = splitter.split(X, y, groups=groups) + else: + splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed) + split_iter = splitter.split(X, y) + scores = np.zeros(len(y), dtype=float) + preds = np.zeros(len(y), dtype=int) + for tr_idx, te_idx in split_iter: + Xtr, Xte = X[tr_idx], X[te_idx] + ytr = y[tr_idx] + if np.unique(ytr).size < 2: + continue + m = clone(model) + m.fit(Xtr, ytr) + preds[te_idx] = m.predict(Xte) + if hasattr(m, "predict_proba"): + scores[te_idx] = m.predict_proba(Xte)[:, 1] + elif hasattr(m, "decision_function"): + scores[te_idx] = m.decision_function(Xte) + else: + scores[te_idx] = preds[te_idx].astype(float) + return y.astype(int), scores, preds + + +def _cv_curves(model, X, y, groups, n_splits, seed, patient_level_cv=True): + if patient_level_cv: + splitter = StratifiedGroupKFold(n_splits=n_splits) + split_iter = splitter.split(X, y, groups=groups) + else: + splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed) + split_iter = splitter.split(X, y) + curves, fold_aucs, fold_accs = [], [], [] + for tr_idx, te_idx in split_iter: + Xtr, Xte = X[tr_idx], X[te_idx] + ytr, yte = y[tr_idx], y[te_idx] + if np.unique(ytr).size < 2 or np.unique(yte).size < 2: + continue + m = clone(model) + m.fit(Xtr, ytr) + if hasattr(m, "predict_proba"): + sc = m.predict_proba(Xte)[:, 1] + elif hasattr(m, "decision_function"): + sc = m.decision_function(Xte) + else: + sc = m.predict(Xte).astype(float) + fpr, tpr, _ = roc_curve(yte, sc, pos_label=1) + curves.append((fpr, tpr, float(roc_auc_score(yte, sc)))) + fold_aucs.append(float(roc_auc_score(yte, sc))) + fold_accs.append(float(accuracy_score(yte, m.predict(Xte)))) + return curves, fold_aucs, fold_accs + + +def _plot_mean_roc(curves, title, path): + if not curves: + return + mean_fpr = np.linspace(0, 1, 200) + tprs, aucs = [], [] + for fpr, tpr, auc_val in curves: + tpr_i = np.interp(mean_fpr, fpr, tpr); tpr_i[0] = 0.0 + tprs.append(tpr_i); aucs.append(auc_val) + mean_tpr = np.mean(tprs, axis=0); mean_tpr[-1] = 1.0 + std_tpr = np.std(tprs, axis=0) + mean_auc = float(np.mean(aucs)); std_auc = float(np.std(aucs)) + fig, ax = plt.subplots(figsize=(5.5, 4.5)) + ax.plot(mean_fpr, mean_tpr, lw=2, label=f"AUC={mean_auc:.3f}±{std_auc:.3f}") + ax.fill_between(mean_fpr, np.maximum(mean_tpr - std_tpr, 0), + np.minimum(mean_tpr + std_tpr, 1), alpha=0.2, color="grey") + ax.plot([0, 1], [0, 1], "k--", lw=1) + ax.set_xlabel("False Positive Rate"); ax.set_ylabel("True Positive Rate") + ax.set_title(title); ax.legend(loc="lower right") + ax.grid(True, alpha=0.3, linestyle="--"); fig.tight_layout() + fig.savefig(path, dpi=170); plt.close(fig) + return mean_auc, std_auc + + +def _plot_overlay(all_curves: dict, title: str, path: Path): + """all_curves: {model_name: (mean_fpr, mean_tpr, mean_auc, std_auc)}""" + fig, ax = plt.subplots(figsize=(7, 5.5)) + cmap = plt.get_cmap("tab10") + for i, (name, (fpr, tpr, mean_auc, std_auc)) in enumerate(all_curves.items()): + ax.plot(fpr, tpr, lw=2, color=cmap(i), label=f"{name} (AUC={mean_auc:.3f}±{std_auc:.3f})") + ax.plot([0, 1], [0, 1], "k--", lw=1) + ax.set_xlabel("False Positive Rate"); ax.set_ylabel("True Positive Rate") + ax.set_title(title); ax.legend(loc="upper left", fontsize="small") + ax.grid(True, alpha=0.3, linestyle="--"); fig.tight_layout() + fig.savefig(path, dpi=170); plt.close(fig) + + +def _print_result(name, aucs, accs): + mu_auc = float(np.mean(aucs)); sd_auc = float(np.std(aucs)) + mu_acc = float(np.mean(accs)); sd_acc = float(np.std(accs)) + print(f" {name:30s} AUC={mu_auc:.3f}±{sd_auc:.3f} ACC={mu_acc:.3f}±{sd_acc:.3f}") + + +# --------------------------------------------------------------------------- +# Classical classifier runners +# --------------------------------------------------------------------------- + +def run_classical( + name: str, + model, + loader: ClinicalLoader, + out_dir: Path, + n_splits: int, + seed: int, + patient_level_cv: bool = True, +) -> dict: + X, y, feat_names, groups = loader.feature_matrix() + curves, fold_aucs, fold_accs = _cv_curves( + model, X, y, groups, n_splits, seed, patient_level_cv=patient_level_cv + ) + + sub = out_dir / name + sub.mkdir(parents=True, exist_ok=True) + res = _plot_mean_roc(curves, f"{name} ROC (mean ± SD)", sub / "roc_mean.png") + mean_auc, std_auc = (res if res else (float("nan"), float("nan"))) + + pd.DataFrame([{ + "model": name, "auc_mean": mean_auc, "auc_std": std_auc, + "acc_mean": float(np.mean(fold_accs)), "acc_std": float(np.std(fold_accs)), + "n_folds": len(fold_aucs), + }]).to_csv(sub / "summary.csv", index=False) + + pd.DataFrame([{ + "fold": i+1, "auc": a, "acc": c + } for i, (a, c) in enumerate(zip(fold_aucs, fold_accs))]).to_csv( + sub / "fold_metrics.csv", index=False + ) + + _print_result(name, fold_aucs, fold_accs) + + # Return curve for overlay + if curves: + mean_fpr = np.linspace(0, 1, 200) + tprs = [np.interp(mean_fpr, fpr, tpr) for fpr, tpr, _ in curves] + mean_tpr = np.mean(tprs, axis=0); mean_tpr[-1] = 1.0 + return {"fpr": mean_fpr, "tpr": mean_tpr, "auc_mean": mean_auc, "auc_std": std_auc} + return {} + + +# --------------------------------------------------------------------------- +# CNN runner +# --------------------------------------------------------------------------- + +def run_cnn( + backbone: str, + image_dir: str, + clinical_dir: str, + label_col: str, + out_dir: Path, + n_splits: int, + seed: int, + epochs: int, + batch_size: int, + lr: float, + freeze_ratio: float, + augment: bool, + device_str: str, + drop_suspects: bool, + preprocessor=None, + img_size: int = 224, + img_loader=None, +) -> dict: + """ + Train a CNN-only (image only, no clinical data) baseline. + + CV strategy: StratifiedGroupKFold on patients (no OD/OS leakage). + Within each outer fold, 20% of training patients are held out as a + validation set for early stopping; the outer test fold is only + evaluated once using the best-val checkpoint. + """ + import torch + import torch.nn as nn + import torch.optim as optim + from torch.utils.data import DataLoader, Dataset + from torchvision import transforms + from PIL import Image + + from v3.classes.backbones import BACKBONES + from v3.classes.image_loader import CachedImageLoader + from v3.classes.utils import choose_device + + device = choose_device(device_str) + spec = BACKBONES.get(backbone) + if spec is None: + raise ValueError(f"Unknown backbone: {backbone!r}. Available: {list(BACKBONES)}") + + # Load patient/eye table from clinical data (labels only — images are the input) + loader_cd = ClinicalLoader(clinical_dir=clinical_dir, drop_suspects=drop_suspects) + df = loader_cd.df[["Patient ID", "eyeID", label_col]].copy() + df = df[df[label_col].isin([0, 1])].reset_index(drop=True) + df["binary_label"] = (df[label_col] == 1).astype(int) + + mean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225] + # If a cropper preprocessor is provided it already resizes to img_size, + # so we skip the Resize in the transform to avoid a second interpolation. + resize_in_tf = preprocessor is None + eval_tf = transforms.Compose([ + *([ transforms.Resize((img_size, img_size)) ] if resize_in_tf else []), + transforms.ToTensor(), + transforms.Normalize(mean, std), + ]) + train_tf = transforms.Compose([ + *([ transforms.Resize((img_size, img_size)) ] if resize_in_tf else []), + transforms.RandomHorizontalFlip(), + transforms.RandomRotation(15), + transforms.ColorJitter(0.2, 0.2, 0.1, 0.05), + transforms.ToTensor(), + transforms.Normalize(mean, std), + ]) if augment else eval_tf + + if img_loader is None: + img_loader = CachedImageLoader(enabled=True, workers=4) + + class EyeDataset(Dataset): + def __init__(self, records, transform): + self.records = records # list of (pid, eye, label) + self.transform = transform + + def warm(self): + paths = [ + str(Path(image_dir) / f"RET{int(pid):03d}{eye.upper()}.jpg") + for pid, eye, _ in self.records + ] + img_loader.warm(paths, preprocessor=preprocessor) + + def __len__(self): return len(self.records) + + def __getitem__(self, idx): + pid, eye, label = self.records[idx] + p = Path(image_dir) / f"RET{int(pid):03d}{eye.upper()}.jpg" + if p.exists(): + img = img_loader.load(p, preprocessor=preprocessor) + else: + img = Image.new("RGB", (img_size, img_size)) + return self.transform(img), int(label) + + class CNNClassifier(nn.Module): + def __init__(self): + super().__init__() + bb_spec = BACKBONES[backbone] + raw_model = bb_spec.ctor(weights=bb_spec.weights_default) + feat_dim, self.backbone = bb_spec.strip(raw_model) + if freeze_ratio > 0: + blocks = bb_spec.blocks(self.backbone) + n_freeze = int(len(blocks) * freeze_ratio) + for blk in blocks[:n_freeze]: + for p in blk.parameters(): + p.requires_grad_(False) + self.head = nn.Linear(feat_dim, 2) + + def forward(self, x): + return self.head(self.backbone(x)) + + def _eval_loader(model, loader): + model.eval() + all_probs, all_y = [], [] + with torch.no_grad(): + for imgs, lbls in loader: + probs = torch.softmax(model(imgs.to(device)), dim=1)[:, 1].cpu().numpy() + all_probs.extend(probs.tolist()) + all_y.extend(lbls.numpy().tolist()) + return np.array(all_y), np.array(all_probs) + + # Build eye-level records and patient-level group array + records_all = list(df[["Patient ID", "eyeID", "binary_label"]].itertuples(index=False, name=None)) + patient_ids = df["Patient ID"].values.astype(int) + labels_arr = df["binary_label"].values.astype(int) + + # Patient-level label for stratification in outer splitter + pat_label_map = df.groupby("Patient ID")["binary_label"].first().to_dict() + patient_labels = np.array([pat_label_map[p] for p in patient_ids]) + + fold_aucs, fold_accs, curves = [], [], [] + outer = StratifiedGroupKFold(n_splits=n_splits) + + for fold, (trainval_idx, te_idx) in enumerate( + outer.split(records_all, patient_labels, groups=patient_ids)): + print(f" [CNN {backbone}] fold {fold+1}/{n_splits}", flush=True) + + # Split trainval patients into train/val (80/20) for early stopping + tv_patients = np.unique(patient_ids[trainval_idx]) + tv_pat_labels = np.array([pat_label_map[p] for p in tv_patients]) + inner = StratifiedGroupKFold(n_splits=5) + tr_pat_set, va_pat_set = next(iter( + (set(tv_patients[ti]), set(tv_patients[vi])) + for ti, vi in [next(inner.split(tv_patients, tv_pat_labels, groups=tv_patients))] + )) + + tr_recs = [records_all[i] for i in trainval_idx if patient_ids[i] in tr_pat_set] + va_recs = [records_all[i] for i in trainval_idx if patient_ids[i] in va_pat_set] + te_recs = [records_all[i] for i in te_idx] + + tr_ds = EyeDataset(tr_recs, train_tf) + va_ds = EyeDataset(va_recs, eval_tf) + te_ds = EyeDataset(te_recs, eval_tf) + for ds in (tr_ds, va_ds, te_ds): + ds.warm() + + tr_loader = DataLoader(tr_ds, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True) + va_loader = DataLoader(va_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True) + te_loader = DataLoader(te_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True) + + model_cnn = CNNClassifier().to(device) + opt = optim.Adam(filter(lambda p: p.requires_grad, model_cnn.parameters()), lr=lr) + + # Class-weighted loss: w_c = N / (N_c * C), matching paper eq. (2) + tr_labels = [r[2] for r in tr_recs] + n_total = len(tr_labels) + n_classes = 2 + class_counts = np.bincount(tr_labels, minlength=n_classes).astype(float) + class_counts = np.maximum(class_counts, 1) # avoid div-by-zero + weights = torch.tensor( + n_total / (class_counts * n_classes), dtype=torch.float32 + ).to(device) + criterion = nn.CrossEntropyLoss(weight=weights) + + for ep in range(epochs): + model_cnn.train() + for imgs, lbls in tr_loader: + imgs, lbls = imgs.to(device), lbls.to(device) + opt.zero_grad() + criterion(model_cnn(imgs), lbls).backward() + opt.step() + + if (ep + 1) % 5 == 0 or ep == epochs - 1: + val_y, val_probs = _eval_loader(model_cnn, va_loader) + val_auc = float(roc_auc_score(val_y, val_probs)) if val_y.size and len(np.unique(val_y)) > 1 else float("nan") + print(f" ep {ep+1:3d}/{epochs} val_auc={val_auc:.3f}", flush=True) + te_y, te_probs = _eval_loader(model_cnn, te_loader) + if te_y.size and len(np.unique(te_y)) > 1: + auc_val = float(roc_auc_score(te_y, te_probs)) + acc_val = float(accuracy_score(te_y, (te_probs >= 0.5).astype(int))) + fold_aucs.append(auc_val) + fold_accs.append(acc_val) + fpr, tpr, _ = roc_curve(te_y, te_probs, pos_label=1) + curves.append((fpr, tpr, auc_val)) + print(f" fold {fold+1} TEST → AUC={auc_val:.3f} ACC={acc_val:.3f}", flush=True) + + name = f"CNN ({backbone})" + sub = out_dir / f"cnn_{backbone}" + sub.mkdir(parents=True, exist_ok=True) + res = _plot_mean_roc(curves, f"{name} ROC (mean ± SD)", sub / "roc_mean.png") + mean_auc, std_auc = (res if res else (float("nan"), float("nan"))) + + pd.DataFrame([{ + "backbone": backbone, "auc_mean": mean_auc, "auc_std": std_auc, + "acc_mean": float(np.mean(fold_accs)) if fold_accs else float("nan"), + "acc_std": float(np.std(fold_accs)) if fold_accs else float("nan"), + "n_folds": len(fold_aucs), + }]).to_csv(sub / "summary.csv", index=False) + + pd.DataFrame([{ + "fold": i+1, "auc": a, "acc": c + } for i, (a, c) in enumerate(zip(fold_aucs, fold_accs))]).to_csv( + sub / "fold_metrics.csv", index=False + ) + + _print_result(name, fold_aucs, fold_accs) + + if curves: + mean_fpr = np.linspace(0, 1, 200) + tprs = [np.interp(mean_fpr, fpr, tpr) for fpr, tpr, _ in curves] + mean_tpr = np.mean(tprs, axis=0); mean_tpr[-1] = 1.0 + return {"fpr": mean_fpr, "tpr": mean_tpr, "auc_mean": mean_auc, "auc_std": std_auc} + return {} + + +# --------------------------------------------------------------------------- +# Main +# --------------------------------------------------------------------------- + +def build_parser(): + ap = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + # Which classifiers to run + ap.add_argument("--knn", action="store_true", help="Run K-Nearest Neighbours") + ap.add_argument("--rf", action="store_true", help="Run Random Forest") + ap.add_argument("--svm", action="store_true", help="Run SVM") + ap.add_argument("--logreg", action="store_true", help="Run Logistic Regression") + ap.add_argument("--cnn", action="store_true", help="Run CNN") + ap.add_argument("--all", action="store_true", help="Run all classifiers") + + # Data paths + ap.add_argument("--image-dir", default="Papila/FundusImages") + ap.add_argument("--clinical-dir", default="Papila/ClinicalData") + ap.add_argument("--label-col", default="Diagnosis") + ap.add_argument("--output-dir", default="analysis_data/papila_reproduce") + ap.add_argument("--tag", default=None, + help="Optional suffix appended to --output-dir (e.g. 'paper_matched').") + + # Clinical data loader options + ap.add_argument("--no-iop-corr", action="store_true", + help="Skip IOP correction (use raw Perkins/Pneumatic values)") + ap.add_argument("--keep-raw-iop", action="store_true", + help="Keep raw IOP column alongside corrected IOP") + ap.add_argument("--no-cat-cols", action="store_true", + help="Exclude categorical columns (Gender, Phakic/Pseudophakic)") + ap.add_argument("--exclude-cols", nargs="*", default=[], + help="Additional columns to exclude from clinical feature matrix") + ap.add_argument("--keep-suspects", action="store_true", + help="Include Suspect (label 2) rows (default: drop them)") + ap.add_argument("--hypertower-loader", action="store_true", + help="Use the v3 HyperTower clinical data bundle (better IOP correction) " + "instead of the standalone ClinicalLoader.") + + # CV + ap.add_argument("--n-splits", type=int, default=5) + ap.add_argument("--seed", type=int, default=42) + ap.add_argument("--paper-cv", action="store_true", + help="Use eye-level StratifiedKFold (matches paper's likely methodology) " + "instead of patient-level GroupKFold (our cleaner default).") + + # CNN image cropping (GT or UNet, same flags as main hypertower) + ap.add_argument("--img-crop-manifest", default=None, + help="Path to crop manifest CSV (enables cropping).") + ap.add_argument("--img-crop-gt", action="store_true", + help="Use GT segmentations to crop (requires --img-crop-manifest).") + ap.add_argument("--img-crop-weights", default=None, + help="UNet weights path for disc cropping (requires --img-crop-manifest).") + ap.add_argument("--img-crop-scale", type=float, default=2.5) + ap.add_argument("--img-crop-size", type=int, default=200, + help="Crop target size in pixels (default 200, matching PAPILA paper).") + ap.add_argument("--img-crop-cache", default="cache_data/phase1_crops") + ap.add_argument("--persist-img-crop-cache", action="store_true") + + # CNN options + ap.add_argument("--backbone", default="refugelike", + help="CNN backbone key (refugelike, resnet50, densenet121, vgg16, " + "efficientnet_b0, inception_v3, mobilenet_v2, refuge_densenet, ...)") + ap.add_argument("--backbones", nargs="+", default=None, + help="Run multiple backbones sequentially, sharing the image cache. " + "Overrides --backbone. e.g. --backbones resnet50 densenet121 vgg16") + ap.add_argument("--epochs", type=int, default=15) + ap.add_argument("--batch-size", type=int, default=16) + ap.add_argument("--lr", type=float, default=1e-4) + ap.add_argument("--freeze-ratio", type=float, default=0.0, + help="Fraction of backbone blocks to freeze (0=finetune all, 1=freeze all)") + ap.add_argument("--augment", action="store_true") + ap.add_argument("--device", default="auto", + choices=["auto", "cpu", "cuda"]) + + # Classical ML hyperparameters + ap.add_argument("--knn-k", type=int, default=5) + ap.add_argument("--rf-n-estimators", type=int, default=500) + ap.add_argument("--svm-c", type=float, default=1.0) + ap.add_argument("--lr-c", type=float, default=1.0) + return ap + + +def main(): + ap = build_parser() + args = ap.parse_args() + + if args.all: + args.knn = args.rf = args.svm = args.logreg = args.cnn = True + + if not any([args.knn, args.rf, args.svm, args.logreg, args.cnn]): + ap.error("Specify at least one classifier: --knn --rf --svm --logreg --cnn (or --all)") + + out_dir = Path(args.output_dir) + if args.tag: + out_dir = out_dir / args.tag + out_dir.mkdir(parents=True, exist_ok=True) + + # Build clinical data loader + if args.hypertower_loader: + from v3.classes.papila_builders import build_papila_data + bundle = build_papila_data( + image_dir=args.image_dir, + clinical_dir=args.clinical_dir, + label_col=args.label_col, + cat_cols=["Gender", "Phakic/Pseudophakic"], + n_splits=args.n_splits, + random_seed=args.seed, + iop_corr_method="ratio", + iop_drop_raw=True, + ) + loader = _BundleLoaderAdapter(bundle, label_col=args.label_col, + drop_suspects=not args.keep_suspects) + print(f"Clinical data: {len(loader.df)} rows [HyperTower loader] " + f"(suspects {'kept' if args.keep_suspects else 'dropped'})") + else: + loader = ClinicalLoader( + clinical_dir=args.clinical_dir, + label_col=args.label_col, + cat_cols=[] if args.no_cat_cols else None, + exclude_cols=list(args.exclude_cols or []), + iop_corr=not args.no_iop_corr, + keep_raw_iop=args.keep_raw_iop, + drop_suspects=not args.keep_suspects, + ) + print(f"Clinical data: {len(loader.df)} rows " + f"(suspects {'kept' if args.keep_suspects else 'dropped'})") + + X, y, feat_names, groups = loader.feature_matrix() + n_patients = len(np.unique(groups)) + print(f"Feature matrix: {X.shape} ({n_patients} patients) class balance: {dict(zip(*np.unique(y, return_counts=True)))}") + + overlay_curves: dict = {} + all_results: list = [] + + t0 = time.time() + + # KNN + if args.knn: + print("\n--- KNN ---") + model = Pipeline([ + ("scale", StandardScaler()), + ("knn", KNeighborsClassifier(n_neighbors=args.knn_k)), + ]) + r = run_classical("KNN", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv) + if r: + overlay_curves["KNN"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"]) + + # Random Forest + if args.rf: + print("\n--- Random Forest ---") + model = RandomForestClassifier( + n_estimators=args.rf_n_estimators, max_features="sqrt", + random_state=args.seed, n_jobs=-1, + ) + r = run_classical("Random Forest", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv) + if r: + overlay_curves["Random Forest"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"]) + + # SVM + if args.svm: + print("\n--- SVM ---") + model = Pipeline([ + ("scale", StandardScaler()), + ("svm", SVC(kernel="rbf", C=args.svm_c, gamma="scale", + probability=True, random_state=args.seed)), + ]) + r = run_classical("SVM", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv) + if r: + overlay_curves["SVM"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"]) + + # Logistic Regression + if args.logreg: + print("\n--- Logistic Regression ---") + model = Pipeline([ + ("scale", StandardScaler()), + ("logreg", LogisticRegression(C=args.lr_c, max_iter=1000, + solver="lbfgs")), + ]) + r = run_classical("Logistic Regression", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv) + if r: + overlay_curves["Logistic Regression"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"]) + + # CNN + if args.cnn: + from v3.classes.croppers import build_image_preprocessor_from_args + from v3.classes.image_loader import CachedImageLoader as _CachedImageLoader + cnn_preprocessor = build_image_preprocessor_from_args(args) + backbones_to_run = args.backbones if args.backbones else [args.backbone] + shared_img_loader = _CachedImageLoader(enabled=True, workers=4) + for backbone in backbones_to_run: + print(f"\n--- CNN ({backbone})" + + (" [cropped]" if cnn_preprocessor else "") + " ---") + r = run_cnn( + backbone=backbone, + image_dir=args.image_dir, + clinical_dir=args.clinical_dir, + label_col=args.label_col, + out_dir=out_dir, + n_splits=args.n_splits, + seed=args.seed, + epochs=args.epochs, + batch_size=args.batch_size, + lr=args.lr, + freeze_ratio=args.freeze_ratio, + augment=args.augment, + device_str=args.device, + drop_suspects=not args.keep_suspects, + preprocessor=cnn_preprocessor, + img_size=args.img_crop_size if cnn_preprocessor else 224, + img_loader=shared_img_loader, + ) + if r: + overlay_curves[f"CNN ({backbone})"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"]) + + # Overlay ROC + if len(overlay_curves) > 1: + _plot_overlay(overlay_curves, "PAPILA Reproduce — Clinical + CNN ROC", out_dir / "roc_overlay.png") + print(f"\nOverlay ROC saved: {out_dir / 'roc_overlay.png'}") + + print(f"\nDone in {time.time()-t0:.1f}s — results in {out_dir}") + + +if __name__ == "__main__": + main() diff --git a/v3/scripts/main/run_cv.py b/v3/scripts/main/run_cv.py new file mode 100644 index 0000000..8f11a76 --- /dev/null +++ b/v3/scripts/main/run_cv.py @@ -0,0 +1,87 @@ +#!/usr/bin/env python +""" +V3 cross-validation runner. + +Outer/inner k-fold: test = current fold, val = next fold, train = rest. +No holdout. No checkpoint saving. + +Usage (single fold-seed, 5-fold, binary, ensemble): + python -m v3.scripts.main.run_cv \ + --run-name my_run \ + --eval-mode binary \ + --tower-mode ensemble \ + --epochs 40 \ + --augment \ + --tune-binary-threshold \ + --in-memory-cache + +Usage (10x5 rep-CV, seeds 100..1000): + python -m v3.scripts.main.run_cv \ + --run-name my_run_10x5 \ + --reps 10 \ + --rep-seed-start 100 \ + --rep-seed-step 100 \ + --eval-mode binary \ + --tower-mode ensemble \ + --epochs 40 \ + --augment \ + --tune-binary-threshold \ + --in-memory-cache +""" +import argparse +import sys +from pathlib import Path + +# Allow running as `python v3/scripts/main/run_cv.py` from repo root +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +from v3.classes.v3_hypertower import V3HyperTower + + +def build_parser() -> argparse.ArgumentParser: + ap = V3HyperTower.build_parser() + ap.description = __doc__ + ap.formatter_class = argparse.RawDescriptionHelpFormatter + ap.add_argument( + "--reps", type=int, default=1, + help="Number of repetitions (each rep uses a different --fold-seed).", + ) + ap.add_argument( + "--rep-seed-start", type=int, default=100, + help="fold-seed for rep 0 (default: 100).", + ) + ap.add_argument( + "--rep-seed-step", type=int, default=100, + help="Increment between rep fold-seeds (default: 100; rep k uses seed start + k*step).", + ) + return ap + + +def main(): + ap = build_parser() + args = ap.parse_args() + + reps = int(args.reps) + seed_start = int(args.rep_seed_start) + seed_step = int(args.rep_seed_step) + base_run_name = args.run_name or "v3_cv" + + for rep in range(reps): + rep_seed = seed_start + rep * seed_step + args.fold_seed = rep_seed + + if reps > 1: + args.run_name = f"{base_run_name}/rep{rep:02d}" + print(f"\n{'='*60}", flush=True) + print(f"Rep {rep+1}/{reps} fold_seed={rep_seed}", flush=True) + print(f"{'='*60}", flush=True) + else: + args.run_name = base_run_name + + tower = V3HyperTower(args) + out_dir = tower.run() + print(f"\nRep {rep+1} output: {out_dir}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/v3/scripts/output_analysis/__init__.py b/v3/scripts/output_analysis/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/v3/scripts/output_analysis/plot_phase1_boxplots.py b/v3/scripts/output_analysis/plot_phase1_boxplots.py new file mode 100644 index 0000000..8f1cfa9 --- /dev/null +++ b/v3/scripts/output_analysis/plot_phase1_boxplots.py @@ -0,0 +1,177 @@ +#!/usr/bin/env python +""" +Plot fold-level AUC box plots for all phase 1 comparison variants. + +Reads fold_metrics.csv files produced by phase_1_papila_reproduce.py +and generates a grouped box plot: one group per classifier, one box +per variant (tag). + +Usage: + python -m v3.scripts.output_analysis.plot_phase1_boxplots \ + --results-dir v3/results/phase1 \ + --tags paper_matched no_leakage hypertower_loader \ + --output v3/results/phase1/auc_boxplot_comparison.png +""" +from __future__ import annotations + +import argparse +import sys +from pathlib import Path + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + + +TAG_LABELS = { + "paper_matched": "Paper-matched\n(eye-level CV)", + "no_leakage": "No leakage\n(patient-level CV)", + "hypertower_loader": "HyperTower loader\n(patient-level CV)", +} + +CLASSIFIER_ORDER = ["KNN", "Random Forest", "SVM", "Logistic Regression"] +CLASSIFIER_SHORT = { + "KNN": "KNN", + "Random Forest": "RF", + "SVM": "SVM", + "Logistic Regression": "LR", +} + +# Colours per variant +VARIANT_COLOURS = [ + "#4878CF", # blue — paper_matched + "#6ACC65", # green — no_leakage + "#D65F5F", # red — hypertower_loader +] + + +def load_fold_aucs(results_dir: Path, tags: list[str]) -> dict: + """ + Returns {tag: {classifier_name: [fold_auc, ...]}} + """ + data: dict = {} + for tag in tags: + tag_dir = results_dir / tag + data[tag] = {} + for clf in CLASSIFIER_ORDER: + fpath = tag_dir / clf / "fold_metrics.csv" + if fpath.exists(): + df = pd.read_csv(fpath) + data[tag][clf] = df["auc"].tolist() + else: + print(f" WARNING: missing {fpath}") + data[tag][clf] = [] + return data + + +def plot_boxplots(data: dict, tags: list[str], output: Path, paper_aucs: dict | None = None): + n_clf = len(CLASSIFIER_ORDER) + n_tags = len(tags) + group_width = 0.8 + box_width = group_width / n_tags * 0.85 + offsets = np.linspace(-group_width / 2 + box_width / 2, + group_width / 2 - box_width / 2, n_tags) + + fig, ax = plt.subplots(figsize=(10, 5.5)) + + for ti, tag in enumerate(tags): + colour = VARIANT_COLOURS[ti % len(VARIANT_COLOURS)] + label = TAG_LABELS.get(tag, tag) + first = True + for ci, clf in enumerate(CLASSIFIER_ORDER): + aucs = data[tag].get(clf, []) + if not aucs: + continue + x = ci + offsets[ti] + bp = ax.boxplot( + aucs, + positions=[x], + widths=box_width, + patch_artist=True, + boxprops=dict(facecolor=colour, alpha=0.75), + medianprops=dict(color="black", linewidth=1.8), + whiskerprops=dict(color=colour, linewidth=1.2), + capprops=dict(color=colour, linewidth=1.2), + flierprops=dict(marker="o", markersize=4, + markerfacecolor=colour, alpha=0.6), + manage_ticks=False, + ) + if first: + bp["boxes"][0].set_label(label) + first = False + + # Paper reference lines (dashed, per-classifier) + if paper_aucs: + for ci, clf in enumerate(CLASSIFIER_ORDER): + if clf in paper_aucs: + ax.hlines(paper_aucs[clf], ci - group_width / 2, ci + group_width / 2, + colors="black", linestyles=":", linewidths=1.2, + label="Paper (PAPILA)" if ci == 0 else "_nolegend_") + + ax.set_xticks(range(n_clf)) + ax.set_xticklabels([CLASSIFIER_SHORT[c] for c in CLASSIFIER_ORDER], fontsize=12) + ax.set_ylabel("AUC (ROC)", fontsize=11) + ax.set_title("Phase 1: Clinical-only classifier AUC by CV strategy", fontsize=12) + ax.set_ylim(0.45, 1.02) + ax.axhline(0.5, color="grey", linestyle="--", linewidth=0.8, alpha=0.5) + ax.grid(axis="y", alpha=0.3, linestyle="--") + ax.legend(loc="lower right", fontsize=9, framealpha=0.9) + fig.tight_layout() + + output.parent.mkdir(parents=True, exist_ok=True) + fig.savefig(output, dpi=180) + plt.close(fig) + print(f"Saved: {output}") + + +def main(): + ap = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument("--results-dir", default="v3/results/phase1") + ap.add_argument("--tags", nargs="+", + default=["paper_matched", "no_leakage", "hypertower_loader"]) + ap.add_argument("--output", default=None, + help="Output PNG path. Default: /auc_boxplot_comparison.png") + ap.add_argument("--no-paper-lines", action="store_true", + help="Omit the dotted paper-reported AUC reference lines.") + args = ap.parse_args() + + results_dir = Path(args.results_dir) + output = Path(args.output) if args.output else results_dir / "auc_boxplot_comparison.png" + + data = load_fold_aucs(results_dir, args.tags) + + # PAPILA paper reported AUCs + paper_aucs = None if args.no_paper_lines else { + "KNN": 0.75, + "Random Forest": 0.64, + "SVM": 0.75, + "Logistic Regression": 0.70, + } + + plot_boxplots(data, args.tags, output, paper_aucs=paper_aucs) + + # Print summary table + print(f"\n{'Classifier':<20}", end="") + for tag in args.tags: + label = tag.replace("_", " ") + print(f" {label:>22}", end="") + print() + print("-" * (20 + 24 * len(args.tags))) + for clf in CLASSIFIER_ORDER: + print(f"{CLASSIFIER_SHORT[clf]:<20}", end="") + for tag in args.tags: + aucs = data[tag].get(clf, []) + if aucs: + print(f" {np.mean(aucs):.3f} ± {np.std(aucs):.3f} ", end="") + else: + print(f" {'—':>22}", end="") + print() + + +if __name__ == "__main__": + main()