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