Repository navigation
[EHN] LaBraM automatic channel reordering and multiple fixes - #931
Conversation
PierreGtch
left a comment
There was a problem hiding this comment.
Thanks @dungscout96 :)
I prefer not adding a ch_names argument as it is a duplicate of chs_info
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.
|
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. |
|
@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.
|
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 ( 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 |
- 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
|
I implemented my comments regarding the However... the weights loading test fails because the positional embedding weights have 64 channels, but the @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. How did we get the weights in braindecode? |
|
I don't remember, I think we have the script at hugging face GitHub historical |
|
But let's regenerate please |
|
get_input channels points to a list with 136 channels.... it's a mess I will try to figure it out and regenerate our huggingface weights |
…n the middle of _adj_position_embedding leading to an incorrect order of positional embeddings
Test PlanThis 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.
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.
|
Thanks for handling the tests updates @bruAristimunha :D |
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.
…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]>

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:
LABRAM_CHANNEL_ORDERconstant defining the 98 standard 10-20 electrode positions used by pretrained LaBraMch_namesparameter toLabram.__init__()for explicit channel name specificationchs_info(MNE format) when availableTest plan
ch_nameschs_infoextraction