EEG Classification with EEGSPDNet#

This tutorial demonstrates how to use EEGSPDNet for motor imagery EEG classification. EEGSPDNet combines channel-specific convolution with SPD matrix learning for robust EEG decoding.

Introduction#

EEGSPDNet [Wilson et al., 2025] is a deep Riemannian network designed specifically for EEG decoding. It extends the SPDNet architecture with:

  • Channel-specific convolution: Learns temporal filters independently for each EEG channel using grouped convolutions

  • Covariance pooling: Computes SPD covariance matrices from the filtered signals

  • Scalable BiMap layers: Progressively reduces dimensionality on the SPD manifold

  • SPD Dropout: Structured dropout that maintains positive definiteness

Setup and Imports#

import warnings

import matplotlib.pyplot as plt
import torch

from braindecode import EEGClassifier
from moabb.datasets import BNCI2014_001
from moabb.paradigms import MotorImagery
from sklearn.metrics import (
    ConfusionMatrixDisplay,
    accuracy_score,
    balanced_accuracy_score,
    confusion_matrix,
)
from sklearn.preprocessing import LabelEncoder
from skorch.callbacks import EpochScoring, GradientNormClipping
from skorch.dataset import ValidSplit

from spd_learn.models import EEGSPDNet


warnings.filterwarnings("ignore")

Loading the Dataset#

We use the BCI Competition IV Dataset 2a for motor imagery classification.

dataset = BNCI2014_001()
paradigm = MotorImagery(n_classes=4)

print(f"Dataset: {dataset.code}")
print(f"Number of subjects: {len(dataset.subject_list)}")
Choosing from all possible events
Dataset: BNCI2014-001
Number of subjects: 9

Understanding EEGSPDNet Architecture#

EEGSPDNet processes EEG in the following stages:

  1. Channel-specific Conv1d: Each channel gets its own set of filters

    \[X_{conv} = \text{GroupedConv1d}(X), \quad X \in \mathbb{R}^{C \times T}\]

    Output shape: (n_chans * n_filters, time - filter_length + 1)

  2. Covariance Pooling: Compute SPD covariance matrix

    \[\Sigma = \frac{1}{T-1} X_{conv} X_{conv}^T\]
  3. BiMap + ReEig blocks: Learn spatial filters while preserving SPD

    \[Y = W^T \Sigma W, \quad Y = U \max(\Lambda, \epsilon) U^T\]
  4. LogEig: Project to tangent space for classification

Creating the EEGSPDNet Model#

Key parameters:

  • n_filters: Number of temporal filters per channel

  • bimap_sizes: Tuple (k, n_layers) defining scaling factor and depth

  • filter_time_length: Length of temporal convolution kernel

  • spd_drop_prob: Dropout probability for SPD dropout layers

n_chans = 22
n_outputs = 4

# Create EEGSPDNet model
model = EEGSPDNet(
    n_chans=n_chans,
    n_outputs=n_outputs,
    n_filters=4,  # 4 filters per channel → 88 total
    bimap_sizes=(2, 2),  # Scale by 2x, 2 BiMap layers: 88→44→22
    filter_time_length=25,  # 100ms filter at 250Hz
    spd_drop_prob=0.0,  # No SPD dropout (can cause instability)
    spd_drop_scaling=True,  # Scale remaining channels
    final_layer_drop_prob=0.5,  # 50% dropout before classifier
)

print("EEGSPDNet Architecture:")
print(model)

# Show BiMap layer dimensions
print("\nBiMap Layer Dimensions:")
print("Input: 88 x 88 (22 channels x 4 filters)")
print("→ BiMap0: 88 → 44")
print("→ BiMap1: 44 → 22")
print("→ LogEig: 22 x 22 → 253 (upper triangular)")
EEGSPDNet Architecture:
EEGSPDNet(
  (conv): Conv1d(22, 88, kernel_size=(25,), stride=(1,), groups=22)
  (cov_pool): CovLayer()
  (spdnet): Sequential(
    (bimap0): ParametrizedBiMap(
      (parametrizations): ModuleDict(
        (weight): ParametrizationList(
          (0): _Orthogonal()
        )
      )
    )
    (reeig0): ReEig()
    (bimap1): ParametrizedBiMap(
      (parametrizations): ModuleDict(
        (weight): ParametrizationList(
          (0): _Orthogonal()
        )
      )
    )
    (reeig1): ReEig()
    (logeig): LogEig()
  )
  (dropout): Dropout(p=0.5, inplace=False)
  (linear): Linear(in_features=253, out_features=4, bias=True)
)

BiMap Layer Dimensions:
Input: 88 x 88 (22 channels x 4 filters)
→ BiMap0: 88 → 44
→ BiMap1: 44 → 22
→ LogEig: 22 x 22 → 253 (upper triangular)

Setting up the Classifier#

batch_size = 32
max_epochs = 100
learning_rate = 1e-4  # Low learning rate for stable SPD learning

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

# Note: SPD networks benefit from gradient clipping to prevent
# divergence during training on the Riemannian manifold.
clf = EEGClassifier(
    model,
    criterion=torch.nn.CrossEntropyLoss,
    optimizer=torch.optim.Adam,
    optimizer__lr=learning_rate,
    train_split=ValidSplit(0.1, stratified=True, random_state=42),
    batch_size=batch_size,
    max_epochs=max_epochs,
    callbacks=[
        (
            "train_acc",
            EpochScoring(
                "accuracy", lower_is_better=False, on_train=True, name="train_acc"
            ),
        ),
        ("gradient_clip", GradientNormClipping(gradient_clip_value=1.0)),
    ],
    device=device,
    verbose=1,
)
Using device: cpu

Training and Evaluation#

subject_id = 1

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

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

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

print(f"\nData shape: {X.shape}")
print(f"Classes: {le.classes_}")

# Split by session
train_idx = meta.query("session == '0train'").index.to_numpy()
test_idx = meta.query("session == '1test'").index.to_numpy()

print(f"Training samples: {len(train_idx)}")
print(f"Test samples: {len(test_idx)}")

# Train
clf.fit(X[train_idx], y[train_idx])

# Evaluate
y_pred_train = clf.predict(X[train_idx])
y_pred_test = clf.predict(X[test_idx])

train_acc = accuracy_score(y[train_idx], y_pred_train)
test_acc = accuracy_score(y[test_idx], y_pred_test)
test_bal_acc = balanced_accuracy_score(y[test_idx], y_pred_test)

print(f"\n{'=' * 50}")
print(f"Results for Subject {subject_id}")
print(f"{'=' * 50}")
print(f"Train Accuracy:    {train_acc * 100:.2f}%")
print(f"Test Accuracy:     {test_acc * 100:.2f}%")
print(f"Test Balanced Acc: {test_bal_acc * 100:.2f}%")
Data shape: (576, 22, 1001)
Classes: ['feet' 'left_hand' 'right_hand' 'tongue']
Training samples: 288
Test samples: 288
  epoch    train_acc    train_loss    valid_acc    valid_loss     dur
-------  -----------  ------------  -----------  ------------  ------
      1       0.2227        1.4459       0.2414        1.3793  1.0941
      2       0.2500        1.4377       0.2759        1.3547  1.0252
      3       0.3047        1.3800       0.4138        1.3388  0.8876
      4       0.3359        1.3674       0.5862        1.3301  0.7725
      5       0.2930        1.3854       0.6552        1.3221  0.7791
      6       0.2617        1.3963       0.6552        1.3157  0.7636
      7       0.3125        1.3727       0.6552        1.3099  0.7662
      8       0.2930        1.3929       0.6552        1.3018  0.7558
      9       0.3047        1.3627       0.6552        1.2952  1.0724
     10       0.3086        1.3871       0.6552        1.2898  1.0607
     11       0.3008        1.3444       0.6552        1.2842  1.0824
     12       0.3516        1.3377       0.6897        1.2787  1.0516
     13       0.3086        1.3546       0.7241        1.2739  1.0541
     14       0.3008        1.3553       0.6897        1.2675  1.0759
     15       0.3125        1.3493       0.6897        1.2621  1.0668
     16       0.3281        1.3464       0.6897        1.2563  1.0515
     17       0.3789        1.3128       0.6897        1.2506  1.0630
     18       0.3867        1.2940       0.7241        1.2437  1.0735
     19       0.3438        1.3263       0.7241        1.2374  1.0609
     20       0.3789        1.3015       0.7241        1.2325  1.0606
     21       0.3984        1.3045       0.7586        1.2274  1.0649
     22       0.3438        1.3054       0.7586        1.2241  1.0617
     23       0.3633        1.3046       0.7241        1.2196  1.0632
     24       0.3945        1.2579       0.6897        1.2146  1.0630
     25       0.3672        1.2830       0.6897        1.2105  1.0453
     26       0.3906        1.2820       0.6897        1.2060  1.0353
     27       0.4141        1.2618       0.6897        1.1996  1.0430
     28       0.4297        1.2449       0.7586        1.1935  1.0398
     29       0.3672        1.2692       0.7586        1.1893  1.0408
     30       0.4258        1.2457       0.7586        1.1837  1.0558
     31       0.4531        1.2541       0.7586        1.1785  1.0282
     32       0.3984        1.2537       0.7931        1.1729  1.0335
     33       0.5352        1.1902       0.7241        1.1693  1.0275
     34       0.4570        1.2343       0.7931        1.1653  1.0363
     35       0.4141        1.2447       0.7931        1.1615  1.0215
     36       0.4727        1.2121       0.7931        1.1583  1.0337
     37       0.4844        1.2003       0.7931        1.1529  1.0316
     38       0.3984        1.2471       0.7931        1.1492  1.0293
     39       0.4922        1.2026       0.6897        1.1448  1.0345
     40       0.4570        1.2071       0.6897        1.1410  1.0335
     41       0.4883        1.1866       0.7931        1.1356  1.0346
     42       0.4648        1.2063       0.7241        1.1318  1.0311
     43       0.5156        1.1756       0.7931        1.1261  1.0619
     44       0.4883        1.1878       0.8276        1.1216  1.0412
     45       0.5469        1.1601       0.6897        1.1178  1.0345
     46       0.4766        1.2130       0.6897        1.1130  1.0375
     47       0.5117        1.1743       0.6897        1.1089  1.0414
     48       0.4922        1.1638       0.6897        1.1060  1.0401
     49       0.4766        1.1739       0.6552        1.1024  1.0457
     50       0.5430        1.1444       0.6207        1.0984  1.0326
     51       0.5117        1.1596       0.6552        1.0950  1.0338
     52       0.5352        1.1331       0.7586        1.0913  1.0319
     53       0.5000        1.1655       0.7241        1.0890  1.0275
     54       0.5117        1.1360       0.8276        1.0841  1.0326
     55       0.5000        1.1372       0.7931        1.0804  1.0313
     56       0.5000        1.1345       0.7586        1.0756  1.0422
     57       0.5117        1.1559       0.7241        1.0724  1.0237
     58       0.5508        1.1182       0.7586        1.0692  1.0272
     59       0.4844        1.1616       0.7586        1.0648  0.9782
     60       0.5352        1.1178       0.8621        1.0592  0.7638
     61       0.5234        1.1049       0.7241        1.0565  0.7401
     62       0.5195        1.1071       0.7586        1.0541  0.7836
     63       0.5938        1.0924       0.8276        1.0512  0.7597
     64       0.5938        1.1068       0.7931        1.0482  0.7475
     65       0.5547        1.0978       0.8276        1.0454  0.7476
     66       0.5859        1.0959       0.8276        1.0413  0.7429
     67       0.5898        1.0951       0.7931        1.0369  0.7276
     68       0.5625        1.0775       0.7931        1.0343  0.7389
     69       0.5547        1.1011       0.7241        1.0325  0.7632
     70       0.5664        1.0808       0.7241        1.0302  1.0050
     71       0.5234        1.0953       0.7586        1.0267  1.0675
     72       0.5781        1.0589       0.7931        1.0199  1.0460
     73       0.6016        1.0495       0.7586        1.0155  1.0615
     74       0.5820        1.0646       0.7586        1.0140  1.0691
     75       0.5508        1.0694       0.7241        1.0122  1.0648
     76       0.6367        1.0341       0.7241        1.0104  1.0604
     77       0.5586        1.0756       0.7241        1.0077  1.0560
     78       0.5625        1.0687       0.7241        1.0034  1.0624
     79       0.5781        1.0502       0.7241        0.9987  1.0633
     80       0.6055        1.0453       0.7241        0.9954  1.0656
     81       0.5820        1.0453       0.7241        0.9920  1.0528
     82       0.5703        1.0501       0.7241        0.9904  1.0574
     83       0.5938        1.0401       0.7241        0.9863  1.0754
     84       0.6562        1.0187       0.7241        0.9840  1.0722
     85       0.5820        1.0406       0.7586        0.9818  1.0620
     86       0.6055        1.0286       0.7931        0.9786  1.0600
     87       0.6328        1.0070       0.7586        0.9776  1.0625
     88       0.6367        1.0095       0.6897        0.9726  1.0684
     89       0.6406        1.0191       0.6897        0.9688  1.0684
     90       0.6055        1.0321       0.7241        0.9681  1.0675
     91       0.5938        1.0354       0.7241        0.9673  1.0683
     92       0.6250        0.9906       0.6897        0.9649  1.0561
     93       0.6172        1.0031       0.6897        0.9613  1.0688
     94       0.6680        0.9940       0.6897        0.9594  1.0817
     95       0.6250        0.9800       0.7241        0.9550  1.0663
     96       0.6445        1.0020       0.7241        0.9540  1.0441
     97       0.5859        1.0044       0.7241        0.9518  1.0638
     98       0.6523        0.9755       0.6897        0.9487  1.0754
     99       0.6406        0.9764       0.6552        0.9459  1.0711
    100       0.6562        0.9723       0.6552        0.9437  1.0572

==================================================
Results for Subject 1
==================================================
Train Accuracy:    73.96%
Test Accuracy:     67.36%
Test Balanced Acc: 67.36%

Visualizing Results#

fig, axes = plt.subplots(1, 3, figsize=(15, 4))

history = clf.history
epochs = range(1, len(history) + 1)

# Loss
ax1 = axes[0]
ax1.plot(epochs, history[:, "train_loss"], "b-", label="Train Loss", linewidth=2)
ax1.plot(epochs, history[:, "valid_loss"], "r--", label="Valid Loss", linewidth=2)
ax1.set_xlabel("Epoch", fontsize=12)
ax1.set_ylabel("Loss", fontsize=12)
ax1.set_title("Training and Validation Loss", fontsize=14)
ax1.legend(fontsize=10)
ax1.grid(True, alpha=0.3)

# Accuracy
ax2 = axes[1]
ax2.plot(epochs, history[:, "train_acc"], "b-", label="Train Acc", linewidth=2)
ax2.plot(epochs, history[:, "valid_acc"], "r--", label="Valid Acc", linewidth=2)
ax2.set_xlabel("Epoch", fontsize=12)
ax2.set_ylabel("Accuracy", fontsize=12)
ax2.set_title("Training and Validation Accuracy", fontsize=14)
ax2.legend(fontsize=10)
ax2.grid(True, alpha=0.3)
ax2.set_ylim([0, 1])

# Confusion Matrix
ax3 = axes[2]
cm = confusion_matrix(y[test_idx], y_pred_test)
disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=le.classes_)
disp.plot(ax=ax3, cmap="Blues", values_format="d")
ax3.set_title(f"Test Confusion Matrix\nAccuracy: {test_acc * 100:.1f}%", fontsize=14)

plt.tight_layout()
plt.show()
Training and Validation Loss, Training and Validation Accuracy, Test Confusion Matrix Accuracy: 67.4%

Comparing with Different Configurations#

Let’s compare different EEGSPDNet configurations to understand the impact of hyperparameters.

configs = {
    "Small (k=2, depth=1)": {"bimap_sizes": (2, 1), "n_filters": 4},
    "Medium (k=2, depth=2)": {"bimap_sizes": (2, 2), "n_filters": 4},
    "Large (k=2, depth=3)": {"bimap_sizes": (2, 3), "n_filters": 6},
}

print("\nModel Size Comparison:")
print("-" * 60)
for name, config in configs.items():
    temp_model = EEGSPDNet(n_chans=n_chans, n_outputs=n_outputs, **config)
    n_params = sum(p.numel() for p in temp_model.parameters())
    print(f"{name}: {n_params:,} parameters")
Model Size Comparison:
------------------------------------------------------------
Small (k=2, depth=1): 10,124 parameters
Medium (k=2, depth=2): 8,144 parameters
Large (k=2, depth=3): 15,398 parameters

Summary#

In this tutorial, we demonstrated how to:

  1. Create and configure an EEGSPDNet model

  2. Understand the architecture’s channel-specific processing

  3. Train and evaluate on motor imagery data

  4. Compare different model configurations

EEGSPDNet’s channel-specific convolution allows it to learn independent temporal features for each electrode, which is particularly useful for spatially distributed brain signals.

Total running time of the script: (1 minutes 42.198 seconds)