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.

BaRISTA pretraining with coordinate, parcel or lobe embeddings, followed by fine-tuning on a new subject.

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_encoder and temporal_pooler split each channel into patches, apply cnn_depth + 1 residual CNN blocks, then project each patch to d_model features. Each block has two dilated convolutions with parameter-free temporal LayerNorm and GELU. Dilation doubles between blocks; channels are encoded independently.

  • spatial_emb adds 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.

  • backbone applies n_layers pre-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_pooling reduces the sequence by a learned linear combination or a mean. final_layer maps 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_indices in input-channel order. NEMAR dataset nm000253 provides Brain Treebank’s indices in electrodes.tsv: x, y, z for coordinates, barista_parcel_index for parcels and barista_lobe_index for 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 to forward(); 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-parcels and braindecode/BaRISTA-lobes. Each repository also holds convert_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 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.

  • spatial_scale (str) – Spatial embedding scale, which sets the embedding tables. Coordinate mode falls back to chs_info when spatial_indices is omitted.

  • spatial_indices (list[int] | list[list[int]] | None) – Dataset-provided embedding indices in input-channel order, used for batches whose forward() 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 from chs_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 of patch_size loses 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 of d_model.

  • cnn_depth (int) – Number of hidden blocks of the dilated CNN temporal encoder; the encoder has cnn_depth + 1 blocks 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, default GELU.

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_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 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) for spatial_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 containing features of shape (batch, d_model) and cls_token=None.

Return type:

torch.Tensor or dict

reset_head(n_outputs)[source]#

Replace the linear classification head for a new n_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.