You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
How do we handle the channels used for pre-training in foundation models? #995
Braindecode ships a growing number of foundation models (BIOT, BENDR, CBraMod,
CodeBrain, EEGPT, LaBraM, LUNA, MEDFormer, REVE, Signal-JEPA, …). Each one
carries its own implicit assumption about the channel set it was pre-trained
on — fixed names, fixed order, fixed 3D locations, bipolar derivations, … —
and each one exposes a different mechanism for the case where the user's
channels don't match. There is no shared contract a user can rely on to
fine-tune or extract features from these models.
This issue is the starting point of a discussion to agree on a common
contract ("if the model advertises X, the user can do Y with guarantee Z").
As a grounding, the inventory below catalogs what exists today.
Issues created
1. Pre-trained weight loading is fragile and inconsistent
Different models have different failure modes when the user's chs_info
doesn't match what the weights were trained on:
BIOT has no documented path for "I don't have those 18 bipolar channels".
EEGPT has two optional learnable projections (chan_proj_type) —
choosing between them is user-facing.
Group C models load the weights fine but the result is not
equivalent to the pre-training regime.
2. Hugging Face Hub distribution doesn't encode the contract
A pre-trained checkpoint on the Hub is a state_dict + config.json. If
the architectural contract ("this weight file is only meaningful for these
channels in this order") is not part of the config, nothing stops a user
from loading the weights onto the wrong channel set. Today, some models
bake the canonical chs_info into the class (LaBraM, BIOT) so the contract
survives a from_pretrained call; others do not.
3. Feature extraction vs fine-tuning is not clearly differentiated
Users have two different workflows:
Fine-tuning: load weights, replace the head, train on a
downstream task. User tolerates some architecture/weight mismatch
because gradients can absorb it.
Feature extraction: load weights, freeze the backbone, compute
embeddings for a linear probe or post-hoc analysis. User expects the
pre-trained representation to be preserved faithfully — any silent
reindexing or zero-padding destroys the representation.
Today both paths share the same API, so a user doing a linear probe can
silently get a worse representation than they think, because a rename/reorder
they didn't know about happened under the hood.
Current heterogeneous solutions in braindecode
Every bullet is something one foundation model does that no other foundation
model does the same way.
LaBraM (pre-refactor) — on_unknown_chs={"zero_pad", "raise"} flag,
zero-padding missing channels from the canonical 128-channel set. Removed
in the current refactor (breaking change).
LaBraM (post-refactor) — strict validation that chs_info matches LABRAM_CHANNEL_ORDER exactly; arbitrary channel sets must go through InterpolatedLaBraM.
Signal-JEPA (PR Improve pre-trained channel embeddings handling in Signal-JEPA #991, merged) — channel_embedding constructor argument
with two modes: "scratch" learns a new spatial embedding from the user's
channels; "pretrain_aligned" looks up the user's channel names in the
62-channel pre-training table and errors if a name is unknown.
BIOT — no migration path. Constructed with whatever n_chans the user
provides; the learnable channel-token embedding is re-initialized. Weights
from an 18-bipolar checkpoint do not transfer cleanly to anything else.
EEGPT — chan_proj_type argument; learnable Conv1d channel
projection applied before the backbone. Also: channel-index lookup via
name → index mapping built at __init__ time.
REVE, LUNA — no ad-hoc channel-handling machinery because the
architecture is position-aware by design. But also no canonical-set
metadata shipped with the pre-trained weights — users have no reference
point for "what channels was this trained on".
BENDR, CBraMod, CodeBrain, MEDFormer — no dedicated channel-handling
API. The user is expected to figure out alignment themselves.
Canonical-set constants — each FM that has one stores it privately
with its own name and structure (LABRAM_CHANNEL_ORDER: list[str], _BIOT_TARGET_CHS_TUPLES: tuple[tuple, ...], EEGPT_19_CHANNELS: list[str], _PRETRAIN_CHS_INFO: list[dict], …). Level of exposure varies from
prefixed-private to re-exported.
Cross-cutting, PR [ENH] Experimental channel interpolation for foundation models #993 — InterpolatedModel(model_cls, target_chs_info)
factory: wrapper approach that prepends a (frozen-by-default) MNE-based
channel-interpolation layer to any backbone. First attempt at a uniform
migration path. Explicitly experimental — it does not replace any of the
per-model mechanisms above.
The fact that this list is long, varied, and growing is the problem. Before
we add the next mechanism, it would help to agree on what the contract
actually is.
Braindecode ships a growing number of foundation models (BIOT, BENDR, CBraMod,
CodeBrain, EEGPT, LaBraM, LUNA, MEDFormer, REVE, Signal-JEPA, …). Each one
carries its own implicit assumption about the channel set it was pre-trained
on — fixed names, fixed order, fixed 3D locations, bipolar derivations, … —
and each one exposes a different mechanism for the case where the user's
channels don't match. There is no shared contract a user can rely on to
fine-tune or extract features from these models.
This issue is the starting point of a discussion to agree on a common
contract ("if the model advertises X, the user can do Y with guarantee Z").
As a grounding, the inventory below catalogs what exists today.
Issues created
1. Pre-trained weight loading is fragile and inconsistent
Different models have different failure modes when the user's
chs_infodoesn't match what the weights were trained on:
on_unknown_chs), producing numerically valid but scientificallyquestionable outputs.
chan_proj_type) —choosing between them is user-facing.
equivalent to the pre-training regime.
2. Hugging Face Hub distribution doesn't encode the contract
A pre-trained checkpoint on the Hub is a
state_dict+config.json. Ifthe architectural contract ("this weight file is only meaningful for these
channels in this order") is not part of the config, nothing stops a user
from loading the weights onto the wrong channel set. Today, some models
bake the canonical
chs_infointo the class (LaBraM, BIOT) so the contractsurvives a
from_pretrainedcall; others do not.3. Feature extraction vs fine-tuning is not clearly differentiated
Users have two different workflows:
downstream task. User tolerates some architecture/weight mismatch
because gradients can absorb it.
embeddings for a linear probe or post-hoc analysis. User expects the
pre-trained representation to be preserved faithfully — any silent
reindexing or zero-padding destroys the representation.
Today both paths share the same API, so a user doing a linear probe can
silently get a worse representation than they think, because a rename/reorder
they didn't know about happened under the hood.
Current heterogeneous solutions in braindecode
Every bullet is something one foundation model does that no other foundation
model does the same way.
on_unknown_chs={"zero_pad", "raise"}flag,zero-padding missing channels from the canonical 128-channel set. Removed
in the current refactor (breaking change).
chs_infomatchesLABRAM_CHANNEL_ORDERexactly; arbitrary channel sets must go throughInterpolatedLaBraM.channel_embeddingconstructor argumentwith two modes:
"scratch"learns a new spatial embedding from the user'schannels;
"pretrain_aligned"looks up the user's channel names in the62-channel pre-training table and errors if a name is unknown.
InterpolatedSignalJEPA:MNE-spline interpolation from arbitrary user channels to the pre-training
set.
n_chansthe userprovides; the learnable channel-token embedding is re-initialized. Weights
from an 18-bipolar checkpoint do not transfer cleanly to anything else.
InterpolatedBIOT: MNE-splineinterpolation, with a TODO about the bipolar-vs-monopolar semantic
mismatch (interpolation is not a bipolar derivation).
chan_proj_typeargument; learnableConv1dchannelprojection applied before the backbone. Also: channel-index lookup via
name → index mapping built at
__init__time.architecture is position-aware by design. But also no canonical-set
metadata shipped with the pre-trained weights — users have no reference
point for "what channels was this trained on".
API. The user is expected to figure out alignment themselves.
with its own name and structure (
LABRAM_CHANNEL_ORDER: list[str],_BIOT_TARGET_CHS_TUPLES: tuple[tuple, ...],EEGPT_19_CHANNELS: list[str],_PRETRAIN_CHS_INFO: list[dict], …). Level of exposure varies fromprefixed-private to re-exported.
InterpolatedModel(model_cls, target_chs_info)factory: wrapper approach that prepends a (frozen-by-default) MNE-based
channel-interpolation layer to any backbone. First attempt at a uniform
migration path. Explicitly experimental — it does not replace any of the
per-model mechanisms above.
The fact that this list is long, varied, and growing is the problem. Before
we add the next mechanism, it would help to agree on what the contract
actually is.