Repository navigation
Add DIVER-1 any-variate iEEG foundation model - #1170
Conversation
Port the DIVER-1 encoder plus a classification head as braindecode.models.DIVER1, with the patch CNN and spectral embedding, channel-metadata embeddings, the sliding-window spatio-temporal conditional positional embedding and the any-variate transformer encoder. Co-authored-by: Cursor <[email protected]>
There was a problem hiding this comment.
🟡 Changes recommended
Moderate issues remain around metadata defaults, public equivariance claims, and focused test coverage.
Get a fresh assessment by requesting another Copilot review.
Pull request overview
Adds the architecture-only DIVER1 any-variate iEEG foundation model and integrates it into Braindecode.
Changes:
- Implements DIVER1 embeddings, metadata handling, transformer encoder, and classification head.
- Registers and publicly exports the model; updates documentation, licensing, and catalog metadata.
- Adds integration and return-feature test coverage.
File summaries
| File | Summary |
|---|---|
test/unit_tests/models/test_return_features.py |
Adds DIVER1 feature-return coverage. |
test/unit_tests/models/test_integration.py |
Adds integration coverage and TorchScript exemption. |
NOTICE.txt |
Adds licensing notice. |
docs/whats_new.rst |
Documents the new model. |
docs/api.rst |
Adds the API listing. |
braindecode/models/util.py |
Registers integration parameters. |
braindecode/models/summary.csv |
Adds model catalog metadata. |
braindecode/models/diver1.py |
Implements DIVER1; review findings concern metadata defaults, permutation-equivariance wording, and missing focused metadata/patch-size tests. |
braindecode/models/__init__.py |
Exports DIVER1 publicly. |
Review details
Suppressed comments (6)
braindecode/models/diver1.py:574
- This unconditional check prevents the advertised any-variate path from accepting a different channel count: with
pooling="mean"anduse_position_emb=False, the remaining operations derive their shapes/IDs fromxand the head has no channel dimension, but this raises before reaching them. Gate this validation on fixed metadata or flatten pooling (or revise the any-variate API claim).
if x.shape[1] != self.n_chans:
braindecode/models/diver1.py:458
chs_info["loc"][:3]is an MNE head-frame position, not an MNI coordinate (seebraindecode/models/util.py:811-813). Multiplying it by 1e3 and treating it as MNI silently feeds the wrong coordinate system into this positional encoder for the documented fallback; require an explicit MNIchan_posor apply and document the appropriate transform.
# MNE stores channel positions in metres; DIVER-1 encodes MNI
# coordinates in millimetres.
rows.append([float(v) * 1e3 for v in list(loc)[:3]])
braindecode/models/diver1.py:499
- When
chs_infois omitted, everykindisNoneand this maps every channel to the learnedEEGembedding. Since DIVER1 is an iEEG model and the metadata-optional constructor is exercised by the registry/integration tests, missing metadata silently selects the scalp-EEG modality instead of a defined unknown/neutral value or the model's iEEG default.
modality = ["iEEG" if _is_intracranial(k) else "EEG" for k in kinds]
braindecode/models/diver1.py:424
- Directly assigning
_n_outputsbypassesEEGModuleMixin._set_n_outputs, so serialized configs can still report the old class count afterreset_head; the replacementLinearis also always CPU/float32, which breaks a reset on a moved or half-precision model. Preserve the old head's metadata, device, and dtype when rebuilding it.
self._n_outputs = n_outputs
self.final_layer = nn.Linear(self.final_layer.in_features, n_outputs)
braindecode/models/diver1.py:439
chan_posis documented as array-like and is accepted here from NumPy/Torch arrays, but the basebuild_model_configskips non-JSON-serializable init arguments and these coordinates only live in non-persistent buffers. Consequently,get_config()/from_config()silently drops explicit array coordinates and reloads a different montage. Store a JSON-serializable copy when explicit positions are supplied.
if chan_pos is not None:
pos = torch.as_tensor(chan_pos, dtype=torch.float32)
braindecode/models/diver1.py:697
- The constructor accepts any positive
patch_size, but this default formula producesstride=0for padded patch sizes below 16, so valid inputs such aspatch_size=1fail in the stride validation. Clamp the default to at least 1 or reject those patch sizes explicitly.
stride = padded // (8 if patch_size >= 100 else 16)
- Files reviewed: 9/10 changed files
- Comments generated: 2
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| coordinate term contributes little). Set ``use_position_emb=False`` to drop | ||
| the coordinate and type terms entirely, which makes the whole model exactly | ||
| channel-permutation equivariant. |
| coords, coords_known = self._resolve_chan_pos(chan_pos) | ||
| self.register_buffer("chan_coords", coords, persistent=False) | ||
| self.register_buffer("chan_coords_known", coords_known, persistent=False) | ||
| type_idx, subtype_idx, subtype_known = self._resolve_chan_types( | ||
| chan_modality, chan_subtype | ||
| ) | ||
| self.register_buffer("chan_type_idx", type_idx, persistent=False) | ||
| self.register_buffer("chan_subtype_idx", subtype_idx, persistent=False) | ||
| self.register_buffer("chan_subtype_known", subtype_known, persistent=False) |
d6e31d5 to
2cec62b
Compare
|
Validation against the authors' reference implementation, ahead of braindecode's own suites:
No checkpoint exists to verify against (the reference ships a placeholder for |
There was a problem hiding this comment.
🟡 Changes recommended
Unresolved model behavior, configuration persistence, coordinate handling, padding, catalog, and test-coverage issues remain.
Get a fresh assessment by requesting another Copilot review.
Review details
Suppressed comments (4)
braindecode/models/diver1.py:136
- With the default
pooling="flatten", this is not true for the model's logits: the final linear layer consumes channel-ordered tokens, so permuting channels changes the output even when the encoder is equivariant. This setting makes the encoder representation equivariant; onlypooling="mean"makes the readout invariant to channel order.
coordinate term contributes little). Set ``use_position_emb=False`` to drop
the coordinate and type terms entirely, which makes the whole model exactly
channel-permutation equivariant.
braindecode/models/diver1.py:925
F.unfoldandF.foldinterpret the padding value on both sides of the temporal axis. Withwindow=7, this producesn_patches + 6windows and gives each output patch a 13-patch effective receptive field, rather than the centered 7-patch window described above. Use half-padding (self.window // 2) here (or otherwise document and test the intended wider receptive field).
"padding": (0, self.window - 1),
braindecode/models/summary.csv:57
- The parameter count in this catalog row does not match the implementation. For the displayed defaults (22 channels, 1000 samples, 4 outputs, hence 2 patches), summing the parameters created by
DIVER1gives 12,716,742, not 12,711,366, so the model-zoo metadata underreports the model size. Please regenerate or update this row.
DIVER1,General,"Prediction, Embedding",500,"n_chans, n_outputs, n_times",12711366,"DIVER1(n_chans=22, n_outputs=4, n_times=1000, sfreq=500)","Attention/Transformer,Foundation Model,Channel",iEEG
test/unit_tests/models/test_return_features.py:89
- The added DIVER1 coverage only smoke-tests the default no-metadata, exact-window path. It does not exercise the
chs_info/explicit metadata resolution,pooling="mean", non-divisible patch padding, or the STCPE/attention ID packing, so substantial regressions in this new encoder could pass the suite; add focused tests for these branches and shape/invariance contracts.
pytest.param(
DIVER1,
N_CHANS,
{"patch_size": 100, "d_model": 64, "n_layers": 2},
True,
id="DIVER1",
- Files reviewed: 8/9 changed files
- Comments generated: 3
- Review effort level: Lite
| if chan_pos is not None: | ||
| pos = torch.as_tensor(chan_pos, dtype=torch.float32) |
| # MNE stores channel positions in metres; DIVER-1 encodes MNI | ||
| # coordinates in millimetres. | ||
| rows.append([float(v) * 1e3 for v in list(loc)[:3]]) |
Co-authored-by: Cursor <[email protected]>
There was a problem hiding this comment.
🔵 Needs a closer look
Unresolved moderate issues remain in DIVER1 configuration, metadata handling, coordinate semantics, channel flexibility, and invariance claims.
Review details
Suppressed comments (5)
Previously missed (2) — in code that hasn't changed since the last review.
braindecode/models/diver1.py:423
reset_headbypassesEEGModuleMixin._set_n_outputs, so invalid output sizes are not validated andget_config()/Hub config retain the oldn_outputs. It also recreates the head on CPU/float32, so calling it after.double()or.to(device)makes the next forward fail. Use the shared helper and preserve the old head's device and dtype.
braindecode/models/diver1.py:577- This guard fixes the channel count to the montage used at construction, and the channel metadata buffers are fixed to the same size. Consequently even
pooling="mean"cannot accept an arbitrary-channel recording; callers must rebuild the model, which contradicts the PR's any-variate/variable-input contract. Either make per-recording metadata part offorwardand derive these tensors dynamically, or narrow the public claim to a model rebuilt per montage.
braindecode/models/diver1.py:376
chan_posis documented as an array-like, but this code only keeps the resolved coordinates in non-persistent buffers.EEGModuleMixin.build_model_configdrops non-JSONable constructor values such as NumPy arrays, soget_config()and Hub save/load silently lose explicitly supplied positions and reload fromchs_infoor unknown coordinates. Preserve a JSON-safe copy of the explicit positions in the serialized config (or restore these buffers through another defined mechanism).
coords, coords_known = self._resolve_chan_pos(chan_pos)
self.register_buffer("chan_coords", coords, persistent=False)
self.register_buffer("chan_coords_known", coords_known, persistent=False)
braindecode/models/diver1.py:457
chs_info[i]["loc"][:3]is an MNE head-frame location, not an MNI coordinate (seebraindecode/models/util.py:729-731). This fallback therefore silently feeds the wrong coordinate system into an embedding documented as MNI; require/transform an explicit MNIchan_pos, or document and use the actual frame instead.
# MNE stores channel positions in metres; DIVER-1 encodes MNI
# coordinates in millimetres.
rows.append([float(v) * 1e3 for v in list(loc)[:3]])
braindecode/models/diver1.py:135
- This overstates the invariance guarantee: with the default
pooling="flatten", the fixed-position linear head assigns different weights to channel slots, so permuting channels changes the logits even when the encoder is equivariant. Qualify the statement to the encoder, or limit the logits claim to mean pooling.
coordinate term contributes little). Set ``use_position_emb=False`` to drop
the coordinate and type terms entirely, which makes the whole model exactly
channel-permutation equivariant.
- Files reviewed: 8/9 changed files
- Comments generated: 0 new
- Review effort level: Lite
bruAristimunha
left a comment
There was a problem hiding this comment.
Small details, and please don't forget to put your name in the citation.cff please.
| The braindecode reimplementation is pure-PyTorch (no ``mup``, no ``jaxtyping``) | ||
| and covers the downstream encoder only. |
| Reimplementation of DIVER-1 (Han et al., 2025), "DIVER-1: Scaling Intracranial | ||
| EEG Foundation Models for Transferable Representations". The architecture is | ||
| transcribed from the authors' reference implementation, whose Transformer | ||
| encoder is adapted from Salesforce's MOIRAI / ``uni2ts`` (Copyright Salesforce, | ||
| Inc.), released under the Apache License, Version 2.0. The license this port | ||
| should carry follows from that provenance but is still to be confirmed with the | ||
| braindecode maintainers. |
There was a problem hiding this comment.
I think Labram is a similar case, with a link list of licenses. I think here we assume the license is from the head, and at most we put some small notes saying that the person needs to check the original implementation.
| # FIFF channel-kind codes (mne.io.constants.FIFF), so chs_info can be read | ||
| # without importing mne. | ||
| _FIFF_SEEG, _FIFF_DBS, _FIFF_ECOG = 802, 803, 902 |
There was a problem hiding this comment.
Is there any specific reason why not import from MNE? In Braindecode, we have a dependency on MNE.
There was a problem hiding this comment.
I will correct it !
| self._n_outputs = n_outputs | ||
| self.final_layer = nn.Linear(self.final_layer.in_features, n_outputs) | ||
|
|
||
| def _resolve_chan_pos(self, chan_pos) -> tuple[torch.Tensor, torch.Tensor]: |
There was a problem hiding this comment.
I think we have an util functions for this.
There was a problem hiding this comment.
Indeed, I will use it but I had to add a fill_missing argument to it that defaults to False. When True, it outputs nan for the channels without info or loc
| "stride": (1, 1), | ||
| "padding": (0, self.window - 1), | ||
| } | ||
| flat = rearrange(x, "batch chans patches dim -> (batch dim) 1 chans patches") |
There was a problem hiding this comment.
define as an einops layer, and you can get back torch script compatibility.
| unfold_kwargs = { | ||
| "kernel_size": (n_chans, self.window), | ||
| "stride": (1, 1), | ||
| "padding": (0, self.window - 1), | ||
| } |
There was a problem hiding this comment.
Let's define this in the init?
| def forward(self, x: torch.Tensor) -> torch.Tensor: | ||
| hidden = self.activation(self.fc_gate(x)) * self.fc1(x) | ||
| return self.dropout2(self.fc2(self.dropout1(hidden))) |
There was a problem hiding this comment.
unstack the operation, could be useful later to inspect the layers
| class _RMSNorm(nn.Module): | ||
| """Root-mean-square layer normalisation with a learned gain. | ||
|
|
||
| :class:`torch.nn.RMSNorm` is only available from PyTorch 2.4, while | ||
| braindecode supports ``torch>=2.0``; this equivalent keeps the model | ||
| importable on older PyTorch, as done in | ||
| :class:`~braindecode.models.ZUNA` and :class:`~braindecode.models.REVE`. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| normalized_shape : int | ||
| Size of the trailing dimension to normalise. | ||
| eps : float | ||
| Term added to the mean square for numerical stability. | ||
| """ | ||
|
|
||
| def __init__(self, normalized_shape: int, eps: float = 1e-5): | ||
| super().__init__() | ||
| self.eps = eps | ||
| self.weight = nn.Parameter(torch.ones(normalized_shape)) | ||
|
|
||
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
There was a problem hiding this comment.
seems the perfect excuse that i was looking to increase the pytorch version :) so, we can delete.
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1170 +/- ##
==========================================
+ Coverage 87.52% 87.71% +0.18%
==========================================
Files 148 149 +1
Lines 17080 17433 +353
==========================================
+ Hits 14950 15292 +342
- Misses 2130 2141 +11 🚀 New features to boost your workflow:
|
…ints The official DIVER-1 code scales attention scores by 1 / head_dim (muP, original_moirai_encoder.py:709) and every released checkpoint was trained that way; the port used PyTorch's default 1 / sqrt(head_dim). Add `mup_attention` (default True) and pass it through both encoders. With the official iEEG Tiny checkpoint (patch 50) the port now matches the reference encoder features exactly (max abs diff 0.0); mup_attention=False reproduces the previous 0.342 gap.
Take master's REVE and ZUNA files (same RMSNorm change, from braindecode#1176) and drop the duplicate API-changes entry for the PyTorch >= 2.4 requirement (braindecode#1174). Affected model tests: 328 passed, 56 skipped.
The official ieeg_pretrained_weights.pt (DIVER-1-0.1s Tiny) is now braindecode/DIVER-1-0.1s-tiny. scripts/convert_diver1_weights.py pins its SHA-256, renames the tensors, checks the encoder features against the official code (identical on CPU) and publishes the encoder with its MIT licence and a model card. The docstring now points to the weights and states their licence.
scripts/convert_diver1_weights.py now converts both release files: the iEEG DIVER-1-0.1s Tiny (braindecode/DIVER-1-0.1s-tiny) and the iEEG and EEG DIVER-1-1s Small (braindecode/DIVER-1-1s-small). Parity with the official code is checked for intracranial and scalp channels (identical on CPU). The saved placeholder montage has zero positions instead of NaN, so config.json is valid JSON on the Hub; the script refuses to publish otherwise.
Resolve util import, changelog, and foundation-test conflicts by preserving both upstream and DIVER additions. Retain upstream model registries and the 32-entry direct TorchScript count. DIVER implementation is unchanged. Validation: targeted model/registry tests 118 passed; utility tests 29 passed; shared suites 1555 passed, 146 skipped, with two pre-existing DIVER reset-head configuration/save failures reproduced on original PR head. Conflict-file pre-commit passed; independent review approved.
Reuse native sequential layers, remove redundant helpers, and condense documentation while preserving checkpoint and numerical behavior. The converter remains available in both DIVER Hugging Face model repositories.
|
|
||
| #: MNE channel types recorded from inside the skull. | ||
| INTRACRANIAL_CH_TYPES = frozenset({"seeg", "dbs", "ecog"}) | ||
|
|
There was a problem hiding this comment.
on top of file this...
| """DIVER-1: an any-variate iEEG foundation model. | ||
|
|
||
| Reimplementation of DIVER-1 (Han et al., 2025), "DIVER-1: Scaling Intracranial | ||
| EEG Foundation Models for Transferable Representations". The architecture is | ||
| transcribed from the authors' reference implementation, whose Transformer | ||
| encoder is adapted from Salesforce's MOIRAI / ``uni2ts`` (Copyright Salesforce, | ||
| Inc.), released under the Apache License, Version 2.0; this file is therefore | ||
| distributed under Apache-2.0 (https://www.apache.org/licenses/LICENSE-2.0). | ||
| The reference repository states no license of its own, so check the terms of the | ||
| original implementation before redistributing. | ||
|
|
||
| Original Authors: Han et al., Seoul National University | ||
| Braindecode Adaptation: Julien Gadonneix | ||
| """ |
There was a problem hiding this comment.
much more compact, and remove these not necessary comments
| # Reference slots: modality EEG=0/iEEG=1; subtype grid=0/strip=1/depth=2. | ||
| # MNE kinds cannot identify strips; -1 marks unknown subtypes. | ||
| _TYPE_TO_SLOTS = {"eeg": (0, -1), "ecog": (1, 0), "seeg": (1, 2), "dbs": (1, 2)} | ||
|
|
There was a problem hiding this comment.
should be in code, not as constant
There was a problem hiding this comment.
Should this be passed as an arg? Because there is no convention on this
| def _check_channel_metadata( | ||
| metadata: torch.Tensor, | ||
| n_chans: int, | ||
| n_modalities: int = 2, | ||
| n_subtypes: int = 3, | ||
| ) -> None: | ||
| """Reject metadata that does not describe the channels of this batch.""" | ||
| if metadata.dim() != 2 or metadata.shape[0] != n_chans or metadata.shape[1] != 5: | ||
| raise ValueError( | ||
| f"chan_metadata must have shape ({n_chans}, 5), one row of " | ||
| f"(x, y, z, modality, sub-modality) per channel of the input, got " | ||
| f"{list(metadata.shape)}." | ||
| ) | ||
| modality = metadata[:, 3] | ||
| if bool(((modality < 0) | (modality >= n_modalities)).any()): | ||
| raise ValueError( | ||
| f"The modality column of chan_metadata must hold slots in " | ||
| f"[0, {n_modalities}), got values from {float(modality.min())} to " | ||
| f"{float(modality.max())}." | ||
| ) | ||
| if bool((metadata[:, 4] >= n_subtypes).any()): | ||
| raise ValueError( | ||
| f"The sub-modality column of chan_metadata must hold slots below " | ||
| f"{n_subtypes}, or a negative value for unknown, got up to " | ||
| f"{float(metadata[:, 4].max())}." | ||
| ) | ||
|
|
There was a problem hiding this comment.
no necessary this check in the top, not sure if it is act necessary.
There was a problem hiding this comment.
You would remove it completely? It could be a useful check when the metadata is passed in the forward function
There was a problem hiding this comment.
just move for the bottom of the file, this what i was thinking
| @staticmethod | ||
| def channel_metadata(chs_info: list[dict]) -> torch.Tensor: | ||
| """Assemble one recording's electrode metadata for :meth:`forward`. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| chs_info : list of dict | ||
| MNE channel information (``info["chs"]``) of the recording, one | ||
| entry per channel, holding its kind and its location. | ||
|
|
||
| Returns | ||
| ------- | ||
| torch.Tensor | ||
| ``(n_chans, 5)`` float tensor whose columns are the three MNI | ||
| coordinates in millimetres, the recording-modality slot and the | ||
| electrode-sub-modality slot, the latter negative when unknown. | ||
|
|
||
| Examples | ||
| -------- | ||
| >>> import mne | ||
| >>> from braindecode.models import DIVER1 | ||
| >>> info = mne.create_info(["A1", "A2"], 500.0, "seeg") | ||
| >>> DIVER1.channel_metadata(info["chs"])[:, 3:].tolist() | ||
| [[1.0, 2.0], [1.0, 2.0]] | ||
| """ | ||
| n_chans = len(chs_info) | ||
| # MNE stores coordinates in metres; DIVER-1 encodes MNI coordinates in | ||
| # millimetres. | ||
| coords = 1e3 * torch.as_tensor( | ||
| extract_channel_locations_from_chs_info( | ||
| chs_info, num_channels=n_chans, fill_missing=True | ||
| ), | ||
| dtype=torch.float32, | ||
| ) | ||
| types = channel_types_from_chs_info(chs_info, num_channels=n_chans) | ||
| undetermined = sorted({t for t in types if t not in _TYPE_TO_SLOTS}) | ||
| if undetermined: | ||
| raise ValueError( | ||
| f"DIVER1 cannot determine the recording modality of every " | ||
| f"channel: the chs_info 'kind' of some resolves to " | ||
| f"{undetermined}, which is neither scalp EEG nor an " | ||
| f"intracranial type ({sorted(INTRACRANIAL_CH_TYPES)}). Set the " | ||
| f"channel kinds in chs_info accordingly." | ||
| ) | ||
| slots = torch.tensor([_TYPE_TO_SLOTS[t] for t in types], dtype=coords.dtype) | ||
| return torch.cat([coords, slots], dim=-1) |
There was a problem hiding this comment.
Nice way, but let's remove from the method, and transform it into a function... it is a really useful function
| for i in range(n_to_extract): | ||
| ch_info = chs_info[i] if i < len(chs_info) else None | ||
| coordinates = None | ||
|
|
||
| if isinstance(ch_info, dict): | ||
| loc = ch_info.get("loc") | ||
| if loc is not None: | ||
| try: | ||
| loc_array = np.asarray(loc, dtype=np.float32) | ||
| except (ValueError, TypeError): | ||
| loc_array = None | ||
| # MNE format: 12-element array with electrode position at indices 0:3 | ||
| if ( | ||
| loc_array is not None | ||
| and loc_array.ndim == 1 | ||
| and loc_array.size >= 3 | ||
| ): | ||
| coordinates = loc_array[:3] | ||
|
|
||
| if coordinates is None: | ||
| if not fill_missing: |
There was a problem hiding this comment.
too big diff here...
Preserve Julien's master merge and EEGMiner real-DFT fixes. Resolve shared functional test append conflict by retaining both rotary and real-DFT coverage.
|
many thank you @julien-gadonneix 🙏🏽 |
Adds
DIVER1(encoder + classification head) from DIVER-1.Model. An any-variate transformer for intracranial EEG. A small patch CNN plus a spectral branch turn each channel into time patches, and attention then runs jointly over all electrodes and time steps, so an arbitrary number and layout of channels is supported. Channel identity comes from
chs_infometadata instead of a fixed montage: sinusoidal electrode coordinates, modality/subtype embeddings, a sliding-window spatio-temporal conditional positional embedding, rotary temporal offsets and a learned same/cross-channel attention bias. The pretraining-only parts (masking, MDRO heads) are not ported, and no public checkpoint exists yet, so this is architecture-only.Checked against the source. The architecture was double-checked against both the paper and the authors' original codebase: transplanting reference weights into this port reproduces the reference forward pass bit-exactly, with matching encoder parameter counts.
Tests.
test_integration.pyandtest_return_features.pypass forDIVER1, as do the categorization / modality / registry suites, andpre-commit(ruff, mypy, codespell, sphinx-lint) is clean. TorchScript is exempted, sinceeinopsand the polymorphicreturn_featuresreturn type are not scriptable, consistent with the existing exemptions.Guidance and advice would be very welcome. Happy to adjust anything.