Skip to content

Adding ZUNA to Braindecode - #1020

Merged
bruAristimunha merged 26 commits into
braindecode:masterfrom
jonathanhuml:zuna
Aug 11, 2026
Merged

bruAristimunha merged 26 commits into
braindecode:masterfrom
jonathanhuml:zuna

Conversation

@jonathanhuml

Copy link
Copy Markdown
Contributor

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 now
ZUNA 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_info contains the necessary position metadata, we also support extracting the positions from there.

jonathanhuml and others added 8 commits May 13, 2026 12:18
…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.
@bruAristimunha

Copy link
Copy Markdown
Collaborator

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.

@jonathanhuml

jonathanhuml commented Jun 24, 2026 •

Copy link
Copy Markdown
Contributor Author

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

@bruAristimunha

Copy link
Copy Markdown
Collaborator

Works for me @jonathanhuml!

@jonathanhuml

Copy link
Copy Markdown
Contributor Author

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.

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.

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.ZUNA module (encoder + classification head, position-aware RoPE, from_pretrained defaulting to braindecode/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.

Comment thread braindecode/models/zuna.py Outdated
Comment thread braindecode/models/zuna.py Outdated
Comment thread braindecode/models/util.py
Comment thread braindecode/models/zuna.py Outdated
Comment thread braindecode/models/zuna.py Outdated
@jonathanhuml

Copy link
Copy Markdown
Contributor Author

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.

jonathanhuml and others added 10 commits August 6, 2026 16:04
- 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.
@bruAristimunha

Copy link
Copy Markdown
Collaborator

Hey @jonathanhuml, once everything is green! I am happy to merge, many thanks for all the effort here!

@bruAristimunha
bruAristimunha merged commit cb9a870 into braindecode:master Aug 11, 2026
13 checks passed
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.

4 participants