Repository navigation
LaBraM: keep the pretrained time embedding at every window length, restore mean pooling - #1155
Merged
Merged
Conversation
Co-authored-by: bruAristimunha <[email protected]>
Copilot
AI
changed the title
[WIP] Fix potential issue with temporal embeddings handling
Clarify LaBram decoder temporal embedding indexing
Sep 8, 2026
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## master #1155 +/- ##
==========================================
+ Coverage 87.71% 87.73% +0.01%
==========================================
Files 149 149
Lines 17433 17452 +19
==========================================
+ Hits 15292 15312 +20
+ Misses 2141 2140 -1 🚀 New features to boost your workflow:
|
Resolutions: - docs/whats_new.rst: keep master's 1.8.1 section (DIVER1, ZUNA, MSCFormer, EEGMiner HPU entries) and re-add the :gh:`1155` Labram enhancement entry at the top of Enhancements; a previous attempt had corrupted the RST title underlines and was discarded. - braindecode/models/labram.py, test/unit_tests/models/test_foundation_models.py: auto-merged, no conflict.
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.
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.
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.
bruAristimunha
marked this pull request as ready for review
October 1, 2026 14:45
Contributor
There was a problem hiding this comment.
Copilot review overview
🟢 Approval recommended
The implementation matches the documented behavior and includes focused regression coverage.
Review effort: Balanced
Findings: None
What changed in this PR
Restores LaBraM’s original temporal embeddings and mean-pooling readout while preserving checkpoint compatibility.
Changes:
- Keeps 16 pretrained absolute time slots across window lengths.
- Restores mean pooling as the default readout.
- Adds compatibility and decoder-indexing tests.
| File | Description |
|---|---|
braindecode/models/labram.py |
Updates embeddings, pooling, and checkpoint loading. |
test/unit_tests/models/test_foundation_models.py |
Adds regression and compatibility coverage. |
docs/whats_new.rst |
Documents behavioral and embedding fixes. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
bruAristimunha
added a commit
to mahirjain01/braindecode
that referenced
this pull request
Oct 1, 2026
Resolves docs/whats_new.rst: keeps master's braindecode#1159 and braindecode#1155 entries and adds the AXON (braindecode#1182) entry above them; author link kept. test_foundation_models.py auto-merged.
bruAristimunha
added a commit
to julien-gadonneix/braindecode
that referenced
this pull request
Oct 1, 2026
Resolutions: - docs/whats_new.rst: master's file kept as is; the MAPA entry (:gh:`1178`) inserted first under Enhancements, above master's braindecode#1159 and braindecode#1155 entries.
bruAristimunha
added a commit
to Fashad-Ahmed/braindecode
that referenced
this pull request
Oct 1, 2026
Resolve docs/whats_new.rst: keep master's braindecode#1159/braindecode#1155 entries and the SleepFM (braindecode#1106) entry. test_foundation_models.py auto-merged (master's LaBraM tests plus the SleepFM tests); _DIRECT_TORCHSCRIPT_MODELS stays 32 on base, PR and master.
bruAristimunha
added a commit
to bruAristimunha/braindecode
that referenced
this pull request
Oct 5, 2026
Resolves the conflict from master's braindecode#1155 (LaBraM time embedding + mean-pooling) landing after this PR's branch point. - docs/whats_new.rst: kept both bug-fix entries (braindecode#1194 qkv routing, braindecode#1159 predict_trials variable-length fix), concatenated. - braindecode/models/labram.py: auto-merged cleanly by git; verified both master's time_embed/use_mean_pooling additions and the PR's self.qkv(x) module-call routing (replacing F.linear on qkv.weight) are present together. - test/unit_tests/models/test_foundation_models.py: auto-merged cleanly, both sides' additions kept. - All other changed files (.github/workflows/tests.yml, braindecode/classifier.py, braindecode/datasets/tuh.py, braindecode/regressor.py, braindecode/training/losses.py, braindecode/training/scoring.py, test/acceptance_tests/*, test/unit_tests/datasets/test_tuh.py, test/unit_tests/test_eegneuralnet.py, test/unit_tests/training/test_losses.py) are master-side changes since the PR's branch point, auto-merged without conflict.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
#1153 asked whether LaBraM's temporal embeddings are handled correctly. In decoder mode they are:
_PatchEmbedaverages the grouped channels, so there is one token per temporal patch (Copilot's clarification and test, kept here). The actual defect was in the default tokenizer mode, together with a readout default that had drifted from the original.What was wrong
temporal_embeddinghad one slot per patch plus one, while the released weights hold the original 16 absolute slots (time_embedinmodeling_finetune.py). They only fit 15-patch windows. On master,Labram.from_pretrained("braindecode/labram-pretrained", n_times=800)fails withsize mismatch for temporal_embedding: [1, 16, 200] vs [1, 5, 200], so any other window length has to drop the pretrained slots and learn them from scratch.use_mean_pooling=Falsethe default so that the pretraining checkpoint loads strictly. That readout is the [CLS] output, which LaBraM's pretraining loss never uses; the docstring saysTrue, and the original fine-tuning script setsuse_mean_pooling=True(fc_normof the mean patch token).What changed
temporal_embeddingkeeps the original 16 absolute slots (patchpstill uses slotp), with extra slots only for windows longer than 16 patches. A load hook copies every slot a checkpoint holds, so the released weights load at every window length, and checkpoints saved with the previous one-slot-per-patch layout load with identical outputs. When a window has more patches than the checkpoint has slots, a warning names the slots that keep their initialization. Decoder mode is unchanged.use_mean_pooling=Trueis the default again. A pretraining checkpoint (withnorm, withoutfc_normor head, like the released weights) loads into the mean-pooling model as in the original fine-tuning script:normis unused andfc_normkeeps its initialization. A checkpoint fine-tuned with the [CLS] readout fails to load into the default model instead of being converted silently; passuse_mean_pooling=Falseto keep that readout.use_mean_pooling=Falseexplicitly. It replaces the finalnormwithIdentityto read the tokens, which only holds for the [CLS] readout.whats_new: the time-embedding fix under Bug fixes, the readout default under API and behavior changes.Copilot's decoder clarification (variable names,
_PatchEmbeddocstring, test) is unchanged.Evidence
modeling_finetune.pyfrom 935963004/LaBraM, both with the released weights; the braindecode side loads the Hub checkpoint with a plain strictload_state_dict. Max |diff| is 0 for every token, for the [CLS] readout and for the mean-pooling readout, at 18 geometries: 22-channel (BCI IV 2a), 19-channel (10-20) and 64-channel montages × 1, 3, 4, 5, 10 and 15 patches.Labram.from_pretrained(..., n_times=800)gives mean pooling, 16 slots and max |diff| 0.pytest test/unit_tests/models -k labram: 81 passed, 9 skipped. Fulltest/unit_tests/models: 2753 passed, 184 skipped, 0 failed (same code as this head; only a docstring andwhats_newwere reworded afterwards)..pre-commit-config.yaml).Fixes #1153