PopulationTransformer (PopT) β braindecode weights
braindecode-native re-host of the official PopulationTransformer (PopT)
popt_brainbert_stft checkpoint, remapped so it loads directly with
braindecode.models.PopulationTransformer.
PopT is a self-supervised population model for intracranial recordings
(iEEG/sEEG). It does not encode a raw time signal: each electrode is represented
by a feature vector β typically the frozen embedding of a per-channel foundation
model such as BrainBERT (stft features, 768-d) β and PopT aggregates across
the electrode population. Each electrode feature is linearly projected and given
a fixed sinusoidal spatial position encoding built from its integer
anatomical coordinates; a CLS token is prepended, a stack of Transformer
encoder layers mixes the population, and the CLS output is the pooled
representation used downstream.
Provenance
| Original code | https://github.com/czlwang/PopulationTransformer |
| Original weights | https://huggingface.co/PopulationTransformer/popt_brainbert_stft |
| Paper | Chau et al. (2024), Population Transformer: Learning Population-level Representations of Neural Activity, arXiv:2406.03044 |
These weights are a format conversion only of the authors' released
checkpoint β no re-training. The input embedding, the spatial positional
encoding and the 6-layer Transformer encoder are carried over bit-for-bit
(verified: all 82 mapped tensors are identical to the source, and the ported
encoder reproduces the upstream PtModelCustom output to < 1e-5). The
upstream masked-modelling heads (cls_head, token_cls_head) are not
carried; the braindecode classification head (final_layer) is
randomly initialised and must be trained/fine-tuned for your task.
Configuration
| param | value |
|---|---|
hidden_dim |
512 |
ffn_dim |
2048 |
n_layers |
6 |
n_heads |
8 |
n_times (feature dim) |
768 |
max_len (coord table) |
5000 |
| activation | GELU |
| parameters (encoder + spec head) | ~20.6M |
Usage
import torch
from braindecode.models import PopulationTransformer
# n_chans = number of electrodes; n_outputs = your task's classes.
model = PopulationTransformer.from_pretrained(
"braindecode/popt-pretrained", n_outputs=2
)
# input = per-electrode features (e.g. frozen BrainBERT stft embeddings),
# shape (batch, n_electrodes, 768). Electrode coordinates are read from
# chs_info when available, otherwise fall back to sequential indices.
x = torch.randn(4, 64, 768)
logits = model(x) # (4, n_outputs)
cls = model(x, return_features=True) # {"features": ..., "cls_token": ...}
The classification head is task-specific: pass your own n_outputs (the head is
re-initialised) and fine-tune. Electrode positions can be provided through
chs_info (their loc) so the spatial encoding reflects the real montage.
Licensing
The upstream PopulationTransformer repository ships no explicit license file,
so the license of these weights is marked unknown. They are re-hosted here
for research use with attribution; if you use them, cite the original work and
respect any terms the authors may later publish. The braindecode code is
BSD-3-Clause.
Citation
@article{chau2024population,
title={Population Transformer: Learning Population-level Representations of Neural Activity},
author={Chau, Geeling and Wang, Christopher and Talukder, Sabera and Subramaniam, Vighnesh and Soedarmadji, Saraswati and Yue, Yisong and Katz, Boris and Barbu, Andrei},
journal={arXiv preprint arXiv:2406.03044},
year={2024}
}
- Downloads last month
- -
Model tree for braindecode/popt-pretrained
Base model
PopulationTransformer/popt_brainbert_stft