Skip to content

Add MSCFormer model (multi-scale convolutional transformer for motor imagery) - #1186

Merged
bruAristimunha merged 6 commits into
braindecode:masterfrom
qinxwew:add-mscformer-model
Sep 29, 2026
Merged

bruAristimunha merged 6 commits into
braindecode:masterfrom
qinxwew:add-mscformer-model

Conversation

@qinxwew

@qinxwew qinxwew commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

Closes #721.

Summary

Adds braindecode.models.MSCFormer, the multi-scale convolutional transformer
network for motor-imagery decoding from Zhao et al. (2025)
(Scientific Reports 15:12935, 10.1038/s41598-025-96611-5).
The model was proposed by its original author in #721, who asked for it to be
integrated into braindecode; I pinged them there before starting and gave it two
weeks, with no reply, so I am opening this PR for the maintainers to judge.

The port is adapted from the author's own reference implementation
(snailpt/MSCFormer, Apache-2.0 —
compatible with the license="apache-2.0" set on the class) and follows the
conventions of the existing CTNet port in this repository (same first author,
same architecture family), so the two models are directly comparable from the
same API.

Files:

  • braindecode/models/mscformer.py — the model
  • braindecode/models/__init__.py, braindecode/models/util.py (registered in
    models_mandatory_parameters, which is what drives the generic model test
    suite), braindecode/models/summary.csv (150,724 parameters), docs/api.rst
  • docs/whats_new.rst

No new test file is needed: test_integration.py and test_models.py are
parameterised over the registry, so registering the model picks up the existing
integration and config cases automatically — including
test_completeness__models_test_cases, which asserts the registry and the test
cases stay in sync.

The architecture is: three parallel temporal-convolution branches with kernel
widths (85, 65, 45) followed by a depth-wise spatial convolution, batch norm and
average pooling, concatenated into a 3 * n_filters_time embedding; a BERT-style
zero class token and a learnable positional encoding; a 5-layer transformer
encoder with post-norm residual connections; and a linear classifier on the
class-token position.

Deviations from the reference implementation

Everything below is deliberate, and only one item is behavioural:

  1. Attention logit scaling — UPDATED after a maintainer fidelity review (see
    the "Replication (braindecode maintainers)" section below). The reference
    divides the attention logits by emb_size ** 0.5 (scaling = self.emb_size ** (1 / 2) in MSCFormer_model.py) for every head, rather than by the standard
    per-head sqrt(head_dim) (with emb_size=48, num_heads=8:
    sqrt(48) ≈ 6.93 vs sqrt(head_dim) = sqrt(6) ≈ 2.45). The port now exposes
    this through a new attention_scale parameter on
    braindecode.modules.MultiHeadAttention (mirroring MIRepNet's
    attention_scale, Add MIRepNet model #1146), and defaults to embed_dim ** -0.5, matching the
    original source exactly
    rather than CTNet's head_dim ** -0.5 convention.
    A state-dict-matched numerical parity check against the original
    implementation (eval mode, fixed seed, batch=3 random input) gives a max abs
    logit diff of 2.4e-7 (float32 noise) with this default, versus 0.0346 with
    the previous head_dim ** -0.5 default. attention_scale stays overridable
    for anyone who wants the other convention. This is documented in the class
    docstring and covered by a new regression test.

  2. torch.zeros(...).cuda() for the class token → x.new_zeros(...), so the
    model runs on CPU/MPS. The reference hard-codes CUDA here.

  3. The x * sqrt(embed_dim) rescaling before the positional encoding is kept
    as-is, since it is the standard transformer input normalisation.

  4. activation_cnn/activation_ffn are exposed as class-valued parameters
    (nn.ELU/nn.GELU, matching the reference) per the interface convention for
    new models, and the classifier is named final_layer so it is picked up
    automatically.

Benchmark

Reproducible benchmark on BCI Competition IV-2a, subject-specific, 5-fold
cross-validation, all 9 subjects, following the reference implementation's
protocol (288 training trials per subject split into contiguous folds, Adam
lr=1e-3, betas=(0.5, 0.999), weight decay 0, batch size 72 concatenated with
216 segmentation-and-reconstruction-augmented trials, number_aug=3,
number_seg=8, model selection on best validation loss, test on the 288
evaluation trials).

Because a single model's number is not interpretable on its own, I ran the
already-merged braindecode.models.CTNet through the exact same harness as a
control, so the comparison is within-harness rather than against numbers copied
from papers.

Subject MSCFormer CTNet Δ
s1 0.7979 ± 0.037 0.7764 ± 0.017 +0.0215
s2 0.4597 ± 0.068 0.4451 ± 0.046 +0.0146
s3 0.9292 ± 0.023 0.9028 ± 0.015 +0.0264
s4 0.7243 ± 0.073 0.7299 ± 0.032 −0.0056
s5 0.6750 ± 0.118 0.6146 ± 0.064 +0.0604
s6 0.6410 ± 0.035 0.5604 ± 0.038 +0.0806
s7 0.8347 ± 0.047 0.7660 ± 0.093 +0.0688
s8 0.8389 ± 0.009 0.7931 ± 0.014 +0.0458
s9 0.7903 ± 0.015 0.7729 ± 0.035 +0.0174
Grand mean 0.7434 (κ 0.658) 0.7068 (κ 0.609) +0.0367

Each cell is the mean accuracy over that subject's 5 folds, ± the population
standard deviation over those folds (ddof=0); κ is the mean Cohen's kappa over
the same folds, and the grand means are over all 45 fold results.

Panel A: per-subject accuracy next to the papers' reported values. Panel B: paired per-fold scatter.

Panel A: per-subject accuracy next to the papers' reported values. Panel B: paired per-fold scatter.

How to read this — including the part that does not flatter the port.

Both models land well below the accuracies their papers report on this dataset
(MSCFormer 82.95%, CTNet 82.52%): −8.61 pp and −11.84 pp respectively. That gap
is systematic across both models rather than specific to MSCFormer, which points
at the harness rather than at the model. The likely contributors are: the
paper's preprocessing is reproduced from its description rather than from shared
code; model selection is on validation loss with an early-stopping patience of
150 epochs; and the runs are seeded and executed on Apple MPS instead of CUDA.

What the harness does reproduce is the relative ordering the paper reports:
MSCFormer outperforms CTNet, by +3.67 pp here versus the paper's own +0.43 pp,
and it does so consistently — better on 8 of 9 subjects and in 32 of 45
subject/fold pairs (panel B). The direction of the effect matches the paper; the
magnitude is larger, which I cannot fully attribute, so I would rather state it
than bury it.

Given #1156 (making model validation a more official gate) is open, I am happy
to rerun this under whatever protocol ends up being standard. The harness is a
single self-contained script that runs either model through the same code path
(same folds, same augmentation, same seed policy, same evaluation), and I am
happy to share it or to open a follow-up PR if the maintainers want it in-tree.

Replication (braindecode maintainers)

We re-audited the port against the original snailpt/MSCFormer source
(commit c627d3af, read-only clone kept at
model-replication-sources/MSCFormer/ with a MANIFEST.yaml recording every
verified line reference) and the published paper, independently of the
author's own benchmark above.

Fidelity fix. The attention-scaling divergence flagged in point 1 above was
the only behavioural difference found. We built an explicit state-dict mapping
between the original PyTorch module and the port and ran a numerical parity
check: eval mode, fixed seed, identical random input, n_chans=22, n_times=1000, pooling_size=44. Max abs logit diff was 0.0346 with the port's previous
default (head_dim ** -0.5) and 2.4e-7 (float32 noise) once MultiHeadAttention
is configured with the original's embed_dim ** -0.5 scale — which is now the
port's default, exposed as an overridable attention_scale parameter (same
pattern as MIRepNet, #1146). A regression test pins this default and the
validation of attention_scale.

Simplification audit. We reviewed the port for dead/duplicated code and for
reuse opportunities against shared braindecode modules with provably identical
output. The private _ResidualAdd / _TransformerEncoderBlock /
_TransformerEncoder / _PositionalEncoding classes duplicate CTNet's
equivalents by design (same convention as every other transformer-style model
port in this codebase — each keeps its own small private helpers rather than
sharing a cross-model base, so unifying them would be out-of-scope stylistic
churn touching unrelated models). No other dead code or needless wrappers were
found; the public API is unchanged apart from the new attention_scale
keyword.

Tests. With the fix applied: quick run (-k "mscformer and not compiled")
15 passed, 1 skipped; full local suite (test_models.py, test_integration.py,
test_return_features.py, -k "not compiled") 1558 passed, 146 skipped, 0
failed
in 320 s; pre-commit clean on all touched files (ruff
format/lint/isort, codespell, mypy, sphinx-lint), no rewrites on a second run.
Environment: macOS, braindecode env Python 3.12, PYTHONPATH=. (package not
pip-installed in that env).

Independent BCI IV-2a replication. Protocol taken from the reference notebook (MSCFormer.ipynb): subject-specific, session T train / session E test, 5 contiguous folds on T (56/56/56/56/64), Adam lr=1e-3, betas=(0.5, 0.999), batch 72 + S&R augmentation (number_seg=8, number_aug=3), a fixed 1000 epochs per fold with the best-validation-accuracy checkpoint, and the mean over the 5 folds. The released code has no patience-based early stopping. It draws S&R augmentation from all of session T, including each fold's validation trials; we kept this to match the source. Both implementations ran with identical splits, batches and seeds (seed 0) on Gaudi HPU (FP32, lazy mode).

Subject Port acc % Original acc % Port κ Original κ
1 81.11 81.04 0.748 0.747
2 54.79 54.51 0.397 0.394
3 87.57 85.69 0.834 0.809
4 78.54 79.44 0.714 0.726
5 64.10 65.21 0.521 0.536
6 59.44 58.47 0.459 0.446
7 83.82 82.92 0.784 0.772
8 81.67 81.32 0.756 0.751
9 82.36 82.92 0.765 0.772
Mean 74.82 74.61 0.664 0.662

The port matches the original code (difference +0.21 pp; per-subject range −1.11 to +1.87 pp). Both are about 8 pp below the paper's 82.95 % and close to the author's 74.34 %. So the gap versus the paper comes from the released training protocol, preprocessing or seeds, not from the port. We ran one seed, and HPU-vs-CPU float32 kernels differ slightly (documented).

Validation

  • pytest test/unit_tests/models/ — MSCFormer is exercised by the parametrised
    integration and config cases; the full model suite is green locally
    (2301 passed, 179 skipped, 0 failed). Like most models it is not in
    _DIRECT_TORCHSCRIPT_MODELS, so no scripting guarantee is claimed here;
    happy to add it to that list if the maintainers want the coverage.
  • Forward pass smoke check: MSCFormer(n_chans=22, n_outputs=4, n_times=1000)
    on 2×22×1000 µV input returns (2, 4); parameter count 150,724 matches the
    summary.csv entry.
  • pre-commit clean on the changed files (ruff format / ruff lint / codespell /
    mypy / sphinx-lint).
  • Parameters registered in summary.csv and in util.py's registry
    (("MSCFormer", ["n_chans", "n_outputs", "n_times"], None)).

Multi-scale convolutional transformer network for motor imagery
classification (Sci Rep 15, 12935), ported from the official
Apache-2.0 implementation following the CTNet adaptation.

- three parallel temporal-conv branches (kernels 85/65/45) with
  depth-wise spatial convolution, concatenated to a 48-dim embedding
- BERT-style class token + learnable positional encoding + 5-block
  post-norm Transformer encoder
- classification from the class-token representation
- registered in util.py / summary.csv / api.rst; integration tests
  picked up via the model registry

Closes braindecode#721 (model part; benchmark to follow)
@codecov

codecov Bot commented Sep 28, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 97.87234% with 2 lines in your changes missing coverage. Please review.
✅ Project coverage is 87.48%. Comparing base (9148f01) to head (eb287fb).

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1186      +/-   ##
==========================================
+ Coverage   87.43%   87.48%   +0.05%     
==========================================
  Files         146      147       +1     
  Lines       16906    17000      +94     
==========================================
+ Hits        14781    14873      +92     
- Misses       2125     2127       +2     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@qinxwew

qinxwew commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor Author

One note on the checks so it does not look like a code problem: test (macos-latest, 3.12) is red because the GitHub-hosted runner was shut down mid-run, not because a test failed. The job log ends with ##[error]The runner has received a shutdown signal at 86% progress, and every test up to that point shows PASSED; the Run Full Test Suite step is reported as cancelled, and no step is marked as failed. The same commit passes on test (macos-latest, 3.13) and on ubuntu 3.12/3.13 and windows 3.12/3.13. I cannot re-trigger that job from a fork (403 Must have admin rights), so a maintainer re-run is the only way to clear it.

Everything else is green, including check-whats-news. Locally the full model suite is green as well (2301 passed, 179 skipped, 0 failed).

bruAristimunha and others added 4 commits September 29, 2026 14:39
Resolve models/__init__.py and docs/whats_new.rst conflicts against master's newly-merged MIRepNet (braindecode#1146) and BaRISTA (braindecode#1173): keep both registry entries, restore alphabetical import/__all__ order, keep both changelog entries.
The MSCFormer port already declares license="apache-2.0" on the model
class and the Apache-2.0 header/adaptation link in the module docstring,
but the NOTICE.txt inventory was missing the corresponding line.
…default

The released snailpt/MSCFormer source scales attention logits by
sqrt(emb_size) (48) for every head, not sqrt(head_dim) (6). The two only
coincide when num_heads == 1; braindecode's shared MultiHeadAttention
defaulted to head_dim**-0.5, which is CTNet's convention but not this
model's original one.

Add an attention_scale parameter (None -> embed_dim**-0.5, following the
same pattern as MIRepNet's attention_scale), keep it overridable, and
validate it must be positive. State-dict-matched numerical parity against
the original implementation (eval mode, batch=3 random input): max abs
logit diff drops from 0.0346 (head_dim**-0.5) to 2.4e-7 (float32 noise)
with the new default. Add a regression test and update the docstring and
whats_new entry.
@qinxwew

qinxwew commented Sep 29, 2026

Copy link
Copy Markdown
Contributor Author

Merged the current master into this branch to clear a conflict in docs/whats_new.rst introduced by the recent merges on master (#1169 and the ZUNA changelog entry). Both entries are kept — the ZUNA entry from master stays above the MSCFormer one. The merge commit touches no code; CI is re-running.

@bruAristimunha

Copy link
Copy Markdown
Collaborator

looks nice for me @qinxwew, thanks for the clean PR!

@bruAristimunha
bruAristimunha merged commit f553e13 into braindecode:master Sep 29, 2026
13 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.

Multi-scale convolutional transformer network for motor imagery brain-computer interface

2 participants