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 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:
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.
Residual convolution blocks refine these features.
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.
A convolutional decoder upsamples back to the input length and adds the encoder feature map of the same resolution at every level.
A final convolution gives
n_outputslogits for every time sample.
Macro Components
SeizureTransformer.encoderOperations. At each level, a “same”-padded
Conv1dwith kernelencoder_kernel_sizes[i]andn_filters[i]output channels, followed byactivation. The result is kept for the decoder, thenSeizureTransformer.poolhalves 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_blocksOperations. 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.transformerOperations.
num_layerspost-normTransformerEncoderLayerblocks withnum_headsheads and adim_feedforwardfeed-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.decoderOperations. At each level,
SeizureTransformer.upsampledoubles the time axis (nearest neighbour), the result is cropped to the length of the matching encoder level, then a “same”-padded convolution with kerneldecoder_kernel_sizes[i]andactivationis applied and the encoder feature map is added.Role. Restore the input resolution for per-sample predictions.
SeizureTransformer.final_layerOperations. A
Conv1dwithfinal_kernel_sizetaps fromn_filters[0]channels ton_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_probsets 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_timessamples 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 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.
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 dividen_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_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 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