Skip to content

[EHN] LaBraM automatic channel reordering and multiple fixes - #931

Merged
bruAristimunha merged 37 commits into
braindecode:masterfrom
dungscout96:model/labram
Feb 6, 2026
Merged

bruAristimunha merged 37 commits into
braindecode:masterfrom
dungscout96:model/labram

Conversation

@dungscout96

Copy link
Copy Markdown
Contributor

Adds automatic channel name handling and reordering to the LaBraM model, ensuring input EEG channels are processed in the correct order expected by pretrained weights.

Key changes:

  • Add LABRAM_CHANNEL_ORDER constant defining the 98 standard 10-20 electrode positions used by pretrained LaBraM
  • Add ch_names parameter to Labram.__init__() for explicit channel name specification
  • Automatically extract channel names from chs_info (MNE format) when available
  • Automatically reorder input channels to match the canonical LaBraM order during forward pass
  • Support case-insensitive channel name matching
  • Warn users about unmatched channels that will be excluded

Test plan

  • Run existing LaBraM tests to ensure no regressions
  • Add new tests for channel reordering functionality:
    • Channel order constant export
    • Initialization with ch_names
    • Forward pass with reordering
    • Correct reordering order verification
    • Case-insensitive matching
    • Warning for unmatched channels
    • chs_info extraction
    • Gradient flow through reordering
    • GPU device compatibility

@PierreGtch PierreGtch left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Thanks @dungscout96 :)

I prefer not adding a ch_names argument as it is a duplicate of chs_info

Comment thread braindecode/models/labram.py Outdated
Comment thread braindecode/models/labram.py Outdated
bruAristimunha and others added 9 commits February 5, 2026 10:19
Co-authored-by: Pierre Guetschel <[email protected]>
Reverts EEGPT channel projection, Conv1dWithConstraint, and filterbank
fix to keep this branch focused on Labram channel reordering only.
Restore code removed in commit 099bd48 that preserves MNE info fields
(description, line_freq, device_info, helium_info, experimenter, proj_name)
when creating filtered raw objects. This avoids merge conflicts when
adding channels in the filterbank function.
Restore whats_new.rst entry for filterbank fix, attributed to Young Truong.
@bruAristimunha

Copy link
Copy Markdown
Collaborator

I have mixed feelings about this, but I am ok to merge because we need to; we are dropping the channels that are not within the intersection when the ch_info is informed.

I think it made me realize that we don't have a convention of what to do or how to handle things when the channels don't align with pre-trained.

Last week, @hubertjb challenged me to be more consistent with reading weights and handling these in-pair channel details. This certainly goes beyond this PR, and we will probably need to look at how other libraries do this and implement the same logic across models.

@bruAristimunha

Copy link
Copy Markdown
Collaborator

@PierreGtch, your input is welcome :)

Remove pre-initialization of _channel_indices and _labram_ch_indices
as regular Python attributes. Instead, register them as None buffers
when channel info is not provided.

This fixes KeyError: "attribute '_channel_indices' already exists"
that occurred when initializing Labram with chs_info, because
register_buffer() raises an error when an attribute already exists.
@PierreGtch

Copy link
Copy Markdown
Collaborator

Ideally, the channels used to pre-train the model should be contained in the weights file and loaded along with the weights

Super ideally, the model does not rely on a fixed set of channel names but directly uses their positions (pos = [ch['loc'] for ch in chs_info]), like Signal-JEPA 😇

In any case, if we drop channels, we should raise a warning or error. Maybe we could add a flag:

on_unknown_ch: Literal['raise', 'warn', 'ignore']
# Behaviour when pre-trained weights are loaded,
# and channels that were not used to pre-train are passed 

Comment thread braindecode/models/labram.py Outdated
PierreGtch and others added 4 commits February 5, 2026 21:39
- Introduced `on_unknown_chs` parameter to manage behavior for unmatched channels.
- Updated `_setup_channel_mapping` and `_get_channel_indices` methods for improved channel validation.
- Fix tests
@PierreGtch

Copy link
Copy Markdown
Collaborator

I implemented my comments regarding the ch_name and on_missing_ch parameters.

However... the weights loading test fails because the positional embedding weights have 64 channels, but the LABRAM_CHANNEL_ORDER lists 120 channels.

@dungscout96 do you know which exact channels were used to obtain the retrained weights in https://huggingface.co/braindecode/Labram-Braindecode/resolve/main/braindecode_labram_base.pt ?

I could not find this information in the original repo 🙃

@dungscout96

dungscout96 commented Feb 6, 2026 •

Copy link
Copy Markdown
Contributor Author

I implemented my comments regarding the ch_name and on_missing_ch parameters.

However... the weights loading test fails because the positional embedding weights have 64 channels, but the LABRAM_CHANNEL_ORDER lists 120 channels.

@dungscout96 do you know which exact channels were used to obtain the retrained weights in https://huggingface.co/braindecode/Labram-Braindecode/resolve/main/braindecode_labram_base.pt ?

I could not find this information in the original repo 🙃

You would hope they make it easier to find these information. After some digging:

SO the pretrained weights should have 128 channels?? It looks like in their paper they selected a 64-channel subset to index into the channel positional embeddings. But still it should be 128 dimension.

Ok I checked the official checkpoint of Labram, it's 128 channel.
image

How did we get the weights in braindecode?

@bruAristimunha

Copy link
Copy Markdown
Collaborator

I don't remember, I think we have the script at hugging face GitHub historical

@bruAristimunha

Copy link
Copy Markdown
Collaborator

But let's regenerate please

@PierreGtch

Copy link
Copy Markdown
Collaborator

get_input channels points to a list with 136 channels.... it's a mess
https://github.com/935963004/LaBraM/blob/c431221e6cfd23dbfa9950e0180682fb322b0548/utils.py#L42-L57

I will try to figure it out and regenerate our huggingface weights

@PierreGtch PierreGtch changed the title LaBraM automatic channel reordering LaBraM automatic channel reordering and multiple fixes Feb 6, 2026
@PierreGtch

Copy link
Copy Markdown
Collaborator

Test Plan

This requires cloning the original Labram repo, so it will not be included in our test suit

# %% ######################################################
# Imports constants and setup
#######################################################
import sys
import torch
import re
from pathlib import Path
import difflib

import torch
from torch import nn
from braindecode.models.labram import Labram, LABRAM_CHANNEL_ORDER

# Add LaBraM repo to path
# > cd ~/Projects/
# > git clone [email protected]:935963004/LaBraM.git
labram_path = Path("~/Projects/LaBraM").expanduser()
sys.path.insert(0, str(labram_path))

from timm.models import create_model
import modeling_finetune  # Register the models to timm

DROPPED_KEYS = ["mask_token", "lm_head.weight", "lm_head.bias"]

# Load checkpoint from GitHub
print("Loading checkpoint from GitHub...")
checkpoint_url = (
    "https://github.com/935963004/LaBraM/raw/main/checkpoints/labram-base.pth"
)
checkpoint = torch.hub.load_state_dict_from_url(
    checkpoint_url,
    map_location="cpu",
    progress=True,
    check_hash=True,
    weights_only=False,
)

# Extract state dict
state_dict = {
    key[8:]: value
    for key, value in checkpoint["model"].items()
    if key.startswith("student.")
}
# Drop unused keys
for k in DROPPED_KEYS:
    del state_dict[k]


# %% ######################################################
# Instantiate the original LaBraM model with pre-trained weights.
#######################################################

# Create model
print("Creating model...")
model = create_model(
    "labram_base_patch200_200",
    # arguments form
    # https://github.com/935963004/LaBraM/blob/c431221e6cfd23dbfa9950e0180682fb322b0548/run_class_finetuning.py#L196C1-L208C6
    # and readme
    pretrained=False,
    num_classes=0,
    drop_rate=0,
    drop_path_rate=0.1,
    attn_drop_rate=0.0,
    drop_block_rate=None,
    use_mean_pooling=False,
    init_scale=0.001,
    use_rel_pos_bias=False,
    use_abs_pos_emb=True,
    init_values=0.1,
    qkv_bias=False,
)
model.eval()

# Load pre-trained weights
model.load_state_dict(state_dict, strict=True)


# %% ######################################################
# Instantiate the Braindecode implementation of LaBraM, rename the weights and load the pre-trained weights.
######################################################

# Instantiate model
chs_info = [{"ch_name": ch_name} for ch_name in LABRAM_CHANNEL_ORDER]
model_braindecode = Labram(
    n_times=3000,
    chs_info=chs_info,
    n_outputs=0,
)
_ = model_braindecode.eval()


# Map keys from original LaBraM to Braindecode implementation
braindecode_state_dict = {}
key_mappings = {}
for k, v in state_dict.items():
    new_k = k

    # simple renames
    if new_k == "pos_embed":
        new_k = "position_embedding"
    elif new_k == "time_embed":
        new_k = "temporal_embedding"

    # blocks mlp fc1/fc2 -> mlp.0 / mlp.2
    new_k = re.sub(r"blocks\.(\d+)\.mlp\.fc1", r"blocks.\1.mlp.0", new_k)
    new_k = re.sub(r"blocks\.(\d+)\.mlp\.fc2", r"blocks.\1.mlp.2", new_k)

    # patch_embed conv/norm mappings:
    m = re.match(r"patch_embed\.conv(\d+)(\..+)$", k)
    if m:
        n = int(m.group(1))
        suffix = m.group(2)
        # map convs to temporal_conv.conv{n}
        new_k = f"patch_embed.temporal_conv.conv{n}{suffix}"

    m2 = re.match(r"patch_embed\.norm(\d+)(\..+)$", k)
    if m2:
        n = int(m2.group(1))
        suffix = m2.group(2)
        new_k = f"patch_embed.temporal_conv.norm{n}{suffix}"

    # fallback: use the simple/new_k computed above
    braindecode_state_dict[new_k] = v
    if new_k != k:
        key_mappings[k] = new_k

print("Key mappings:")
max_len = max(len(old_key) for old_key in key_mappings.keys())
for old_key, new_key in key_mappings.items():
    print(f"{old_key.ljust(max_len)} -> {new_key}")

# Load mapped state dict into model
model_braindecode.load_state_dict(braindecode_state_dict, strict=True)


# %%########################################################
# Compare the outputs of the original and Braindecode implementations for the same input.
##########################################################

rng = torch.Generator().manual_seed(42)
x = torch.randn(1, 128, 3000, generator=rng)

with torch.no_grad():
    y1 = model(x.view(1, 128, 15, 200), return_all_tokens=True)
    y2 = model_braindecode(x, return_all_tokens=True)  # return_patch_tokens=True)

assert y1.shape == y2.shape
assert (y1 == y2).all()

# %%########################################################
# Push weights to hub
#########################################################

model_braindecode.push_to_hub(
    repo_id="braindecode/labram-pretrained", commit_message="Several fixes (see PR#931)"
)


# %%########################################################
# Now load form the newly pushed HF checkpoint
#########################################################

model_hf = Labram.from_pretrained(
    "braindecode/labram-pretrained",
    cache_dir=Path("~/.cache/huggingface/").expanduser(),
)


# %%########################################################
# Finally, compare the outputs of the original and HF model for the same input.
##########################################################

rng = torch.Generator().manual_seed(42)
x = torch.randn(1, 128, 3000, generator=rng)

with torch.no_grad():
    y1 = model(x.view(1, 128, 15, 200), return_all_tokens=True)
    y2 = model_hf(x, return_all_tokens=True)  # return_patch_tokens=True)

assert y1.shape == y2.shape
assert (y1 == y2).all()


# %%########################################################
# Sandbox
#########################################################

def format_model_str(model):
    return list(
        map(lambda line: line.replace(" ", ""), repr(model).splitlines(keepends=True))
    )


model_str = format_model_str(model)
model_braindecode_str = format_model_str(model_braindecode)

# diff = difflib.unified_diff(
diff = difflib.ndiff(
    model_str,
    model_braindecode_str,
    # fromfile="Original LaBraM",
    # tofile="Braindecode LaBraM",
)
print("".join(diff))

Fix AttributeError: 'Labram' object has no attribute 'emb_dim' by
using self.embed_dim and the n_outputs parameter instead of
self.n_outputs.
@bruAristimunha bruAristimunha changed the title LaBraM automatic channel reordering and multiple fixes [EHN] LaBraM automatic channel reordering and multiple fixes Feb 6, 2026
Replace manual torch.hub download + key renaming with the
from_pretrained class method, consistent with Luna and REVE tests.
Weights are cached under ~/mne_data/labram_pretrained/ for CI persistence.
@bruAristimunha
bruAristimunha merged commit ff03fe3 into braindecode:master Feb 6, 2026
13 checks passed
@PierreGtch

Copy link
Copy Markdown
Collaborator

Thanks for handling the tests updates @bruAristimunha :D

@PierreGtch PierreGtch mentioned this pull request Mar 4, 2026
bruAristimunha added a commit that referenced this pull request Sep 30, 2026
The original LaBraM fine-tunes on fc_norm(mean of the patch tokens)
(use_mean_pooling=True in modeling_finetune.py and in
run_class_finetuning.py), and its pretraining loss never uses [CLS].
braindecode documents the same default, but #931 set the code default to
False so that the pretraining checkpoint loaded strictly. Since then the
default readout was the untrained [CLS] token.

The default is True again. A pretraining checkpoint (per-token norm, no
fc_norm, no head), such as the released weights, loads into a
mean-pooling model as in the original fine-tuning script: norm is unused
and fc_norm keeps its initialization, also when the state dict was
already filtered to the model's keys. A checkpoint with a head was
fine-tuned with its own readout and still fails to load instead of being
converted. Checkpoints saved by braindecode record use_mean_pooling and
load unchanged. With the released weights the readout equals the
original fine-tuning readout at 1 to 15 patches.
bruAristimunha added a commit that referenced this pull request Oct 1, 2026
…store mean pooling (#1155)

* Initial plan

* Clarify Labram decoder temporal patch semantics

Co-authored-by: bruAristimunha <[email protected]>

* Fix release-note and lint checks for PR #1155

* Load LaBraM's pretrained time embedding at every window length

The original LaBraM keeps 16 absolute time slots (time_embed in
modeling_finetune.py) and a window of P patches uses slots 0..P-1. The
released weights hold those 16 slots. braindecode sized the embedding to
the window (P + 1 slots), so the pretrained slots only loaded for 15 s
windows: any other window failed with a size mismatch, and benchmarks
worked around it by training the time embedding from scratch.

The tokenizer now keeps max(16, P) slots and uses slots 0..P-1, as the
original. Loading takes every slot a checkpoint holds and keeps the
model's own values for the others, so checkpoints saved with P + 1 slots
still load with identical outputs; a warning names the slots a longer
window uses that the checkpoint does not provide. The decoder mode is
unchanged. With the released weights, tokens and readouts equal the
original code at 1 to 15 patches on 19-, 22- and 64-channel montages.

* Restore LaBraM's mean-pooling readout as the default

The original LaBraM fine-tunes on fc_norm(mean of the patch tokens)
(use_mean_pooling=True in modeling_finetune.py and in
run_class_finetuning.py), and its pretraining loss never uses [CLS].
braindecode documents the same default, but #931 set the code default to
False so that the pretraining checkpoint loaded strictly. Since then the
default readout was the untrained [CLS] token.

The default is True again. A pretraining checkpoint (per-token norm, no
fc_norm, no head), such as the released weights, loads into a
mean-pooling model as in the original fine-tuning script: norm is unused
and fc_norm keeps its initialization, also when the state dict was
already filtered to the model's keys. A checkpoint with a head was
fine-tuned with its own readout and still fails to load instead of being
converted. Checkpoints saved by braindecode record use_mean_pooling and
load unchanged. With the released weights the readout equals the
original fine-tuning readout at 1 to 15 patches.

* Add the LaBraM time-embedding and readout fixes to whats_new

* Pin the readout in the LaBraM decoder temporal-embedding test

The test replaces the final norm with Identity to read the tokens as they
enter the readout. That only holds for the [CLS] readout, which is no
longer the default, so state it explicitly.

---------

Co-authored-by: copilot-swe-agent[bot] <[email protected]>
Co-authored-by: bruAristimunha <[email protected]>
Co-authored-by: Bru <[email protected]>
Co-authored-by: Bruno Aristimunha <[email protected]>
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.

3 participants