Skip to content

[ENH] Experimental channel interpolation for foundation models - #993

Merged
PierreGtch merged 29 commits into
braindecode:masterfrom
PierreGtch:channel-projection
Apr 19, 2026
Merged

PierreGtch merged 29 commits into
braindecode:masterfrom
PierreGtch:channel-projection

Conversation

@PierreGtch

@PierreGtch PierreGtch commented Apr 19, 2026 •

Copy link
Copy Markdown
Collaborator

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 — an nn.Module that holds an (n_target, n_source) matrix built at init from MNE's Raw.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.
  • Frozen matrix by default (trainable=False → non-persistent buffer, recomputed from chs_info). Set trainable=True to promote it to nn.Parameter.
  • EEG-only validation (rejects channels with an explicit non-EEG kind); loc required only when MNE is invoked.

New factory

  • braindecode.models.InterpolatedModel(model_cls, target_chs_info) — returns a subclass of model_cls that prepends the interpolation layer in forward and rebinds self._chs_info / self._n_chans post-super so the user-facing view reflects the user's channels while the backbone sees its canonical set. Round-trips through get_config / from_config with the user's chs_info.

Shipped variants

  • braindecode.models.InterpolatedSignalJEPA — 62-channel canonical set. Additive: coexists with the existing channel_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 from LABRAM_CHANNEL_ORDER. Positions: standard_1005 for monopolar names, midpoints for bipolar / intermediate / aliased names. A TODO flags 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

  • Labram now requires chs_info names to match LABRAM_CHANNEL_ORDER exactly (case-insensitive). The on_unknown_chs parameter and the forward-time ch_names argument are removed. Migrate to InterpolatedLaBraM for arbitrary channel sets.

Style / internals

  • Labram's _LABRAM_TARGET_CHS_TUPLES is the single source of truth for the 128 canonical positions; _LABRAM_TARGET_CHS_INFO and LABRAM_CHANNEL_ORDER are derived via comprehensions (no duplication).
  • New _build_chs_info_from_montage(names, montage) helper in braindecode/models/interpolated.py.

Why not also handle bipolar properly?

Bipolar channels (FP1-F7, etc.) represent V(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 with TODO comments and can be addressed in a follow-up PR by a dedicated BipolarDerivationLayer.

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.
  • BIOT itself is unchanged; use InterpolatedBIOT(chs_info=user_chs_info, ...) when loading the 18-channel pre-trained checkpoints with arbitrary input channels.

Test plan

  • 15 unit tests for ChannelInterpolationLayer in test/unit_tests/models/test_interpolation.py (permutation short-circuit, case-insensitivity, always vs name_match hybrid, validators, trainable buffer/parameter semantics, mocked assertion that MNE is not called on full name coverage).
  • 14 unit tests for the factory + shipped variants in test/unit_tests/models/test_interpolated.py (target required, rebind semantics, get_config round-trip, shipped LaBraM / SignalJEPA / BIOT behavior, _build_chs_info_from_montage helper, regression test that pretrain_aligned still works).
  • 19 obsolete Labram on_unknown_chs / ch_names tests removed.

Docs

  • docs/api.rst: added ChannelInterpolationLayer, InterpolatedModel, InterpolatedLaBraM, InterpolatedSignalJEPA, InterpolatedBIOT.
  • docs/whats_new.rst: enhancements + API-change entries.

PierreGtch and others added 21 commits April 19, 2026 10:45
- 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.
@PierreGtch PierreGtch changed the title Channel projection layer [ENH] Experimental channel interpolation for foundation models Apr 19, 2026
…__, skip signature inspection tests, register summary rows
…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.
@PierreGtch
PierreGtch marked this pull request as ready for review April 19, 2026 19:11
Copilot AI review requested due to automatic review settings April 19, 2026 19:11

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

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 ChannelInterpolationLayer plus an InterpolatedModel(...) factory to prepend channel projection to existing backbones.
  • Ship InterpolatedLaBraM, InterpolatedSignalJEPA, and InterpolatedBIOT wrappers and wire them into the model registry/docs.
  • Refactor Labram to 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.

Comment on lines +389 to +404
# 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)."
)

Copilot AI Apr 19, 2026

Copy link

Choose a reason for hiding this comment

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

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).

Copilot uses AI. Check for mistakes.
Comment on lines +741 to +744
# 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

Copilot AI Apr 19, 2026

Copy link

Choose a reason for hiding this comment

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

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.

Suggested change
# 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

Copilot uses AI. Check for mistakes.
Comment on lines +55 to +70
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.

Copilot AI Apr 19, 2026

Copy link

Choose a reason for hiding this comment

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

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.

Copilot uses AI. Check for mistakes.
Comment on lines +81 to +95
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"
)

Copilot AI Apr 19, 2026

Copy link

Choose a reason for hiding this comment

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

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.

Suggested change
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)

Copilot uses AI. Check for mistakes.
@PierreGtch
PierreGtch merged commit 80bc327 into braindecode:master Apr 19, 2026
15 of 17 checks passed
PierreGtch added a commit to PierreGtch/braindecode that referenced this pull request Apr 19, 2026
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.
bruAristimunha added a commit that referenced this pull request May 19, 2026
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.
bruAristimunha added a commit that referenced this pull request May 20, 2026
- 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.
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