Repository navigation
Add BaRISTA iEEG model - #1173
Conversation
There was a problem hiding this comment.
💡 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".
| torch.Tensor | ||
| Rotated tensor, of the shape and dtype of ``x``. | ||
| """ | ||
| return x * cos.to(x.dtype) + rotate_half(x) * sin.to(x.dtype) |
There was a problem hiding this comment.
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 👍 / 👎.
| if ( | ||
| positions is None | ||
| or len(positions) != self.n_chans | ||
| or not positions.isfinite().all() |
There was a problem hiding this comment.
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 👍 / 👎.
There was a problem hiding this comment.
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
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.
| pytest.param( | ||
| BaRISTA, | ||
| N_CHANS, | ||
| {"chs_info": _chs(), "patch_size": 200}, | ||
| False, |
There was a problem hiding this comment.
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
Open (3)
Resolved since last review (1)
There was a problem hiding this comment.
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
Open (5)
BaRISTA uses RMSNorm unavailable under supported PyTorch versions · New Add coverage for spatial_indices metadata validation and behavior · New Add coverage for non-default spatial scale paths and validation What's New links to reverted PR instead of implementing PR · New Fix Hub example to provide channel metadata for BaRISTA
Resolved since last review (1)
| 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}.") |
Codecov Report❌ Patch coverage is 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:
|
There was a problem hiding this comment.
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
Open (5)
BaRISTA uses RMSNorm unavailable under supported PyTorch versions Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation What's New links to reverted PR instead of implementing PR Fix Hub example to provide channel metadata for BaRISTA
There was a problem hiding this comment.
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
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 |
Co-authored-by: Bru <[email protected]>
There was a problem hiding this comment.
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
Open (7)
Constructor-provided indices are checked for finiteness and range, but indices supplied to… · New This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… · New Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::… · New Correct overstated documented parameter count Fix Hub example to provide channel metadata for BaRISTA
| pytest.param( | ||
| BaRISTA, | ||
| N_CHANS, | ||
| {"chs_info": _chs(), "patch_size": 200}, | ||
| False, | ||
| id="BaRISTA", | ||
| ), |
| - 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`_). |
There was a problem hiding this comment.
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
Open (21)
pooch.os_cache()returns a string in the supported Pooch versions, so applying/to it raises… Align learned pooling artifact and conversion reporting · New The generated files omitfinal_layerthroughwithout_head, so a normal strict state-dict load… Caller-suppliedspatial_indicesare used without moving them to the model's device. A common GPU… When--pooling learnedis selected, the generated model still uses the same README and repository… The parameter count does not match the implementation for the documented constructor. With 22… Indices supplied toforwardbypass the range validation performed by_build_spatial_embedding.… Per-recordingspatial_indicesare used without moving them to the model's device. Whenxand… Per-recording indices supplied toforwardbypass the range checks performed by… The added generic return-feature case exercises only the default coordinate montage. The new… Constructor-provided indices are range-checked, but indices supplied toforwardreach… Forward-suppliedspatial_indicesis not moved to the model's device. Unlikedefault_indices, it… The added generic test exercises only the default coordinate configuration with fixed learned… This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… Constructor-provided indices are checked for finiteness and range, but indices supplied to… Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation Correct the typo in the notice fromwritentowrittenand remove the suppression comment. If… Correct the spelling of 'writen' to 'written' and remove the corresponding codespell suppression. This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::…
And 1 more that still need to be addressed.
Resolved since last review (1)
- 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.
There was a problem hiding this comment.
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
Open (13)
pooch.os_cache()returns a string in the supported Pooch versions, so applying/to it raises… The generated files omitfinal_layerthroughwithout_head, so a normal strict state-dict load… The parameter count does not match the implementation for the documented constructor. With 22… The added generic return-feature case exercises only the default coordinate montage. The new… The added generic test exercises only the default coordinate configuration with fixed learned… This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation Update PR description to reflect the added BaRISTA test module · New Correct the typo in the notice fromwritentowrittenand remove the suppression comment. If… Correct the spelling of 'writen' to 'written' and remove the corresponding codespell suppression. This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::… Correct overstated documented parameter count
Resolved since last review (9)
Align learned pooling artifact and conversion reporting Caller-suppliedspatial_indicesare used without moving them to the model's device. A common GPU… When--pooling learnedis selected, the generated model still uses the same README and repository… Indices supplied toforwardbypass the range validation performed by_build_spatial_embedding.… Per-recordingspatial_indicesare used without moving them to the model's device. Whenxand… Per-recording indices supplied toforwardbypass the range checks performed by… Constructor-provided indices are range-checked, but indices supplied toforwardreach… Forward-suppliedspatial_indicesis not moved to the model's device. Unlikedefault_indices, it… Constructor-provided indices are checked for finiteness and range, but indices supplied to…
- 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.
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
Moderate issues remain with coordinate-frame validation, registry configuration, and the stale test-description statement.
Review effort: Lite
Findings: 1
Open (12)
pooch.os_cache()returns a string in the supported Pooch versions, so applying/to it raises… The parameter count does not match the implementation for the documented constructor. With 22… The added generic return-feature case exercises only the default coordinate montage. The new… The added generic test exercises only the default coordinate configuration with fixed learned… This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation Update PR description to reflect the added BaRISTA test module Correct the typo in the notice fromwritentowrittenand remove the suppression comment. If… Correct the spelling of 'writen' to 'written' and remove the corresponding codespell suppression. This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::… Correct overstated documented parameter count
Resolved since last review (1)
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.
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
Unresolved moderate findings remain for variable-window test coverage and parameter-count metadata, plus a release-version directive nit.
Review effort: Lite
Findings: 6
Open (11)
The parameter count does not match the implementation for the documented constructor. With 22… The added generic return-feature case exercises only the default coordinate montage. The new… The added generic test exercises only the default coordinate configuration with fixed learned… This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation Update PR description to reflect the added BaRISTA test module Correct the typo in the notice fromwritentowrittenand remove the suppression comment. If… Correct the spelling of 'writen' to 'written' and remove the corresponding codespell suppression. This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::… Correct overstated documented parameter count
Resolved since last review (1)
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
Resolve the coordinate-frame validation issue in barista.py before approval.
Review effort: Lite
Findings: 6
Open (8)
The parameter count does not match the implementation for the documented constructor. With 22… The added generic return-feature case exercises only the default coordinate montage. The new… The added generic test exercises only the default coordinate configuration with fixed learned… This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::… Correct overstated documented parameter count
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
Coordinate-frame validation, the generated Hub example, and parameter-count metadata need correction.
Review effort: Lite
Findings: 6
Open (8)
The parameter count does not match the implementation for the documented constructor. With 22… The added generic return-feature case exercises only the default coordinate montage. The new… The added generic test exercises only the default coordinate configuration with fixed learned… This only adds BaRISTA to the shared return-features contract test. The BaRISTA-specific spatial… Add coverage for spatial_indices metadata validation and behavior Add coverage for non-default spatial scale paths and validation This entry is underCurrent 1.8.1, while the BaRISTA class docstring declares `.. versionadded::… Correct overstated documented parameter count
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.
* 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]>



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 reusesPatchTokenizer,GatedLinearUnitand 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 revision83b27375eba60e9eba9da4e7dd8fb283baace376, 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 throughBaRISTA.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:
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.