PierreGtch commited on
Commit
101fdd1
·
verified ·
1 Parent(s): eac267f

Upload encoder weights, config and model card

Browse files
Files changed (4) hide show
  1. README.md +125 -0
  2. config.json +39 -0
  3. metadata.json +23 -0
  4. model.safetensors +3 -0
README.md ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: cc-by-4.0
3
+ pipeline_tag: feature-extraction
4
+ tags:
5
+ - eeg
6
+ - self-supervised-learning
7
+ - foundation-model
8
+ - pytorch
9
+ - masked-autoencoder
10
+ - open-eeg-bench
11
+ ---
12
+
13
+ # eeg-fm-masking_mae_r9cm_L1
14
+
15
+ Pretrained EEG encoder from the paper *What masking geometry works best for EEG foundation models? A controlled evaluation across MAE and JEPA*
16
+ ([website](https://pierregtch.github.io/eeg-fm-masking) · [code](https://github.com/PierreGtch/eeg-fm-masking) · [all 58 models](https://huggingface.co/collections/PierreGtch/eeg-fm-masking-6ab912b6a03bba1348fc7366)).
17
+
18
+ It is one of **58 encoders trained under an identical recipe** where only the masking
19
+ geometry changes: 5 spatial radii × 6 temporal lengths × 2 frameworks
20
+ (the r = all, L = 33 cell, which would mask the whole window, does not exist).
21
+ This model is a **masked autoencoder (MAE)**: the encoder sees the unmasked patches and a light decoder reconstructs the raw signal of the masked patches.
22
+
23
+ | | |
24
+ |---|---|
25
+ | Framework | MAE |
26
+ | Mask spatial radius `r` | 9 cm |
27
+ | Mask temporal length `L` | 1 patch |
28
+ | Masker parameter `pct_unmasked` | 0.45 |
29
+ | Checkpoint | epoch 10 of 10 (`v9`, the one evaluated in the paper) |
30
+ | Encoder parameters | 12.69 M |
31
+ | Training run | [`yncl6get`](https://wandb.ai/pierregtch/chan-inv-clf/runs/yncl6get) |
32
+
33
+ The paper recommends r = 9 cm, L = 2: see [`eeg-fm-masking_mae_r9cm_L2`](https://huggingface.co/PierreGtch/eeg-fm-masking_mae_r9cm_L2) and [`eeg-fm-masking_jepa_r9cm_L2`](https://huggingface.co/PierreGtch/eeg-fm-masking_jepa_r9cm_L2).
34
+
35
+ ## What is in this repo
36
+
37
+ * `model.safetensors`: the **encoder only** (patch tokeniser `feature_encoder.*` + transformer
38
+ `model.*`), i.e. exactly the tensors loaded for the downstream evaluation of the paper.
39
+ The MAE decoder is
40
+ not included; for JEPA the published weights are the student encoder, as evaluated in the paper.
41
+ * `config.json`: the keyword arguments of `ContextualEncoderBenchmarkWrapper` (architecture +
42
+ input scaling). Pass it unchanged as `model_kwargs`.
43
+ * `metadata.json`: masking parameters, training-run id, checkpoint epoch/step.
44
+
45
+ ## Input requirements
46
+
47
+ * **Sampling rate: 200 Hz.** The signal is cut into 1 s patches (200 samples, 20-sample overlap).
48
+ * **Units: volts.** The wrapper multiplies by `factor = 1e+06` and applies
49
+ per-window `median_std_clip` scaling (clip at σ = 15) itself;
50
+ do not standardise the data beforehand.
51
+ * **Channel positions in metres** (MNE `info["chs"][i]["loc"][:3]`). The model is montage-agnostic:
52
+ any number and set of channels works, as long as every channel has a 3D position.
53
+
54
+ ## Usage
55
+
56
+ Install the code: `pip install git+https://github.com/PierreGtch/eeg-fm-masking`.
57
+
58
+ ### Loading the model
59
+
60
+ ```python
61
+ import json
62
+ from huggingface_hub import hf_hub_download
63
+ from safetensors.torch import load_file
64
+ from eeg_fm_masking.oeb.wrapper import ContextualEncoderBenchmarkWrapper
65
+
66
+ repo = "PierreGtch/eeg-fm-masking_mae_r9cm_L1"
67
+ config = json.load(open(hf_hub_download(repo, "config.json")))
68
+ model = ContextualEncoderBenchmarkWrapper(
69
+ n_chans=n_chans, n_times=n_times, n_outputs=n_outputs, sfreq=200.0,
70
+ chs_info=chs_info, # MNE channel info (info["chs"]), positions in metres
71
+ **config,
72
+ )
73
+ model.load_state_dict(load_file(hf_hub_download(repo, "model.safetensors")), strict=False)
74
+ ```
75
+
76
+ `strict=False` only leaves out the dataset-dependent parts (channel-position buffer and
77
+ classification head).
78
+
79
+ ### Evaluation / fine-tuning with OpenEEGBench
80
+
81
+ ```python
82
+ import json
83
+ from huggingface_hub import hf_hub_download
84
+ from open_eeg_bench.backbone import PretrainedBackbone
85
+
86
+ repo = "PierreGtch/eeg-fm-masking_mae_r9cm_L1"
87
+ backbone = PretrainedBackbone(
88
+ model_cls="eeg_fm_masking.oeb.wrapper.ContextualEncoderBenchmarkWrapper",
89
+ hub_repo=repo,
90
+ model_kwargs=json.load(open(hf_hub_download(repo, "config.json"))),
91
+ )
92
+ ```
93
+
94
+ ## Downstream results (OpenEEGBench, frozen encoder + ridge probe)
95
+
96
+ Frozen encoder, ridge regression/classification on the flattened contextual features,
97
+ 12 datasets × 5 seeds. Balanced accuracy for classification, R² for `seed-vig`.
98
+
99
+ | Dataset | Metric | Score (mean ± sd) | Seeds |
100
+ |---|---|---|---|
101
+ | arithmetic_zyma2019 | balanced acc. | 0.699 ± 0.011 | 5 |
102
+ | bcic2020-3 | balanced acc. | 0.267 ± 0.029 | 5 |
103
+ | bcic2a | balanced acc. | 0.453 ± 0.005 | 5 |
104
+ | chbmit | balanced acc. | 0.876 ± 0.050 | 5 |
105
+ | faced | balanced acc. | 0.312 ± 0.007 | 5 |
106
+ | isruc-sleep | balanced acc. | 0.662 ± 0.004 | 5 |
107
+ | mdd_mumtaz2016 | balanced acc. | 0.816 ± 0.012 | 5 |
108
+ | physionet | balanced acc. | 0.579 ± 0.004 | 5 |
109
+ | seed-v | balanced acc. | 0.283 ± 0.001 | 5 |
110
+ | seed-vig | R² | -0.163 ± 0.093 | 5 |
111
+ | tuab | balanced acc. | 0.805 ± 0.003 | 5 |
112
+ | tuev | balanced acc. | 0.906 ± 0.026 | 5 |
113
+
114
+ ## Training
115
+
116
+ * **Data:** the openly-licensed subset of the REVE pre-training corpus (323 recordings),
117
+ so that the weights can be redistributed.
118
+ * **Schedule:** 10 epochs, 2 × H100, batch size 600 per GPU,
119
+ learning rate 0.00024 (warm-up 3080 steps, final 1e-06),
120
+ weight decay 0.01.
121
+
122
+ ## License and citation
123
+
124
+ Weights released under **CC-BY-4.0**; code under MIT ([GitHub](https://github.com/PierreGtch/eeg-fm-masking)).
125
+ If you use these models, please cite the paper (reference on the [GitHub page](https://github.com/PierreGtch/eeg-fm-masking)).
config.json ADDED
@@ -0,0 +1,39 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "clip_sigma": 15.0,
3
+ "factor": 1000000.0,
4
+ "feature_encoder": {
5
+ "dim": 512,
6
+ "modelName": "LinearPatchEmbedding",
7
+ "patch_overlap": 20,
8
+ "patch_size": 200
9
+ },
10
+ "masker": {
11
+ "length_blocks": 1,
12
+ "n_target_blocks": null,
13
+ "pct_unmasked": 0.44999999999999996,
14
+ "radius_blocks": 0.09,
15
+ "scalp_surface": 0.0942477796076938,
16
+ "vectorized": true
17
+ },
18
+ "pos_encoder": {
19
+ "max_seconds": 600.0,
20
+ "max_x": 0.15,
21
+ "modelName": "AdditivePositionalEncoder",
22
+ "sfreq_features": 1.1111111111111112,
23
+ "spat_dim": 384,
24
+ "time_dim": 128
25
+ },
26
+ "scaler": "median_std_clip",
27
+ "shared_feature_encoder": true,
28
+ "transformer": {
29
+ "activation": "gelu",
30
+ "bias": false,
31
+ "d_model": 512,
32
+ "dim_feedforward": 1365,
33
+ "dropout": 0.0,
34
+ "glu": true,
35
+ "nhead": 8,
36
+ "norm": "rms_norm",
37
+ "num_layers": 4
38
+ }
39
+ }
metadata.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "framework": "mae",
3
+ "mask_radius_m": 0.09,
4
+ "mask_radius_label": "9cm",
5
+ "mask_length_patches": 1,
6
+ "pct_unmasked": 0.45,
7
+ "wandb_run_id": "yncl6get",
8
+ "checkpoint_version": "v9",
9
+ "wandb_run_name": "mask_sweep_full_mae_r009_l01_af69b36",
10
+ "wandb_run_url": "https://wandb.ai/pierregtch/chan-inv-clf/runs/yncl6get",
11
+ "epoch": 9,
12
+ "global_step": 30064,
13
+ "n_params": 12687872,
14
+ "dropped_prefixes": [
15
+ "predictor"
16
+ ],
17
+ "features_shape": [
18
+ 2,
19
+ 19,
20
+ 11,
21
+ 512
22
+ ]
23
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0bc057558486b44ee63407ab7317b533d0511f7df5217560058c0b9778c0d9a1
3
+ size 50771472