braindecode.models.CSBrain#
- class braindecode.models.CSBrain(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, patch_size=200, dim_feedforward=800, n_layer=12, nhead=8, activation=<class 'torch.nn.modules.activation.GELU'>, emb_dim=200, temporal_kernel_sizes=(1, 3, 5), drop_prob=0.1, head_drop_prob=None, brain_regions=None, channel_order=None, head_hidden_dim=None, return_encoder_output=False)[source]#
Cross-scale Spatiotemporal Brain Foundation Model from Zhou et al. (2025) [zhou2025csbrain].
Foundation Model Attention/Transformer
CSBrain is an EEG foundation model pre-trained with masked patch reconstruction, designed to decode brain activity across scales. It combines three mechanisms on top of CBraMod-style 200-sample patching:
Cross-scale temporal embedding: multi-scale convolutions (kernels 1/3/5 over the patch axis) fold brief bursts and slow rhythms into the same token vocabulary;
Region embedding: per-region convolutions with circular padding mix each anatomical region’s electrodes;
Structured sparse attention (SSA): inter-window attention over sliding windows of 5 patches, plus inter-region attention restricted by a group mask so each electrode attends to at most one electrode per other region, avoiding spurious long-range dependencies.
Channel names (
chs_info) are mapped to five anatomical regions (frontal / parietal / temporal / occipital / central) and reordered to be contiguous; unrecognised names fall into the central region. Withoutchs_info(and withoutbrain_regions) the region embedding is skipped and the inter-region attention is unmasked.- 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.
patch_size (
int) – Temporal patch size in samples (200 samples = 1 second at 200 Hz).dim_feedforward (
int) – Dimension of the feedforward network in the encoder layers.n_layer (
int) – Number of encoder layers.nhead (
int) – Number of attention heads.activation (
type[Module]) – Activation function used in the encoder feedforward blocks.emb_dim (
int) – Output dimension of the final projection applied to the encoder output.temporal_kernel_sizes (
Sequence[int]) – Kernel sizes of the cross-scale temporal embedding convolutions.drop_prob (
float) – Dropout probability of the backbone (patch embedding and encoder layers), and of the task head unlesshead_drop_probis set.head_drop_prob (
float|None) – Dropout probability of the task head.Noneusesdrop_prob. The reference fine-tuning models keep the backbone at 0.1 and set only the head dropout (--dropout 0.3for BCIC IV-2a), and they replaceproj_outbynn.Identity()after loading the pretrained weights:drop_prob=0.1, head_drop_prob=0.3, thenmodel.proj_out = nn.Identity().brain_regions (
Sequence[int] |None) – Explicit region id per input channel (0 frontal, 1 parietal, 2 temporal, 3 occipital, 4 central), taking precedence over the name-based derivation. Pass this to reproduce a dataset-specific region layout, e.g. the one used by the authors’ released fine-tuning checkpoints.channel_order (
Sequence[int] |None) – Permutation of the input channels applied before the region modules (the reference’ssorted_indices). It must group the channels by ascending region id; inside a region it sets the electrode ring of the circular region convolution and the round-robin attention groups.Nonekeeps the input order inside each region. Most reference fine-tuning models (e.g. CHB-MIT, Siena, SEED-V) use a hand-made topological order, so they need this to match exactly.head_hidden_dim (
int|None) – Width of the first hidden layer of the task head.Noneusesn_patch * emb_dim, the width of most reference fine-tuning heads (e.g. 800 for 4 s and 2000 for 10 s windows at 200 Hz). Some reference heads differ, e.g. SEED-V (1 s windows) uses 800; pass it here to load those checkpoints. The first head layer hasn_chans * n_patch * emb_dim * head_hidden_dimweights, which grows quadratically with the window length by default (about 2.3e9 for 64 channels and 30 s), so set a smaller width for long windows.return_encoder_output (
bool) – If False (default), the projected encoder output is flattened and passed through the task head to produce class logits of sizen_outputs. If True, return the encoder output features.
- Raises:
ValueError – If some input signal-related parameters are not specified: and can not be inferred.
Notes
The released checkpoints use other module names. To load one, drop the
module./backbone.prefix and renameencoder.layers.toencoder.,TemEmbedEEGLayer.totemporal_embed.,BrainEmbedEEGLayer.toregion_embed.,linear1toff_block.0andlinear2toff_block.3. Like the reference fine-tuning code, keep only the tensors whose name and shape match and load withstrict=False: the task head and the region convolutions of regions missing from the montage stay at their initial values.References
[zhou2025csbrain]Zhou, Y., Wu, J., Ren, Z., Yao, Z., Lu, W., Peng, K., Zheng, Q., Song, C., Ouyang, W., & Gou, C. (2025). CSBrain: A Cross-scale Spatiotemporal Brain Foundation Model for EEG Decoding. Advances in Neural Information Processing Systems (NeurIPS 2025, Spotlight). https://arxiv.org/abs/2506.23075
[csbraincode]Released implementation: yuchen2199/CSBrain
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 CSBrain # Train your model model = CSBrain(n_chans=22, n_outputs=4, n_times=1000) # ... training code ... # Push to the Hub model.push_to_hub( repo_id="username/my-csbrain-model", commit_message="Initial model upload", )
Loading a model from the Hub:
from braindecode.models import CSBrain # Load pretrained model model = CSBrain.from_pretrained("username/my-csbrain-model") # Load with a different number of outputs (head is rebuilt automatically) model = CSBrain.from_pretrained("username/my-csbrain-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 = CSBrain.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, mask=None, return_features=False)[source]#
Define the computation performed at every call.
Should be overridden by all subclasses.
Note
Although the recipe for forward pass needs to be defined within this function, one should call the
Moduleinstance afterwards instead of this since the former takes care of running the registered hooks while the latter silently ignores them.- Parameters:
x – The description is missing.
mask – The description is missing.
return_features – The description is missing.
- reset_head(n_outputs)[source]#
Replace the classification head for a new number of outputs.
This is called automatically by
from_pretrained()when the user passes ann_outputsthat differs from the saved config. Override in subclasses that need a model-specific head structure. Implementations keep changed constructor arguments in sync withself._update_init_kwargs, so that a saved model can be loaded back. Implementations requiring positive outputs can also useself._set_n_outputsto validate and record the new value.- Parameters:
n_outputs (int) – New number of output classes.
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.