Skip to content

Improve pre-trained channel embeddings handling in Signal-JEPA - #991

Merged
PierreGtch merged 20 commits into
braindecode:masterfrom
PierreGtch:sjepa_channel-embeddings
Apr 18, 2026
Merged

PierreGtch merged 20 commits into
braindecode:masterfrom
PierreGtch:sjepa_channel-embeddings

Conversation

@PierreGtch

@PierreGtch PierreGtch commented Apr 17, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Makes it easy to load pre-trained Signal-JEPA channel positional encoding weights when fine-tuning on a channel set different from the 62-channel pre-training set.

Key changes

  • Add channel_embedding parameter to SignalJEPA and SignalJEPA_Contextual, with two modes:
    • "scratch" (default): train channel embeddings from scratch for the user's chs_info.
    • "pretrain_aligned": reuse the 62-channel pre-training layout so pre-trained embedding rows can be loaded directly, then index into the subset matching the user's chs_info (case-insensitive).
  • Publish two HuggingFace checkpoints:
    • braindecode/signal-jepa — full checkpoint including the 62-row channel embedding matrix (pretrain_aligned ready).
    • braindecode/signal-jepa_without-chans — backbone only, for users who want to train channel embeddings from scratch.
  • Update all downstream variants (SignalJEPA_PreLocal, SignalJEPA_Contextual, SignalJEPA_PostLocal) to honour the new parameter and to correctly transfer the pre-trained _PosEncoder buffer in from_pretrained.
  • Update docstrings with Pretrained Weights / Usage sections, and refresh the fine-tuning and model-loading examples to use the new hub IDs.

Breaking changes

  • SignalJEPA.forward and SignalJEPA_Contextual.forward no longer accept ch_idxs; channel indices are now resolved internally from chs_info.

Tests

  • 30 unit tests in test/unit_tests/models/test_signal_jepa.py covering the resolver helper, _PosEncoder buffer, channel_embedding modes, forward-API change, from_pretrained instance transfer and HF round-trip.
  • Integration tests in test/integration_tests/test_pretrained_hub_models.py updated for the new hub IDs and downstream kwargs.

Test plan

  • Unit tests pass (pytest test/unit_tests/models/test_signal_jepa.py)
  • Integration tests pass against live HuggingFace hub
  • Fine-tuning example runs end-to-end with the new hub ID
  • pre-commit run --all-files passes

🤖 Generated with Claude Code

PierreGtch and others added 10 commits April 17, 2026 00:40
Temporary constant with 62 EEG channels named from a standard extended
10-20 layout. Locations are ALL zero placeholders — they MUST be
replaced with the authoritative coordinates from the Signal-JEPA
pre-training run before publishing the HuggingFace checkpoint. The
TODO(PRETRAIN_COORDS) comment marks the block for that update.

Unblocks Tasks 2-7 of the channel_embedding implementation plan
(docs/superpowers/plans/2026-04-16-signal-jepa-channel-embedding.md).

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
Pure function that maps (channel_embedding, chs_info) to the tuple
(effective_chs_info, channel_locations, ch_idxs). All validation lives
here so the model constructor can stay concise and the logic is
unit-testable in isolation.

Case-insensitive name matching for 'pretrain_aligned' mode. Explicit
ValueError with actionable fallback message when a user channel is
outside the pre-training set.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
_PosEncoder now requires ch_idxs at __init__ and stores it via
register_buffer(..., persistent=False). forward() uses this buffer when
its ch_idxs kwarg is None. Renames the init kwarg ch_locs ->
channel_locations for clarity (the locations describe the table rows,
not the input channels).

Non-persistent is deliberate: ch_idxs is derived from chs_info +
channel_embedding at __init__, so saving it in state_dict would
duplicate config.json and conflict with user-overridden chs_info at
from_pretrained time.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
New kwarg channel_embedding: Literal['scratch', 'pretrain_aligned'].
Runs the resolver before super().__init__, so when the user passes
chs_info=None with channel_embedding='pretrain_aligned', the 62
pre-training channels become the effective chs_info (and hence n_chans).

Stores self._channel_embedding so build_model_config serializes it
into config.json for HF round-trips.

The feature_encoder-only path (_init_transformer=False, used by
PreLocal and PostLocal via _BaseSignalJEPA) skips the resolver
entirely.

SignalJEPA.__init__ now forwards the new kwarg to its super. Other
subclasses (Contextual, PreLocal, PostLocal) will be updated in Task 5.

Fixed _PRETRAIN_CHS_INFO locations from all-zeros to realistic 3D
coordinates to avoid division-by-zero in positional encoding.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
Replace the synthetic topographic coordinates (introduced during Task 4
implementation to work around division-by-zero in _ChannelEmbedding
reset_parameters when all locs are identical) with obviously-fake
systematic placeholders: loc = [0.001*i, 0.0, 0.0] for the i-th channel.

These are clearly NOT real EEG coordinates, so there is no risk of the
model being silently used with incorrect locations. The TODO(PRETRAIN_COORDS)
comment continues to mark the block for replacement before publishing
the HuggingFace checkpoint.

Channel names and order are unchanged.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
Both public classes now accept channel_embedding at __init__ and
forward it to _BaseSignalJEPA. The ch_idxs argument is removed from
their forward() signatures — it is now resolved once at __init__ and
stored in _PosEncoder.default_ch_idxs (non-persistent buffer).

PreLocal and PostLocal are unchanged: they don't build a pos_encoder.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
The method set_fixed_ch_names did not exist on _PosEncoder; the call
was guarded by 'if chs_info is not None', so it silently produced an
AttributeError whenever a caller used the transfer-learning path of
SignalJEPA_Contextual.from_pretrained with chs_info set.

The functionality it attempted to provide (binding the model to a
specific channel list) is now covered by the channel_embedding
parameter plus the chs_info passed directly to the constructor.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
Verifies that:
  - channel_embedding round-trips through config.json (via
    build_model_config + track_model_init_kwargs)
  - default_ch_idxs is NOT in state_dict (non-persistent buffer)
  - user-overridden chs_info at from_pretrained time correctly rebuilds
    the ch_idxs mapping while still loading the 62-row embedding weights

Gated on huggingface_hub availability.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
Replaces the systematic placeholder locs (0.001 * i on the x axis)
with the authoritative 62 channels used to pre-train Signal-JEPA.

The channel set is derived from the MOABB Lee2019_SSVEP dataset,
subject 1, session '1', run '1train': the first 62 ch_names (the
last 5 of the dataset's 67 channels are 4 EMG + 1 stim). Positions
come from MNE's standard_1005 montage.

The channel NAMES and ORDER in this list are now the definitive
contract with the HuggingFace checkpoint braindecode/signal-jepa;
changing either would silently mismatch the _ChannelEmbedding row
order of the published weights.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
When SignalJEPA_Contextual.from_pretrained(src_instance, chs_info=...)
was called with a pretrain_aligned source and a different destination
chs_info, the deepcopy()'d pos_encoder carried the source's
default_ch_idxs unchanged — mapping the user's channels to the WRONG
rows of the 62-entry embedding table. No runtime error, just silently
incorrect results.

Fix: after the deepcopy, re-resolve default_ch_idxs against the
destination chs_info when the source was built in pretrain_aligned
mode, and propagate _channel_embedding onto the new model so
subsequent save_pretrained round-trips use the right config.

Add regression test test_from_pretrained_transfer_pretrain_aligned_different_subset
that constructs a source on pretrain rows [0, 1, 2] and transfers to
destination rows [5, 0, 10], asserting the recomputed buffer and
preserved embedding weights.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
Copilot AI review requested due to automatic review settings April 17, 2026 14:38

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.

Pull request overview

This PR improves Signal-JEPA’s handling of channel embeddings when using pre-trained checkpoints by introducing an explicit “pretrain-aligned” channel set/mapping and making the channel index mapping an internal (non-persistent) buffer rather than a forward argument.

Changes:

  • Added _PRETRAIN_CHS_INFO (62-channel canonical pretraining order) plus _resolve_channel_embedding_config() and a channel_embedding init parameter to support "scratch" vs "pretrain_aligned" behavior.
  • Updated _PosEncoder to store a non-persistent default_ch_idxs buffer and updated SignalJEPA/SignalJEPA_Contextual to drop the ch_idxs kwarg from forward().
  • Expanded unit tests to cover config resolution, buffer persistence behavior, transfer learning via from_pretrained, and HF save/load round-trips.

Reviewed changes

Copilot reviewed 2 out of 2 changed files in this pull request and generated 4 comments.

File Description
braindecode/models/signal_jepa.py Implements pretrain-aligned channel embedding configuration, adds canonical pretraining channel list, updates positional encoder behavior, and adjusts transfer logic.
test/unit_tests/models/test_signal_jepa.py Adds/updates tests validating the new channel embedding modes, buffer behavior, API change (no ch_idxs in forward), transfer behavior, and HF serialization.

💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.

Comment thread braindecode/models/signal_jepa.py
Comment thread braindecode/models/signal_jepa.py Outdated
Comment thread braindecode/models/signal_jepa.py Outdated
Comment thread test/unit_tests/models/test_signal_jepa.py Outdated
'hub_id' is not a real kwarg; the first positional arg of from_pretrained
is pretrained_model_name_or_path. Rephrase as a descriptive 'checkpoint'
label to avoid pointing users to a non-existent parameter.
build_model_config captures __init__ kwargs directly, not instance
attributes, so the placement of this assignment relative to super().__init__
is irrelevant for config tracking. The assignment speaks for itself.
The list comprehension was 96 chars, over the 88-char ruff line-length
configured in pyproject.toml. pre-commit's ruff hooks don't currently
run on test/, which is why it slipped through.
Copilot AI review requested due to automatic review settings April 17, 2026 16:28

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.

Pull request overview

Copilot reviewed 6 out of 6 changed files in this pull request and generated 6 comments.


💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.

Comment thread docs/whats_new.rst
Comment thread braindecode/models/signal_jepa.py Outdated
Comment thread braindecode/models/signal_jepa.py Outdated
Comment thread braindecode/models/signal_jepa.py Outdated
Comment thread braindecode/models/signal_jepa.py Outdated
Comment thread examples/advanced_training/plot_finetune_foundation_model.py
PierreGtch and others added 3 commits April 18, 2026 20:43
_ChannelEmbedding sliced loc[3:6] to extract electrode positions from
MNE's 12-element loc arrays, but loc[3:6] is the *reference* electrode
position (typically [0, 0, 0] for standard montages), not the actual
electrode position at loc[0:3]. This collapsed every channel to the
origin, produced max_abs_coordinate=0, and caused a 0/0 divide-by-zero
in _pos_encode_contineous that filled the channel embedding weight
(and therefore the entire forward pass) with NaN.

Slice loc[:3] instead. Add a regression test that constructs the
channel embedding from a real mne.Info with standard_1020 and checks
that neither the embedding weight nor the forward output contains NaN.
Co-authored-by: Copilot <[email protected]>
Copilot AI review requested due to automatic review settings April 18, 2026 19:22

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.

Pull request overview

Copilot reviewed 6 out of 6 changed files in this pull request and generated 2 comments.

Comments suppressed due to low confidence (1)

braindecode/models/signal_jepa.py:1344

  • If all provided channel coordinates are identical (common when MNE montage is missing and loc[:3] is all zeros), then global_min == global_max == 0 and max_abs_coordinate becomes 0. This causes _pos_encode_contineous(..., x_max=0) to divide by zero and initialize the embedding weights with NaNs. Consider guarding against max_abs_coordinate == 0 (raise a clear error suggesting setting a montage / providing non-degenerate locations, or fall back to a safe initialization).
        channel_mins, channel_maxs = zip(*self.coordinate_ranges)
        global_min = min(channel_mins)
        global_max = max(channel_maxs)
        self.max_abs_coordinate = max(abs(global_min), abs(global_max))
        self.embedding_dim_per_coordinate = embedding_dim // len(self.coordinate_ranges)

💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.

Comment on lines 123 to 129
model = SignalJEPA_PreLocal.from_pretrained(
"braindecode/SignalJEPA-PreLocal-pretrained",
"braindecode/signal-jepa_without-chans",
n_chans=19,
n_times=256,
n_outputs=len(classes),
strict=False,
)

Copilot AI Apr 18, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This section still reads as if the Hub checkpoint already includes the downstream classification layers, but the code now loads the shared SSL checkpoint (signal-jepa_without-chans) and uses strict=False (implying the task head is initialized locally and not coming from the checkpoint). Please update the surrounding tutorial text to explain that only the pretrained backbone weights are loaded here and the classification head is newly initialized (hence strict=False).

Copilot uses AI. Check for mistakes.
Comment on lines +32 to +45
{"ch_name": "Fp1", "loc": [-0.0294367, 0.0839171, -0.00699]},
{"ch_name": "Fp2", "loc": [0.0298723, 0.0848959, -0.00708]},
{"ch_name": "F7", "loc": [-0.0702629, 0.0424743, -0.01142]},
{"ch_name": "F3", "loc": [-0.0502438, 0.0531112, 0.042192]},
{"ch_name": "Fz", "loc": [0.0003122, 0.058512, 0.066462]},
{"ch_name": "F4", "loc": [0.0518362, 0.0543048, 0.040814]},
{"ch_name": "F8", "loc": [0.0730431, 0.0444217, -0.012]},
{"ch_name": "FC5", "loc": [-0.0772149, 0.0186433, 0.02446]},
{"ch_name": "FC1", "loc": [-0.0340619, 0.0260111, 0.079987]},
{"ch_name": "FC2", "loc": [0.0347841, 0.0264379, 0.078808]},
{"ch_name": "FC6", "loc": [0.0795341, 0.0199357, 0.024438]},
{"ch_name": "T7", "loc": [-0.0841611, -0.0160187, -0.009346]},
{"ch_name": "C3", "loc": [-0.0653581, -0.0116317, 0.064358]},
{"ch_name": "Cz", "loc": [0.0004009, -0.009167, 0.100244]},

Copilot AI Apr 18, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

_PRETRAIN_CHS_INFO stores loc as Python lists. Because EEGModuleMixin.__init__ treats a list-of-dicts with loc as list as serialized Hub config, constructing SignalJEPA(channel_embedding='pretrain_aligned', chs_info=None) will trigger _deserialize_chs_info() and emit a warning every time. Consider storing loc as tuples (or np arrays) in _PRETRAIN_CHS_INFO to avoid the noisy/misleading warning in the common pretrained path.

Suggested change
{"ch_name": "Fp1", "loc": [-0.0294367, 0.0839171, -0.00699]},
{"ch_name": "Fp2", "loc": [0.0298723, 0.0848959, -0.00708]},
{"ch_name": "F7", "loc": [-0.0702629, 0.0424743, -0.01142]},
{"ch_name": "F3", "loc": [-0.0502438, 0.0531112, 0.042192]},
{"ch_name": "Fz", "loc": [0.0003122, 0.058512, 0.066462]},
{"ch_name": "F4", "loc": [0.0518362, 0.0543048, 0.040814]},
{"ch_name": "F8", "loc": [0.0730431, 0.0444217, -0.012]},
{"ch_name": "FC5", "loc": [-0.0772149, 0.0186433, 0.02446]},
{"ch_name": "FC1", "loc": [-0.0340619, 0.0260111, 0.079987]},
{"ch_name": "FC2", "loc": [0.0347841, 0.0264379, 0.078808]},
{"ch_name": "FC6", "loc": [0.0795341, 0.0199357, 0.024438]},
{"ch_name": "T7", "loc": [-0.0841611, -0.0160187, -0.009346]},
{"ch_name": "C3", "loc": [-0.0653581, -0.0116317, 0.064358]},
{"ch_name": "Cz", "loc": [0.0004009, -0.009167, 0.100244]},
{"ch_name": "Fp1", "loc": (-0.0294367, 0.0839171, -0.00699)},
{"ch_name": "Fp2", "loc": (0.0298723, 0.0848959, -0.00708)},
{"ch_name": "F7", "loc": (-0.0702629, 0.0424743, -0.01142)},
{"ch_name": "F3", "loc": (-0.0502438, 0.0531112, 0.042192)},
{"ch_name": "Fz", "loc": (0.0003122, 0.058512, 0.066462)},
{"ch_name": "F4", "loc": (0.0518362, 0.0543048, 0.040814)},
{"ch_name": "F8", "loc": (0.0730431, 0.0444217, -0.012)},
{"ch_name": "FC5", "loc": (-0.0772149, 0.0186433, 0.02446)},
{"ch_name": "FC1", "loc": (-0.0340619, 0.0260111, 0.079987)},
{"ch_name": "FC2", "loc": (0.0347841, 0.0264379, 0.078808)},
{"ch_name": "FC6", "loc": (0.0795341, 0.0199357, 0.024438)},
{"ch_name": "T7", "loc": (-0.0841611, -0.0160187, -0.009346)},
{"ch_name": "C3", "loc": (-0.0653581, -0.0116317, 0.064358)},
{"ch_name": "Cz", "loc": (0.0004009, -0.009167, 0.100244)},

Copilot uses AI. Check for mistakes.
@PierreGtch
PierreGtch merged commit 6332ba0 into braindecode:master Apr 18, 2026
17 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