Skip to content

Add BaRISTA model - #1171

Merged
bruAristimunha merged 1 commit into
braindecode:masterfrom
julien-gadonneix:add-barista-model
Sep 21, 2026
Merged

bruAristimunha merged 1 commit into
braindecode:masterfrom
julien-gadonneix:add-barista-model

Conversation

@julien-gadonneix

Copy link
Copy Markdown
Collaborator

Architecture verified against original codebase and paper.
Tests to load pretrained weights and reproduce results on the way.

@codecov

codecov Bot commented Sep 18, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 87.30769% with 33 lines in your changes missing coverage. Please review.
✅ Project coverage is 86.58%. Comparing base (f6144d3) to head (fe68be1).

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1171      +/-   ##
==========================================
+ Coverage   86.57%   86.58%   +0.01%     
==========================================
  Files         142      143       +1     
  Lines       16303    16563     +260     
==========================================
+ Hits        14114    14341     +227     
- Misses       2189     2222      +33     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@julien-gadonneix

Copy link
Copy Markdown
Collaborator Author

BaRISTA pretrained weight loading

It corresponds. The released checkpoint is a flat state dict of the pretraining model, and nothing but module naming separates it from the braindecode port: after renaming, every ported tensor matches a model tensor by name and shape, with zero left over. Coverage is 839,144 of 839,286 parameters — 100.0% to one decimal.

What gets renamed

Count Reference name Ported name
60 backbone.layers.N.attention.* backbone.layers.N.self_attn.*
24 tokenizer.temporal_encoder.feature_extractor.net.N.* temporal_encoder.blocks.N.*
1 tokenizer.temporal_pooler.final_layer.weight temporal_pooler.weight
1 tokenizer.spatial_encoder.subcomponent_embeddings.0.weight spatial_emb.tables.0.weight

What is new

  • final_layer.weight / final_layer.bias (2, 64), (2,) — the task head, sized by n_outputs, so necessarily absent from a pretraining checkpoint.
  • token_pooling.weight (1, 12) — the learned read-out over the 12 time patches, shaped by the window.

What is not used

12 tensors, one per transformer layer: backbone.layers.N.attention.rotary_emb.inv_freq. These are rotary position frequencies, which the port recomputes from the layer geometry rather than storing, so dropping them loses nothing.

@julien-gadonneix

Copy link
Copy Markdown
Collaborator Author

Paper result reproduction
image
Paper result on sentence onset classification: 0.853 (0.007 sem)
Reproduction: 0.910 (0.003 sem)

Difference may come from the fact that they enforce balanced classes over all, whereas in the reproduction, the classes are imbalanced but I use a training sampler to have balanced training batches and ensure that in test and validation splits, all classes are present.

@julien-gadonneix
julien-gadonneix marked this pull request as ready for review September 21, 2026 12:08
Copilot AI lite review requested due to automatic review settings September 21, 2026 12:08

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 was unable to review this pull request because the user who requested the review has reached their quota limit.

@bruAristimunha

Copy link
Copy Markdown
Collaborator

Not related to the issue with docs! Seems very good for me! many thanks @julien-gadonneix 🙏🏽

@bruAristimunha
bruAristimunha merged commit 16ee1d4 into braindecode:master Sep 21, 2026
11 of 12 checks passed
bruAristimunha added a commit that referenced this pull request Sep 21, 2026
@bruAristimunha

Copy link
Copy Markdown
Collaborator

My bad, I got confused with another model, I will open again the PR, and I will finish this one for you.

@julien-gadonneix
julien-gadonneix deleted the add-barista-model branch September 21, 2026 21:15
bruAristimunha added a commit that referenced this pull request Sep 29, 2026
* Restore BaRISTA model for review (#1171)

* Simplify BaRISTA and fix model contract edge cases

* Read BaRISTA spatial indices from dataset metadata

* Preserve shared class targets with unique MNE event IDs

* Link BaRISTA release note to the restored PR

* changes to allow for different n_chans (#1175)

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

* Simplify BaRISTA spatial setup and license note

* Add BaRISTA figure and verify conversion of released weights

* Publish the converted BaRISTA encoders to the Hub

Make the converter the single recipe that ships with the weights: it now
writes a model card and copies itself and NOTICE.txt into each Hub
directory, and --push-to uploads them through BaRISTA.push_to_hub.

The published encoders pool by mean so one file serves any montage and
window length, and the class docstring points at the three repositories.

* Publish BaRISTA encoders without the untrained head

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

* fix: validate BaRISTA forward indices and align the coordinate fallback

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

* BaRISTA: keep the conversion script on the Hub, link the license

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.

* Keep BaRISTA notice concise with upstream license link

* Integrate BaRISTA coverage into existing model tests

---------

Co-authored-by: Julien GADONNEIX <[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