Skip to content

Commit b5f1a2b

Browse files
authored
Add EEGPT backbone (#40)
* add EEGPT * Changelog
1 parent e99fb3c commit b5f1a2b

3 files changed

Lines changed: 37 additions & 0 deletions

File tree

CHANGELOG.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
77

88
## [Unreleased]
99

10+
### Added
11+
- Add the EEGPT backbone ([#40](https://github.com/braindecode/OpenEEGBench/pull/40)).
1012

1113
## [0.5.0] - 2026-05-28
1214

open_eeg_bench/backbone_utils.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,20 @@
1+
import numpy as np
2+
from braindecode.models import EEGPT, InterpolatedModel
3+
from braindecode.models.eegpt import EEGPT_CHANNELS
4+
from mne.channels import make_standard_montage
5+
6+
montage = make_standard_montage("standard_1020")
7+
ch_pos = {
8+
ch.upper(): (ch, loc) for ch, loc in montage.get_positions()["ch_pos"].items()
9+
}
10+
_EEGPT_TARGET_CHS_INFO = [
11+
{
12+
"ch_name": ch_pos[ch.upper()][0],
13+
"kind": "eeg",
14+
"loc": ch_pos[ch.upper()][1],
15+
}
16+
for ch in EEGPT_CHANNELS
17+
]
18+
InterpolatedEEGPT = InterpolatedModel(
19+
EEGPT, _EEGPT_TARGET_CHS_INFO, name="InterpolatedEEGPT"
20+
)

open_eeg_bench/default_configs/backbones.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -84,11 +84,26 @@ def reve(**overrides) -> PretrainedBackbone:
8484
return PretrainedBackbone(**defaults)
8585

8686

87+
def eegpt(**overrides) -> PretrainedBackbone:
88+
defaults = dict(
89+
# model_cls="braindecode.models.EEGPT",
90+
# model_cls="braindecode.models.InterpolatedEEGPT",
91+
model_cls="open_eeg_bench.backbone_utils.InterpolatedEEGPT",
92+
model_kwargs={"chan_proj_type": "none", "n_chans_target": 19},
93+
peft_ff_modules=["qkv", "fc1", "fc2"],
94+
normalization=WindowZScore(),
95+
hub_repo="braindecode/eegpt-pretrained",
96+
)
97+
defaults.update(overrides)
98+
return PretrainedBackbone(**defaults)
99+
100+
87101
ALL_BACKBONES = {
88102
"biot": biot,
89103
"labram": labram,
90104
"bendr": bendr,
91105
"cbramod": cbramod,
92106
"signal_jepa": signal_jepa,
93107
"reve": reve,
108+
"eegpt": eegpt,
94109
}

0 commit comments

Comments
 (0)