moved_repo_first_update

This commit is contained in:
rpotter6298
2026-02-24 10:39:48 +01:00
commit 9894a23f09
98 changed files with 35387 additions and 0 deletions
+102
View File
@@ -0,0 +1,102 @@
from .network_manager import (
FoldResult,
LoaderBundle,
NetworkManager,
PatientSplit,
)
from .split_manager import (
PatientFirstSplitManager,
SplitPlan,
build_patient_split_plans,
)
from .profiles import (
DatasetProfile,
SimpleDatasetProfile,
SlotDescriptor,
PapilaProfile,
build_papila_profile,
)
from .loader_factory import SlotLoaderFactory
from .slot_dataset import SlotDataset, slot_collate
from .papila_data import PapilaData
from .papila_builders import build_papila_data
from .data_bundle import DataBundle
from .dataset import ClinicalDataset
from .config_builder import (
ConfigAssembly,
assemble_config,
load_config,
resolve_imports,
)
from .filters import RegexFilter, ColumnFilter, apply_regex_filters, apply_column_filters
from .transforms import (
ImageTransformConfig,
backbone_transform_config,
build_backbone_transform,
build_imagenet_transform,
ResizeTransform,
CenterCropTransform,
ROICropTransform,
JitterBundleTransform,
UnetMaskProvider,
TRANSFORM_REGISTRY,
build_transform_chain,
)
from .model_builder import V2ModelBundle, build_model_bundle
from .towers import ImageTower, MDTower, SiameseImageTower, build_backbone
from .bridges import Bridge, VoteBridge
from .v2_hypertower import V2HyperTower, V2ModeComparisonOps, V2ModeComparator
from .hypertower_logger import HypertowerLogger
__all__ = [
"NetworkManager",
"PatientSplit",
"LoaderBundle",
"FoldResult",
"PatientFirstSplitManager",
"SplitPlan",
"build_patient_split_plans",
"DatasetProfile",
"SimpleDatasetProfile",
"SlotDescriptor",
"PapilaProfile",
"build_papila_profile",
"PapilaData",
"build_papila_data",
"DataBundle",
"ClinicalDataset",
"SlotLoaderFactory",
"SlotDataset",
"slot_collate",
"ConfigAssembly",
"assemble_config",
"load_config",
"resolve_imports",
"RegexFilter",
"ColumnFilter",
"apply_regex_filters",
"apply_column_filters",
"ImageTransformConfig",
"backbone_transform_config",
"build_backbone_transform",
"build_imagenet_transform",
"ResizeTransform",
"CenterCropTransform",
"ROICropTransform",
"JitterBundleTransform",
"UnetMaskProvider",
"TRANSFORM_REGISTRY",
"build_transform_chain",
"V2ModelBundle",
"build_model_bundle",
"ImageTower",
"MDTower",
"SiameseImageTower",
"build_backbone",
"Bridge",
"VoteBridge",
"V2HyperTower",
"V2ModeComparisonOps",
"V2ModeComparator",
"HypertowerLogger",
]
+93
View File
@@ -0,0 +1,93 @@
from __future__ import annotations
import torch
import torch.nn as nn
from classes.SE_attention import SEBlock, SEGateLogger
class Bridge(nn.Module):
def __init__(
self,
img_dim,
meta_dim,
num_classes,
fusion_dim=256,
mode="fused",
use_se: bool = True,
se_reduction: int = 16,
se_pre_norm: bool = True,
):
super().__init__()
self.mode = mode
self.use_se = use_se
# project towers to equal width
self.W_img = nn.Linear(img_dim, fusion_dim)
self.W_md = nn.Linear(meta_dim, fusion_dim)
# optional: layernorm before SE
self.ln_img = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity()
self.ln_md = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity()
# SE gate on the fused vector
self.se = SEBlock(fusion_dim, reduction=se_reduction, residual=True) if use_se else None
self.se_log = SEGateLogger(enabled=use_se, track_channels=False, dim=fusion_dim)
# heads
self.classifier_fused = nn.Sequential(
nn.ReLU(),
nn.Dropout(0.5),
nn.Linear(fusion_dim, num_classes),
)
self.classifier_img = nn.Linear(img_dim, num_classes)
self.classifier_md = nn.Linear(meta_dim, num_classes)
def reset_se_stats(self):
"""Call at epoch start."""
if getattr(self, "se_log", None):
self.se_log.reset()
def get_se_stats(self, reset: bool = True):
"""Call after eval. Returns dict or None."""
if getattr(self, "se_log", None) and self.se_log.enabled:
return self.se_log.get(reset=reset)
return None
def forward(self, img_feats, md_feats):
out_img = None if self.mode == "metadata_only" else self.classifier_img(img_feats)
out_md = None if self.mode == "image_only" else self.classifier_md(md_feats)
if self.mode == "fused":
hi = self.ln_img(self.W_img(img_feats)) # image features
hm = self.ln_md(self.W_md(md_feats)) # metadata features
fused = hi * hm # elementwise product
# apply SE gates
if self.se is not None:
fused, gates = self.se(fused)
if self.se_log.enabled:
self.se_log.accumulate(gates)
if self.se is not None and self.training and self.se_log.enabled:
if not hasattr(self, "_dbg_seen"):
self._dbg_seen = 0
if self._dbg_seen < 3: # print only a few times
print("[SE] gate mean this batch:", gates.mean().item())
self._dbg_seen += 1
out_f = self.classifier_fused(fused)
return out_f, out_img, out_md
# if ablation modes:
if self.mode == "image_only":
return out_img, out_img, None
if self.mode == "metadata_only":
return out_md, None, out_md
class VoteBridge(nn.Module):
def __init__(self, num_classes):
super().__init__()
self.vote_combiner = nn.Linear(num_classes * 2, num_classes) # two sets of logits
def forward(self, out_img, out_md):
votes = torch.cat([out_img, out_md], dim=1)
return self.vote_combiner(votes)
+276
View File
@@ -0,0 +1,276 @@
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional
import json
from classes.v2.papila_data import PapilaData
@dataclass
class ImportSpec:
id: str
class_name: str
params: Dict[str, Any]
@dataclass
class DataSourceSpec:
node_id: str
label: str
output_type: str
source: Optional[Dict[str, Any]]
source_ref: Optional[Dict[str, Any]]
@dataclass
class TransformSpec:
node_id: str
label: str
transform_type: str
params: Dict[str, Any]
@dataclass
class LoaderSpec:
node_id: str
label: str
input_type: str
input_index: str
input_key: str
output_key: str
transforms: List[TransformSpec]
data_source: Optional[DataSourceSpec]
@dataclass
class TowerSpec:
node_id: str
label: str
tower_type: str
params: Dict[str, Any]
@dataclass
class BridgeSpec:
node_id: str
label: str
method: str
params: Dict[str, Any]
@dataclass
class ClassifierSpec:
node_id: str
label: str
@dataclass
class ConfigAssembly:
raw: Dict[str, Any]
imports: Dict[str, ImportSpec]
data_sources: Dict[str, DataSourceSpec]
transforms: Dict[str, TransformSpec]
loaders: Dict[str, LoaderSpec]
towers: Dict[str, TowerSpec]
bridges: Dict[str, BridgeSpec]
classifiers: Dict[str, ClassifierSpec]
def load_config(path: Path) -> Dict[str, Any]:
payload = json.loads(Path(path).read_text())
if not isinstance(payload, dict):
raise ValueError("Config JSON must be an object.")
return payload
def assemble_config(path: Path) -> ConfigAssembly:
config = load_config(path)
meta = config.get("meta", {})
imports = _build_imports(meta.get("imports", []))
nodes = {node["id"]: node for node in config.get("nodes", [])}
edges = config.get("edges", [])
data_sources: Dict[str, DataSourceSpec] = {}
transforms: Dict[str, TransformSpec] = {}
loaders: Dict[str, LoaderSpec] = {}
towers: Dict[str, TowerSpec] = {}
bridges: Dict[str, BridgeSpec] = {}
classifiers: Dict[str, ClassifierSpec] = {}
for node in nodes.values():
ntype = node.get("type")
if ntype == "data":
data_sources[node["id"]] = DataSourceSpec(
node_id=node["id"],
label=node.get("label", ""),
output_type=node.get("outputType", ""),
source=node.get("source"),
source_ref=node.get("sourceRef"),
)
elif ntype == "transform":
transforms[node["id"]] = TransformSpec(
node_id=node["id"],
label=node.get("label", ""),
transform_type=node.get("transformType", ""),
params=_extract_transform_params(node),
)
elif ntype == "loader":
loaders[node["id"]] = LoaderSpec(
node_id=node["id"],
label=node.get("label", ""),
input_type=node.get("inputType", ""),
input_index=node.get("inputIndex", ""),
input_key=node.get("inputKey", ""),
output_key=node.get("outputKey", ""),
transforms=[],
data_source=None,
)
elif ntype in ("image_tower", "metadata_tower"):
towers[node["id"]] = TowerSpec(
node_id=node["id"],
label=node.get("label", ""),
tower_type=node.get("towerType", "image" if ntype == "image_tower" else "metadata"),
params=_extract_tower_params(node),
)
elif ntype == "bridge":
bridges[node["id"]] = BridgeSpec(
node_id=node["id"],
label=node.get("label", ""),
method=node.get("bridgeMethod", "fusion"),
params=_extract_bridge_params(node),
)
elif ntype == "classifier":
classifiers[node["id"]] = ClassifierSpec(
node_id=node["id"],
label=node.get("label", ""),
)
# attach transforms + data sources to loaders by walking upstream
for loader_id, loader in loaders.items():
chain = _upstream_chain(loader_id, nodes, edges)
for node_id in reversed(chain):
if node_id in transforms:
loader.transforms.append(transforms[node_id])
if node_id in data_sources:
loader.data_source = data_sources[node_id]
return ConfigAssembly(
raw=config,
imports=imports,
data_sources=data_sources,
transforms=transforms,
loaders=loaders,
towers=towers,
bridges=bridges,
classifiers=classifiers,
)
def resolve_imports(assembly: ConfigAssembly) -> Dict[str, Any]:
resolved: Dict[str, Any] = {}
for import_id, spec in assembly.imports.items():
if spec.class_name == "PapilaData":
params = spec.params
resolved[import_id] = PapilaData.from_dirs(
image_dir=params.get("image_dir", "Papila/FundusImages"),
clinical_dir=params.get("clinical_dir", "Papila/ClinicalData"),
label_col=params.get("label_col", "Diagnosis"),
cat_cols=params.get("cat_cols", ["Gender", "Phakic/Pseudophakic"]),
)
else:
raise ValueError(f"Unsupported import class {spec.class_name!r}")
return resolved
def _build_imports(entries: Iterable[Dict[str, Any]]) -> Dict[str, ImportSpec]:
specs: Dict[str, ImportSpec] = {}
for entry in entries or []:
import_id = entry.get("id")
if not import_id:
continue
specs[import_id] = ImportSpec(
id=import_id,
class_name=entry.get("className", ""),
params=entry.get("params", {}) or {},
)
return specs
def _extract_transform_params(node: Dict[str, Any]) -> Dict[str, Any]:
return {
"transformType": node.get("transformType"),
"roiMaskSource": node.get("roiMaskSource"),
"roiScale": node.get("roiScale"),
"roiTargetSize": node.get("roiTargetSize"),
"roiFallback": node.get("roiFallback"),
"centerCropSize": node.get("centerCropSize"),
"jitterHFlip": node.get("jitterHFlip"),
"jitterVFlip": node.get("jitterVFlip"),
"jitterRotation": node.get("jitterRotation"),
"jitterColorEnabled": node.get("jitterColorEnabled"),
"jitterColor": node.get("jitterColor"),
"resizeSize": node.get("resizeSize"),
}
def _extract_tower_params(node: Dict[str, Any]) -> Dict[str, Any]:
if node.get("towerType") == "metadata":
return {
"hidden_dim": node.get("mdHiddenDim"),
"dropout": node.get("mdDropout"),
"use_se": node.get("mdUseSe"),
"se_reduction": node.get("mdSeReduction"),
"se_pre_norm": node.get("mdSePreNorm"),
"freeze_ratio": node.get("mdFreezeRatio"),
}
return {
"backbone": node.get("imageBackbone"),
"freeze_ratio": node.get("imageFreezeRatio"),
"augment": node.get("imageAugment"),
"geometry_dim": node.get("imageGeometryDim"),
"use_se": node.get("imageUseSe"),
"se_reduction": node.get("imageSeReduction"),
"se_pre_norm": node.get("imageSePreNorm"),
}
def _extract_bridge_params(node: Dict[str, Any]) -> Dict[str, Any]:
return {
"fusion_dim": node.get("bridgeFusionDim"),
"use_se": node.get("bridgeUseSe"),
"se_reduction": node.get("bridgeSeReduction"),
"se_pre_norm": node.get("bridgeSePreNorm"),
}
def _edge_from(edge: Dict[str, Any]) -> Optional[str]:
return edge.get("from") or edge.get("source")
def _edge_to(edge: Dict[str, Any]) -> Optional[str]:
return edge.get("to") or edge.get("target")
def _upstream_chain(start_id: str, nodes: Dict[str, Dict[str, Any]], edges: List[Dict[str, Any]]) -> List[str]:
chain: List[str] = []
visited = set()
current = start_id
while True:
if current in visited:
break
visited.add(current)
incoming = [edge for edge in edges if _edge_to(edge) == current]
if not incoming:
break
# prefer first incoming edge for now
current = _edge_from(incoming[0])
if not current:
break
chain.append(current)
node = nodes.get(current)
if node and node.get("type") == "data":
break
return chain
+241
View File
@@ -0,0 +1,241 @@
from __future__ import annotations
from pathlib import Path
from typing import Callable, Dict, Iterable, List, Optional, Tuple
import numpy as np
import pandas as pd
class DataBundle:
"""
Generic, torch-free container for metadata and file/label bookkeeping.
Keeps feature typing, vectorization, and patient-level splits generic.
Dataset-specific preprocessing (e.g., eye canonicalization) should live
in the dataset builder (e.g., papila_builders in v2).
"""
def __init__(
self,
*,
image_dir: str,
clinical_dir: Optional[str] = None,
label_col: str,
patient_col: str = "Patient ID",
cat_cols: Optional[Iterable[str]] = None,
max_unique_for_cat: int = 4,
n_splits: int = 5,
random_seed: int = 42,
filename_template: str = "RET{pid:03d}{eye}.jpg",
image_path_fn: Optional[Callable[[pd.Series], Path]] = None,
) -> None:
self.image_dir = Path(image_dir)
self.label_col = label_col
self.patient_col = patient_col
self.max_unique_for_cat = max_unique_for_cat
self.n_splits = n_splits
self.filename_template = filename_template
self.image_path_fn = image_path_fn
self.clinical_dir = Path(clinical_dir) if clinical_dir else None
# Internal state
self.frames: List[pd.DataFrame] = []
self.df: pd.DataFrame = pd.DataFrame()
self.scalar_cols: List[str] = []
self.cat_cols: List[str] = list(cat_cols) if cat_cols is not None else []
self.scalar_stats: Dict[str, Dict[str, float]] = {}
self.cat_maps: Dict[str, Dict[object, int]] = {}
self.feature_dim: int = 0
self.folds: Dict[int, Dict[str, List[object]]] = {}
self.random_seed = int(random_seed)
# ------------------- Public API -------------------
def add_df(
self,
df: pd.DataFrame,
*,
id_column: Optional[str] = None,
exclude_cols: Optional[Iterable[str]] = None,
) -> None:
"""
Add a dataframe and re-run typing, stats, and K-fold indices.
QC rules:
- Must have patient ID column; if not provided under that name, specify id_column.
"""
df = df.copy()
self._ensure_patient_id(df, id_column)
if self.label_col not in df.columns:
raise ValueError(f"label_col '{self.label_col}' not found in added dataframe")
self.frames.append(df)
self._refresh_master_df(exclude_cols=exclude_cols)
self._infer_or_validate_feature_types(exclude_cols=exclude_cols)
self._compute_numeric_stats()
self._build_cat_maps()
self._compute_feature_dim()
self._build_kfold_indices()
def get_split_ids(self, fold: int) -> Tuple[List[object], List[object]]:
rec = self.folds.get(fold)
if not rec:
raise KeyError(f"Fold {fold} not available. Built folds: {sorted(self.folds.keys())}")
return rec["train_ids"], rec["test_ids"]
def get_split_dfs(self, fold: int) -> Tuple[pd.DataFrame, pd.DataFrame]:
train_ids, test_ids = self.get_split_ids(fold)
train_df = self.df[self.df[self.patient_col].isin(train_ids)].reset_index(drop=True)
test_df = self.df[self.df[self.patient_col].isin(test_ids)].reset_index(drop=True)
return train_df, test_df
def vectorize_row(self, row: pd.Series) -> np.ndarray:
"""Return a numpy feature vector (torch-free)."""
feats: List[float] = []
miss: List[float] = []
# numeric
for col in self.scalar_cols:
v = pd.to_numeric(row.get(col), errors="coerce")
if pd.isna(v):
miss.append(1.0)
v = self.scalar_stats[col]["median"]
else:
miss.append(0.0)
lo = self.scalar_stats[col]["min"]
hi = self.scalar_stats[col]["max"]
feats.append((float(v) - lo) / (hi - lo) if hi > lo else 0.0)
# categorical
for col in self.cat_cols:
mapping = self.cat_maps[col]
one = [0.0] * len(mapping)
key = row.get(col)
one[mapping.get(key, 0)] = 1.0 # 0 is <UNK>
feats.extend(one)
# numeric missing flags
feats.extend(miss)
return np.asarray(feats, dtype=np.float32)
def get_image_path(self, row: pd.Series) -> Path:
if self.image_path_fn is not None:
return Path(self.image_path_fn(row))
pid = int(row[self.patient_col])
eye = row.get("eyeID", "")
if eye in ("OS", "OD"):
eye_str = eye
else:
eye_str = str(eye)
return self.image_dir / self.filename_template.format(pid=pid, eye=eye_str)
def encode_metadata(self, row: pd.Series) -> np.ndarray:
return self.vectorize_row(row)
def get_label(self, row: pd.Series) -> int:
return int(row[self.label_col])
# ------------------- Internal helpers -------------------
def _ensure_patient_id(self, df: pd.DataFrame, id_column: Optional[str]) -> None:
if self.patient_col in df.columns:
return
if id_column and id_column in df.columns:
df.rename(columns={id_column: self.patient_col}, inplace=True)
return
candidates = [
c
for c in df.columns
if c.lower().replace(" ", "") in {"patientid", "patient", "pid"}
]
if len(candidates) == 1:
df.rename(columns={candidates[0]: self.patient_col}, inplace=True)
return
raise ValueError(
f"A '{self.patient_col}' column is required; provide id_column=... if it has a different name."
)
def _refresh_master_df(self, exclude_cols: Optional[Iterable[str]] = None) -> None:
self.df = pd.concat(self.frames, axis=0, ignore_index=True)
if exclude_cols:
self.df = self.df.drop(columns=[c for c in exclude_cols if c in self.df.columns])
def _infer_or_validate_feature_types(self, exclude_cols: Optional[Iterable[str]] = None) -> None:
excluded = set(exclude_cols or []) | {self.label_col, self.patient_col}
feature_candidates = [c for c in self.df.columns if c not in excluded]
cats = set(self.cat_cols) if self.cat_cols else set()
scalars = set()
for c in feature_candidates:
if c in cats:
continue
s = self.df[c]
as_num = pd.to_numeric(s, errors="coerce")
num_missing = as_num.isna().mean()
num_unique = s.dropna().nunique()
if as_num.notna().any() and num_missing < 1.0 and num_unique > self.max_unique_for_cat:
scalars.add(c)
else:
if num_unique <= self.max_unique_for_cat or as_num.isna().mean() > 0.0:
cats.add(c)
else:
scalars.add(c)
self.cat_cols = sorted(cats)
self.scalar_cols = sorted(scalars)
def _compute_numeric_stats(self) -> None:
self.scalar_stats.clear()
for col in self.scalar_cols:
s = pd.to_numeric(self.df[col], errors="coerce")
vals = s.dropna().astype(float).values
if vals.size == 0:
lo, hi, med = 0.0, 1.0, 0.0
else:
lo, hi = float(np.min(vals)), float(np.max(vals))
med = float(np.median(vals))
if hi <= lo:
hi = lo + 1.0
self.scalar_stats[col] = {"min": lo, "max": hi, "median": med}
def _build_cat_maps(self) -> None:
self.cat_maps.clear()
for col in self.cat_cols:
cats = [v for v in self.df[col].dropna().unique().tolist()]
try:
cats = sorted(cats)
except Exception:
pass
mapping = {"<UNK>": 0}
for i, v in enumerate(cats, start=1):
mapping[v] = i
self.cat_maps[col] = mapping
def _compute_feature_dim(self) -> None:
self.feature_dim = len(self.scalar_cols) + sum(len(m) for m in self.cat_maps.values()) + len(self.scalar_cols)
# ------------------- K-fold on unique patients -------------------
def _build_kfold_indices(self) -> None:
pats = self.df[self.patient_col].unique().tolist()
labels_by_pat: Dict[object, object] = {}
for pid, grp in self.df.groupby(self.patient_col):
lab = grp[self.label_col].dropna()
if len(lab) == 0:
labels_by_pat[pid] = 0
else:
labels_by_pat[pid] = lab.mode().iloc[0]
y_pat = np.array([labels_by_pat[p] for p in pats])
try:
from sklearn.model_selection import StratifiedGroupKFold
sgkf = StratifiedGroupKFold(
n_splits=self.n_splits, shuffle=True, random_state=self.random_seed
)
split_iter = sgkf.split(X=pats, y=y_pat, groups=pats)
except Exception:
from sklearn.model_selection import StratifiedKFold
skf = StratifiedKFold(
n_splits=self.n_splits, shuffle=True, random_state=self.random_seed
)
split_iter = skf.split(X=np.zeros(len(pats)), y=y_pat)
self.folds.clear()
for i, (train_idx, test_idx) in enumerate(split_iter):
train_ids = [pats[j] for j in train_idx]
test_ids = [pats[j] for j in test_idx]
self.folds[i] = {"train_ids": train_ids, "test_ids": test_ids}
+57
View File
@@ -0,0 +1,57 @@
from torch.utils.data import Dataset
from PIL import Image
import numpy as np
import torch
class ClinicalDataset(Dataset):
"""Generic dataset wrapping a DataBundle-like instance.
Returns (img_tensor, meta_tensor, label)."""
def __init__(
self,
clinical_data,
img_transform,
meta_transform=None,
image_preprocessor=None,
geometry_provider=None,
geometry_dim: int = 0,
):
self.clinical = clinical_data
self.transform_image = img_transform
self.meta_transform = meta_transform or (lambda x: x)
self.image_preprocessor = image_preprocessor
self.geometry_provider = geometry_provider
self.geometry_dim = geometry_dim if geometry_provider is not None else 0
def __len__(self):
return len(self.clinical.df)
def __getitem__(self, idx: int):
row = self.clinical.df.iloc[idx]
# load & transform image
img_path = self.clinical.get_image_path(row)
orig_img = Image.open(img_path).convert("RGB")
img = orig_img
if self.image_preprocessor is not None:
img = self.image_preprocessor(img, img_path)
img_t = self.transform_image(img)
# encode & transform metadata
meta = self.clinical.encode_metadata(row)
meta_t = self.meta_transform(meta)
# label
label = self.clinical.get_label(row)
if self.geometry_dim > 0:
features = None
if self.geometry_provider is not None and hasattr(self.geometry_provider, "geometry_features"):
features = self.geometry_provider.geometry_features(orig_img, img_path)
if features is None:
geom_vec = torch.zeros(self.geometry_dim, dtype=torch.float32)
else:
features = np.asarray(features, dtype=np.float32)
if features.shape[0] != self.geometry_dim:
geom_vec = torch.zeros(self.geometry_dim, dtype=torch.float32)
else:
geom_vec = torch.from_numpy(features)
return img_t, meta_t, geom_vec, label
return img_t, meta_t, label
+119
View File
@@ -0,0 +1,119 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Iterable, List, Sequence, Tuple, Union
import re
import pandas as pd
@dataclass
class RegexFilter:
pattern: str
flags: int = 0
def apply_paths(self, paths: Sequence[str]) -> Tuple[List[str], List[str]]:
if not self.pattern:
return list(paths), []
try:
regex = re.compile(self.pattern, self.flags)
except re.error as err:
return list(paths), [f'Invalid regex "{self.pattern}": {err}']
filtered = [p for p in paths if regex.search(p)]
return filtered, []
@dataclass
class ColumnFilter:
column: str
operator: str
value: str
case_insensitive: bool = True
def apply_df(self, df: pd.DataFrame) -> Tuple[pd.DataFrame, List[str]]:
warnings: List[str] = []
if not self.column:
return df, ["Column filter missing column name."]
columns = list(df.columns)
col_index = _resolve_column_index(columns, self.column, warnings)
if col_index is None:
return df, warnings
col_name = columns[col_index]
if self.value is None or self.value == "":
return df, [f'Column filter "{self.column}" missing value.']
series = df[col_name]
mask = series.apply(
lambda cell: compare_cell(
cell, self.value, self.operator, case_insensitive=self.case_insensitive
)
)
return df[mask], warnings
FilterSpec = Union[RegexFilter, ColumnFilter]
def apply_regex_filters(paths: Sequence[str], filters: Iterable[RegexFilter]) -> Tuple[List[str], List[str]]:
filtered = list(paths)
warnings: List[str] = []
for filt in filters:
filtered, warn = filt.apply_paths(filtered)
warnings.extend(warn)
return filtered, warnings
def apply_column_filters(df: pd.DataFrame, filters: Iterable[ColumnFilter]) -> Tuple[pd.DataFrame, List[str]]:
filtered = df
warnings: List[str] = []
for filt in filters:
filtered, warn = filt.apply_df(filtered)
warnings.extend(warn)
return filtered, warnings
def compare_cell(cell, raw_value: str, operator: str, case_insensitive: bool = True) -> bool:
cell_str = "" if cell is None else str(cell).strip()
value_str = "" if raw_value is None else str(raw_value).strip()
if case_insensitive:
cell_str = cell_str.lower()
value_str = value_str.lower()
if operator == "=":
return cell_str == value_str
if operator == "!=":
return cell_str != value_str
cell_num = _to_float(cell_str)
value_num = _to_float(value_str)
if cell_num is None or value_num is None:
return False
if operator == ">":
return cell_num > value_num
if operator == ">=":
return cell_num >= value_num
if operator == "<":
return cell_num < value_num
if operator == "<=":
return cell_num <= value_num
return False
def _resolve_column_index(columns: Sequence[str], column: str, warnings: List[str]) -> int | None:
try:
return columns.index(column)
except ValueError:
lower = column.lower()
matches = [idx for idx, col in enumerate(columns) if str(col).lower() == lower]
if matches:
if len(matches) > 1:
warnings.append(
f'Column "{column}" matched multiple headers; using "{columns[matches[0]]}".'
)
return matches[0]
warnings.append(f'Column "{column}" not found.')
return None
def _to_float(value: str) -> float | None:
try:
return float(value)
except (TypeError, ValueError):
return None
+128
View File
@@ -0,0 +1,128 @@
from __future__ import annotations
import csv
import json
import logging
from pathlib import Path
from typing import Optional
DEFAULT_OPTIONAL_EPOCH_COLS = [
"pct_fused",
"pct_img",
"pct_md",
"phase",
"se_mean",
"se_std",
"se_pct_lt_0.2",
"se_pct_gt_0.8",
"holdout_loss",
"holdout_acc_fused",
"holdout_acc_img",
"holdout_acc_md",
"holdout_auc_fused",
"holdout_auc_img",
"holdout_auc_md",
"best_monitor",
"best_so_far",
"best_epoch",
"early_best_so_far",
"early_bad_epochs",
"early_improved",
"early_monitor",
"holdout_best_monitor",
"holdout_best_so_far",
"holdout_best_epoch",
]
class HypertowerLogger:
"""
Shared logging utility for V2 tower workflows.
- train.log line logging
- epoch_log.csv row logging with stable header
- lightweight JSON/array artifact helpers
"""
def __init__(
self,
*,
run_dir: Path,
train_log_path: Optional[Path] = None,
epoch_log_path: Optional[Path] = None,
logger_name: Optional[str] = None,
) -> None:
self.run_dir = Path(run_dir).resolve()
self.run_dir.mkdir(parents=True, exist_ok=True)
self.train_log_path = Path(train_log_path) if train_log_path else (self.run_dir / "train.log")
self.epoch_log_path = Path(epoch_log_path) if epoch_log_path else (self.run_dir / "epoch_log.csv")
self._logger_name = logger_name or f"hypertower.{id(self)}"
self.logger = logging.getLogger(self._logger_name)
self.logger.setLevel(logging.INFO)
self.logger.handlers = []
fh = logging.FileHandler(str(self.train_log_path))
fh.setFormatter(logging.Formatter("%(asctime)s - %(message)s"))
self.logger.addHandler(fh)
self.logger.propagate = False
self._epoch_log_fp = None
self._epoch_log_writer = None
self._epoch_log_fields: list[str] | None = None
def info(self, msg: str) -> None:
self.logger.info(msg)
def warning(self, msg: str) -> None:
self.logger.warning(msg)
def error(self, msg: str) -> None:
self.logger.error(msg)
def write_epoch_row(
self,
row: dict,
*,
path: str | Path | None = None,
optional_cols: Optional[list[str]] = None,
) -> None:
optional = optional_cols if optional_cols is not None else DEFAULT_OPTIONAL_EPOCH_COLS
if self._epoch_log_writer is None:
fieldnames = list(dict.fromkeys([*row.keys(), *optional]))
target_path = Path(path) if path is not None else self.epoch_log_path
target_path.parent.mkdir(parents=True, exist_ok=True)
self._epoch_log_fp = open(target_path, "w", newline="", encoding="utf-8")
self._epoch_log_writer = csv.DictWriter(self._epoch_log_fp, fieldnames=fieldnames)
self._epoch_log_writer.writeheader()
self._epoch_log_fields = fieldnames
assert self._epoch_log_fields is not None
assert self._epoch_log_writer is not None
assert self._epoch_log_fp is not None
for key in self._epoch_log_fields:
row.setdefault(key, None)
self._epoch_log_writer.writerow({k: row.get(k) for k in self._epoch_log_fields})
self._epoch_log_fp.flush()
def write_json(self, path: str | Path, payload: dict) -> None:
target = Path(path)
if not target.is_absolute():
target = self.run_dir / target
target.parent.mkdir(parents=True, exist_ok=True)
target.write_text(json.dumps(payload, indent=2), encoding="utf-8")
def close(self) -> None:
if self._epoch_log_fp is not None:
try:
self._epoch_log_fp.close()
except Exception:
pass
self._epoch_log_fp = None
self._epoch_log_writer = None
self._epoch_log_fields = None
for handler in list(self.logger.handlers):
try:
handler.close()
except Exception:
pass
self.logger.removeHandler(handler)
+165
View File
@@ -0,0 +1,165 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Callable, Optional
from torch.utils.data import DataLoader
from .network_manager import LoaderBundle, PatientSplit
from .slot_dataset import SlotDataset, slot_collate
from .profiles.base import SlotDescriptor, SimpleDatasetProfile
def _default_slot_descriptors(patient_col: str, label_col: str) -> dict[str, SlotDescriptor]:
return {
"id_1": SlotDescriptor(
key="id_1",
kind="id",
description=f"Patient identifier column ({patient_col})",
required=True,
shape_hint="scalar",
),
"label_1": SlotDescriptor(
key="label_1",
kind="label",
description=f"Label column ({label_col})",
required=True,
shape_hint="scalar",
),
"image_1": SlotDescriptor(
key="image_1",
kind="image",
description="Primary image slot",
required=False,
shape_hint="HWC or CHW",
),
"matrix_1": SlotDescriptor(
key="matrix_1",
kind="matrix",
description="Primary matrix slot",
required=False,
shape_hint="[feature_dim]",
),
}
def _row_to_sample(
row: Any,
*,
clinical: Any,
patient_col: str,
label_col: str,
) -> dict[str, Any]:
return {
"id_1": row[patient_col],
"label_1": row[label_col],
"image_1": clinical.get_image_path(row) if hasattr(clinical, "get_image_path") else None,
"matrix_1": clinical.vectorize_row(row) if hasattr(clinical, "vectorize_row") else None,
}
@dataclass
class SlotLoaderFactory:
"""
Generic loader factory that emits dict batches keyed by slot names.
"""
image_transform: Optional[Callable] = None
matrix_transform: Optional[Callable] = None
num_workers: int = 0
def build(
self,
*,
clinical: Any,
split: PatientSplit,
args: Any,
fold: int,
profile: Optional[Any] = None,
) -> LoaderBundle:
batch_size = int(getattr(args, "batch_size", 8))
slot_desc = self._resolve_slot_descriptors(clinical=clinical, profile=profile)
train_samples = self._build_samples(split.train, clinical, profile, slot_desc)
val_samples = self._build_samples(split.val, clinical, profile, slot_desc)
holdout_samples = (
self._build_samples(split.holdout, clinical, profile, slot_desc)
if split.holdout is not None
else None
)
train_loader = DataLoader(
SlotDataset(
train_samples,
slot_desc,
image_transform=self.image_transform,
matrix_transform=self.matrix_transform,
),
batch_size=batch_size,
shuffle=True,
num_workers=self.num_workers,
collate_fn=slot_collate,
)
val_loader = DataLoader(
SlotDataset(
val_samples,
slot_desc,
image_transform=self.image_transform,
matrix_transform=self.matrix_transform,
),
batch_size=batch_size,
shuffle=False,
num_workers=self.num_workers,
collate_fn=slot_collate,
)
holdout_loader = None
if holdout_samples is not None:
holdout_loader = DataLoader(
SlotDataset(
holdout_samples,
slot_desc,
image_transform=self.image_transform,
matrix_transform=self.matrix_transform,
),
batch_size=batch_size,
shuffle=False,
num_workers=self.num_workers,
collate_fn=slot_collate,
)
return LoaderBundle(train=train_loader, val=val_loader, holdout=holdout_loader)
@staticmethod
def _resolve_slot_descriptors(
*,
clinical: Any,
profile: Optional[Any],
) -> dict[str, SlotDescriptor]:
if profile is not None and hasattr(profile, "slot_descriptors"):
return profile.slot_descriptors()
patient_col = getattr(clinical, "patient_col", "Patient ID")
label_col = getattr(clinical, "label_col", "Diagnosis")
return _default_slot_descriptors(patient_col, label_col)
@staticmethod
def _build_samples(
df,
clinical: Any,
profile: Optional[Any],
slot_desc: dict[str, SlotDescriptor],
) -> list[dict[str, Any]]:
if df is None or df.empty:
return []
if profile is not None and hasattr(profile, "build_samples"):
return profile.build_samples(df=df, clinical=clinical)
patient_col = getattr(profile, "patient_col", None) if profile is not None else None
label_col = getattr(profile, "label_col", None) if profile is not None else None
pcol = patient_col or "Patient ID"
lcol = label_col or getattr(clinical, "label_col", "Diagnosis")
samples = []
for _, row in df.iterrows():
sample = _row_to_sample(row, clinical=clinical, patient_col=pcol, label_col=lcol)
for key in slot_desc.keys():
sample.setdefault(key, None)
samples.append(sample)
return samples
+148
View File
@@ -0,0 +1,148 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Callable, Optional
import torch
from torch import nn
from classes.v2.bridges import Bridge, VoteBridge
from classes.v2.towers import ImageTower, MDTower
from .config_builder import ConfigAssembly
from .transforms import build_transform_chain
@dataclass
class V2ModelBundle:
image_tower: Optional[ImageTower]
metadata_tower: Optional[MDTower]
bridge: Optional[nn.Module]
classifier: Optional[nn.Module]
image_transform: Optional[Callable]
matrix_transform: Optional[Callable]
def build_model_bundle(
assembly: ConfigAssembly,
clinical: Any,
*,
device: Optional[torch.device] = None,
strict: bool = True,
) -> V2ModelBundle:
"""
Build torch modules and input transforms from a V2 config assembly.
"""
image_tower_spec = _pick_tower(assembly, "image")
md_tower_spec = _pick_tower(assembly, "metadata")
bridge_spec = _pick_bridge(assembly)
image_loader = _pick_loader(assembly, input_type="image")
clinical_core = getattr(clinical, "clinical", clinical)
num_classes = _infer_num_classes(clinical)
img_tower = None
if image_tower_spec is not None:
img_tower = ImageTower(
backbone=image_tower_spec.params.get("backbone", "efficientnet_b0"),
freeze_ratio=float(image_tower_spec.params.get("freeze_ratio", 0.0) or 0.0),
use_se=bool(image_tower_spec.params.get("use_se", False)),
se_reduction=int(image_tower_spec.params.get("se_reduction", 16) or 16),
se_pre_norm=bool(image_tower_spec.params.get("se_pre_norm", True)),
augment=bool(image_tower_spec.params.get("augment", True)),
geometry_dim=int(image_tower_spec.params.get("geometry_dim", 0) or 0),
)
if device is not None:
img_tower = img_tower.to(device)
md_tower = None
if md_tower_spec is not None:
md_tower = MDTower(
clinical_core,
hidden_dim=int(md_tower_spec.params.get("hidden_dim", 128) or 128),
dropout=float(md_tower_spec.params.get("dropout", 0.1) or 0.1),
use_se=bool(md_tower_spec.params.get("use_se", False)),
se_reduction=int(md_tower_spec.params.get("se_reduction", 16) or 16),
se_pre_norm=bool(md_tower_spec.params.get("se_pre_norm", True)),
)
if device is not None:
md_tower = md_tower.to(device)
bridge = None
if bridge_spec is not None and img_tower is not None and md_tower is not None:
if bridge_spec.method == "consensus":
bridge = VoteBridge(num_classes=num_classes)
else:
bridge = Bridge(
img_dim=img_tower.out_dim,
meta_dim=md_tower.out_dim,
num_classes=num_classes,
fusion_dim=int(bridge_spec.params.get("fusion_dim", 256) or 256),
mode="fused",
use_se=bool(bridge_spec.params.get("use_se", True)),
se_reduction=int(bridge_spec.params.get("se_reduction", 16) or 16),
se_pre_norm=bool(bridge_spec.params.get("se_pre_norm", True)),
)
if device is not None:
bridge = bridge.to(device)
classifier = None
if assembly.classifiers:
classifier = nn.Identity()
if device is not None:
classifier = classifier.to(device)
image_transform = None
if image_loader is not None and image_tower_spec is not None:
image_transform = build_transform_chain(
image_loader.transforms,
backbone_name=image_tower_spec.params.get("backbone", "efficientnet_b0"),
augment=bool(image_tower_spec.params.get("augment", True)),
strict=strict,
)
return V2ModelBundle(
image_tower=img_tower,
metadata_tower=md_tower,
bridge=bridge,
classifier=classifier,
image_transform=image_transform,
matrix_transform=None,
)
def _pick_tower(assembly: ConfigAssembly, tower_type: str):
matches = [tower for tower in assembly.towers.values() if tower.tower_type == tower_type]
if not matches:
return None
if len(matches) > 1:
raise ValueError(f"Multiple {tower_type} towers found; only one is supported for now.")
return matches[0]
def _pick_bridge(assembly: ConfigAssembly):
if not assembly.bridges:
return None
if len(assembly.bridges) > 1:
raise ValueError("Multiple bridges found; only one is supported for now.")
return next(iter(assembly.bridges.values()))
def _pick_loader(assembly: ConfigAssembly, input_type: str):
matches = [loader for loader in assembly.loaders.values() if loader.input_type == input_type]
if not matches:
return None
if len(matches) > 1:
raise ValueError(f"Multiple loaders with input_type={input_type!r} found.")
return matches[0]
def _infer_num_classes(clinical: Any) -> int:
df = getattr(clinical, "df", None)
label_col = getattr(clinical, "label_col", None)
if df is None and hasattr(clinical, "clinical"):
df = clinical.clinical.df
label_col = clinical.clinical.label_col
if df is None or label_col is None or label_col not in df.columns:
return 2
return int(df[label_col].dropna().nunique())
+200
View File
@@ -0,0 +1,200 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Optional, Protocol
import pandas as pd
@dataclass
class PatientSplit:
"""Patient-disjoint split definition for a fold."""
train: pd.DataFrame
val: pd.DataFrame
holdout: Optional[pd.DataFrame] = None
@dataclass
class LoaderBundle:
"""All loaders needed by a training run."""
train: Any
val: Any
holdout: Optional[Any] = None
@dataclass
class FoldResult:
"""Normalized fold output from trainer implementations."""
fold: int
metrics: dict[str, Any]
artifacts: dict[str, Any]
class SplitManager(Protocol):
def build_plans(
self,
*,
clinical: Any,
args: Any,
profile: Optional[Any] = None,
) -> list[PatientSplit]:
...
class GraphFactory(Protocol):
def build(
self,
*,
clinical: Any,
args: Any,
fold: int,
profile: Optional[Any] = None,
) -> Any:
...
class LoaderFactory(Protocol):
def build(
self,
*,
clinical: Any,
split: PatientSplit,
args: Any,
fold: int,
profile: Optional[Any] = None,
) -> LoaderBundle:
...
class Trainer(Protocol):
def fit(
self,
*,
graph: Any,
loaders: LoaderBundle,
args: Any,
fold: int,
profile: Optional[Any] = None,
) -> FoldResult:
...
class NetworkManager:
"""
V2 orchestration entrypoint.
This class is intentionally small and modular:
- split policy is delegated to a SplitManager
- graph assembly is delegated to a GraphFactory
- dataloaders are delegated to a LoaderFactory
- train/eval/checkpoint lifecycle is delegated to a Trainer
"""
def __init__(
self,
*,
clinical: Any,
args: Any,
split_manager: SplitManager,
graph_factory: GraphFactory,
loader_factory: LoaderFactory,
trainer: Trainer,
profile: Optional[Any] = None,
) -> None:
self.clinical = clinical
self.args = args
self.split_manager = split_manager
self.graph_factory = graph_factory
self.loader_factory = loader_factory
self.trainer = trainer
self.profile = profile
self._split_plans: Optional[list[PatientSplit]] = None
def run_fold(self, fold: int) -> FoldResult:
plans = self._get_split_plans()
if fold < 0 or fold >= len(plans):
raise IndexError(f"Requested fold {fold} but only {len(plans)} fold plans are available")
split = plans[fold]
self._validate_patient_disjointness(split)
self._validate_labels(split)
graph = self.graph_factory.build(
clinical=self.clinical,
args=self.args,
fold=fold,
profile=self.profile,
)
loaders = self.loader_factory.build(
clinical=self.clinical,
split=split,
args=self.args,
fold=fold,
profile=self.profile,
)
return self.trainer.fit(
graph=graph,
loaders=loaders,
args=self.args,
fold=fold,
profile=self.profile,
)
def run_all_folds(self, n_splits: Optional[int] = None) -> list[FoldResult]:
plans = self._get_split_plans()
max_folds = len(plans)
if n_splits is None:
n = max_folds
else:
n = int(n_splits)
if n < 1:
raise ValueError("n_splits must be >= 1")
if n > max_folds:
raise ValueError(f"Requested {n} folds but only {max_folds} fold plans are available")
return [self.run_fold(fold) for fold in range(n)]
def _get_split_plans(self) -> list[PatientSplit]:
if self._split_plans is None:
self._split_plans = self.split_manager.build_plans(
clinical=self.clinical,
args=self.args,
profile=self.profile,
)
if not self._split_plans:
raise ValueError("SplitManager returned no fold plans")
return self._split_plans
def _validate_patient_disjointness(self, split: PatientSplit) -> None:
train_ids = self._patient_ids(split.train)
val_ids = self._patient_ids(split.val)
holdout_ids = self._patient_ids(split.holdout) if split.holdout is not None else set()
if train_ids & val_ids:
overlap = sorted(train_ids & val_ids)[:10]
raise ValueError(f"Patient leakage between train/val: {overlap}")
if train_ids & holdout_ids:
overlap = sorted(train_ids & holdout_ids)[:10]
raise ValueError(f"Patient leakage between train/holdout: {overlap}")
if val_ids & holdout_ids:
overlap = sorted(val_ids & holdout_ids)[:10]
raise ValueError(f"Patient leakage between val/holdout: {overlap}")
def _validate_labels(self, split: PatientSplit) -> None:
label_col = getattr(self.clinical, "label_col", None)
if not label_col:
return
for name, df in (("train", split.train), ("val", split.val), ("holdout", split.holdout)):
if df is None:
continue
if label_col not in df.columns:
raise ValueError(f"{name} split is missing label column {label_col!r}")
@staticmethod
def _patient_ids(df: Optional[pd.DataFrame]) -> set[Any]:
if df is None or df.empty:
return set()
if "Patient ID" not in df.columns:
raise ValueError("Split dataframes must include 'Patient ID'")
return set(df["Patient ID"].tolist())
+152
View File
@@ -0,0 +1,152 @@
from __future__ import annotations
from typing import Dict, List
import numpy as np
import pandas as pd
from classes.v2.data_bundle import DataBundle
# ---- Pachymetry → IOP correction (per PAPILA Table 3) ----
_PACHY_TABLE: Dict[int, int] = {
475: +5,
485: +4,
495: +4,
505: +3,
515: +2,
525: +1,
535: +1,
545: 0,
555: -1,
565: -1,
575: -2,
585: -3,
595: -4,
605: -4,
615: -5,
}
_PACHY_KEYS = np.array(sorted(_PACHY_TABLE.keys()))
def _nearest_pachy_key(x: float) -> int:
idx = int(np.argmin(np.abs(_PACHY_KEYS - float(x))))
return int(_PACHY_KEYS[idx])
def _pick_iop(row: pd.Series) -> float:
"""Prefer Pneumatic, else Perkins; may return NaN."""
raw = row["Pneumatic"] if not pd.isna(row.get("Pneumatic", np.nan)) else row.get("Perkins", np.nan)
return float(raw) if not pd.isna(raw) else np.nan
def _correct_iop(raw_iop: float, pachy: float) -> float:
"""Return corrected IOP using nearest pachymetry bin; if pachy missing, return raw."""
if pd.isna(raw_iop):
return np.nan
if pd.isna(pachy):
return float(raw_iop)
key = _nearest_pachy_key(float(pachy))
return float(raw_iop) + float(_PACHY_TABLE[key])
def _apply_iop_and_drop_md(df: pd.DataFrame) -> pd.DataFrame:
"""Add IOP_raw/IOP_corr and drop VF_MD if present (in-place safe)."""
df["IOP_raw"] = df.apply(_pick_iop, axis=1)
pachy = df.get("Pachymetry", pd.Series(np.nan, index=df.index))
df["IOP_corr"] = [
_correct_iop(r, p) for r, p in zip(df["IOP_raw"].values, pachy.values)
]
if "VF_MD" in df.columns:
df.drop(columns=["VF_MD"], inplace=True)
return df
def _canonicalize_eye_column(df: pd.DataFrame) -> None:
if "eyeID" in df.columns:
src = "eyeID"
else:
src = None
for c in df.columns:
if "eye" in c.lower():
src = c
break
if src is None:
df["eyeID"] = "OS"
return
s = df[src]
def norm(v):
if pd.isna(v):
return None
x = str(v).strip().upper()
if x in {"OS", "L", "LEFT", "0"}:
return "OS"
if x in {"OD", "R", "RIGHT", "1"}:
return "OD"
try:
num = int(float(x))
return "OD" if num % 2 == 1 else "OS"
except Exception:
return None
mapped = s.map(norm)
uniq = {u for u in mapped.dropna().unique().tolist()}
if not uniq.issubset({"OS", "OD"}):
raise ValueError(f"eyeID must be binary; found values {sorted(uniq)}")
df["eyeID"] = mapped.fillna("OS")
def build_papila_data(
*,
image_dir: str,
clinical_dir: str,
label_col: str,
cat_cols: List[str],
n_splits: int = 5,
random_seed: int = 42,
) -> DataBundle:
"""
Build a DataBundle for PAPILA with dataset-specific preprocessing:
- load OD/OS Excel sheets
- normalize Patient ID
- canonicalize eyeID
- compute IOP_raw / IOP_corr, drop VF_MD
- build feature typing & folds
"""
bundle = DataBundle(
image_dir=image_dir,
clinical_dir=clinical_dir,
label_col=label_col,
patient_col="Patient ID",
cat_cols=cat_cols,
n_splits=n_splits,
random_seed=random_seed,
filename_template="RET{pid:03d}{eye}.jpg",
)
od = pd.read_excel(f"{clinical_dir}/patient_data_od.xlsx", header=1)
od["eyeID"] = "OD"
os = pd.read_excel(f"{clinical_dir}/patient_data_os.xlsx", header=1)
os["eyeID"] = "OS"
for frame in (od, os):
if "Patient ID" not in frame.columns and "ID" in frame.columns:
frame.rename(columns={"ID": "Patient ID"}, inplace=True)
frame["Patient ID"] = frame["Patient ID"].astype(str).str.extract(r"(\d+)")[0].astype(int)
_canonicalize_eye_column(frame)
bundle.add_df(od, id_column="ID")
bundle.add_df(os, id_column="ID")
for i in range(len(bundle.frames)):
bundle.frames[i] = _apply_iop_and_drop_md(bundle.frames[i])
bundle._refresh_master_df()
bundle._infer_or_validate_feature_types()
bundle._compute_numeric_stats()
bundle._build_cat_maps()
bundle._compute_feature_dim()
bundle._build_kfold_indices()
return bundle
+61
View File
@@ -0,0 +1,61 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Iterable, Optional
import pandas as pd
from classes.v2.data_bundle import DataBundle
from classes.v2.papila_builders import build_papila_data
@dataclass
class PapilaData:
"""
V2-friendly wrapper around the DataBundle pipeline.
Keeps all formatting/normalization behavior from build_papila_clinical,
but exposes a minimal surface area for the V2 engine.
"""
clinical: DataBundle
patient_col: str = "Patient ID"
@property
def df(self) -> pd.DataFrame:
return self.clinical.df
@property
def label_col(self) -> str:
return self.clinical.label_col
@property
def feature_dim(self) -> int:
return self.clinical.feature_dim
def get_image_path(self, row: pd.Series):
return self.clinical.get_image_path(row)
def vectorize_row(self, row: pd.Series):
return self.clinical.vectorize_row(row)
@classmethod
def from_dirs(
cls,
*,
image_dir: str,
clinical_dir: str,
label_col: str,
cat_cols: Iterable[str],
n_splits: int = 5,
random_seed: int = 42,
) -> "PapilaData":
clinical = build_papila_data(
image_dir=image_dir,
clinical_dir=clinical_dir,
label_col=label_col,
cat_cols=list(cat_cols),
n_splits=n_splits,
random_seed=random_seed,
)
return cls(clinical=clinical)
+10
View File
@@ -0,0 +1,10 @@
from .base import DatasetProfile, SimpleDatasetProfile, SlotDescriptor
from .papila import PapilaProfile, build_papila_profile
__all__ = [
"DatasetProfile",
"SimpleDatasetProfile",
"SlotDescriptor",
"PapilaProfile",
"build_papila_profile",
]
+53
View File
@@ -0,0 +1,53 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Protocol, Any
import pandas as pd
@dataclass(frozen=True)
class SlotDescriptor:
"""
Metadata for a generic batch slot key (e.g., image_1, matrix_1).
"""
key: str
kind: str
description: str
required: bool = True
shape_hint: str | None = None
class DatasetProfile(Protocol):
"""
Dataset-specific wiring that stays outside the generic V2 engine.
"""
name: str
patient_col: str
label_col: str
def slot_descriptors(self) -> dict[str, SlotDescriptor]:
...
def semantic_aliases(self) -> dict[str, str]:
...
def build_samples(self, *, df: pd.DataFrame, clinical: Any) -> list[dict[str, Any]]:
...
@dataclass(frozen=True)
class SimpleDatasetProfile:
name: str
patient_col: str
label_col: str
slots: dict[str, SlotDescriptor]
aliases: dict[str, str]
def slot_descriptors(self) -> dict[str, SlotDescriptor]:
return dict(self.slots)
def semantic_aliases(self) -> dict[str, str]:
return dict(self.aliases)
+155
View File
@@ -0,0 +1,155 @@
from __future__ import annotations
import pandas as pd
from dataclasses import dataclass
from .base import SimpleDatasetProfile, SlotDescriptor
@dataclass(frozen=True)
class PapilaProfile(SimpleDatasetProfile):
sample_mode: str = "patient" # "patient" | "eye"
def build_samples(self, *, df: pd.DataFrame, clinical) -> list[dict[str, object]]:
samples: list[dict[str, object]] = []
patient_col = self.patient_col
label_col = self.label_col
mode = (self.sample_mode or "patient").lower()
if mode not in {"patient", "eye"}:
raise ValueError(f"Unsupported sample_mode '{self.sample_mode}'. Expected 'patient' or 'eye'.")
if mode == "eye":
for _, row in df.iterrows():
pid = row[patient_col]
label = row[label_col]
image_1 = clinical.get_image_path(row) if hasattr(clinical, "get_image_path") else None
matrix_1 = clinical.vectorize_row(row) if hasattr(clinical, "vectorize_row") else None
samples.append(
{
"id_1": pid,
"label_1": label,
"image_1": image_1,
"matrix_1": matrix_1,
}
)
return samples
for pid, grp in df.groupby(patient_col):
label_series = grp[label_col]
if label_series.empty:
continue
mode_vals = label_series.mode()
label = mode_vals.iloc[0] if not mode_vals.empty else label_series.iloc[0]
def _row_for_eye(eye: str):
if "eyeID" not in grp.columns:
return None
match = grp[grp["eyeID"].astype(str).str.upper() == eye]
if match.empty:
return None
return match.iloc[0]
row_od = _row_for_eye("OD")
row_os = _row_for_eye("OS")
row_any = grp.iloc[0]
image_1 = clinical.get_image_path(row_od) if row_od is not None else None
image_2 = clinical.get_image_path(row_os) if row_os is not None else None
matrix_1 = clinical.vectorize_row(row_od) if row_od is not None else None
matrix_2 = clinical.vectorize_row(row_os) if row_os is not None else None
if image_1 is None and hasattr(clinical, "get_image_path"):
image_1 = clinical.get_image_path(row_any)
if matrix_1 is None and hasattr(clinical, "vectorize_row"):
matrix_1 = clinical.vectorize_row(row_any)
samples.append(
{
"id_1": pid,
"label_1": label,
"image_1": image_1,
"image_2": image_2,
"matrix_1": matrix_1,
"matrix_2": matrix_2,
}
)
return samples
def build_papila_profile(
*,
patient_col: str = "Patient ID",
label_col: str = "Diagnosis",
sample_mode: str = "patient",
) -> PapilaProfile:
"""
PAPILA-specific semantic map for generic V2 slot keys.
The engine remains slot-based (image_1/image_2/matrix_1/...).
PAPILA meaning is captured here so run config stays dataset-local.
"""
slots = {
"id_1": SlotDescriptor(
key="id_1",
kind="id",
description=f"Patient identifier column ({patient_col})",
required=True,
shape_hint="scalar",
),
"label_1": SlotDescriptor(
key="label_1",
kind="label",
description=f"Diagnosis label column ({label_col})",
required=True,
shape_hint="scalar",
),
"image_1": SlotDescriptor(
key="image_1",
kind="image",
description="Fundus image slot 1 (PAPILA: OD / right eye)",
required=False,
shape_hint="HWC or CHW",
),
"image_2": SlotDescriptor(
key="image_2",
kind="image",
description="Fundus image slot 2 (PAPILA: OS / left eye)",
required=False,
shape_hint="HWC or CHW",
),
"matrix_1": SlotDescriptor(
key="matrix_1",
kind="matrix",
description="Clinical metadata feature vector",
required=False,
shape_hint="[feature_dim]",
),
"matrix_2": SlotDescriptor(
key="matrix_2",
kind="matrix",
description="Optional auxiliary tabular vector (reserved for experiments)",
required=False,
shape_hint="[feature_dim_2]",
),
}
aliases = {
"id_1": "patient_id",
"label_1": "diagnosis",
"image_1": "od_fundus",
"image_2": "os_fundus",
"matrix_1": "clinical_metadata",
"matrix_2": "aux_metadata",
}
return PapilaProfile(
name="papila",
patient_col=patient_col,
label_col=label_col,
slots=slots,
aliases=aliases,
sample_mode=sample_mode,
)
+98
View File
@@ -0,0 +1,98 @@
from __future__ import annotations
from typing import Any, Callable, Optional
from pathlib import Path
from PIL import Image
import numpy as np
import torch
from torch.utils.data import Dataset
from torchvision import transforms
from .profiles.base import SlotDescriptor
def slot_collate(batch: list[dict[str, Any]]) -> dict[str, Any]:
if not batch:
return {}
keys = batch[0].keys()
out: dict[str, Any] = {}
for key in keys:
vals = [item.get(key) for item in batch]
if all(isinstance(v, torch.Tensor) for v in vals):
try:
out[key] = torch.stack(vals, dim=0)
except Exception:
out[key] = vals
else:
out[key] = vals
return out
class SlotDataset(Dataset):
"""
Dataset that yields dicts of slot-keyed values.
Sample records are expected to be dicts with keys matching slot descriptors.
Image slots accept filesystem paths; matrix slots accept array-like values.
"""
def __init__(
self,
samples: list[dict[str, Any]],
slot_descriptors: dict[str, SlotDescriptor],
*,
image_transform: Optional[Callable[[Image.Image], torch.Tensor]] = None,
matrix_transform: Optional[Callable[[Any], torch.Tensor]] = None,
image_preprocessor: Optional[Callable[..., Image.Image]] = None,
) -> None:
self.samples = samples
self.slot_descriptors = slot_descriptors
self.image_transform = image_transform or transforms.ToTensor()
self.matrix_transform = matrix_transform or self._default_matrix_transform
self.image_preprocessor = image_preprocessor
def __len__(self) -> int:
return len(self.samples)
def __getitem__(self, idx: int) -> dict[str, Any]:
record = self.samples[idx]
out: dict[str, Any] = {}
for key, desc in self.slot_descriptors.items():
val = record.get(key)
if desc.kind == "image":
out[key] = self._load_image(val, required=desc.required)
elif desc.kind == "matrix":
out[key] = self._load_matrix(val, required=desc.required)
else:
out[key] = val
return out
def _load_image(self, value: Any, *, required: bool) -> Optional[torch.Tensor]:
if value is None:
if required:
raise ValueError("Missing required image slot")
return None
path = Path(value)
img = Image.open(path).convert("RGB")
if self.image_preprocessor is not None:
try:
img = self.image_preprocessor(img, path)
except TypeError:
img = self.image_preprocessor(img)
return self.image_transform(img)
def _load_matrix(self, value: Any, *, required: bool) -> Optional[torch.Tensor]:
if value is None:
if required:
raise ValueError("Missing required matrix slot")
return None
return self.matrix_transform(value)
@staticmethod
def _default_matrix_transform(value: Any) -> torch.Tensor:
if isinstance(value, torch.Tensor):
return value.float()
if isinstance(value, np.ndarray):
return torch.from_numpy(value.astype(np.float32, copy=False))
return torch.as_tensor(value, dtype=torch.float32)
+197
View File
@@ -0,0 +1,197 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Iterable, Optional
import numpy as np
import pandas as pd
from sklearn.model_selection import KFold, StratifiedKFold
from .network_manager import PatientSplit
@dataclass(frozen=True)
class SplitPlan:
train_patient_ids: set[Any]
val_patient_ids: set[Any]
holdout_patient_ids: set[Any]
def build_patient_split_plans(
patient_ids: Iterable[Any],
patient_labels: Iterable[Any],
*,
n_splits: int,
seed: int,
holdout_per_class: int = 0,
holdout_seed: int = 123,
) -> list[SplitPlan]:
"""
Core vector-based splitter.
Inputs are one row per patient:
- patient_ids: unique patient IDs
- patient_labels: one label per patient
"""
ids = np.asarray(list(patient_ids))
labels = np.asarray(list(patient_labels))
if ids.ndim != 1 or labels.ndim != 1:
raise ValueError("patient_ids and patient_labels must be 1D arrays")
if ids.size != labels.size:
raise ValueError(f"Length mismatch: ids={ids.size}, labels={labels.size}")
if ids.size == 0:
raise ValueError("No patients available for splitting")
if len(set(ids.tolist())) != ids.size:
raise ValueError("patient_ids must be unique (one label per patient)")
if n_splits < 2:
raise ValueError("n_splits must be >= 2")
holdout_ids: set[Any] = set()
if holdout_per_class > 0:
rng = np.random.default_rng(holdout_seed)
for label in np.unique(labels):
idx = np.where(labels == label)[0]
if idx.size == 0:
continue
n = min(holdout_per_class, idx.size)
chosen = rng.choice(idx, size=n, replace=False)
holdout_ids.update(ids[chosen].tolist())
keep_mask = ~np.isin(ids, list(holdout_ids))
cv_ids = ids[keep_mask]
cv_labels = labels[keep_mask]
if cv_ids.size < n_splits:
raise ValueError(
f"Not enough patients ({cv_ids.size}) for n_splits={n_splits} after holdout removal"
)
use_stratified = _can_stratify(cv_labels, n_splits)
if use_stratified:
splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
splits = list(splitter.split(cv_ids, cv_labels))
else:
splitter = KFold(n_splits=n_splits, shuffle=True, random_state=seed)
splits = list(splitter.split(cv_ids))
plans: list[SplitPlan] = []
for train_idx, val_idx in splits:
plans.append(
SplitPlan(
train_patient_ids=set(cv_ids[train_idx].tolist()),
val_patient_ids=set(cv_ids[val_idx].tolist()),
holdout_patient_ids=set(holdout_ids),
)
)
return plans
class PatientFirstSplitManager:
"""
Patient-level splitter for V2.
Behavior:
- Optional binary filtering happens first (labels in {0,1} only).
- Optional holdout is sampled at the patient level (never per-eye rows).
- K-fold split is built on remaining patients.
- Returned dataframes contain all rows for each selected patient.
"""
def __init__(
self,
*,
patient_col: str = "Patient ID",
label_col: Optional[str] = None,
) -> None:
self.patient_col = patient_col
self.label_col = label_col
def build_plans(
self,
*,
clinical: Any,
args: Any,
profile: Optional[Any] = None,
) -> list[PatientSplit]:
profile_label_col = getattr(profile, "label_col", None) if profile is not None else None
profile_patient_col = getattr(profile, "patient_col", None) if profile is not None else None
patient_col = profile_patient_col or self.patient_col
label_col = self.label_col or profile_label_col or getattr(clinical, "label_col", None)
if label_col is None:
raise ValueError("Could not resolve label column from SplitManager or clinical.label_col")
if not hasattr(clinical, "df"):
raise ValueError("Clinical object must expose a dataframe at .df")
df_full = clinical.df.copy()
self._validate_columns(df_full, label_col, patient_col=patient_col)
eval_mode = str(getattr(args, "eval_mode", "multiclass")).lower()
if eval_mode == "binary":
df_full = df_full[df_full[label_col].isin([0, 1])].reset_index(drop=True)
holdout_per_class = int(getattr(args, "holdout_per_class", 0) or 0)
holdout_seed = int(getattr(args, "holdout_seed", 123))
n_splits = int(getattr(args, "n_splits", 5))
fold_seed = int(getattr(args, "fold_seed", 42))
patient_table = self._patient_label_table(df_full, label_col, patient_col=patient_col)
plans = build_patient_split_plans(
patient_ids=patient_table[patient_col].to_numpy(),
patient_labels=patient_table["_label"].to_numpy(),
n_splits=n_splits,
seed=fold_seed,
holdout_per_class=holdout_per_class,
holdout_seed=holdout_seed,
)
out: list[PatientSplit] = []
for plan in plans:
train_df = (
df_full[df_full[patient_col].isin(plan.train_patient_ids)]
.reset_index(drop=True)
)
val_df = (
df_full[df_full[patient_col].isin(plan.val_patient_ids)]
.reset_index(drop=True)
)
holdout_df = None
if plan.holdout_patient_ids:
holdout_df = (
df_full[df_full[patient_col].isin(plan.holdout_patient_ids)]
.reset_index(drop=True)
)
out.append(PatientSplit(train=train_df, val=val_df, holdout=holdout_df))
return out
def _validate_columns(self, df: pd.DataFrame, label_col: str, patient_col: Optional[str] = None) -> None:
pcol = patient_col or self.patient_col
if pcol not in df.columns:
raise ValueError(f"Missing required patient column: {pcol!r}")
if label_col not in df.columns:
raise ValueError(f"Missing required label column: {label_col!r}")
def _patient_label_table(
self,
df: pd.DataFrame,
label_col: str,
patient_col: Optional[str] = None,
) -> pd.DataFrame:
pcol = patient_col or self.patient_col
grouped = (
df.groupby(pcol, as_index=False)[label_col]
.agg(lambda x: x.mode().iloc[0] if not x.mode().empty else x.iloc[0])
.rename(columns={label_col: "_label"})
.sort_values(pcol)
.reset_index(drop=True)
)
if grouped.empty:
raise ValueError("No patients available for splitting")
return grouped
def _can_stratify(labels: np.ndarray, n_splits: int) -> bool:
if labels.size == 0:
return False
unique, counts = np.unique(labels, return_counts=True)
if len(unique) < 2:
return False
return bool(np.all(counts >= n_splits))
+279
View File
@@ -0,0 +1,279 @@
from __future__ import annotations
import math
from typing import Optional
import torch
from torch import nn
from torchvision import transforms
from classes.backbones import BACKBONES, list_names, load_backbone_weights
from classes.SE_attention import SEBlock
from classes.v2.data_bundle import DataBundle
def build_backbone(name: str, freeze_ratio: float = 0.0, augment: bool = True):
"""
Operational builder:
- instantiate with DEFAULT weights
- strip classifier → features
- apply ratio-based freezing over coarse blocks
- return (model, out_dim, transform)
"""
key = (name or "").lower()
if key not in BACKBONES:
raise ValueError(f"Unsupported backbone '{name}'. Valid options: {list_names()}")
spec = BACKBONES[key]
m = spec.ctor(weights=spec.weights_default)
out_dim, m = spec.strip(m)
load_backbone_weights(key, m)
# transforms: use the weights mean/std, but keep your augmentation pipeline
mean = getattr(spec.weights_default, "meta", {}).get("mean", (0.485, 0.456, 0.406))
std = getattr(spec.weights_default, "meta", {}).get("std", (0.229, 0.224, 0.225))
crop = 299 if key == "inception_v3" else 224
if augment:
transform = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(crop),
transforms.RandomHorizontalFlip(),
transforms.RandomVerticalFlip(),
transforms.RandomRotation(15),
transforms.ColorJitter(0.1, 0.1, 0.1, 0.05),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
]
)
else:
transform = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(crop),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
]
)
# ratio-based freezing: freeze earliest floor(N * freeze_ratio) blocks
fr = max(0.0, min(1.0, float(freeze_ratio)))
blocks = spec.blocks(m)
n = len(blocks)
freeze_n = int(math.floor(n * fr))
for b in blocks[:freeze_n]:
for p in b.parameters():
p.requires_grad = False
return m, out_dim, transform
class ImageTower(nn.Module):
"""
Vision backbone → pooled features.
- backbone: one of list_names() (default 'efficientnet_b0')
- always DEFAULT torchvision weights
- freeze_ratio ∈ [0,1] freezes earliest floor(N*freeze_ratio) blocks
- returns [N, out_dim] features from backbone forward
"""
def __init__(
self,
backbone: str = "efficientnet_b0",
freeze_ratio: float = 0.0,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
augment: bool = True,
geometry_dim: int = 0,
):
super().__init__()
self.backbone, base_dim, self.transform = build_backbone(
backbone, freeze_ratio, augment=augment
)
self._name = backbone
# Keep ordered blocks for dynamic freezing/thawing
key = (self._name or "").lower()
self._spec = BACKBONES[key]
self._blocks = self._spec.blocks(self.backbone)
# Optional tower-level SE over the final feature vector
self.base_dim = base_dim
self.geometry_dim = max(0, int(geometry_dim))
self.out_dim = self.base_dim + self.geometry_dim
self.tower_ln = nn.LayerNorm(self.base_dim) if se_pre_norm else nn.Identity()
self.tower_se = (
SEBlock(self.base_dim, reduction=se_reduction, residual=True)
if use_se
else None
)
def forward(
self, x: torch.Tensor, geometry: Optional[torch.Tensor] = None
) -> torch.Tensor:
y = self.backbone(x)
# sanity: pooled features, not logits
assert y.dim() == 2 and y.size(1) == self.base_dim, (
f"Expected features [N,{self.base_dim}], got {tuple(y.shape)}"
)
if self.tower_se is not None:
y, _ = self.tower_se(self.tower_ln(y))
if self.geometry_dim > 0:
if geometry is None or geometry.numel() == 0:
geom = torch.zeros(
y.size(0), self.geometry_dim, device=y.device, dtype=y.dtype
)
else:
if geometry.dim() == 1:
geom = geometry.unsqueeze(0)
else:
geom = geometry
geom = geom.to(device=y.device, dtype=y.dtype)
if geom.size(0) != y.size(0):
raise ValueError(
f"Geometry batch size mismatch: {geom.size(0)} vs {y.size(0)}"
)
if geom.size(1) != self.geometry_dim:
raise ValueError(
f"Expected geometry dim {self.geometry_dim}, got {geom.size(1)}"
)
y = torch.cat([y, geom], dim=1)
return y
def set_freeze_ratio(self, ratio: float):
"""Dynamically freeze earliest floor(N*ratio) backbone blocks."""
r = max(0.0, min(1.0, float(ratio)))
n = len(self._blocks)
freeze_n = int(math.floor(n * r))
# Unfreeze all first
for b in self._blocks:
for p in b.parameters():
p.requires_grad = True
# Freeze earliest blocks
for b in self._blocks[:freeze_n]:
for p in b.parameters():
p.requires_grad = False
class SiameseImageTower(nn.Module):
"""
Shared-weight bilateral image tower.
Runs OD and OS images through a single shared backbone, then returns
cat([f_mean, f_delta]) where:
f_mean = (f_od + f_os) / 2 -- shared bilateral representation
f_delta = f_od - f_os -- asymmetry, signed OD-relative
out_dim = 2 * backbone_out_dim
When x_os is None (single-eye fallback):
f_mean = f_od
f_delta = zeros
so the module degrades gracefully when only one eye is available.
The shared backbone means both eyes contribute to every gradient update,
effectively doubling the training signal for the visual pathway without
doubling parameters.
"""
def __init__(
self,
backbone: str = "efficientnet_b0",
freeze_ratio: float = 0.0,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
augment: bool = True,
):
super().__init__()
self._tower = ImageTower(
backbone=backbone,
freeze_ratio=freeze_ratio,
use_se=use_se,
se_reduction=se_reduction,
se_pre_norm=se_pre_norm,
augment=augment,
geometry_dim=0,
)
self.out_dim = self._tower.out_dim * 2
self.transform = self._tower.transform
def forward(
self,
x_od: torch.Tensor,
x_os: Optional[torch.Tensor] = None,
) -> torch.Tensor:
f_od = self._tower(x_od)
if x_os is None:
f_mean = f_od
f_delta = torch.zeros_like(f_od)
else:
f_os = self._tower(x_os)
f_mean = (f_od + f_os) * 0.5
f_delta = f_od - f_os
return torch.cat([f_mean, f_delta], dim=1)
def set_freeze_ratio(self, ratio: float) -> None:
"""Delegates to the shared inner tower."""
self._tower.set_freeze_ratio(ratio)
class MDTower(nn.Module):
"""MLP over DataBundle.vectorize_row outputs (convert to torch inside tower)."""
def __init__(
self,
clinical_data: DataBundle,
hidden_dim: int = 128,
dropout: float = 0.1,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
):
super().__init__()
self.feature_dim = clinical_data.feature_dim
self.out_dim = hidden_dim
# two-block MLP so we can optionally freeze/thaw per block
self.block0 = nn.Sequential(
nn.Linear(self.feature_dim, hidden_dim),
nn.LayerNorm(hidden_dim),
nn.ReLU(inplace=True),
nn.Dropout(dropout),
)
self.block1 = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(inplace=True),
)
self.net = nn.Sequential(self.block0, self.block1)
self.tower_ln = nn.LayerNorm(hidden_dim) if se_pre_norm else nn.Identity()
self.tower_se = (
SEBlock(hidden_dim, reduction=se_reduction, residual=True)
if use_se
else None
)
def forward(self, meta_np_or_torch) -> torch.Tensor:
if isinstance(meta_np_or_torch, torch.Tensor):
x = meta_np_or_torch
else:
x = torch.as_tensor(meta_np_or_torch, dtype=torch.float32)
h = self.net(x)
if self.tower_se is not None:
h, _ = self.tower_se(self.tower_ln(h))
return h
def set_freeze_ratio(self, ratio: float):
"""Optionally freeze earliest blocks of the MLP."""
r = max(0.0, min(1.0, float(ratio)))
# Unfreeze all
for p in self.block0.parameters():
p.requires_grad = True
for p in self.block1.parameters():
p.requires_grad = True
# Freeze earliest blocks based on ratio threshold
if r >= 0.5:
for p in self.block0.parameters():
p.requires_grad = False
if r >= 1.0:
for p in self.block1.parameters():
p.requires_grad = False
+315
View File
@@ -0,0 +1,315 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable, Iterable, Optional, Tuple, Union
import numpy as np
from PIL import Image
from torchvision import transforms
from classes.backbones import BACKBONES
IMAGENET_MEAN: Tuple[float, float, float] = (0.485, 0.456, 0.406)
IMAGENET_STD: Tuple[float, float, float] = (0.229, 0.224, 0.225)
@dataclass
class ImageTransformConfig:
"""
Mirrors the hypertower v1 preprocessing:
- Resize(256)
- CenterCrop(crop)
- Optional augmentations (H/V flip, rotation, color jitter)
- ToTensor + Normalize(mean/std)
"""
crop_size: int = 224
resize_size: int = 256
mean: Tuple[float, float, float] = IMAGENET_MEAN
std: Tuple[float, float, float] = IMAGENET_STD
augment: bool = True
rotation_deg: int = 15
color_jitter: Tuple[float, float, float, float] = (0.1, 0.1, 0.1, 0.05)
hflip: bool = True
vflip: bool = True
def build(self) -> transforms.Compose:
ops = [
transforms.Resize(self.resize_size),
transforms.CenterCrop(self.crop_size),
]
if self.augment:
if self.hflip:
ops.append(transforms.RandomHorizontalFlip())
if self.vflip:
ops.append(transforms.RandomVerticalFlip())
if self.rotation_deg:
ops.append(transforms.RandomRotation(self.rotation_deg))
if self.color_jitter:
ops.append(transforms.ColorJitter(*self.color_jitter))
ops.extend(
[
transforms.ToTensor(),
transforms.Normalize(mean=self.mean, std=self.std),
]
)
return transforms.Compose(ops)
def backbone_transform_config(backbone_name: str, augment: bool = True) -> ImageTransformConfig:
"""
Build a transform config that matches v1 ImageTower/backbone preprocessing.
Uses DEFAULT weights mean/std and InceptionV3 crop size when relevant.
"""
key = (backbone_name or "").lower()
if key not in BACKBONES:
raise ValueError(f"Unsupported backbone '{backbone_name}'.")
spec = BACKBONES[key]
mean = getattr(spec.weights_default, "meta", {}).get("mean", IMAGENET_MEAN)
std = getattr(spec.weights_default, "meta", {}).get("std", IMAGENET_STD)
crop = 299 if key == "inception_v3" else 224
return ImageTransformConfig(crop_size=crop, mean=mean, std=std, augment=augment)
def build_backbone_transform(backbone_name: str, augment: bool = True) -> transforms.Compose:
return backbone_transform_config(backbone_name, augment=augment).build()
def build_imagenet_transform(augment: bool = True, crop_size: int = 224) -> transforms.Compose:
return ImageTransformConfig(crop_size=crop_size, augment=augment).build()
@dataclass
class ResizeTransform:
size: Union[int, Tuple[int, int]] = 256
interpolation: int = Image.BILINEAR
def __post_init__(self) -> None:
self._op = transforms.Resize(self.size, interpolation=self.interpolation)
def __call__(self, image: Image.Image) -> Image.Image:
return self._op(image)
@dataclass
class CenterCropTransform:
size: Union[int, Tuple[int, int]] = 224
def __post_init__(self) -> None:
self._op = transforms.CenterCrop(self.size)
def __call__(self, image: Image.Image) -> Image.Image:
return self._op(image)
class UnetMaskProvider:
"""
Placeholder for a UNet-powered mask provider.
This will be replaced once a UNet tower is wired in.
"""
def __call__(self, image: Image.Image, image_path: Optional[str] = None):
raise NotImplementedError("UNet mask provider is not wired yet.")
@dataclass
class ROICropTransform:
"""
Crop an image using a binary mask (GT or UNet).
Expects a mask of the same spatial size as the image; nonzero pixels are ROI.
"""
mask_source: str = "gt" # "gt" | "unet"
mask_provider: Optional[Callable[[Image.Image, Optional[str]], np.ndarray]] = None
scale: float = 2.5
target_size: Optional[Tuple[int, int]] = (224, 224)
fallback_to_original: bool = True
def __post_init__(self) -> None:
if self.mask_source not in {"gt", "unet"}:
raise ValueError(f"mask_source must be 'gt' or 'unet', got '{self.mask_source}'.")
def __call__(
self,
image: Image.Image,
mask: Optional[Union[np.ndarray, Image.Image]] = None,
image_path: Optional[str] = None,
) -> Image.Image:
resolved_mask = mask
if resolved_mask is None and self.mask_provider is not None:
resolved_mask = self.mask_provider(image, image_path)
if resolved_mask is None:
if self.fallback_to_original:
return image
raise ValueError("ROI crop requested but no mask provided.")
mask_arr = (
np.asarray(resolved_mask)
if not isinstance(resolved_mask, Image.Image)
else np.array(resolved_mask)
)
if mask_arr.ndim == 3:
mask_arr = mask_arr[..., 0]
mask_arr = mask_arr > 0
if not np.any(mask_arr):
return image if self.fallback_to_original else image
ys, xs = np.where(mask_arr)
y_min, y_max = ys.min(), ys.max()
x_min, x_max = xs.min(), xs.max()
cx = (x_min + x_max) / 2.0
cy = (y_min + y_max) / 2.0
width = (x_max - x_min + 1)
height = (y_max - y_min + 1)
size = max(width, height) * float(self.scale)
left = int(round(cx - size / 2))
right = int(round(cx + size / 2))
upper = int(round(cy - size / 2))
lower = int(round(cy + size / 2))
left = max(0, left)
upper = max(0, upper)
right = min(image.width, right)
lower = min(image.height, lower)
crop = image.crop((left, upper, right, lower))
if self.target_size is not None:
crop = crop.resize(self.target_size, Image.BILINEAR)
return crop
@dataclass
class JitterBundleTransform:
"""
Augmentations bundle: flips, rotation, color jitter.
"""
hflip: bool = True
vflip: bool = True
rotation_deg: int = 15
color_jitter: Optional[Tuple[float, float, float, float]] = (0.1, 0.1, 0.1, 0.05)
def __post_init__(self) -> None:
ops = []
if self.hflip:
ops.append(transforms.RandomHorizontalFlip())
if self.vflip:
ops.append(transforms.RandomVerticalFlip())
if self.rotation_deg:
ops.append(transforms.RandomRotation(self.rotation_deg))
if self.color_jitter:
ops.append(transforms.ColorJitter(*self.color_jitter))
self._op = transforms.Compose(ops) if ops else None
def __call__(self, image: Image.Image) -> Image.Image:
if self._op is None:
return image
return self._op(image)
TRANSFORM_REGISTRY = {
"resize": ResizeTransform,
"roi_crop": ROICropTransform,
"center_crop": CenterCropTransform,
"jitter_bundle": JitterBundleTransform,
}
def _parse_color_jitter(value: Optional[Union[str, Iterable[float]]]) -> Optional[Tuple[float, float, float, float]]:
if value is None:
return None
if isinstance(value, str):
parts = [p.strip() for p in value.split(",") if p.strip()]
if not parts:
return None
try:
nums = [float(p) for p in parts]
except ValueError:
return None
if len(nums) == 1:
return (nums[0], nums[0], nums[0], nums[0])
if len(nums) >= 4:
return (nums[0], nums[1], nums[2], nums[3])
return tuple(nums + [nums[-1]] * (4 - len(nums))) # pad to length 4
try:
vals = list(value)
except TypeError:
return None
if not vals:
return None
vals = [float(v) for v in vals]
if len(vals) == 1:
return (vals[0], vals[0], vals[0], vals[0])
if len(vals) >= 4:
return (vals[0], vals[1], vals[2], vals[3])
return tuple(vals + [vals[-1]] * (4 - len(vals)))
def build_transform_chain(
transform_specs: Iterable[object],
*,
backbone_name: str,
augment: bool = True,
mask_provider: Optional[Callable[[Image.Image, Optional[str]], np.ndarray]] = None,
strict: bool = True,
) -> transforms.Compose:
"""
Build an image transform pipeline from a list of transform specs plus the
standard ToTensor + Normalize steps. This mirrors the V1 preprocessing
but uses the explicit transform nodes from config.
"""
ops: list[Callable[[Image.Image], Image.Image]] = []
for spec in transform_specs:
transform_type = getattr(spec, "transform_type", None)
params = getattr(spec, "params", None)
if transform_type is None and isinstance(spec, dict):
transform_type = spec.get("transformType") or spec.get("transform_type")
params = spec
params = params or {}
if transform_type == "resize":
size = params.get("resizeSize", 256)
ops.append(ResizeTransform(size=size))
elif transform_type == "center_crop":
size = params.get("centerCropSize", 224)
ops.append(CenterCropTransform(size=size))
elif transform_type == "jitter_bundle":
if not augment:
continue
jitter = JitterBundleTransform(
hflip=bool(params.get("jitterHFlip", True)),
vflip=bool(params.get("jitterVFlip", True)),
rotation_deg=int(params.get("jitterRotation", 15) or 0),
color_jitter=_parse_color_jitter(params.get("jitterColor"))
if params.get("jitterColorEnabled", True)
else None,
)
ops.append(jitter)
elif transform_type == "roi_crop":
roi = ROICropTransform(
mask_source=params.get("roiMaskSource", "gt"),
mask_provider=mask_provider,
scale=float(params.get("roiScale", 2.5)),
target_size=(int(params.get("roiTargetSize", 224)), int(params.get("roiTargetSize", 224)))
if params.get("roiTargetSize") is not None
else None,
fallback_to_original=bool(params.get("roiFallback", True)),
)
if roi.mask_provider is None and roi.mask_source == "unet":
if strict:
raise ValueError("ROI crop requires a mask provider for 'unet' source.")
ops.append(roi)
else:
if strict:
raise ValueError(f"Unsupported transform type: {transform_type!r}")
# Always end with tensor + normalize, using backbone defaults
cfg = backbone_transform_config(backbone_name, augment=augment)
ops.extend(
[
transforms.ToTensor(),
transforms.Normalize(mean=cfg.mean, std=cfg.std),
]
)
return transforms.Compose(ops)
File diff suppressed because it is too large Load Diff