braindecode.models.PopulationTransformer#

class braindecode.models.PopulationTransformer(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, *, hidden_dim=512, ffn_dim=2048, n_layers=6, n_heads=8, max_len=5000, coord_units='m', shift_coords=False, activation=<class 'torch.nn.modules.activation.GELU'>, drop_prob=0.1, channel_strategy='native', channel_strategy_kwargs=None)[source]#

PopulationTransformer (PopT) from Chau et al. (2024) [PopT2024].

Foundation Model Attention/Transformer

PopT is a self-supervised population model for intracranial recordings (sEEG/iEEG). It does not encode a raw time signal; instead each electrode is represented by a feature vector — typically the frozen embedding of a per-channel foundation model such as BrainBERT — and PopT aggregates across electrodes. Every electrode feature is linearly projected and given a fixed sinusoidal spatial position encoding built from its integer anatomical coordinates (one embedding per X/Y/Z axis plus a sequence id). A CLS token is prepended, a stack of standard Transformer encoder layers mixes the population, and the CLS output is the pooled representation used for downstream decoding. Pre-training is by masked / replaced-token modelling over the electrode population.

Following the braindecode convention, the per-electrode feature vector plays the role of the n_times axis, so the model keeps the standard (batch, n_chans, n_times) input signature: n_chans is the number of electrodes and n_times is the upstream feature dimension (768 for BrainBERT stft features). Electrode coordinates are read from chs_info (their loc) and discretised to absolute integer indices inside the model, as upstream feeds them; when no positions are available the electrodes fall back to distinct sequential indices.

The CLS output goes through a single linear layer, as in the upstream fine-tuning model (PtDownstreamModel.linear_out, one logit trained with binary cross-entropy there; n_outputs=1 reproduces it).

The defaults are the released popt_brainbert_stft configuration: hidden_dim=512, ffn_dim=2048, n_heads=8, n_layers=6, used on n_times=768 BrainBERT features (~20M parameters).

Important

Pre-trained weights available. The official checkpoint is released by the authors and loads directly:

model = PopulationTransformer.from_pretrained(
    "braindecode/popt-pretrained", n_outputs=2
)

It uses the default configuration; n_chans and n_outputs may be changed freely, as the population is pooled through the CLS token and the classification head is task-specific (the checkpoint carries no trained fine-tuning head).

Warning

Evaluate with a time-blocked split. The paper’s downstream results use a random 80/10/10 split over word-aligned 5 s windows. Words are a fraction of a second apart, so almost every test window overlaps a training window, and labels that drift slowly in time leak into training. Re-running the paper setup (7 subjects, 3 seeds) with contiguous blocks of time, and dropping training windows that overlap the test set, pretrained PopT goes from 0.79 to 0.51 ROC-AUC on Pitch (chance) and from 0.89 to 0.64 on Volume. Onset (0.86 to 0.84) and Speech (0.90 to 0.84) hold up, and pretraining still beats training from scratch on Onset, Speech and Volume. The model and weights are not affected; the issue is only in the evaluation. When fine-tuning, split by blocks of time.

Added in version 1.8.2.

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.

  • hidden_dim (int) – Transformer model width D. Must be divisible by 8. Default 512, as the released model.

  • ffn_dim (int) – Inner dimension of the Transformer feed-forward blocks. Default 2048, as the released model.

  • n_layers (int) – Number of Transformer encoder layers. Default 6, as the released model.

  • n_heads (int) – Number of attention heads. Default 8, as the released model.

  • max_len (int) – Size of the coordinate table (largest addressable integer coordinate). Default 5000, as upstream.

  • coord_units (str) – How chs_info positions become integer coordinates. "m" (default) treats them as MNE metres and rounds them to millimetres. "raw" rounds the positions as they are: use it when x/y/z already hold the Brain Treebank integer (left, inferior, posterior) coordinates, as NEMAR nm000253 stores them. Either way the indices are absolute, not shifted, as upstream feeds them (pt_supervised_task_coords.py). They match the pretrained checkpoint only if the positions are already in the upstream (left, inferior, posterior) space; MNE head-frame positions (e.g. a standard montage) are not that space. Indices outside [0, max_len - 1] are clamped, with a warning. You can also pass coords to forward() directly.

  • shift_coords (bool) – If True, shift each axis so that its smallest index is 0. Default False. Upstream does not shift, and the shift changes the position encoding, and so the output of the pretrained model; use it only for positions with negative values when training from scratch.

  • activation (type[Module]) – Feed-forward activation, given as a class. Default GELU.

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

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

[PopT2024]

Chau, G., Wang, C., Talukder, S., Subramaniam, V., Soedarmadji, S., Yue, Y., Katz, B., & Barbu, A. (2024). Population Transformer: Learning Population-level Representations of Neural Activity. arXiv preprint arXiv:2406.03044.

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 PopulationTransformer

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

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

Loading a model from the Hub:

from braindecode.models import PopulationTransformer

# Load pretrained model
model = PopulationTransformer.from_pretrained("username/my-populationtransformer-model")

# Load with a different number of outputs (head is rebuilt automatically)
model = PopulationTransformer.from_pretrained("username/my-populationtransformer-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 = PopulationTransformer.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, coords=None, seq_id=None, return_features=False, key_padding_mask=None)[source]#

Aggregate a population of electrode features.

Parameters:
  • x (Tensor) – Per-electrode features of shape (batch, n_chans, n_times), where n_times is the upstream feature dimension.

  • coords (Tensor | None) – Integer coordinates of shape (batch, n_chans, 3). Defaults to the coordinates derived from chs_info at construction, broadcast over the batch.

  • seq_id (Tensor | None) – Integer sequence ids of shape (batch, n_chans). Defaults to zero (single population).

  • return_features (bool) – If True, return {"features": cls, "cls_token": cls} (the pooled CLS representation) instead of the class logits.

  • key_padding_mask (Tensor | None) – Boolean (batch, n_chans) mask, True for padded electrodes, so recordings with different electrode sets can share a batch (upstream src_key_padding_mask). The CLS token is never masked.

Returns:

Class logits of shape (batch, n_outputs), or the feature dict when return_features is set.

Return type:

torch.Tensor or dict

load_state_dict(state_dict, *args, **kwargs)[source]#

Also accept the untrained final_layer.{norm,fc} head of the HF mirror.

Parameters:
  • state_dict – The description is missing.

  • *args – The description is missing.

  • **kwargs – The description is missing.

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.