braindecode.models.BrainBERT#

class braindecode.models.BrainBERT(hidden_dim=192, ffn_dim=384, n_layers=2, n_heads=4, nperseg=400, noverlap=350, idx_freq_cutoff=40, stft_clip=10, stft_zscore_before_clip=True, pool_n_frames=10, activation=<class 'torch.nn.modules.activation.GELU'>, drop_prob=0.1, n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None)[source]#

BrainBERT from Wang et al. (2023) [BrainBERT2023].

Foundation Model Attention/Transformer

BrainBERT architecture

BrainBERT is a self-supervised foundation model for intracranial signals (sEEG/iEEG). It takes as input not the waveform but its spectrogram: a short-time Fourier transform maps each channel onto a series of time frames, and each frame is one token. A linear projection and a fixed sinusoidal positional encoding are followed by a stack of standard Transformer encoder layers. The model is pre-trained by masked spectrogram modelling, i.e. by reconstructing masked time/frequency patches, and the representation of each frame thus serves downstream decoding.

In the present implementation, we compute the STFT inside forward (the spectrogram submodule), so that the model still takes the standard (batch, n_chans, n_times) input, whereas the upstream reference takes a pre-computed spectrogram as input. The input encoding and the encoder match the upstream MaskedTFModel to 1e-5.

The released checkpoint was trained on signals sampled at 2048 Hz and re-referenced with a Laplacian, with nperseg=400, noverlap=350 and the first idx_freq_cutoff=40 frequency bins. The defaults below give a modest, ready-to-run model; the released (“large”) model uses hidden_dim=768, ffn_dim=3072, n_heads=12 and n_layers=6 (about 43M parameters), and passing these values reproduces it.

The pooling follows the published downstream protocol. Upstream processes one electrode at a time, averages the pool_n_frames=10 encoder outputs centred on the window (outputs[:, middle-5:middle+5].mean in preprocessors/spec_pretrained.py) and applies a linear probe; the average over every frame appears in that file only as a commented-out alternative. We use the average over the central frames and add a mean over channels, which is the identity for n_chans=1: a single-channel BrainBERT thus reproduces the upstream feature exactly, and several channels remain supported as the braindecode generalisation. Pass pool_n_frames=None to average every frame. To obtain one output per electrode as upstream does, build the model with n_chans=1 and stack the electrodes in the batch.

Important

Pre-trained weights are available. The checkpoint released by the authors is available from the Hugging Face Hub:

model = BrainBERT.from_pretrained(
    "braindecode/brainbert-pretrained", n_outputs=2
)

It has the “large” configuration above; n_chans, n_times and n_outputs may change freely, since the frames are pooled and the head is specific to the task (pass these, not chs_info or input_window_seconds, which conflict with the saved config). The upstream repository provides no LICENSE file, so we mark the licence of the weights as unknown rather than assume a permissive one.

Note

As in the upstream front-end, a channel whose z-scored spectrogram is flat (a dead channel) is set to ones, and a NaN sample zeroes that channel’s whole spectrogram for the window, without a warning.

Added in version 1.9.

Parameters:
  • hidden_dim (int) – Transformer model width D. Default 192. The released model uses 768.

  • ffn_dim (int) – Inner dimension of the Transformer feed-forward blocks. Default 384. The released model uses 3072.

  • n_layers (int) – Number of Transformer encoder layers. Default 2. Released model: 6.

  • n_heads (int) – Number of attention heads. Default 4. Released model: 12.

  • nperseg (int) – STFT window length in samples. Default 400.

  • noverlap (int) – STFT overlap in samples. Default 350 (hop of 50).

  • idx_freq_cutoff (int) – Number of low-frequency STFT bins kept; the Transformer input_dim. This is a bin index, not a frequency in Hz: the default 40 bins reach about 200 Hz at 2048 Hz with nperseg=400. Default 40.

  • stft_clip (int) – Boundary frames trimmed from each end of the spectrogram. Default 10, as in the upstream preprocessors/stft.py used with the released checkpoint (the demo notebook uses 5; see stft_zscore_before_clip).

  • stft_zscore_before_clip (bool) – Whether the spectrogram is z-scored before the boundary frames are trimmed. Default True, which together with stft_clip=10 reproduces the recipe behind the published numbers. Set to False with stft_clip=5 to reproduce the upstream demo notebook instead; the two recipes are close but not equal, and they yield sequences of different length.

  • pool_n_frames (int | None) – Number of encoder frames, centred on the window, averaged into the pooled representation. Default 10, as upstream. None averages all frames.

  • activation (type[Module]) – Transformer feed-forward activation class. Default nn.GELU (as pretrained).

  • drop_prob (float) – Dropout probability. Default 0.1.

  • 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.

Raises:

ValueError – If some input signal-related parameters are not specified: and can not 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

[BrainBERT2023]

Wang, C., Subramaniam, V., Yaari, A.U., Kreiman, G., Katz, B., Cases, I. and Barbu, A., 2023. BrainBERT: Self-supervised representation learning for intracranial recordings. In International Conference on Learning Representations, ICLR. Code: czlwang/BrainBERT

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 BrainBERT

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

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

Loading a model from the Hub:

from braindecode.models import BrainBERT

# Load pretrained model
model = BrainBERT.from_pretrained("username/my-brainbert-model")

# Load with a different number of outputs (head is rebuilt automatically)
model = BrainBERT.from_pretrained("username/my-brainbert-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 = BrainBERT.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]#

Decode a batch of signals.

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

  • return_features (bool) – If True, return the pooled encoder embedding instead of the class logits, as {"features": pooled, "cls_token": None} (braindecode foundation-model convention). BrainBERT pools the centre frames and then the channels, and has no class token, hence cls_token is None. A scripted model (torch.jit.script) ignores this flag and always returns the logits.

Returns:

Class logits of shape (batch, n_outputs), or the feature dict {"features", "cls_token"} when return_features is set.

Return type:

torch.Tensor or dict

reset_head(n_outputs)[source]#

Swap the classification head for a new number of outputs.

Parameters:

n_outputs (int) – New number of output classes.

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.