Repository navigation
Add CSBrain cross-scale spatiotemporal brain foundation model - #1196
Conversation
Reimplementation of CSBrain (Zhou et al., NeurIPS 2025 Spotlight, arXiv:2506.23075) as an EEGModuleMixin model, verified bit-exact against the authors' reference implementation with their released pretrained checkpoint (max abs output diff 0.0 on matched inputs). - cross-scale temporal embedding (k=1/3/5), per-region circular-padded embedding and structured sparse attention (inter-window + masked inter-region) - channel names mapped to five anatomical regions, or an explicit brain_regions sequence (e.g. to reproduce the authors' per-dataset layouts); without channel info the model degrades to a single region - standard registrations: __init__, summary.csv, util.py, api.rst, architecture figure; targeted unit tests for region derivation, attention-mask structure and degenerate construction Closes braindecode#1077
c75645a to
8a472cd
Compare
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1196 +/- ##
==========================================
+ Coverage 88.32% 88.46% +0.13%
==========================================
Files 155 156 +1
Lines 18617 18885 +268
==========================================
+ Hits 16443 16706 +263
- Misses 2174 2179 +5 🚀 New features to boost your workflow:
|
# Conflicts: # docs/whats_new.rst
|
Merged the latest |
…model_compiled The port also Kaiming-initialised every Conv2d (temporal/region embedding, patch stem, positional conv), which the reference _weights_init does not. The residual stream then grows ~3x per layer and the default 12-layer model emits logits of ~1e4-1e5 at init, so Inductor's float32 reassociation error (~0.1 absolute) exceeds allclose(atol=1e-4) on some seeds/platforms (test_model_compiled[CSBrain] on macOS/3.13). No graph breaks are involved. Also mark the cls_token=None return with nosec B105 (Codacy/Bandit false positive, same as BrainBERT/Brant/MIRepNet).
The reference fine-tuning heads are flatten -> n_patch * 200 -> 200 -> n_outputs (e.g. 2000 hidden units for the 10 s FACED/CHB-MIT/Siena windows); the port fixed the first hidden layer at 4 * emb_dim, which matches only 4 s windows. The lazy head (no n_times) keeps 4 * emb_dim. Also correct two comments: the name-based region rule reproduces the reference PhysioNet-MI/FACED/SHU-MI/... layouts, while its BCI-IV-2a, SEED-V, SEED-VIG and Siena layouts put FC in central and PO in parietal; and without chs_info the region embedding is skipped (unmasked attention).
Keep every case of the per-model file in a CSBrain section of the shared suite, and add regression tests for the reference init (bounded residual stream), the window-dependent head width, the PhysioNet-MI name layout and the BCI-IV-2a brain_regions order. The two behaviour tests fail on c561f8a.
The reference fine-tuning models (CHB-MIT, Siena, SEED-V, ...) reorder the channels with a hand-made sorted_indices that sets the electrode ring inside each region; the stable grouping of the port differs, which changes the circular region convolution and the attention groups (random-weight parity vs the official code: 90.9 abs logit diff on Siena, 7025 on the CHB-MIT backbone). channel_order takes that permutation (validated: a permutation that groups the regions ascending) and gives 0.0 difference on both.
Review follow-up. The window-dependent head only applied when n_chans and n_times were passed explicitly; with chs_info or input_window_seconds + sfreq it fell back to the lazy 4 * emb_dim head (800 wide for 10 s windows, so the reference checkpoints could not load). Resolve both through the mixin. head_hidden_dim overrides the first hidden width for reference heads that are not n_patch * 200 wide (SEED-V: 800 for 1 s windows) and bounds the head for long windows. brain_regions is checked against n_chans; reset_head uses _set_n_outputs so the saved config follows the head.
|
Integration gate (braindecode maintainers) Target: paper Table 2, PhysioNet-MI (0.6304 ± 0.0090) and BCIC-IV-2a (0.5657 ± 0.0071, test bal. acc., subjects 8-9), reproduced from the released weights + public MOABB data via NeuralBench, head Protocol (2a): the data preparation the paper used (CBraMod Verified: forward parity port-vs-original is exact (max-abs diff 0.0) on both datasets; Results (updated 2026-10-06):
Verdict: the port matches the reference implementation; the 2a paper number does not reproduce upstream. Running the authors' released Small API note, not blocking: braindecode's single Licence: CSBrain ships no LICENSE file upstream, so no checkpoint mirror on the Hub until that is settled. |
|
Thanks for running the integration gate — the exact parity check ( On the LICENSE question — agreed, and here is where it stands:
Once #7 lands, I have the original One concrete gap worth flagging on the checkpoint side. On
Which would you prefer? (1) is a bit more code but makes the released weights drop-in; (2) keeps the PR lean. I'll go either way. |
# Conflicts: # test/unit_tests/models/test_models.py
…SBrain Review follow-up on the integration gate note: the reference fine-tuning scripts keep the backbone at 0.1 and raise only the head dropout, which a single drop_prob cannot express; document the wrapper/subclass recipe.
|
Thanks @bruAristimunha for the updated gate — good to see 2a move once the fine-tuning config matches Two follow-ups on the new head:
|
|
CI note on |
bruAristimunha
left a comment
There was a problem hiding this comment.
Thanks for the port and the patience! Replication on Voyager via NeuralBench: on BCIC-IV-2a this port gives 0.515, matching the authors' own released code run with their own data preparation (0.506 +/- 0.007); the paper's 0.566 is not reached by the released code itself (no train/test leak, no test-set selection), so the port is faithful to the released system. PhysioNet-MI 0.6135 vs 0.6304 (-2.7%). CI green.
Implements CSBrain [zhou2025csbrain] (NeurIPS 2025 Spotlight) as an
EEGModuleMixinmodel, as claimed in #1077.What
models/csbrain.py: CBraMod-style 200-sample patching + three stacked mechanisms per encoder layer (12 layers by default):chs_info) are mapped to five anatomical regions (frontal/parietal/temporal/occipital/central, the authors' canonical full-montage rule) and reordered to be contiguous; an explicitbrain_regionsparameter overrides the name derivation — this is how the authors' per-dataset layouts (e.g. their BCI-IV-2a config) are reproduced exactly. Without channel information the model degenerates to a single region and keeps full attention instead of failing.models/__init__.py,summary.csv(10,962,402 params atCSBrain(n_outputs=2)),util.py,api.rst, architecture figure.Correctness verification
Bit-exact against the reference implementation. Loading the authors' released pretrained checkpoint (
CSBrain.pth, Google Drive) into both the reference implementation (yuchen2199/CSBrain, with their BCI-IV-2a region layout) and this port — key-remapped (encoder.layers.*→encoder.*,TemEmbedEEGLayer.*→temporal_embed.*,BrainEmbedEEGLayer.*→region_embed.*) and shape-filtered, 289/295 tensors — identical inputs produce max abs output diff 0.0. The six skipped tensors are the temporal-region blocks with no matching channel group in the 2a layout; the authors' own fine-tuning code filters the same way.Benchmark (BCI IV-2a, same harness as #1186)
Subject-specific 5-fold CV over the 288 training trials of each subject's T session, evaluation on the 288 E-session trials — identical folds/seeds/test set as the MSCFormer/CTNet rows below. CSBrain is fine-tuned with the authors' recipe (pretrained init, AdamW lr 1e-4 wd 0.01, cosine to 1e-6, batch 64, label smoothing 0.1, drop 0.3, ~600 updates, selection on best validation accuracy; their BCI-IV-2a region layout via
brain_regions).Per-subject (CSBrain, mean±std over 5 folds): s1 .442±.033 · s2 .330±.006 · s3 .590±.028 · s4 .399±.045 · s5 .326±.026 · s6 .354±.026 · s7 .551±.020 · s8 .421±.016 · s9 .533±.024 (kappa .38 mean).
Reading the numbers honestly:
acc_0.57726in the filename, so ~0.58 is the authors' own level on 2a — the port adds nothing on top of that (bit-exactness) and subtracts nothing.License note
The upstream repository ships no LICENSE file at the time of writing (I opened yuchen2199/CSBrain#7 asking the authors to add one). This file is an independent reimplementation in the
EEGModuleMixinidiom following the paper (and the already-ported CBraMod for the shared patch-embedding design); the reference implementation was used only as a numerical cross-check (the bit-exactness section above). If the authors add a permissive license, a follow-up can switch the attribution to "adapted from" with a NOTICE.txt entry.Reproduction
bench_mscformer_2a.py/bench_csbrain_2a.py(same 5-fold subject-specific protocol, CSV-resumable; fine-tuning settings documented in the script headers)results_2a.csv(model, subject, fold, acc, kappa, best_epoch, seconds)Checklist
final_layernaming,activationclass-default parameter,drop_prob__init__,summary.csv,util.py,api.rst, architecture figure)docs/whats_new.rstentry (added with the PR number)Closes #1077