SPDIM: Source-Free Domain Adaptation on SPD Manifolds#

This example reproduces the SPDIM pipeline [Li et al., 2025] for source-free unsupervised domain adaptation (SFUDA) on SPD manifolds, using the geometric operations in spd_learn.functional.

Introduction#

SPDIM (SPD Information Maximization) [Li et al., 2025] is a source-free unsupervised domain adaptation method for EEG-based BCIs that operates directly on the SPD manifold. It adapts a pre-trained TSMNet model to a new session/subject using only unlabeled target data.

The Problem: Label Shift on SPD Manifolds#

In EEG-based BCIs, covariance features \(C_i \in \mathcal{S}_{++}^D\) suffer from conditional shift (distribution shifts across domains) and label shift (the class priors change resulting in class-imbalance).

A natural baseline is the Recentering Transform (RCT) [Zanini et al., 2017], which centers SPD features around the identity by applying the congruence transformation:

\[\tilde{C}_i = \bar{C}_j^{-1/2} \, C_i \, \bar{C}_j^{-1/2}\]

where \(\bar{C}_j\) is the Fréchet mean of the target domain. However, the paper’s Proposition 2 shows that RCT only compensates conditional shift when the label priors are identical across domains. Under label shift, \(\bar{C}_j\) is biased toward the over-represented class, causing RCT to misalign.

SPDIM Overview#

SPDIM addresses this by learning the transport parameters via an Information Maximization (IM) loss that does not require labels. SPDIM learns a full SPD reference matrix \(\Phi_j\) that replaces the standard centering (Eq. 19 in the paper):

\[\tilde{C}_i = \Phi_j^{1/2} \, \bar{C}_j^{-1/2} \, C_i \, \bar{C}_j^{-1/2} \, \Phi_j^{1/2}\]

It initializes \(\Phi_j\) with the target Fréchet mean and optimizes it as an SPD-constrained parameter via torch.nn.utils.parametrize and SymmetricPositiveDefinite.

SPDIM optimizes the IM loss (Eq. 21):

\[\mathcal{L}_{\mathrm{IM}} = \underbrace{H(Y | X)}_{\text{conditional entropy}} - \underbrace{H(\bar{Y})}_{\text{marginal entropy}}\]

Here, \(H(Y \mid X)\) is the conditional entropy of the model predictions for each target sample. \(H(Y)\) is the marginal entropy of the predicted labels, estimated here by \(H(\bar{Y})\), the entropy of the average predictive distribution across the target set. This encourages confident predictions (low \(H(Y \mid X)\)) while maintaining class diversity (high \(H(Y)\)).

Setup and Imports#

import math
import warnings

import matplotlib.pyplot as plt
import numpy as np
import torch


warnings.filterwarnings("ignore")

SPDIM Geometric Operations#

The Fréchet mean and geodesic distances used by SPDIM are available directly from spd_learn.functional.

Loading the Dataset#

We use BNCI2015_001 (2-class motor imagery: right hand vs feet), the same dataset used in the SPDIM paper. This dataset has 12 subjects, each with two sessions.

We demonstrate cross-session transfer on Subject 7, which exhibits meaningful session-to-session variability.

  • Source domain: Session A (training with labels)

  • Target domain: Session B (adaptation without labels)

from braindecode.datasets import create_from_X_y
from moabb.datasets import BNCI2015_001
from moabb.paradigms import MotorImagery
from sklearn.preprocessing import LabelEncoder

from spd_learn.functional import frechet_mean


dataset = BNCI2015_001()
paradigm = MotorImagery(
    n_classes=2,
    events=["right_hand", "feet"],
    fmin=4,
    fmax=36,
    tmin=1.0,
    tmax=4.0,
    resample=256,
)

subject_id = 7

cache_config = dict(
    save_raw=True,
    save_epochs=True,
    save_array=True,
    use=True,
    overwrite_raw=False,
    overwrite_epochs=False,
    overwrite_array=False,
)

X, labels, meta = paradigm.get_data(
    dataset=dataset,
    subjects=[subject_id],
    cache_config=cache_config,
)

le = LabelEncoder()
y = le.fit_transform(labels)

sfreq = int(paradigm.resample)

# Split by session
sessions = meta["session"].unique()
source_session, target_session = sessions[0], sessions[1]

source_idx = meta.query(f"session == '{source_session}'").index.to_numpy()
target_idx = meta.query(f"session == '{target_session}'").index.to_numpy()

X_source, y_source = X[source_idx], y[source_idx]
X_target, y_target = X[target_idx], y[target_idx]

# Create braindecode dataset for target domain
target_ds = create_from_X_y(
    X_target,
    y_target,
    drop_last_window=True,
    sfreq=sfreq,
)

print(f"Dataset: {dataset.code}")
print(f"Subject {subject_id}: Session {source_session} -> {target_session}")
print(f"Source domain: {len(X_source)} samples")
print(f"Target domain: {len(target_ds)} samples")
print(f"Classes: {le.classes_}")
This is nemar-py 0.3.0.
Preparing to download nm000140 from https://data.nemar.org/
Retrieving 5 of 349 manifest files (16 concurrent downloads).
primary backend failed; falling back to next layer: 'README.md' is not annexed (no sha256/md5 checksum on the manifest entry); the S3 backend has nothing to fetch.

Overall:   0%|          | 0.00/12.4k [00:00<?, ?B/s]
Overall:   3%|▎         | 366/12.4k [00:00<00:05, 2.18kB/s]
Overall:  42%|████▏     | 5.23k/12.4k [00:00<00:00, 12.3kB/s]
Overall:  42%|████▏     | 5.23k/12.4k [00:00<00:00, 11.2kB/s]
Finished downloading nm000140 v1.0.2.
This is nemar-py 0.3.0.
Preparing to download nm000140 from https://data.nemar.org/
Retrieving 6 of 349 manifest files (16 concurrent downloads).

S3:   0%|          | 0.00/124M [00:00<?, ?B/s]
S3:   1%|          | 1.00M/124M [00:00<01:47, 1.20MB/s]
S3:   2%|▏         | 3.00M/124M [00:01<00:33, 3.76MB/s]
S3:   4%|▍         | 5.00M/124M [00:01<00:20, 6.15MB/s]
S3:   9%|▉         | 11.0M/124M [00:01<00:07, 15.1MB/s]
S3:  17%|█▋        | 21.0M/124M [00:01<00:03, 30.0MB/s]
S3:  28%|██▊       | 35.0M/124M [00:01<00:01, 48.9MB/s]
S3:  40%|████      | 50.0M/124M [00:01<00:01, 67.6MB/s]
S3:  51%|█████     | 63.0M/124M [00:01<00:00, 82.6MB/s]
S3:  59%|█████▊    | 73.0M/124M [00:01<00:00, 84.8MB/s]
S3:  68%|██████▊   | 84.0M/124M [00:02<00:00, 91.8MB/s]
S3:  76%|███████▋  | 95.0M/124M [00:02<00:00, 96.9MB/s]
S3:  85%|████████▌ | 106M/124M [00:02<00:00, 98.5MB/s]
S3:  97%|█████████▋| 121M/124M [00:02<00:00, 102MB/s]
S3: 100%|██████████| 124M/124M [00:02<00:00, 53.0MB/s]
Finished downloading nm000140 v1.0.2.
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/README'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/participants.tsv'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/participants.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_space-CapTrak_electrodes.tsv'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_space-CapTrak_coordsystem.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_space-CapTrak_electrodes.json'...
The provided raw data contains annotations, but you did not pass an "event_id" mapping from annotation descriptions to event codes. We will generate arbitrary event codes. To specify custom event codes, please pass "event_id".
Used Annotations descriptions: ['feet', 'right_hand']
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_events.tsv'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_events.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/dataset_description.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.json'...
Copying data files to sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.edf
Found no extension for raw file, assuming "BTi" format and appending extension .pdf
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_channels.tsv'...
Converting data files to EDF format
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/sub-7_ses-0A_scans.tsv'...
Wrote /home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/sub-7_ses-0A_scans.tsv entry with eeg/sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.edf.
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-0A/eeg/sub-7_ses-0A_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/participants.tsv'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/participants.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_space-CapTrak_electrodes.tsv'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_space-CapTrak_coordsystem.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_space-CapTrak_electrodes.json'...
The provided raw data contains annotations, but you did not pass an "event_id" mapping from annotation descriptions to event codes. We will generate arbitrary event codes. To specify custom event codes, please pass "event_id".
Used Annotations descriptions: ['feet', 'right_hand']
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_events.tsv'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_events.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/dataset_description.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.json'...
Copying data files to sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.edf
Found no extension for raw file, assuming "BTi" format and appending extension .pdf
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_channels.tsv'...
Converting data files to EDF format
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/sub-7_ses-1B_scans.tsv'...
Wrote /home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/sub-7_ses-1B_scans.tsv entry with eeg/sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.edf.
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/sub-7/ses-1B/eeg/sub-7_ses-1B_task-imagery_run-0_desc-0b679affc63b69abcbdaf494a48abf16_eeg.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/dataset_description.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/dataset_description.json'...
Writing '/home/runner/mne_data/MNE-BIDS-bnci2015-001/dataset_description.json'...
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Creating RawArray with float64 data, n_channels=13, n_times=768
    Range : 0 ... 767 =      0.000 ...     2.996 secs
Ready.
Dataset: BNCI2015-001
Subject 7: Session 0A -> 1B
Source domain: 200 samples
Target domain: 200 samples
Classes: ['feet' 'right_hand']

Simulating Label Shift in Target Domain#

A key contribution of SPDIM is handling label shift — when the class priors differ between source and target domains.

Following the paper’s get_label_ratio protocol with ratio_level=0.2: we keep all samples of the last class and subsample the other class(es) to 20%. This creates a 5:1 class imbalance, making the Fréchet mean biased toward the majority class.

As shown by the paper’s Proposition 2, this biased mean causes RCT to misalign: the recentered features no longer align with the source domain’s learned decision boundary.

ratio_level = 0.2

# SPDIM protocol: subsample all classes except the last to ratio_level
rng = np.random.RandomState(42)
classes = sorted(np.unique(y_target))
subsample_inds = np.sort(
    np.concatenate(
        [
            rng.choice(
                np.flatnonzero(y_target == c),
                size=math.ceil(
                    np.sum(y_target == c)
                    * (ratio_level if i < len(classes) - 1 else 1.0)
                ),
                replace=False,
            )
            for i, c in enumerate(classes)
        ]
    )
)

target_shifted_ds = target_ds.split(by={"shifted": subsample_inds.tolist()})["shifted"]

# Keep arrays for SPDIM adaptation methods
X_target_shifted = X_target[subsample_inds]
y_target_shifted = y_target[subsample_inds]

print(f"\nAfter label shift (ratio_level={ratio_level}):")
print(f"  Target samples: {len(target_shifted_ds)}")
for c in np.unique(y_target_shifted):
    n = (y_target_shifted == c).sum()
    print(f"  Class {le.classes_[c]}: {n} ({100 * n / len(y_target_shifted):.0f}%)")
After label shift (ratio_level=0.2):
  Target samples: 120
  Class feet: 20 (17%)
  Class right_hand: 100 (83%)

Training TSMNet on Source Domain#

We train TSMNet using braindecode’s EEGClassifier wrapper with:

  • AdamW optimizer with lr=1e-3 and weight_decay=1e-4

  • Gradient clipping (max norm 1.0) for stable SPD optimization

  • Validation split (10%) for early stopping

from braindecode import EEGClassifier
from skorch.callbacks import GradientNormClipping
from skorch.dataset import ValidSplit

from spd_learn.models import TSMNet


n_chans = X_source.shape[1]
n_outputs = len(le.classes_)

torch.manual_seed(42)

device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"\nUsing device: {device}")

model = TSMNet(
    n_chans=n_chans,
    n_outputs=n_outputs,
    n_temp_filters=4,
    temp_kernel_length=25,
    n_spatiotemp_filters=40,
    n_bimap_filters=20,
    reeig_threshold=1e-4,
)

clf = EEGClassifier(
    model,
    criterion=torch.nn.CrossEntropyLoss,
    optimizer=torch.optim.AdamW,
    optimizer__lr=1e-4,
    optimizer__weight_decay=1e-4,
    train_split=ValidSplit(0.1, stratified=True, random_state=42),
    batch_size=32,
    max_epochs=30,  # Reduced from 200 for faster documentation build
    callbacks=[
        ("gradient_clip", GradientNormClipping(gradient_clip_value=5.0)),
    ],
    device=device,
    verbose=1,
)

print("\n" + "=" * 50)
print("Training TSMNet on Source Domain")
print("=" * 50)
clf.fit(X_source, y_source)
Using device: cpu

==================================================
Training TSMNet on Source Domain
==================================================
  epoch    train_loss    valid_acc    valid_loss     dur
-------  ------------  -----------  ------------  ------
      1        0.6967       0.3000        0.7093  0.2710
      2        0.6954       0.3500        0.7065  0.2624
      3        0.6912       0.3500        0.7045  0.2736
      4        0.6892       0.4500        0.7027  0.2387
      5        0.6850       0.4500        0.7012  0.2482
      6        0.6826       0.4500        0.6999  0.2892
      7        0.6841       0.4500        0.6988  0.3684
      8        0.6795       0.4000        0.6976  0.3814
      9        0.6743       0.4000        0.6967  0.3929
     10        0.6737       0.4000        0.6957  0.3625
     11        0.6721       0.6000        0.6951  0.3792
     12        0.6685       0.5500        0.6940  0.2745
     13        0.6681       0.3500        0.6930  0.2558
     14        0.6631       0.3500        0.6920  0.2507
     15        0.6649       0.4000        0.6909  0.2471
     16        0.6640       0.5000        0.6896  0.2624
     17        0.6563       0.4500        0.6884  0.2486
     18        0.6573       0.5000        0.6873  0.2442
     19        0.6570       0.5000        0.6861  0.2457
     20        0.6539       0.5500        0.6855  0.2763
     21        0.6554       0.6000        0.6854  0.2396
     22        0.6529       0.6500        0.6846  0.2477
     23        0.6476       0.6000        0.6834  0.2748
     24        0.6495       0.6500        0.6826  0.2396
     25        0.6443       0.6000        0.6810  0.2431
     26        0.6460       0.5500        0.6791  0.2755
     27        0.6431       0.5500        0.6780  0.2485
     28        0.6424       0.5500        0.6769  0.2366
     29        0.6378       0.6000        0.6763  0.2677
     30        0.6344       0.6500        0.6756  0.2510
<class 'braindecode.classifier.EEGClassifier'>[initialized](
  module_=TSMNet(
    (cnn): Sequential(
      (0): Conv2d(1, 4, kernel_size=(1, 25), stride=(1, 1), padding=same, padding_mode=reflect)
      (1): Conv2d(4, 40, kernel_size=(13, 1), stride=(1, 1))
      (2): Flatten(start_dim=2, end_dim=-1)
    )
    (covpool): CovLayer()
    (spdnet): Sequential(
      (0): ParametrizedBiMap(
        (parametrizations): ModuleDict(
          (weight): ParametrizationList(
            (0): _Orthogonal()
          )
        )
      )
      (1): ReEig()
    )
    (spdbnorm): ParametrizedSPDBatchNormMeanVar(
      (parametrizations): ModuleDict(
        (weight): ParametrizationList(
          (0): PositiveDefiniteScalar()
        )
        (bias): ParametrizationList(
          (0): SymmetricPositiveDefinite()
        )
      )
    )
    (logeig): Sequential(
      (0): LogEig()
      (1): Flatten(start_dim=1, end_dim=-1)
    )
    (head): Linear(in_features=210, out_features=2, bias=True)
  ),
)
In a Jupyter environment, please rerun this cell to show the HTML representation or trust the notebook.
On GitHub, the HTML representation is unable to render, please try loading this page with nbviewer.org.


Baseline: No Adaptation#

Evaluate the source-trained model on target domain without adaptation.

from sklearn.metrics import balanced_accuracy_score


underlying_model = clf.module_

y_pred_source = clf.predict(X_source)
source_bacc = balanced_accuracy_score(y_source, y_pred_source)

y_pred_target_no_adapt = clf.predict(target_shifted_ds)
no_adapt_bacc = balanced_accuracy_score(y_target_shifted, y_pred_target_no_adapt)

print(f"\n{'=' * 50}")
print("Baseline: No Adaptation")
print(f"{'=' * 50}")
print(f"Source Balanced Accuracy: {source_bacc * 100:.2f}%")
print(f"Target Balanced Accuracy: {no_adapt_bacc * 100:.2f}%")
print(f"Performance Drop: {(source_bacc - no_adapt_bacc) * 100:.2f}%")
==================================================
Baseline: No Adaptation
==================================================
Source Balanced Accuracy: 82.00%
Target Balanced Accuracy: 87.00%
Performance Drop: -5.00%

Helper Functions#

For SPDIM, we intercept TSMNet’s forward pass before SPDBatchNormMeanVar and replace the centering step with learnable geodesic transport, while keeping the variance normalization and rebiasing from the trained BN layer.

from spd_learn.modules import LogEig


def extract_spd_features(model, X_data, batch_size=32):
    """Extract SPD features before batch normalization."""
    model.eval()
    if not isinstance(X_data, torch.Tensor):
        dtype = next(model.parameters()).dtype
        X_data = torch.tensor(X_data, dtype=dtype)
    dev = next(model.parameters()).device
    X_data = X_data.to(dev)

    spd_list = []
    with torch.no_grad():
        for i in range(0, len(X_data), batch_size):
            batch = X_data[i : i + batch_size]
            x = model.cnn(batch[:, None, ...])
            x = model.covpool(x)
            x = model.spdnet(x)
            spd_list.append(x.cpu())

    return torch.cat(spd_list, dim=0)


def spdim_forward(model, X_spd, adapter=None):
    """SPDIM test-time forward pass.

    Matches the original SPDIM test-time pipeline
    (``transp_geosedic_identity_transp``): geodesic transport + LogEig
    + classifier, without dispersion normalization.

    1. Geodesic transport: A^{-t/2} X A^{-t/2}
    2. LogEig (tangent space mapping)
    3. Classifier
    """
    # Geodesic transport (replaces BN centering)
    if adapter is not None:
        X_transported = adapter(X_spd)

    # LogEig + classifier
    logeig = LogEig(upper=True, flatten=True)
    X_tangent = logeig(X_transported)
    dev = next(model.parameters()).device
    dtype = next(model.parameters()).dtype
    logits = model.head(X_tangent.to(dev, dtype=dtype))
    return logits

SFUDA Step 1: Refit BN Statistics (RCT Baseline)#

The Recentering Transform (RCT) [Zanini et al., 2017] baseline recomputes the Fréchet mean and variance on target SPD features using the full Karcher flow. This corresponds to setting \(\varphi = 1\) (standard centering) in the geodesic transport.

Under label shift, Proposition 2 predicts that this will degrade performance because the biased Fréchet mean shifts features away from the source decision boundary.

# Save original running stats
orig_running_mean = underlying_model.spdbnorm.running_mean.clone()
orig_running_var = underlying_model.spdbnorm.running_var.clone()


def refit_spdbn_frechet(model, X_data, batch_size=32):
    """Refit SPDBatchNormMeanVar using the Fréchet mean (SPDIM style)."""
    X_spd = extract_spd_features(model, X_data, batch_size=batch_size)
    mean, distances = frechet_mean(X_spd, max_iter=50, return_distances=True)
    variance = distances.square().mean(dim=0, keepdim=True).squeeze()
    with torch.no_grad():
        model.spdbnorm.running_mean.copy_(mean)
        model.spdbnorm.running_var.fill_(variance.item())


print(f"\n{'=' * 50}")
print("SFUDA Step 1: Refit BN Statistics (RCT)")
print(f"{'=' * 50}")

refit_spdbn_frechet(underlying_model, X_target_shifted)

target_frechet_mean = underlying_model.spdbnorm.running_mean.clone()

rct_pred = clf.predict(target_shifted_ds)
rct_bacc = balanced_accuracy_score(y_target_shifted, rct_pred)
print(f"RCT Balanced Accuracy: {rct_bacc * 100:.2f}%")
print(f"Improvement over baseline: {(rct_bacc - no_adapt_bacc) * 100:+.2f}%")

# Restore original stats
underlying_model.spdbnorm.running_mean.copy_(orig_running_mean)
underlying_model.spdbnorm.running_var.copy_(orig_running_var)
==================================================
SFUDA Step 1: Refit BN Statistics (RCT)
==================================================
RCT Balanced Accuracy: 70.50%
Improvement over baseline: -16.50%

tensor([[5.0514]])

Information Maximization Loss#

The IM loss encourages confident predictions (low conditional entropy) while maintaining class diversity (high marginal entropy).

import torch.nn.functional as F


def im_loss(logits, temperature=2.0):
    """Information Maximization loss (matching SPDIM paper)."""
    p = F.softmax(logits / temperature, dim=1)
    ce = -(p * torch.log(p + 1e-5)).sum(dim=1).mean()
    p_bar = p.mean(dim=0)
    me = -(p_bar * torch.log(p_bar + 1e-5)).sum()
    return ce - me

SPDIM(bias) Strategy#

SPDIM(bias) (Eq. 19) learns a full SPD reference mean that replaces the (biased) Fréchet mean in the geodesic transport. With \(D(D+1)/2\) degrees of freedom (vs 1 scalar for geodesic), it can compensate both conditional and label shift.

We initialize the learnable mean with the target Fréchet mean and keep it on the SPD manifold via SymmetricPositiveDefinite.

Learnable SPD Recenter Module#

from torch.nn.utils.parametrize import register_parametrization

from spd_learn.functional import matrix_inv_sqrt
from spd_learn.modules.manifold import SymmetricPositiveDefinite


class SPDLearnableRecenter(torch.nn.Module):
    def __init__(
        self,
        num_features,
        device=None,
        dtype=None,
    ):
        super().__init__()
        self.num_features = num_features

        self.bias = torch.nn.Parameter(
            torch.empty(1, num_features, num_features, device=device, dtype=dtype),
        )

        self.reset_parameters()
        register_parametrization(self, "bias", SymmetricPositiveDefinite())

    @torch.no_grad()
    def reset_parameters(self) -> None:
        self.bias.zero_()
        self.bias[0].fill_diagonal_(1.0)

    def forward(self, input):
        bias_inv_sqrt = matrix_inv_sqrt.apply(self.bias)
        output = bias_inv_sqrt @ input @ bias_inv_sqrt
        return output

SPDIM(bias) Optimization#

print(f"\n{'=' * 50}")
print("SPDIM(bias): Learnable SPD Mean")
print(f"{'=' * 50}")

X_spd_target = extract_spd_features(underlying_model, X_target_shifted, batch_size=32)

print(f"SPD reference initialized. Shape: {target_frechet_mean.shape}")

adapter = SPDLearnableRecenter(target_frechet_mean.shape[-1])
adapter.bias = target_frechet_mean.clone()

optimizer_bias = torch.optim.Adam(adapter.parameters(), lr=0.05)
n_epochs_bias = 30  # Reduced from 200 for faster documentation build
losses_bias = []
best_loss_bias = float("inf")
best_bias = target_frechet_mean.clone().detach()
for epoch in range(n_epochs_bias):
    optimizer_bias.zero_grad()
    logits = spdim_forward(
        underlying_model,
        X_spd_target,
        adapter,
    )
    loss = im_loss(logits, temperature=2.0)
    loss.backward()
    optimizer_bias.step()

    current_loss = loss.item()
    losses_bias.append(current_loss)
    if current_loss < best_loss_bias:
        best_loss_bias = current_loss
        best_bias = adapter.bias.clone().detach()

    if (epoch + 1) % 10 == 0 or epoch == 0:
        print(f"  Epoch {epoch + 1:3d}/{n_epochs_bias}: loss={current_loss:.4f}")

# Evaluate with best parameters
with torch.no_grad():
    adapter.bias = best_bias
    logits = spdim_forward(
        underlying_model,
        X_spd_target,
        adapter,
    )
    y_pred_bias = logits.argmax(dim=1).cpu().numpy()

bias_bacc = balanced_accuracy_score(y_target_shifted, y_pred_bias)
print(f"\nSPDIM(bias) Balanced Accuracy: {bias_bacc * 100:.2f}%")
print(f"Improvement over baseline: {(bias_bacc - no_adapt_bacc) * 100:+.2f}%")
==================================================
SPDIM(bias): Learnable SPD Mean
==================================================
SPD reference initialized. Shape: torch.Size([1, 20, 20])
  Epoch   1/30: loss=-0.0028
  Epoch  10/30: loss=-0.0044
  Epoch  20/30: loss=-0.0051
  Epoch  30/30: loss=-0.0054

SPDIM(bias) Balanced Accuracy: 81.00%
Improvement over baseline: -6.00%

Results Summary#

results = {
    "No Adaptation": no_adapt_bacc,
    "RCT (Refit BN)": rct_bacc,
    "SPDIM(bias)": bias_bacc,
}

print(f"\n{'=' * 60}")
print(f"Results Summary (Subject {subject_id}, Label Shift ratio={ratio_level})")
print(f"{'=' * 60}")
print(f"{'Method':<25} {'Bal. Accuracy':>14} {'vs Baseline':>14}")
print("-" * 57)
for method, acc in results.items():
    if method == "No Adaptation":
        print(f"{method:<25} {acc * 100:>12.2f}% {'-':>14}")
    else:
        imp = acc - no_adapt_bacc
        print(f"{method:<25} {acc * 100:>12.2f}% {imp * 100:>+12.2f}%")
print("-" * 57)
print("Chance level: 50.00% (2 classes)")

best_method = max(results.keys(), key=lambda k: results[k])
print(f"\nBest method: {best_method} ({results[best_method] * 100:.2f}%)")
============================================================
Results Summary (Subject 7, Label Shift ratio=0.2)
============================================================
Method                     Bal. Accuracy    vs Baseline
---------------------------------------------------------
No Adaptation                    87.00%              -
RCT (Refit BN)                   70.50%       -16.50%
SPDIM(bias)                      81.00%        -6.00%
---------------------------------------------------------
Chance level: 50.00% (2 classes)

Best method: No Adaptation (87.00%)

Visualizing Results#

fig, axes = plt.subplots(1, 2, figsize=(14, 5))

# 1. Bar chart
ax1 = axes[0]
methods = list(results.keys())
accuracies = [results[m] * 100 for m in methods]
colors = ["#e74c3c", "#3498db", "#2ecc71"]
bars = ax1.bar(
    range(len(methods)),
    accuracies,
    color=colors,
    edgecolor="black",
    linewidth=1.5,
)
ax1.set_xticks(range(len(methods)))
ax1.set_xticklabels(methods, rotation=35, ha="right", fontsize=9)
ax1.set_ylabel("Balanced Accuracy (%)", fontsize=12)
ax1.set_title("Domain Adaptation Comparison", fontsize=14)
ax1.set_ylim([0, 100])
ax1.axhline(y=50, color="gray", linestyle="--", alpha=0.5, label="Chance (50%)")
ax1.axhline(
    y=source_bacc * 100,
    color="blue",
    linestyle=":",
    alpha=0.5,
    label=f"Source ({source_bacc * 100:.1f}%)",
)
for bar, acc in zip(bars, accuracies):
    ax1.text(
        bar.get_x() + bar.get_width() / 2,
        bar.get_height() + 1.5,
        f"{acc:.1f}%",
        ha="center",
        va="bottom",
        fontsize=8,
        fontweight="bold",
    )
ax1.legend(loc="lower right", fontsize=8)

# 2. SPDIM(bias) loss curve
ax2 = axes[1]
ax2.plot(range(1, len(losses_bias) + 1), losses_bias, "r-", linewidth=2)
ax2.set_xlabel("Epoch", fontsize=12)
ax2.set_ylabel("IM Loss", fontsize=12)
ax2.set_title("SPDIM(bias) Optimization", fontsize=14)
ax2.grid(True, alpha=0.3)

plt.tight_layout()
plt.show()
Domain Adaptation Comparison, SPDIM(bias) Optimization

Training History#

fig, ax = plt.subplots(1, 1, figsize=(8, 5))
history = clf.history
epochs_hist = range(1, len(history) + 1)
ax.plot(epochs_hist, history[:, "train_loss"], "b-", label="Train Loss", linewidth=2)
ax.plot(epochs_hist, history[:, "valid_loss"], "r--", label="Valid Loss", linewidth=2)
ax.set_xlabel("Epoch", fontsize=12)
ax.set_ylabel("Loss", fontsize=12)
ax.set_title("TSMNet Training History", fontsize=14)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
TSMNet Training History

Discussion#

In this example we reproduced the SPDIM pipeline for source-free domain adaptation on SPD manifolds. The results illustrate the paper’s theoretical predictions.

Why RCT degrades under label shift#

The Recentering Transform (RCT) [Zanini et al., 2017] computes the Fréchet mean of the target domain and uses it to center the SPD features. Under label shift, the Fréchet mean is biased toward the over-represented class (here, feet at 83% of samples).

As predicted by Proposition 2 of the paper, this biased mean causes the recentered features to no longer align with the source domain’s learned decision boundary, resulting in degraded accuracy. Without label shift, RCT typically improves accuracy.

Key implementation details#

The following details match the original SPDIM code:

  • Temperature = 2.0 in the IM loss softmax (T=0.8 for multi-class tasks).

  • Best-model tracking: Returns the parameter with lowest IM loss.

  • Test-time BN: Only geodesic transport (no dispersion normalization), matching the original SPDIM test-time pipeline.

References#

[1]

Reinmar J Kobler, Jun-ichiro Hirayama, Qibin Zhao, and Motoaki Kawanabe. Spd domain-specific batch normalization to crack interpretable unsupervised domain adaptation in eeg. In Advances in Neural Information Processing Systems, volume 35, 6219–6235. 2022. URL: https://proceedings.neurips.cc/paper_files/paper/2022/hash/28ef7ee7cd3e03093acc39e1272411b7-Abstract-Conference.html.

[2] (1,2,3)

Paolo Zanini, Marco Congedo, Christian Jutten, Salem Said, and Yannick Berthoumieu. Transfer learning: a riemannian geometry framework with applications to brain–computer interfaces. IEEE Transactions on Biomedical Engineering, 65(5):1107–1116, 2017.

[3] (1,2)

Shanglin Li, Motoaki Kawanabe, and Reinmar J Kobler. SPDIM: source-free unsupervised conditional and label shift adaptation in EEG. In The Thirteenth International Conference on Learning Representations. 2025. URL: https://openreview.net/forum?id=CoQw1dXtGb.

# Cleanup
plt.close("all")

Total running time of the script: (0 minutes 22.911 seconds)