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). ACLStoken is prepended, a stack of standard Transformer encoder layers mixes the population, and theCLSoutput 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_timesaxis, so the model keeps the standard(batch, n_chans, n_times)input signature:n_chansis the number of electrodes andn_timesis the upstream feature dimension (768 for BrainBERTstftfeatures). Electrode coordinates are read fromchs_info(theirloc) 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
CLSoutput 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=1reproduces it).The defaults are the released
popt_brainbert_stftconfiguration:hidden_dim=512,ffn_dim=2048,n_heads=8,n_layers=6, used onn_times=768BrainBERT 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_chansandn_outputsmay be changed freely, as the population is pooled through theCLStoken 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 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.
hidden_dim (
int) – Transformer model widthD. 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) – Howchs_infopositions become integer coordinates."m"(default) treats them as MNE metres and rounds them to millimetres."raw"rounds the positions as they are: use it whenx/y/zalready 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 passcoordstoforward()directly.shift_coords (
bool) – IfTrue, shift each axis so that its smallest index is 0. DefaultFalse. 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. DefaultGELU.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"(callmodel.channel_layer.fitfirst) or"latent"map the montage ofchs_info(or of thechs_infogiven toforward) onto the backbone’s channels with aChannelLayer. Saved in the config.model(x)andmodel.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_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 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), wheren_timesis the upstream feature dimension.coords (
Tensor|None) – Integer coordinates of shape(batch, n_chans, 3). Defaults to the coordinates derived fromchs_infoat 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) – IfTrue, return{"features": cls, "cls_token": cls}(the pooledCLSrepresentation) instead of the class logits.key_padding_mask (
Tensor|None) – Boolean(batch, n_chans)mask,Truefor padded electrodes, so recordings with different electrode sets can share a batch (upstreamsrc_key_padding_mask). TheCLStoken is never masked.
- Returns:
Class logits of shape
(batch, n_outputs), or the feature dict whenreturn_featuresis 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.