Repository navigation
Keep EEGPT channel IDs out of the state dict so pretrained weights load on any montage - #1195
Conversation
chans_id is derived from chs_info at construction but was a persistent buffer, so it was saved in checkpoints and loaded back: - the released braindecode/eegpt-pretrained weights (62 IDs) failed to load with a size mismatch on any other montage, and with the default channel projection (19 IDs), i.e. EEGPT.from_pretrained() with no arguments; - on a montage with the same channel count, loading silently replaced the model's channel IDs with the checkpoint's. Register chans_id as non-persistent and drop the key when loading older checkpoints that still carry it.
braindecode#1194 adds its entry at the top of the same list; keeping the two apart lets them merge in either order without a conflict.
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 (3)
What changed in this PR
Ensure EEGPT pretrained checkpoints load across different channel montages by preventing montage-derived chans_id from being persisted/overridden during checkpoint load.
Changes:
- Make
EEGPT.chans_ida non-persistent buffer so it is excluded fromstate_dict. - Drop legacy
chans_identries duringEEGPTcheckpoint loading to keep a model’s montage-derived IDs. - Add regression tests and a release-note entry documenting the new behavior.
| File | Description |
|---|---|
braindecode/models/eegpt.py |
Marks chans_id as non-persistent and ignores legacy chans_id keys on load to avoid montage mismatch/override. |
test/unit_tests/models/test_models.py |
Adds regression tests ensuring chans_id isn’t saved and legacy checkpoints don’t override montage-derived IDs. |
docs/whats_new.rst |
Documents the checkpoint-loading behavior change for EEGPT channel IDs. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs): | ||
| # Checkpoints saved before ``chans_id`` became non-persistent (including | ||
| # the released EEGPT weights) still carry it. Drop it, so the IDs keep | ||
| # matching this model's montage and a different montage can load. | ||
| state_dict.pop(prefix + "chans_id", None) | ||
| super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) |
| ["C3", "CZ", "C4"], # fewer channels than the checkpoint | ||
| ["P4", "PZ", "CP2", "C4", "CZ", "C3"], # same channels, other order |
| model = _eegpt(["C3", "CZ", "C4"]) | ||
| assert "chans_id" not in model.state_dict() | ||
|
|
||
|
|
||
| @pytest.mark.parametrize( | ||
| "target_names", | ||
| [ | ||
| ["C3", "CZ", "C4"], # fewer channels than the checkpoint | ||
| ["P4", "PZ", "CP2", "C4", "CZ", "C3"], # same channels, other order | ||
| ], | ||
| ) | ||
| def test_eegpt_checkpoint_keeps_the_model_channel_ids(target_names): | ||
| # Checkpoints saved before chans_id became non-persistent, like the released | ||
| # EEGPT weights, still store it. Loading one must neither fail on another | ||
| # montage nor replace the IDs of the model's channels by the checkpoint's. | ||
| source = _eegpt(["C3", "CZ", "C4", "CP2", "PZ", "P4"]) |
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 (5)
This mutates the caller-providedstate_dictin-place, which can be surprising if the same dict is… · New Theall(...)assertion is vacuously true whenmissing_keysis empty, so this test could pass… · New The test channel names use uppercaseCZ/PZ. Ifprepare_chan_idsever becomes case-sensitive… The test channel names use uppercaseCZ/PZ. Ifprepare_chan_idsever becomes case-sensitive… Consider matchingtorch.nn.Module._load_from_state_dict’s full explicit signature (e.g.,…
| def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs): | ||
| # Checkpoints saved before ``chans_id`` became non-persistent (including | ||
| # the released EEGPT weights) still carry it. Drop it, so the IDs keep | ||
| # matching this model's montage and a different montage can load. |
|
|
||
| incompatible = target.load_state_dict(state_dict, strict=False) | ||
|
|
||
| assert not incompatible.unexpected_keys |
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## master #1195 +/- ##
=======================================
Coverage 87.71% 87.72%
=======================================
Files 149 149
Lines 17433 17436 +3
=======================================
+ Hits 15292 15295 +3
Misses 2141 2141 🚀 New features to boost your workflow:
|


Summary
EEGPT.chans_idis derived fromchs_infowhen the model is built, but it is registered as a persistent buffer. It is therefore saved in checkpoints and copied back on load, which causes two failures with the releasedbraindecode/eegpt-pretrainedweights:size mismatch for chans_id. This includesEEGPT.from_pretrained("braindecode/eegpt-pretrained")with no arguments: the default channel projection (conv1d_constraint) gives 19 IDs.Reproduced with the released checkpoint (revision
e41cb3ae):from_pretrained(...)[1, 62]vs[1, 19]chan_proj_type="none", saved 62 channelschan_proj_type="none", 9 channels[1, 62]vs[1, 9]chan_proj_type="none", the 62 channels reversed[0, 1, 2, …][61, 60, 59, …]Changes
chans_idwithpersistent=False, so it is no longer written to state dicts.EEGPT._load_from_state_dictdrops achans_idkey if present, so checkpoints saved before this change, including the released weights, load understrict=True.Testing
New tests in
test_models.py:test_eegpt_chans_id_is_not_saved;test_eegpt_checkpoint_keeps_the_model_channel_ids, which loads an old-style state dict (withchans_id) into a model with fewer channels and into one with the same channels in another order. It checks that the model keeps its own IDs and that the pretrained channel-embedding table still loads;test_eegpt_checkpoint_loads_with_channel_projection, where the only missing keys are those ofchan_proj.All 4 cases fail on master, with the size mismatch or the replaced IDs.
Released checkpoint: for the saved 62-channel montage with
chan_proj_type="none"and a seeded head, logits and encoder features are bit-identical to master. The classifier probe is not in the checkpoint.pytest test_models.py test_integration.py test_interpolated.py test_huggingface.py test_return_features.py -k eegpt: 66 passed, 7 skipped. The skips predate this PR: TorchScript, plus activation, dropout and embedding checks.New or changed behavior is covered by relevant regression tests
Style checks recorded:
pre-commit run --files braindecode/models/eegpt.py test/unit_tests/models/test_models.py docs/whats_new.rstpasseddocs/whats_new.rstupdatedNotes for reviewers
chans_id. Older braindecode versions will report it as missing when loading them; withstrict=False(thefrom_pretraineddefault) they rebuild it fromchs_info.chs_infouses sequential IDs, as before. The checkpoint no longer overrides them; before, it failed to load because the shapes differed ((n_chans,)vs(1, 62)).EEGPT.load_state_dict.