| """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" |
|
|
| |
| 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)), |
| } |
|
|