Skip to content

Add BaRISTA iEEG model - #1173

Merged
bruAristimunha merged 17 commits into
masterfrom
reopen-1171-barista
Sep 29, 2026
Merged

bruAristimunha merged 17 commits into
masterfrom
reopen-1171-barista

Conversation

@bruAristimunha

@bruAristimunha bruAristimunha commented Sep 21, 2026 •

Copy link
Copy Markdown
Collaborator

Adds BaRISTA, an intracranial EEG encoder/classifier with temporal patch tokenization and joint space-time rotary attention. Restores @julien-gadonneix's contribution from #1171 after its accidental merge and revert in #1172.

Spatial indices come from dataset metadata. The merged nemarDatasets/nm000253#3 supplies Brain Treebank's coordinate, parcel and lobe indices. The model contains no atlas-name tables or runtime dataset downloads. MNE-coordinate binning is an optional fallback; released weights should use the dataset's indices.

Includes #1175: callers can supply each recording's indices to forward; mean pooling accepts different channel counts and window lengths. Learned pooling requires a fixed total token count. Spatial tables and optional defaults are built by one helper. The model reuses PatchTokenizer, GatedLinearUnit and native RMSNorm; #1174 supplies the PyTorch >= 2.4 minimum.

The conversion script lives on the Hub, next to each converted encoder (coords, parcels, lobes — convert_barista_weights.py), not in this repository. It downloads the three official checkpoints from pinned source revision 83b27375eba60e9eba9da4e7dd8fb283baace376, verifies SHA-256 hashes, renames tensors, fuses the gated projections, checks encoder tokens against the released forward equations and writes the Hub directories, notices and a JSON report. Every learned encoder tensor must load; only downstream pooling/classification tensors may be missing. The models load through BaRISTA.from_pretrained(...) with non-strict loading: the releases do not include the downstream pooling and classification layers, which stay randomly initialised.

Adds the authors' original overview figure and rewrites the model documentation around its inputs, conversion steps and validation limits. License. BaRISTA coverage is integrated into the existing test_models.py; no separate test module is added.

Validation:

  • All three releases: 149 learned encoder tensors loaded for coordinates, 147 each for parcels and lobes. CPU float32 maximum token differences from the independent equation check were 4.41e-6, 8.52e-6 and 6.20e-6 on PyTorch 2.4.0 (atol=1e-5, rtol=1e-4).
  • The actual upstream model classes produced identical tokens after explicit CPU adaptations for xformers attention and CUDA-only spatial allocation. This does not validate the CUDA xformers kernel.
  • PyTorch 2.6.0 also passed with mean pooling, two channels and 1,025 samples. Strict local Hub reloads reproduce logits exactly. Missing encoder tensors, unexpected tensors and changed rotary frequencies are rejected.
  • Spatial helper simplification: eight configurations matched the previous implementation exactly; variable montages, scripting, gradients and twelve invalid-input cases passed. Dataset-index config/Hub/script/export checks passed.
  • Three existing BaRISTA feature tests, all configured pre-commit checks, standalone converter lint and a focused Sphinx documentation build passed.
  • Applied avoid-ai-writing and tom-neurips-review: removed repetitive Operations/Role phrasing, corrected the fixed-token-count description, and made the checkpoint validation limits explicit.

The original releases contain no downstream head. Converted heads are newly initialized and need fine-tuning. No raw dataset evaluation, downstream accuracy, GPU or mixed-precision equivalence is claimed; masked pretraining is outside this implementation.

Latest test consolidation: test_models.py — 754 passed, 8 skipped; pre-commit checks passed.

Copilot AI lite review requested due to automatic review settings September 21, 2026 14:47

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: 9c21bc434f

ℹ️ 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".

Comment thread braindecode/models/barista.py Outdated
Comment thread braindecode/functional/functions.py Outdated
torch.Tensor
Rotated tensor, of the shape and dtype of ``x``.
"""
return x * cos.to(x.dtype) + rotate_half(x) * sin.to(x.dtype)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Move rotary tables onto the input device

When this newly exported helper is used with a CUDA, MPS, or other accelerator tensor, directly passing the tables returned by rotary_positional_encoding fails with a device-mismatch error: those tables are created on CPU, and cos.to(x.dtype)/sin.to(x.dtype) change only their dtype. Convert them to both x.device and x.dtype so the two public helpers compose correctly outside BaRISTA's registered-buffer path.

Useful? React with 👍 / 👎.

Comment thread braindecode/models/barista.py Outdated
Comment on lines +684 to +687
if (
positions is None
or len(positions) != self.n_chans
or not positions.isfinite().all()

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Reject zero locations on individual channels

When one channel has an all-zero loc while another channel has a valid position, extract_channel_locations_from_chs_info accepts the array because it rejects zeros only when the entire stacked array is zero, and this condition then checks only length and finiteness. The missing channel is consequently treated as lying at the head origin and receives a learned coordinate embedding, silently corrupting the intended spatial encoding instead of producing the documented missing/degenerate-location error.

Useful? React with 👍 / 👎.

@bruAristimunha bruAristimunha changed the title Add BaRISTA model (restore #1171 for review) Add BaRISTA iEEG model Sep 21, 2026

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

Head resetting can break device/config consistency, and the third-party license notice requires fuller attribution.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 2 Medium severity · 1 Low severity

Open (3)
What changed in this PR

Adds the BaRISTA iEEG foundation model, including spatial embeddings, temporal encoding, rotary attention, and public API integration.

Changes:

  • Adds BaRISTA architecture and model registration.
  • Adds GLU feed-forward and rotary positional encoding utilities.
  • Updates documentation, licensing, catalog metadata, and smoke tests.
File Description
braindecode/​models/​barista.py BaRISTA implementation
braindecode/​models/​util.py Model test/config registration
braindecode/​functional/​functions.py Rotary encoding helpers
braindecode/​modules/​blocks.py GLU feed-forward block
braindecode/​models/​__init__.py Public model export
braindecode/​modules/​__init__.py Public module export
braindecode/​functional/​__init__.py Functional exports
test/​unit_tests/​models/​test_return_features.py BaRISTA smoke coverage
NOTICE.txt USC license attribution
docs/​whats_new.rst Release note
docs/​api.rst API documentation entries
braindecode/​models/​summary.csv Model catalog entry

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread braindecode/models/barista.py Outdated
Comment on lines +81 to +85
pytest.param(
BaRISTA,
N_CHANS,
{"chs_info": _chs(), "patch_size": 200},
False,
Comment thread NOTICE.txt Outdated
Copilot AI review requested due to automatic review settings September 21, 2026 15:07

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

The inherited Hub construction example omits required chs_info for BaRISTA’s coordinate-based default.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 Medium severity · 2 Low severity

Open (3)
Resolved since last review (1)

Comment thread braindecode/models/barista.py
Copilot AI review requested due to automatic review settings September 21, 2026 15:27

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

Native RMSNorm depends on the unmerged PyTorch 2.4 requirement change, and spatial-index behavior lacks focused committed tests.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 High severity · 2 Medium severity · 2 Low severity

Open (5)
Resolved since last review (1)

Comment thread braindecode/models/barista.py
Comment thread braindecode/models/barista.py Outdated
Comment on lines +373 to +376
indices = torch.as_tensor(spatial_indices)
expected = (self.n_chans, 3) if is_coords else (self.n_chans,)
if tuple(indices.shape) != expected:
raise ValueError(f"spatial_indices must have shape {expected}.")
Comment thread docs/whats_new.rst Outdated
@codecov

codecov Bot commented Sep 21, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 92.09302% with 17 lines in your changes missing coverage. Please review.
✅ Project coverage is 87.11%. Comparing base (f8c9918) to head (8b895a0).
⚠️ Report is 4 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1173      +/-   ##
==========================================
+ Coverage   86.98%   87.11%   +0.13%     
==========================================
  Files         144      146       +2     
  Lines       16570    16870     +300     
==========================================
+ Hits        14413    14697     +284     
- Misses       2157     2173      +16     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

Copilot AI review requested due to automatic review settings September 21, 2026 18:12

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🔵 Needs a closer look

It depends on an unmerged PyTorch minimum bump, generates an invalid default Hub example, and lacks focused tests for spatial-index modes.

Review effort: Lite
Findings: 1 High severity · 2 Medium severity · 2 Low severity

Open (5)

Copilot AI review requested due to automatic review settings September 21, 2026 19:07

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

The parameter count is incorrect and explicit spatial-index behavior lacks repository regression coverage.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 2 Medium severity · 2 Low severity

Open (4)
Resolved since last review (2)

Model,Application,Type,Sampling Frequency (Hz),Hyperparameters,#Parameters,get_#Parameters,Categorization,Modality
ATCNet,General,Prediction,250,"n_chans, n_outputs, n_times",113732,"ATCNet(n_chans=22, n_outputs=4, n_times=1000)","Convolution,Recurrent,Attention/Transformer",EEG
AttentionBaseNet,Motor Imagery,Prediction,250,"n_chans, n_outputs, n_times",3692,"AttentionBaseNet(n_chans=22, n_outputs=4, n_times=1000)","Convolution,Attention/Transformer",EEG
BaRISTA,General,"Prediction, Embedding",2048,"chs_info, n_outputs, n_times",870324,"BaRISTA(n_chans=22, chs_info=<user>, n_outputs=4, n_times=6144, sfreq=2048)","Attention/Transformer,Foundation Model,Channel",iEEG
Copilot AI review requested due to automatic review settings September 21, 2026 19:13

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Warning

Copilot couldn't run its full agentic review because it didn't start before the timeout. Make sure your repository has a runner available, or add a copilot-code-review.yml file specifying one with the runs-on attribute. See the docs for more details.

Copilot review overview

Review effort: Lite
Findings: 4 Medium severity · 3 Low severity

Open (7)

Comment thread braindecode/models/barista.py
Comment on lines +81 to +87
pytest.param(
BaRISTA,
N_CHANS,
{"chs_info": _chs(), "patch_size": 200},
False,
id="BaRISTA",
),
Comment thread docs/whats_new.rst
Comment on lines +31 to +36
- Add :class:`braindecode.models.BaRISTA`, an intracranial EEG foundation model
whose spatial encoding scale is a free choice: electrodes are tokenized
channel-wise and space enters as a single learned embedding selected by the
electrode coordinate, its atlas parcel or its lobe, before a joint
space-time transformer encoder with rotary temporal embeddings
(:gh:`1173` by `Julien Gadonneix`_).
Copilot AI review requested due to automatic review settings September 22, 2026 21:04

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

Unresolved moderate issues remain in device handling, test coverage, registry metadata, parameter counts, and checkpoint loading consistency.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 High severity · 16 Medium severity · 4 Low severity

Open (21)

And 1 more that still need to be addressed.

Resolved since last review (1)

Comment thread scripts/convert_barista_weights.py Outdated
- Move spatial indices passed to forward onto the embedding device and
  range-check them in eager mode (clear ValueError instead of IndexError;
  export/TorchScript/compile paths unchanged).
- Reorder the MNE coordinate fallback to the (left, inferior, posterior)
  columns of the released tables; negated RAS alone gave (L, P, I).
- Drop the converter's --pooling learned option, which saved a randomly
  initialised read-out: the releases carry no pooling weights.
- Docstring opening follows the "<Name> from <Author> et al" convention.
- Add test_barista.py: region scales, per-forward montages, invalid
  indices, the learned-pooling grid check and the coordinate order.
Copilot AI review requested due to automatic review settings September 24, 2026 22:05

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

Two moderate findings remain regarding logits validation and variable-window test coverage.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 High severity · 7 Medium severity · 5 Low severity

Open (13)
Resolved since last review (9)

Comment thread test/unit_tests/models/test_barista.py Outdated
bruAristimunha added a commit to julien-gadonneix/braindecode that referenced this pull request Sep 24, 2026
- Turn the MAPA_DKT_REGIONS string into Sphinx #: comments (the bare
  string failed check-docstring-first) and keep the api.rst entry on one
  line (docstrfmt).
- License header is Apache-2.0, matching the upstream code it transcribes.
- Changelog cites this PR (braindecode#1178), not braindecode#1173.
- Docstring: window normalization is one of two numerical departures from
  the reference (the per-window STFT with reflect padding is the other);
  normalization="none" feeds the raw STFT magnitude, not the reference's
  inputs; pooling="flatten" reads the normed four-tap concatenation, not
  the paper's pre-norm block-12 read-out.
Copilot AI review requested due to automatic review settings September 29, 2026 09:26

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

The converter now lives only in the braindecode/BaRISTA-{coords,parcels,lobes}
Hub repositories (synced with this PR's version); the docstring points there
and the license note is a single link.
Copilot AI review requested due to automatic review settings September 29, 2026 11:42

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Copilot AI review requested due to automatic review settings September 29, 2026 11:51

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Copilot AI review requested due to automatic review settings September 29, 2026 11:56

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

@bruAristimunha
bruAristimunha merged commit d6c562f into master Sep 29, 2026
14 checks passed
bruAristimunha added a commit to qinxwew/braindecode that referenced this pull request Sep 29, 2026
Resolve models/__init__.py and docs/whats_new.rst conflicts against master's newly-merged MIRepNet (braindecode#1146) and BaRISTA (braindecode#1173): keep both registry entries, restore alphabetical import/__all__ order, keep both changelog entries.
bruAristimunha added a commit that referenced this pull request Oct 6, 2026
* onboard MAPA

* fix: make MAPA pass pre-commit and correct its fidelity notes

- Turn the MAPA_DKT_REGIONS string into Sphinx #: comments (the bare
  string failed check-docstring-first) and keep the api.rst entry on one
  line (docstrfmt).
- License header is Apache-2.0, matching the upstream code it transcribes.
- Changelog cites this PR (#1178), not #1173.
- Docstring: window normalization is one of two numerical departures from
  the reference (the per-window STFT with reflect padding is the other);
  normalization="none" feeds the raw STFT magnitude, not the reference's
  inputs; pooling="flatten" reads the normed four-tap concatenation, not
  the paper's pre-norm block-12 read-out.

* enable single init and forward with different subjects

* MAPA: normalization='session' takes a session-normalized spectrogram

Co-authored-by: Cursor <[email protected]>

* allow to receive spectrogram bands directly

* fix: MAPA layout cache owns its indices; reuse functional.rotate_pairs

- _token_layout cached the caller's sensor-indices tensor by reference,
  so an in-place edit of that tensor compared equal to the cached copy and
  returned a stale layout. Store a clone; covered by
  test_mapa_metadata_cache_owns_its_snapshot.
- Replace the private _rotate_half with braindecode.functional.rotate_pairs,
  which master added for the shared rotary code (bitwise identical).
- Header: link the upstream repository next to the Apache-2.0 notice, as
  LUNA/ZUNA do.

* fix: MAPA.reset_head keeps get_config() in sync

Master's generic test_reset_head_updates_config /
test_reset_head_model_reloads_after_saving require every model with a
custom reset_head to call _update_init_kwargs, otherwise from_config and
from_pretrained rebuild the old head size.

* test: MAPA's released checkpoint reproduces the authors' features

A network-marked test downloads mapa_vits384.pt from the authors' Hub
repository at a pinned commit (sha256 equal to their GitHub release v0.1.0),
loads it with only the classification head missing, and compares the pooled
features of two deterministic raw windows, frontend included, with values
computed by the authors' code (bentang18/MAPA at bf2b49e).

* test: say where MAPA's checkpoint test fits its z-score and how to run it

The reference values fit the authors' robust z-score on each window, as
the default normalization="window" does; the docstring now says so and
notes that CI's unit-test jobs do not pass --run-network.

* refactor: trim MAPA docstrings and reuse braindecode blocks

Rewrite the class docstring in the layout of the other foundation models,
shorten the private helpers' docstrings, merge duplicated input checks, and
reuse rescale_parameter and FeedForwardBlock. FeedForwardBlock names its
layers 0-3, so a mapping renames the released checkpoint's mlp.fc1/fc2 keys
on load. Outputs and initialization are unchanged.

Co-authored-by: Cursor <[email protected]>

* MAINT simplify MAPA: inline constants, fold layout/parse/stft helpers, strict sfreq, fix reset_head device/dtype and rotary device

* MAINT move MAPA DKT region vocabulary into a shared models/util helper

* MAINT derive MAPA DKT region table from MNE's FreeSurfer LUT

util.dkt_region_slots() now derives the 74-slot DKT region table from
mne.read_freesurfer_lut() ids at call time instead of typing out the
parcel/structure names: the 62 cortical parcels come from excluding five
DK ids (unknown, corpuscallosum, bankssts, frontalpole, temporalpole)
out of the contiguous 1000-1035 block and sorting what remains
alphabetically per hemisphere; the 12 subcortical slots come from the
six aseg ids in 10-18 that have a bilateral Right- counterpart. The one
piece that cannot be derived from MNE -- the released slot order of
those six aseg structures, which is neither ascending-id nor
alphabetical order -- is kept as a tiny position permutation, not a
name list, documented as sourced from the released checkpoint. A test
freezes the exact previously-hard-coded 74-name tuple and asserts the
new helper reproduces it exactly.

---------

Co-authored-by: Bru <[email protected]>
Co-authored-by: Cursor <[email protected]>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants