Repository navigation
Route EEGDINO and LaBraM attention through their qkv module so LoRA and hooks apply - #1194
Conversation
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.
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.
There was a problem hiding this comment.
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
Open (4)
Now thatself.qkv(x)is called, any bias configured on theself.qkvmodule itself will be… · New Similar to EEGDINO: switching from `nn.functional.linear(..., weight=self.qkv.weight,… · New The second set of hooks (the ones that modify the projection) aren’t removed. Even though the model… · New As in the EEGDINO test, the second set of forward hooks isn’t removed. Capturing those handles and… · New
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.qkvforward 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.
| 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) |
| 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) |
| 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) |
There was a problem hiding this comment.
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: 4
Open (7)
This change will start applyingself.qkv's ownbiasparameter (if the Linear was constructed… · Newtorch.equalrequires bit-identical results; this can be brittle across devices/backends (e.g.,… · New Similar to EEGDINO: switching from `nn.functional.linear(..., weight=self.qkv.weight,… Now thatself.qkv(x)is called, any bias configured on theself.qkvmodule itself will be… The second set of forward hooks isn't removed (no handles are kept), which can leave hooks attached… · New As in the EEGDINO test, the second set of forward hooks isn’t removed. Capturing those handles and… The second set of hooks (the ones that modify the projection) aren’t removed. Even though the model…
| ) | ||
| 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) |
| 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 Report✅ All modified and coverable lines are covered by tests. 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:
|
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.
…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.
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
The separate fp32 bias additions promote mixed-precision QKV activations back to fp32.
Review effort: Balanced
Findings: 6
Open (9)
Bias addition promotes autocast projection output to fp32 · New Standalone bias changes return_qkv output dtype under autocast · Newtorch.equalrequires bit-identical results; this can be brittle across devices/backends (e.g.,… This change will start applyingself.qkv's ownbiasparameter (if the Linear was constructed… Similar to EEGDINO: switching from `nn.functional.linear(..., weight=self.qkv.weight,… Now thatself.qkv(x)is called, any bias configured on theself.qkvmodule itself will be… The second set of forward hooks isn't removed (no handles are kept), which can leave hooks attached… As in the EEGDINO test, the second set of forward hooks isn’t removed. Capturing those handles and… The second set of hooks (the ones that modify the projection) aren’t removed. Even though the model…
There was a problem hiding this comment.
Copilot review overview
🟢 Approval recommended
The implementation preserves existing parameters and outputs while adequately testing the corrected hook behavior.
Review effort: Balanced
Findings: 4
Open (7)
torch.equalrequires bit-identical results; this can be brittle across devices/backends (e.g.,… This change will start applyingself.qkv's ownbiasparameter (if the Linear was constructed… Similar to EEGDINO: switching from `nn.functional.linear(..., weight=self.qkv.weight,… Now thatself.qkv(x)is called, any bias configured on theself.qkvmodule itself will be… The second set of forward hooks isn't removed (no handles are kept), which can leave hooks attached… As in the EEGDINO test, the second set of forward hooks isn’t removed. Capturing those handles and… The second set of hooks (the ones that modify the projection) aren’t removed. Even though the model…
Resolved since last review (2)
Resolved docs/whats_new.rst by keeping both changelog entries.
There was a problem hiding this comment.
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: 6
Open (10)
The second set of forward hooks (lines 56–57) is never removed, which can leak state into later… · New Like the EEGDINO test, the second set of hooks registered in lines 632–633 is not removed. Please… · Newtorch.equalrequires bit-identical results; this can be brittle across devices/backends (e.g.,… This change will start applyingself.qkv's ownbiasparameter (if the Linear was constructed… Similar to EEGDINO: switching from `nn.functional.linear(..., weight=self.qkv.weight,… Now thatself.qkv(x)is called, any bias configured on theself.qkvmodule itself will be…torch.equalrequires bit-identical outputs and can be brittle if backend nondeterminism or minor… · New The second set of forward hooks isn't removed (no handles are kept), which can leave hooks attached… As in the EEGDINO test, the second set of forward hooks isn’t removed. Capturing those handles and… The second set of hooks (the ones that modify the projection) aren’t removed. Even though the model…
| 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) |
| 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) |
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.


Summary
The attention blocks of
EEGDINOandLaBraMread the weight of theirqkvlinear layer without calling the layer:Anything that wraps or hooks
attn.qkvis therefore skipped silently. That includes LoRA withtarget_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 onqkvfires in 0 of 12 layers for both models.Changes
EEGDINO._AttentionandLaBraM._Attentioncallself.qkv(x)and add the q/v bias afterwards (LaBraM only whenqkv_bias=True).Testing
New regression tests:
test_eegdino.py::test_attention_calls_qkv_module;test_foundation_models.py::test_labram_attention_calls_qkv_module, forqkv_biasFalse and True.With non-zero q/v biases, each test checks three things: a hook on every
attn.qkvfires 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 withassert 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:
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.rstpasseddocs/whats_new.rstupdatedNotes for reviewers
braindecode.modelsthat readself.qkv.weightdirectly.