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_embed

Operations. MLP on each channel’s position and orientation plus an EEG/MAG/GRAD type embedding, then RMSNorm. Role. Montage-agnostic channel identity.

BrainTokenizer.encoder

Operations. SEANet encodes each (channel, window); n_neuro learned queries attend over the sensor-conditioned channels. Role. (batch, n_neuro, n_windows, n_tokens, emb_dim) latents.

BrainTokenizer.quantizer

Operations. Residual vector quantization with EMA codebooks. Role. num_quantizers codebook indices per token.

BrainTokenizer.final_layer

Operations. 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_neuro sources, keyed by sensor geometry.

  • Spectral: learned implicitly by the convolutions.

Additional Mechanisms

  • Windows of window_length samples; a shorter input is zero-padded and an incomplete non-overlapping tail is dropped (zero-filled in the reconstruction).

  • tokenize() runs in eval mode 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 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) – 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.

  • 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 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_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 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.

forward(x)[source]#

Reconstruct x.

Parameters:

x (Tensor) – The description is missing.

Return type:

Tensor

tokenize(x, overlap_ratio=0.0)[source]#

Quantized features and codebook indices of x, in eval mode.

Parameters:
  • x (Tensor) – (batch, n_chans, n_times).

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

Returns:

  • feat (torch.Tensor) – (batch, n_neuro, n_windows * n_tokens, emb_dim).

  • indices (torch.Tensor) – (batch, n_neuro, n_windows * n_tokens, num_quantizers).