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
BrainTokenizerfollowed 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.tokenizerOperations.
BrainTokenizer.tokenize()with windows overlapping byoverlap_ratio, plus the learned source embeddings. Role. Frozen feature extractor (no gradients, no codebook updates).BrainOmni.projectionOperations.
Linear(emb_dim, lm_dim), identity when equal.BrainOmni.blocksOperations. Half of the features attend over time (RoPE), the other half over the
n_neurosources, 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_layerOperations.
Dropout(0.1) -> Linear -> activation -> Linearon the flattenedn_neuro * lm_dimfeatures. Role. Classification head.
Temporal, Spatial, and Spectral Encoding
Temporal: RoPE attention over the token sequence of the windows.
Spatial: attention over the
n_neurosources.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-pretrainedandbraindecode/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 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.
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.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_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 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