Skip to content

Channel layer: one channel_strategy for every pretrained model - #1241

Merged
bruAristimunha merged 64 commits into
braindecode:masterfrom
bruAristimunha:feat/channel-layer
Oct 7, 2026
Merged

bruAristimunha merged 64 commits into
braindecode:masterfrom
bruAristimunha:feat/channel-layer

Conversation

@bruAristimunha

@bruAristimunha bruAristimunha commented Oct 6, 2026 •

Copy link
Copy Markdown
Collaborator

One channel layer inside every pretrained model, HF-tokenizer style: the user picks how an arbitrary montage reaches the backbone with channel_strategy=... (saved in the config, optional per-call chs_info).

model = EEGPT(chs_info=raw.info["chs"], n_times=1024, n_outputs=2, channel_strategy="field")
y = model(x)                     # or model(x, chs_info=other_chs); model.forward(x) does the same
  • braindecode/modules/channels.py (one file): resolve each channel (exact name → alias → position, EEG-type check; names without positions take standard_1005 positions in the head frame that raw.set_montage gives), then map it with one strategy: native (default, today's behaviour), exact, zero, nearest, idw, spline, field, source (minimum-norm on a sphere fitted to standard_1005, optionally trainable), wiener, region, latent.
  • The 19 pretrained EEG classes take channel_strategy / channel_strategy_kwargs. With native nothing is added: same code path, state dict and outputs as before (max-abs 0.0), TorchScript unchanged. SleepFM and SleepFMStager (polysomnography grouped by modality) accept only native.
  • Under a strategy, both model(x) and model.forward(x) apply the layer once; the saved config keeps the user's input montage, so from_config / from_pretrained rebuild the same model. A target that cannot be copied or reconstructed warns; if no target is observed the model raises instead of feeding zeros.
  • Per-montage maps are built once (16-montage cache on the module's device, built outside inference_mode); steady state is one device op with no host transfer.
  • Removes InterpolatedModel, the five Interpolated* classes, ChannelInterpolationLayer and ChannelTokenizer (experimental since 1.5.0): InterpolatedBENDR(chs_info=…) → BENDR(chs_info=…, channel_strategy="spline").
  • Tests: test_channels.py (strategies, fidelity on a head model different from the inversion one, device/dtype, cache, inference mode, pickle/deepcopy) and test_pretrained_compat.py (geometry contract for every pretrained model, native checkpoints load strictly under a strategy).

Replaces modules/channel_tokenizer.py (braindecode#1235). BENDR's non-canonical path
now uses ChannelTokenizer(ChannelTarget('montage', ...), 'spline', reg=0),
the former fixed_order behaviour; the canonical path builds no tokenizer
(bit-identical).
MNE's sphere EEG formula returns NaN for an electrode exactly above the
sphere centre (biosemi64 Cz at (0, 0, r)), so the source strategy produced
NaN maps for any biosemi64 input that needed a filled target. Shift such
points by 1 um.
Each model declares a montage ChannelTarget (BIOT: the 18 monopolar
electrodes behind its bipolar derivations, which forward forms as
V(A) - V(B); CodeBrain: 10-20 19 channels; MIRepNet: the 45-channel
use_channels_names template of the original code) and takes
channel_strategy / channel_strategy_kwargs. Native is unchanged
(bit-identical, same state_dict) but warns (FutureWarning) when the
montage is not the canonical one.
EEGDINO fills its 19 slots from any montage under a channel strategy.
CBraMod, Brant and BrainBERT are free: sensor strategies pass x through,
source/latent feed parcels/latents. CBraMod masks the patches of channels
the layer marks as not observed. Brant and BrainBERT stay scriptable: the
layer is eager-only and native drops it.
LatentStrategy.build returns a pass-through map (no extra tensors) for a
positions target without a training montage; apply then failed on the
missing 'used' entry. Return x unchanged, as the base strategy does.
WienerStrategy read its cov / dense_positions buffers with .double().numpy(),
which raises a raw TypeError on MPS and on CUDA. Copy them to the CPU first,
and keep fitted buffers on the strategy's device. New test: fit, move to
cuda/mps, forward a montage not built yet (spline, trainable source, latent,
wiener); skipped without an accelerator.
ChannelTokenizer cleared its per-montage cache only in fit(); loading another
fitted wiener state kept serving the old maps. Clear it in
_load_from_state_dict.
…ce (M2)

source returned L_t S for reconstructed rows, with S killing constants: those
rows summed to 0 while copies summed to 1, so a common-mode / reference signal
reached copies but not reconstructed channels. field had the same mix (row
sums 0.09, -0.56, ...). Reconstructed rows now estimate target minus the mean
of the used inputs, plus that mean: every row sums to 1 and the whole output
stays in the input's reference.

The source fidelity test is parametrised over two truth references (average
of the 19 sites, and Fp1): 0.888 / 0.761, bounds 0.92 / 0.79. New test:
constant input -> constant output for nearest, idw, spline, field, source
(physics and trainable).
… axis

biosemi64 Cz at (0, 0, 0.095) lies on the line through the sphere centre and a
grid dipole; MNE's sphere formula returns NaN rows there, which spread to every
'source' fill row. Move such electrodes by 1 um and recompute once.
channel_strategy / channel_strategy_kwargs; under a strategy the position
embedding is indexed by ChannelEncoding.channel_ids (finishing braindecode#1227's names
migration), forward takes chs_info (ch_names as a shortcut), and reconstructed
channels are masked as attention keys. Native path unchanged.
ids into the 62-name EEGPT_CHANNELS table; with chan_proj_type='none' the
encoder reads ChannelEncoding.channel_ids and masks reconstructed channels as
attention keys (SDPA key-padding mask), otherwise the channel projection maps
the reconstructed vocabulary to the 19 standard channels. Native unchanged.
ids into the montage vocabulary (lazy class attribute, so the Hub file is only
fetched when a strategy needs it); the channel embedding uses
ChannelEncoding.channel_ids. Native unchanged and never fetches the vocabulary
for its contract.
Under a strategy the pre-training channel table is used and rows are looked
up with ChannelEncoding.channel_ids; feature encoder and contextual head are
sized to the channels the layer produces. The zero-span positional encoding
fix is kept. PostLocal / PreLocal stay native-only (no channel embedding).
MVPFormer has no channel vocabulary (slots in input order): contract 'free'.
Sensor strategies pass channels through, 'source' / 'latent' feed parcels /
latents into the slots; the concat head is sized to them. Native unchanged.
Contract, native == exact on the canonical montage (bit-identical), strict
load, config round trip, from_pretrained with a strategy, 10 strategies x
G1-G4 per model with declared errors, key-padding mask reaches LaBraM and
EEGPT attention, LaBraM NOT_YET compat cells under a strategy, sphere-axis
lead-field regression.
Spec section 4: once per built montage (maps are cached), warn when a
reconstructed row has |w|_1 > 2 and when targets have support < 0.5 (no used
input within 21 mm). The copy and support blocks of ChannelStrategy.build move
into _copies / _fill_support helpers.
…nd-trip (m3)

The fresh channel_tokenizer keys warning told wiener users to train the empty
covariance buffers; it now says to call channel_tokenizer.fit() for
strategies that are not trainable (uses ChannelStrategy.trainable). New
test_base cases: the warning wording, and a fitted wiener model through
save_pretrained / from_pretrained ((0, 0) buffer resized on load) gives equal
outputs on a montage that needs reconstruction.
Registering _test_trainable at import mutated the global strategy registry, so
a second import (xdist workers, importlib.reload) raised 'already registered'.
A fixture registers it per test and pops it on teardown.
_ids_subset matched every unmatched input channel to the vocabulary entry
within 15 mm, so a digitised known name (FCz 5 mm from Cz) was relabelled
silently. Apply the rule of _match: only names unknown to standard_1005 with
a user position fall back to position.
…t (m5)

A 'positions' ChannelTarget without chs_info is a pass-through, so source,
spline, latent, ... returned the input unchanged without a word. Warn once at
construction (not for exact / zero) and point to ChannelTarget(chs_info=...).
A montage first seen under torch.inference_mode() (e.g. a validation
batch) cached inference tensors, and the next training step on that
montage raised "Inference tensors cannot be saved for backward" (latent,
trainable source). The one build path now runs with inference mode off.
Copilot AI balanced review requested due to automatic review settings October 7, 2026 13:35

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.

Comment thread braindecode/models/base.py Outdated
def _apply_channel_layer(model, args, kwargs):
"""Forward pre-hook: run the channel layer on ``x`` (``chs_info=`` per call)."""
x, _ = model.channel_layer(args[0], kwargs.pop("chs_info", None))
return (x, *args[1:]), kwargs
Comment thread braindecode/modules/channels.py Outdated
Copilot AI balanced review requested due to automatic review settings October 7, 2026 14:49
…one only zeros silently

- Under a channel strategy, `model.forward` is an instance attribute bound to
  `EEGModuleMixin._channel_forward` (replaces the forward pre-hook), so
  `model(x)` and `model.forward(x)` both map x exactly once (`chs_info=` /
  LaBraM `ch_names=` per call, x positional or by keyword). Native models have
  no instance attribute: bit-identical, TorchScript unchanged.
  `get_output_shape` calls the class forward (the backbone on its target).
- ChannelLayer: a target missing from the input that has no position cannot be
  reconstructed: warn naming it (e.g. BENDR SCALE). If no target is observed and
  none is reconstructed, raise a ValueError instead of returning zeros.
- Tests: forward == call (layer runs once), deepcopy and pickle round trip,
  no-observed-target declared, SCALE warning.

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.

🟡 Changes recommended

Unresolved mapping, supported-input, and checkpoint-loading regressions prevent safe approval.

7 open findings
7 resolved since last review
Previously missed (3)

In code that hasn't changed since last review

Medium severity Size summary inputs from supplied channel-layer geometry

braindecode/​models/​base.py:356

Storing the target montage here makes get_torchinfo_statistics() create backbone-sized input, while the pre-hook expects the recording's channel count. An eight-channel BENDR with channel_strategy='spline' reports n_chans=20, so both the summary and str(model) fail with a montage-length error. Preserve target geometry for backbone construction, but size the summary input from channel_layer.chs_info when supplied (in get_torchinfo_statistics, lines 739–750).

Medium severity Revalidate mutable montage contents on every call

braindecode/​modules/​channels.py:288

This identity shortcut bypasses invalidation when the supplied montage list is reordered or its channel locations change in place. A subsequent forward using the same list and correspondingly reordered data silently uses the old mapping, although the content-based key below would detect the change. Compare the key on every call so mutable chs_info cannot leave stale channel assignments.

Medium severity Handle single-sample latent statistics without NaNs

braindecode/​modules/​channels.py:434

With a one-sample input and missing latent targets, the variance divides by zero and the temporal-difference mean is taken over an empty dimension. These NaN statistics propagate through attention, making the output non-finite. Handle (batch, channels, 1) with zero variance and difference statistics while preserving the calculation for longer inputs.

🧠 Review effort: Balanced

Comment thread braindecode/models/base.py Outdated
Comment on lines +103 to +107
if args:
return (model.channel_layer(args[0], chs)[0], *args[1:]), kwargs
x = next(iter(inspect.signature(model.forward).parameters)) # x, X, eeg...
kwargs[x] = model.channel_layer(kwargs[x], chs)[0]
return args, kwargs
Comment on lines +665 to +667
for k, v in super().state_dict().items():
if k.startswith("channel_layer.") and k not in new_state_dict:
new_state_dict[k] = v
self._key = self._chs = None
for k in ("cov", "dense_pos"): # fitted buffers change size
if hasattr(self, k) and prefix + k in state_dict:
setattr(self, k, torch.empty_like(state_dict[prefix + k]))
Comment on lines +427 to +428
self._use(chs) # steady state: no host work, every map already on device
out = self.weight @ x
Copilot AI balanced review requested due to automatic review settings October 7, 2026 15:21
The head-frame fix put the electrodes about 15 mm in front of MNE's default
sphere centre (0, 0, 0.04). The template head is now the sphere fitted to the
dense standard_1005 in the head frame (r0 (-1.0, 14.8, 39.2) mm, R 98.8 mm).
The hard-coded Berg-Scherg parameters are unchanged: they depend only on the
relative radii and conductivities (re-derived, identical).

Source error on another head (19 10-20 targets from 8 inputs, 3 random
dipoles): realistic sample BEM 1.289 -> 1.031, 4-shell sphere 0.771 -> 0.698.
New test bounds it on another 4-shell head (0.759; the old sphere: 0.820).

Mark the pickle round-trip in test_channels.py as nosec (Codacy B403/B301):
it unpickles the model it has just pickled.

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.

🟡 Changes recommended

Unresolved signal-mapping, metadata-alignment, cache-invalidation, and concurrency issues can cause incorrect outputs or runtime failures.

9 open findings
1 resolved since last review
Previously missed (4)

In code that hasn't changed since last review

Medium severity Keep input and backbone shapes separate

braindecode/​models/​base.py:341

Replacing the input metadata with the target montage makes input_shape report the backbone's channel count. BENDR built from eight input channels reports (1, 20, n_times), but forwarding that shape fails because the channel layer expects eight channels. get_torchinfo_statistics() also uses self.n_chans, so str(model) fails too. Keep separate input and backbone shapes: use the input montage for public input-shape reporting and summaries, and the target montage for backbone initialization and dummy forwards. Add a subset-montage regression.

Medium severity Exclude bipolar inputs from proximity matching

braindecode/​modules/​channels.py:321

A measured bipolar name is unknown to the standard montage, so this proximity fallback can copy a voltage difference into a monopolar target. For example, F3-F1 located at its midpoint is within 15 mm of both F3 and F1 and is copied into both rows, even under exact. Exclude bipolar inputs from proximity-only matching; exact matches to a bipolar target should still be copied.

Medium severity Exclude bipolar channels from spatial reconstruction

braindecode/​modules/​channels.py:349

This reconstruction source mask includes measured bipolar channels whenever they have a finite location. With a partially observed bipolar montage, spline/field/source therefore treat V(A)-V(B) as the potential at one sensor location when filling missing electrode rows, producing physically incorrect reconstructions. Keep bipolar measurements available for direct copies, but exclude them from the electrode samples used for spatial reconstruction.

Medium severity Reject inputs shorter than two samples

braindecode/​modules/​channels.py:444

A one-sample input produces NaNs whenever latent reconstruction is needed: the variance divides by zero and the temporal difference has an empty mean. The public layer currently accepts this input without a minimum-length check, and NaNs propagate through attention into the output. Reject windows shorter than two samples before computing these statistics, or define finite single-sample statistics.

🧠 Review effort: Balanced

Comment on lines +383 to +387
args = (self.channel_layer(args[0], chs_info)[0], *args[1:])
else: # x by keyword: the backbone's first input (x, X, eeg...)
x = list(inspect.signature(type(self).forward).parameters)[1]
kwargs[x] = self.channel_layer(kwargs[x], chs_info)[0]
return type(self).forward(self, *args, **kwargs)
Comment on lines +119 to +123
fit = mock.patch.object( # deterministic fit: reuse its result
mne.bem,
"_fwd_eeg_fit_berg_scherg",
lambda m, *_: m.update(zip(("mu", "lambda"), _MU_LAMBDA), nfit=3) or 0.0,
)
Comment on lines +285 to +288
"""Make the maps of ``chs`` the buffers (identity, then LRU, then build)."""
# ponytail: identity fast path; a montage list mutated in place is not seen.
if chs is self._chs:
return
Copilot AI balanced review requested due to automatic review settings October 7, 2026 15:32

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.

🟡 Changes recommended

Unresolved mapping, checkpoint compatibility, and input-handling defects can produce incorrect signals or runtime failures.

11 open findings
Previously missed (6)

In code that hasn't changed since last review

Medium severity Keep input and backbone geometries separate

braindecode/​models/​base.py:341

Replacing the input geometry with the target breaks introspection when their channel counts differ. An eight-channel BENDR with channel_strategy='spline' reports input_shape=(1, 20, n_times), but its channel layer rejects that tensor because it expects eight channels. get_torchinfo_statistics() uses the same count (base.py:741-751), so str(model) and print(model) also fail. Keep input and backbone geometry separate, and use the input montage for public input shapes and summaries.

Medium severity Align channel metadata with remapped signals

braindecode/​models/​base.py:383

Only the signal is remapped; channel-aligned forward arguments retain their input order. For example, DIVER1 with reversed signals, chs_info, and chan_metadata maps the signals to target order but leaves metadata reversed, attaching spatial embeddings to the wrong electrodes. Different channel counts also fail metadata validation. LUNA's channel_locations and REVE's pos have the same alignment problem. Transform these arguments using model-specific rules, or clearly reject unsupported overrides under a channel strategy.

Medium severity Preserve native zero-fill for coordinate-less montages

braindecode/​models/​bendr.py:313

This changes the default native behavior for coordinate-less montages. A names-only subset such as Fz/Cz/Pz previously copied those channels and zero-filled the rest. The new layer infers standard positions, attempts spline reconstruction, and rejects the three-channel input. Preserve the coordinate-less zero-fill fallback under native, leaving inferred-position reconstruction to an explicitly selected strategy.

Medium severity Remove identity shortcut from montage cache validation

braindecode/​modules/​channels.py:293

Reusing a montage list after modifying it in place leaves the cached map stale. For example, after chs.reverse(), forwarding x.flip(1) with chs_info=chs silently assigns signals to the wrong electrodes because this identity check returns before examining the names. Remove the identity shortcut so the content-based key detects changes to names, types, and coordinates.

Medium severity Exclude bipolar inputs from monopolar reconstruction

braindecode/​modules/​channels.py:354

Positioned bipolar inputs are treated as monopolar electrodes here. For example, LaBraM's FP1-F7 midpoint metadata makes reconstruction use V(FP1)-V(F7) as the potential at that midpoint, producing incorrect missing-channel estimates. These inputs can also enter the position-based copying path above. Preserve exact copies of measured bipolar targets, but exclude bipolar inputs from monopolar positional matching and reconstruction, or explicitly model their endpoint differences.

Medium severity Handle single-sample inputs without NaN statistics

braindecode/​modules/​channels.py:451

For a (batch, channels, 1) input with a target to reconstruct, the variance divides by zero and the difference statistic averages an empty tensor. These NaNs propagate through latent attention into the output. ChannelLayer accepts this shape without a minimum-length check; use zero variance and difference statistics for single-sample inputs.

🧠 Review effort: Balanced

GitHub https://github.com/ycq091044/BIOT (accessed 2024-02-13)
"""

_channel_target = BIOT_CHANNEL_ORDER
f"Input has {x.shape[1]} channels but the montage has {len(chs)}; they must match."
)
self._use(chs) # steady state: no host work, every map already on device
out = self.weight @ x
Copilot AI balanced review requested due to automatic review settings October 7, 2026 17:05

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.

if layer is not None:
self.channel_layer = layer
# Instance attribute: model(x) and model.forward(x) both map x.
setattr(self, "forward", self._channel_forward)
]
kwargs = channel_strategy_kwargs or {}
layer = ChannelLayer(target, channel_strategy, chs_info, **kwargs)
chs_info, n_chans = layer.target, None
Comment on lines +291 to +293
# ponytail: identity fast path; a montage list mutated in place is not seen.
if chs is self._chs:
return
low = np.array([n.lower() for n in names])
alias = np.nan_to_num(cdist(tstd, std), nan=np.inf) < 1e-6
dist = np.nan_to_num(cdist(tpos, pos), nan=np.inf)
near = np.where(np.isnan(std).any(1), dist, np.inf)
f"Strategy 'exact': target channels {[m for m, t in zip(mono, todo) if t]} are not in the input montage {names}; use a reconstructing strategy (e.g. 'spline') or supply them."
)
seen = np.abs(D) @ missing == 0
if not seen.any() and (self.strategy == "zero" or not todo.any()):
UserWarning,
stacklevel=4,
)
use = keep & np.isfinite(pos).all(1)
f"fit() is for 'wiener' and needs X (n_samples, {len(pos)}); got {self.strategy!r}, X {X.shape}."
)
dev = self.get_buffer("cov").device
self.cov = torch.tensor(np.cov(X.T), dtype=torch.float32, device=dev)
Comment on lines +450 to +451
var = (xf - xf.mean(-1, keepdim=True)).square().sum(-1) / (x.shape[-1] - 1)
stats = torch.stack([var, xf.diff(dim=-1).abs().mean(-1)], -1)
Copilot AI balanced review requested due to automatic review settings October 7, 2026 17:54

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.

🔵 Needs a closer look

Model summaries fail when the configured input montage has a different channel count from the backbone target.

19 open findings
Previously missed (2)

In code that hasn't changed since last review

Medium severity Size summary inputs from the configured input montage

braindecode/​models/​base.py:341

For a fixed-target model built from a smaller input montage, this makes self.n_chans describe the backbone rather than the input. get_torchinfo_statistics() still forwards (1, self.n_chans, self.n_times) through the channel layer (base.py:740–755). Consequently, str(model) and model summaries fail for BENDR with an eight-channel input and channel_strategy="spline": the dummy input has 20 channels, but the layer expects eight. Use len(channel_layer.chs_info) to size the summary input when an input montage is configured, while retaining the target count for backbone construction.

Medium severity Reuse source reconstruction results for fixed and trainable maps

braindecode/​modules/​channels.py:384

When a trainable source strategy needs to reconstruct channels, _mixing() already calls _source() with these positions. This second call repeats the MNE forward solution and inverse calculation. _sphere_head() caches the template, not these per-montage calculations, so every uncached montage pays that cost twice. Compute (Lt, S) once and reuse it for both the fixed map and trainable correction.

🧠 Review effort: Balanced

@bruAristimunha
bruAristimunha merged commit 030d24f into braindecode:master Oct 7, 2026
13 checks passed
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Oct 7, 2026
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