Skip to content

Add AXON: axis-factorized EEG foundation model - #1182

Merged
bruAristimunha merged 15 commits into
braindecode:masterfrom
mahirjain01:add-axon
Oct 9, 2026
Merged

bruAristimunha merged 15 commits into
braindecode:masterfrom
mahirjain01:add-axon

Conversation

@mahirjain01

Copy link
Copy Markdown
Contributor

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 temporal path: each token attends to its own electrode across time;
  • a spatial path: each token attends to all electrodes at the same time step.

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).

from braindecode.models import AXON

model = AXON.from_pretrained("NeuroDX/axon-eeg", chs_info=raw.info["chs"], n_outputs=4, n_times=800)

Design

  • Any montage, any channel order. Electrodes are identified by their 3D position, taken from chs_info[i]["loc"] (MNE head coordinates). Channels without a position are looked up by name in the standard 10-20 / 10-05 montages, using resolve_montage_name. Positions are a non-persistent buffer, so the same weights load onto any montage.
  • Required arguments: chs_info and n_outputs. The model expects 200 Hz and warns for other rates. It raises an error if n_times is shorter than one patch.
  • Input scaling: each channel is z-scored within each window inside the model. The z-score is scale-free, so volts and microvolts give the same output. It is computed in float32, so float16 models stay finite on flat channels.
  • Standard braindecode API: return_features=True returns {"features", "tokens", "cls_token": None}; encode() returns the token grid; reset_head() replaces the head.
  • Scope: only the encoder and a classification head are included. The pretraining decoder is not.

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

  • Same outputs as the original research code. I compared AXON with the exact model code that produced the paper's results, on the same checkpoint.
    • Random inputs, 3 montages (19, 7 and 22 shuffled channels) × 3 window lengths: largest token difference 3.6e-6.
    • 32 real 64-channel motor-imagery windows, with positions looked up from channel names only: largest difference 5.5e-6. The looked-up positions are identical to the ones the evaluation pipeline used.
  • The weights downloaded from the Hub give the same features.
  • Tests:
    • 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.
    • The model-wide checks (registry, summary.csv, badges): 138 passed, 1 skipped (GPU only).
  • pre-commit run is clean on all changed files.

Notes for reviewers

  • The network test downloads the 475 MB weight file.

mahirjain01 and others added 5 commits September 28, 2026 11:25
Adds braindecode.models.AXON with pretrained weights on the Hugging Face Hub (NeuroDX/axon-eeg), unit tests, and docs entries.
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

codecov Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 97.24138% with 4 lines in your changes missing coverage. Please review.
✅ Project coverage is 89.28%. Comparing base (cb1ca88) to head (30294f8).
⚠️ Report is 1 commits behind head on master.

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:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@mahirjain01

Copy link
Copy Markdown
Contributor Author

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 tmp_path, so it adds roughly 1.4 GB to the runner's disk during the test session. If that is a concern, I would be happy to add a smaller AXON configuration for those tests, in whatever form you prefer.

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.
@bruAristimunha

bruAristimunha commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

Integration gate (braindecode maintainers)

Target: paper Table 1 (BAC, mean ± std over 3 seeds), NeuralBench replication of the public MannasAI/axon-eeg checkpoint (epoch 10; Table 10 shows this is the validation-selected epoch behind Table 1). Protocol: AdamW, the paper's LR schedule, 30 epochs, seeds 33/87/145, LP trains only final_layer; FT uses the HF-card recipe (encoder LR 0.1× head LR, 3 head-only epochs, axis_gate/scale_gate frozen). Gate: within 5% of the paper, or within the paper's own spread if that is larger. Parity to an independent fp64 reimplementation is exact (logits ≤1.2e-16).

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.

cell port paper gap within gate?
LP bcic 0.294 ± .006 0.292 ± .007 +0.7% ✅
FT bcic 0.455 ± .031 0.428 ± .033 +6.3% (Δ 0.027 < paper std) ✅
FT motor 0.624 ± .006 0.625 ± .005 −0.2% ✅
LP motor 0.421 ± .001 0.455 ± .005 −7.5% ❌
LP adftd 0.658 ± .016 0.538 ± .006 +22.3% ❌
FT adftd 0.563 ± .057 0.604 ± .022 −6.8% ❌
LP hmc 0.693 ± .003 0.650 ± .002 +6.6% ❌
FT hmc running 0.732 ± .009 — pending

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).

@bruAristimunha bruAristimunha added model Adds a new model needs-replication Model PR: paper number must be replicated (NeuralBench) before merge labels Oct 5, 2026
@mahirjain01

Copy link
Copy Markdown
Contributor Author

Hi @bruAristimunha, thank you for the review. Below are the splits:

1. Subject splits

These 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)

  • train: 001, 002, 003, 004, 005, 006, 007, 008, 009, 011, 012, 013, 017, 018, 020, 022, 023, 024, 025, 027, 028, 029, 030, 032, 033, 035, 036, 037, 039, 043, 045, 046, 047, 049, 050, 051, 052, 053, 054, 055, 057, 058, 059, 060, 061, 063, 064, 066, 067, 068, 069, 070, 071, 072, 075, 076, 077, 081, 082, 084, 085, 086, 087, 088
  • validation: 010, 016, 031, 034, 040, 042, 048, 062, 074, 080, 083
  • test: 014, 015, 019, 021, 026, 038, 041, 044, 056, 065, 073, 078, 079

hmc (102 / 23 / 23 subjects)

  • train: 1, 2, 3, 4, 6, 7, 8, 9, 15, 16, 17, 19, 22, 23, 25, 27, 28, 31, 35, 36, 37, 38, 40, 41, 42, 43, 44, 46, 47, 49, 50, 51, 54, 55, 57, 58, 59, 62, 63, 65, 66, 67, 68, 69, 71, 74, 76, 77, 78, 79, 80, 81, 83, 84, 85, 86, 89, 90, 92, 93, 94, 95, 96, 97, 98, 100, 101, 103, 104, 105, 106, 107, 108, 109, 112, 113, 114, 115, 116, 117, 118, 119, 120, 122, 123, 126, 127, 128, 129, 132, 134, 136, 137, 139, 140, 144, 145, 148, 149, 150, 152, 154
  • validation: 10, 13, 20, 21, 29, 32, 33, 34, 39, 56, 72, 73, 75, 82, 87, 102, 110, 111, 130, 131, 142, 146, 151
  • test: 5, 11, 12, 18, 24, 30, 45, 48, 60, 61, 70, 88, 91, 99, 121, 124, 125, 133, 138, 141, 143, 147, 153

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_mv_img

  • Motor imagery runs only: R04, R06, R08, R10, R12, R14. Execution runs and baselines are excluded.
  • Events T1/T2 only (T0 rest is excluded), giving 4 classes: R04/R08/R12 → left fist / right fist; R06/R10/R14 → both fists / both feet.
  • One 4 s window per cue, starting at cue onset (90 windows per subject).

Preprocessing for all datasets: 0.1 Hz high-pass (plus a 100 Hz low-pass when the native rate exceeds 200 Hz), mains notch where the data are not already notched (60 Hz for motor, 50 Hz for hmc), resampling to 200 Hz in µV, then the per-window, per-channel z-score inside the model. Electrode positions are MNE standard_1020 positions in the head frame, in cm, which is what chs_info[i]["loc"] gives in the braindecode port.

3. The adftd difference (0.538 vs 0.492)

This difference does not come from the split; both numbers use the split above. It comes from the downstream epoch-selection rule. For each task, Table 1 reports the better, by test BAC, of two candidate epochs: the best-validation-BAC epoch and the lowest-validation-loss epoch. The model card selects on validation BAC only. Per task:

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.

@mahirjain01

Copy link
Copy Markdown
Contributor Author

One NeuralBench detail that may matter

If the replication feeds AXON NeuralBench's channel_positions: for datasets without layout_or_montage_name, these come from each recording's stored locations and are not always in the head frame. For Miltiadous2023Dice they match MNE's raw standard_1020 coordinates, about 5 cm from the head frame. On NeuralBench's dementia task this changed AXON's fine-tuned test BAC from 41.3 ± 4.0 to 48.3 ± 0.1 (3 seeds) when we used head-frame positions looked up from the channel names instead.

Resolved docs/whats_new.rst by keeping both entries.
@mahirjain01

Copy link
Copy Markdown
Contributor Author

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 test_foundation_models.py and test_return_features.py. I am happy to resolve it whenever that is convenient for you.

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.

@bruAristimunha bruAristimunha left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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!

@bruAristimunha
bruAristimunha merged commit a619f4b into braindecode:master Oct 9, 2026
16 checks passed
bruAristimunha added a commit to lindicaphxag-tech/braindecode that referenced this pull request Oct 9, 2026
…raindecode#1253) into EEG-CLIP; EEGCLIP text side in _UNUSED_IN_FORWARD
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

model Adds a new model needs-replication Model PR: paper number must be replicated (NeuralBench) before merge

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants