"""Leakage-safe split builders for A1 baseline evaluation protocols.""" from __future__ import annotations from dataclasses import asdict, dataclass from typing import Any import pandas as pd @dataclass(frozen=True) class CrossRunFold: """One leave-one-run-out fold for Protocol A.""" protocol: str subject: str fold_id: str test_run: int train_runs: str n_train_runs: int n_test_runs: int @dataclass(frozen=True) class WithinRunBlockedSplit: """One blocked temporal split with HRF safety gap for Protocol B.""" protocol: str subject: str run: int split_id: str n_volumes: int test_start_tr: int test_end_tr_exclusive: int gap_tr: int train_left_start_tr: int train_left_end_tr_exclusive: int train_right_start_tr: int train_right_end_tr_exclusive: int n_train_volumes: int n_test_volumes: int n_gap_excluded_volumes: int @dataclass(frozen=True) class WithinRunSkipped: """Skipped run when blocked split constraints cannot be satisfied.""" subject: str run: int n_volumes: int reason: str def build_cross_run_folds(manifest_df: pd.DataFrame) -> pd.DataFrame: """Build Protocol A folds (leave-one-run-out within subject).""" if manifest_df.empty: return pd.DataFrame() required_columns = {"subject", "run"} missing = required_columns.difference(manifest_df.columns) if missing: raise ValueError(f"Manifest missing required columns: {sorted(missing)}") rows: list[CrossRunFold] = [] for subject, subject_df in manifest_df.groupby("subject"): runs = sorted({int(run) for run in subject_df["run"].tolist()}) if len(runs) < 2: continue for test_run in runs: train_runs = [run for run in runs if run != test_run] rows.append( CrossRunFold( protocol="A_cross_run", subject=str(subject), fold_id=f"{subject}_test_run{test_run}", test_run=int(test_run), train_runs=",".join(str(run) for run in train_runs), n_train_runs=len(train_runs), n_test_runs=1, ) ) output_df = pd.DataFrame([asdict(row) for row in rows]) if not output_df.empty: output_df = output_df.sort_values(["subject", "test_run"]).reset_index(drop=True) return output_df def _compute_within_run_blocked_bounds( n_volumes: int, test_fraction: float, gap_tr: int, min_train_volumes: int, min_test_volumes: int, ) -> tuple[dict[str, int] | None, str | None]: if n_volumes <= 0: return None, "non_positive_volume_count" if not (0.0 < test_fraction < 1.0): return None, "invalid_test_fraction" if gap_tr < 0: return None, "negative_gap" n_test = max(min_test_volumes, int(round(n_volumes * test_fraction))) n_test = min(n_test, n_volumes) if n_test >= n_volumes: return None, "test_block_covers_entire_run" # Center block so train windows exist on both sides when possible. test_start = max(0, (n_volumes - n_test) // 2) test_end = min(n_volumes, test_start + n_test) left_train_start = 0 left_train_end = max(0, test_start - gap_tr) right_train_start = min(n_volumes, test_end + gap_tr) right_train_end = n_volumes n_train = (left_train_end - left_train_start) + (right_train_end - right_train_start) n_gap = (test_start - left_train_end) + (right_train_start - test_end) if n_train < min_train_volumes: return None, "insufficient_train_volumes_after_gap" if n_test < min_test_volumes: return None, "insufficient_test_volumes" if left_train_end < left_train_start or right_train_end < right_train_start: return None, "invalid_train_segment_bounds" bounds = { "test_start_tr": int(test_start), "test_end_tr_exclusive": int(test_end), "train_left_start_tr": int(left_train_start), "train_left_end_tr_exclusive": int(left_train_end), "train_right_start_tr": int(right_train_start), "train_right_end_tr_exclusive": int(right_train_end), "n_train_volumes": int(n_train), "n_test_volumes": int(n_test), "n_gap_excluded_volumes": int(n_gap), } return bounds, None def build_within_run_blocked_splits( manifest_df: pd.DataFrame, test_fraction: float = 0.2, gap_tr: int = 8, min_train_volumes: int = 40, min_test_volumes: int = 20, ) -> tuple[pd.DataFrame, pd.DataFrame]: """Build Protocol B blocked temporal splits for each subject-run.""" if manifest_df.empty: return pd.DataFrame(), pd.DataFrame() required_columns = {"subject", "run", "n_volumes"} missing = required_columns.difference(manifest_df.columns) if missing: raise ValueError(f"Manifest missing required columns: {sorted(missing)}") split_rows: list[WithinRunBlockedSplit] = [] skipped_rows: list[WithinRunSkipped] = [] for row in manifest_df.itertuples(index=False): subject = str(getattr(row, "subject")) run = int(getattr(row, "run")) n_volumes = int(getattr(row, "n_volumes")) bounds, reason = _compute_within_run_blocked_bounds( n_volumes=n_volumes, test_fraction=test_fraction, gap_tr=gap_tr, min_train_volumes=min_train_volumes, min_test_volumes=min_test_volumes, ) if bounds is None: skipped_rows.append( WithinRunSkipped( subject=subject, run=run, n_volumes=n_volumes, reason=str(reason), ) ) continue split_rows.append( WithinRunBlockedSplit( protocol="B_within_run_blocked", subject=subject, run=run, split_id=f"{subject}_run{run}_blocked", n_volumes=n_volumes, test_start_tr=bounds["test_start_tr"], test_end_tr_exclusive=bounds["test_end_tr_exclusive"], gap_tr=int(gap_tr), train_left_start_tr=bounds["train_left_start_tr"], train_left_end_tr_exclusive=bounds["train_left_end_tr_exclusive"], train_right_start_tr=bounds["train_right_start_tr"], train_right_end_tr_exclusive=bounds["train_right_end_tr_exclusive"], n_train_volumes=bounds["n_train_volumes"], n_test_volumes=bounds["n_test_volumes"], n_gap_excluded_volumes=bounds["n_gap_excluded_volumes"], ) ) split_df = pd.DataFrame([asdict(row) for row in split_rows]) skipped_df = pd.DataFrame([asdict(row) for row in skipped_rows]) if not split_df.empty: split_df = split_df.sort_values(["subject", "run"]).reset_index(drop=True) if not skipped_df.empty: skipped_df = skipped_df.sort_values(["subject", "run"]).reset_index(drop=True) return split_df, skipped_df def summarize_split_counts( manifest_df: pd.DataFrame, cross_run_df: pd.DataFrame, within_run_df: pd.DataFrame, within_run_skipped_df: pd.DataFrame, ) -> dict[str, Any]: """Return compact split summary for reporting JSON outputs.""" return { "n_manifest_rows": int(len(manifest_df)), "n_subjects_manifest": int(manifest_df["subject"].nunique()) if not manifest_df.empty else 0, "n_protocol_a_folds": int(len(cross_run_df)), "n_protocol_b_splits": int(len(within_run_df)), "n_protocol_b_skipped": int(len(within_run_skipped_df)), }