Repository navigation
Improve pre-trained channel embeddings handling in Signal-JEPA - #991
Conversation
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]>
There was a problem hiding this comment.
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 achannel_embeddinginit parameter to support"scratch"vs"pretrain_aligned"behavior. - Updated
_PosEncoderto store a non-persistentdefault_ch_idxsbuffer and updatedSignalJEPA/SignalJEPA_Contextualto drop thech_idxskwarg fromforward(). - 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.
'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.
There was a problem hiding this comment.
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.
_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]>
There was a problem hiding this comment.
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), thenglobal_min == global_max == 0andmax_abs_coordinatebecomes 0. This causes_pos_encode_contineous(..., x_max=0)to divide by zero and initialize the embedding weights with NaNs. Consider guarding againstmax_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.
| 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, | ||
| ) |
There was a problem hiding this comment.
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).
| {"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]}, |
There was a problem hiding this comment.
_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.
| {"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)}, |
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
channel_embeddingparameter toSignalJEPAandSignalJEPA_Contextual, with two modes:"scratch"(default): train channel embeddings from scratch for the user'schs_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'schs_info(case-insensitive).braindecode/signal-jepa— full checkpoint including the 62-row channel embedding matrix (pretrain_alignedready).braindecode/signal-jepa_without-chans— backbone only, for users who want to train channel embeddings from scratch.SignalJEPA_PreLocal,SignalJEPA_Contextual,SignalJEPA_PostLocal) to honour the new parameter and to correctly transfer the pre-trained_PosEncoderbuffer infrom_pretrained.Breaking changes
SignalJEPA.forwardandSignalJEPA_Contextual.forwardno longer acceptch_idxs; channel indices are now resolved internally fromchs_info.Tests
test/unit_tests/models/test_signal_jepa.pycovering the resolver helper,_PosEncoderbuffer,channel_embeddingmodes, forward-API change,from_pretrainedinstance transfer and HF round-trip.test/integration_tests/test_pretrained_hub_models.pyupdated for the new hub IDs and downstream kwargs.Test plan
pytest test/unit_tests/models/test_signal_jepa.py)pre-commit run --all-filespasses🤖 Generated with Claude Code