Repository navigation
Add AXON: axis-factorized EEG foundation model - #1182
Conversation
Adds braindecode.models.AXON with pretrained weights on the Hugging Face Hub (NeuroDX/axon-eeg), unit tests, and docs entries.
# Conflicts: # docs/whats_new.rst
Resolutions: - braindecode/models/__init__.py, braindecode/models/util.py, docs/api.rst, braindecode/models/summary.csv: keep both sides, AXON entries kept in alphabetical position next to master's new models. - docs/whats_new.rst: drop the conflict markers that were committed on the PR head (b2c3586), keep master's 1.8.1 section and place the :gh:`1182` AXON entry first under Enhancements; keep the `Mahir Jain` author link. - test/unit_tests/models/test_foundation_models.py: keep master's DIVER-1, ZUNA and STEEGFormer tests plus the AXON hub-loading test; fold the bespoke test/unit_tests/models/test_axon.py into the shared file (maintainer convention: no per-model test files) under a "Tests for AXON Model" section with the helpers renamed _axon_*. - test/unit_tests/models/test_return_features.py: auto-merged, AXON param kept next to EEGPT. - test_integration.py: _DIRECT_TORCHSCRIPT_MODELS count unchanged (32) on base, head and master; AXON is not a direct-TorchScript model.
- Pass license="apache-2.0" to EEGModuleMixin (the kwarg otherwise defaults to bsd-3-clause) and point the header at the reference implementation/weights; list braindecode/models/axon.py under the Apache-2.0 section of NOTICE.txt, following the LUNA/ZUNA convention (no bundled license text). - reset_head now calls _update_init_kwargs so get_config()/save_pretrained reflect the new n_outputs, as every other model with a custom head does on master; covered by test_axon_reset_head_updates_config.
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1182 +/- ##
==========================================
+ Coverage 89.23% 89.28% +0.05%
==========================================
Files 158 159 +1
Lines 19604 19746 +142
==========================================
+ Hits 17493 17631 +138
- Misses 2111 2115 +4 🚀 New features to boost your workflow:
|
|
Hi @bruAristimunha, CI status: 12 of 14 checks pass. The two failing jobs appear unrelated to the AXON changes. Would you be able to re-run the two failed jobs? On the Windows failure: AXON's three generic Hugging Face tests each save the full 118M-parameter model (about 475 MB) to Thank you. |
Resolves docs/whats_new.rst: keeps master's braindecode#1159 and braindecode#1155 entries and adds the AXON (braindecode#1182) entry above them; author link kept. test_foundation_models.py auto-merged.
|
Integration gate (braindecode maintainers) Target: paper Table 1 (BAC, mean ± std over 3 seeds), NeuralBench replication of the public Splits and recording mix (updated 2026-10-06): the subject splits now follow the EEG-FM-Bench dataset builders (xw1216/EEG-FM-Bench @325398d7), which match Table 8 on all six tasks. adftd uses the deterministic label-balanced split (64/11/13 subjects). hmc uses the seed-42 subject split (103/24/24). motor uses subjects 1–69 / 70–88 / 89–109, all 109 subjects, with a common average reference.
Status: not yet within the gate. Escalated to the maintainers for a decision; no change is requested in this PR. The model code and checkpoint are not the issue: every cell is at most one dataset-protocol detail away from the paper. adftd depends heavily on which subjects land in the test set: on the same checkpoint, LP reads 0.434 on one 13-subject split and 0.658 on another. The HF card's LP adftd (0.492) also differs from Table 1 (0.538). |
|
Hi @bruAristimunha, thank you for the review. Below are the splits: 1. Subject splitsThese come from the processed datasets the Table 1 evaluation read. The splits are subject-disjoint and fixed: all three downstream seeds use the same split, and only the training seed changes. adftd and hmc subject IDs (train / validation / test)adftd (64 / 11 / 13 subjects)
hmc (102 / 23 / 23 subjects)
motor_mv_img (PhysioNet EEG Motor Movement/Imagery, all 109 subjects): train = subjects 1–69, validation = 70–88, test = 89–109. 2. Run selection behind
|
| motor | bcic | workload | hmc | siena | adftd | mean | |
|---|---|---|---|---|---|---|---|
| LP, Table 1 | 0.457 | 0.292 | 0.673 | 0.650 | 0.867 | 0.538 | 0.579 |
| LP, validation only | 0.456 | 0.283 | 0.673 | 0.650 | 0.866 | 0.492 | 0.570 |
| FT, Table 1 | 0.625 | 0.428 | 0.685 | 0.732 | 0.859 | 0.604 | 0.656 |
| FT, validation only | 0.624 | 0.428 | 0.685 | 0.732 | 0.859 | 0.604 | 0.655 |
For a replication that selects on validation, the validation-only rows are the matching targets. We will state the selection rule explicitly in the next arXiv version.
|
One NeuralBench detail that may matter If the replication feeds AXON NeuralBench's |
Resolved docs/whats_new.rst by keeping both entries.
|
Hi @bruAristimunha, thank you for the update, and for #1251 on the Windows disk issue. Two points in the summary may help the maintainers' decision. Both use data already in this thread. 1. hmc split. The current hmc cell uses the builder's seed-42 split over 151 subjects (103 / 24 / 24). The Table 1 evaluation read processed hmc data with 148 subjects (IDs 14, 26, 52, 53, 64 and 135 are absent), split 102 / 23 / 23; the exact lists are in my comment of 5 October. LP hmc 0.693 is therefore measured on a different test set from Table 1's 0.650. 2. adftd: 0.492 (model card) vs 0.538 (Table 1). These two numbers come from the same split and the same three runs. They differ only in the epoch-selection rule (section 3 of the same comment). adftd LP is very sensitive to the epoch: in our runs, test BAC ranges from 0.33 to 0.68 across the 30 epochs, with 11 validation and 13 test subjects. Per seed, the best-validation-BAC epoch gives 0.474 / 0.491 / 0.511 on test, and the lowest-validation-loss epoch gives 0.545 / 0.534 / 0.534; Table 1 kept the latter. The adftd split (64 / 11 / 13, label-balanced) and the motor setup (subjects 1–69 / 70–88 / 89–109, common average reference) are the ones we used. The branch now conflicts with master in |
Conflicts: test_foundation_models.py (kept master's CodeBrain/SleepFM tests and the AXON block; dropped the stale DIVER-1 header master removed) and test_return_features.py (kept AXON; InterpolatedLaBraM entry was removed on master).
The temporal sinusoid table is built on the input's device and in at least float32 (sinusoidal_positional_encoding gains optional device/dtype arguments whose defaults keep every other caller unchanged), and the per-window z-score runs in float64 for a float64 model instead of being forced to float32.
Conflict: docs/whats_new.rst (kept the NeuroRVQTokenizer entry from braindecode#1223 and the AXON entry).
warn_if_sfreq_differs(model_name, sfreq, pretrained_sfreq) warns (does not raise) when sfreq is given and differs from the checkpoint's rate. AXON keeps its warning; outputs unchanged.
# Conflicts: # docs/whats_new.rst
bruAristimunha
left a comment
There was a problem hiding this comment.
Verified: AXON passes master's model contract, integration and pretrained-compat suites in our CPU test jobs, its forward follows the model's device and dtype, and outputs are bit-identical to the reviewed head, including with the released Hub weights. The author-protocol replication rerun is recorded on our side. Thanks @mahirjain01 for AXON and for publishing the exact splits and preprocessing!
…raindecode#1253) into EEG-CLIP; EEGCLIP text side in _UNUSED_IN_FORWARD
What this adds
braindecode.models.AXON, a transformer encoder for EEG pretrained with masked autoencoding (arXiv:2609.08788).Each window is cut into one token per electrode per 1-second patch. Every layer runs two attention paths in parallel:
A small per-token gate mixes the two paths.
Pretrained weights are on the Hugging Face Hub: NeuroDX/axon-eeg (Apache-2.0, 118.6M parameters).
Design
chs_info[i]["loc"](MNE head coordinates). Channels without a position are looked up by name in the standard 10-20 / 10-05 montages, usingresolve_montage_name. Positions are a non-persistent buffer, so the same weights load onto any montage.chs_infoandn_outputs. The model expects 200 Hz and warns for other rates. It raises an error ifn_timesis shorter than one patch.return_features=Truereturns{"features", "tokens", "cls_token": None};encode()returns the token grid;reset_head()replaces the head.Files
braindecode/models/axon.py: the model.braindecode/models/__init__.py,util.py(models_mandatory_parameters),summary.csv: registration.docs/api.rst,docs/whats_new.rst: docs.test/unit_tests/models/test_axon.py: 8 AXON-specific tests.test/unit_tests/models/test_return_features.py: AXON added to the feature-dict tests.test/unit_tests/models/test_foundation_models.py: one network test that loads the Hub weights.Verification
pytest test/unit_tests/models -k AXON --run-network: 33 passed. This covers TorchScript,torch.export,torch.compile, the Hugging Face round trip, config and return-features.summary.csv, badges): 138 passed, 1 skipped (GPU only).pre-commit runis clean on all changed files.Notes for reviewers