braindecode.models.BrainOmni#

class braindecode.models.BrainOmni(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, *, emb_dim=256, n_neuro=16, window_length=512, overlap_ratio=0.25, n_filters=32, ratios=(8, 4, 2), kernel_size=5, last_kernel_size=5, tokenizer_num_heads=4, codebook_dim=256, codebook_size=512, num_quantizers=4, rotation_trick=True, tokenizer_drop_prob=0.0, lm_dim=256, num_heads=8, depth=12, drop_prob=0.1, activation=<class 'torch.nn.modules.activation.SELU'>)[source]#

BrainOmni from Xiao et al. (2025) [brainomni].

Foundation Model Attention/Transformer

Architecture Overview

A frozen BrainTokenizer followed by factored spatial-temporal attention blocks and a classification head:

(batch, n_chans, n_times) -> BrainTokenizer.tokenize -> projection ->
spatial-temporal blocks -> mean over time -> (batch, n_outputs)

Macro Components

BrainOmni.tokenizer

Operations. BrainTokenizer.tokenize() with windows overlapping by overlap_ratio, plus the learned source embeddings. Role. Frozen feature extractor (no gradients, no codebook updates).

BrainOmni.projection

Operations. Linear(emb_dim, lm_dim), identity when equal.

BrainOmni.blocks

Operations. Half of the features attend over time (RoPE), the other half over the n_neuro sources, then a feed-forward layer. The last block is part of the pretrained stack but unused downstream, as released. Role. Space-time contextualization.

BrainOmni.final_layer

Operations. Dropout(0.1) -> Linear -> activation -> Linear on the flattened n_neuro * lm_dim features. Role. Classification head.

Temporal, Spatial, and Spectral Encoding

  • Temporal: RoPE attention over the token sequence of the windows.

  • Spatial: attention over the n_neuro sources.

  • Spectral: inherited from the tokenizer’s convolutions.

Additional Mechanisms

  • The block outputs are L2-normalized before pooling, as released.

  • The RoPE cache holds cosines only for the first 240 positions, as in the released checkpoints; longer sequences rebuild it from freqs.

Important

Weights converted from the released tiny and base checkpoints are on the Hugging Face Hub at braindecode/brainomni-tiny-pretrained and braindecode/brainomni-base-pretrained; the head is not pretrained:

from braindecode.models import BrainOmni
model = BrainOmni.from_pretrained("braindecode/brainomni-tiny-pretrained", chs_info=raw.info["chs"], n_outputs=2)

Input is expected at 256 Hz, preprocessed as in the authors’ code.

Added in version 1.8.

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.

  • emb_dim (int) – Tokenizer embedding dimension.

  • n_neuro (int) – Number of latent source tokens.

  • window_length (int) – Samples per tokenizer window.

  • overlap_ratio (float) – Overlap between tokenizer windows, in [0, 1).

  • n_filters (int) – Base number of SEANet filters.

  • ratios (tuple[int, ...]) – SEANet downsampling ratios.

  • kernel_size (int) – SEANet kernel size.

  • last_kernel_size (int) – Kernel size of the first and last SEANet convolutions.

  • tokenizer_num_heads (int) – Heads of the tokenizer cross-attention.

  • codebook_dim (int) – Codebook dimension.

  • codebook_size (int) – Entries per codebook.

  • num_quantizers (int) – Number of residual VQ stages.

  • rotation_trick (bool) – Rotation trick in the tokenizer quantizer.

  • tokenizer_drop_prob (float) – Tokenizer attention dropout.

  • lm_dim (int) – Transformer dimension.

  • num_heads (int) – Transformer heads (even: half temporal, half spatial).

  • depth (int) – Number of transformer blocks, the last one unused.

  • drop_prob (float) – Transformer dropout.

  • activation (type[Module]) – Activation of the classification head.

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

[brainomni]

Xiao, Q., Cui, Z., Zhang, C., Chen, S., Wu, W., Thwaites, A., Woolgar, A., Zhou, B., Zhang, C. (2025). BrainOmni: A Brain Foundation Model for Unified EEG and MEG Signals. NeurIPS 2025. https://arxiv.org/abs/2505.18185

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 BrainOmni

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

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

Loading a model from the Hub:

from braindecode.models import BrainOmni

# Load pretrained model
model = BrainOmni.from_pretrained("username/my-brainomni-model")

# Load with a different number of outputs (head is rebuilt automatically)
model = BrainOmni.from_pretrained("username/my-brainomni-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 = BrainOmni.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

encode(x)[source]#

L2-normalized (batch, n_neuro, n_tokens, lm_dim) embedding.

Parameters:

x (Tensor) – The description is missing.

Return type:

Tensor

forward(x, return_features=False)[source]#

Classify x or return the pooled features.

Parameters:
  • x (Tensor) – The description is missing.

  • return_features (bool) – The description is missing.

reset_head(n_outputs)[source]#

Replace the classification head for n_outputs classes.

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.