braindecode.models.Guetschel2026#

class braindecode.models.Guetschel2026(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, embed_dim=512, depth=4, num_heads=8, dim_feedforward=1365, patch_size=200, patch_overlap=20, pos_half_range=0.15, activation=<class 'torch.nn.modules.activation.GELU'>, drop_prob=0.0, normalization='median_std_clip', input_scale=1000000.0, clip_sigma=15.0, random_projection=None, random_projection_seed=0, channel_strategy='native', channel_strategy_kwargs=None)[source]#

Encoder of the EEG masking-geometry study from Guetschel et al. (2026) [guetschel2026].

Official website of the study: https://pierregtch.github.io/eeg-fm-masking/

Foundation Model Attention/Transformer Channel

Added in version 1.8.2.

Figure 1 of Guetschel et al. (2026): shared MAE/JEPA pre-training pipeline (A) and the block-masking geometries (B)

Figure 1 of [guetschel2026]. (A) MAE and JEPA share one pre-training pipeline; this class is its tokeniser and encoder. (B) The masks are blocks of spatial radius \(r\) and temporal length \(L\).#

Which mask should an EEG foundation model learn from? The study answers with a controlled sweep: one backbone, pre-trained 58 times on the same data with the same recipe, changing only the pretext (MAE or JEPA) and the geometry of the mask. This class is that backbone: all 58 checkpoints load into it.

Both pretexts agree on the best mask, blocks of radius 9 cm and length 2 patches. With it, the frozen features reach the level of REVE-Base under a linear probe on the 12 datasets of OpenEEGBench, with 12.7 M parameters and a fraction of REVE’s pre-training compute.

Architecture

The backbone follows REVE-Small, with a simpler positional encoding:

  • feature_encoder Patch embedding. Each channel is cut into overlapping 1 s patches (200 samples, 20 of overlap), and one linear layer embeds each patch into 512 features.

  • model.pos_encoder Positional encoding. Fixed sinusoids of the electrode position \((x, y, z)\) and of the patch index are added to the embeddings. Any montage works, as long as the channels have positions.

  • model.transformer Transformer encoder. Four pre-norm layers (RMSNorm, 8 heads, GEGLU feed-forward) attend over all the channel × patch tokens.

  • final_layer Head. Flatten, then a linear layer, with an optional fixed random projection in between (see below). It is not part of the checkpoints.

The 58 checkpoints

Each checkpoint is a Hugging Face repository, named after its three parameters (hub_repo_id() builds the name):

PierreGtch/eeg-fm-masking_{pretext}_r{radius}_L{length}

Pretext

mae or jepa

Radius \(r\)

one (a single channel), 6cm, 9cm, 12cm, all (every channel)

Length \(L\)

1, 2, 4, 8, 16 or 33 patches (33 is the whole 30 s window)

Every combination exists except \(r\) = all with \(L\) = 33, which would mask the whole window: 2 × 29 = 58 checkpoints.

  • Recommended: mae_r9cm_L2 or jepa_r9cm_L2. Many other masks are nearly as good.

  • To avoid for downstream use: masks that are too local (\(r\) = one), too global (\(r\) = all) or, in most cases, too long (\(L\) = 8, 16). JEPA collapses at \(r\) = all. These checkpoints are released to study how the mask shapes the representations.

  • Intermediate epochs: model.safetensors is the end of epoch 10, the one evaluated in the paper; the folders epoch_01/ to epoch_09/ hold the earlier epochs of the same run.

License (MIT, the code). The weights are released under CC-BY-4.0 (collection, project page).

Usage

from braindecode.models import Guetschel2026

raw.set_montage("standard_1020")  # channel positions, in metres
model = Guetschel2026.from_pretrained(
    Guetschel2026.hub_repo_id("mae", "9cm", 2),
    chs_info=raw.info["chs"],
    n_times=1000,  # 5 s at 200 Hz
    n_outputs=4,
    # filename="epoch_05/model.safetensors",  # an intermediate epoch
)

# features: (batch, n_chans, n_patches, 512)
features = model(x, return_features=True)["features"]

# linear probing: train only the head, which starts from random weights
for name, p in model.named_parameters():
    p.requires_grad = name.startswith("final_layer.")

Random projection head

The paper probes the frozen features as OpenEEGBench does: it projects the flattened features to 5000 dimensions with a Gaussian random projection, then fits a linear model on top. Pass random_projection=5000 to put the same projection between the flatten and the linear layer of the head:

  • The projection is drawn once, from random_projection_seed, with the same distribution as scikit-learn’s GaussianRandomProjection. It is stored as a buffer: saved with the model, never trained.

  • It is large: random_projection × n_chans × n_patches × 512 values, about 0.9 GB in float32 for 5000 components, 22 channels and 4 s windows.

Warning

Input requirements

  • Sampling rate: 200 Hz, as in pre-training.

  • Units: volts, without standardisation. The model scales each window itself, as in pre-training (microvolts, then division by the median channel standard deviation, clipped at 15).

  • Channel positions: in metres, in the MNE head frame, as given by raw.set_montage(...). Every channel must be 5 to 20 cm from the origin, which catches positions in centimetres or millimetres. For channels with standard names but no positions, pass channel_strategy="exact".

  • Window: at least 200 samples; trailing samples that do not fill a patch are dropped.

Note

Differences from the reference implementation. The backbone gives bit-identical features. Around it:

  • the head takes the actual number of overlapping patches, (n_times - 200) // 180 + 1. The original wrapper sizes its head for n_times // 200 patches, which ignores the overlap, so that head only fits short windows (not 2000 or 6000 samples, for instance); OpenEEGBench replaces it, so the paper’s results do not depend on it;

  • the random projection head is new;

  • channel positions are validated;

  • attention dropout is off in eval mode.

Parameters:
  • n_outputs (int) – Number of outputs of the model. This is the number of classes in the case of classification.

  • n_chans (int) – Number of EEG channels.

  • chs_info (list of dict) – Information about each individual EEG channel. This should be filled with info["chs"]. Refer to mne.Info for more details.

  • n_times (int) – Number of time samples of the input window.

  • input_window_seconds (float) – Length of the input window in seconds.

  • sfreq (float) – Sampling frequency of the EEG recordings.

  • embed_dim (int) – Width of the tokens. Must be divisible by 8 and by num_heads.

  • depth (int) – Number of transformer layers.

  • num_heads (int) – Number of attention heads.

  • dim_feedforward (int) – Width of each half of the GEGLU feed-forward block.

  • patch_size (int) – Number of samples of a patch (1 s at 200 Hz).

  • patch_overlap (int) – Number of samples shared by two consecutive patches.

  • pos_half_range (float) – Half range, in metres, of the electrode coordinates mapped to \([0, 1]\) before the sinusoidal encoding.

  • activation (type[Module]) – Activation of the gate of the feed-forward block.

  • drop_prob (float) – Dropout probability of the attention and feed-forward branches.

  • normalization (str) – Input scaling. "median_std_clip" divides each window by the median over channels of the channel standard deviations and clips it at clip_sigma; "none" only multiplies by input_scale. Use "none" with input_scale=1.0 for data you already scaled.

  • input_scale (float) – Factor applied to the input first (volts to microvolts).

  • clip_sigma (float) – Clipping bound of "median_std_clip".

  • random_projection (int | None) – Size of the fixed random projection inserted in the head (5000 in the paper), or None for no projection. Memory grows linearly with it (see “Random projection head” above).

  • random_projection_seed (int) – Seed of the random projection.

  • channel_strategy (str) – How any montage reaches the backbone (pretrained models only; see Channel strategies: any montage in). "native" keeps the model as it is. "exact", "zero", "nearest", "idw", "spline", "field", "source", "region", "wiener" (call model.channel_layer.fit first) or "latent" map the montage of chs_info (or of the chs_info given to forward) onto the backbone’s channels with a ChannelLayer. Saved in the config. model(x) and model.forward(x) both apply the layer.

  • channel_strategy_kwargs (dict | None) – Options of the strategy (e.g. {"reg": 1e-2} for "spline").

Raises:

ValueError – If the channel positions are missing, not in metres (every channel must be 5 to 20 cm from the origin) or not in the MNE head frame, if an architecture argument is invalid, or if the signal-related parameters are missing and cannot be inferred.

Notes

If some input signal-related parameters are not specified, there will be an attempt to infer them from the other parameters.

References

[guetschel2026] (1,2)

Guetschel, P., Aristimunha, B., El Ouahidi, Y., Delorme, A., Moreau, T., & Tangermann, M. (2026). What masking geometry works best for EEG foundation models? arXiv:2609.33487. https://arxiv.org/abs/2609.33487 Code: https://github.com/PierreGtch/eeg-fm-masking. Project page: https://pierregtch.github.io/eeg-fm-masking/.

Hugging Face Hub integration

When the optional huggingface_hub package is installed, all models automatically gain the ability to be pushed to and loaded from the Hugging Face Hub. Install with:

pip install braindecode[hub]

Pushing a model to the Hub:

from braindecode.models import Guetschel2026

# Train your model
model = Guetschel2026(n_chans=22, n_outputs=4, n_times=1000)
# ... training code ...

# Push to the Hub
model.push_to_hub(
    repo_id="username/my-guetschel2026-model",
    commit_message="Initial model upload",
)

Loading a model from the Hub:

from braindecode.models import Guetschel2026

# Load pretrained model
model = Guetschel2026.from_pretrained("username/my-guetschel2026-model")

# Load with a different number of outputs (head is rebuilt automatically)
model = Guetschel2026.from_pretrained("username/my-guetschel2026-model", n_outputs=4)

Extracting features and replacing the head:

import torch

x = torch.randn(1, model.n_chans, model.n_times)
# Extract encoder features (consistent dict across all models)
out = model(x, return_features=True)
features = out["features"]

# Replace the classification head
model.reset_head(n_outputs=10)

Saving and restoring full configuration:

import json

config = model.get_config()            # all __init__ params
with open("config.json", "w") as f:
    json.dump(config, f)

model2 = Guetschel2026.from_config(config)    # reconstruct (no weights)

All model parameters (both EEG-specific and model-specific such as dropout rates, activation functions, number of filters) are automatically saved to the Hub and restored when loading.

See Loading and Adapting Pretrained Foundation Models for a complete tutorial.

Methods

forward(x, return_features=False)[source]#

Forward pass.

Parameters:
  • x (Tensor) – Input of shape (batch, n_chans, n_times), in volts.

  • return_features (bool) – If True, return the features instead of the logits.

Returns:

The logits, of shape (batch, n_outputs), or, with return_features=True, {"features": z, "cls_token": None} with z of shape (batch, n_chans, n_patches, embed_dim).

Return type:

Union[Tensor, Dict[str, Optional[Tensor]]]

classmethod hub_repo_id(pretext, mask_radius, mask_length)[source]#

Hub repository of one of the 58 released checkpoints.

Parameters:
  • pretext (str) – Pre-training objective.

  • mask_radius (str) – Radius of the masked blocks of channels.

  • mask_length (int) – Length of the masked blocks in patches. ("all", 33) was never trained, because it would mask the whole window.

Returns:

"PierreGtch/eeg-fm-masking_{pretext}_r{mask_radius}_L{mask_length}".

Return type:

str

Raises:

ValueError – If a value is not in the lists above, or for ("all", 33).

Examples

>>> Guetschel2026.hub_repo_id("mae", "9cm", 2)
'PierreGtch/eeg-fm-masking_mae_r9cm_L2'
load_state_dict(state_dict, *args, **kwargs)[source]#

Load a state dict whose backbone matches this model exactly.

Arguments are passed on to torch.nn.Module.load_state_dict(). Keys under final_layer.* (and channel_layer.*) may be absent or extra, because the pretrained checkpoints carry no head. Any other missing or unexpected key raises an error.

Raises:

RuntimeError – If a backbone key is missing or unexpected, even with strict=False (from_pretrained() loads with strict=False).

Parameters:
  • state_dict – The description is missing.

  • *args – The description is missing.

  • **kwargs – The description is missing.

reset_head(n_outputs)[source]#

Replace the last linear layer, keeping the random projection.

Parameters:

n_outputs (int) – New number of outputs.

Return type:

None

Examples

>>> from braindecode.models import BENDR
>>> model = BENDR(n_chans=22, n_times=1000, n_outputs=4)
>>> model.reset_head(10)
>>> model.n_outputs
10

Added in version 1.4.