Repository navigation
Adding ZUNA to Braindecode - #1020
Conversation
…ntions ZUNA worked CUDA-only because flex_attention is unstable for autograd on CPU; this collapses the whole API onto SDPA and trims the inference port to a maintainable size. Model (braindecode/models/zuna.py, 861 -> 400 lines): - Drop flex_attention path, BlockMask, _create_document_mask, packed multi-document layout, and the use_flex / attention_impl plumbing. SDPA per-batch is mathematically equivalent (each sample is its own document) and supports CPU + fp16 + autograd. - Replace custom _RMSNorm with torch.nn.RMSNorm. - Drop _ZUNAEncoderArgs dataclass, _build_zuna_encoder_args validation, InitStdFactor enum, and the six zero-arg config helpers in favour of module constants plus one patchable _encoder_config(). - Drop unused dataclass knobs: init_base_std, init_std_factor, encoder_hidden_dim, n_kv_heads, ffn_dim_multiplier, multiple_of, dropout_type / dropout_vec, encoder_latent_downsample_factor, tok_idx_type. - Drop _apply_channel_mask + dropped_channels API (caller-side concern). - Drop _discretize_channel_positions string extremes_type; inline the constant ZUNA_POS_HALF_RANGE = 0.12. - Drop _repeat_kv / GQA plumbing (never used: n_kv_heads always equals n_heads in the published config). - Cache positions from chs_info at construction; resolve once. - Replace 4-lookup + cat RoPE gather with one indexed read + flatten. - Bucket positions in fp32 so fp16 inference doesn't shift bucket edges. - Distinguish 'no coords' vs 'names without montage' errors. - Split forward shape validation so ndim / n_chans / n_times mismatches get targeted messages. Param count matches Zyphra/ZUNA published checkpoint exactly (172,069,668 params for n_chans=22, n_outputs=4). Tests (test/unit_tests/models/test_models.py, 290 -> 151 lines): - Drop _zuna_forward inference_mode wrapper - no longer needed. - Drop @pytest.mark.skipif(not zuna.HAS_FLEX) on every test. - Drop redundant channel_mask / dropped_channels tests with the API. - Drop _zuna_published_config snapshot test (superseded by inline constants). - Parametrize n_times rejection over (1279, 1281). - Add test_zuna_requires_montage_when_names_only. - Drop CUDA-only skips in test_integration.py for ZUNA.
…ndecode style Follow the BrainOmni/REVE convention where every architecture knob is a documented __init__ argument instead of frozen module constants. - ZUNA.__init__ now takes dim, n_layers, n_heads, head_dim, fine_time_pts, latent_dim, max_seqlen, rope_theta, rope_dim, pos_bins, pos_half_range, norm_eps, drop_prob and activation (defaults reproduce the published Zyphra/ZUNA config); drop the _encoder_config() indirection. - Remove the ZUNA_* module constants; values are now literal defaults. - Apache-2.0 header (upstream is Apache-2.0, not BSD-3) + adaptation credit. - Ship a torch>=2.0-compatible _RMSNorm (nn.RMSNorm needs torch>=2.4). - reset_head keeps get_config()/hub config in sync. - Tests configure a small model via constructor kwargs; un-skip the activation/drop_prob/export checks; fix stale flex_attention comments.
|
I am mostly struggling to get some results a little above the change level with this model @jonathanhuml, any input is welcome. I don't know what to test more. |
Hey @bruAristimunha! I appreciate you testing this out, sorry it's been a pain. The 5 second trial limitation is definitely a major constraint here, we are in the midst of a new training run that relaxes this assumption to arbitrary trial lengths (up to 30 seconds) and the preprocessing in general, as well as diversifying the tasks. Could you give me some time (at least a week) to chop the new version up into braindecode's API? We are seeing much better results in Meta's new NeuralBench repo with this updated version. This seems similar to Open EEG Bench so hopefully that should generalize. I will send you a message on Discord next week to update you on my progress if that works |
|
Works for me @jonathanhuml! |
Thank you @bruAristimunha! |
Resolve conflicts from new-model additions on master: - summary.csv: adopt master schema (new Modality column, Prediction rename, added models) and re-add the ZUNA row with Modality=EEG - util.py: keep both master's EEGDINO/DANCE signal-param entries and ZUNA's - test_integration.py: keep both ZUNA and STEEGFormer in not_working_models - test_models.py: keep both 'inspect' and 'warnings' imports The stale branch was failing CI on a pytest collection error (IdMaker.__init__ arity); merging master picks up the fix.
The changelog-check CI was red: a new user-visible model needs a whats_new entry and ZUNA had none. Add the entry plus the missing 'Jon Huml' contributor link target.
Re-host the upstream Zyphra/ZUNA encoder checkpoint at braindecode/ZUNA (a bit-identical mirror, apache-2.0) so the library has a stable, permanent location instead of depending on the third-party repo. - ZUNA.from_pretrained now defaults pretrained_model_name_or_path to 'braindecode/ZUNA' (weights normalised to the standard model.safetensors, so no filename argument is needed) - Only the encoder is pretrained; the classification head stays randomly initialised and must be fine-tuned (documented) - Add TestZUNAPretrained to the pretrained-hub integration suite - Update the changelog entry
Avoids the signature divergence from the parent from_pretrained (Codacy / pylint arguments-differ) while keeping identical behaviour: with no path given, it defaults to the braindecode/ZUNA re-host; any explicit repo id or local path (positional or keyword) is passed straight through.
There was a problem hiding this comment.
Pull request overview
Adds the ZUNA EEG foundation-model encoder to braindecode.models, including dynamic channel-position handling in forward, a Braindecode classification head, Hugging Face Hub loading defaults, and accompanying docs/test coverage.
Changes:
- Introduce the new
braindecode.models.ZUNAmodule (encoder + classification head, position-aware RoPE,from_pretraineddefaulting tobraindecode/ZUNA). - Add unit/integration tests covering ZUNA construction, forward paths, feature returns, montage-based name resolution, and Hub checkpoint loading.
- Register and document ZUNA in the public models API, summary tables, and release notes.
Reviewed changes
Copilot reviewed 9 out of 10 changed files in this pull request and generated 5 comments.
Show a summary per file
| File | Description |
|---|---|
braindecode/models/zuna.py |
New ZUNA model implementation (encoder-only feature extractor + classification head) with position/name resolution and Hub loading. |
braindecode/models/__init__.py |
Exports ZUNA from the models package. |
braindecode/models/util.py |
Adds ZUNA to the model config/mandatory-parameter registry. |
braindecode/models/summary.csv |
Adds ZUNA entry to the model summary table. |
test/unit_tests/models/test_models.py |
Adds unit tests for ZUNA API surface and forward/feature behaviors. |
test/unit_tests/models/test_integration.py |
Excludes ZUNA from torch.compile and TorchScript integration tests with documented reasons. |
test/integration_tests/test_pretrained_hub_models.py |
Adds an integration test that loads ZUNA from the Hub and runs a forward pass. |
docs/api.rst |
Adds ZUNA to the generated API docs listing. |
docs/whats_new.rst |
Adds release notes entry for ZUNA and author link. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
|
I'm coming back to this @bruAristimunha! Sorry for the delay here, we had a lot of things to catch up on with our model release. |
- Collapse duplicate window/shape validations into single checks, keeping the error messages the tests pin. - Replace the manual _repeat_kv expand+reshape with SDPA's native enable_gqa (torch>=2.5 needed only for GQA configs; the published ZUNA1.1 config has n_kv_heads == n_heads and is unaffected). - Compress reset_head config sync into a loop and shorten comments.
Replace raw view/reshape/transpose/repeat_interleave chains with named einops rearrange/repeat patterns (head split/merge, register-token interleave, 4D RoPE table gather, patch tokenisation), and use the einops Rearrange layer instead of nn.Flatten in final_layer, matching the rest of the model zoo. No parameter or state-dict key changes.
Model: - Reject rope_dim != 4 at construction instead of crashing in forward. - Drop partial/non-finite chs_info positions so forward falls back to channel_names instead of silently corrupting RoPE buckets. - Guard torch.compiler.is_compiling behind a helper (torch 2.0 floor). - Forward **kwargs and raise when no upstream key matches in load_state_dict; remove the dead _apply_upstream_config machinery. - RMSNorm casts to the weight dtype so model.half() works. - Bound the tok_idx cache and key it by montage source. - Replace the matrix rotary table with native torch.polar / view_as_complex (numerics verified identical). - Inline the module-level constants into the code that uses them. Tests: move the ZUNA block from test_models.py into a dedicated test_zuna.py, dropping signature/import tautologies covered by the shared parametrized suites, and add coverage for the fixes above.
|
Hey @jonathanhuml, once everything is green! I am happy to merge, many thanks for all the effort here! |
SUMMARY
We propose to add ZUNA, a masked diffusion autoencoder trained to perform masked channel infilling and superresolution for arbitrary electrode numbers and positions in EEG signals.
While the original encoder-decoder model has 380M parameters, this port is a feature extractor that only exposes the latents and does not perform reconstruction. The implementation also includes basic support for channel masking and dropped-channel inference, either by channel index or by channel name. This is intended to preserve some of the practical behavior expected from a model trained around masked channel infilling.
By discarding the decoder, the total model size is about 170M parameters. We are totally happy to add the decoder back in later: let us know what BrainDecode could most benefit from! We did this to keep the file as lightweight and readable as possible. We are currently training a new version and will likely submit another PR soon, so we can definitely integrate this into the next version if so desired.
We have tried to keep as close as possible to the requirements and base.py structure in Braindecode. We have two main design decisions that would probably be helpful to highlight for the Braindecode team for any potential feedback:
ZUNA currently depends on PyTorch
flex_attention. This is only available in PyTorch>=2.5, while Braindecode currently supports PyTorch>=2.0. To avoid making ZUNA break imports for users on older PyTorch versions, this PR uses a soft import pattern similar to the Hugging Face integration. As a result, Braindecode can still be imported normally, but instantiating or running ZUNA requires a compatible PyTorch version. ZUNA could theoretically support other attention variants at the cost of less efficient GPU usage, but we have implemented the flex-only variant for nowZUNA is montage-agnostic but requires 3D channel positions during the forward pass. These positions are not inherent to the model weights, so this PR allows users to provide them dynamically rather than binding them at model initialization. When
chs_infocontains the necessary position metadata, we also support extracting the positions from there.