Skip to content

LaBraM: keep the pretrained time embedding at every window length, restore mean pooling - #1155

Merged
bruAristimunha merged 9 commits into
masterfrom
copilot/fix-temporal-embeddings-issue
Oct 1, 2026
Merged

bruAristimunha merged 9 commits into
masterfrom
copilot/fix-temporal-embeddings-issue

Conversation

Copilot AI commented Sep 8, 2026 •

Copy link
Copy Markdown
Contributor

#1153 asked whether LaBraM's temporal embeddings are handled correctly. In decoder mode they are: _PatchEmbed averages 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

  • Time embedding sized to the window. temporal_embedding had one slot per patch plus one, while the released weights hold the original 16 absolute slots (time_embed in modeling_finetune.py). They only fit 15-patch windows. On master, Labram.from_pretrained("braindecode/labram-pretrained", n_times=800) fails with size 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.
  • Readout default. [EHN] LaBraM automatic channel reordering and multiple fixes #931 made use_mean_pooling=False the default so that the pretraining checkpoint loads strictly. That readout is the [CLS] output, which LaBraM's pretraining loss never uses; the docstring says True, and the original fine-tuning script sets use_mean_pooling=True (fc_norm of the mean patch token).

What changed

  • In tokenizer mode, temporal_embedding keeps the original 16 absolute slots (patch p still uses slot p), 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=True is the default again. A pretraining checkpoint (with norm, without fc_norm or head, like the released weights) loads into the mean-pooling model as in the original fine-tuning script: norm is unused and fc_norm keeps its initialization. A checkpoint fine-tuned with the [CLS] readout fails to load into the default model instead of being converted silently; pass use_mean_pooling=False to keep that readout.
  • The decoder temporal-embedding test now passes use_mean_pooling=False explicitly. It replaces the final norm with Identity to 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, _PatchEmbed docstring, test) is unchanged.

Evidence

  • Identical to the original code. braindecode vs modeling_finetune.py from 935963004/LaBraM, both with the released weights; the braindecode side loads the Hub checkpoint with a plain strict load_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.
  • Tests. 6 new tests (9 cases). 6 cases fail on the previous head of this PR (9224f3e) and pass here; the other 3 guard compatibility (15-patch windows, old one-slot-per-patch checkpoints, rejecting [CLS]-fine-tuned checkpoints). pytest test/unit_tests/models -k labram: 81 passed, 9 skipped. Full test/unit_tests/models: 2753 passed, 184 skipped, 0 failed (same code as this head; only a docstring and whats_new were reworded afterwards).
  • No new ruff findings (ruff 0.14.9, as in .pre-commit-config.yaml).

Fixes #1153

Copilot AI changed the title [WIP] Fix potential issue with temporal embeddings handling Clarify LaBram decoder temporal embedding indexing Sep 8, 2026
Copilot AI requested a review from bruAristimunha September 8, 2026 09:49
@codecov

codecov Bot commented Sep 21, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 87.73%. Comparing base (25773c0) to head (06520d3).
⚠️ Report is 2 commits behind head on master.

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:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

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 bruAristimunha changed the title Clarify LaBram decoder temporal embedding indexing LaBraM: keep the pretrained time embedding at every window length, restore mean pooling Sep 30, 2026
@bruAristimunha
bruAristimunha marked this pull request as ready for review October 1, 2026 14:45
Copilot AI balanced review requested due to automatic review settings October 1, 2026 14:45

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.

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
bruAristimunha merged commit 0af5317 into master Oct 1, 2026
15 checks passed
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.
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.

labram.py temporal embeddings potential issue

3 participants