Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/source/whats_new.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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`_).
Expand All @@ -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
Expand Down
39 changes: 32 additions & 7 deletions moabb/evaluations/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
):
Expand All @@ -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:
Expand Down
134 changes: 124 additions & 10 deletions moabb/evaluations/evaluations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
-----
Expand All @@ -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
Expand Down Expand Up @@ -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 = {}
Expand Down
Loading
Loading