Skip to content
Merged
Changes from 2 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
32 changes: 9 additions & 23 deletions examples/how_to_benchmark/plot_cross_subject_transfer_rpa.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,10 +31,9 @@
import pandas as pd
from matplotlib.patches import Rectangle
from pyriemann.estimation import Covariances
from pyriemann.preprocessing import Whitening
from pyriemann.tangentspace import TangentSpace
from pyriemann.transfer import TLCenter
from sklearn import config_context
from sklearn.base import BaseEstimator, TransformerMixin
from sklearn.linear_model import LogisticRegression
from sklearn.pipeline import make_pipeline

Expand Down Expand Up @@ -121,7 +120,7 @@
# ``transform``.


class RiemannianAlignment(TransformerMixin, BaseEstimator):
class RiemannianAlignment(TLCenter):
"""Recenter source and target covariance matrices by domain."""

def fit(self, X, y=None, *, subjects=None, X_target_unlabeled=None):
Expand All @@ -130,30 +129,17 @@ def fit(self, X, y=None, *, subjects=None, X_target_unlabeled=None):
"RiemannianAlignment needs `subjects` and `X_target_unlabeled` metadata."
)

X = np.asarray(X)
subjects = np.asarray(subjects)
self.source_whiteners_ = {
subject: Whitening(metric="riemann").fit(X[subjects == subject])
for subject in np.unique(subjects)
}
self.target_whitener_ = Whitening(metric="riemann").fit(
np.asarray(X_target_unlabeled)
)
X_target_unlabeled = np.asarray(X_target_unlabeled)
X_ = np.vstack((np.asarray(X), X_target_unlabeled))
sub = np.repeat(self.target_domain, X_target_unlabeled.shape[0])
subjects_ = np.concatenate((np.asarray(subjects), sub))
self = super().fit(X_, subjects_ + "/0")
return self

def fit_transform(self, X, y=None, *, subjects=None, X_target_unlabeled=None):
"""Fit domain references and align the source training trials."""
self.fit(X, y, subjects=subjects, X_target_unlabeled=X_target_unlabeled)
subjects = np.asarray(subjects)
X_aligned = np.empty_like(X)
for subject, whitener in self.source_whiteners_.items():
mask = subjects == subject
X_aligned[mask] = whitener.transform(X[mask])
return X_aligned

def transform(self, X):
"""Align unseen trials with the target reference."""
return self.target_whitener_.transform(X)
return super().transform(X)


###############################################################################
Expand All @@ -175,7 +161,7 @@ def transform(self, X):
# ``transform_input``.

with config_context(enable_metadata_routing=True):
alignment = RiemannianAlignment().set_fit_request(
alignment = RiemannianAlignment("target_unlabeled").set_fit_request(
subjects=True, X_target_unlabeled=True
)

Expand Down
Loading