Repository navigation
Add BrainOmni + BrainTokenizer (unified EEG/MEG foundation model) - #1043
bruAristimunha wants to merge 70 commits into
Conversation
Adds _BrainSensorModule, _ForwardSolution, _BackwardSolution, _BrainTokenizerEncoder, and _BrainTokenizerDecoder to brainomni.py, faithfully ported from upstream BrainOmni/model_utils/module.py with _ prefixes and upstream state-dict key alignment preserved.
Adds scripts/convert_brainomni_weights.py to convert OpenTSLab/BrainOmni upstream .pt checkpoints into braindecode hub format (config.json + safetensors). Includes exhaustive key-coverage assertions, geometry-buffer awareness (pos/sensor_type), round-trip validation via save_pretrained / from_pretrained, optional numerical parity mode, and an offline --self-test that validates the full remap logic without network access.
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: b8ad7899dd
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
| coil = str(ch.get("coil_type", "")) | ||
| if "PLANAR" in coil or "GRAD" in coil: | ||
| return _SENSOR_GRAD |
There was a problem hiding this comment.
Classify integer-coded gradiometers correctly
When chs_info comes from real MNE objects, MEG channels usually have integer kind and integer coil_type values; planar gradiometers are not passed as kind="grad". Converting coil_type to a string here means values like 3012 never contain "PLANAR" or "GRAD", so those gradiometers are classified as magnetometers and _orientation_of uses loc[9:12] instead of the grad orientation in loc[3:6], giving incorrect sensor embeddings for MEG data.
Useful? React with 👍 / 👎.
There was a problem hiding this comment.
Pull request overview
Adds BrainOmni and BrainTokenizer foundation models to braindecode, including model registration, documentation, tests, and a maintainer conversion script to remap upstream BrainOmni checkpoints into braindecode’s Hub format.
Changes:
- Introduces
braindecode.models.BrainTokenizer(VQ-VAE tokenizer) andbraindecode.models.BrainOmni(downstream classifier) with EEG/MEG geometry handling. - Adds extensive unit tests for geometry derivation, core blocks, tokenizer behavior, and BrainOmni classifier behavior.
- Adds a conversion script for upstream checkpoints and updates docs/model registry/CI integration expectations (compile/TorchScript skips).
Reviewed changes
Copilot reviewed 9 out of 9 changed files in this pull request and generated 4 comments.
Show a summary per file
| File | Description |
|---|---|
braindecode/models/brainomni.py |
New BrainOmni + BrainTokenizer implementations, geometry derivation, attention/SEANet/VQ components |
braindecode/models/__init__.py |
Exposes new models at braindecode.models.* |
braindecode/models/util.py |
Registers init-parameter expectations + marks BrainTokenizer as non-classification output |
braindecode/models/summary.csv |
Adds model-zoo metadata entries for BrainOmni/BrainTokenizer |
test/unit_tests/models/test_foundation_models.py |
Adds unit tests for BrainOmni/BrainTokenizer + helpers |
test/unit_tests/models/test_integration.py |
Skips torch.compile/TorchScript for BrainOmni(/Tokenizer) with rationale |
scripts/convert_brainomni_weights.py |
New conversion/validation script for upstream checkpoints |
docs/api.rst |
Adds BrainOmni/BrainTokenizer to API docs listings |
docs/whats_new.rst |
Changelog entry announcing the new models |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| coil = str(ch.get("coil_type", "")) | ||
| if "PLANAR" in coil or "GRAD" in coil: | ||
| return _SENSOR_GRAD | ||
| return _SENSOR_MAG |
| Pre-trained weights for ``BrainTokenizer`` are available on | ||
| `HuggingFace <https://huggingface.co/OpenTSLab/BrainOmni>`_ | ||
| (see the repository for the *tiny* and *base* tokenizer checkpoints). | ||
| Load with ``BrainTokenizer.from_pretrained("OpenTSLab/BrainOmni")``. |
…te, terser BrainOmni
…dows/... instead of B/C/N)
Resolve conflicts by keeping both sides' additions (BrainOmni/BrainTokenizer and EEGDINO) in summary.csv, util.py, whats_new.rst, and test_modules.py. Generalize PatchTokenizer with overlap/stride support and dynamic forward-time padding, and reuse it for BrainTokenizer._unfold.
- EEGModuleMixin: render a braindecode model card (description, usage, citation) on save_pretrained/push_to_hub, and warn when a strict=False load leaves missing/unexpected keys (the silent mis-load on the from_pretrained path). - RotaryPositionalEmbedding: clarify that n_dim is the full attention dim (load-bearing for checkpoint parity), not the per-head dim. - quantization: credit lucidrains vector-quantize-pytorch (MIT) and note the upstream-checkpoint state-dict parity constraint.
…onventions) Apply the braindecode model-convention review to BrainOmni/BrainTokenizer: - Docstrings (B8): canonical "<Name> from Xiao et al. (2025) [brainomni]_" headers + ".. versionadded:: 1.6.1". - Architecture docs (B8b): add Architecture Overview / Macro Components (Model.attr + Operations/Role) / Temporal-Spatial-Spectral / Additional Mechanisms rubric sections to both classes. - Weight init (B9): shared _init_weights applied via self.apply, faithful to upstream (Linear/Embedding -> trunc_normal_ std=0.02, RMSNorm -> 1.0); weight_norm convolutions left untouched. - Pretrained weights (B10): point from_pretrained at the converted braindecode/brainomni-pretrained and braindecode/braintokenizer-pretrained repos (raw upstream OpenTSLab keys, decoder.* vs final_layer.*, would silently mis-load under strict=False). - Signature (B3/A2): "# braindecode parameters" / "# model-specific parameters" comment separators + "*" keyword-only model params; del all six signal args. - Correct the frozen-backbone description: there is no train() override; the tokenizer is hard-frozen via torch.no_grad in BrainTokenizer.tokenize. Tests: 163 passed, 9 skipped (unchanged from baseline).
# Conflicts: # braindecode/models/summary.csv # braindecode/models/util.py # braindecode/modules/blocks.py # test/unit_tests/models/test_foundation_models.py # test/unit_tests/models/test_integration.py # test/unit_tests/models/test_modules.py
|
I completed a maximum-scope and paper-gate audit of this draft. The exact live The smallest defensible replacement stack is:
The qualifying target is BrainOmni-tiny, PhysioNet-MI Table 2 balanced accuracy Please drop the global base Audited head: |
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1043 +/- ##
==========================================
+ Coverage 87.78% 87.96% +0.17%
==========================================
Files 149 151 +2
Lines 17461 18324 +863
==========================================
+ Hits 15329 16119 +790
- Misses 2132 2205 +73 🚀 New features to boost your workflow:
|
Resolve the six conflicts with master 55afc12 by keeping both sides: - models/__init__.py: import and export BrainOmni/BrainTokenizer next to BrainBERT/Brant. - models/util.py: keep the BrainOmni/BrainTokenizer mandatory-parameter entries next to DIVER1, keep the BrainOmni geometry helpers next to master's resolve_channel_indices, and take master's positions_from_chs_info signature. - docs/api.rst: list BrainOmni after BrainBERT. - docs/whats_new.rst: master's entries followed by the BrainOmni entry. - test_return_features.py: sorted imports (BaRISTA, BrainBERT, BrainOmni, Brant). - test_foundation_models.py: union of imports (master dropped reve.RMSNorm, unused here), master's new DIVER1/ZUNA/STEEGFormer tests followed by the BrainOmni block.
…dBlock)
- RotaryPositionalEmbedding: rotate pairs with functional.rotate_pairs in
float32 instead of a cached complex buffer. No buffers, so no cast or
state-dict repair logic; the per-head frequency split is unchanged.
- The loader drops the released RoPE keys: the export stores freqs rounded
to bfloat16 and the complex cache without its sine part. The port now uses
the exact float32 rotation the upstream model computes at construction.
The tiny-checkpoint signature is updated to that upstream path.
- Replace the private _FeedForward with modules.FeedForwardBlock (same
Linear-SELU-Linear-Dropout maths; key remap ff.layer.{0,2} -> ff.{0,3}).
- license="mit" on both classes, _set_n_outputs in reset_head, and the
per-file MIT entries in NOTICE.txt; LICENSES/ is kept.
Random-weight and released-checkpoint parity against OpenTSLab/BrainOmni
340d6b5a: encoder max |diff| <= 5.1e-7, tokenizer 0, 0 index mismatches.
….decoder rename Review (REVIEW.md) MINOR: the ff/aggregate_mlp renaming and the dropped RoPE and pretraining keys were only covered by network-marked tests. NIT: the tokenizer.decoder. rename matched anywhere in a key.
The first forward of a fresh BrainOmni runs K-means over 4096 x 512 x 256 broadcast differences (~2 GB per intermediate); test_local_push_and_pull_roundtrip peaked at 5.2 GB RSS and the ubuntu CI runners died with a shutdown signal in that test twice. Chunking the samples keeps every distance bit-identical (checked) and peaks at 1.5 GB.
Summary
Adds two models derived from BrainOmni: A Brain Foundation Model for Unified
EEG and MEG Signals (Xiao et al., NeurIPS 2025;
arXiv:2505.18185; upstream
OpenTSLab/BrainOmni, MIT):
BrainTokenizer— the SEANet/channel-compression/residual-VQ tokenizer,registered as a non-classification model.
BrainOmni— the downstream classifier: frozen tokenizer inference,factored spatio-temporal transformer, released mean-pooling path, and
classification head.
Key contracts:
checkpoints. Immutable artifact revisions and SHA-256 values are pinned in
the integration tests; no checkpoint is copied or re-hosted by this PR.
released numerical signatures. The deterministic RoPE cache is regenerated
from its authoritative frequency buffer because the DeepSpeed export lost
its complex phase.
chs_infofor EEG, MAG, planar GRAD, and axialGRAD. Real MNE types use public
mne.channel_type; lightweightch_typemetadata and real-MNE metadata both round-trip through config and local Hub
serialization with exact state/output parity.
Distributed EMA initialization and sufficient statistics are synchronized
across ranks.
tokenizer's prior mode; projection, transformer blocks, and head remain
trainable during downstream fine-tuning.
PatchTokenizerAPI remains valid, short and oddSEANet cases follow the source padding policy, invalid public configurations
fail early, and the implementation remains compatible with the declared
PyTorch 2.4 floor, using native RMSNorm directly.
The released source preprocessing differs from the paper's prose at its final
normalization step. Model documentation records that discrepancy and tells
users to reproduce the released source pipeline when using released weights.
Provenance and licensing
340d6b5aba886af76b217272cdb3251651e9bf16.9a4d3c70495370397ccfbfd6d2496f25647545a5.LICENSES/BrainOmni-MIT.txtpreserves the OpenTSLab and Meta/EnCodec MITgrants;
LICENSES/vector-quantize-pytorch-MIT.txtpreserves the additionalderived quantizer grant;
NOTICE.txtrecords both components.MIT files are packaged. No external weight is bundled.
The official OpenTSLab repositories contain raw
.ptstate dictionaries, notBraindecode
config.jsonplus safetensors repositories. The documentationtherefore shows
from_opentslab_config(...)followed by strictload_state_dict(...); it does not claim the upstream repository works withBraindecode
from_pretrained(...).Verification
models; tokenizer/tiny numerical signatures match the pinned source.
update.
Hub round trips for both public models.
compile/export, reset-head, and invalid-boundary suites.
packaged-license gates.
c8b01ab478220a2ea8460bbc064c6f1bd61753c1found no actionable issue.
The PR remains a draft while fresh exact-SHA CI runs. Publishing converted
Braindecode-format weights, if desired later, is a separate maintainer action
and is not part of this PR.