post-restructure
This commit is contained in:
@@ -33,6 +33,7 @@ from .transforms import (
|
||||
ImageTransformConfig,
|
||||
backbone_transform_config,
|
||||
build_backbone_transform,
|
||||
build_eval_transform,
|
||||
build_imagenet_transform,
|
||||
ResizeTransform,
|
||||
CenterCropTransform,
|
||||
@@ -45,6 +46,7 @@ from .transforms import (
|
||||
from .model_builder import V2ModelBundle, build_model_bundle
|
||||
from .towers import ImageTower, MDTower, 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
|
||||
|
||||
@@ -79,6 +81,7 @@ __all__ = [
|
||||
"ImageTransformConfig",
|
||||
"backbone_transform_config",
|
||||
"build_backbone_transform",
|
||||
"build_eval_transform",
|
||||
"build_imagenet_transform",
|
||||
"ResizeTransform",
|
||||
"CenterCropTransform",
|
||||
@@ -95,6 +98,8 @@ __all__ = [
|
||||
"build_backbone",
|
||||
"Bridge",
|
||||
"VoteBridge",
|
||||
"SingleEyeHT",
|
||||
"BilateralHT",
|
||||
"V2HyperTower",
|
||||
"V2ModeComparisonOps",
|
||||
"V2ModeComparator",
|
||||
|
||||
@@ -0,0 +1,395 @@
|
||||
"""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 classes.geometry_features import compute_geometry_features, disc_cup_from_mask_image
|
||||
from classes.refuge_classification import _geometry_from_mask
|
||||
from classes.unet_segmenter import UNetSegmenter
|
||||
|
||||
|
||||
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.to_tensor(resized).unsqueeze(0).to(self.segmenter.device)
|
||||
|
||||
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
|
||||
@@ -55,3 +55,53 @@ class ClinicalDataset(Dataset):
|
||||
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])
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from .network_manager import LoaderBundle, PatientSplit
|
||||
@@ -163,3 +164,60 @@ class SlotLoaderFactory:
|
||||
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,
|
||||
batch_size: int,
|
||||
shuffle: bool,
|
||||
num_workers: int,
|
||||
) -> DataLoader:
|
||||
ds = SlotDataset(
|
||||
samples,
|
||||
slots,
|
||||
image_transform=image_transform,
|
||||
image_preprocessor=image_preprocessor,
|
||||
)
|
||||
return DataLoader(
|
||||
ds,
|
||||
batch_size=batch_size,
|
||||
shuffle=shuffle,
|
||||
num_workers=num_workers,
|
||||
collate_fn=slot_collate,
|
||||
)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
"""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,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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:
|
||||
if y_true.size == 0:
|
||||
return 0.5
|
||||
grid = np.linspace(0.0, 1.0, 1001)
|
||||
best_t, best_acc = 0.5, -1.0
|
||||
for t in grid:
|
||||
pred = (p1 >= t).astype(int)
|
||||
acc = float((pred == y_true).mean())
|
||||
if acc > best_acc or (acc == best_acc and abs(t - 0.5) < abs(best_t - 0.5)):
|
||||
best_acc, best_t = acc, float(t)
|
||||
return best_t
|
||||
|
||||
|
||||
def multiclass_acc_with_bias(y_true: np.ndarray, probs: np.ndarray, bias: np.ndarray) -> float:
|
||||
if y_true.size == 0:
|
||||
return float("nan")
|
||||
logits = np.log(np.clip(probs, 1e-8, 1.0)) + bias.reshape(1, -1)
|
||||
return float((np.argmax(logits, axis=1) == y_true).mean())
|
||||
|
||||
|
||||
def tune_multiclass_bias(y_true: np.ndarray, probs: np.ndarray, *, iters: int = 2) -> np.ndarray:
|
||||
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
|
||||
@@ -0,0 +1,539 @@
|
||||
"""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 classes.v2.bridges import Bridge
|
||||
from classes.v2.towers import ImageTower, MDTower
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Model classes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class SingleEyeHT(nn.Module):
|
||||
"""
|
||||
ImageTower + MDTower + 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,
|
||||
md_hidden_dim: int = 128,
|
||||
fusion_dim: int = 256,
|
||||
):
|
||||
super().__init__()
|
||||
self.img_tower = ImageTower(
|
||||
backbone=backbone,
|
||||
freeze_ratio=freeze_ratio,
|
||||
augment=augment,
|
||||
use_se=False,
|
||||
)
|
||||
self.md_tower = MDTower(
|
||||
clinical_data=clinical_data,
|
||||
hidden_dim=md_hidden_dim,
|
||||
use_se=False,
|
||||
)
|
||||
self.bridge = Bridge(
|
||||
img_dim=self.img_tower.out_dim,
|
||||
meta_dim=self.md_tower.out_dim,
|
||||
num_classes=num_classes,
|
||||
fusion_dim=fusion_dim,
|
||||
mode="fused",
|
||||
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 = self.img_tower(x)
|
||||
md_feats = self.md_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 metadata towers combine OD/OS embeddings
|
||||
- standard Bridge fuses joint image + joint metadata embeddings
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
backbone: str,
|
||||
freeze_ratio: float,
|
||||
augment: bool,
|
||||
clinical_data,
|
||||
num_classes: int,
|
||||
md_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_md_tower = MDTower(
|
||||
clinical_data=clinical_data,
|
||||
hidden_dim=md_hidden_dim,
|
||||
use_se=False,
|
||||
)
|
||||
img_dim = self.eye_img_tower.out_dim
|
||||
md_dim = self.eye_md_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_md_tower(meta_od)
|
||||
img_os = self.eye_img_tower(x_os)
|
||||
md_os = self.eye_md_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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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:
|
||||
if phase == "tower_warmup":
|
||||
_set_requires_grad(model.img_tower, True)
|
||||
_set_requires_grad(model.md_tower, True)
|
||||
_set_requires_grad(model.bridge.classifier_img, True)
|
||||
_set_requires_grad(model.bridge.classifier_md, 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 == "fused_warmup":
|
||||
_set_requires_grad(model.img_tower, False)
|
||||
_set_requires_grad(model.md_tower, False)
|
||||
_set_requires_grad(model.bridge.classifier_img, False)
|
||||
_set_requires_grad(model.bridge.classifier_md, 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_md_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_md_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,
|
||||
) -> 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 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)
|
||||
img_feats = model.img_tower(x)
|
||||
md_feats = model.md_tower(m)
|
||||
|
||||
if phase == "tower_warmup":
|
||||
logits_i = model.bridge.classifier_img(img_feats)
|
||||
logits_m = model.bridge.classifier_md(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 random() < bcd_prob:
|
||||
if random() < 0.5:
|
||||
logits = model.bridge.classifier_img(img_feats)
|
||||
else:
|
||||
logits = model.bridge.classifier_md(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,
|
||||
) -> 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 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"),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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(
|
||||
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_single_components(
|
||||
model: SingleEyeHT,
|
||||
loader: DataLoader,
|
||||
device: torch.device,
|
||||
*,
|
||||
aggregate_patient: bool,
|
||||
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
|
||||
"""
|
||||
Collect fused/img/md probabilities for SingleEyeHT.
|
||||
- aggregate_patient=False: eye-level (OD/OS as independent samples)
|
||||
- aggregate_patient=True : patient-level (average OD/OS per head)
|
||||
"""
|
||||
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)
|
||||
|
||||
def _per_eye_probs(x, m):
|
||||
img_feats = model.img_tower(x.to(device))
|
||||
md_feats = model.md_tower(m.to(device))
|
||||
out_f, out_i, out_m = model.bridge(img_feats, md_feats)
|
||||
return (
|
||||
F.softmax(out_f, dim=1),
|
||||
F.softmax(out_i, dim=1),
|
||||
F.softmax(out_m, dim=1),
|
||||
)
|
||||
|
||||
pf_od, pi_od, pm_od = _per_eye_probs(x1, m1)
|
||||
pf_os, pi_os, pm_os = _per_eye_probs(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())
|
||||
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()]
|
||||
|
||||
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),
|
||||
)
|
||||
|
||||
|
||||
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)
|
||||
@@ -0,0 +1,96 @@
|
||||
"""Result dataclasses and serialisation helpers for V2 fold outputs."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Primitive helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Dataclasses
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@dataclass
|
||||
class FoldResult:
|
||||
mode: str
|
||||
fold: int
|
||||
# Epoch where each model hit its peak val AUC
|
||||
best_epoch_single: int # SingleEyeHT — selected by ensemble val AUC
|
||||
best_epoch_bilat: int # BilateralHT — selected by bilateral val AUC
|
||||
# Classic (eye-level eval of SingleEyeHT; n = 2 * ensemble_val_n)
|
||||
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 same 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
|
||||
# Holdout metrics (evaluated at best val epoch; nan if no holdout)
|
||||
classic_holdout_auc: float
|
||||
classic_holdout_acc: float
|
||||
ensemble_holdout_auc: float
|
||||
ensemble_holdout_acc: float
|
||||
bilat_holdout_auc: float
|
||||
bilat_holdout_acc: float
|
||||
holdout_n: int # number of holdout bilateral samples
|
||||
# Training sample counts
|
||||
single_train_n: int
|
||||
bilat_train_n: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class FoldArtifacts:
|
||||
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]
|
||||
@@ -76,6 +76,11 @@ def build_backbone_transform(backbone_name: str, augment: bool = True) -> transf
|
||||
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()
|
||||
|
||||
|
||||
@@ -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
|
||||
+1058
-3742
File diff suppressed because it is too large
Load Diff
@@ -11,7 +11,7 @@ REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
if str(REPO_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(REPO_ROOT))
|
||||
|
||||
from classes.v2.v2_hypertower import build_parser, run_mode
|
||||
from classes.v2.v2_hypertower import V2HyperTower
|
||||
|
||||
|
||||
def parse_args():
|
||||
@@ -35,7 +35,7 @@ def parse_args():
|
||||
|
||||
def main():
|
||||
seq_args, remaining = parse_args()
|
||||
base_parser = build_parser()
|
||||
base_parser = V2HyperTower.build_parser()
|
||||
first_run = True
|
||||
for eval_mode in seq_args.eval_modes:
|
||||
for tower_mode in seq_args.tower_modes:
|
||||
@@ -45,7 +45,7 @@ def main():
|
||||
if not first_run:
|
||||
cli.append("--persist-img-crop-cache")
|
||||
args = base_parser.parse_args(cli)
|
||||
run_mode(args)
|
||||
V2HyperTower(args).run()
|
||||
first_run = False
|
||||
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
#!/usr/bin/env python3
|
||||
"""CLI wrapper that delegates to classes.frontend.Multifold with V2 loaders."""
|
||||
"""CLI wrapper for the V2 hypertower pipeline using V2HyperTower directly."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -7,30 +7,17 @@ from pathlib import Path
|
||||
import sys
|
||||
|
||||
# ensure repo root on path
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
if str(REPO_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(REPO_ROOT))
|
||||
|
||||
import classes.frontend as frontend
|
||||
from classes.v2.v2_hypertower import V2HyperTower
|
||||
|
||||
|
||||
def run_cli(cli_args=None):
|
||||
parser = frontend.Multifold.build_parser()
|
||||
parser.set_defaults(warmup_tower_epochs=None, warmup_fused_epochs=None)
|
||||
parser.add_argument(
|
||||
"--sample-mode",
|
||||
choices=["eye", "patient"],
|
||||
default="eye",
|
||||
help="Build samples per eye (row-level) or per patient (multi-slot).",
|
||||
)
|
||||
parser = V2HyperTower.build_parser()
|
||||
args = parser.parse_args(cli_args)
|
||||
|
||||
# Monkeypatch the HyperTower class used inside Multifold.
|
||||
frontend.HyperTower = V2HyperTower
|
||||
|
||||
runner = frontend.Multifold(args)
|
||||
runner.run()
|
||||
V2HyperTower(args).run()
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -44,13 +44,10 @@ from classes.v2.data_bundle import DataBundle
|
||||
from classes.v2.papila_builders import build_papila_data
|
||||
from classes.v2.profiles.papila import build_papila_profile
|
||||
from classes.v2.split_manager import PatientFirstSplitManager
|
||||
from classes.v2.v2_hypertower import (
|
||||
SingleEyeHT,
|
||||
_score_arrays,
|
||||
build_eval_transform,
|
||||
filter_bilateral_samples,
|
||||
make_loader,
|
||||
)
|
||||
from classes.v2.loader_factory import filter_bilateral_samples, make_loader
|
||||
from classes.v2.metrics import _score_arrays
|
||||
from classes.v2.models import SingleEyeHT
|
||||
from classes.v2.transforms import build_eval_transform
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Label display helpers
|
||||
|
||||
Reference in New Issue
Block a user