Repository navigation
Add MSCFormer model (multi-scale convolutional transformer for motor imagery) - #1186
Conversation
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 Report❌ Patch coverage is 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:
|
|
One note on the checks so it does not look like a code problem: Everything else is green, including |
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.
# Conflicts: # docs/whats_new.rst
|
Merged the current |
|
looks nice for me @qinxwew, thanks for the clean PR! |
Closes #721.
Summary
Adds
braindecode.models.MSCFormer, the multi-scale convolutional transformernetwork 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 theconventions 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 modelbraindecode/models/__init__.py,braindecode/models/util.py(registered inmodels_mandatory_parameters, which is what drives the generic model testsuite),
braindecode/models/summary.csv(150,724 parameters),docs/api.rstdocs/whats_new.rstNo new test file is needed:
test_integration.pyandtest_models.pyareparameterised 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 testcases 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_timeembedding; a BERT-stylezero 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:
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)inMSCFormer_model.py) for every head, rather than by the standardper-head
sqrt(head_dim)(withemb_size=48,num_heads=8:sqrt(48) ≈ 6.93vssqrt(head_dim) = sqrt(6) ≈ 2.45). The port now exposesthis through a new
attention_scaleparameter onbraindecode.modules.MultiHeadAttention(mirroring MIRepNet'sattention_scale, Add MIRepNet model #1146), and defaults toembed_dim ** -0.5, matching theoriginal source exactly rather than CTNet's
head_dim ** -0.5convention.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.5default.attention_scalestays overridablefor anyone who wants the other convention. This is documented in the class
docstring and covered by a new regression test.
torch.zeros(...).cuda()for the class token →x.new_zeros(...), so themodel runs on CPU/MPS. The reference hard-codes CUDA here.
The
x * sqrt(embed_dim)rescaling before the positional encoding is keptas-is, since it is the standard transformer input normalisation.
activation_cnn/activation_ffnare exposed as class-valued parameters(
nn.ELU/nn.GELU, matching the reference) per the interface convention fornew models, and the classifier is named
final_layerso it is picked upautomatically.
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 with216 segmentation-and-reconstruction-augmented trials,
number_aug=3,number_seg=8, model selection on best validation loss, test on the 288evaluation trials).
Because a single model's number is not interpretable on its own, I ran the
already-merged
braindecode.models.CTNetthrough the exact same harness as acontrol, so the comparison is within-harness rather than against numbers copied
from papers.
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 overthe 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.
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/MSCFormersource(commit
c627d3af, read-only clone kept atmodel-replication-sources/MSCFormer/with aMANIFEST.yamlrecording everyverified 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 previousdefault (
head_dim ** -0.5) and 2.4e-7 (float32 noise) onceMultiHeadAttentionis configured with the original's
embed_dim ** -0.5scale — which is now theport's default, exposed as an overridable
attention_scaleparameter (samepattern 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/_PositionalEncodingclasses duplicate CTNet'sequivalents 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_scalekeyword.
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, 0failed in 320 s;
pre-commitclean on all touched files (ruffformat/lint/isort, codespell, mypy, sphinx-lint), no rewrites on a second run.
Environment: macOS,
braindecodeenv Python 3.12,PYTHONPATH=.(package notpip-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), Adamlr=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).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 parametrisedintegration 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.
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.csventry.pre-commitclean on the changed files (ruff format / ruff lint / codespell /mypy / sphinx-lint).
summary.csvand inutil.py's registry(
("MSCFormer", ["n_chans", "n_outputs", "n_times"], None)).