799 lines
25 KiB
Python
799 lines
25 KiB
Python
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())
|