Skip to content

[models] Add SeizureTransformer - #1236

Merged
bruAristimunha merged 2 commits into
braindecode:masterfrom
raghav-rathi:add-seizure-transformer
Oct 6, 2026
Merged

bruAristimunha merged 2 commits into
braindecode:masterfrom
raghav-rathi:add-seizure-transformer

Conversation

@raghav-rathi

@raghav-rathi raghav-rathi commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

Model information

Implementation fidelity

  • Reference implementation: https://github.com/keruiwu/SeizureTransformer at cf83f59 under MIT, plus the competition checkpoint from the authors' Docker image yujjio/seizure_transformer:latest, file wu_2025/model.pth with SHA-256 79b14e4715fef055ba252ae3f7a072325907e3d5d11bf91a8070d4220995f65d. I wrote the class from the layer table in the paper (Table 13) and checked it against their code. No code is copied from their repo.

  • Deviations:

    • Returns logits of shape (batch, n_outputs, n_times) instead of sigmoid probabilities of shape (batch, n_times)
    • One drop_prob for the residual blocks, the positional encoding and the Transformer. Upstream uses 0.1 for all three
    • Odd lengths use max pooling with ceil_mode=True and a crop in the decoder instead of padding with -1e10 and precomputed crops. The outputs are the same, see below, and inputs shorter than n_times also work
    • The positional table comes from braindecode.functional.sinusoidal_positional_encoding. On my machine it equals the stored one exactly but on others it can differ by up to 1.5e-5 because float32 sin varies between platforms. This does not change the outputs
    • The checkpoint also stores a spare copy of the Transformer layer that the forward pass never uses, because nn.TransformerEncoder deep-copies the layer it gets. This model leaves it out, so it has 37,848,897 parameters instead of 41,001,281
  • Checkpoint/parity evidence: After renaming, all 216 tensors the network uses load with strict=True. Largest absolute difference of the logits against the reference code on the same input

    input float64 CPU float32 CPU float32 GPU, TF32 off
    random 4 x 19 x 15360 9.8e-15 5.2e-6 7.6e-6
    random, lengths 15000, 15001, 7777 and 1001 9.8e-15 or less
    Siena sub-00 run-00, 44 windows 1.8e-5

    On the GPU, PyTorch's default TF32 convolutions shift the logits of both implementations away from the float64 result, so the GPU column has TF32 off.

Checklist

Implementation (braindecode/models/seizure_transformer.py)

  • Class inherits from EEGModuleMixin before nn.Module, license="mit", attribution in the header and in NOTICE.txt
  • Signal parameters forwarded to super().__init__(...)
  • Input and output shapes documented and tested, (batch, n_chans, n_times) to (batch, n_outputs, n_times) per-sample logits
  • self.final_layer is the last child module and there is no final sigmoid or softmax
  • activation=nn.ELU and activation_res=nn.ReLU class defaults and drop_prob
  • Docstring describes the architecture, parameters, paper and deviations, with [Wu2025]_ references
  • No new runtime dependencies

Registration and documentation

  • Exported in braindecode/models/__init__.py
  • Test case in models_mandatory_parameters with n_times=1024 to keep CI fast, and an entry in non_classification_models
  • Row in braindecode/models/summary.csv with 37,848,897 parameters
  • API entry in docs/api.rst
  • Architecture figure at docs/_static/model/seizure_transformer_arch.png from the authors' repo under MIT, listed in NOTICE.txt
  • Entry in docs/whats_new.rst

Validation and compatibility

  • Shared model tests plus three tests in test_models.py. They check one prediction per input sample for even, odd and shorter inputs, and errors for too-long inputs and invalid settings
  • pre-commit run --from-ref origin/master --to-ref HEAD passes with all hooks
  • Shared components changed: none
  • Save and load after head changes: test_reset_head_updates_config and test_reset_head_model_reloads_after_saving pass, and save_pretrained then from_pretrained of a trained model gives identical outputs
  • reset_head replaces the output convolution, no sentinel values

Validation evidence

  • pytest test/unit_tests/models/ -k "SeizureTransformer or seizure_transformer" on CPU gives 25 passed and 5 skipped. The skips are the non-classification exemptions and the embedding-parameter test. This includes torch.compile, torch.export, TorchScript and the new contract tests from Add registry-wide model contract gates #1208
  • The whole test_model_contract.py passes with this model registered, 146 passed and 2 skipped
  • Full suite with the tests.yml command pytest -vv --durations=0 -n 16 --dist worksteal test/ on CPU with Linux, Python 3.12 and torch 2.14 gave 4084 passed, 261 skipped and 0 failed. That run was on master from earlier today. After rebasing on the latest master I reran the tests above and pre-commit
  • Training through braindecode with EEGRegressor, BCEWithLogitsLoss and RAdam at lr 1e-4 and weight decay 2e-5 as in the paper. From scratch on 160 Siena windows from 11 subjects, 12 epochs on one RTX 4090 in 12 s. Training loss went from 0.714 to 0.356 and per-sample AUROC on 34 windows from 3 held-out subjects went from 0.588 to 0.871, or 0.855 for the same run on CPU. This only shows that training works, it is not a benchmark
  • Reduced precision on Siena with the released weights. Under float16 autocast all eight scores below stay the same. Under bfloat16 autocast the event scores stay the same and sample F1 goes from 0.3747 to 0.3737, because some event edges move a little, by up to 0.12 s in the recordings I checked closely. Mixed precision training with bfloat16, or float16 with GradScaler, follows the float32 loss curve, and the largest activation in the network is about 50, far below the float16 limit
  • Not run: macOS, Windows, Python 3.13 and the docs build

Benchmark reproduction

  • Benchmark on a public dataset: Siena, through braindecode.datasets.SIENA
  • Compared with the published values: the SzCORE leaderboard entry for this checkpoint, results/yujjio-seizure-transformer-latest/siena.json in esl-epfl/szcore
  • Protocol documented below
  • Limits explained below

All 41 Siena recordings were loaded with braindecode.datasets.SIENA and prepared the way the authors do it. Each channel is z-scored over the whole recording and cut into 60 s windows with the last one zero padded. Each window then gets a causal order 3 Butterworth band-pass from 0.5 to 120 Hz and notch filters at 1 Hz and 60 Hz. This class with the converted checkpoint predicted on the CPU and on one RTX 4090. After the authors' post-processing, which is a 0.8 threshold, opening and closing with 5 samples and removing events shorter than 2 s, szcore-evaluation scored the result the same way the leaderboard does.

SzCORE metric published braindecode
event F1 0.705967 0.705967
event sensitivity 0.684524 0.684524
event precision 0.886752 0.886752
event false positives per day 0.9914 0.9914
sample F1 0.374688 0.374688
sample sensitivity 0.269446 0.269446
sample precision 0.904780 0.904780
sample false positives per day 11.654 11.654

All eight values match the published ones on both devices, event precision only differs in the 16th decimal. Run under the same settings the reference code and this class also write identical event files for all 41 recordings, on the CPU and on the GPU with TF32 off.

Siena was part of the training data for this checkpoint, so this shows that braindecode reproduces the released model and not how well it generalises. The paper's held-out results use TUSZ, which needs a TUH data agreement, plus SeizeIT1 and Dianalund. I did not run those.

Notes for reviewers

  • The converted checkpoint is ready. Would you like it on the Hub under braindecode/ like the other ports? I'd ask the authors first and then add the from_pretrained example and the Hub test entry, here or in a follow-up PR
  • I used .. versionadded:: 1.8.2 like DIVER1 and BaRISTA
  • Happy to follow up with a seizure detection tutorial on Siena with SzCORE scoring if that helps

SeizureTransformer (Wu, Zhao and Yener, 2025) labels every sample of an
EEG window with a U-shaped network: a convolutional encoder, residual
convolution blocks, a Transformer encoder and a convolutional decoder
with skip connections. It won the 2025 SzCORE seizure detection
challenge.

The braindecode class is written from the paper and returns per-sample
logits of shape (batch, n_outputs, n_times). The authors' released
checkpoint converts to it: all 216 tensors the network uses load
strictly, and the outputs match the reference code to 1e-14 in float64.

Closes braindecode#1199
@bruAristimunha

Copy link
Copy Markdown
Collaborator

Integration gate (braindecode maintainers)

Thanks @raghav-rathi for a careful port and an unusually complete PR body: the checkpoint hash, key map and parity table made this quick to check.

Target. The paper's own cells (TUSZ, SeizeIT1, Dianalund) sit behind DUAs or are private. So the gate is the one public score of the released competition checkpoint (yujjio/seizure_transformer:latest, model.pth sha256 79b14e47…95f65d): the SzCORE benchmark on Siena (BIDS, Zenodo 10640762), event-based F1 0.7060 ± 0.3815 over 14 subjects. Declared deviation: Siena is in the checkpoint's training data, so this checks fidelity to the released system, not generalisation.

Protocol. Inference only, all 41 recordings, run on Voyager from one replicate.yaml (NeuralBench fork, experiments/replications/seizuretransformer/). The original column runs the authors' wu_2025 package from the Docker image verbatim. The port column runs the same preprocessing and post-processing with the network swapped for this PR's SeizureTransformer (head c95e107, 216 tensors, strict=True). Both are scored with szcore-evaluation.

event F1 sample F1
port (this PR) 0.7060 ± 0.3815 0.3747
original code 0.7060 ± 0.3815 0.3747
SzCORE benchmark (current evaluator) 0.7060 ± 0.3815 0.3747
2025 evaluator: port / published 0.8236 / 0.8213 0.4371 / 0.4371

The port and the original wrote byte-identical annotation TSVs for 41/41 recordings. The current-evaluator numbers match the benchmark to every digit. On a random 4×19×15360 input, the CPU probability difference against the authors' architecture is 7.5e-8 in fp32.

Verdict: REPLICATED (gap 0 %, accepted gap 5 %). Not blocking: the six red CI jobs were cancelled by runner starvation, not failed. The recomputed positional table differs from the stored one by 1.5e-5 (float32 sin), not "exactly", which has no effect on outputs.

@raghav-rathi

Copy link
Copy Markdown
Contributor Author

Thanks a lot for running the whole replication on your side! You're right about the positional table. On my machine it came out exactly equal but float32 sin isn't the same on every platform so I updated that line in the description

@codecov

codecov Bot commented Oct 5, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 88.04%. Comparing base (5e00a5b) to head (0555d69).
⚠️ Report is 4 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1236      +/-   ##
==========================================
+ Coverage   87.99%   88.04%   +0.05%     
==========================================
  Files         151      152       +1     
  Lines       17628    17713      +85     
==========================================
+ Hits        15511    15596      +85     
  Misses       2117     2117              
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@bruAristimunha bruAristimunha left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the port! Replication on Voyager via NeuralBench: the released SzCORE checkpoint through this port scores event-based F1 0.7060 on Siena (41 recordings), identical to the authors' code and to the published SzCORE benchmark to full float precision; weight parity 7.5e-8 (fp32), strict load of all 216 tensors. I bumped versionadded to 1.9.0. CI green.

@bruAristimunha
bruAristimunha merged commit bc80538 into braindecode:master Oct 6, 2026
21 of 22 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Add SeizureTransformer, the 2025 SzCORE seizure detection winner

2 participants