From 83a3231820ac00536e8dce31e6f1f027c6de7b9a Mon Sep 17 00:00:00 2001 From: slhhuang Date: Thu, 3 Sep 2026 00:48:29 -0400 Subject: [PATCH 1/6] Add patient-level k-fold split generation --- scripts/make_kfold_splits.py | 278 +++++++++++++++++++++++++++++++++++ tests/test_kfold_splits.py | 100 +++++++++++++ 2 files changed, 378 insertions(+) create mode 100644 scripts/make_kfold_splits.py create mode 100644 tests/test_kfold_splits.py diff --git a/scripts/make_kfold_splits.py b/scripts/make_kfold_splits.py new file mode 100644 index 0000000..7a92839 --- /dev/null +++ b/scripts/make_kfold_splits.py @@ -0,0 +1,278 @@ +"""Create reproducible patient-level k-fold splits for the paired LVEF cohort. + +E13 requirements: +- Split by unique subject_id, never by study/record. +- Each subject belongs to exactly one outer test fold. +- Within each fold, assign every row to train/val/test. +- Verify zero subject overlap within every fold. +- Preserve approximately the canonical 70/10/20 train/val/test proportions. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +from pathlib import Path +from typing import Any + +import numpy as np +import pandas as pd + + +def file_sha256(path: Path) -> str: + """Compute SHA256 hash for an input file.""" + h = hashlib.sha256() + with path.open("rb") as f: + for chunk in iter(lambda: f.read(1024 * 1024), b""): + h.update(chunk) + return h.hexdigest() + + +def hash_values(values: list[Any]) -> str: + """Hash sorted values for reproducibility without relying on file order.""" + text = "\n".join(str(x) for x in sorted(values, key=lambda y: str(y))) + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def read_cohort(path: Path) -> pd.DataFrame: + """Read a CSV or Parquet cohort file.""" + if path.suffix.lower() == ".csv": + return pd.read_csv(path) + if path.suffix.lower() == ".parquet": + return pd.read_parquet(path) + raise ValueError(f"Unsupported input format: {path.suffix}") + + +def write_cohort(df: pd.DataFrame, path: Path) -> None: + """Write a CSV or Parquet cohort file.""" + path.parent.mkdir(parents=True, exist_ok=True) + + if path.suffix.lower() == ".csv": + df.to_csv(path, index=False) + return + + if path.suffix.lower() == ".parquet": + df.to_parquet(path, index=False) + return + + raise ValueError(f"Unsupported output format: {path.suffix}") + + +def verify_no_subject_overlap(splits: dict[str, list[Any]]) -> None: + """Assert that train, validation, and test subjects do not overlap.""" + train = set(splits["train"]) + val = set(splits["val"]) + test = set(splits["test"]) + + if train & val: + raise AssertionError("Subject leakage detected: train and val overlap.") + if train & test: + raise AssertionError("Subject leakage detected: train and test overlap.") + if val & test: + raise AssertionError("Subject leakage detected: val and test overlap.") + + +def make_subject_folds( + subjects: list[Any], + *, + n_folds: int = 5, + val_frac: float = 0.10, + seed: int = 42, +) -> list[dict[str, list[Any]]]: + """Create patient-level outer folds with an inner validation holdout.""" + if n_folds < 2: + raise ValueError("n_folds must be at least 2.") + + if not 0.0 < val_frac < 1.0: + raise ValueError("val_frac must be between 0 and 1.") + + subjects = sorted(subjects, key=lambda x: str(x)) + + if len(subjects) < n_folds: + raise ValueError( + f"Need at least {n_folds} subjects for {n_folds}-fold CV; got {len(subjects)}." + ) + + rng = np.random.default_rng(seed) + shuffled = list(rng.permutation(subjects)) + + test_folds = [list(x) for x in np.array_split(shuffled, n_folds)] + folds: list[dict[str, list[Any]]] = [] + + for fold_idx, test_subjects in enumerate(test_folds): + test_set = set(test_subjects) + remaining = [s for s in shuffled if s not in test_set] + + # val_frac is expressed relative to the full cohort. + # With 5 folds, test is ~20%; selecting 10% of the full cohort + # for validation leaves approximately 70% for training. + n_val = max(1, int(round(len(subjects) * val_frac))) + + fold_rng = np.random.default_rng(seed + fold_idx + 1) + remaining_shuffled = list(fold_rng.permutation(remaining)) + + val_subjects = remaining_shuffled[:n_val] + train_subjects = remaining_shuffled[n_val:] + + split = { + "train": train_subjects, + "val": val_subjects, + "test": test_subjects, + } + verify_no_subject_overlap(split) + folds.append(split) + + return folds + + +def verify_test_fold_coverage( + subjects: list[Any], + folds: list[dict[str, list[Any]]], +) -> None: + """Verify that every subject appears in exactly one outer test fold.""" + test_subjects = [subject for fold in folds for subject in fold["test"]] + + if len(test_subjects) != len(set(test_subjects)): + raise AssertionError("A subject appears in more than one outer test fold.") + + if set(test_subjects) != set(subjects): + raise AssertionError("Outer test folds do not cover every subject exactly once.") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--input", + type=Path, + default=Path("data/processed/echo_hubert_manifest.parquet"), + help="Joined EchoJEPA + HuBERT manifest.", + ) + parser.add_argument( + "--out-dir", + type=Path, + default=Path("data/processed/kfold"), + help="Directory where fold-specific manifests and metadata are written.", + ) + parser.add_argument("--n-folds", type=int, default=5) + parser.add_argument("--val-frac", type=float, default=0.10) + parser.add_argument("--seed", type=int, default=42) + args = parser.parse_args() + + cohort = read_cohort(args.input) + + if "subject_id" not in cohort.columns: + raise KeyError("Input cohort must contain a subject_id column.") + + if cohort["subject_id"].isna().any(): + raise ValueError("Input cohort contains missing subject_id values.") + + subjects = cohort["subject_id"].drop_duplicates().tolist() + + folds = make_subject_folds( + subjects, + n_folds=args.n_folds, + val_frac=args.val_frac, + seed=args.seed, + ) + + verify_test_fold_coverage(subjects, folds) + + args.out_dir.mkdir(parents=True, exist_ok=True) + fold_metadata = [] + + for fold_idx, splits in enumerate(folds): + verify_no_subject_overlap(splits) + + split_lookup = { + subject_id: split for split, subject_ids in splits.items() for subject_id in subject_ids + } + + fold_df = cohort.copy() + fold_df["split"] = fold_df["subject_id"].map(split_lookup) + + if fold_df["split"].isna().any(): + raise AssertionError(f"Fold {fold_idx}: some cohort rows were not assigned a split.") + + output_splits = { + split: (fold_df.loc[fold_df["split"] == split, "subject_id"].drop_duplicates().tolist()) + for split in ("train", "val", "test") + } + verify_no_subject_overlap(output_splits) + + fold_path = args.out_dir / f"fold_{fold_idx}.parquet" + write_cohort(fold_df, fold_path) + + subject_split_df = pd.DataFrame( + [ + { + "subject_id": subject_id, + "split": split, + } + for split, subject_ids in splits.items() + for subject_id in subject_ids + ] + ).sort_values(["split", "subject_id"]) + + subject_split_path = args.out_dir / f"fold_{fold_idx}_subjects.csv" + subject_split_df.to_csv(subject_split_path, index=False) + + row_counts = { + split: int((fold_df["split"] == split).sum()) for split in ("train", "val", "test") + } + subject_counts = { + split: int(fold_df.loc[fold_df["split"] == split, "subject_id"].nunique()) + for split in ("train", "val", "test") + } + + fold_metadata.append( + { + "fold": fold_idx, + "manifest_path": str(fold_path), + "subject_splits_path": str(subject_split_path), + "row_counts": row_counts, + "subject_counts": subject_counts, + "split_subject_id_hashes": { + split: hash_values(subject_ids) for split, subject_ids in splits.items() + }, + "leakage_check": { + "train_val_overlap": 0, + "train_test_overlap": 0, + "val_test_overlap": 0, + }, + } + ) + + print(f"Fold {fold_idx}: rows={row_counts}, subjects={subject_counts}") + + metadata = { + "task": "E13_patient_level_kfold", + "input_path": str(args.input), + "input_sha256": file_sha256(args.input), + "seed": args.seed, + "n_folds": args.n_folds, + "val_frac": args.val_frac, + "n_rows": int(len(cohort)), + "n_subjects": int(len(subjects)), + "subject_id_hash": hash_values(subjects), + "outer_test_coverage": { + "n_unique_test_subjects": int( + len({subject for fold in folds for subject in fold["test"]}) + ), + "expected_subjects": int(len(subjects)), + "each_subject_tested_once": True, + }, + "folds": fold_metadata, + } + + metadata_path = args.out_dir / "kfold_manifest.json" + metadata_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + + print(f"Wrote {args.n_folds} fold manifests to: {args.out_dir}") + print(f"Wrote k-fold metadata to: {metadata_path}") + print("Leakage check passed for every fold.") + print("Every subject appears in exactly one outer test fold.") + + +if __name__ == "__main__": + main() diff --git a/tests/test_kfold_splits.py b/tests/test_kfold_splits.py new file mode 100644 index 0000000..33f667a --- /dev/null +++ b/tests/test_kfold_splits.py @@ -0,0 +1,100 @@ +import sys +from pathlib import Path + +SCRIPTS_DIR = Path(__file__).resolve().parents[1] / "scripts" +sys.path.insert(0, str(SCRIPTS_DIR)) + +from make_kfold_splits import ( # noqa: E402 + make_subject_folds, + verify_no_subject_overlap, + verify_test_fold_coverage, +) + +def test_five_fold_sizes_match_70_10_20(): + subjects = list(range(100)) + + folds = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=42, + ) + + assert len(folds) == 5 + + for fold in folds: + assert len(fold["train"]) == 70 + assert len(fold["val"]) == 10 + assert len(fold["test"]) == 20 + + +def test_no_subject_leakage_within_folds(): + subjects = list(range(100)) + + folds = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=42, + ) + + for fold in folds: + verify_no_subject_overlap(fold) + + +def test_every_subject_is_tested_once(): + subjects = list(range(100)) + + folds = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=42, + ) + + verify_test_fold_coverage(subjects, folds) + + test_subjects = [subject for fold in folds for subject in fold["test"]] + + assert len(test_subjects) == 100 + assert len(set(test_subjects)) == 100 + + +def test_same_seed_is_reproducible(): + subjects = list(range(100)) + + folds_a = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=42, + ) + + folds_b = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=42, + ) + + assert folds_a == folds_b + + +def test_different_seed_changes_assignment(): + subjects = list(range(100)) + + folds_a = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=42, + ) + + folds_b = make_subject_folds( + subjects, + n_folds=5, + val_frac=0.10, + seed=43, + ) + + assert folds_a != folds_b From 2c3987f2e013bc9f97ed2df695cd8b0770c857a7 Mon Sep 17 00:00:00 2001 From: slhhuang Date: Sun, 6 Sep 2026 16:41:32 -0400 Subject: [PATCH 2/6] Add k-fold training and evaluation runner --- scripts/run_kfold_cv.py | 292 +++++++++++++++++++++++++++++++++++++ tests/test_kfold_cv.py | 193 ++++++++++++++++++++++++ tests/test_kfold_splits.py | 1 + 3 files changed, 486 insertions(+) create mode 100644 scripts/run_kfold_cv.py create mode 100644 tests/test_kfold_cv.py diff --git a/scripts/run_kfold_cv.py b/scripts/run_kfold_cv.py new file mode 100644 index 0000000..af710d8 --- /dev/null +++ b/scripts/run_kfold_cv.py @@ -0,0 +1,292 @@ +#!/usr/bin/env python3 +"""Run patient-level k-fold cross-validation for E13. + +For each fold, this script will: +1. Train the existing ECG, echo, concat, and fused probes. +2. Evaluate the fused checkpoint under full, echo_dropped, and ecg_dropped conditions. +3. Save fold-specific metrics and predictions. +4. Aggregate results across folds. + +The existing train_probes.py and evaluate_missing_modality.py paths are reused so +the k-fold experiment stays comparable with the canonical single-split run. +""" + +from __future__ import annotations + +import argparse +import json +import subprocess +import sys +from pathlib import Path +from typing import Any + +import numpy as np + +from primed_ai.probes import manifest as manifest_io +from primed_ai.probes.common import auroc, regression_metrics + +CONDITIONS = ("full", "echo_dropped", "ecg_dropped") + + +def run_command(command: list[str]) -> None: + """Run one pipeline command and stop immediately if it fails.""" + print("\nRunning:") + print(" ".join(command)) + subprocess.run(command, check=True) + + +def read_json(path: Path) -> dict[str, Any]: + """Read a JSON result file.""" + with path.open("r", encoding="utf-8") as f: + return json.load(f) + + +def train_fold( + fold_manifest: Path, + fold_dir: Path, + *, + epochs: int, + seed: int, + fusion_dim: int, +) -> Path: + """Train all four probes for one fold and return the fused checkpoint path.""" + probes_dir = fold_dir / "probes" + + command = [ + sys.executable, + "scripts/train_probes.py", + "--manifest", + str(fold_manifest), + "--probe", + "all", + "--out-dir", + str(probes_dir), + "--epochs", + str(epochs), + "--fusion-dim", + str(fusion_dim), + "--seed", + str(seed), + ] + + run_command(command) + + checkpoint = probes_dir / "fused" / "cross_attn_fused.pt" + + if not checkpoint.is_file(): + raise FileNotFoundError(f"Expected fused checkpoint was not created: {checkpoint}") + + return checkpoint + + +def evaluate_fold( + fold_manifest: Path, + checkpoint: Path, + fold_dir: Path, + *, + fusion_dim: int, + seed: int, + n_bootstrap: int, +) -> Path: + """Evaluate one fold's fused checkpoint under all modality conditions.""" + results_dir = fold_dir / "results" + results_dir.mkdir(parents=True, exist_ok=True) + + output_path = results_dir / "missing_modality.json" + + manifest_df = manifest_io.load(fold_manifest) + echo_dim, ecg_dim = manifest_io.dims(manifest_df) + + command = [ + sys.executable, + "scripts/evaluate_missing_modality.py", + "--manifest", + str(fold_manifest), + "--checkpoint", + str(checkpoint), + "--output", + str(output_path), + "--embed-dim", + str(fusion_dim), + "--echo-dim", + str(echo_dim), + "--ecg-dim", + str(ecg_dim), + "--seed", + str(seed), + "--n-bootstrap", + str(n_bootstrap), + ] + + run_command(command) + + if not output_path.is_file(): + raise FileNotFoundError(f"Expected missing-modality result was not created: {output_path}") + + return output_path + + +def aggregate_results(result_paths: list[Path]) -> dict[str, Any]: + """Combine per-fold metrics and out-of-fold predictions.""" + fold_payloads = [read_json(path) for path in result_paths] + + summary: dict[str, Any] = { + "n_folds": len(fold_payloads), + "conditions": {}, + } + + for condition in CONDITIONS: + per_fold = [] + pooled_lvef = [] + pooled_prediction = [] + pooled_ef40 = [] + + for fold_idx, payload in enumerate(fold_payloads): + metrics = payload["test"][condition] + predictions = payload["predictions"][condition] + + per_fold.append( + { + "fold": fold_idx, + "n_test": len(predictions["lvef"]), + "mae": metrics["mae"], + "ef40_auroc": metrics["ef40_auroc"], + "bootstrap": payload["bootstrap"].get(condition, {}), + } + ) + + pooled_lvef.extend(predictions["lvef"]) + pooled_prediction.extend(predictions["prediction"]) + pooled_ef40.extend(predictions["ef_le_40"]) + + fold_mae = np.asarray([row["mae"] for row in per_fold], dtype=float) + fold_auroc = np.asarray( + [row["ef40_auroc"] for row in per_fold], + dtype=float, + ) + + pooled_lvef_array = np.asarray(pooled_lvef, dtype=float) + pooled_prediction_array = np.asarray(pooled_prediction, dtype=float) + pooled_ef40_array = np.asarray(pooled_ef40, dtype=bool) + + pooled_metrics = regression_metrics( + pooled_lvef_array, + pooled_prediction_array, + ) + pooled_auroc = auroc( + pooled_ef40_array, + -pooled_prediction_array, + ) + + summary["conditions"][condition] = { + "per_fold": per_fold, + "across_fold": { + "mae_mean": round(float(np.mean(fold_mae)), 4), + "mae_std": round(float(np.std(fold_mae, ddof=1)), 4), + "ef40_auroc_mean": round(float(np.mean(fold_auroc)), 4), + "ef40_auroc_std": round(float(np.std(fold_auroc, ddof=1)), 4), + }, + "pooled_out_of_fold": { + "n": len(pooled_lvef_array), + "mae": pooled_metrics["mae"], + "ef40_auroc": round(float(pooled_auroc), 4), + }, + } + + return summary + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--folds-dir", + type=Path, + required=True, + help="Directory containing fold_0.parquet, fold_1.parquet, etc.", + ) + parser.add_argument( + "--out-dir", + type=Path, + default=Path("results/kfold_cv"), + help="Directory for fold checkpoints and evaluation results.", + ) + parser.add_argument("--n-folds", type=int, default=5) + parser.add_argument("--epochs", type=int, default=50) + parser.add_argument("--fusion-dim", type=int, default=256) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--n-bootstrap", type=int, default=1000) + args = parser.parse_args() + + args.out_dir.mkdir(parents=True, exist_ok=True) + + result_paths: list[Path] = [] + + for fold_idx in range(args.n_folds): + print(f"\n{'=' * 60}") + print(f"Fold {fold_idx}") + print(f"{'=' * 60}") + + fold_manifest = args.folds_dir / f"fold_{fold_idx}.parquet" + + if not fold_manifest.is_file(): + raise FileNotFoundError(f"Fold manifest not found: {fold_manifest}") + + fold_dir = args.out_dir / f"fold_{fold_idx}" + fold_dir.mkdir(parents=True, exist_ok=True) + + fold_seed = args.seed + + checkpoint = train_fold( + fold_manifest, + fold_dir, + epochs=args.epochs, + seed=fold_seed, + fusion_dim=args.fusion_dim, + ) + + result_path = evaluate_fold( + fold_manifest, + checkpoint, + fold_dir, + fusion_dim=args.fusion_dim, + seed=fold_seed, + n_bootstrap=args.n_bootstrap, + ) + + result_paths.append(result_path) + + summary = aggregate_results(result_paths) + + summary["config"] = { + "folds_dir": str(args.folds_dir), + "out_dir": str(args.out_dir), + "n_folds": args.n_folds, + "epochs": args.epochs, + "fusion_dim": args.fusion_dim, + "seed": args.seed, + "n_bootstrap": args.n_bootstrap, + } + + summary_path = args.out_dir / "kfold_results.json" + + with summary_path.open("w", encoding="utf-8") as f: + json.dump(summary, f, indent=2) + + print(f"\nWrote k-fold results to: {summary_path}") + + for condition in CONDITIONS: + results = summary["conditions"][condition] + across = results["across_fold"] + pooled = results["pooled_out_of_fold"] + + print(f"\n{condition}") + print(f" Across-fold MAE: {across['mae_mean']:.4f} ± {across['mae_std']:.4f}") + print( + f" Across-fold AUROC: {across['ef40_auroc_mean']:.4f} ± {across['ef40_auroc_std']:.4f}" + ) + print(f" Pooled OOF MAE: {pooled['mae']:.4f}") + print(f" Pooled OOF AUROC: {pooled['ef40_auroc']:.4f}") + + +if __name__ == "__main__": + main() diff --git a/tests/test_kfold_cv.py b/tests/test_kfold_cv.py new file mode 100644 index 0000000..81864cc --- /dev/null +++ b/tests/test_kfold_cv.py @@ -0,0 +1,193 @@ +import json +import sys +from pathlib import Path + +SCRIPTS_DIR = Path(__file__).resolve().parents[1] / "scripts" +sys.path.insert(0, str(SCRIPTS_DIR)) + +import run_kfold_cv as kfold_cv # noqa: E402 + +CONDITIONS = ("full", "echo_dropped", "ecg_dropped") + + +def _payload(lvef, prediction): + ef_le_40 = [value <= 40 for value in lvef] + + return { + "test": { + condition: { + "mae": 5.0, + "ef40_auroc": 1.0, + } + for condition in CONDITIONS + }, + "bootstrap": { + condition: { + "mae_ci_low": 4.0, + "mae_ci_high": 6.0, + } + for condition in CONDITIONS + }, + "predictions": { + condition: { + "lvef": lvef, + "prediction": prediction, + "ef_le_40": ef_le_40, + } + for condition in CONDITIONS + }, + } + + +def test_aggregate_results_combines_out_of_fold_predictions(tmp_path): + fold_0 = tmp_path / "fold_0.json" + fold_1 = tmp_path / "fold_1.json" + + fold_0.write_text( + json.dumps(_payload([30.0, 60.0], [35.0, 55.0])), + encoding="utf-8", + ) + fold_1.write_text( + json.dumps(_payload([20.0, 70.0], [25.0, 65.0])), + encoding="utf-8", + ) + + result = kfold_cv.aggregate_results([fold_0, fold_1]) + + assert result["n_folds"] == 2 + + for condition in CONDITIONS: + condition_result = result["conditions"][condition] + + assert len(condition_result["per_fold"]) == 2 + + across = condition_result["across_fold"] + assert across["mae_mean"] == 5.0 + assert across["mae_std"] == 0.0 + assert across["ef40_auroc_mean"] == 1.0 + assert across["ef40_auroc_std"] == 0.0 + + pooled = condition_result["pooled_out_of_fold"] + assert pooled["n"] == 4 + assert pooled["mae"] == 5.0 + assert pooled["ef40_auroc"] == 1.0 + + +def test_main_runs_each_fold_without_real_training(tmp_path, monkeypatch): + folds_dir = tmp_path / "folds" + out_dir = tmp_path / "results" + folds_dir.mkdir() + + for fold_idx in range(2): + (folds_dir / f"fold_{fold_idx}.parquet").touch() + + trained = [] + evaluated = [] + + def fake_train_fold( + fold_manifest, + fold_dir, + *, + epochs, + seed, + fusion_dim, + ): + trained.append( + { + "manifest": fold_manifest.name, + "fold_dir": fold_dir.name, + "epochs": epochs, + "seed": seed, + "fusion_dim": fusion_dim, + } + ) + return fold_dir / "probes" / "fused" / "cross_attn_fused.pt" + + def fake_evaluate_fold( + fold_manifest, + checkpoint, + fold_dir, + *, + fusion_dim, + seed, + n_bootstrap, + ): + evaluated.append( + { + "manifest": fold_manifest.name, + "fold_dir": fold_dir.name, + "seed": seed, + "fusion_dim": fusion_dim, + "n_bootstrap": n_bootstrap, + } + ) + return fold_dir / "results" / "missing_modality.json" + + fake_summary = { + "n_folds": 2, + "conditions": { + condition: { + "per_fold": [], + "across_fold": { + "mae_mean": 0.0, + "mae_std": 0.0, + "ef40_auroc_mean": 0.0, + "ef40_auroc_std": 0.0, + }, + "pooled_out_of_fold": { + "n": 0, + "mae": 0.0, + "ef40_auroc": 0.0, + }, + } + for condition in kfold_cv.CONDITIONS + }, + } + + monkeypatch.setattr(kfold_cv, "train_fold", fake_train_fold) + monkeypatch.setattr(kfold_cv, "evaluate_fold", fake_evaluate_fold) + monkeypatch.setattr( + kfold_cv, + "aggregate_results", + lambda result_paths: fake_summary, + ) + + monkeypatch.setattr( + sys, + "argv", + [ + "run_kfold_cv.py", + "--folds-dir", + str(folds_dir), + "--out-dir", + str(out_dir), + "--n-folds", + "2", + "--epochs", + "3", + "--fusion-dim", + "16", + "--seed", + "42", + "--n-bootstrap", + "10", + ], + ) + + kfold_cv.main() + + assert [row["manifest"] for row in trained] == [ + "fold_0.parquet", + "fold_1.parquet", + ] + assert len(evaluated) == 2 + + assert all(row["seed"] == 42 for row in trained) + assert all(row["seed"] == 42 for row in evaluated) + + summary_path = out_dir / "kfold_results.json" + assert summary_path.is_file() + + summary = json.loads(summary_path.read_text(encoding="utf-8")) + assert summary["config"]["n_folds"] == 2 + assert summary["config"]["seed"] == 42 diff --git a/tests/test_kfold_splits.py b/tests/test_kfold_splits.py index 33f667a..057ff8a 100644 --- a/tests/test_kfold_splits.py +++ b/tests/test_kfold_splits.py @@ -10,6 +10,7 @@ verify_test_fold_coverage, ) + def test_five_fold_sizes_match_70_10_20(): subjects = list(range(100)) From f722cbff7e50359ff5fedf784160ea5e1cec90a0 Mon Sep 17 00:00:00 2001 From: slhhuang Date: Sun, 6 Sep 2026 17:52:49 -0400 Subject: [PATCH 3/6] Address k-fold review feedback --- scripts/make_kfold_splits.py | 19 +++++++++ tests/test_kfold_splits.py | 79 +++++++++++++++++++++++++++++++++--- 2 files changed, 93 insertions(+), 5 deletions(-) diff --git a/scripts/make_kfold_splits.py b/scripts/make_kfold_splits.py index 7a92839..5d43c1a 100644 --- a/scripts/make_kfold_splits.py +++ b/scripts/make_kfold_splits.py @@ -18,6 +18,7 @@ import numpy as np import pandas as pd +from check_ef40_prevalence import normalize_ef_le_40 def file_sha256(path: Path) -> str: @@ -87,6 +88,11 @@ def make_subject_folds( if not 0.0 < val_frac < 1.0: raise ValueError("val_frac must be between 0 and 1.") + if val_frac >= 1.0 - (1.0 / n_folds): + raise ValueError( + "val_frac is too large for the requested number of folds; " + "the training split would be empty." + ) subjects = sorted(subjects, key=lambda x: str(x)) if len(subjects) < n_folds: @@ -189,6 +195,10 @@ def main() -> None: } fold_df = cohort.copy() + + if "split" in fold_df.columns: + fold_df["split_canonical"] = fold_df["split"] + fold_df["split"] = fold_df["subject_id"].map(split_lookup) if fold_df["split"].isna().any(): @@ -220,11 +230,19 @@ def main() -> None: row_counts = { split: int((fold_df["split"] == split).sum()) for split in ("train", "val", "test") } + subject_counts = { split: int(fold_df.loc[fold_df["split"] == split, "subject_id"].nunique()) for split in ("train", "val", "test") } + ef40_bool = normalize_ef_le_40(fold_df["ef_le_40"]) + + ef40_counts = { + split: int(ef40_bool.loc[fold_df["split"] == split].sum()) + for split in ("train", "val", "test") + } + fold_metadata.append( { "fold": fold_idx, @@ -232,6 +250,7 @@ def main() -> None: "subject_splits_path": str(subject_split_path), "row_counts": row_counts, "subject_counts": subject_counts, + "ef_le_40_counts": ef40_counts, "split_subject_id_hashes": { split: hash_values(subject_ids) for split, subject_ids in splits.items() }, diff --git a/tests/test_kfold_splits.py b/tests/test_kfold_splits.py index 057ff8a..b735cea 100644 --- a/tests/test_kfold_splits.py +++ b/tests/test_kfold_splits.py @@ -1,10 +1,10 @@ +import json import sys -from pathlib import Path -SCRIPTS_DIR = Path(__file__).resolve().parents[1] / "scripts" -sys.path.insert(0, str(SCRIPTS_DIR)) - -from make_kfold_splits import ( # noqa: E402 +import pandas as pd +import pytest +from make_kfold_splits import ( + main, make_subject_folds, verify_no_subject_overlap, verify_test_fold_coverage, @@ -99,3 +99,72 @@ def test_different_seed_changes_assignment(): ) assert folds_a != folds_b + + +def test_val_frac_cannot_empty_training_split(): + subjects = list(range(100)) + + with pytest.raises(ValueError, match="training split would be empty"): + make_subject_folds( + subjects, + n_folds=2, + val_frac=0.5, + seed=42, + ) + + +def test_main_writes_ef40_counts_and_preserves_canonical_split( + tmp_path, + monkeypatch, +): + input_path = tmp_path / "cohort.parquet" + out_dir = tmp_path / "kfold" + + cohort = pd.DataFrame( + { + "subject_id": list(range(100)), + "split": ["train"] * 70 + ["val"] * 10 + ["test"] * 20, + "ef_le_40": [i % 4 == 0 for i in range(100)], + } + ) + cohort.to_parquet(input_path, index=False) + + monkeypatch.setattr( + sys, + "argv", + [ + "make_kfold_splits.py", + "--input", + str(input_path), + "--out-dir", + str(out_dir), + "--n-folds", + "5", + "--val-frac", + "0.10", + "--seed", + "42", + ], + ) + + main() + + metadata = json.loads((out_dir / "kfold_manifest.json").read_text(encoding="utf-8")) + + fold_df = pd.read_parquet(out_dir / "fold_0.parquet") + fold_metadata = metadata["folds"][0] + + assert "split_canonical" in fold_df.columns + + expected_canonical = cohort.sort_values("subject_id")["split"].tolist() + actual_canonical = fold_df.sort_values("subject_id")["split_canonical"].tolist() + assert actual_canonical == expected_canonical + + for split in ("train", "val", "test"): + expected_count = int( + fold_df.loc[ + fold_df["split"] == split, + "ef_le_40", + ].sum() + ) + assert fold_metadata["ef_le_40_counts"][split] == expected_count From be618736648cbe21b5713b5758cbf7f53c3cd6fa Mon Sep 17 00:00:00 2001 From: slhhuang Date: Sun, 6 Sep 2026 18:19:33 -0400 Subject: [PATCH 4/6] Reuse split helper functions --- scripts/make_kfold_splits.py | 61 +++++------------------------------- 1 file changed, 7 insertions(+), 54 deletions(-) diff --git a/scripts/make_kfold_splits.py b/scripts/make_kfold_splits.py index 5d43c1a..9b7c43e 100644 --- a/scripts/make_kfold_splits.py +++ b/scripts/make_kfold_splits.py @@ -11,7 +11,6 @@ from __future__ import annotations import argparse -import hashlib import json from pathlib import Path from typing import Any @@ -19,59 +18,13 @@ import numpy as np import pandas as pd from check_ef40_prevalence import normalize_ef_le_40 - - -def file_sha256(path: Path) -> str: - """Compute SHA256 hash for an input file.""" - h = hashlib.sha256() - with path.open("rb") as f: - for chunk in iter(lambda: f.read(1024 * 1024), b""): - h.update(chunk) - return h.hexdigest() - - -def hash_values(values: list[Any]) -> str: - """Hash sorted values for reproducibility without relying on file order.""" - text = "\n".join(str(x) for x in sorted(values, key=lambda y: str(y))) - return hashlib.sha256(text.encode("utf-8")).hexdigest() - - -def read_cohort(path: Path) -> pd.DataFrame: - """Read a CSV or Parquet cohort file.""" - if path.suffix.lower() == ".csv": - return pd.read_csv(path) - if path.suffix.lower() == ".parquet": - return pd.read_parquet(path) - raise ValueError(f"Unsupported input format: {path.suffix}") - - -def write_cohort(df: pd.DataFrame, path: Path) -> None: - """Write a CSV or Parquet cohort file.""" - path.parent.mkdir(parents=True, exist_ok=True) - - if path.suffix.lower() == ".csv": - df.to_csv(path, index=False) - return - - if path.suffix.lower() == ".parquet": - df.to_parquet(path, index=False) - return - - raise ValueError(f"Unsupported output format: {path.suffix}") - - -def verify_no_subject_overlap(splits: dict[str, list[Any]]) -> None: - """Assert that train, validation, and test subjects do not overlap.""" - train = set(splits["train"]) - val = set(splits["val"]) - test = set(splits["test"]) - - if train & val: - raise AssertionError("Subject leakage detected: train and val overlap.") - if train & test: - raise AssertionError("Subject leakage detected: train and test overlap.") - if val & test: - raise AssertionError("Subject leakage detected: val and test overlap.") +from make_splits import ( + file_sha256, + hash_values, + read_cohort, + verify_no_subject_overlap, + write_cohort, +) def make_subject_folds( From 9310fc0983e999495801179ce25548e2789a393b Mon Sep 17 00:00:00 2001 From: Kevin Zhou <86838731+kevzho@users.noreply.github.com> Date: Fri, 11 Sep 2026 01:41:16 -0400 Subject: [PATCH 5/6] Record k-fold run provenance and prevalence checks --- data/processed/kfold/kfold_manifest.json | 168 +++++++++++++++++++++++ scripts/run_kfold_cv.py | 102 ++++++++++++++ tests/test_kfold_cv.py | 59 +++++++- 3 files changed, 323 insertions(+), 6 deletions(-) create mode 100644 data/processed/kfold/kfold_manifest.json diff --git a/data/processed/kfold/kfold_manifest.json b/data/processed/kfold/kfold_manifest.json new file mode 100644 index 0000000..67e9a71 --- /dev/null +++ b/data/processed/kfold/kfold_manifest.json @@ -0,0 +1,168 @@ +{ + "task": "E13_patient_level_kfold", + "input_path": "../PRIMED-AI/data/processed/echo_hubert_manifest.parquet", + "input_sha256": "cfb7664f5e95a3b8419ec6107ff71c8eb04fed15c9a14465f9cbc6b8ff2481ef", + "seed": 42, + "n_folds": 5, + "val_frac": 0.1, + "n_rows": 1184, + "n_subjects": 992, + "subject_id_hash": "332c922821b24ad294eaa1d2a709c9ad1f49d19452461b53cfcfa5d4ec63fb0c", + "outer_test_coverage": { + "n_unique_test_subjects": 992, + "expected_subjects": 992, + "each_subject_tested_once": true + }, + "folds": [ + { + "fold": 0, + "manifest_path": "data/processed/kfold/fold_0.parquet", + "subject_splits_path": "data/processed/kfold/fold_0_subjects.csv", + "row_counts": { + "train": 827, + "val": 120, + "test": 237 + }, + "subject_counts": { + "train": 694, + "val": 99, + "test": 199 + }, + "ef_le_40_counts": { + "train": 169, + "val": 27, + "test": 50 + }, + "split_subject_id_hashes": { + "train": "8ee0f4fa88378a23550c6e6476b3f5e7fa0aac2a0c40b7c58db6df254080a6ea", + "val": "e48d61e667e6c9589313f6a79eabbecf6006c3e712610fc90b9a6e114c6c4255", + "test": "d0601efc868996c574f64dac5991104781bae05b61ec37ee5196d32a2f6cf754" + }, + "leakage_check": { + "train_val_overlap": 0, + "train_test_overlap": 0, + "val_test_overlap": 0 + } + }, + { + "fold": 1, + "manifest_path": "data/processed/kfold/fold_1.parquet", + "subject_splits_path": "data/processed/kfold/fold_1_subjects.csv", + "row_counts": { + "train": 839, + "val": 121, + "test": 224 + }, + "subject_counts": { + "train": 694, + "val": 99, + "test": 199 + }, + "ef_le_40_counts": { + "train": 169, + "val": 27, + "test": 50 + }, + "split_subject_id_hashes": { + "train": "6f6086044efeeba116bcd79948251d8bce4b0148d968dd000b87b042004fab13", + "val": "d17f335a1481229c6f8d0156ae836c87ff390b144a35df75d407fc9c573d27ec", + "test": "938e34450bdb0df0661ad91834bcaadd7f6b9c184c2e7c0fd07328f225cf263f" + }, + "leakage_check": { + "train_val_overlap": 0, + "train_test_overlap": 0, + "val_test_overlap": 0 + } + }, + { + "fold": 2, + "manifest_path": "data/processed/kfold/fold_2.parquet", + "subject_splits_path": "data/processed/kfold/fold_2_subjects.csv", + "row_counts": { + "train": 832, + "val": 122, + "test": 230 + }, + "subject_counts": { + "train": 695, + "val": 99, + "test": 198 + }, + "ef_le_40_counts": { + "train": 176, + "val": 28, + "test": 42 + }, + "split_subject_id_hashes": { + "train": "1f00e0fd9058a35b7bfa1af9eea807caf57c4491d51cbad0e09756caf3087909", + "val": "3a4e057f62a35c71c8ff51b3ddf49c26599fc55ced0c815cdb54b1d6de829d79", + "test": "864288073998b584692d3b3418a859a7ec452138025043c7c59e3e85ce9c1e22" + }, + "leakage_check": { + "train_val_overlap": 0, + "train_test_overlap": 0, + "val_test_overlap": 0 + } + }, + { + "fold": 3, + "manifest_path": "data/processed/kfold/fold_3.parquet", + "subject_splits_path": "data/processed/kfold/fold_3_subjects.csv", + "row_counts": { + "train": 831, + "val": 119, + "test": 234 + }, + "subject_counts": { + "train": 695, + "val": 99, + "test": 198 + }, + "ef_le_40_counts": { + "train": 176, + "val": 21, + "test": 49 + }, + "split_subject_id_hashes": { + "train": "1cd0f63b6409d232e582509fe274c1ffbbe67cf75ddae55561b57729cc723c6b", + "val": "48d421e0928be43ba2c4bbb5f8af4c6fc3a471d8b4a9d046ffaaef33d5ca23bc", + "test": "97d0a80a99710b07370d434348c0741486edd8217aaa9f713485255138e4b86c" + }, + "leakage_check": { + "train_val_overlap": 0, + "train_test_overlap": 0, + "val_test_overlap": 0 + } + }, + { + "fold": 4, + "manifest_path": "data/processed/kfold/fold_4.parquet", + "subject_splits_path": "data/processed/kfold/fold_4_subjects.csv", + "row_counts": { + "train": 814, + "val": 111, + "test": 259 + }, + "subject_counts": { + "train": 695, + "val": 99, + "test": 198 + }, + "ef_le_40_counts": { + "train": 173, + "val": 18, + "test": 55 + }, + "split_subject_id_hashes": { + "train": "f3ef3109e0df6d729fd8b10de24fa26061875b70424189f564f12a138ed1113b", + "val": "9c91c1a6e336f982a3b855501d83ae3a187517a42dd15a5b8a2028208dbb9da1", + "test": "ad8f72fcfce6c9327f5d59c55e2f47a65671b267e92a05efbe94f515c25a9853" + }, + "leakage_check": { + "train_val_overlap": 0, + "train_test_overlap": 0, + "val_test_overlap": 0 + } + } + ] +} \ No newline at end of file diff --git a/scripts/run_kfold_cv.py b/scripts/run_kfold_cv.py index af710d8..38a0245 100644 --- a/scripts/run_kfold_cv.py +++ b/scripts/run_kfold_cv.py @@ -17,10 +17,12 @@ import json import subprocess import sys +import warnings from pathlib import Path from typing import Any import numpy as np +from make_splits import file_sha256 from primed_ai.probes import manifest as manifest_io from primed_ai.probes.common import auroc, regression_metrics @@ -41,6 +43,65 @@ def read_json(path: Path) -> dict[str, Any]: return json.load(f) +def load_kfold_metadata(path: Path, *, n_folds: int) -> dict[int, dict[str, Any]]: + """Read and validate the split metadata used by a k-fold run.""" + if not path.is_file(): + raise FileNotFoundError( + f"K-fold metadata not found: {path}. Run make_kfold_splits.py first." + ) + payload = read_json(path) + if payload.get("n_folds") != n_folds: + raise ValueError( + f"K-fold metadata has n_folds={payload.get('n_folds')}, expected {n_folds}." + ) + + metadata_by_fold = {int(entry["fold"]): entry for entry in payload.get("folds", [])} + expected = set(range(n_folds)) + if set(metadata_by_fold) != expected: + raise ValueError( + "K-fold metadata does not contain exactly the requested fold IDs: " + f"expected {sorted(expected)}, got {sorted(metadata_by_fold)}." + ) + return metadata_by_fold + + +def check_validation_prevalence( + metadata_by_fold: dict[int, dict[str, Any]], + *, + min_val_positives: int, + allow_low_val_positives: bool, +) -> None: + """Stop before training when a fold's validation AUROC is too underpowered.""" + if min_val_positives < 1: + raise ValueError("min_val_positives must be at least 1.") + + low_folds = [] + for fold_idx, metadata in sorted(metadata_by_fold.items()): + try: + n_positive = int(metadata["ef_le_40_counts"]["val"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError( + f"Fold {fold_idx} metadata is missing ef_le_40_counts.val. " + "Regenerate folds with make_kfold_splits.py." + ) from exc + if n_positive < min_val_positives: + low_folds.append((fold_idx, n_positive)) + + if not low_folds: + return + + message = ( + "Validation EF<=40 counts are below the requested minimum " + f"({min_val_positives}): " + + ", ".join(f"fold {fold}: {count}" for fold, count in low_folds) + + ". Checkpoint selection may be unstable." + ) + if allow_low_val_positives: + warnings.warn(message, stacklevel=2) + return + raise ValueError(message + " Re-run with --allow-low-val-positives to override.") + + def train_fold( fold_manifest: Path, fold_dir: Path, @@ -215,11 +276,35 @@ def main() -> None: parser.add_argument("--fusion-dim", type=int, default=256) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--n-bootstrap", type=int, default=1000) + parser.add_argument( + "--kfold-manifest", + type=Path, + help="Split metadata from make_kfold_splits.py (default: FOLDS_DIR/kfold_manifest.json).", + ) + parser.add_argument( + "--min-val-positives", + type=int, + default=10, + help="Minimum EF<=40 validation examples required before training each fold.", + ) + parser.add_argument( + "--allow-low-val-positives", + action="store_true", + help="Warn instead of stopping when a fold has too few EF<=40 validation examples.", + ) args = parser.parse_args() args.out_dir.mkdir(parents=True, exist_ok=True) + kfold_metadata_path = args.kfold_manifest or args.folds_dir / "kfold_manifest.json" + metadata_by_fold = load_kfold_metadata(kfold_metadata_path, n_folds=args.n_folds) + check_validation_prevalence( + metadata_by_fold, + min_val_positives=args.min_val_positives, + allow_low_val_positives=args.allow_low_val_positives, + ) result_paths: list[Path] = [] + fold_provenance: list[dict[str, Any]] = [] for fold_idx in range(args.n_folds): print(f"\n{'=' * 60}") @@ -254,6 +339,16 @@ def main() -> None: ) result_paths.append(result_path) + fold_provenance.append( + { + "fold": fold_idx, + "manifest_path": str(fold_manifest), + "manifest_sha256": file_sha256(fold_manifest), + "fused_checkpoint_path": str(checkpoint), + "fused_checkpoint_sha256": file_sha256(checkpoint), + "missing_modality_result_path": str(result_path), + } + ) summary = aggregate_results(result_paths) @@ -265,6 +360,13 @@ def main() -> None: "fusion_dim": args.fusion_dim, "seed": args.seed, "n_bootstrap": args.n_bootstrap, + "min_val_positives": args.min_val_positives, + "allow_low_val_positives": args.allow_low_val_positives, + } + summary["provenance"] = { + "kfold_manifest_path": str(kfold_metadata_path), + "kfold_manifest_sha256": file_sha256(kfold_metadata_path), + "folds": fold_provenance, } summary_path = args.out_dir / "kfold_results.json" diff --git a/tests/test_kfold_cv.py b/tests/test_kfold_cv.py index 81864cc..35897d9 100644 --- a/tests/test_kfold_cv.py +++ b/tests/test_kfold_cv.py @@ -1,11 +1,8 @@ import json import sys -from pathlib import Path -SCRIPTS_DIR = Path(__file__).resolve().parents[1] / "scripts" -sys.path.insert(0, str(SCRIPTS_DIR)) - -import run_kfold_cv as kfold_cv # noqa: E402 +import pytest +import run_kfold_cv as kfold_cv CONDITIONS = ("full", "echo_dropped", "ecg_dropped") @@ -81,6 +78,18 @@ def test_main_runs_each_fold_without_real_training(tmp_path, monkeypatch): for fold_idx in range(2): (folds_dir / f"fold_{fold_idx}.parquet").touch() + (folds_dir / "kfold_manifest.json").write_text( + json.dumps( + { + "n_folds": 2, + "folds": [ + {"fold": fold_idx, "ef_le_40_counts": {"val": 10}} for fold_idx in range(2) + ], + } + ), + encoding="utf-8", + ) + trained = [] evaluated = [] @@ -101,7 +110,10 @@ def fake_train_fold( "fusion_dim": fusion_dim, } ) - return fold_dir / "probes" / "fused" / "cross_attn_fused.pt" + checkpoint = fold_dir / "probes" / "fused" / "cross_attn_fused.pt" + checkpoint.parent.mkdir(parents=True, exist_ok=True) + checkpoint.write_bytes(b"synthetic checkpoint") + return checkpoint def fake_evaluate_fold( fold_manifest, @@ -191,3 +203,38 @@ def fake_evaluate_fold( summary = json.loads(summary_path.read_text(encoding="utf-8")) assert summary["config"]["n_folds"] == 2 assert summary["config"]["seed"] == 42 + assert len(summary["provenance"]["kfold_manifest_sha256"]) == 64 + assert [row["fold"] for row in summary["provenance"]["folds"]] == [0, 1] + assert all(len(row["manifest_sha256"]) == 64 for row in summary["provenance"]["folds"]) + assert all(len(row["fused_checkpoint_sha256"]) == 64 for row in summary["provenance"]["folds"]) + + +def test_main_refuses_underpowered_validation_fold(tmp_path, monkeypatch): + folds_dir = tmp_path / "folds" + folds_dir.mkdir() + (folds_dir / "kfold_manifest.json").write_text( + json.dumps( + { + "n_folds": 2, + "folds": [ + {"fold": 0, "ef_le_40_counts": {"val": 9}}, + {"fold": 1, "ef_le_40_counts": {"val": 10}}, + ], + } + ), + encoding="utf-8", + ) + monkeypatch.setattr( + sys, + "argv", + [ + "run_kfold_cv.py", + "--folds-dir", + str(folds_dir), + "--n-folds", + "2", + ], + ) + + with pytest.raises(ValueError, match="fold 0: 9"): + kfold_cv.main() From dc27f1d6a5ccb93af5ceeeac68646b4c3149a6ac Mon Sep 17 00:00:00 2001 From: Kevin Zhou <86838731+kevzho@users.noreply.github.com> Date: Tue, 22 Sep 2026 21:33:57 -0400 Subject: [PATCH 6/6] Publish canonical k-fold results and verify fold hashes --- TECHNICAL.md | 10 + docs/results/README.md | 9 + docs/results/SHA256SUMS | 2 + .../results}/kfold_manifest.json | 9 +- docs/results/kfold_results.json | 332 ++++++++++++++++++ scripts/make_kfold_splits.py | 14 +- scripts/run_kfold_cv.py | 33 +- tests/test_kfold_cv.py | 49 ++- tests/test_kfold_splits.py | 8 + 9 files changed, 454 insertions(+), 12 deletions(-) rename {data/processed/kfold => docs/results}/kfold_manifest.json (90%) create mode 100644 docs/results/kfold_results.json diff --git a/TECHNICAL.md b/TECHNICAL.md index efaa78c..b603fd7 100644 --- a/TECHNICAL.md +++ b/TECHNICAL.md @@ -226,6 +226,16 @@ model. This tests the deployed fused checkpoint under missing input. Learned nul and branch dropout are deferred unless the missing-modality rerun shows masking is unstable. +**Patient-level five-fold validation (E13).** On the corrected 1,184-row manifest, each +subject appears in exactly one outer test fold; the five finite-row test sets pool to +1,172 out-of-fold predictions after the same 12-row non-finite embedding exclusion used +by training. Full-input performance is stable across folds: MAE 9.3587 ± 0.2619 and +EF≤40 AUROC 0.8261 ± 0.0346 (pooled OOF 9.3537 / 0.8165). Missing-input behavior is less +stable: echo-dropped MAE is 17.8287 ± 4.4433 and ECG-dropped MAE is 11.7005 ± 1.6594. +The single holdout remains the locked paper split, while k-fold results quantify its +sampling uncertainty. Aggregate results and split certificates are in +`docs/results/kfold_results.json` and `docs/results/kfold_manifest.json`. + ### 7.2 Fairness audit Post-hoc stratification of probe predictions — no additional training required. diff --git a/docs/results/README.md b/docs/results/README.md index b4b4a35..18a6e75 100644 --- a/docs/results/README.md +++ b/docs/results/README.md @@ -63,3 +63,12 @@ manifest against a controlled clip-16 manifest. It records validation metrics, r attention spread, manifest/checkpoint hashes, and the decision to keep pooled embeddings canonical. Local manifests, checkpoints, run metadata, and per-example outputs remain gitignored. + +## Patient-level k-fold validation + +`kfold_manifest.json` is the aggregate-only E13 (#78) split certificate for the corrected +1,184-row manifest. It records input and per-fold parquet hashes, row/subject/EF≤40 counts, +whole-split subject hashes, complete outer-test coverage, and zero-overlap checks; it does +not contain subject IDs. `kfold_results.json` reports per-fold bootstrap intervals, +across-fold mean±SD, pooled out-of-fold metrics, and manifest/checkpoint provenance. It +contains no predictions or patient/study identifiers. diff --git a/docs/results/SHA256SUMS b/docs/results/SHA256SUMS index a916728..d2a976a 100644 --- a/docs/results/SHA256SUMS +++ b/docs/results/SHA256SUMS @@ -8,6 +8,8 @@ e6b98a7e4a6e1e06ada75ee0d312cc8824dd1de9b1c165c7ef7d98bd5ce545de docs/results/f cac171e29b456daf1735c1bf4935cd3f2afb2d14b974a1edb60fd635f08656e5 docs/results/failure_report.fused.json 0e67db6f8db9c35dbe58d71eeb5976145355bbefa14241a4d41f372cbbdf5d4f docs/results/cohort_sensitivity.json 19d8b785fd6fc91082b9a16dc9dab4ab48bbaa70da2ca52f8b685cdf013cdc59 docs/results/pooling_ablation.json +7d7efb1ddd86bffb36bd56041452f4689a4442073cb390d59c66b9cd46407b8a docs/results/kfold_manifest.json +f2ba52d2a023b4dac275febc3a2e42355cb8ca81568032bec6d67e243efc8091 docs/results/kfold_results.json 81694c9ba20fbf4af596c942c0fa3acf76275248142204e02324681f7c7a982d data/processed/echo_hubert_manifest.parquet 05d286c189e6249239c2ce73f51daf047277c728a84ecf2b42342f1d95a09802 probes/ecg/ecg_only.joblib 03d32973ea025e0c8612535c39cede2e022b5cdee85740166468fad754b0b024 probes/echo/echo_only.pt diff --git a/data/processed/kfold/kfold_manifest.json b/docs/results/kfold_manifest.json similarity index 90% rename from data/processed/kfold/kfold_manifest.json rename to docs/results/kfold_manifest.json index 67e9a71..cf884dc 100644 --- a/data/processed/kfold/kfold_manifest.json +++ b/docs/results/kfold_manifest.json @@ -1,6 +1,6 @@ { "task": "E13_patient_level_kfold", - "input_path": "../PRIMED-AI/data/processed/echo_hubert_manifest.parquet", + "input_path": "data/processed/echo_hubert_manifest.parquet", "input_sha256": "cfb7664f5e95a3b8419ec6107ff71c8eb04fed15c9a14465f9cbc6b8ff2481ef", "seed": 42, "n_folds": 5, @@ -17,6 +17,7 @@ { "fold": 0, "manifest_path": "data/processed/kfold/fold_0.parquet", + "manifest_sha256": "d8bc0ef7908d51a4a6441a914bc34d151fdc752545e8c0a0318cc6a11673267e", "subject_splits_path": "data/processed/kfold/fold_0_subjects.csv", "row_counts": { "train": 827, @@ -47,6 +48,7 @@ { "fold": 1, "manifest_path": "data/processed/kfold/fold_1.parquet", + "manifest_sha256": "18cff1546c50af462234e3929cee913685b26d5241a33ea35da369a1832e12be", "subject_splits_path": "data/processed/kfold/fold_1_subjects.csv", "row_counts": { "train": 839, @@ -77,6 +79,7 @@ { "fold": 2, "manifest_path": "data/processed/kfold/fold_2.parquet", + "manifest_sha256": "822d297cc7abcd522312277614c30b1927922fef937a6c464e61b503bae68942", "subject_splits_path": "data/processed/kfold/fold_2_subjects.csv", "row_counts": { "train": 832, @@ -107,6 +110,7 @@ { "fold": 3, "manifest_path": "data/processed/kfold/fold_3.parquet", + "manifest_sha256": "d22781038e72e5a8a121f3fd73f355fb65d25aac0d0d94ae2684a88223f5f044", "subject_splits_path": "data/processed/kfold/fold_3_subjects.csv", "row_counts": { "train": 831, @@ -137,6 +141,7 @@ { "fold": 4, "manifest_path": "data/processed/kfold/fold_4.parquet", + "manifest_sha256": "872eff6080377b056b9fcc91fa5a62dc0dd1772323cee50e792485ae9599a450", "subject_splits_path": "data/processed/kfold/fold_4_subjects.csv", "row_counts": { "train": 814, @@ -165,4 +170,4 @@ } } ] -} \ No newline at end of file +} diff --git a/docs/results/kfold_results.json b/docs/results/kfold_results.json new file mode 100644 index 0000000..719694a --- /dev/null +++ b/docs/results/kfold_results.json @@ -0,0 +1,332 @@ +{ + "n_folds": 5, + "conditions": { + "full": { + "per_fold": [ + { + "fold": 0, + "n_test": 235, + "mae": 9.4522, + "ef40_auroc": 0.7754, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 8.5989, + "mae_ci_high": 10.3985, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.7005, + "ef40_auroc_ci_high": 0.8498 + } + }, + { + "fold": 1, + "n_test": 223, + "mae": 9.5901, + "ef40_auroc": 0.8101, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 8.5445, + "mae_ci_high": 10.6686, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.7255, + "ef40_auroc_ci_high": 0.8829 + } + }, + { + "fold": 2, + "n_test": 230, + "mae": 8.9868, + "ef40_auroc": 0.8444, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 8.0032, + "mae_ci_high": 9.8819, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.775, + "ef40_auroc_ci_high": 0.9042 + } + }, + { + "fold": 3, + "n_test": 228, + "mae": 9.5729, + "ef40_auroc": 0.8652, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 8.6498, + "mae_ci_high": 10.517, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.8014, + "ef40_auroc_ci_high": 0.9226 + } + }, + { + "fold": 4, + "n_test": 256, + "mae": 9.1917, + "ef40_auroc": 0.8353, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 8.3232, + "mae_ci_high": 10.1103, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.7708, + "ef40_auroc_ci_high": 0.892 + } + } + ], + "across_fold": { + "mae_mean": 9.3587, + "mae_std": 0.2619, + "ef40_auroc_mean": 0.8261, + "ef40_auroc_std": 0.0346 + }, + "pooled_out_of_fold": { + "n": 1172, + "mae": 9.3537, + "ef40_auroc": 0.8165 + } + }, + "echo_dropped": { + "per_fold": [ + { + "fold": 0, + "n_test": 235, + "mae": 23.4454, + "ef40_auroc": 0.7146, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 21.8951, + "mae_ci_high": 25.0147, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.6449, + "ef40_auroc_ci_high": 0.7839 + } + }, + { + "fold": 1, + "n_test": 223, + "mae": 20.9422, + "ef40_auroc": 0.6825, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 19.3608, + "mae_ci_high": 22.4739, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.5972, + "ef40_auroc_ci_high": 0.7611 + } + }, + { + "fold": 2, + "n_test": 230, + "mae": 15.7548, + "ef40_auroc": 0.7061, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 14.6241, + "mae_ci_high": 16.974, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.6221, + "ef40_auroc_ci_high": 0.7791 + } + }, + { + "fold": 3, + "n_test": 228, + "mae": 16.875, + "ef40_auroc": 0.7605, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 15.5736, + "mae_ci_high": 18.2847, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.6879, + "ef40_auroc_ci_high": 0.8295 + } + }, + { + "fold": 4, + "n_test": 256, + "mae": 12.126, + "ef40_auroc": 0.546, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 11.2108, + "mae_ci_high": 13.1429, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.4491, + "ef40_auroc_ci_high": 0.6407 + } + } + ], + "across_fold": { + "mae_mean": 17.8287, + "mae_std": 4.4433, + "ef40_auroc_mean": 0.6819, + "ef40_auroc_std": 0.0811 + }, + "pooled_out_of_fold": { + "n": 1172, + "mae": 17.7092, + "ef40_auroc": 0.6188 + } + }, + "ecg_dropped": { + "per_fold": [ + { + "fold": 0, + "n_test": 235, + "mae": 10.898, + "ef40_auroc": 0.5989, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 9.8705, + "mae_ci_high": 11.9522, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.4913, + "ef40_auroc_ci_high": 0.7058 + } + }, + { + "fold": 1, + "n_test": 223, + "mae": 11.7042, + "ef40_auroc": 0.7184, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 10.538, + "mae_ci_high": 12.8138, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.6281, + "ef40_auroc_ci_high": 0.8038 + } + }, + { + "fold": 2, + "n_test": 230, + "mae": 11.3347, + "ef40_auroc": 0.4623, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 10.3058, + "mae_ci_high": 12.4323, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.3453, + "ef40_auroc_ci_high": 0.5739 + } + }, + { + "fold": 3, + "n_test": 228, + "mae": 10.0967, + "ef40_auroc": 0.6582, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 9.0036, + "mae_ci_high": 11.1817, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.553, + "ef40_auroc_ci_high": 0.7648 + } + }, + { + "fold": 4, + "n_test": 256, + "mae": 14.4689, + "ef40_auroc": 0.796, + "bootstrap": { + "n_bootstrap": 1000, + "mae_ci_low": 13.4415, + "mae_ci_high": 15.5524, + "n_auroc_replicates": 1000, + "n_auroc_discarded": 0, + "ef40_auroc_ci_low": 0.7169, + "ef40_auroc_ci_high": 0.8706 + } + } + ], + "across_fold": { + "mae_mean": 11.7005, + "mae_std": 1.6594, + "ef40_auroc_mean": 0.6468, + "ef40_auroc_std": 0.1263 + }, + "pooled_out_of_fold": { + "n": 1172, + "mae": 11.7612, + "ef40_auroc": 0.6254 + } + } + }, + "config": { + "folds_dir": "data/processed/kfold", + "out_dir": "results/kfold_cv", + "n_folds": 5, + "epochs": 50, + "fusion_dim": 256, + "seed": 42, + "n_bootstrap": 1000, + "min_val_positives": 10, + "allow_low_val_positives": false + }, + "provenance": { + "kfold_manifest_path": "docs/results/kfold_manifest.json", + "kfold_manifest_sha256": "7d7efb1ddd86bffb36bd56041452f4689a4442073cb390d59c66b9cd46407b8a", + "folds": [ + { + "fold": 0, + "manifest_path": "data/processed/kfold/fold_0.parquet", + "manifest_sha256": "d8bc0ef7908d51a4a6441a914bc34d151fdc752545e8c0a0318cc6a11673267e", + "fused_checkpoint_path": "results/kfold_cv/fold_0/probes/fused/cross_attn_fused.pt", + "fused_checkpoint_sha256": "b04483fe09334a033a5de0cf2a30282813b83de449492604fb2e54076657e27e", + "missing_modality_result_path": "results/kfold_cv/fold_0/results/missing_modality.json" + }, + { + "fold": 1, + "manifest_path": "data/processed/kfold/fold_1.parquet", + "manifest_sha256": "18cff1546c50af462234e3929cee913685b26d5241a33ea35da369a1832e12be", + "fused_checkpoint_path": "results/kfold_cv/fold_1/probes/fused/cross_attn_fused.pt", + "fused_checkpoint_sha256": "8f7ba91a190d0b494776a4250610f36a1a7189efd18758113d26266076009837", + "missing_modality_result_path": "results/kfold_cv/fold_1/results/missing_modality.json" + }, + { + "fold": 2, + "manifest_path": "data/processed/kfold/fold_2.parquet", + "manifest_sha256": "822d297cc7abcd522312277614c30b1927922fef937a6c464e61b503bae68942", + "fused_checkpoint_path": "results/kfold_cv/fold_2/probes/fused/cross_attn_fused.pt", + "fused_checkpoint_sha256": "d10c89c1614d38ecaf4ff8b557fe91a287987422b9b467de98d245438642d89d", + "missing_modality_result_path": "results/kfold_cv/fold_2/results/missing_modality.json" + }, + { + "fold": 3, + "manifest_path": "data/processed/kfold/fold_3.parquet", + "manifest_sha256": "d22781038e72e5a8a121f3fd73f355fb65d25aac0d0d94ae2684a88223f5f044", + "fused_checkpoint_path": "results/kfold_cv/fold_3/probes/fused/cross_attn_fused.pt", + "fused_checkpoint_sha256": "b9de0094aa0982c8a856a127d943b6d806b54baaeec47005c0da81b08082176d", + "missing_modality_result_path": "results/kfold_cv/fold_3/results/missing_modality.json" + }, + { + "fold": 4, + "manifest_path": "data/processed/kfold/fold_4.parquet", + "manifest_sha256": "872eff6080377b056b9fcc91fa5a62dc0dd1772323cee50e792485ae9599a450", + "fused_checkpoint_path": "results/kfold_cv/fold_4/probes/fused/cross_attn_fused.pt", + "fused_checkpoint_sha256": "d59848fa9de6aaddded22304f4be26e6e17930861455f36bf4f0960c251d6e88", + "missing_modality_result_path": "results/kfold_cv/fold_4/results/missing_modality.json" + } + ] + } +} diff --git a/scripts/make_kfold_splits.py b/scripts/make_kfold_splits.py index 9b7c43e..d8da372 100644 --- a/scripts/make_kfold_splits.py +++ b/scripts/make_kfold_splits.py @@ -113,6 +113,11 @@ def main() -> None: default=Path("data/processed/kfold"), help="Directory where fold-specific manifests and metadata are written.", ) + parser.add_argument( + "--publish-metadata", + type=Path, + help="Optional aggregate-only copy of kfold_manifest.json for docs/results.", + ) parser.add_argument("--n-folds", type=int, default=5) parser.add_argument("--val-frac", type=float, default=0.10) parser.add_argument("--seed", type=int, default=42) @@ -200,6 +205,7 @@ def main() -> None: { "fold": fold_idx, "manifest_path": str(fold_path), + "manifest_sha256": file_sha256(fold_path), "subject_splits_path": str(subject_split_path), "row_counts": row_counts, "subject_counts": subject_counts, @@ -237,11 +243,17 @@ def main() -> None: "folds": fold_metadata, } + metadata_text = json.dumps(metadata, indent=2) + "\n" metadata_path = args.out_dir / "kfold_manifest.json" - metadata_path.write_text(json.dumps(metadata, indent=2), encoding="utf-8") + metadata_path.write_text(metadata_text, encoding="utf-8") + if args.publish_metadata: + args.publish_metadata.parent.mkdir(parents=True, exist_ok=True) + args.publish_metadata.write_text(metadata_text, encoding="utf-8") print(f"Wrote {args.n_folds} fold manifests to: {args.out_dir}") print(f"Wrote k-fold metadata to: {metadata_path}") + if args.publish_metadata: + print(f"Published aggregate metadata to: {args.publish_metadata}") print("Leakage check passed for every fold.") print("Every subject appears in exactly one outer test fold.") diff --git a/scripts/run_kfold_cv.py b/scripts/run_kfold_cv.py index 38a0245..750aaeb 100644 --- a/scripts/run_kfold_cv.py +++ b/scripts/run_kfold_cv.py @@ -43,7 +43,7 @@ def read_json(path: Path) -> dict[str, Any]: return json.load(f) -def load_kfold_metadata(path: Path, *, n_folds: int) -> dict[int, dict[str, Any]]: +def load_kfold_metadata(path: Path, *, n_folds: int, folds_dir: Path) -> dict[int, dict[str, Any]]: """Read and validate the split metadata used by a k-fold run.""" if not path.is_file(): raise FileNotFoundError( @@ -62,6 +62,23 @@ def load_kfold_metadata(path: Path, *, n_folds: int) -> dict[int, dict[str, Any] "K-fold metadata does not contain exactly the requested fold IDs: " f"expected {sorted(expected)}, got {sorted(metadata_by_fold)}." ) + + for fold_idx, metadata in sorted(metadata_by_fold.items()): + fold_path = folds_dir / f"fold_{fold_idx}.parquet" + if not fold_path.is_file(): + raise FileNotFoundError(f"Fold manifest not found: {fold_path}") + expected_sha256 = metadata.get("manifest_sha256") + if not expected_sha256: + raise ValueError( + f"Fold {fold_idx} metadata is missing manifest_sha256. " + "Regenerate folds with make_kfold_splits.py." + ) + actual_sha256 = file_sha256(fold_path) + if actual_sha256 != expected_sha256: + raise ValueError( + f"Fold {fold_idx} manifest hash mismatch: metadata has " + f"{expected_sha256}, but {fold_path} has {actual_sha256}." + ) return metadata_by_fold @@ -296,7 +313,11 @@ def main() -> None: args.out_dir.mkdir(parents=True, exist_ok=True) kfold_metadata_path = args.kfold_manifest or args.folds_dir / "kfold_manifest.json" - metadata_by_fold = load_kfold_metadata(kfold_metadata_path, n_folds=args.n_folds) + metadata_by_fold = load_kfold_metadata( + kfold_metadata_path, + n_folds=args.n_folds, + folds_dir=args.folds_dir, + ) check_validation_prevalence( metadata_by_fold, min_val_positives=args.min_val_positives, @@ -319,13 +340,13 @@ def main() -> None: fold_dir = args.out_dir / f"fold_{fold_idx}" fold_dir.mkdir(parents=True, exist_ok=True) - fold_seed = args.seed - + # Reuse one seed deliberately so fold assignment, not RNG drift, is the + # experimental variable across the five training runs. checkpoint = train_fold( fold_manifest, fold_dir, epochs=args.epochs, - seed=fold_seed, + seed=args.seed, fusion_dim=args.fusion_dim, ) @@ -334,7 +355,7 @@ def main() -> None: checkpoint, fold_dir, fusion_dim=args.fusion_dim, - seed=fold_seed, + seed=args.seed, n_bootstrap=args.n_bootstrap, ) diff --git a/tests/test_kfold_cv.py b/tests/test_kfold_cv.py index 35897d9..fef32b4 100644 --- a/tests/test_kfold_cv.py +++ b/tests/test_kfold_cv.py @@ -83,7 +83,14 @@ def test_main_runs_each_fold_without_real_training(tmp_path, monkeypatch): { "n_folds": 2, "folds": [ - {"fold": fold_idx, "ef_le_40_counts": {"val": 10}} for fold_idx in range(2) + { + "fold": fold_idx, + "ef_le_40_counts": {"val": 10}, + "manifest_sha256": kfold_cv.file_sha256( + folds_dir / f"fold_{fold_idx}.parquet" + ), + } + for fold_idx in range(2) ], } ), @@ -212,13 +219,23 @@ def fake_evaluate_fold( def test_main_refuses_underpowered_validation_fold(tmp_path, monkeypatch): folds_dir = tmp_path / "folds" folds_dir.mkdir() + for fold_idx in range(2): + (folds_dir / f"fold_{fold_idx}.parquet").touch() (folds_dir / "kfold_manifest.json").write_text( json.dumps( { "n_folds": 2, "folds": [ - {"fold": 0, "ef_le_40_counts": {"val": 9}}, - {"fold": 1, "ef_le_40_counts": {"val": 10}}, + { + "fold": 0, + "ef_le_40_counts": {"val": 9}, + "manifest_sha256": kfold_cv.file_sha256(folds_dir / "fold_0.parquet"), + }, + { + "fold": 1, + "ef_le_40_counts": {"val": 10}, + "manifest_sha256": kfold_cv.file_sha256(folds_dir / "fold_1.parquet"), + }, ], } ), @@ -238,3 +255,29 @@ def test_main_refuses_underpowered_validation_fold(tmp_path, monkeypatch): with pytest.raises(ValueError, match="fold 0: 9"): kfold_cv.main() + + +def test_load_kfold_metadata_rejects_stale_fold_manifest(tmp_path): + folds_dir = tmp_path / "folds" + folds_dir.mkdir() + fold_path = folds_dir / "fold_0.parquet" + fold_path.write_bytes(b"current fold") + metadata_path = folds_dir / "kfold_manifest.json" + metadata_path.write_text( + json.dumps( + { + "n_folds": 1, + "folds": [ + { + "fold": 0, + "manifest_sha256": "0" * 64, + "ef_le_40_counts": {"val": 10}, + } + ], + } + ), + encoding="utf-8", + ) + + with pytest.raises(ValueError, match="Fold 0 manifest hash mismatch"): + kfold_cv.load_kfold_metadata(metadata_path, n_folds=1, folds_dir=folds_dir) diff --git a/tests/test_kfold_splits.py b/tests/test_kfold_splits.py index b735cea..15e2dc0 100644 --- a/tests/test_kfold_splits.py +++ b/tests/test_kfold_splits.py @@ -9,6 +9,7 @@ verify_no_subject_overlap, verify_test_fold_coverage, ) +from make_splits import file_sha256 def test_five_fold_sizes_match_70_10_20(): @@ -144,15 +145,22 @@ def test_main_writes_ef40_counts_and_preserves_canonical_split( "0.10", "--seed", "42", + "--publish-metadata", + str(tmp_path / "published" / "kfold_manifest.json"), ], ) main() metadata = json.loads((out_dir / "kfold_manifest.json").read_text(encoding="utf-8")) + published_metadata = json.loads( + (tmp_path / "published" / "kfold_manifest.json").read_text(encoding="utf-8") + ) + assert published_metadata == metadata fold_df = pd.read_parquet(out_dir / "fold_0.parquet") fold_metadata = metadata["folds"][0] + assert fold_metadata["manifest_sha256"] == file_sha256(out_dir / "fold_0.parquet") assert "split_canonical" in fold_df.columns