Repository navigation
[models] Add TMSA-Net - #1209
bruAristimunha merged 9 commits into
Conversation
|
Exact-head validation is green on |
|
Current-head validation is now complete on |
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1209 +/- ##
==========================================
+ Coverage 89.02% 89.20% +0.18%
==========================================
Files 157 159 +2
Lines 19398 19727 +329
==========================================
+ Hits 17269 17598 +329
Misses 2129 2129 🚀 New features to boost your workflow:
|
|
Integration gate (braindecode maintainers) Target: BCI-IV-2a ( Already verified: numerical parity is exact (logits/input-gradient max-abs diff 0.0, 45/45 state entries, 20303 params), so the port is the same function as TMSA-Net (#1209), BCIC-IV-2a, 9 subjects x 5 seeds, port vs original code (matched init, same cached windows):
Best-val, like-for-like on the 45 cells where both codes finished: port 0.6749 vs original 0.6781 (diff -0.0032). Verdict (5% gate vs paper): PARTIAL: BCIC-IV-2a 9-subject mean (best-val, primary; 45/45 port, 45/45 original cells) port 0.6749 vs paper 0.8322 = -18.9% (outside 5%); like-for-like on the 45 cells where both codes finished: port 0.6749 vs original 0.6781 (diff -0.0032). Documented cause: declared substitute protocol. The paper's 2a preprocessing is a private .mat set (4-38 Hz MOABB windows substituted); the paper selects its number as the max TEST accuracy over 2000 epochs (we select on a 20% validation hold-out of session 1, which also removes 20% of the training data); its seed is random and unpublished. No action needed from you. If the exact 2a preprocessing script or a fixed seed exists outside the published repo, it would let us tighten the substitute protocol. |
Resolved NOTICE.txt docs/whats_new.rst by keeping both entries.
NOTICE.txt: master's MIT file list plus braindecode/models/tmsanet.py.
Fold the feature extractor, transformer wrapper and local-key conv module into their callers and use einops layers for the head split/merge and transposes. Drop guards torch already enforces; keep embed_dim >= num_heads, which otherwise builds an empty attention. The attention stays an explicit softmax so outputs, gradients and the dropout RNG order match the reference exactly. Tests: one test of the three released configurations, including the 19 -> 16 head bottleneck; the generic suites cover the rest. Shorter whats_new entry.
bruAristimunha
left a comment
There was a problem hiding this comment.
Verified: TMSANet is bit-identical to the authors' code, passes master's model contract and integration suites including the new CPU accelerator and low-precision checks, and with online Euclidean alignment (a declared deviation) the BCI IV 2a result lands within 5% of the paper. Thanks @lindicaphxag-tech for TMSA-Net!
…raindecode#1253) into EEG-CLIP; EEGCLIP text side in _UNUSED_IN_FORWARD
Part of #1201.
Summary
This adds TMSA-Net (Zhao & Zhu, 2025), one of the proposed first ports in #1201, as a Braindecode motor-imagery model.
The port keeps the released computation rather than replacing the custom attention with
nn.MultiheadAttention.In particular, the reference implementation uses
so the default BCI-IV-2a setting
embed_dim=19, num_heads=4is intentionally 19 → 16 → 19. The local-key and global-key attention branches are computed separately and summed before the output projection.Braindecode integration
TMSANetfollowsEEGModuleMixinand the standardn_chans / n_times / n_outputscontract.final_layeris exposed for the standard classifier/head integration.radixargument is not exposed: in the released source it only multiplies the input-channel count, while Braindecode'sn_chansalready denotes the actual number of input channels. All released configurations useradix=1.Reference fidelity
Pinned reference:
Whit3Zhao/TMSA-Net@c60882db35eeff860a5014df7b0f54dda6601c65A deterministic CPU parity assay maps all 45 model-state entries from the released implementation into this port. On the default BCI-IV-2a configuration and the same input:
Current-head parity run (
82b61e569ea9dcf0f90efd3eaf07ce05d1c1b128):https://github.com/lindicaphxag-tech/lindicaphxag-tech/actions/runs/37189419406
The parity check is against the released source implementation; this PR does not claim pretrained-checkpoint parity.
Validation
Focused regressions cover:
The current PR head passes 13 focused TMSA-Net tests, 12 TMSANet-selected Braindecode integration cases, and the touched-file pre-commit checks:
https://github.com/lindicaphxag-tech/lindicaphxag-tech/actions/runs/37189437290
The parity and focused validation above both target the current PR head; no pretrained-checkpoint result is implied.