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 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(thespectrogramsubmodule), 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 upstreamMaskedTFModelto 1e-5.The released checkpoint was trained on signals sampled at 2048 Hz and re-referenced with a Laplacian, with
nperseg=400,noverlap=350and the firstidx_freq_cutoff=40frequency bins. The defaults below give a modest, ready-to-run model; the released (“large”) model useshidden_dim=768,ffn_dim=3072,n_heads=12andn_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=10encoder outputs centred on the window (outputs[:, middle-5:middle+5].meaninpreprocessors/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 forn_chans=1: a single-channel BrainBERT thus reproduces the upstream feature exactly, and several channels remain supported as the braindecode generalisation. Passpool_n_frames=Noneto average every frame. To obtain one output per electrode as upstream does, build the model withn_chans=1and 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_timesandn_outputsmay change freely, since the frames are pooled and the head is specific to the task (pass these, notchs_infoorinput_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
NaNsample zeroes that channel’s whole spectrogram for the window, without a warning.Added in version 1.9.
- Parameters:
hidden_dim (
int) – Transformer model widthD. 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 Transformerinput_dim. This is a bin index, not a frequency in Hz: the default 40 bins reach about 200 Hz at 2048 Hz withnperseg=400. Default 40.stft_clip (
int) – Boundary frames trimmed from each end of the spectrogram. Default 10, as in the upstreampreprocessors/stft.pyused with the released checkpoint (the demo notebook uses 5; seestft_zscore_before_clip).stft_zscore_before_clip (
bool) – Whether the spectrogram is z-scored before the boundary frames are trimmed. DefaultTrue, which together withstft_clip=10reproduces the recipe behind the published numbers. Set toFalsewithstft_clip=5to 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.Noneaverages all frames.activation (
type[Module]) – Transformer feed-forward activation class. Defaultnn.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 tomne.Infofor 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_hubpackage 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) – IfTrue, 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, hencecls_tokenisNone. 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"}whenreturn_featuresis set.- Return type:
torch.Tensor or dict