braindecode.models.SeizureTransformer#

class braindecode.models.SeizureTransformer(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, *, n_filters=(32, 64, 128, 256, 512), encoder_kernel_sizes=(11, 9, 7, 7, 5), decoder_kernel_sizes=(3, 5, 5, 7, 7), res_kernel_sizes=(3, 3, 3, 3, 2, 3, 2), final_kernel_size=11, num_layers=8, num_heads=4, dim_feedforward=2048, drop_prob=0.1, activation=<class 'torch.nn.modules.activation.ELU'>, activation_res=<class 'torch.nn.modules.activation.ReLU'>)[source]#

SeizureTransformer from Wu et al. (2025) [Wu2025].

Convolution Attention/Transformer

SeizureTransformer architecture

SeizureTransformer with its default configuration, reproduced from [Wu2025] (MIT).#

Architecture Overview

SeizureTransformer labels every time sample of an EEG window instead of the window as a whole, so seizure onsets and offsets come straight out of the network without sliding-window inference [Wu2025]. It won the 2025 seizure detection challenge, scored with the SzCORE framework [Dan2025]. The network is U-shaped:

  1. A convolutional encoder halves the time axis at each level while it widens the features. With the defaults, a 60 s window at 256 Hz becomes 480 time steps of 512 features.

  2. Residual convolution blocks refine these features.

  3. Sinusoidal positions are added and a Transformer encoder relates every time step to the whole window. Its output is added back to the residual-block output.

  4. A convolutional decoder upsamples back to the input length and adds the encoder feature map of the same resolution at every level.

  5. A final convolution gives n_outputs logits for every time sample.

Macro Components

SeizureTransformer.encoder

Operations. At each level, a “same”-padded Conv1d with kernel encoder_kernel_sizes[i] and n_filters[i] output channels, followed by activation. The result is kept for the decoder, then SeizureTransformer.pool halves the time axis (max pooling, rounding odd lengths up).

Role. Down-sample the long input and extract local features at increasingly coarse time scales.

SeizureTransformer.res_blocks

Operations. Pre-activation residual blocks, one per entry of res_kernel_sizes: twice batch normalisation, activation_res, channel dropout (Dropout1d) and a “same”-padded convolution, plus an identity shortcut.

Role. Refine the down-sampled features with small kernels before global attention, as in EQTransformer [Mousavi2020].

SeizureTransformer.transformer

Operations. num_layers post-norm TransformerEncoderLayer blocks with num_heads heads and a dim_feedforward feed-forward width, applied after fixed sinusoidal positions and dropout.

Role. Capture dependencies across the whole window and give the model most of its capacity: about 25 of its 38 million parameters with the defaults.

SeizureTransformer.decoder

Operations. At each level, SeizureTransformer.upsample doubles the time axis (nearest neighbour), the result is cropped to the length of the matching encoder level, then a “same”-padded convolution with kernel decoder_kernel_sizes[i] and activation is applied and the encoder feature map is added.

Role. Restore the input resolution for per-sample predictions.

SeizureTransformer.final_layer

Operations. A Conv1d with final_kernel_size taps from n_filters[0] channels to n_outputs.

Role. Map each time sample to its logits. The output has shape (batch, n_outputs, n_times).

Usage

The model returns logits. With n_outputs=1, torch.sigmoid() gives the seizure probability of each sample, and the paper trains with binary cross-entropy [Wu2025]. The authors’ competition model uses the defaults and expects this input, which the model does not prepare itself:

  • 19 channels in the order Fp1, F3, C3, P3, O1, F7, T3, T5, Fz, Cz, Pz, Fp2, F4, C4, P4, O2, F8, T4, T6 (average reference, as in the SzCORE BIDS datasets such as SIENA);

  • each channel z-scored over the whole recording, then resampled to 256 Hz;

  • non-overlapping 60 s windows (n_times=15360), the last one padded with zeros;

  • in each window, a causal third-order Butterworth band-pass from 0.5 to 120 Hz, then IIR notch filters at 1 Hz and 60 Hz (quality factor 30).

The authors turn the probabilities into seizure events with a 0.8 threshold, a binary opening then a binary closing with a 5-sample structuring element, and the removal of events shorter than 2 s.

Differences from the Reference Implementation

  • The output is logits with a channel axis, (batch, n_outputs, n_times), instead of sigmoid probabilities of shape (batch, n_times).

  • A single drop_prob sets the dropout of the residual blocks, the positions and the Transformer (all 0.1 in the reference).

  • Odd lengths are handled by rounding the pooled length up and cropping in the decoder, so any input of at most n_times samples gives one prediction per sample. The positional table is computed instead of stored. The reference checkpoint also stores an unused copy of the Transformer layer that is not part of this model.

Added in version 1.9.0.

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.

  • n_filters (tuple[int, ...]) – Output channels of the encoder levels. The decoder mirrors them, and the last value is the Transformer width.

  • encoder_kernel_sizes (tuple[int, ...]) – Kernel size of the convolution at each encoder level, from the input down.

  • decoder_kernel_sizes (tuple[int, ...]) – Kernel size of the convolution at each decoder level, from the deepest level up.

  • res_kernel_sizes (tuple[int, ...]) – Kernel sizes of the residual blocks, one block per entry.

  • final_kernel_size (int) – Kernel size of the output convolution.

  • num_layers (int) – Number of Transformer encoder layers.

  • num_heads (int) – Number of attention heads. Must divide n_filters[-1].

  • dim_feedforward (int) – Hidden width of the Transformer feed-forward blocks.

  • drop_prob (float) – Dropout probability of the residual blocks, the positional encoding and the Transformer layers.

  • activation (type[Module]) – Non-linearity after the encoder and decoder convolutions.

  • activation_res (type[Module]) – Non-linearity inside the residual blocks.

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

[Wu2025] (1,2,3,4)

Wu, K., Zhao, Z., & Yener, B. (2025). Large EEG-U-Transformer for time-step level detection without pre-training. arXiv preprint arXiv:2504.00336. Code: https://github.com/keruiwu/SeizureTransformer (MIT).

[Dan2025]

Dan, J., Pale, U., Amirshahi, A., Cappelletti, W., Ingolfsson, T. M., Wang, X., Cossettini, A., Bernini, A., Benini, L., Beniczky, S., Atienza, D., & Ryvlin, P. (2025). SzCORE: Seizure Community Open-Source Research Evaluation framework for the validation of electroencephalography-based automated seizure detection algorithms. Epilepsia, 66(S3), 14-24. https://doi.org/10.1111/epi.18113

[Mousavi2020]

Mousavi, S. M., Ellsworth, W. L., Zhu, W., Chuang, L. Y., & Beroza, G. C. (2020). Earthquake transformer: an attentive deep-learning model for simultaneous earthquake detection and phase picking. Nature Communications, 11, 3952.

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 SeizureTransformer

# Train your model
model = SeizureTransformer(n_chans=22, n_outputs=4, n_times=1000)
# ... training code ...

# Push to the Hub
model.push_to_hub(
    repo_id="username/my-seizuretransformer-model",
    commit_message="Initial model upload",
)

Loading a model from the Hub:

from braindecode.models import SeizureTransformer

# Load pretrained model
model = SeizureTransformer.from_pretrained("username/my-seizuretransformer-model")

# Load with a different number of outputs (head is rebuilt automatically)
model = SeizureTransformer.from_pretrained("username/my-seizuretransformer-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 = SeizureTransformer.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]#

Predict logits for every time sample.

Parameters:

x (Tensor) – Input of shape (batch, n_chans, n_times). Inputs shorter than the configured n_times are accepted.

Returns:

Logits of shape (batch, n_outputs, n_times).

Return type:

Tensor

reset_head(n_outputs)[source]#

Replace the output convolution for a new number of outputs.

Parameters:

n_outputs (int) – New number of output classes.

Return type:

None

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.