csai / code /a1_pipeline /splits.py
Mohith202's picture
Add core ROI masks, visualization script, and model profiles
8c5a642
Raw
History Blame Contribute Delete
7.78 kB
"""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)),
}