Repository navigation
Geometry-compatibility contract test for every model with released weights (stacked on #1226) - #1228
Conversation
…od, LUNA; ZUNA default on_non_divisible='pad'
05b5186 to
aba4cf1
Compare
aba4cf1 to
5403186
Compare
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: 7
Open (8)
This relies onPatchTokenizer._prepare_input, which is a private method (leading underscore) and… · New This assertion is brittle becausestate_dict()includes buffers as well as parameters; a… · Newmne.set_log_level(\"ERROR\")sets a global logging level at import time, which can affect… · New MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… · New MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… · New MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… · Newbaseis a mutable list of dicts that’s shared across multiple parametrized cases… · Newself.rearrangepreviously referred to an einopsRearrangelayer but is now aPatchTokenizer.… · New
What changed in this PR
Adds a comprehensive, declarative “geometry compatibility” contract test for all models with released weights, and updates patch-tokenization behavior to support non-divisible windows via padding by default.
Changes:
- Added
test_pretrained_compat.pyto exercise every pretrained-weight model across a grid of channel geometries and window lengths (with strictxfailfor not-yet-supported cells). - Updated foundation-model tests and docs to reflect default padding for non-divisible windows (
on_non_divisible="pad"). - Integrated
PatchTokenizerinto multiple models (e.g., Labram/LUNA/CBraMod) and updated ZUNA defaults/docs accordingly.
| File | Description |
|---|---|
| test/unit_tests/models/test_pretrained_compat.py | New contract test covering model geometry/window compatibility across many scenarios. |
| test/unit_tests/models/test_foundation_models.py | Updates/extends tests to validate default padding behavior and warning/override behavior. |
| docs/whats_new.rst | Documents the new compatibility test and the on_non_divisible default/behavior changes. |
| braindecode/models/zuna.py | Changes default on_non_divisible to "pad" and updates docstring accordingly. |
| braindecode/models/luna.py | Introduces PatchTokenizer to support non-divisible windows; updates patch embedding & token prep. |
| braindecode/models/labram.py | Adds PatchTokenizer-backed handling for non-divisible windows; removes divisibility check from _PatchEmbed. |
| braindecode/models/cbramod.py | Replaces einops Rearrange patching with PatchTokenizer; updates head sizing for padded case. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| x_signal = self.tokenizer._prepare_input(x_signal) | ||
| if mask is not None: | ||
| mask = self.tokenizer._prepare_input(mask) |
| with torch.no_grad(): | ||
| out = model(x, ch_names=_TEN_TWENTY) if cls is Labram else model(x) | ||
| assert torch.isfinite(out).all() | ||
| assert not any("tokenizer" in k for k in model.state_dict()) |
| chs = [] | ||
| for n in names: | ||
| loc = np.zeros(12) | ||
| key = upper.get(n.upper()) | ||
| if key is not None: | ||
| loc[:3] = pos[key] | ||
| chs.append({"ch_name": n, "kind": kind, "loc": loc}) | ||
| return chs | ||
|
|
||
|
|
||
| def chs_coords_only(n, kind="eeg"): | ||
| chs = [] | ||
| for i in range(n): | ||
| th = 2 * math.pi * i / n | ||
| loc = np.zeros(12) | ||
| loc[:3] = [0.08 * math.cos(th), 0.08 * math.sin(th), 0.03] | ||
| chs.append({"ch_name": f"E{i + 1}", "kind": kind, "loc": loc}) | ||
| return chs | ||
|
|
||
|
|
||
| def chs_names_no_loc(names): | ||
| return [{"ch_name": n, "kind": "eeg", "loc": np.zeros(12)} for n in names] |
| def chs_from_montage(names, montage="standard_1005", kind="eeg"): | ||
| pos = _montage(montage).get_positions()["ch_pos"] | ||
| upper = {k.upper(): k for k in pos} | ||
| chs = [] | ||
| for n in names: | ||
| loc = np.zeros(12) | ||
| key = upper.get(n.upper()) | ||
| if key is not None: | ||
| loc[:3] = pos[key] | ||
| chs.append({"ch_name": n, "kind": kind, "loc": loc}) | ||
| return chs | ||
|
|
||
|
|
||
| def chs_coords_only(n, kind="eeg"): | ||
| chs = [] | ||
| for i in range(n): | ||
| th = 2 * math.pi * i / n | ||
| loc = np.zeros(12) | ||
| loc[:3] = [0.08 * math.cos(th), 0.08 * math.sin(th), 0.03] | ||
| chs.append({"ch_name": f"E{i + 1}", "kind": kind, "loc": loc}) | ||
| return chs | ||
|
|
||
|
|
||
| def chs_names_no_loc(names): | ||
| return [{"ch_name": n, "kind": "eeg", "loc": np.zeros(12)} for n in names] |
| chs = [] | ||
| for n in names: | ||
| loc = np.zeros(12) | ||
| key = upper.get(n.upper()) | ||
| if key is not None: | ||
| loc[:3] = pos[key] | ||
| chs.append({"ch_name": n, "kind": kind, "loc": loc}) | ||
| return chs | ||
|
|
||
|
|
||
| def chs_coords_only(n, kind="eeg"): | ||
| chs = [] | ||
| for i in range(n): | ||
| th = 2 * math.pi * i / n | ||
| loc = np.zeros(12) | ||
| loc[:3] = [0.08 * math.cos(th), 0.08 * math.sin(th), 0.03] | ||
| chs.append({"ch_name": f"E{i + 1}", "kind": kind, "loc": loc}) | ||
| return chs | ||
|
|
||
|
|
||
| def chs_names_no_loc(names): | ||
| return [{"ch_name": n, "kind": "eeg", "loc": np.zeros(12)} for n in names] |
| base = geos["G1"]["chs_info"] | ||
| if sf: | ||
| geos["G5a"] = dict(chs_info=base, sfreq=sf, n_times=int(sf)) | ||
| geos["G5b"] = dict(chs_info=base, sfreq=sf, n_times=int(sf * 30)) | ||
| geos["G5c"] = dict(chs_info=base, sfreq=sf, n_times=nt + 37) |
| # padding/cropping a non-divisible time axis at forward time. | ||
| self.rearrange = PatchTokenizer( | ||
| patch_size=patch_size, | ||
| n_times=self._n_times if self._n_times is not None else patch_size, | ||
| on_non_divisible=on_non_divisible, |
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: 2
Open (14)
_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()… · New_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()… · New The docstring/shape comment for_SegmentPatch.forward()still statesn_times//patch_size, but… · New Setting MNE’s log level at import time changes global process state for the entire test session and… · Newchs_infois a mutable list of dicts, and the same object is reused across multiple geometries… · Newbaseis a mutable list of dicts that’s shared across multiple parametrized cases… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by…mne.set_log_level(\"ERROR\")sets a global logging level at import time, which can affect… This assertion is brittle becausestate_dict()includes buffers as well as parameters; a… This relies onPatchTokenizer._prepare_input, which is a private method (leading underscore) and…self.rearrangenow holds aPatchTokenizer, not an einopsRearrangelayer. This name is… · Newself.rearrangepreviously referred to an einopsRearrangelayer but is now aPatchTokenizer.…
| def __init__( | ||
| self, embed_dim: int = 64, patch_size: int = 40, on_non_divisible: str = "pad" | ||
| ) -> None: | ||
| super().__init__() | ||
| self.patch_size = patch_size | ||
| self.embed_dim = embed_dim | ||
| self.tokenizer = PatchTokenizer( | ||
| patch_size=patch_size, n_times=patch_size, on_non_divisible=on_non_divisible | ||
| ) |
| output: (B, C*S, D) where S = T//patch_size, D = embed_dim | ||
| """ | ||
| x = rearrange(x, "B C (S P) -> B (C S) P", P=self.patch_size) | ||
| x = rearrange(self.tokenizer(x), "B C S P -> B (C S) P") |
| Returns: | ||
| -------- | ||
| X_patch: Tensor | ||
| [batch, n_chans, n_times//patch_size, patch_size] |
| base = geos["G1"]["chs_info"] | ||
| if sf: | ||
| geos["G5a"] = dict(chs_info=base, sfreq=sf, n_times=int(sf)) | ||
| geos["G5b"] = dict(chs_info=base, sfreq=sf, n_times=int(sf * 30)) | ||
| geos["G5c"] = dict(chs_info=base, sfreq=sf, n_times=nt + 37) | ||
| return geos |
| # Shared tokenizer: (batch, n_chans, n_times) -> (batch, n_chans, n_patch, patch_size), | ||
| # padding/cropping a non-divisible time axis at forward time. | ||
| self.rearrange = PatchTokenizer( | ||
| patch_size=patch_size, | ||
| n_times=self._n_times if self._n_times is not None else patch_size, | ||
| on_non_divisible=on_non_divisible, | ||
| ) | ||
| self._on_non_divisible = on_non_divisible |
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
REVE is skipped, sampling-rate coverage is absent, and some rejection checks can report false positives.
Review effort: Balanced
Findings: 2
Open (17)
_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()…_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()… Remove EEGDINO-G2 from expected failures · New Ensure REVE cases run offline in CI · Newchs_infois a mutable list of dicts, and the same object is reused across multiple geometries… Setting MNE’s log level at import time changes global process state for the entire test session and… The docstring/shape comment for_SegmentPatch.forward()still statesn_times//patch_size, but…baseis a mutable list of dicts that’s shared across multiple parametrized cases… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by…mne.set_log_level(\"ERROR\")sets a global logging level at import time, which can affect… This assertion is brittle becausestate_dict()includes buffers as well as parameters; a… This relies onPatchTokenizer._prepare_input, which is a private method (leading underscore) and… Add alternate sampling-rate test coverage · Newself.rearrangenow holds aPatchTokenizer, not an einopsRearrangelayer. This name is…self.rearrangepreviously referred to an einopsRearrangelayer but is now aPatchTokenizer.…
| ("BENDR", "G4"), | ||
| ("BENDR", "G2"), | ||
| ("BENDR", "G3"), | ||
| ("BENDR", "G3b"), # ChannelTokenizer(fixed_order) | ||
| ( | ||
| "EEGDINO", | ||
| "G2", | ||
| ), # ChannelTokenizer(index_slots): > 19 channels needs a declared error | ||
| ( | ||
| "SignalJEPA", | ||
| "G3b", | ||
| ), # names without coordinates: output is NaN (division by zero) today |
| strict=True, reason="not migrated yet (see design doc)" | ||
| ) | ||
| ) | ||
| yield pytest.param(name, gname, gkw, id=f"{name}-{gname}", marks=marks) |
| geos["G5a"] = dict(chs_info=base, sfreq=sf, n_times=int(sf)) | ||
| geos["G5b"] = dict(chs_info=base, sfreq=sf, n_times=int(sf * 30)) | ||
| geos["G5c"] = dict(chs_info=base, sfreq=sf, n_times=nt + 37) |
mne.set_log_level('ERROR') at import silenced mne.utils.warn for the whole pytest worker, so 15 pytest.warns tests in test_models.py failed with DID NOT WARN when collected after this file.
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: 2
Open (20)
_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()…_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()…baseis the same mutablechs_infolist reused acrossG1/G5a/G5b/G5c. If any model… · Newwarnings.simplefilter('ignore')suppresses all warnings from model construction/forward, which… · New Forwant == 'raise', the test currently accepts anyValueError/RuntimeError, which can allow… · New Ensure REVE cases run offline in CI Remove EEGDINO-G2 from expected failureschs_infois a mutable list of dicts, and the same object is reused across multiple geometries… Setting MNE’s log level at import time changes global process state for the entire test session and… The docstring/shape comment for_SegmentPatch.forward()still statesn_times//patch_size, but…baseis a mutable list of dicts that’s shared across multiple parametrized cases… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by…mne.set_log_level(\"ERROR\")sets a global logging level at import time, which can affect… This assertion is brittle becausestate_dict()includes buffers as well as parameters; a… This relies onPatchTokenizer._prepare_input, which is a private method (leading underscore) and… Add alternate sampling-rate test coverageself.rearrangenow holds aPatchTokenizer, not an einopsRearrangelayer. This name is…self.rearrangepreviously referred to an einopsRearrangelayer but is now aPatchTokenizer.…
| base = geos["G1"]["chs_info"] | ||
| if sf: | ||
| geos["G5a"] = dict(chs_info=base, sfreq=sf, n_times=int(sf)) | ||
| geos["G5b"] = dict(chs_info=base, sfreq=sf, n_times=int(sf * 30)) | ||
| geos["G5c"] = dict(chs_info=base, sfreq=sf, n_times=nt + 37) |
| with warnings.catch_warnings(): | ||
| warnings.simplefilter("ignore") | ||
| model = spec["cls"](**kw).eval() | ||
| with torch.no_grad(): | ||
| return model(torch.randn(1, n_ch, gkw["n_times"])) |
| if want == "raise": | ||
| with pytest.raises((ValueError, RuntimeError)): | ||
| build_and_forward() |
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## master #1228 +/- ##
==========================================
+ Coverage 88.41% 88.42% +0.01%
==========================================
Files 156 156
Lines 19009 19009
==========================================
+ Hits 16806 16809 +3
+ Misses 2203 2200 -3 🚀 New features to boost your workflow:
|
Inline the 19-line TEN_TWENTY list as one split(), and fold the NOT_YET branch of test_geometry_contract into the forward branch (it ran the same assertion). Parametrize ids and strict xfail markers are unchanged.
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
EEGDINO asserts the wrong xfail target, REVE is skipped by default, and sampling-rate geometries are not generated.
Review effort: Balanced
Findings: 2
Open (19)
_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()…_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()… Forwant == 'raise', the test currently accepts anyValueError/RuntimeError, which can allow…warnings.simplefilter('ignore')suppresses all warnings from model construction/forward, which…baseis the same mutablechs_infolist reused acrossG1/G5a/G5b/G5c. If any model… Ensure REVE cases run offline in CI Remove EEGDINO-G2 from expected failureschs_infois a mutable list of dicts, and the same object is reused across multiple geometries… The docstring/shape comment for_SegmentPatch.forward()still statesn_times//patch_size, but…baseis a mutable list of dicts that’s shared across multiple parametrized cases… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by…mne.set_log_level(\"ERROR\")sets a global logging level at import time, which can affect… This assertion is brittle becausestate_dict()includes buffers as well as parameters; a… This relies onPatchTokenizer._prepare_input, which is a private method (leading underscore) and… Add alternate sampling-rate test coverageself.rearrangenow holds aPatchTokenizer, not an einopsRearrangelayer. This name is…self.rearrangepreviously referred to an einopsRearrangelayer but is now aPatchTokenizer.…
Resolved since last review (1)
Resolved docs/whats_new.rst by keeping both entries.
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
The EEGDINO expectation is inverted, and released-weight models are missing from the advertised comprehensive contract.
Review effort: Balanced
Findings: 2
Open (20)
_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()…_PatchEmbedNetworkinitializesPatchTokenizerwithn_times=patch_size, butforward()… Cover NeuroRVQ and BrainOmni in geometry compatibility tests · New Forwant == 'raise', the test currently accepts anyValueError/RuntimeError, which can allow…warnings.simplefilter('ignore')suppresses all warnings from model construction/forward, which…baseis the same mutablechs_infolist reused acrossG1/G5a/G5b/G5c. If any model… Ensure REVE cases run offline in CI Remove EEGDINO-G2 from expected failureschs_infois a mutable list of dicts, and the same object is reused across multiple geometries… The docstring/shape comment for_SegmentPatch.forward()still statesn_times//patch_size, but…baseis a mutable list of dicts that’s shared across multiple parametrized cases… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… MNEchs_infoentries typically storekindas a FIFF integer constant (as produced by… This assertion is brittle becausestate_dict()includes buffers as well as parameters; a… This relies onPatchTokenizer._prepare_input, which is a private method (leading underscore) and… Remove unsupported alternate sampling-rate coverage claim · New Add alternate sampling-rate test coverageself.rearrangenow holds aPatchTokenizer, not an einopsRearrangelayer. This name is…self.rearrangepreviously referred to an einopsRearrangelayer but is now aPatchTokenizer.…
Resolved since last review (1)
| channels="coords", | ||
| coords_checked=False, | ||
| ), | ||
| } |
| names without coordinates, short / long / non-divisible windows, other | ||
| sampling rates -- and must either forward or raise the *declared* error. |
test_pretrained_compat.py (add/add after braindecode#1228's squash merge): keep this PR's side, i.e. braindecode#1228's file with the six NOT_YET cells and the fixed_order raise removed, since ChannelTokenizer makes them pass. whats_new: keep both; drop the EEGDINO clause (this PR does not change EEGDINO).



Summary
Stacked on #1226 (review the last commit only). Part of #1227.
One test that checks every model with released weights on a grid of input geometries — the mechanism timm uses (
test_model_load_pretrainedoverlist_models(pretrained=True)), applied to EEG's axes: canonical montage, permuted order, a 64-channel montage outside the 10-20 vocabulary, coordinates-only channels, names without coordinates, 1 s / 30 s / non-divisible windows. Random weights, no download.Changes
test/unit_tests/models/test_pretrained_compat.py: 19 classes × 9 geometries. Expected outcome per cell is derived from one declared channel strategy per model (COMPAT), not hard-coded; cells not handled yet arexfail(strict=True)(LaBraM outside its vocabulary, BENDR on any other order, EEGDINO > 19 channels, SignalJEPA NaN with names-without-coords), so each fix flips a marker.docs/whats_new.rst.Testing
pytest test/unit_tests/models/test_pretrained_compat.py→ 128 passed, 8 skipped (REVE: network-marked by conftest), 9 xfailed (~19 s on CPU).ruffclean.Regression test (this is one)
Style checks recorded
docs/whats_new.rstupdatedNotes for reviewers
Found while writing it: SignalJEPA returns NaN when channel names have zero
loc(division by zero in the channel embedding) — tracked in #1227, not fixed here.