braindecode.models.TMSANet#
- class braindecode.models.TMSANet(n_outputs=None, n_chans=None, n_times=None, chs_info=None, input_window_seconds=None, sfreq=None, embed_dim=19, pool_size=50, pool_stride=15, num_heads=4, fc_ratio=2, depth=1, drop_prob=0.5, att_drop_prob=0.5, fc_drop_prob=0.5, activation=<class 'torch.nn.modules.activation.GELU'>)[source]#
TMSA-Net from Zhao and Zhu (2025) [tmsanet].
Convolution Attention/Transformer
TMSA-Net combines multi-scale temporal convolutions, a spatial convolution across EEG channels, and a Transformer whose attention sums a global-key branch and a local-key branch (keys from multi-scale 1D convolutions) for motor-imagery classification.
As in the released code, the head width is
embed_dim // num_heads, so the defaultembed_dim=19with four heads projects queries, keys and values through 16 dimensions (19 -> 16 -> 19).- 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.
n_times (int) – Number of time samples of the input window.
chs_info (list of dict) – Information about each individual EEG channel. This should be filled with
info["chs"]. Refer tomne.Infofor more details.input_window_seconds (float) – Length of the input window in seconds.
sfreq (float) – Sampling frequency of the EEG recordings.
embed_dim (
int) – Embedding width after the temporal/spatial feature extractor. The released source uses 19 for BCI Competition IV 2a, 6 for BCI Competition IV 2b, and 10 for HGD.pool_size (
int) – Kernel size of the temporal average pooling layer.pool_stride (
int) – Stride of the temporal average pooling layer.num_heads (
int) – Number of attention heads.fc_ratio (
int) – Expansion ratio of the Transformer feed-forward block.depth (
int) – Number of Transformer encoder blocks.drop_prob (
float) – Dropout probability before the Transformer and after the local-key convolutions.att_drop_prob (
float) – Dropout probability on the attention weights (0.7 for HGD in the released source).fc_drop_prob (
float) – Dropout probability in the feed-forward block.activation (
type[Module]) – Activation after the spatial convolution and in the feed-forward block.
- Raises:
ValueError – If some input signal-related parameters are not specified: and can not be inferred.
Notes
Ported from
Whit3Zhao/TMSA-Net@c60882db35eeff860a5014df7b0f54dda6601c65; the referenceradixargument only multiplies the channel count and is not exposed.References
[tmsanet]Zhao, Q., Zhu, W. TMSA-Net: A novel attention mechanism for improved motor imagery EEG signal processing. Biomedical Signal Processing and Control 102, 107189 (2025). https://doi.org/10.1016/j.bspc.2024.107189
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 TMSANet # Train your model model = TMSANet(n_chans=22, n_outputs=4, n_times=1000) # ... training code ... # Push to the Hub model.push_to_hub( repo_id="username/my-tmsanet-model", commit_message="Initial model upload", )
Loading a model from the Hub:
from braindecode.models import TMSANet # Load pretrained model model = TMSANet.from_pretrained("username/my-tmsanet-model") # Load with a different number of outputs (head is rebuilt automatically) model = TMSANet.from_pretrained("username/my-tmsanet-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 = TMSANet.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)[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.