post-restructure

This commit is contained in:
rpotter6298
2026-02-26 12:22:47 +01:00
parent 8980cf5f9b
commit fd1452c187
13 changed files with 2495 additions and 3769 deletions
+5
View File
@@ -33,6 +33,7 @@ from .transforms import (
ImageTransformConfig, ImageTransformConfig,
backbone_transform_config, backbone_transform_config,
build_backbone_transform, build_backbone_transform,
build_eval_transform,
build_imagenet_transform, build_imagenet_transform,
ResizeTransform, ResizeTransform,
CenterCropTransform, CenterCropTransform,
@@ -45,6 +46,7 @@ from .transforms import (
from .model_builder import V2ModelBundle, build_model_bundle from .model_builder import V2ModelBundle, build_model_bundle
from .towers import ImageTower, MDTower, SiameseImageTower, build_backbone from .towers import ImageTower, MDTower, SiameseImageTower, build_backbone
from .bridges import Bridge, VoteBridge from .bridges import Bridge, VoteBridge
from .models import SingleEyeHT, BilateralHT
from .v2_hypertower import V2HyperTower, V2ModeComparisonOps, V2ModeComparator from .v2_hypertower import V2HyperTower, V2ModeComparisonOps, V2ModeComparator
from .hypertower_logger import HypertowerLogger from .hypertower_logger import HypertowerLogger
@@ -79,6 +81,7 @@ __all__ = [
"ImageTransformConfig", "ImageTransformConfig",
"backbone_transform_config", "backbone_transform_config",
"build_backbone_transform", "build_backbone_transform",
"build_eval_transform",
"build_imagenet_transform", "build_imagenet_transform",
"ResizeTransform", "ResizeTransform",
"CenterCropTransform", "CenterCropTransform",
@@ -95,6 +98,8 @@ __all__ = [
"build_backbone", "build_backbone",
"Bridge", "Bridge",
"VoteBridge", "VoteBridge",
"SingleEyeHT",
"BilateralHT",
"V2HyperTower", "V2HyperTower",
"V2ModeComparisonOps", "V2ModeComparisonOps",
"V2ModeComparator", "V2ModeComparator",
+395
View File
@@ -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
+50
View File
@@ -55,3 +55,53 @@ class ClinicalDataset(Dataset):
geom_vec = torch.from_numpy(features) geom_vec = torch.from_numpy(features)
return img_t, meta_t, geom_vec, label return img_t, meta_t, geom_vec, label
return img_t, meta_t, 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])
+58
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Callable, Optional from typing import Any, Callable, Optional
import torch
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from .network_manager import LoaderBundle, PatientSplit from .network_manager import LoaderBundle, PatientSplit
@@ -163,3 +164,60 @@ class SlotLoaderFactory:
sample.setdefault(key, None) sample.setdefault(key, None)
samples.append(sample) samples.append(sample)
return samples 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)
+231
View File
@@ -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
+539
View File
@@ -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)
+96
View File
@@ -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]
+5
View File
@@ -76,6 +76,11 @@ def build_backbone_transform(backbone_name: str, augment: bool = True) -> transf
return backbone_transform_config(backbone_name, augment=augment).build() 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: def build_imagenet_transform(augment: bool = True, crop_size: int = 224) -> transforms.Compose:
return ImageTransformConfig(crop_size=crop_size, augment=augment).build() return ImageTransformConfig(crop_size=crop_size, augment=augment).build()
+47
View File
@@ -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
View File
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: if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT)) 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(): def parse_args():
@@ -35,7 +35,7 @@ def parse_args():
def main(): def main():
seq_args, remaining = parse_args() seq_args, remaining = parse_args()
base_parser = build_parser() base_parser = V2HyperTower.build_parser()
first_run = True first_run = True
for eval_mode in seq_args.eval_modes: for eval_mode in seq_args.eval_modes:
for tower_mode in seq_args.tower_modes: for tower_mode in seq_args.tower_modes:
@@ -45,7 +45,7 @@ def main():
if not first_run: if not first_run:
cli.append("--persist-img-crop-cache") cli.append("--persist-img-crop-cache")
args = base_parser.parse_args(cli) args = base_parser.parse_args(cli)
run_mode(args) V2HyperTower(args).run()
first_run = False first_run = False
+4 -17
View File
@@ -1,5 +1,5 @@
#!/usr/bin/env python3 #!/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 from __future__ import annotations
@@ -7,30 +7,17 @@ from pathlib import Path
import sys import sys
# ensure repo root on path # 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: if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT)) sys.path.insert(0, str(REPO_ROOT))
import classes.frontend as frontend
from classes.v2.v2_hypertower import V2HyperTower from classes.v2.v2_hypertower import V2HyperTower
def run_cli(cli_args=None): def run_cli(cli_args=None):
parser = frontend.Multifold.build_parser() parser = V2HyperTower.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).",
)
args = parser.parse_args(cli_args) args = parser.parse_args(cli_args)
V2HyperTower(args).run()
# Monkeypatch the HyperTower class used inside Multifold.
frontend.HyperTower = V2HyperTower
runner = frontend.Multifold(args)
runner.run()
def main(): 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.papila_builders import build_papila_data
from classes.v2.profiles.papila import build_papila_profile from classes.v2.profiles.papila import build_papila_profile
from classes.v2.split_manager import PatientFirstSplitManager from classes.v2.split_manager import PatientFirstSplitManager
from classes.v2.v2_hypertower import ( from classes.v2.loader_factory import filter_bilateral_samples, make_loader
SingleEyeHT, from classes.v2.metrics import _score_arrays
_score_arrays, from classes.v2.models import SingleEyeHT
build_eval_transform, from classes.v2.transforms import build_eval_transform
filter_bilateral_samples,
make_loader,
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Label display helpers # Label display helpers