Repository navigation
[ENH] Experimental channel interpolation for foundation models - #993
Conversation
…ng/loading it in the state dict.
- Expand mode docstring to be self-contained (no 'see design spec'). - Document permissive handling of missing/None 'kind' in _assert_eeg_only. - Replace defensive assert with RuntimeError in _compute_interpolation_matrix_mne. - Add Parameters section to _compute_interpolation_matrix_mne docstring. - Add monkeypatch test proving MNE is not called on full name coverage.
Labram now validates that chs_info (when provided) matches LABRAM_CHANNEL_ORDER exactly, removing the old channel-selection / on_unknown_chs machinery. Arbitrary channel sets are served by the new InterpolatedLaBraM wrapper (InterpolatedModel(Labram, _LABRAM_TARGET_CHS_INFO)) which is exported from braindecode.models and registered in models_mandatory_parameters.
…__, skip signature inspection tests, register summary rows
…eatures Labram test
…melcase
Without the explicit name, the factory defaults to
"Interpolated{model_cls.__name__}" which produces "InterpolatedLabram" from
the Labram class. models_dict and summary.csv both key on the exposed
"InterpolatedLaBraM" spelling, so the mismatch broke
test_make_model_config_instantiation and test_completeness_summary_table.
The Interpolated* factory declares n_chans explicitly so skorch can
discover it, but intentionally does not use it (the backbone's channel
count is derived from target_chs_info). When save_pretrained serializes
the __init__ kwargs, n_chans is emitted as null. The test must treat
config.get('n_chans') is None the same as an absent key, and fall back
to comparing against len(config['chs_info']).
Also relax test_config_contains_all_parameters to require only
braindecode_version plus one of {n_chans, chs_info}, since Interpolated*
models do not re-expose sfreq/n_outputs at the top level of the config.
There was a problem hiding this comment.
Pull request overview
Adds an experimental channel-interpolation mechanism to project arbitrary user channel sets into a foundation model’s canonical channel space (via an MNE-derived interpolation matrix), and ships interpolated variants for LaBraM, SignalJEPA, and BIOT. It also updates LaBraM to require the canonical 128-channel order and adjusts docs/tests accordingly.
Changes:
- Introduce
ChannelInterpolationLayerplus anInterpolatedModel(...)factory to prepend channel projection to existing backbones. - Ship
InterpolatedLaBraM,InterpolatedSignalJEPA, andInterpolatedBIOTwrappers and wire them into the model registry/docs. - Refactor
Labramto enforce canonical channel naming/order (breaking change) and update unit/integration/HF-config tests.
Reviewed changes
Copilot reviewed 19 out of 19 changed files in this pull request and generated 4 comments.
Show a summary per file
| File | Description |
|---|---|
braindecode/modules/interpolation.py |
New interpolation layer + MNE matrix builder and validations. |
braindecode/modules/__init__.py |
Exports ChannelInterpolationLayer. |
braindecode/models/interpolated.py |
Adds InterpolatedModel factory and _build_chs_info_from_montage helper. |
braindecode/models/labram.py |
Defines canonical 128-ch positions, enforces canonical chs_info, adds InterpolatedLaBraM. |
braindecode/models/signal_jepa.py |
Adds InterpolatedSignalJEPA variant using the pretrain montage. |
braindecode/models/biot.py |
Adds canonical 18-ch montage constants, makes index non-persistent, adds InterpolatedBIOT. |
braindecode/models/reve.py |
Makes REVE embedding buffer non-persistent. |
braindecode/models/__init__.py |
Exposes interpolated variants and the factory in the public API. |
braindecode/models/util.py |
Updates model test parametrization to support callable signal_params and adds interpolated models’ required params. |
braindecode/models/summary.csv |
Adds summary entries for interpolated model variants. |
docs/api.rst |
Documents the new interpolation layer and interpolated model variants. |
docs/whats_new.rst |
Adds enhancement + API-change entries for interpolation + LaBraM breaking change. |
test/unit_tests/models/test_interpolation.py |
New unit tests for ChannelInterpolationLayer and MNE-matrix helper. |
test/unit_tests/models/test_interpolated.py |
New unit tests for InterpolatedModel and shipped interpolated variants. |
test/unit_tests/models/test_return_features.py |
Updates return-features tests to use interpolated LaBraM and canonical LaBraM chs_info. |
test/unit_tests/models/test_models.py |
Updates LaBraM defaults and expected token counts for 128 channels. |
test/unit_tests/models/test_integration.py |
Integrates interpolated models into integration suite and supports callable signal_params. |
test/unit_tests/models/test_huggingface.py |
Adjusts config assertions to allow interpolated models to identify channels via chs_info. |
test/unit_tests/models/test_foundation_models.py |
Updates LaBraM tests for canonical 128 channels and removes obsolete channel-mapping tests. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| # Require chs_info to match LABRAM_CHANNEL_ORDER exactly (case-insensitive). | ||
| # Arbitrary channel sets should go through InterpolatedLaBraM. | ||
| try: | ||
| _chs_info = self.chs_info | ||
| except ValueError: | ||
| _chs_info = None | ||
| if _chs_info is not None: | ||
| user_names = [ch["ch_name"] for ch in _chs_info] # type: ignore[index] | ||
| canonical = LABRAM_CHANNEL_ORDER | ||
| if [n.lower() for n in user_names] != [n.lower() for n in canonical]: | ||
| raise ValueError( | ||
| f"Labram requires chs_info to match LABRAM_CHANNEL_ORDER exactly " | ||
| f"({len(canonical)} channels, specific order). Got {len(user_names)} " | ||
| f"channels. For arbitrary channel sets, use InterpolatedLaBraM " | ||
| f"(from braindecode.models import InterpolatedLaBraM)." | ||
| ) |
There was a problem hiding this comment.
Labram only validates chs_info when it is provided. If a user instantiates Labram with chs_info=None and n_chans not equal to the canonical 128, the model will still initialize but forward() later assumes the full canonical channel set (via LABRAM_CHANNEL_ORDER), leading to shape/indexing errors. Consider enforcing at init time that either (a) chs_info is provided and matches LABRAM_CHANNEL_ORDER, or (b) n_chans == len(LABRAM_CHANNEL_ORDER) when chs_info is omitted (and raise a clear error otherwise).
| # After __init__ validation, x already has chs_info == LABRAM_CHANNEL_ORDER, | ||
| # so input_chans is always arange(len(canonical) + 1) (CLS + all canonical). | ||
| input_chans = torch.arange( | ||
| len(LABRAM_CHANNEL_ORDER) + 1, device=x.device, dtype=torch.long |
There was a problem hiding this comment.
forward() always builds input_chans as arange(len(LABRAM_CHANNEL_ORDER) + 1), which assumes the input has exactly the canonical 128 channels. If x.shape[1] != len(LABRAM_CHANNEL_ORDER), position_embedding[:, input_chans] will be inconsistent with the actual token count and can error or mis-index. Either validate x.shape[1] against the canonical length here (raising a clear ValueError), or derive input_chans from the actual channels being used.
| # After __init__ validation, x already has chs_info == LABRAM_CHANNEL_ORDER, | |
| # so input_chans is always arange(len(canonical) + 1) (CLS + all canonical). | |
| input_chans = torch.arange( | |
| len(LABRAM_CHANNEL_ORDER) + 1, device=x.device, dtype=torch.long | |
| expected_n_chans = len(LABRAM_CHANNEL_ORDER) | |
| actual_n_chans = x.shape[1] | |
| if actual_n_chans != expected_n_chans: | |
| raise ValueError( | |
| "LaBraM forward expected input with " | |
| f"{expected_n_chans} channels matching LABRAM_CHANNEL_ORDER, " | |
| f"but got {actual_n_chans} channels." | |
| ) | |
| # After validating the runtime input shape, input_chans is the CLS token | |
| # plus all canonical channel positions. | |
| input_chans = torch.arange( | |
| expected_n_chans + 1, device=x.device, dtype=torch.long |
| src_chs_info : list of dict | ||
| Source (user) channel info; each dict must have ``"ch_name"`` and | ||
| ``"loc"`` keys (MNE-style). | ||
| tgt_chs_info : list of dict | ||
| Target channel info; same structure. | ||
| mode : {"always", "name_match"} | ||
| How the matrix is built. Default ``"always"``. | ||
|
|
||
| * ``"always"``: every row of ``W`` is computed via | ||
| :func:`mne.io.Raw.interpolate_to` using the 3D positions. | ||
| * ``"name_match"``: for each target channel whose ``ch_name`` | ||
| (case-insensitive) also appears in ``src_chs_info``, the | ||
| corresponding row of ``W`` is a one-hot vector selecting that | ||
| source channel (its 3D position is ignored). Remaining rows, | ||
| if any, are filled via MNE. If every target name has a source | ||
| match, MNE is not invoked and no ``"loc"`` is required. |
There was a problem hiding this comment.
The src_chs_info / tgt_chs_info parameter docs state that each dict must have a "loc" key, but in mode="name_match" with full name coverage the implementation explicitly short-circuits and does not require loc. Please update the docstring to reflect that loc is only required when an MNE-based matrix is actually computed.
| def test_always_mode_uses_mne_even_when_names_match(): | ||
| # Identical src and tgt by name — in name_match this would be identity, | ||
| # in always mode it uses the MNE matrix (NOT identity). | ||
| # Use at least 4 channels to satisfy MNE's minimum digitization requirement. | ||
| names = ["Fz", "Cz", "Pz", "C3", "C4"] | ||
| src = [_montage_ch(n) for n in names] | ||
| tgt = [_montage_ch(n) for n in names] | ||
| layer = ChannelInterpolationLayer(src, tgt, mode="always") | ||
| assert layer.matrix.shape == (5, 5) | ||
| # MNE on identical positions will approximate identity but likely not | ||
| # be exactly identity. Check non-trivial off-diagonal structure. | ||
| off_diag = layer.matrix - torch.diag(torch.diagonal(layer.matrix)) | ||
| assert torch.any(off_diag.abs() > 1e-6), ( | ||
| "expected non-trivial MNE-computed matrix, got pure diagonal" | ||
| ) |
There was a problem hiding this comment.
test_always_mode_uses_mne_even_when_names_match asserts that the MNE-computed matrix has a non-trivial off-diagonal (>1e-6) for identical source/target montages. For identical positions, a valid interpolation can be (near-)identity (off-diagonals can be exactly 0 or only numerical noise depending on MNE version/settings), so this check can be brittle. A more robust way to test the intent would be to monkeypatch _compute_interpolation_matrix_mne (or the MNE call) and assert it was invoked in mode="always" rather than asserting specific numerical structure.
| def test_always_mode_uses_mne_even_when_names_match(): | |
| # Identical src and tgt by name — in name_match this would be identity, | |
| # in always mode it uses the MNE matrix (NOT identity). | |
| # Use at least 4 channels to satisfy MNE's minimum digitization requirement. | |
| names = ["Fz", "Cz", "Pz", "C3", "C4"] | |
| src = [_montage_ch(n) for n in names] | |
| tgt = [_montage_ch(n) for n in names] | |
| layer = ChannelInterpolationLayer(src, tgt, mode="always") | |
| assert layer.matrix.shape == (5, 5) | |
| # MNE on identical positions will approximate identity but likely not | |
| # be exactly identity. Check non-trivial off-diagonal structure. | |
| off_diag = layer.matrix - torch.diag(torch.diagonal(layer.matrix)) | |
| assert torch.any(off_diag.abs() > 1e-6), ( | |
| "expected non-trivial MNE-computed matrix, got pure diagonal" | |
| ) | |
| def test_always_mode_uses_mne_even_when_names_match(monkeypatch): | |
| # Identical src and tgt by name — in name_match this would be identity, | |
| # but in always mode it should still invoke the MNE-backed path. | |
| import braindecode.modules.interpolation as interpolation_module | |
| names = ["Fz", "Cz", "Pz", "C3", "C4"] | |
| src = [_montage_ch(n) for n in names] | |
| tgt = [_montage_ch(n) for n in names] | |
| expected = torch.eye(5, dtype=torch.float32) * 2.0 | |
| calls = {} | |
| def fake_compute_interpolation_matrix_mne(src_arg, tgt_arg, method="spline"): | |
| calls["src"] = src_arg | |
| calls["tgt"] = tgt_arg | |
| calls["method"] = method | |
| return expected | |
| monkeypatch.setattr( | |
| interpolation_module, | |
| "_compute_interpolation_matrix_mne", | |
| fake_compute_interpolation_matrix_mne, | |
| ) | |
| layer = ChannelInterpolationLayer(src, tgt, mode="always") | |
| assert layer.matrix.shape == (5, 5) | |
| assert calls["src"] == src | |
| assert calls["tgt"] == tgt | |
| assert calls["method"] == "spline" | |
| torch.testing.assert_close(layer.matrix, expected) |
Mirrors the InterpolatedLaBraM / InterpolatedBIOT pattern introduced in braindecode#993: uses the InterpolatedModel factory with a 20-channel target composed of the 19 pre-training EEG channels plus a SCALE placeholder positioned at the centroid of those 19 points. The SCALE row of the interpolation matrix is therefore a spatial spline of the user's EEG (not the dn3 MappingDeep1010 amplitude statistic); this is documented in the source comment and is the same kind of approximation that InterpolatedBIOT uses for its 18 bipolar derivations. Side effect of introducing the derived constants: renames the _BENDR_TARGET_CHS constant (added earlier in this PR) to _BENDR_TARGET_CHS_TUPLES, matching the LaBraM naming convention (_TARGET_CHS_TUPLES → _TARGET_CHS_INFO → CHANNEL_ORDER). Exposes InterpolatedBENDR via braindecode.models.__init__, registers it in models_mandatory_parameters, adds it to summary.csv, docs/api.rst, and the activation/drop_prob/torchscript skip lists in test_integration.py (matching the other Interpolated* wrappers). Adds two targeted tests in test_interpolated.py and an InterpolatedBENDR row in test_return_features.py.
PR #993 (1.5.0) removed the per-batch ch_names kwarg from Labram.forward and tightened the __init__ chs_info check into a ValueError. Downstream wrappers (neuroai's _LabramChannelWrapper) were still calling Labram with non-canonical chs_info and a per-batch ch_names subset, producing 7 test failures. Restore the kwarg as keyword-only (via *); when provided it case-insensitively maps each name to LABRAM_CHANNEL_ORDER and uses those indices for the position-embedding bank. When None, fall back to the 1.5.0 arange-over-canonical behavior. The __init__ canonical check is downgraded to a UserWarning so wrappers can build the inner Labram with their union channel set and resolve the subset per batch. Update test_labram_rejects_non_canonical_chs to assert the warning (renamed to test_labram_warns_on_non_canonical_chs). Document the restoration under 1.5.1 'API and behavior changes' in whats_new.
- forward(): move `*` after the return_* flags so only `ch_names` is keyword-only. `return_patch_tokens`, `return_all_tokens` and `return_features` stay positional, preserving back-compat for callers that used positional args before #993. - forward(): when `ch_names is None`, validate `x.shape[1] == len(LABRAM_CHANNEL_ORDER)` up front and raise a clear ValueError pointing to `ch_names=` or `InterpolatedLaBraM` instead of failing later with a confusing shape mismatch inside `forward_features`. - forward() docstring: clarify that `ch_names` is keyword-only and only honored when `neural_tokenizer=True`; in decoder mode the position embedding is sequential. - examples/plot_channel_interpolation: drop the redundant in-body `import warnings` (already imported at module top) and rewrite the "silently mis-align" wording — the model warns at construction and forward now raises ValueError on a non-canonical channel count. - test_foundation_models: add focused tests for the ch_names path — subset forward, case-insensitive matching, unknown-channel ValueError, length mismatch, the new None+non-canonical guard, and back-compat for positional return_* flags.
Summary
Adds an experimental channel interpolation feature that lets arbitrary user channel sets be projected to a pre-trained foundation model's canonical channel space via an MNE-backed (frozen by default) interpolation matrix. Compatible with linear probing: the projection matrix preserves pre-trained weights and does not require SGD to be re-learned.
What's in
New primitive
braindecode.modules.ChannelInterpolationLayer— annn.Modulethat holds an(n_target, n_source)matrix built at init from MNE'sRaw.interpolate_to. Two modes:"always"— full MNE spline interpolation on 3D positions."name_match"— rows for target channels whose name (case-insensitive) is present in the source are filled as one-hot vectors; remaining rows via MNE. If all target names match, MNE is never invoked and the matrix is a pure permutation.trainable=False→ non-persistent buffer, recomputed fromchs_info). Settrainable=Trueto promote it tonn.Parameter.kind);locrequired only when MNE is invoked.New factory
braindecode.models.InterpolatedModel(model_cls, target_chs_info)— returns a subclass ofmodel_clsthat prepends the interpolation layer inforwardand rebindsself._chs_info/self._n_chanspost-super so the user-facing view reflects the user's channels while the backbone sees its canonical set. Round-trips throughget_config/from_configwith the user'schs_info.Shipped variants
braindecode.models.InterpolatedSignalJEPA— 62-channel canonical set. Additive: coexists with the existingchannel_embedding="pretrain_aligned"added in Improve pre-trained channel embeddings handling in Signal-JEPA #991 (subset-by-name) and handles the more general case (arbitrary channels via interpolation).braindecode.models.InterpolatedLaBraM— 128-channel canonical set derived fromLABRAM_CHANNEL_ORDER. Positions:standard_1005for monopolar names, midpoints for bipolar / intermediate / aliased names. ATODOflags the bipolar-midpoint simplification (a faithful bipolar derivation would need a dedicated layer, deferred).braindecode.models.InterpolatedBIOT— 18-channel TCP + SHHS canonical set from the original BIOT repo. All channels are bipolar / differential; same midpoint simplification +TODO.Breaking changes
Labramnow requireschs_infonames to matchLABRAM_CHANNEL_ORDERexactly (case-insensitive). Theon_unknown_chsparameter and the forward-timech_namesargument are removed. Migrate toInterpolatedLaBraMfor arbitrary channel sets.Style / internals
Labram's_LABRAM_TARGET_CHS_TUPLESis the single source of truth for the 128 canonical positions;_LABRAM_TARGET_CHS_INFOandLABRAM_CHANNEL_ORDERare derived via comprehensions (no duplication)._build_chs_info_from_montage(names, montage)helper inbraindecode/models/interpolated.py.Why not also handle bipolar properly?
Bipolar channels (
FP1-F7, etc.) representV(A) − V(B), which MNE spatial interpolation cannot recover from monopolar input. We approximate the position as the midpoint for the interpolation matrix — good enough to feed a foundation model that relies on learned position embeddings indexed by name (Labram), but not physically correct for BIOT whose input is expected to be bipolar. This limitation is flagged withTODOcomments and can be addressed in a follow-up PR by a dedicatedBipolarDerivationLayer.Migration
Labram(on_unknown_chs=...)→InterpolatedLaBraM(chs_info=user_chs_info, ...).SignalJEPA(channel_embedding="pretrain_aligned")still works unchanged.InterpolatedSignalJEPA(chs_info=user_chs_info, ...)is the new, more general alternative.InterpolatedBIOT(chs_info=user_chs_info, ...)when loading the 18-channel pre-trained checkpoints with arbitrary input channels.Test plan
ChannelInterpolationLayerintest/unit_tests/models/test_interpolation.py(permutation short-circuit, case-insensitivity,alwaysvsname_matchhybrid, validators,trainablebuffer/parameter semantics, mocked assertion that MNE is not called on full name coverage).test/unit_tests/models/test_interpolated.py(target required, rebind semantics,get_configround-trip, shipped LaBraM / SignalJEPA / BIOT behavior,_build_chs_info_from_montagehelper, regression test thatpretrain_alignedstill works).on_unknown_chs/ch_namestests removed.Docs
docs/api.rst: addedChannelInterpolationLayer,InterpolatedModel,InterpolatedLaBraM,InterpolatedSignalJEPA,InterpolatedBIOT.docs/whats_new.rst: enhancements + API-change entries.