diff --git a/docs/source/whats_new.rst b/docs/source/whats_new.rst index 6061228da..e635595d7 100644 --- a/docs/source/whats_new.rst +++ b/docs/source/whats_new.rst @@ -23,6 +23,7 @@ Version 1.8 (Source - GitHub) Enhancements ~~~~~~~~~~~~ +- Allow :class:`~moabb.evaluations.CrossSubjectEvaluation` to accept an optional top-level ``splitter`` instance, enabling transfer-learning protocols to reuse MOABB's existing caching, parallel execution, and result handling while preserving the default protocol (:gh:`1088` by `lindicaphxag-tech`_). - Add Leelakittisin2025 sit-stand transition imagery, PerezBlanco2026 wrist motor-execution, and Vagaja2023 VR motor-imagery datasets (:pr:`1199`) (by `Bruno Aristimunha`_). - Add MartinezPeon2025, MILimbEEG dataset loaders with synthetic regression coverage ({gh}`1198` by `Bruno Aristimunha`_). - Add Pan2023, Pan2025, PoloHortiguela2025 dataset loaders with synthetic regression coverage ({gh}`1197` by `Bruno Aristimunha`_). @@ -42,6 +43,7 @@ Enhancements API changes ~~~~~~~~~~~ +- :class:`moabb.evaluations.CrossSubjectEvaluation` now rejects custom splitter outputs with more than three items instead of silently treating all middle items as calibration data; custom splitters must yield ``(train, test)`` or ``(train, cal, test)`` (:gh:`1207` by `lindicaphxag-tech`_). - :class:`moabb.evaluations.CrossSessionEvaluation` custom cross-validation folds that hold out multiple sessions now emit one result row per held-out session instead of one aggregate row per fold. The default leave-one-session-out diff --git a/moabb/evaluations/base.py b/moabb/evaluations/base.py index b8fa875ae..983d02d49 100644 --- a/moabb/evaluations/base.py +++ b/moabb/evaluations/base.py @@ -756,22 +756,46 @@ def _build_eval_config(self, param_grid): "param_grid": None, # overridden per-task below if needed } - @staticmethod - def _preview_splits(splitter, y, metadata): + def _preview_splits(self, splitter, y, metadata): """Materialize folds up front with optional splitter metadata.""" preview = [] - # ``*cal`` absorbs the optional calibration slice from a transfer - # splitter; a plain 2-tuple splitter gives cal == [] (no calibration). - for cv_ind, (train_idx, *cal, test_idx) in enumerate(splitter.split(y, metadata)): - calib_idx = cal[0] if cal else train_idx[:0] + for cv_ind, split in enumerate(splitter.split(y, metadata)): + split = tuple(split) + if len(split) == 2: + train_idx, test_idx = split + calib_idx = np.asarray(train_idx)[:0] + elif len(split) == 3: + train_idx, calib_idx, test_idx = split + else: + raise ValueError( + "Cross-validation splitters must yield either " + "(train, test) or (train, calibration, test); " + f"fold {cv_ind} yielded {len(split)} slices." + ) + + self._validate_fold_indices( + train_idx, calib_idx, test_idx, n_samples=len(metadata), cv_ind=cv_ind + ) + split_metadata = None if hasattr(splitter, "get_metadata"): split_metadata = splitter.get_metadata() if split_metadata is not None: - split_metadata = dict(split_metadata) + split_metadata = deepcopy(dict(split_metadata)) preview.append((cv_ind, train_idx, calib_idx, test_idx, split_metadata)) return preview + def _validate_fold_indices( + self, train_idx, calib_idx, test_idx, *, n_samples, cv_ind + ): + """Hook for evaluation-specific validation of materialized fold indices.""" + return None + + def _validate_test_fold_metadata(self, test_metadata): + """Validate metadata assumptions made by the parallel task builder.""" + if test_metadata.empty: + raise ValueError("Cross-validation split produced an empty test fold.") + def _build_task_list( self, dataset, X, y, metadata, splitter, work_plan, pipelines, param_grid ): @@ -782,6 +806,7 @@ def _build_task_list( for cv_ind, train_idx, calib_idx, test_idx, split_meta in fold_preview: test_meta = metadata.iloc[test_idx] + self._validate_test_fold_metadata(test_meta) subject = test_meta["subject"].iloc[0] if subject not in work_plan: diff --git a/moabb/evaluations/evaluations.py b/moabb/evaluations/evaluations.py index 272aa5465..c8f845563 100644 --- a/moabb/evaluations/evaluations.py +++ b/moabb/evaluations/evaluations.py @@ -3,7 +3,12 @@ import numpy as np from sklearn.base import clone -from sklearn.model_selection import GroupKFold, LeaveOneGroupOut, StratifiedKFold +from sklearn.model_selection import ( + BaseCrossValidator, + GroupKFold, + LeaveOneGroupOut, + StratifiedKFold, +) from sklearn.preprocessing import LabelEncoder from tqdm import tqdm @@ -498,6 +503,17 @@ class CrossSubjectEvaluation(BaseEvaluation): ``"roc_auc"`` metrics. Cannot be combined with manual ``calibration_size`` or ``calibration_labeled`` in ``cv_kwargs``, except for the default ``TRAIN`` mode. + splitter : BaseCrossValidator or None + Optional top-level cross-subject splitter. It must follow MOABB's + ``split(y, metadata)`` contract and yield unique, in-range, + one-dimensional positional integer indices into ``y`` and + ``metadata`` for either train/test or train/calibration/test slices. + The slices must be pairwise disjoint. Each test fold must contain exactly + one subject, matching + MOABB's per-subject result-row semantics. When provided, it replaces + ``CrossSubjectSplitter`` and cannot + be combined with ``cv_class``, ``cv_kwargs``, ``n_splits``, + ``groups``, or a non-default ``cs_mode``. Defaults to ``None``. Notes ----- @@ -509,14 +525,54 @@ class CrossSubjectEvaluation(BaseEvaluation): _score_per_session = True _needs_all_subjects = True - def __init__(self, *args, cs_mode=CrossSubjectMode.TRAIN, **kwargs): - cv_kwargs = dict(kwargs.get("cv_kwargs") or {}) - internal_cv_keys = frozenset() - + def __init__( + self, + *args, + cs_mode=CrossSubjectMode.TRAIN, + splitter: Optional[BaseCrossValidator] = None, + **kwargs, + ): if cs_mode is None: cs_mode = CrossSubjectMode.TRAIN - cs_mode = CrossSubjectMode(cs_mode) + + if splitter is not None: + if not isinstance(splitter, BaseCrossValidator): + raise TypeError("splitter must be a sklearn BaseCrossValidator instance.") + + conflicts = [] + if kwargs.get("cv_class") is not None: + conflicts.append("cv_class") + if kwargs.get("cv_kwargs"): + conflicts.append("cv_kwargs") + if kwargs.get("n_splits") is not None: + conflicts.append("n_splits") + if kwargs.get("groups") is not None: + conflicts.append("groups") + if cs_mode != CrossSubjectMode.TRAIN: + conflicts.append("cs_mode") + if conflicts: + names = ", ".join(conflicts) + raise ValueError( + f"splitter cannot be combined with protocol options: {names}." + ) + + self.splitter = splitter + self.cs_mode = cs_mode + self.trialwise = False + additional_columns = list(kwargs.get("additional_columns") or ()) + for column in getattr(splitter, "metadata_columns", ()): + if column not in additional_columns: + additional_columns.append(column) + kwargs["additional_columns"] = additional_columns + super().__init__(*args, **kwargs) + self._cv_internal_keys = frozenset() + self._cv_explicit_keys = frozenset() + return + + self.splitter = None + cv_kwargs = dict(kwargs.get("cv_kwargs") or {}) + internal_cv_keys = frozenset() self.cs_mode = cs_mode # Manual cv_kwargs still work when the default train-only blockwise @@ -547,14 +603,72 @@ def __init__(self, *args, cs_mode=CrossSubjectMode.TRAIN, **kwargs): self._cv_internal_keys = internal_cv_keys self._cv_explicit_keys = frozenset(self.cv_kwargs) - internal_cv_keys + def _validate_fold_indices( + self, train_idx, calib_idx, test_idx, *, n_samples, cv_ind + ): + if self.splitter is None: + return + + named_indices = { + "train": np.asarray(train_idx), + "calibration": np.asarray(calib_idx), + "test": np.asarray(test_idx), + } + for name, indices in named_indices.items(): + if indices.ndim != 1 or not np.issubdtype(indices.dtype, np.integer): + raise TypeError( + "A top-level CrossSubjectEvaluation splitter must return " + f"one-dimensional integer positional indices; fold {cv_ind} " + f"{name} indices have shape {indices.shape} and dtype " + f"{indices.dtype}." + ) + if indices.size and (indices.min() < 0 or indices.max() >= n_samples): + raise ValueError( + "A top-level CrossSubjectEvaluation splitter returned " + f"out-of-range {name} indices in fold {cv_ind} for " + f"{n_samples} samples." + ) + if np.unique(indices).size != indices.size: + raise ValueError( + "A top-level CrossSubjectEvaluation splitter returned " + f"duplicate {name} indices in fold {cv_ind}." + ) + + for left, right in ( + ("train", "calibration"), + ("train", "test"), + ("calibration", "test"), + ): + if np.intersect1d(named_indices[left], named_indices[right]).size: + raise ValueError( + "A top-level CrossSubjectEvaluation splitter must return " + "disjoint train/calibration/test slices; " + f"fold {cv_ind} has overlapping {left} and {right} indices." + ) + + def _validate_test_fold_metadata(self, test_metadata): + super()._validate_test_fold_metadata(test_metadata) + if self.splitter is None: + return + test_subjects = test_metadata["subject"].unique() + if len(test_subjects) != 1: + raise ValueError( + "A top-level CrossSubjectEvaluation splitter must hold out " + "exactly one subject per test fold because MOABB records one " + "subject identity per result row; got test subjects " + f"{test_subjects.tolist()}." + ) + def _create_splitter(self): - """Create the CrossSubjectSplitter for parallel evaluation. + """Create the top-level splitter for parallel evaluation. + An explicit ``splitter`` is used directly. Otherwise, ``calibration_size`` and ``calibration_labeled`` passed via - ``cv_kwargs`` turn each fold into a transfer split: - - ``(train, calibration, test)``. + ``cv_kwargs`` configure the default ``CrossSubjectSplitter``. """ + if self.splitter is not None: + return self.splitter + if self.n_splits is None: default_class = LeaveOneGroupOut default_kwargs = {} diff --git a/moabb/tests/test_evaluations.py b/moabb/tests/test_evaluations.py index 24439aa24..a6b6ba782 100644 --- a/moabb/tests/test_evaluations.py +++ b/moabb/tests/test_evaluations.py @@ -13,6 +13,7 @@ from sklearn.discriminant_analysis import LinearDiscriminantAnalysis as LDA from sklearn.dummy import DummyClassifier as Dummy from sklearn.model_selection import ( + BaseCrossValidator, GroupKFold, GroupShuffleSplit, LeaveOneGroupOut, @@ -65,6 +66,143 @@ def __init__(self, n_splits=2, /, **kwargs): super().__init__(n_splits=n_splits, **kwargs) +class HeldOutSubjectSplitter(BaseCrossValidator): + """Minimal top-level splitter that holds out one chosen subject.""" + + metadata_columns = ("held_out_subject",) + + def __init__(self, subject): + self.subject = subject + self._metadata = None + + def split(self, y, metadata): + del y + positions = np.arange(len(metadata)) + test_mask = metadata["subject"].to_numpy() == self.subject + self._metadata = {"held_out_subject": self.subject} + yield positions[~test_mask], positions[test_mask] + + def get_n_splits(self, *args, **kwargs): + return 1 + + def get_metadata(self): + return self._metadata + + +class ThreeWayHeldOutSubjectSplitter(BaseCrossValidator): + """Hold out one subject and split its labels into calibration/test slices.""" + + metadata_columns = ("calibration_size", "calibration_labeled") + + def __init__(self, subject): + self.subject = subject + self._metadata = None + + def split(self, y, metadata): + positions = np.arange(len(metadata)) + target = positions[metadata["subject"].to_numpy() == self.subject] + train = positions[metadata["subject"].to_numpy() != self.subject] + + calibration_parts = [] + test_parts = [] + target_y = np.asarray(y)[target] + for label in np.unique(target_y): + label_positions = target[target_y == label] + split_at = max(1, len(label_positions) // 2) + split_at = min(split_at, len(label_positions) - 1) + calibration_parts.append(label_positions[:split_at]) + test_parts.append(label_positions[split_at:]) + + calibration = np.sort(np.concatenate(calibration_parts)) + test = np.sort(np.concatenate(test_parts)) + self._metadata = { + "calibration_size": len(calibration) / (len(calibration) + len(test)), + "calibration_labeled": False, + } + yield train, calibration, test + + def get_n_splits(self, *args, **kwargs): + return 1 + + def get_metadata(self): + return self._metadata + + +class MultiSubjectTestSplitter(BaseCrossValidator): + """Deliberately invalid splitter with two subjects in one test fold.""" + + def split(self, y, metadata): + del y + positions = np.arange(len(metadata)) + test_mask = metadata["subject"].isin([2, 3]).to_numpy() + yield positions[~test_mask], positions[test_mask] + + def get_n_splits(self, *args, **kwargs): + return 1 + + +class ReusedNestedMetadataSplitter(BaseCrossValidator): + """Reuse and mutate one nested metadata object across folds.""" + + metadata_columns = ("fold_trace",) + + def __init__(self): + self._metadata = {"fold_trace": {"fold": None, "history": []}} + + def split(self, y, metadata): + del y + positions = np.arange(len(metadata)) + midpoint = max(1, len(positions) // 2) + folds = ( + (positions[midpoint:], positions[:midpoint]), + (positions[:midpoint], positions[midpoint:]), + ) + for fold, (train, test) in enumerate(folds): + self._metadata["fold_trace"]["fold"] = fold + self._metadata["fold_trace"]["history"].append(fold) + yield train, test + + def get_n_splits(self, *args, **kwargs): + return 2 + + def get_metadata(self): + return self._metadata + + +class MalformedTopLevelSplitter(BaseCrossValidator): + """Top-level splitter used to exercise fold-index validation.""" + + def __init__(self, failure): + self.failure = failure + + def split(self, y, metadata): + del y + positions = np.arange(len(metadata)) + test_mask = metadata["subject"].to_numpy() == 2 + train = positions[~test_mask] + test = positions[test_mask] + + if self.failure == "overlap": + train = np.concatenate([train, test[:1]]) + yield train, test + elif self.failure == "duplicate": + yield np.concatenate([train, train[:1]]), test + elif self.failure == "float": + yield train.astype(float), test + elif self.failure == "out_of_range": + bad_test = test.copy() + bad_test[0] = len(metadata) + yield train, bad_test + elif self.failure == "four_way": + empty = train[:0] + yield train, empty, empty, test + else: + raise AssertionError(f"unknown failure mode {self.failure!r}") + + def get_n_splits(self, *args, **kwargs): + return 1 + + def _group_run(metadata): return metadata["run"].to_numpy() @@ -639,6 +777,180 @@ def test_custom_cv_receives_compatible_defaults_and_overrides(tmp_path): assert splitter.random_state == 17 +def test_cross_subject_accepts_top_level_splitter(tmp_path): + splitter = HeldOutSubjectSplitter(subject=2) + evaluation = ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[dataset], + hdf5_path=tmp_path, + splitter=splitter, + ) + + assert evaluation._create_splitter() is splitter + assert "held_out_subject" in evaluation.additional_columns + + _, y, metadata = FakeImageryParadigm().get_data(dataset) + folds = list(evaluation._create_splitter().split(y, metadata)) + + assert len(folds) == 1 + train, test = folds[0] + assert set(metadata.loc[train, "subject"]) == {1} + assert set(metadata.loc[test, "subject"]) == {2} + assert splitter.get_metadata() == {"held_out_subject": 2} + + +@pytest.mark.parametrize( + "protocol_kwargs", + [ + {"cv_class": GroupKFold}, + {"cv_kwargs": {"random_state": 17}}, + {"n_splits": 2}, + {"groups": "session"}, + {"cs_mode": ev.CrossSubjectMode.TRAIN_TRIALWISE}, + ], +) +def test_cross_subject_top_level_splitter_rejects_protocol_conflicts( + tmp_path, protocol_kwargs +): + with pytest.raises(ValueError, match="splitter cannot be combined"): + ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[dataset], + hdf5_path=tmp_path, + splitter=HeldOutSubjectSplitter(subject=2), + **protocol_kwargs, + ) + + +def test_cross_subject_top_level_splitter_rejects_multi_subject_test_fold(tmp_path): + ds = FakeDataset(["left_hand", "right_hand"], n_subjects=3, n_sessions=2, seed=18) + evaluation = ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[ds], + hdf5_path=tmp_path, + overwrite=True, + n_jobs=1, + splitter=MultiSubjectTestSplitter(), + ) + pipe = make_pipeline(Covariances("oas"), CSP(8), LDA()) + + with pytest.raises(ValueError, match="exactly one subject"): + evaluation.process(OrderedDict([("P", pipe)])) + + +@pytest.mark.parametrize( + "failure, message", + [ + ("overlap", "disjoint"), + ("duplicate", "duplicate train"), + ("float", "integer positional indices"), + ("out_of_range", "out-of-range test"), + ("four_way", "either \\(train, test\\) or \\(train, calibration, test\\)"), + ], +) +def test_cross_subject_top_level_splitter_rejects_invalid_fold_indices( + tmp_path, failure, message +): + splitter = MalformedTopLevelSplitter(failure) + evaluation = ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[dataset], + hdf5_path=tmp_path, + splitter=splitter, + ) + _, y, metadata = FakeImageryParadigm().get_data(dataset) + + with pytest.raises((TypeError, ValueError), match=message): + evaluation._preview_splits(splitter, y, metadata) + + +def test_cross_subject_top_level_splitter_metadata_is_snapshotted(tmp_path): + splitter = ReusedNestedMetadataSplitter() + evaluation = ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[dataset], + hdf5_path=tmp_path, + splitter=splitter, + ) + _, y, metadata = FakeImageryParadigm().get_data(dataset) + + preview = evaluation._preview_splits(splitter, y, metadata) + + assert preview[0][4]["fold_trace"] == {"fold": 0, "history": [0]} + assert preview[1][4]["fold_trace"] == {"fold": 1, "history": [0, 1]} + assert preview[0][4]["fold_trace"] is not preview[1][4]["fold_trace"] + + +def test_cross_subject_top_level_splitter_indices_are_positional(tmp_path): + splitter = HeldOutSubjectSplitter(subject=2) + _, y, metadata = FakeImageryParadigm().get_data(dataset) + metadata = metadata.copy() + metadata.index = np.arange(100, 100 + len(metadata)) + + train, test = next(splitter.split(y, metadata)) + + assert np.array_equal(test, np.flatnonzero(metadata["subject"].to_numpy() == 2)) + assert set(metadata.iloc[test]["subject"]) == {2} + assert not np.intersect1d(train, test).size + + +def test_cross_subject_top_level_splitter_type_is_validated(tmp_path): + with pytest.raises(TypeError, match="BaseCrossValidator"): + ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[dataset], + hdf5_path=tmp_path, + splitter=object(), + ) + + +def test_cross_subject_top_level_splitter_runs_end_to_end(tmp_path): + ds = FakeDataset(["left_hand", "right_hand"], n_subjects=3, n_sessions=2, seed=18) + splitter = HeldOutSubjectSplitter(subject=3) + evaluation = ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[ds], + hdf5_path=tmp_path, + overwrite=True, + n_jobs=1, + splitter=splitter, + ) + pipe = make_pipeline(Covariances("oas"), CSP(8), LDA()) + + results = evaluation.process(OrderedDict([("P", pipe)])) + + # Result tables serialize dataset metadata values as strings. + assert set(results["subject"]) == {"3"} + assert set(results["held_out_subject"]) == {3} + + +def test_cross_subject_top_level_three_way_splitter_routes_calibration(tmp_path): + from sklearn import config_context + + _TRANSFER_CAPTURE.clear() + with config_context(enable_metadata_routing=True): + step = _TransferRecorder().set_fit_request(subjects=True, X_target_unlabeled=True) + pipe = make_pipeline(Covariances("oas"), step, CSP(8), LDA()) + ds = FakeDataset(["left_hand", "right_hand"], n_subjects=3, n_sessions=2, seed=19) + splitter = ThreeWayHeldOutSubjectSplitter(subject=3) + evaluation = ev.CrossSubjectEvaluation( + paradigm=FakeImageryParadigm(), + datasets=[ds], + hdf5_path=tmp_path, + overwrite=True, + n_jobs=1, + splitter=splitter, + ) + + results = evaluation.process(OrderedDict([("T", pipe)])) + + assert len(results) > 0 + assert set(results["subject"]) == {"3"} + assert _TRANSFER_CAPTURE, "transfer step was never fitted" + assert all(capture["n_subjects"] > 0 for capture in _TRANSFER_CAPTURE) + assert all(capture["n_target"] > 0 for capture in _TRANSFER_CAPTURE) + + @pytest.fixture(scope="module") def small_evaluation_data(): _, y, metadata = FakeImageryParadigm().get_data(dataset)