braindecode.models.BaRISTA#
- class braindecode.models.BaRISTA(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, *, spatial_scale='coords', spatial_indices=None, coord_bins=200, patch_size=512, d_model=64, n_layers=12, num_heads=4, mlp_ratio=4, cnn_depth=4, cnn_channels=5, cnn_kernel_size=3, pooling='learned', drop_prob=0.1, activation=<class 'torch.nn.modules.activation.GELU'>)[source]#
BaRISTA from Oganesian et al (2025) [Oganesian2025].
Attention/Transformer Foundation Model Channel
Added in version 1.8.2.
BaRISTA encodes intracranial EEG with a shared temporal CNN and joint space-time attention. Its spatial embeddings use electrode coordinates, atlas parcels or lobes. This class provides the encoder and a classification head; it does not implement the masked pretraining objective.
Overview from the authors’ repository. The figure includes pretraining, which is outside this implementation.#
Architecture Overview
For patch \(i\) of channel \(j\), the temporal CNN and projection \(\mathcal{F}\) produce a token. The spatial embedding is added before attention:
\[\mathbf{S}_{ij} = \mathcal{F}(\mathbf{P}_{ij}) + \mathbf{E}_{sp(j)}.\]Tokens are ordered by patch, then channel:
\[\mathbf{S} = [\mathbf{S}_{11}, \ldots, \mathbf{S}_{1C}, \mathbf{S}_{21}, \ldots, \mathbf{S}_{nC}].\]Each attention layer therefore connects all electrodes and time patches.
Macro Components
patch_tokenizer,temporal_encoderandtemporal_poolersplit each channel into patches, applycnn_depth + 1residual CNN blocks, then project each patch tod_modelfeatures. Each block has two dilated convolutions with parameter-free temporal LayerNorm and GELU. Dilation doubles between blocks; channels are encoded independently.spatial_embadds a learned vector to each electrode’s tokens. Coordinate mode sums three embedding lookups. Parcel and lobe modes use one table, so electrodes assigned to the same region share an embedding.backboneappliesn_layerspre-norm blocks with RMSNorm, rotary self-attention and a GELU-gated feed-forward network by default. All channels in a patch share its rotary position.token_poolingreduces the sequence by a learned linear combination or a mean.final_layermaps the resulting vector to class logits. Learned pooling fixes the number of tokens; mean pooling permits it to vary between recordings.
Temporal, Spatial, and Spectral Encoding
Temporal encoding uses non-overlapping patches and patch-index rotary embeddings. Samples beyond the last complete patch are dropped.
Spatial encoding uses dataset-provided coordinate or region indices.
spatial_scale="none"disables it.The CNN operates on waveforms. There is no explicit spectral transform.
Additional Mechanisms
Supply
spatial_indicesin input-channel order. NEMAR datasetnm000253provides Brain Treebank’s indices inelectrodes.tsv:x, y, zfor coordinates,barista_parcel_indexfor parcels andbarista_lobe_indexfor lobes. Region index 0 denotes an unknown region and contributes no spatial embedding. The coordinate fallback bins finite, same-frame MNE positions onto a centred 1 mm grid. Use the dataset’s indices with the released weights: the MNE fallback does not recover Brain Treebank’s coordinate convention.With
pooling="mean", channel counts and window lengths can vary between batches, provided each window contains a full patch. Pass each recording’s indices toforward(); otherwise it uses the constructor’s indices or MNE positions. All samples in one batch share a montage. The encoder uses PyTorch attention on separate batch items, corresponding to the reference’s block-diagonal attention mask.Pre-trained weights
The three released encoders are published as
braindecode/BaRISTA-coords,braindecode/BaRISTA-parcelsandbraindecode/BaRISTA-lobes. Each repository also holdsconvert_barista_weights.py, the script that produced it: it downloads the release from a pinned source revision, checks the SHA-256 hash, renames tensors, combines the gated projections and checks encoder tokens against the released forward equations on float32 CPU inputs (explicit PyTorch attention in place of xformers; downstream accuracy and mixed precision are not tested). Load one and supply the montage indices of the batch:model = BaRISTA.from_pretrained("braindecode/BaRISTA-parcels", n_chans=64) logits = model(x, spatial_indices=parcel_indices)
They pool by mean, so one encoder serves any montage and window length. The releases contain no downstream head: pooling and classifier weights in the converted models are newly initialized and require fine-tuning.
License (non-commercial).
- 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.
spatial_scale (
str) – Spatial embedding scale, which sets the embedding tables. Coordinate mode falls back tochs_infowhenspatial_indicesis omitted.spatial_indices (
list[int] |list[list[int]] |None) – Dataset-provided embedding indices in input-channel order, used for batches whoseforward()does not pass its own. Shape(n_chans, 3)for coordinates, with values in[0, coord_bins); shape(n_chans,)for parcels or lobes, with values in[0, 121)or[0, 21)respectively. Region index 0 denotes unknown. Use the dataset’s BaRISTA index mapping, not arbitrary atlas label numbers.coord_bins (
int) – Number of slots per coordinate axis, default 200. Also the grid width in millimetres when deriving indices fromchs_info.patch_size (
int) – Number of samples per temporal patch, default 512 (250 ms at the paper’s 2048 Hz). Windows are tokenized into whole patches, so a window that is not a multiple ofpatch_sizeloses its trailing samples, as in the reference.d_model (
int) – Token embedding dimension.n_layers (
int) – Number of transformer encoder blocks.num_heads (
int) – Number of attention heads.mlp_ratio (
int) – Hidden dimension of the feed-forward blocks, as a multiple ofd_model.cnn_depth (
int) – Number of hidden blocks of the dilated CNN temporal encoder; the encoder hascnn_depth + 1blocks in total, the last one mapping back to a univariate signal.cnn_channels (
int) – Number of feature maps of the hidden blocks of the dilated CNN.cnn_kernel_size (
int) – Convolution width of the dilated CNN.pooling (
str) – Token aggregation before the head."learned"reproduces the paper’s finetuning protocol, a bias-free linear combination of the tokens, and requires the same total token count as at construction."mean"averages tokens and accepts different montages and window lengths.drop_prob (
float) – Dropout rate used in the encoder.activation (
type[Module]) – Activation layer class of the feed-forward blocks, defaultGELU.
- 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
[Oganesian2025]Oganesian, L. L., Hashemi, S. & Shanechi, M. M. (2025). BaRISTA: Brain scale informed spatiotemporal representation of human intracranial neural activity. Advances in Neural Information Processing Systems 38. https://arxiv.org/abs/2512.12135
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 BaRISTA # Train your model model = BaRISTA(n_chans=22, n_outputs=4, n_times=1000) # ... training code ... # Push to the Hub model.push_to_hub( repo_id="username/my-barista-model", commit_message="Initial model upload", )
Loading a model from the Hub:
from braindecode.models import BaRISTA # Load pretrained model model = BaRISTA.from_pretrained("username/my-barista-model") # Load with a different number of outputs (head is rebuilt automatically) model = BaRISTA.from_pretrained("username/my-barista-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 = BaRISTA.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, spatial_indices=None, return_features=False)[source]#
Encode an iEEG batch into class logits.
- Parameters:
x (
Tensor) – Input of shape(batch, n_chans, n_times).spatial_indices (
Optional[Tensor]) – Embedding indices of this batch’s montage, of shape(n_chans, 3)forspatial_scale="coords"and(n_chans,)otherwise. Every sample of the batch shares them. Defaults to the montage resolved at construction, which only fits the construction-time channel count.return_features (
bool) – Return the pooled embedding instead of logits.
- Returns:
Logits of shape
(batch, n_outputs), or a dictionary containingfeaturesof shape(batch, d_model)andcls_token=None.- Return type:
torch.Tensor or dict