Skip to content

Route EEGDINO and LaBraM attention through their qkv module so LoRA and hooks apply - #1194

Merged
bruAristimunha merged 5 commits into
braindecode:masterfrom
bruAristimunha:fix/attention-qkv-module
Oct 5, 2026
Merged

bruAristimunha merged 5 commits into
braindecode:masterfrom
bruAristimunha:fix/attention-qkv-module

Conversation

@bruAristimunha

Copy link
Copy Markdown
Collaborator

Summary

The attention blocks of EEGDINO and LaBraM read the weight of their qkv linear layer without calling the layer:

qkv = F.linear(x, self.qkv.weight, bias)

Anything that wraps or hooks attn.qkv is therefore skipped silently. That includes LoRA with target_modules=["qkv"]: PEFT replaces the module, the forward never calls it, and fine-tuning trains adapters that have no effect on the output. On master, a forward hook on qkv fires in 0 of 12 layers for both models.

Changes

  • EEGDINO._Attention and LaBraM._Attention call self.qkv(x) and add the q/v bias afterwards (LaBraM only when qkv_bias=True).
  • No parameter or state-dict change, so released checkpoints load as before.

Testing

  • New regression tests:

    • test_eegdino.py::test_attention_calls_qkv_module;
    • test_foundation_models.py::test_labram_attention_calls_qkv_module, for qkv_bias False and True.

    With non-zero q/v biases, each test checks three things: a hook on every attn.qkv fires once per forward, the output with the hook is unchanged, and a hook that changes the projection changes the model output. On master, all three test cases fail with assert 0 == 12.

  • Outputs are unchanged apart from float rounding. With the same weights (non-zero q/v biases) on master and on this branch, the largest differences are:

    • EEGDINO logits: 7.5e-9;
    • LaBraM last-block activations: 7.2e-7, on values up to 4.1;
    • LaBraM with qkv_bias=False: bit-identical.
  • pytest test_eegdino.py test_foundation_models.py test_integration.py test_models.py test_return_features.py test_huggingface.py -k "eegdino or labram": 92 passed, 9 skipped. The skips predate this PR: one network-only test, plus TorchScript, export and interpolated-model cases.

  • New or changed behavior is covered by relevant regression tests

  • Style checks recorded: pre-commit run --files braindecode/models/eegdino.py braindecode/models/labram.py test/unit_tests/models/test_eegdino.py test/unit_tests/models/test_foundation_models.py docs/whats_new.rst passed

  • docs/whats_new.rst updated

Notes for reviewers

  • Found while fine-tuning EEGDINO with LoRA in NeuralBench, which currently works around it by patching the attention forward.
  • These are the only two models in braindecode.models that read self.qkv.weight directly.

Both BEiT-style attentions computed F.linear(x, self.qkv.weight, bias),
reading the weight of the qkv Linear without calling it. Forward hooks
and adapters that wrap or hook that module (e.g. LoRA on target_modules
qkv) were therefore silently skipped. Call self.qkv(x) and add the q/v
bias afterwards; outputs change only by float rounding.
Copilot AI balanced review requested due to automatic review settings September 30, 2026 07:53
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Sep 30, 2026
braindecode#1194 adds its entry at the top of the same list; keeping the two apart
lets them merge in either order without a conflict.

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.

Warning

Copilot couldn't run its full agentic review because it didn't start before the timeout. Make sure your repository has a runner available, or add a copilot-code-review.yml file specifying one with the runs-on attribute. See the docs for more details.

Copilot review overview

Review effort: Lite
Findings: 2 Medium severity · 2 Low severity

Open (4)
What changed in this PR

This PR ensures EEGDINO and Labram attention blocks execute their qkv projection via the attn.qkv module (instead of directly reading qkv.weight), so module wrappers/adapters (e.g., LoRA) and forward hooks correctly take effect.

Changes:

  • Update EEGDINO and Labram attention to call self.qkv(x) and apply q/v bias afterward.
  • Add regression tests verifying attn.qkv forward hooks fire and can modify outputs.
  • Document the behavioral fix in docs/whats_new.rst.
File Description
braindecode/​models/​eegdino.py Route attention QKV computation through self.qkv(x) so hooks/adapters are honored.
braindecode/​models/​labram.py Same routing fix for Labram attention; apply optional qkv bias afterward.
test/​unit_tests/​models/​test_eegdino.py Add regression test asserting attn.qkv forward hooks fire and can affect outputs.
test/​unit_tests/​models/​test_foundation_models.py Add parametrized regression test for Labram with/without qkv_bias.
docs/​whats_new.rst Add release note describing the fix and expected output changes (rounding only).

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment thread braindecode/models/eegdino.py Outdated
Comment on lines +479 to +482
bias = torch.cat((self.q_bias, torch.zeros_like(self.v_bias), self.v_bias))
qkv = F.linear(x, self.qkv.weight, bias)
# Call the ``qkv`` module rather than ``F.linear`` on its weight, so
# hooks and adapters attached to it (e.g. LoRA) take part in the forward.
qkv = self.qkv(x) + bias
qkv = nn.functional.linear(input=x, weight=self.qkv.weight, bias=qkv_bias)
# Call the ``qkv`` module rather than ``linear`` on its weight, so hooks
# and adapters attached to it (e.g. LoRA) take part in the forward.
qkv = self.qkv(x)
Comment on lines +51 to +58
assert torch.equal(model(x), reference)
assert len(calls) == len(model.encoder_layers)
for handle in handles:
handle.remove()

for layer in model.encoder_layers: # an adapter that changes the projection
layer.attn.qkv.register_forward_hook(lambda _m, _i, out: 2 * out)
assert not torch.allclose(model(x), reference)
Comment on lines +436 to +443
assert torch.equal(model(x, ch_names=ch_names), reference)
assert len(calls) == len(attentions)
for handle in handles:
handle.remove()

for attn in attentions: # an adapter that changes the projection
attn.qkv.register_forward_hook(lambda _m, _i, out: 2 * out)
assert not torch.allclose(model(x, ch_names=ch_names), reference)
Copilot AI balanced review requested due to automatic review settings September 30, 2026 08:10

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.

)
qkv = nn.functional.linear(input=x, weight=self.qkv.weight, bias=qkv_bias)
# Call the ``qkv`` module rather than ``linear`` on its weight, so hooks
# and adapters attached to it (e.g. LoRA) take part in the forward.
layer.attn.qkv.register_forward_hook(lambda *_: calls.append(None))
for layer in model.encoder_layers
]
assert torch.equal(model(x), reference)
Comment on lines +436 to +443
assert torch.equal(model(x, ch_names=ch_names), reference)
assert len(calls) == len(attentions)
for handle in handles:
handle.remove()

for attn in attentions: # an adapter that changes the projection
attn.qkv.register_forward_hook(lambda _m, _i, out: 2 * out)
assert not torch.allclose(model(x, ch_names=ch_names), reference)
@codecov

codecov Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 87.98%. Comparing base (2426ddc) to head (da714f9).
⚠️ Report is 7 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1194      +/-   ##
==========================================
+ Coverage   87.84%   87.98%   +0.14%     
==========================================
  Files         151      151              
  Lines       17602    17612      +10     
==========================================
+ Hits        15462    15496      +34     
+ Misses       2140     2116      -24     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

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.
bruAristimunha added a commit that referenced this pull request Oct 4, 2026
…ad on any montage (#1195)

* Keep EEGPT channel IDs out of the state dict

chans_id is derived from chs_info at construction but was a persistent
buffer, so it was saved in checkpoints and loaded back:
- the released braindecode/eegpt-pretrained weights (62 IDs) failed to
  load with a size mismatch on any other montage, and with the default
  channel projection (19 IDs), i.e. EEGPT.from_pretrained() with no
  arguments;
- on a montage with the same channel count, loading silently replaced
  the model's channel IDs with the checkpoint's.

Register chans_id as non-persistent and drop the key when loading older
checkpoints that still carry it.

* Add whats_new entry for #1195

* Move the #1195 whats_new entry to the end of the bug-fix list

#1194 adds its entry at the top of the same list; keeping the two apart
lets them merge in either order without a conflict.
Copilot AI balanced review requested due to automatic review settings October 5, 2026 09:06

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Comment thread braindecode/models/eegdino.py Outdated
Comment thread braindecode/models/labram.py Outdated
Copilot AI balanced review requested due to automatic review settings October 5, 2026 09:17

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

@bruAristimunha bruAristimunha added the maintenance Bug fix / refactor / tests — not a new model label Oct 5, 2026
Resolved docs/whats_new.rst by keeping both changelog entries.
Copilot AI balanced review requested due to automatic review settings October 5, 2026 11:17
@bruAristimunha
bruAristimunha merged commit df7d686 into braindecode:master Oct 5, 2026
11 of 12 checks passed

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.

Comment on lines +51 to +58
assert torch.equal(model(x), reference)
assert len(calls) == len(model.encoder_layers)
for handle in handles:
handle.remove()

for layer in model.encoder_layers: # an adapter that changes the projection
layer.attn.qkv.register_forward_hook(lambda _m, _i, out: 2 * out)
assert not torch.allclose(model(x), reference)
Comment on lines +627 to +634
assert torch.equal(model(x, ch_names=ch_names), reference)
assert len(calls) == len(attentions)
for handle in handles:
handle.remove()

for attn in attentions: # an adapter that changes the projection
attn.qkv.register_forward_hook(lambda _m, _i, out: 2 * out)
assert not torch.allclose(model(x, ch_names=ch_names), reference)
attn.qkv.register_forward_hook(lambda *_: calls.append(None))
for attn in attentions
]
assert torch.equal(model(x, ch_names=ch_names), reference)
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Oct 5, 2026
Conflict resolution:
- docs/whats_new.rst: kept both the braindecode#1232 from_pretrained geometry-kwargs bug-fix entry and master's braindecode#1207/braindecode#1212 Deep4Net/ShallowFBCSPNet bug-fix entries.
- test/unit_tests/models/test_eegdino.py: kept both the braindecode#1232 geometry-kwargs test block and master's braindecode#1194 qkv-hook attention test block.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

maintenance Bug fix / refactor / tests — not a new model

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants