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
+43
View File
@@ -0,0 +1,43 @@
#!/usr/bin/env bash
# Quick smoke-test for run_multifold: runs two 1-epoch configs.
set -euo pipefail
ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
cd "$ROOT_DIR"
MANIFEST="manifest.csv"
IMAGENET_WEIGHTS="models/unet_segmenter/norm_imagenet/best.pt"
if [[ ! -f "$MANIFEST" ]]; then
echo "Missing $MANIFEST; run scripts/main/refuge/build_manifest.py first." >&2
exit 1
fi
if [[ ! -f "$IMAGENET_WEIGHTS" ]]; then
echo "Missing $IMAGENET_WEIGHTS; train the imagenet-normalized UNet first." >&2
exit 1
fi
COMMON_ARGS=(
--backbone resnet50
--fusion-mode fused
--epochs 1
--batch-size 4
--img-crop-manifest "$MANIFEST"
--img-crop-weights "$IMAGENET_WEIGHTS"
--img-crop-normalize imagenet
--shortname smoketest
--holdout-per-class 12
)
echo "[smoketest] Binary eval, fused head"
python scripts/run_multifold.py \
"${COMMON_ARGS[@]}" \
--eval_mode binary \
--run-id smoketest_binary
echo "→ Results under analysis_data/smoketest/smoketest_binary"
echo "[smoketest] Multiclass eval, fused head"
python scripts/run_multifold.py \
"${COMMON_ARGS[@]}" \
--eval_mode multiclass \
--run-id smoketest_multiclass
echo "→ Results under analysis_data/smoketest/smoketest_multiclass"
@@ -0,0 +1,798 @@
from __future__ import annotations
import argparse
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict, List, Optional
import sys
import random
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
import numpy as np
import torch
from torch import nn
from classes.frontend import Multifold
from classes.bridge import Bridge, VoteBridge
from classes.dataset import ClinicalDataset
from classes.hypertower import _ClinicalView
from classes.image_tower import ImageTower
from classes.md_tower import MDTower
from classes.papila_builders import build_papila_clinical
from classes.v2 import (
PatientSplit,
SlotLoaderFactory,
SlotDataset,
slot_collate,
assemble_config,
build_model_bundle,
build_papila_profile,
resolve_imports,
)
from classes.v2.split_manager import PatientFirstSplitManager
def build_v1_defaults() -> Dict[str, Any]:
parser = Multifold.build_parser()
args = parser.parse_args([])
clinical = build_papila_clinical(
image_dir=args.image_dir,
clinical_dir=args.clinical_dir,
label_col=args.label_col,
cat_cols=list(args.cat_cols),
n_splits=args.n_splits,
random_seed=args.fold_seed,
)
return {
"args": args,
"clinical": clinical,
"image_dir": args.image_dir,
"clinical_dir": args.clinical_dir,
"label_col": args.label_col,
"cat_cols": list(args.cat_cols),
"image_transform": {
"resize": 256,
"center_crop": 224,
"hflip": True,
"vflip": True,
"rotation": 15,
"color_jitter": (0.1, 0.1, 0.1, 0.05),
},
"image_tower": {
"backbone": args.backbone,
"freeze_ratio": args.freeze_ratio,
"augment": args.img_augment,
"geometry_dim": 0,
"use_se": False, # se_where default is bridge
"se_reduction": args.se_reduction_tower,
"se_pre_norm": args.se_pre_norm_tower,
},
"md_tower": {
"hidden_dim": 128,
"dropout": 0.1,
"use_se": False,
"se_reduction": args.se_reduction_tower,
"se_pre_norm": args.se_pre_norm_tower,
"freeze_ratio": 0.0,
},
"bridge": {
"method": "fusion" if args.fusion_mode == "fused" else "consensus",
"fusion_dim": 256,
"use_se": args.use_se,
"se_reduction": args.se_reduction,
"se_pre_norm": args.se_pre_norm,
},
}
def build_v2_from_config(path: Path) -> Dict[str, Any]:
assembly = assemble_config(path)
imports = resolve_imports(assembly)
if not imports:
raise ValueError("Config did not include any imports.")
clinical = next(iter(imports.values()))
image_loader = _find_loader(assembly, input_type="image")
matrix_loader = _find_loader(assembly, input_type="matrix")
return {
"assembly": assembly,
"clinical": clinical,
"image_loader": image_loader,
"matrix_loader": matrix_loader,
"image_transform_chain": [t.transform_type for t in image_loader.transforms],
}
def _find_loader(assembly, input_type: str):
matches = [
loader for loader in assembly.loaders.values() if loader.input_type == input_type
]
if not matches:
raise ValueError(f"No loader with input_type={input_type!r} found in config.")
if len(matches) > 1:
raise ValueError(f"Multiple loaders with input_type={input_type!r} found.")
return matches[0]
def compare_configs(v1: Dict[str, Any], v2: Dict[str, Any]) -> List[str]:
diffs: List[str] = []
# data sources
v1_rows, v1_cols = v1["clinical"].df.shape
v2_rows, v2_cols = v2["clinical"].df.shape
if v1_rows != v2_rows or v1_cols != v2_cols:
diffs.append(
f"Clinical DF shape mismatch: v1={v1_rows}x{v1_cols}, v2={v2_rows}x{v2_cols}"
)
# loader presence
if not v2.get("image_loader"):
diffs.append("Missing image loader in v2 config.")
if not v2.get("matrix_loader"):
diffs.append("Missing metadata loader in v2 config.")
# transform chain expectations
expected_chain = ["resize", "center_crop", "jitter_bundle"]
if v2.get("image_transform_chain") != expected_chain:
diffs.append(
f"Image transform chain mismatch: v1 expects {expected_chain}, v2 has {v2.get('image_transform_chain')}"
)
# image tower settings
v1_img = v1["image_tower"]
v2_img = _extract_tower(assembly=v2["assembly"], tower_type="image")
_compare_dict(diffs, "ImageTower", v1_img, v2_img)
# metadata tower settings
v1_md = v1["md_tower"]
v2_md = _extract_tower(assembly=v2["assembly"], tower_type="metadata")
_compare_dict(diffs, "MDTower", v1_md, v2_md)
# bridge settings
v2_bridge = _extract_bridge(v2["assembly"])
_compare_dict(diffs, "Bridge", v1["bridge"], v2_bridge)
if not v2["assembly"].classifiers:
diffs.append("Missing classifier node in v2 config.")
# splits
v1_train, v1_val = v1["clinical"].get_split_dfs(0)
sm = PatientFirstSplitManager()
args = SimpleNamespace(
n_splits=v1["args"].n_splits,
fold_seed=v1["args"].fold_seed,
holdout_per_class=v1["args"].holdout_per_class,
holdout_seed=v1["args"].holdout_seed,
eval_mode=v1["args"].eval_mode,
)
splits = sm.build_plans(clinical=v2["clinical"], args=args, profile=None)
v2_train = splits[0].train
v2_val = splits[0].val
if len(v1_train) != len(v2_train) or len(v1_val) != len(v2_val):
diffs.append(
f"Split sizes mismatch: v1 train/val={len(v1_train)}/{len(v1_val)}, "
f"v2 train/val={len(v2_train)}/{len(v2_val)}"
)
return diffs
def _extract_tower(*, assembly, tower_type: str) -> Dict[str, Any]:
towers = [
tower for tower in assembly.towers.values() if tower.tower_type == tower_type
]
if not towers:
raise ValueError(f"No {tower_type} tower found in v2 config.")
if len(towers) > 1:
raise ValueError(f"Multiple {tower_type} towers found in v2 config.")
return towers[0].params
def _extract_bridge(assembly) -> Dict[str, Any]:
if not assembly.bridges:
raise ValueError("No bridge node found in v2 config.")
if len(assembly.bridges) > 1:
raise ValueError("Multiple bridge nodes found in v2 config.")
bridge = next(iter(assembly.bridges.values()))
payload = dict(bridge.params)
payload["method"] = bridge.method
return payload
def _compare_dict(diffs: List[str], label: str, v1: Dict[str, Any], v2: Dict[str, Any]) -> None:
for key, v1_val in v1.items():
v2_val = v2.get(key)
if isinstance(v1_val, tuple):
v1_val = list(v1_val)
if isinstance(v2_val, tuple):
v2_val = list(v2_val)
if v1_val != v2_val:
diffs.append(f"{label} mismatch for {key}: v1={v1_val} v2={v2_val}")
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument(
"--config",
type=Path,
default=Path("hypertower_v2_config.json"),
help="Path to v2 config JSON",
)
parser.add_argument("--samples", type=int, default=8, help="Number of samples to compare")
parser.add_argument("--seed", type=int, default=1234, help="Seed used for deterministic comparisons")
parser.add_argument(
"--image-compare",
choices=["shape", "value"],
default="value",
help="Compare image tensors by shape only or by value",
)
parser.add_argument(
"--no-data-compare",
action="store_true",
help="Skip the data/loader comparison step",
)
parser.add_argument(
"--sample-mode",
choices=["eye", "patient"],
default="eye",
help="Sample mode for V2 loaders (eye-level or patient-level).",
)
parser.add_argument("--train-epochs", type=int, default=2, help="Epochs to run in train comparison.")
parser.add_argument("--train-folds", type=int, default=2, help="Folds to run in train comparison.")
parser.add_argument("--train-batch-size", type=int, default=8, help="Batch size for train comparison.")
parser.add_argument("--max-batches", type=int, default=10, help="Max batches per epoch (train/val).")
parser.add_argument("--loss-tol", type=float, default=0.5, help="Tolerance for loss diffs.")
parser.add_argument("--acc-tol", type=float, default=0.15, help="Tolerance for accuracy diffs.")
parser.add_argument(
"--device",
choices=["auto", "cpu", "cuda"],
default="cpu",
help="Device to use for training comparison.",
)
parser.add_argument(
"--no-train-compare",
action="store_true",
help="Skip the training comparison step.",
)
args = parser.parse_args()
v1 = build_v1_defaults()
v2 = build_v2_from_config(args.config)
diffs = compare_configs(v1, v2)
if diffs:
print("Differences detected:")
for diff in diffs:
print(f"- {diff}")
return 1
print("V1 vs V2 config comparison: OK (settings and loaders match).")
if not args.no_data_compare:
data_diffs = compare_initial_data(
v1,
v2,
samples=args.samples,
seed=args.seed,
compare_mode=args.image_compare,
)
if data_diffs:
print("Differences detected in initial data:")
for diff in data_diffs:
print(f"- {diff}")
return 1
print("Initial data comparison: OK (image/meta/label inputs match).")
if not args.no_train_compare:
train_diffs = compare_training_runs(
v1,
v2,
epochs=args.train_epochs,
folds=args.train_folds,
batch_size=args.train_batch_size,
max_batches=args.max_batches,
seed=args.seed,
loss_tol=args.loss_tol,
acc_tol=args.acc_tol,
sample_mode=args.sample_mode,
device=args.device,
)
if train_diffs:
print("Differences detected in training comparison:")
for diff in train_diffs:
print(f"- {diff}")
return 1
print("Training comparison: OK (metrics within tolerance).")
return 0
def _seed_all(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def _build_v1_modules(v1: Dict[str, Any]) -> Dict[str, Any]:
args = v1["args"]
clinical = v1["clinical"]
img_tower = ImageTower(
backbone=args.backbone,
freeze_ratio=args.freeze_ratio,
use_se=False,
se_reduction=args.se_reduction_tower,
se_pre_norm=args.se_pre_norm_tower,
augment=args.img_augment,
geometry_dim=0,
)
md_tower = MDTower(
clinical,
hidden_dim=128,
dropout=0.1,
use_se=False,
se_reduction=args.se_reduction_tower,
se_pre_norm=args.se_pre_norm_tower,
)
bridge = None
if args.fusion_mode == "vote":
bridge = VoteBridge(num_classes=args.num_classes)
else:
bridge = Bridge(
img_dim=img_tower.out_dim,
meta_dim=md_tower.out_dim,
num_classes=args.num_classes,
fusion_dim=256,
mode="fused",
use_se=args.use_se,
se_reduction=args.se_reduction,
se_pre_norm=args.se_pre_norm,
)
return {"image_tower": img_tower, "metadata_tower": md_tower, "bridge": bridge}
def compare_initial_data(
v1: Dict[str, Any],
v2: Dict[str, Any],
*,
samples: int = 8,
seed: int = 1234,
compare_mode: str = "value",
) -> List[str]:
diffs: List[str] = []
v1_modules = _build_v1_modules(v1)
v2_modules = build_model_bundle(v2["assembly"], v2["clinical"])
# Build consistent train split for both datasets
v1_train, _ = v1["clinical"].get_split_dfs(0)
split = PatientSplit(train=v1_train, val=v1_train.iloc[:0], holdout=None)
# V1 dataset
v1_view = _ClinicalView(v1["clinical"], v1_train)
v1_ds = ClinicalDataset(v1_view, v1_modules["image_tower"].transform)
# V2 dataset
if v2_modules.image_transform is None:
diffs.append("V2 image transform could not be built from config.")
return diffs
loader_factory = SlotLoaderFactory(image_transform=v2_modules.image_transform)
v2_loaders = loader_factory.build(
clinical=v2["clinical"],
split=split,
args=SimpleNamespace(batch_size=1),
fold=0,
profile=None,
)
v2_ds = v2_loaders.train.dataset
total = min(samples, len(v1_ds), len(v2_ds))
for idx in range(total):
_seed_all(seed + idx)
v1_item = v1_ds[idx]
_seed_all(seed + idx)
v2_item = v2_ds[idx]
if len(v1_item) == 4:
v1_img, v1_meta, _, v1_label = v1_item
else:
v1_img, v1_meta, v1_label = v1_item
v2_img = v2_item.get("image_1")
v2_meta = v2_item.get("matrix_1")
v2_label = v2_item.get("label_1")
if v2_label is None or int(v2_label) != int(v1_label):
diffs.append(f"Label mismatch at idx {idx}: v1={int(v1_label)} v2={v2_label}")
if v2_meta is None:
diffs.append(f"Missing v2 metadata at idx {idx}")
else:
if not torch.allclose(v1_meta, v2_meta, atol=1e-6, rtol=0.0):
max_diff = float((v1_meta - v2_meta).abs().max().item())
diffs.append(f"Metadata mismatch at idx {idx}: max_abs_diff={max_diff:.6f}")
if v2_img is None:
diffs.append(f"Missing v2 image at idx {idx}")
else:
if tuple(v1_img.shape) != tuple(v2_img.shape):
diffs.append(
f"Image shape mismatch at idx {idx}: v1={tuple(v1_img.shape)} v2={tuple(v2_img.shape)}"
)
elif compare_mode == "value":
max_diff = float((v1_img - v2_img).abs().max().item())
if max_diff > 1e-5:
diffs.append(f"Image tensor mismatch at idx {idx}: max_abs_diff={max_diff:.6f}")
return diffs
def compare_training_runs(
v1: Dict[str, Any],
v2: Dict[str, Any],
*,
epochs: int,
folds: int,
batch_size: int,
max_batches: int,
seed: int,
loss_tol: float,
acc_tol: float,
sample_mode: str,
device: str,
) -> List[str]:
diffs: List[str] = []
if sample_mode != "eye":
diffs.append("Training compare only supports sample_mode='eye' for parity with v1.")
return diffs
if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
_seed_all(seed)
# Build profile for V2 dataset
profile = build_papila_profile(
patient_col="Patient ID",
label_col=v1["args"].label_col,
sample_mode=sample_mode,
)
# Build splits
sm = PatientFirstSplitManager()
split_args = SimpleNamespace(
n_splits=v1["args"].n_splits,
fold_seed=v1["args"].fold_seed,
holdout_per_class=v1["args"].holdout_per_class,
holdout_seed=v1["args"].holdout_seed,
eval_mode=v1["args"].eval_mode,
)
plans = sm.build_plans(clinical=v2["clinical"], args=split_args, profile=profile)
folds = min(folds, len(plans))
for fold in range(folds):
# Build V1 modules per fold (seeded)
_seed_all(seed + fold * 1000 + 1)
v1_modules = _build_v1_modules(v1)
_move_modules(v1_modules, device)
# Build V2 modules per fold (seeded to match V1 init)
_seed_all(seed + fold * 1000 + 1)
v2_bundle = build_model_bundle(v2["assembly"], v2["clinical"])
if v2_bundle.bridge is None:
diffs.append("V2 model bundle missing bridge.")
return diffs
if isinstance(v2_bundle.bridge, VoteBridge):
diffs.append("V2 bridge is VoteBridge; training compare only supports fusion bridge.")
return diffs
if v2_bundle.image_transform is None:
diffs.append("V2 image transform missing; cannot run training compare.")
return diffs
v2_modules = {
"image_tower": v2_bundle.image_tower,
"metadata_tower": v2_bundle.metadata_tower,
"bridge": v2_bundle.bridge,
}
_move_modules(v2_modules, device)
split = plans[fold]
v1_train = split.train
v1_val = split.val
v1_train_ds = _build_v1_dataset(v1, v1_train, v1_modules["image_tower"].transform)
v1_val_ds = _build_v1_dataset(v1, v1_val, v1_modules["image_tower"].transform)
v2_train_ds = _build_v2_dataset(v2, v1_train, profile, v2_bundle.image_transform)
v2_val_ds = _build_v2_dataset(v2, v1_val, profile, v2_bundle.image_transform)
# Optimizers
v1_opt = torch.optim.Adam(
list(v1_modules["image_tower"].parameters())
+ list(v1_modules["metadata_tower"].parameters())
+ list(v1_modules["bridge"].parameters()),
lr=float(v1["args"].lr),
)
v2_opt = torch.optim.Adam(
list(v2_modules["image_tower"].parameters())
+ list(v2_modules["metadata_tower"].parameters())
+ list(v2_modules["bridge"].parameters()),
lr=float(v1["args"].lr),
)
criterion = nn.CrossEntropyLoss()
for epoch in range(epochs):
_seed_all(seed + fold * 100 + epoch)
v1_train_metrics = _run_epoch_v1(
v1_modules,
v1_train_ds,
v1_opt,
criterion,
device,
batch_size=batch_size,
max_batches=max_batches,
train=True,
seed=seed + fold * 100 + epoch,
)
v2_train_metrics = _run_epoch_v2(
v2_modules,
v2_train_ds,
v2_opt,
criterion,
device,
batch_size=batch_size,
max_batches=max_batches,
train=True,
seed=seed + fold * 100 + epoch,
)
v1_val_metrics = _run_epoch_v1(
v1_modules,
v1_val_ds,
None,
criterion,
device,
batch_size=batch_size,
max_batches=max_batches,
train=False,
seed=seed + fold * 100 + epoch + 777,
)
v2_val_metrics = _run_epoch_v2(
v2_modules,
v2_val_ds,
None,
criterion,
device,
batch_size=batch_size,
max_batches=max_batches,
train=False,
seed=seed + fold * 100 + epoch + 777,
)
print(
f"[fold {fold} epoch {epoch}] "
f"v1 train loss={v1_train_metrics['loss']:.4f} acc={v1_train_metrics['acc']:.4f} | "
f"v2 train loss={v2_train_metrics['loss']:.4f} acc={v2_train_metrics['acc']:.4f}"
)
print(
f"[fold {fold} epoch {epoch}] "
f"v1 val loss={v1_val_metrics['loss']:.4f} acc={v1_val_metrics['acc']:.4f} | "
f"v2 val loss={v2_val_metrics['loss']:.4f} acc={v2_val_metrics['acc']:.4f}"
)
diffs.extend(
_compare_epoch_metrics(
fold,
epoch,
v1_train_metrics,
v2_train_metrics,
v1_val_metrics,
v2_val_metrics,
loss_tol,
acc_tol,
)
)
return diffs
def _compare_epoch_metrics(
fold: int,
epoch: int,
v1_train: Dict[str, float],
v2_train: Dict[str, float],
v1_val: Dict[str, float],
v2_val: Dict[str, float],
loss_tol: float,
acc_tol: float,
) -> List[str]:
diffs: List[str] = []
for split_name, a, b in (
("train", v1_train, v2_train),
("val", v1_val, v2_val),
):
loss_diff = abs(a["loss"] - b["loss"])
acc_diff = abs(a["acc"] - b["acc"])
if loss_diff > loss_tol:
diffs.append(
f"Fold {fold} epoch {epoch} {split_name} loss diff {loss_diff:.4f} (v1={a['loss']:.4f} v2={b['loss']:.4f})"
)
if acc_diff > acc_tol:
diffs.append(
f"Fold {fold} epoch {epoch} {split_name} acc diff {acc_diff:.4f} (v1={a['acc']:.4f} v2={b['acc']:.4f})"
)
return diffs
def _move_modules(modules: Dict[str, Any], device: str) -> None:
for module in modules.values():
if module is not None and hasattr(module, "to"):
module.to(device)
def _build_v1_dataset(v1: Dict[str, Any], df, image_transform) -> ClinicalDataset:
view = _ClinicalView(v1["clinical"], df)
return ClinicalDataset(view, image_transform)
def _build_v2_dataset(v2: Dict[str, Any], df, profile, image_transform) -> SlotDataset:
samples = profile.build_samples(df=df, clinical=v2["clinical"])
return SlotDataset(
samples,
profile.slot_descriptors(),
image_transform=image_transform,
)
def _make_loader(
dataset,
*,
batch_size: int,
shuffle: bool,
seed: int,
collate_fn=None,
) -> torch.utils.data.DataLoader:
g = torch.Generator()
g.manual_seed(seed)
return torch.utils.data.DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
generator=g,
collate_fn=collate_fn,
)
def _run_epoch_v1(
modules: Dict[str, Any],
dataset: ClinicalDataset,
optimizer: Optional[torch.optim.Optimizer],
criterion: nn.Module,
device: str,
*,
batch_size: int,
max_batches: int,
train: bool,
seed: int,
) -> Dict[str, float]:
_seed_all(seed)
loader = _make_loader(dataset, batch_size=batch_size, shuffle=train, seed=seed)
image_tower = modules["image_tower"]
md_tower = modules["metadata_tower"]
bridge = modules["bridge"]
image_tower.train(train)
md_tower.train(train)
bridge.train(train)
total_loss = 0.0
total_correct = 0
total_count = 0
context = torch.enable_grad() if train else torch.no_grad()
with context:
for step, batch in enumerate(loader):
if step >= max_batches:
break
if len(batch) == 4:
imgs, metas, _, labels = batch
else:
imgs, metas, labels = batch
imgs = imgs.to(device)
metas = metas.to(device)
labels = labels.to(device)
if optimizer is not None:
optimizer.zero_grad()
img_feats = image_tower(imgs)
md_feats = md_tower(metas)
out_fused, _, _ = bridge(img_feats, md_feats)
loss = criterion(out_fused, labels)
if optimizer is not None:
loss.backward()
optimizer.step()
total_loss += float(loss.detach().item()) * labels.size(0)
total_correct += (out_fused.argmax(dim=1) == labels).sum().item()
total_count += labels.size(0)
if total_count == 0:
return {"loss": float("nan"), "acc": float("nan")}
return {"loss": total_loss / total_count, "acc": total_correct / total_count}
def _run_epoch_v2(
modules: Dict[str, Any],
dataset: SlotDataset,
optimizer: Optional[torch.optim.Optimizer],
criterion: nn.Module,
device: str,
*,
batch_size: int,
max_batches: int,
train: bool,
seed: int,
) -> Dict[str, float]:
_seed_all(seed)
loader = _make_loader(
dataset,
batch_size=batch_size,
shuffle=train,
seed=seed,
collate_fn=slot_collate,
)
image_tower = modules["image_tower"]
md_tower = modules["metadata_tower"]
bridge = modules["bridge"]
image_tower.train(train)
md_tower.train(train)
bridge.train(train)
total_loss = 0.0
total_correct = 0
total_count = 0
context = torch.enable_grad() if train else torch.no_grad()
with context:
for step, batch in enumerate(loader):
if step >= max_batches:
break
imgs = batch.get("image_1")
metas = batch.get("matrix_1")
labels = batch.get("label_1")
if imgs is None or metas is None or labels is None:
continue
if not torch.is_tensor(imgs) or not torch.is_tensor(metas):
continue
imgs = imgs.to(device)
metas = metas.to(device)
labels = torch.as_tensor(labels, device=device)
if optimizer is not None:
optimizer.zero_grad()
img_feats = image_tower(imgs)
md_feats = md_tower(metas)
out_fused, _, _ = bridge(img_feats, md_feats)
loss = criterion(out_fused, labels)
if optimizer is not None:
loss.backward()
optimizer.step()
total_loss += float(loss.detach().item()) * labels.size(0)
total_correct += (out_fused.argmax(dim=1) == labels).sum().item()
total_count += labels.size(0)
if total_count == 0:
return {"loss": float("nan"), "acc": float("nan")}
return {"loss": total_loss / total_count, "acc": total_correct / total_count}
if __name__ == "__main__":
raise SystemExit(main())
+135
View File
@@ -0,0 +1,135 @@
#!/usr/bin/env python3
"""Describe PAPILA splits and V2 slot-based loaders."""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
from types import SimpleNamespace
import torch
# Ensure repo root is importable when running as: python3 scripts/test_v2_papila_loaders.py
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from classes import build_papila_clinical
from classes.v2 import build_papila_profile, PatientFirstSplitManager, SlotLoaderFactory
def _describe_split(name: str, df, label_col: str) -> str:
if df is None or df.empty:
return f"{name}: empty"
patient_ids = set(df["Patient ID"].tolist())
class_counts = dict(df.groupby(label_col).size().to_dict())
return (
f"{name}: patients={len(patient_ids)} rows={len(df)} "
f"class_rows={class_counts}"
)
def _describe_batch(batch: dict) -> list[str]:
lines = []
for key, val in batch.items():
if isinstance(val, torch.Tensor):
lines.append(f"{key}: tensor shape={tuple(val.shape)} dtype={val.dtype}")
elif isinstance(val, list):
non_none = next((v for v in val if v is not None), None)
lines.append(
f"{key}: list len={len(val)} sample_type={type(non_none).__name__ if non_none is not None else 'None'}"
)
else:
lines.append(f"{key}: {type(val).__name__}")
return lines
def main() -> None:
ap = argparse.ArgumentParser(description="Describe PAPILA splits + V2 slot-based loaders.")
ap.add_argument("--image-dir", default="Papila/FundusImages")
ap.add_argument("--clinical-dir", default="Papila/ClinicalData")
ap.add_argument("--label-col", default="Diagnosis")
ap.add_argument(
"--cat-cols",
nargs="*",
default=["Gender", "Phakic/Pseudophakic"],
help="Categorical columns for PAPILA builder.",
)
ap.add_argument("--n-splits", type=int, default=5)
ap.add_argument("--fold-seed", type=int, default=42)
ap.add_argument("--holdout-per-class", type=int, default=1)
ap.add_argument("--holdout-seed", type=int, default=123)
ap.add_argument("--batch-size", type=int, default=4)
ap.add_argument("--num-workers", type=int, default=0)
ap.add_argument("--fold", type=int, default=0, help="Which fold to inspect in detail.")
ap.add_argument(
"--sample-mode",
choices=["patient", "eye"],
default="patient",
help="Build samples per patient (multi-slot) or per eye (row-level).",
)
args = ap.parse_args()
clinical = build_papila_clinical(
image_dir=args.image_dir,
clinical_dir=args.clinical_dir,
label_col=args.label_col,
cat_cols=args.cat_cols,
n_splits=args.n_splits,
random_seed=args.fold_seed,
)
profile = build_papila_profile(
patient_col="Patient ID",
label_col=args.label_col,
sample_mode=args.sample_mode,
)
print("=== PAPILA profile slots ===")
for key, desc in profile.slot_descriptors().items():
print(f"{key}: kind={desc.kind} required={desc.required} desc={desc.description}")
print("aliases:", profile.semantic_aliases())
split_args = SimpleNamespace(
eval_mode="multiclass",
holdout_per_class=args.holdout_per_class,
holdout_seed=args.holdout_seed,
n_splits=args.n_splits,
fold_seed=args.fold_seed,
)
split_manager = PatientFirstSplitManager(patient_col="Patient ID", label_col=args.label_col)
plans = split_manager.build_plans(clinical=clinical, args=split_args, profile=profile)
print("\n=== Split summaries ===")
for i, split in enumerate(plans):
print(f"fold {i}:")
print(" " + _describe_split("train", split.train, args.label_col))
print(" " + _describe_split("val", split.val, args.label_col))
print(" " + _describe_split("holdout", split.holdout, args.label_col))
if args.fold < 0 or args.fold >= len(plans):
raise SystemExit(f"Requested fold {args.fold} but only {len(plans)} folds are available")
split = plans[args.fold]
loader_factory = SlotLoaderFactory(num_workers=args.num_workers)
loaders = loader_factory.build(
clinical=clinical,
split=split,
args=SimpleNamespace(batch_size=args.batch_size),
fold=args.fold,
profile=profile,
)
print(f"\n=== Loader inspection (fold {args.fold}) ===")
for name, loader in (("train", loaders.train), ("val", loaders.val), ("holdout", loaders.holdout)):
if loader is None:
print(f"{name}: None")
continue
print(f"{name}: batches={len(loader)} batch_size={loader.batch_size}")
batch = next(iter(loader))
for line in _describe_batch(batch):
print(f" {line}")
if __name__ == "__main__":
main()
+218
View File
@@ -0,0 +1,218 @@
#!/usr/bin/env python3
"""Tiny smoke test for classes.v2.split_manager."""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
from types import SimpleNamespace
import numpy as np
import pandas as pd
# Ensure repo root is importable when running as: python3 scripts/test_v2_split_manager.py
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from classes import build_papila_clinical
from classes.v2 import build_papila_profile
from classes.v2.split_manager import PatientFirstSplitManager, build_patient_split_plans
class _ClinicalStub:
def __init__(self, df: pd.DataFrame, label_col: str = "Diagnosis") -> None:
self.df = df
self.label_col = label_col
def _make_fake_df(n_patients: int, n_classes: int, label_col: str) -> pd.DataFrame:
rows = []
for pid in range(1, n_patients + 1):
label = (pid - 1) % n_classes
for eye in ("OD", "OS"):
rows.append(
{
"Patient ID": pid,
"eyeID": eye,
label_col: label,
"dummy_feature": float(pid),
}
)
return pd.DataFrame(rows)
def _summarize_fold(
fold: int,
split,
label_col: str,
expected_rows_per_patient: int | None = None,
) -> str:
train_ids = set(split.train["Patient ID"].tolist())
val_ids = set(split.val["Patient ID"].tolist())
holdout_ids = set(split.holdout["Patient ID"].tolist()) if split.holdout is not None else set()
if train_ids & val_ids:
raise RuntimeError(f"Fold {fold}: train/val overlap detected")
if train_ids & holdout_ids:
raise RuntimeError(f"Fold {fold}: train/holdout overlap detected")
if val_ids & holdout_ids:
raise RuntimeError(f"Fold {fold}: val/holdout overlap detected")
if expected_rows_per_patient is not None:
# Used only for synthetic data where we know OD+OS are both present.
for name, df in (("train", split.train), ("val", split.val), ("holdout", split.holdout)):
if df is None or df.empty:
continue
counts = df.groupby("Patient ID").size().unique().tolist()
if counts != [expected_rows_per_patient]:
raise RuntimeError(f"Fold {fold}: {name} has broken per-patient row grouping: {counts}")
train_cls = dict(split.train.groupby(label_col).size().to_dict())
val_cls = dict(split.val.groupby(label_col).size().to_dict())
hold_cls = dict(split.holdout.groupby(label_col).size().to_dict()) if split.holdout is not None else {}
return (
f"fold={fold} "
f"train_patients={len(train_ids)} val_patients={len(val_ids)} holdout_patients={len(holdout_ids)} "
f"train_rows={len(split.train)} val_rows={len(split.val)} holdout_rows={0 if split.holdout is None else len(split.holdout)} "
f"train_class_rows={train_cls} val_class_rows={val_cls} holdout_class_rows={hold_cls}"
)
def _confirm_holdout_consistency_and_exclusion(splits) -> None:
holdout_sets: list[set] = []
val_union: set = set()
for split in splits:
holdout_ids = set(split.holdout["Patient ID"].tolist()) if split.holdout is not None else set()
holdout_sets.append(holdout_ids)
val_union.update(split.val["Patient ID"].tolist())
# A) Holdout should be the same patients across all folds.
baseline = holdout_sets[0] if holdout_sets else set()
for i, holdout_ids in enumerate(holdout_sets):
if holdout_ids != baseline:
raise RuntimeError(
f"Holdout mismatch: fold 0 has {sorted(baseline)}, fold {i} has {sorted(holdout_ids)}"
)
# B) Holdout patients should never appear in any validation/test fold.
overlap = baseline & val_union
if overlap:
raise RuntimeError(f"Holdout patients found in val/test sets: {sorted(overlap)}")
print(
"Holdout checks: OK "
f"(constant across folds, holdout_patients={len(baseline)}, overlap_with_any_val=0)"
)
def main() -> None:
ap = argparse.ArgumentParser(description="Smoke test PatientFirstSplitManager with synthetic data.")
ap.add_argument("--dataset", choices=["papila", "synthetic"], default="papila")
ap.add_argument(
"--patients",
type=int,
default=30,
help="Synthetic mode only: number of fake patients to generate.",
)
ap.add_argument(
"--synthetic-classes",
type=int,
default=3,
help="Synthetic mode only: number of classes to generate.",
)
ap.add_argument("--n-splits", type=int, default=5)
ap.add_argument("--holdout-per-class", type=int, default=1)
ap.add_argument("--fold-seed", type=int, default=42)
ap.add_argument("--holdout-seed", type=int, default=123)
ap.add_argument("--image-dir", default="Papila/FundusImages")
ap.add_argument("--clinical-dir", default="Papila/ClinicalData")
ap.add_argument("--label-col", default="Diagnosis")
ap.add_argument(
"--sample-mode",
choices=["patient", "eye"],
default="patient",
help="Build samples per patient (multi-slot) or per eye (row-level).",
)
ap.add_argument(
"--cat-cols",
nargs="*",
default=["Gender", "Phakic/Pseudophakic"],
help="Categorical columns for PAPILA builder.",
)
args = ap.parse_args()
if args.synthetic_classes < 2:
raise SystemExit("--synthetic-classes must be >= 2")
expected_rows_per_patient: int | None = None
if args.dataset == "papila":
clinical = build_papila_clinical(
image_dir=args.image_dir,
clinical_dir=args.clinical_dir,
label_col=args.label_col,
cat_cols=args.cat_cols,
n_splits=args.n_splits,
random_seed=args.fold_seed,
)
df = clinical.df.copy()
print(
f"Loaded PAPILA dataframe: rows={len(df)} patients={df['Patient ID'].nunique()} "
f"labels={dict(df.groupby(args.label_col).size().to_dict())}"
)
else:
df = _make_fake_df(args.patients, args.synthetic_classes, args.label_col)
clinical = _ClinicalStub(df=df, label_col=args.label_col)
expected_rows_per_patient = 2
print(
f"Loaded synthetic dataframe: rows={len(df)} patients={df['Patient ID'].nunique()} "
f"labels={dict(df.groupby(args.label_col).size().to_dict())}"
)
n_classes = int(df[args.label_col].nunique())
print(f"Detected classes from dataframe: n_classes={n_classes}")
split_args = SimpleNamespace(
eval_mode="multiclass",
holdout_per_class=args.holdout_per_class,
holdout_seed=args.holdout_seed,
n_splits=args.n_splits,
fold_seed=args.fold_seed,
)
manager = PatientFirstSplitManager(patient_col="Patient ID", label_col=args.label_col)
profile = build_papila_profile(
patient_col="Patient ID",
label_col=args.label_col,
sample_mode=args.sample_mode,
)
splits = manager.build_plans(clinical=clinical, args=split_args, profile=profile)
print("=== Adapter split manager output ===")
for fold, split in enumerate(splits):
print(
_summarize_fold(
fold,
split,
label_col=args.label_col,
expected_rows_per_patient=expected_rows_per_patient,
)
)
_confirm_holdout_consistency_and_exclusion(splits)
# Also smoke-test the pure vector API directly.
patient_labels = df.groupby("Patient ID")[args.label_col].first()
plans = build_patient_split_plans(
patient_ids=patient_labels.index.to_numpy(),
patient_labels=patient_labels.to_numpy(),
n_splits=args.n_splits,
seed=args.fold_seed,
holdout_per_class=args.holdout_per_class,
holdout_seed=args.holdout_seed,
)
print(f"\nVector API produced {len(plans)} fold plans.")
print("OK: split manager smoke test passed.")
if __name__ == "__main__":
main()