Repository navigation
from_pretrained: take input geometry from the caller, not the Hub config (stacked on #1226) - #1232
Conversation
…od, LUNA; ZUNA default on_non_divisible='pad'
97bcf03 to
acef441
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: 1
Open (4)
The current pinning logic triggers even when the caller explicitly passes geometry keys withNone… · Newself.rearrangeis now aPatchTokenizer, not an einopsRearrange, so the attribute name is… · New Similar to LUNA, this relies onPatchTokenizer._prepare_input, which is a private method. To… · New This calls a privatePatchTokenizermethod (_prepare_input), which is brittle (private APIs can… · New
What changed in this PR
Fixes Hugging Face from_pretrained geometry precedence so caller-provided input geometry (channels/time/sfreq) doesn’t conflict with geometry loaded from config.json, and extends foundation-model patching behavior to support non-divisible windows via PatchTokenizer (padding/cropping) with updated defaults, tests, and docs.
Changes:
- Override
EEGModuleMixin.from_pretrained(HF-only) to pin derived geometry arguments based on caller inputs. - Introduce/propagate
on_non_divisible+PatchTokenizerusage in Labram/CBraMod/LUNA/ZUNA paths and update defaults. - Add regression tests and update release notes.
| File | Description |
|---|---|
braindecode/models/base.py |
Pins geometry kwargs in from_pretrained to prevent config-supplied derived args from colliding with caller overrides. |
braindecode/models/labram.py |
Uses PatchTokenizer to handle non-divisible time axes and updates patch count logic. |
braindecode/models/cbramod.py |
Replaces einops Rearrange patching with PatchTokenizer and adjusts head sizing for padded windows. |
braindecode/models/luna.py |
Adds on_non_divisible support and tokenization/padding through PatchTokenizer. |
braindecode/models/zuna.py |
Updates on_non_divisible default to "pad" and aligns docstring. |
test/unit_tests/models/test_eegdino.py |
Adds regression test ensuring from_pretrained takes geometry from the caller (not config). |
test/unit_tests/models/test_foundation_models.py |
Updates tests for new padding default and adds shared non-divisible-window behavior checks. |
docs/whats_new.rst |
Documents the bug fix and the new default padding behavior. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| if "chs_info" in kwargs and "n_chans" not in kwargs: | ||
| chs_info = kwargs["chs_info"] | ||
| kwargs["n_chans"] = None if chs_info is None else len(chs_info) | ||
| elif "n_chans" in kwargs and "chs_info" not in kwargs: | ||
| kwargs["chs_info"] = None | ||
| if ("n_times" in kwargs or "sfreq" in kwargs) and ( | ||
| "input_window_seconds" not in kwargs | ||
| ): |
| # 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, | ||
| ) |
| X_patch: Tensor | ||
| [batch, n_chans, n_times//patch_size, patch_size] | ||
| """ | ||
| x = self.tokenizer._prepare_input(x) |
| x_signal = self.tokenizer._prepare_input(x_signal) | ||
| if mask is not None: | ||
| mask = self.tokenizer._prepare_input(mask) |
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 (7)
The pinning logic keys off kwarg presence rather than whether the caller actually provided a… · New The current pinning logic triggers even when the caller explicitly passes geometry keys withNone… This introduces a dependency onPatchTokenizer._prepare_input, which is a private API (leading… · Newstate_dict()includes buffers as well as parameters, so this assertion can fail if… · New This calls a privatePatchTokenizermethod (_prepare_input), which is brittle (private APIs can… Similar to LUNA, this relies onPatchTokenizer._prepare_input, which is a private method. To…self.rearrangeis now aPatchTokenizer, not an einopsRearrange, so the attribute name is…
| # config cannot supply it. | ||
| if "chs_info" in kwargs and "n_chans" not in kwargs: | ||
| chs_info = kwargs["chs_info"] | ||
| kwargs["n_chans"] = None if chs_info is None else len(chs_info) | ||
| elif "n_chans" in kwargs and "chs_info" not in kwargs: | ||
| kwargs["chs_info"] = None | ||
| if ("n_times" in kwargs or "sfreq" in kwargs) and ( | ||
| "input_window_seconds" not in kwargs |
| 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()) |
Conflict resolution: - docs/whats_new.rst: kept both the braindecode#1232 from_pretrained geometry-kwargs bug-fix entry and master's braindecode#1207/braindecode#1212 Deep4Net/ShallowFBCSPNet bug-fix entries. - test/unit_tests/models/test_eegdino.py: kept both the braindecode#1232 geometry-kwargs test block and master's braindecode#1194 qkv-hook attention test block.
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: 3
Open (9)
As written, a caller explicitly passingchs_info=None(common pattern to mean “don’t override /… · New The pinning logic keys off kwarg presence rather than whether the caller actually provided a… The current pinning logic triggers even when the caller explicitly passes geometry keys withNone…state_dict()includes buffers as well as parameters, so this assertion can fail if… This introduces a dependency onPatchTokenizer._prepare_input, which is a private API (leading… This calls a privatePatchTokenizermethod (_prepare_input), which is brittle (private APIs can… Similar to LUNA, this relies onPatchTokenizer._prepare_input, which is a private method. To…self.rearrangeis now aPatchTokenizer, not an einopsRearrange, so the attribute name is… The sentence “The derived argument is now pinned from the caller's one.” is grammatically awkward… · New
| if "chs_info" in kwargs and "n_chans" not in kwargs: | ||
| chs_info = kwargs["chs_info"] | ||
| kwargs["n_chans"] = None if chs_info is None else len(chs_info) |
| derived argument is now pinned from the caller's one. This unblocks loading EEGPT, | ||
| STEEGFormer, Brant and MVPFormer checkpoints on a montage other than their | ||
| pretraining dataset's (:gh:`1232` by `Bruno Aristimunha`_). |
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## master #1232 +/- ##
=======================================
Coverage 87.99% 87.99%
=======================================
Files 151 151
Lines 17628 17638 +10
=======================================
+ Hits 15511 15521 +10
Misses 2117 2117 🚀 New features to boost your workflow:
|
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: 4
Open (12)
Settingn_chans=Nonewhen the caller explicitly passeschs_info=Nonecan unintentionally… · New As written, a caller explicitly passingchs_info=None(common pattern to mean “don’t override /… The pinning logic keys off kwarg presence rather than whether the caller actually provided a… The current pinning logic triggers even when the caller explicitly passes geometry keys withNone… Derivingn_chansvialen(chs_info)is not robust ifchs_infocan be anmne.Info(or similar… · New The newfrom_pretrainedbehavior has an important edge case when callers passchs_info=None… · Newstate_dict()includes buffers as well as parameters, so this assertion can fail if… This introduces a dependency onPatchTokenizer._prepare_input, which is a private API (leading… This calls a privatePatchTokenizermethod (_prepare_input), which is brittle (private APIs can… Similar to LUNA, this relies onPatchTokenizer._prepare_input, which is a private method. To…self.rearrangeis now aPatchTokenizer, not an einopsRearrange, so the attribute name is… The sentence “The derived argument is now pinned from the caller's one.” is grammatically awkward…
| # derived geometry arg so it cannot clash with the caller's one | ||
| # (e.g. caller chs_info vs saved n_chans) in __init__'s checks. | ||
| if "chs_info" in kwargs and "n_chans" not in kwargs: | ||
| chs_info = kwargs["chs_info"] | ||
| kwargs["n_chans"] = None if chs_info is None else len(chs_info) |
| chs_info = kwargs["chs_info"] | ||
| kwargs["n_chans"] = None if chs_info is None else len(chs_info) |
| assert EEGDINO.from_pretrained(save_dir, n_outputs=6)(x).shape == (1, 6) | ||
|
|
||
|
|
||
| def test_from_pretrained_takes_geometry_from_the_caller_not_the_config(tmp_path): |
…rs on_non_divisible (#1250) * FIX from_pretrained with explicit None geometry; LaBraM decoder honours on_non_divisible - from_pretrained(d, chs_info=None) / (d, n_chans=None) load the saved geometry again: the pinning checks the value, not the key (#1232). - Labram(neural_tokenizer=False) pads / crops / raises on a window not divisible by patch_size, like the tokenizer mode (#1226). - PatchTokenizer.prepare_input is public (_prepare_input kept as alias); LaBraM and LUNA call it. * MAINT PatchTokenizer: drop the unused _prepare_input alias * DOC whats_new entry for #1250 * FIX from_pretrained: drop explicit None geometry so config.json fills it



Summary
Model.from_pretrained(repo, chs_info=my_19_channels)fails withn_chans=62 different from chs_info … lengthfor EEGPT, STEEGFormer, Brant, MVPFormer (and BIOT/BENDR/… on any montage). Found by thefrom_pretrained-level audit in #1227 (66 of 100 failing cells).Root cause is not in the models:
PyTorchModelHubMixin.from_pretrainedfills every__init__argument the caller omitted fromconfig.json, so the caller'schs_infomeets the checkpoint'sn_chans, andEEGModuleMixin.__init__'s consistency check fires. Same for a caller'sn_times/sfreqagainst the savedinput_window_seconds.Stacked on #1226 (review the last commit). Part of #1227.
Changes
EEGModuleMixin.from_pretrained(HF-only block): one guard before delegating —chs_infogiven → pinn_chans=len(chs_info);n_chansgiven → pinchs_info=None;n_timesorsfreqgiven → pininput_window_seconds=None(derivable). The config can no longer supply the derived argument.test_eegdino.py: localsave_pretrained→from_pretrainedwith a differentchs_info,n_chans,n_times,sfreq(RED on master with the exact production error, GREEN here).docs/whats_new.rst(Bug fixes).Testing
pytest test/unit_tests/models/test_eegdino.py test/unit_tests/models/test_signal_jepa.py→ 40 passed.pytest test/unit_tests/models/test_huggingface.py test/unit_tests/models/test_base.py test/unit_tests/models/test_return_features.py→ 372 passed, 25 skipped.pytest test/unit_tests/models/test_foundation_models.py -k "pretrained or hub or roundtrip"→ 4 passed, 6 skipped.pre-commit run --files <changed>clean. Not run: full suite (CI).Regression test
Style checks recorded
docs/whats_new.rstupdatedNotes for reviewers
Behaviour change only when the caller passes geometry: previously the omitted sibling came from the config (and usually collided); now it is derived/
None. Omitting all geometry still loads exactly the saved config. Whether the weights fit the new geometry is the model's business (fixed-order models still raise, as they should).