Skip to content

Keep EEGPT channel IDs out of the state dict so pretrained weights load on any montage - #1195

Merged
bruAristimunha merged 3 commits into
braindecode:masterfrom
bruAristimunha:fix/eegpt-chans-id-buffer
Oct 4, 2026
Merged

bruAristimunha merged 3 commits into
braindecode:masterfrom
bruAristimunha:fix/eegpt-chans-id-buffer

Conversation

@bruAristimunha

Copy link
Copy Markdown
Collaborator

Summary

EEGPT.chans_id is derived from chs_info when 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 released braindecode/eegpt-pretrained weights:

  1. The checkpoint does not load unless the model has exactly the same 62 channels. The checkpoint stores 62 IDs, so any other montage fails with size mismatch for chans_id. This includes EEGPT.from_pretrained("braindecode/eegpt-pretrained") with no arguments: the default channel projection (conv1d_constraint) gives 19 IDs.
  2. With the same channel count, loading silently gives the model the wrong channel IDs. The checkpoint's 62 IDs replace the model's own, so a montage with the same channels in a different order looks up the wrong channel embedding for every channel, without any warning.

Reproduced with the released checkpoint (revision e41cb3ae):

from_pretrained(...) master this PR
no arguments (default channel projection) size mismatch [1, 62] vs [1, 19] loads
chan_proj_type="none", saved 62 channels loads loads, same logits and features
chan_proj_type="none", 9 channels size mismatch [1, 62] vs [1, 9] loads, IDs of the 9 channels
chan_proj_type="none", the 62 channels reversed loads, but with the checkpoint's IDs [0, 1, 2, …] loads, IDs [61, 60, 59, …]

Changes

  • Register chans_id with persistent=False, so it is no longer written to state dicts.
  • EEGPT._load_from_state_dict drops a chans_id key if present, so checkpoints saved before this change, including the released weights, load under strict=True.
  • The docstring notes that the IDs are rebuilt from the montage and are not stored.

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 (with chans_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 of chan_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.rst passed

  • docs/whats_new.rst updated

Notes for reviewers

  • Checkpoints saved after this change no longer contain chans_id. Older braindecode versions will report it as missing when loading them; with strict=False (the from_pretrained default) they rebuild it from chs_info.
  • A model built without chs_info uses sequential IDs, as before. The checkpoint no longer overrides them; before, it failed to load because the shapes differed ((n_chans,) vs (1, 62)).
  • Found while fine-tuning EEGPT on several montages in NeuralBench, which currently works around it by subclassing EEGPT.load_state_dict.

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.
Copilot AI balanced review requested due to automatic review settings September 30, 2026 07:57
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.

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.

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 Low severity

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_id a non-persistent buffer so it is excluded from state_dict.
  • Drop legacy chans_id entries during EEGPT checkpoint 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.

Comment on lines +399 to +404
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)
Comment on lines +699 to +700
["C3", "CZ", "C4"], # fewer channels than the checkpoint
["P4", "PZ", "CP2", "C4", "CZ", "C3"], # same channels, other order
Comment on lines +692 to +707
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"])
Copilot AI balanced review requested due to automatic review settings September 30, 2026 08:14

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.

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 Medium severity · 3 Low severity

Open (5)

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

codecov Bot commented Sep 30, 2026

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 87.72%. Comparing base (017809e) to head (1c143de).

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:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@bruAristimunha
bruAristimunha merged commit ccbbea3 into braindecode:master Oct 4, 2026
14 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.

2 participants