File size: 7,775 Bytes
8c5a642
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
"""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)),
    }