You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
Add BrainBERT: foundation model for intracranial (sEEG/iEEG) signals - #1104
Towards #1097 — adds a braindecode-native port of BrainBERT (Wang et al., BrainBERT: Self-supervised representation learning for intracranial recordings,
ICLR 2023), a self-supervised foundation model for intracranial (sEEG/iEEG)
signals. Suggested on #1097.
(B, C, T) → STFT spectrogram → linear projection + sinusoidal positional encoding → Transformer encoder → pool over channels & frames → head.
Following the convention established for Brant (#1100), the short-time
Fourier transform front-end is computed insideforward (module _STFTSpectrogram) so the model keeps the standard (batch, n_chans, n_times)
input signature, whereas the upstream reference consumes a pre-computed
spectrogram.
What's in this PR
braindecode/models/brainbert.py — BrainBERT(EEGModuleMixin, nn.Module)
with modest ready-to-run defaults (~0.65M params) and the released large
config documented (hidden_dim=768, ffn_dim=3072, n_heads=12, n_layers=6,
~43M), plus from_pretrained, reset_head and return_features.
braindecode/modules/brainbert_modules.py — STFT front-end, input embedding,
sinusoidal positional encoding, spectrogram-prediction head (weight parity)
and classification head.
test/unit_tests/models/test_brainbert.py — contract tests + a parity
gate.
Numerical fidelity
The in-model STFT reproduces the upstream scipy magnitude spectrogram to 1.5e-6 (bit-exact in float32).
Parity gate: with the upstream reference available (BRAINBERT_SRC), the
input encoding + Transformer are checked bit-exact (atol=1e-5) against
the upstream MaskedTFModel; the gate is skipped otherwise so CI stays green
without the external dependency.
Standalone braindecode modules ported from the upstream BrainBERT reference:
the in-model STFT spectrogram front-end (reproducing the scipy front-end), the
input embedding with sinusoidal positional encoding, the spectrogram-prediction
head (kept for weight parity) and the classification head.
BrainBERT(EEGModuleMixin, nn.Module): self-supervised foundation model for
intracranial (sEEG/iEEG) signals (Wang et al., ICLR 2023). The STFT is computed
inside forward so the model keeps the standard (batch, n_chans, n_times) input
signature. Modest ready-to-run defaults, with the released large config
documented; supports from_pretrained, reset_head and return_features.
Contract tests plus a parity gate: the in-model STFT is checked against the
scipy reference, and the input encoding + Transformer are checked bit-exact
against the upstream MaskedTFModel when BRAINBERT_SRC is available (skipped
otherwise so CI stays green without the external dependency).
The reason will be displayed to describe this comment to others. Learn more.
Pull request overview
Note
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.
Adds a braindecode-native port of BrainBERT (ICLR 2023) for intracranial (sEEG/iEEG) signals, including its in-model STFT front-end, Transformer encoder stack, and associated registration + docs + tests.
Changes:
Introduce BrainBERT model and supporting modules (STFT spectrogram, input embedding, heads).
Register BrainBERT in the model registry and model summary metadata.
Add unit tests including a SciPy STFT equivalence check and an optional upstream parity gate.
Reviewed changes
Copilot reviewed 7 out of 7 changed files in this pull request and generated 7 comments.
Show a summary per file
File
Description
braindecode/models/brainbert.py
Adds the BrainBERT model implementation (STFT-in-forward, Transformer, pooling + head, return_features).
BrainBERT.forward returns a Dict (features) or a Tensor (logits) depending
on return_features, and torch.jit.script rejects this polymorphic return
type. Add BrainBERT to not_working_models, matching Brant/EEGDINO/MVPFormer
which are excluded for the same reason.
The reason will be displayed to describe this comment to others. Learn more.
Pull request overview
Copilot reviewed 8 out of 8 changed files in this pull request and generated 4 comments.
Comments suppressed due to low confidence (2)
test/unit_tests/models/test_brainbert.py:98
The PR description claims the in-model STFT matches the upstream scipy spectrogram to ~1.5e-6, but this test currently allows atol=1e-4 (and the default rtol). Tightening the tolerance would better protect the numerical-fidelity guarantee against regressions.
assert np.allclose(ref, ours, atol=1e-4)
test/unit_tests/models/test_brainbert.py:110
_import_upstream_or_skip() prepends BRAINBERT_SRC to sys.path but never removes it. When the parity gate is enabled, this can leak into the rest of the test session and cause confusing import shadowing/order issues. It’s safer to restore sys.path after the import attempt.
sys.path.insert(0, src)
try:
import models as upstream_models # noqa: F401 (populates registry)
except ImportError as exc: # pragma: no cover - depends on external code
pytest.skip(f"upstream BrainBERT not importable: {exc}")
I fed identical inputs through the upstream BrainBERT stack (BrainBERT/models/masked_tf_model.py, from the official repo) and the braindecode port with the released stft_large weights copied over.
Port parity (gated in CI):
Check
Upstream ref
braindecode port
max|diff|
STFT front-end
scipy.signal.stft
_STFTSpectrogram
~1.5e-6
Transformer encoder
MaskedTFModel
BrainBERT
bit-exact (atol 1e-5)
from_pretrained
official stft_large (43.19M)
braindecode/brainbert-pretrained (43.19M)
6.4e-5
The from_pretrained residual (6.4e-5, vs 0.0 for the freshly-built encoder) comes solely from the sinusoidal positional encoding being regenerated on CPU vs baked on GPU at pretraining time; it grows with position and stays negligible at real sequence lengths. Reproduce with BRAINBERT_SRC=<upstream> pytest test/unit_tests/models/test_brainbert.py; the parity gate skips cleanly without BRAINBERT_SRC. Weights re-hosted at braindecode/brainbert-pretrained: BrainBERT.from_pretrained("braindecode/brainbert-pretrained", n_outputs=2).
Downstream replication (Brain Treebank, CC-BY): beyond activation parity, I ran the paper's linear-probing protocol end-to-end (frozen features → Linear(768→1), ROC-AUC per electrode) on sub_3/sub_4, all valid Laplacian electrodes, on the 4 paper tasks. Swapping the upstream extractor for the braindecode port is a no-op on the metric — over 76 shared electrode×task pairs, max|ΔAUC| = 0.0023, mean|ΔAUC| = 0.0004. Absolute levels match the paper on good electrodes (best-electrode onset 0.84, speech 0.80).
The reason will be displayed to describe this comment to others. Learn more.
I reviewed this exact candidate against the BrainBERT paper, released source,
dataset records, and checkpoint metadata. I am requesting changes because the
mandatory paper-result gate is not yet legally or scientifically runnable.
The easiest faithful gate is Table 2: frozen-STFT BrainBERT sentence-onset
ROC-AUC 0.66 +/- 0.03, aggregated as the mean and population standard
deviation across the ten electrodes selected by the separate time-domain
baseline. Current blockers are material:
the upstream source/checkpoint has no declared reuse terms, and the current
Braindecode checkpoint is an unverified rehost with no proved immutable
lineage to the official mutable archive;
the PR pools all channels and frames with LayerNorm, whereas the paper's
downstream path uses one electrode, exactly ten center frames, and a plain
linear head;
the released sources disagree about the preprocessing used for the
published row, so choosing one silently would invent a protocol.
Please freeze the exact selected electrodes, recordings, preprocessing,
checkpoint bytes/hash/license, and aggregation contract, then reproduce this
row on one exact candidate SHA within a preregistered tolerance before merge.
For scope, reduce the library contribution to the faithful single-electrode
model/checkpoint boundary with focused tests; keep dataset acquisition and the
full reproduction experiment external and opt-in. No signal or checkpoint
bytes were downloaded because the permission/protocol gate is incomplete.
Paper-replication gate at exact head 7b1cff6: hold.
The exact released checkpoint now runs on the NEMAR Brain Treebank adaptation, and a bounded subject-03/run-01 frozen onset probe reached test ROC-AUC 0.7569192871 (BCE 0.5903241634). This is useful checkpoint/data-path evidence, not the paper's 0.82 +/- 0.07 result: it omits joint fine-tuning, top-10 electrode selection, seven-session aggregation, and the original HDF5 equivalence. The upstream source also publishes no license while the current Hub metadata claims BSD-3-Clause.
A current-master correction candidate closes the NaN STFT path and removes a false stft_clip equivalence claim, but it remains local because the paper/result and rights gates are still open. Do not merge on architecture tests alone.
Brings the branch up to 1.8.1. The only conflict was the changelog: the
BrainBERT entry moves from the released 1.7.0 section into the current one,
after the entries already there.
Four changes, all answering the same objection: the port ran a plausible
variant of BrainBERT rather than the one the released checkpoint was
evaluated with.
STFT recipe. Upstream ships two, and they are not equivalent.
`preprocessors/stft.py` z-scores first and trims 10 frames per side;
`notebooks/demo.ipynb` trims 5 first and then z-scores. It is the former
that `conf/preprocessor/stft_pretrained.yaml` reaches through
`preprocessors/spec_pretrained.py`, so it is the former that produced the
published downstream numbers. The port implemented the notebook one. The
default is now the checkpoint's recipe, the other stays reachable through
`stft_clip` and `stft_zscore_before_clip`, and the docstring no longer
claims an equivalence that does not hold: on filtered noise the two
correlate at 0.999 but differ by up to 0.48 z-unit per bin and by 10
frames of sequence length.
NaN handling. Upstream zeroes NaNs that survive the statistics and
replaces a fully constant window by ones; without either, one bad sample
poisons the whole attention window. Both are ported, branch-free so
`torch.export` still traces the forward.
Pooling. Upstream averages the 10 encoder frames centred on the window of
a single electrode (`outputs[:, middle-5:middle+5].mean`); averaging every
frame appears in that same file only as a commented-out alternative. The
port averaged all frames and all channels. It now pools the centre frames
and then the channels, which is the identity at n_chans=1, so a
single-channel model reproduces the upstream feature exactly. A window too
short for the centre window is refused at construction rather than pooled
differently in silence.
Probe. `models/linear_wav_baseline.py` is `nn.Linear(input_dim, 1)` and
nothing else, so the head's LayerNorm is gone.
Also: `freq_cutoff` is renamed `idx_freq_cutoff`, since it is a bin index
and not a frequency (40 bins reach about 200 Hz, not 40); the positional
encoding refuses a sequence longer than its table instead of truncating
it; `normalizing` rejects unknown values; and the summary row moves into
alphabetical order with a 1 s window at 2048 Hz, the window the checkpoint
expects.
Towards #1097
`EEGModuleMixin.__init_subclass__` defaults `license` to `bsd-3-clause`
when a model does not declare one, and that value is what a
`push_to_hub` writes onto the model card. The upstream BrainBERT
repository ships no LICENSE file, so BSD-3 is a claim nobody can back.
The class now declares `license="unknown"`, which is what the re-hosted
weights at `braindecode/brainbert-pretrained` already say, and a test
pins it so the default cannot creep back in.
This mechanism only reached master in #1134, after this branch was
opened, which is why the omission was invisible until now.
Towards #1097
…ence
The previous STFT test reproduced the demo-notebook recipe and asserted
the port matched it, which is exactly the claim that was wrong: it
validated the port against the recipe the released checkpoint does not
use. Both recipes are now written out from the upstream sources and
checked, and a third test pins that they are *not* interchangeable, so a
future simplification cannot quietly collapse them.
The tolerance drops from 1e-4 to 5e-6. The residual is the float32 Hann
window (scipy builds it in float64) and measures 7.1e-7; at 1e-4 the test
could not have distinguished the two recipes it now separates.
New coverage: the pooled feature equals upstream's
`outputs[:, middle-5:middle+5].mean` on the same encoder output and
differs from the all-frames mean; the head carries no LayerNorm; a flat
channel stays finite; `normalizing="none"`, `clip=0` and an unknown
`normalizing` all behave; the declared licence is `unknown`.
Parity gate: `sys.path` is restored whatever happens, and the gate skips
rather than passing vacuously if `models` resolves anywhere other than
BRAINBERT_SRC. Verified green against a fresh clone of
github.com/czlwang/BrainBERT, not only skipped.
Towards #1097
The three settings that silently change what the pretrained encoder sees
are the sampling rate (an STFT bin is a fraction of sfreq, so the same 40
bins reach 205 Hz at 2048 Hz and 100 Hz at 1000 Hz), the normalisation
order (upstream ships two recipes) and the window length (the protocol
pools 10 centred frames). The example works through all three on BCI
Competition IV dataset 4, then reproduces the published downstream
protocol: frozen encoder, bare linear probe.
It also probes an untrained copy of the same architecture as a control.
The untrained encoder reaches 0.580 ROC-AUC against 0.629 for the
pretrained one and 0.500 for chance, so most of the headroom above
chance comes from the architecture rather than from the pretraining --
reporting the pretrained number alone would have overclaimed.
`activation` was typed `type[nn.Module]` and instantiated with
`activation()`. That is braindecode's house spelling (BIOT, LaBraM,
CBraMod all use it) and it stays the default, but `activation()` raises
a TypeError on the other forms `nn.TransformerEncoderLayer` itself
accepts: a string such as "gelu", a ready-made module, or a bare
callable like `torch.nn.functional.gelu`.
Fold the four forms in `_as_transformer_activation` instead: a class is
instantiated, everything else is forwarded untouched, and a class that
is not an nn.Module subclass is refused with a clear message rather
than failing later inside the encoder layer.
pool_n_frames=0 passed the "enough frames to pool" guard, then forward
sliced middle:middle and averaged an empty tensor: the model returned NaN
features rather than failing. Negative values built too. Both are refused
at construction now.
reset_head assigned _n_outputs directly, which accepted 0 and left the
init kwargs and the Hub config reporting the previous count, so a model
re-serialized after reset_head advertised a head it no longer had. It
goes through EEGModuleMixin._set_n_outputs, the helper added in 1.8 for
exactly this, which validates and keeps both configs in step.
Widening the annotation to a union starting with str broke the shared
model config round-trip: pydantic matched the str branch first, so a
serialized "torch.nn.modules.activation.GELU" came back as that string
instead of the class, and test_config's json-serialization check failed
on BrainBERT alone.
The annotation goes back to type[nn.Module], which is what every other
braindecode constructor declares and what the config coder knows how to
decode. The runtime normalisation stays, so a string, a module instance
or a bare callable is still accepted -- it just is not the documented,
serializable form. A test pins the annotation, because test_config skips
when pydantic is absent and that is how this got through locally.
from_pretrained("braindecode/brainbert-pretrained") resolves to whatever
main points at, and a Hub branch is mutable. Upstream distributes the
weights from a Google Drive folder that is mutable too, so nothing in the
chain was immutable: a later push would silently change what the model
loads, and with it any number published against it, without a single line
of this repository changing.
BRAINBERT_WEIGHTS_REVISION pins the commit, BRAINBERT_WEIGHTS_SHA256
records the digests of both weight files so the bytes can be checked
without trusting the Hub to have served the right commit, and every
documented call site passes the revision.
Two tests guard it: one pins the constants so a change has to be
deliberate, the other walks every from_pretrained call in the module and
fails on one without a revision -- an unpinned copy-pasteable snippet is
how an unpinned load spreads.
Unavailable checkpoint is advertised as an official pretrained model
braindecode/models/brainbert.py:40
The PR description leaves re-hosting braindecode/brainbert-pretrained unchecked, but this new constant and the model docs advertise that repository as an available official checkpoint (and the tutorial unconditionally loads it). On the described PR state, those copy-paste paths cannot work. Either add and verify the artifact/provenance before exposing this API, or keep this contribution architecture-only and remove/guard the pretrained-load claims.
Replacement head ignores model device and dtype
braindecode/models/brainbert.py:332
When this method is called after the model has been moved to CUDA or a non-default dtype, the replacement head is created on CPU with float32 parameters. The next forward then fails on a device/dtype mismatch. Preserve the old final_layer.fc.weight device and dtype when constructing the new head, as the other reset_head implementations do.
❌ Patch coverage is 88.70968% with 14 lines in your changes missing coverage. Please review.
✅ Project coverage is 86.79%. Comparing base (4e332c4) to head (d8bfc0d). ⚠️ Report is 1 commits behind head on master.
- Add a key mapping from the official checkpoint's input_encoding.* to
input_embedding.*. stft_large_pretrained.pth (sha256 d5ba03df) now loads
with only final_layer.* missing and the rebuilt pe table unexpected, and
every loaded tensor is identical to the official file.
- reset_head keeps the old head's device and dtype.
- An odd hidden_dim raises a clear ValueError instead of a shape error in
the sinusoidal encoding.
- forward refuses inputs that give fewer spectrogram frames than the
pooled centre frames, instead of silently pooling fewer or returning NaN.
- Fix the reset_head config test, which read a non-existent _init_kwargs
attribute (_braindecode_init_kwargs).
hey @bruAristimunha@julien-gadonneix ! Paper-replication gate: two-stage protocol now runs end to end on the 7 sessions, 947 electrodes, 4 tasks, frozen BrainBERT probe read on electrodes selected by the linear baseline. 🙌
- final_layer is a bare nn.Linear; the unused pretraining head is dropped
(the re-hosted checkpoint loads unchanged, strict=False skips its keys)
- the private STFT and input-embedding helpers live below the model class;
braindecode/modules/brainbert_modules.py is removed
- sfreq is no longer required and the unused normalizing option is gone
- forward returns the logits under torch.jit.script, so BrainBERT leaves
the TorchScript skip list and scripts directly
- the bespoke test file is replaced by one STFT-vs-scipy check in
test_models.py; BrainBERT joins the shared return_features suite, which
now takes n_times from models_mandatory_parameters
- BrainBERT is listed in docs/api.rst; the tutorial is removed
Claude-Session: https://claude.ai/code/session_016335Fx1o9ZdJnwKbzsYKgX
- reshapes and pooling are einops Rearrange/Reduce layers built in
__init__, so forward is a sequence of named layers with shape comments
- the BRAINBERT_WEIGHTS_* constants are gone; the docstring shows the Hub
repo id as a literal, like the other models
- the class docstring links the paper's architecture figure and states how
to obtain one output per electrode as upstream does
- outputs, scripting and export are unchanged
Claude-Session: https://claude.ai/code/session_016335Fx1o9ZdJnwKbzsYKgX
API listing falsely claims unavailable pretrained weights
docs/api.rst:70
The API listing says BrainBERT is shipped "with pre-trained weights", but the PR checklist says the official weights are not re-hosted in this change. This advertises an unavailable artifact to documentation users; remove the weight claim or include the artifact before merging.
Release note incorrectly claims pretrained weights are included
docs/whats_new.rst:42
This release note claims the model comes with pretrained weights, contradicting the PR checklist's unchecked weight re-hosting item. Please describe the addition as an architecture implementation without re-hosted weights, or add the promised verified artifact.
…sitive
- forward raises on a channel-count mismatch (it silently re-batched) and
on non-floating inputs; the constructor names the minimum n_times and
rejects invalid STFT settings; the positional table overflow is a
ValueError; sfreq != 2048 Hz warns as in BIOT
- enable_nested_tensor=False removes the activation warning (no mask is
ever passed); `# nosec B105` marks the cls_token=None Bandit false
positive as in ZUNA
- docstring: which from_pretrained overrides work, the .pth load recipe,
NaN and dead-channel behaviour; versionadded 1.9
- tests: test_models.py carries only the BrainBERT block (STFT parity at
2048 and 6000 samples incl. a flat signal, mapping targets); the
return_features suite reads n_times via _get_signal_params; BrainBERT is
in the Hub integration registry
- normalizing stays an instance attribute for pipelines that apply their
own statistics; features are unchanged versus the previous head
Claude-Session: https://claude.ai/code/session_01G2YTx6U99AkGoHswQ5ZhHX
This public docstring claims that the released checkpoint is available at braindecode/brainbert-pretrained, but the PR description explicitly leaves re-hosting that artifact unchecked and this change adds no weights. As written, the documented from_pretrained example points users at an artifact this PR does not provide; either ship/verify the artifact or describe this contribution as architecture-only.
"""The authors' ``input_encoding.*`` keys map onto real port parameters."""
model = BrainBERT(n_chans=1, n_outputs=2, n_times=2048)
state = model.state_dict()
assert model.mapping and all(target in state for target in model.mapping.values())
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
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.
Summary
Towards #1097 — adds a braindecode-native port of BrainBERT (Wang et al.,
BrainBERT: Self-supervised representation learning for intracranial recordings,
ICLR 2023), a self-supervised foundation model for intracranial (sEEG/iEEG)
signals. Suggested on #1097.
Upstream code and released weights:
Architecture
(B, C, T) → STFT spectrogram → linear projection + sinusoidal positional encoding → Transformer encoder → pool over channels & frames → head.Following the convention established for Brant (#1100), the short-time
Fourier transform front-end is computed inside
forward(module_STFTSpectrogram) so the model keeps the standard(batch, n_chans, n_times)input signature, whereas the upstream reference consumes a pre-computed
spectrogram.
What's in this PR
braindecode/models/brainbert.py—BrainBERT(EEGModuleMixin, nn.Module)with modest ready-to-run defaults (~0.65M params) and the released large
config documented (
hidden_dim=768, ffn_dim=3072, n_heads=12, n_layers=6,~43M), plus
from_pretrained,reset_headandreturn_features.braindecode/modules/brainbert_modules.py— STFT front-end, input embedding,sinusoidal positional encoding, spectrogram-prediction head (weight parity)
and classification head.
models/__init__.py,models/util.py,summary.csv,docs/whats_new.rst.test/unit_tests/models/test_brainbert.py— contract tests + a paritygate.
Numerical fidelity
1.5e-6 (bit-exact in float32).
BRAINBERT_SRC), theinput encoding + Transformer are checked bit-exact (
atol=1e-5) againstthe upstream
MaskedTFModel; the gate is skipped otherwise so CI stays greenwithout the external dependency.
Checklist
models/__init__.py,models/util.py,summary.csv.test/unit_tests/models/test_brainbert.py).(
braindecode/brainbert-pretrained).examples/.