braindecode.models.BrainTokenizer#
- class braindecode.models.BrainTokenizer(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, 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, drop_prob=0.0)[source]#
BrainTokenizer from Xiao et al. (2025) [brainomni].
Foundation Model Attention/Transformer
Architecture Overview
A VQ-VAE over raw EEG/MEG conditioned on sensor geometry:
(batch, n_chans, n_times) -> windows -> SEANet encoder -> cross-attention to n_neuro sources -> residual VQ -> cross-attention back to n_chans -> SEANet decoder -> (batch, n_chans, n_times)
Macro Components
BrainTokenizer.sensor_embedOperations. MLP on each channel’s position and orientation plus an EEG/MAG/GRAD type embedding, then RMSNorm. Role. Montage-agnostic channel identity.
BrainTokenizer.encoderOperations. SEANet encodes each
(channel, window);n_neurolearned queries attend over the sensor-conditioned channels. Role.(batch, n_neuro, n_windows, n_tokens, emb_dim)latents.BrainTokenizer.quantizerOperations. Residual vector quantization with EMA codebooks. Role.
num_quantizerscodebook indices per token.BrainTokenizer.final_layerOperations. The sensor embeddings query the quantized sources, then the SEANet decoder rebuilds each window. Role. Reconstruction.
Temporal, Spatial, and Spectral Encoding
Temporal: SEANet’s strided convolutions and LSTM downsample each window by
prod(ratios).Spatial: cross-attention between channels and
n_neurosources, keyed by sensor geometry.Spectral: learned implicitly by the convolutions.
Additional Mechanisms
Windows of
window_lengthsamples; a shorter input is zero-padded and an incomplete non-overlapping tail is dropped (zero-filled in the reconstruction).tokenize()runs inevalmode without gradients, so the codebooks are never EMA-updated.
Important
Weights converted from the released checkpoint are on the Hugging Face Hub at
braindecode/braintokenizer-pretrained:from braindecode.models import BrainTokenizer model = BrainTokenizer.from_pretrained("braindecode/braintokenizer-pretrained", chs_info=raw.info["chs"])
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) – Embedding dimension.n_neuro (
int) – Number of latent source tokens.window_length (
int) – Samples per analysis window.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 cross-attention blocks.codebook_dim (
int) – Codebook dimension.codebook_size (
int) – Entries per codebook.num_quantizers (
int) – Number of residual VQ stages.rotation_trick (
bool) – Pass encoder gradients through the rotation trick instead of the straight-through estimator.drop_prob (
float) – Attention dropout.
- 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 BrainTokenizer # Train your model model = BrainTokenizer(n_chans=22, n_outputs=4, n_times=1000) # ... training code ... # Push to the Hub model.push_to_hub( repo_id="username/my-braintokenizer-model", commit_message="Initial model upload", )
Loading a model from the Hub:
from braindecode.models import BrainTokenizer # Load pretrained model model = BrainTokenizer.from_pretrained("username/my-braintokenizer-model") # Load with a different number of outputs (head is rebuilt automatically) model = BrainTokenizer.from_pretrained("username/my-braintokenizer-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 = BrainTokenizer.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_decode(x)[source]#
Return the reconstruction, commitment loss and codebook indices.
A dropped non-overlapping tail is zero-filled.
- Parameters:
x (
Tensor) – The description is missing.