moved_repo_first_update
This commit is contained in:
Executable
+43
@@ -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())
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user