Skip to content

Fix reset_head leaving the saved config on the previous head - #1181

Merged
bruAristimunha merged 7 commits into
braindecode:masterfrom
raghav-rathi:fix/1179-reset-head-config
Sep 29, 2026
Merged

bruAristimunha merged 7 commits into
braindecode:masterfrom
raghav-rathi:fix/1179-reset-head-config

Conversation

@raghav-rathi

Copy link
Copy Markdown
Contributor

Fixes #1179.

reset_head(n) rebuilt the head but left n_outputs at its old value in the saved config, so save_pretrained followed by from_pretrained failed with a state-dict size mismatch. This hit 18 classes: BENDR, BIOT, CBraMod, EEGDINO, EEGPT, Labram, MetaNeuromotorHand, MVPFormer, REVE, STEEGFormer, ZUNA, the three SignalJEPA classifiers, and the Interpolated BENDR, BIOT, EEGPT and LaBraM wrappers, which inherit reset_head from their backbone.

Two more problems on the same path:

  • BENDR built with final_layer=False, and CBraMod or EEGDINO built with return_encoder_output=True, gain a head in reset_head, but the saved config still described a feature extractor. Reloading raised no error and dropped the trained head: EEGDINO came back returning (2, 200) features instead of (2, 4) class scores.
  • The new head was built in train mode even when the model was in eval, in all 20 reset_head implementations. With dropout in the head (EEGPT, EEGDINO), two forward passes on the same input differed while model.training was False (by up to 0.34 for EEGPT and 0.04 for EEGDINO). from_pretrained(..., n_outputs=k) returned a model in that state, because it calls reset_head after huggingface_hub has put the model in eval.

Changes

  • The 14 implementations that wrote self._n_outputs directly now call _set_n_outputs, the helper added for this in Refactor shared EMG model support components #1145.
  • New private helper _update_init_kwargs records other constructor arguments that reset_head changes, and _set_n_outputs now uses it. BENDR records final_layer=True; CBraMod and EEGDINO record return_encoder_output=False.
  • Every reset_head puts the new head in self.training mode. DANCE also rebuilds its DETR class head and VEMG2Pose its decoder; both get the same treatment. Only the new modules change mode, so a submodule the user froze keeps its mode.
  • The base reset_head docstring now states both requirements for new models.
  • Behaviour change: reset_head(0) now raises ValueError in these models, as it already did in the six models using _set_n_outputs. Labram's timm-style reset_classifier(0) is unchanged.

Tests

In test/unit_tests/models/test_return_features.py. The first two are parametrized over every registered model that overrides reset_head (24 cases, the interpolated wrappers included), so new models are picked up automatically:

  • test_reset_head_updates_config_and_keeps_eval_mode: the new n_outputs is in get_config(), and no module is left in train mode on an eval model.
  • test_reset_head_model_reloads_after_saving: after reset_head, save_pretrained then from_pretrained gives back identical weights.
  • test_from_pretrained_new_n_outputs_is_deterministic_in_eval (EEGDINO, EEGPT, the heads with dropout): from_pretrained(..., n_outputs=k) returns an eval model whose output is the same on two calls.
  • test_reset_head_on_feature_extractor_reloads_as_classifier: the three feature-extractor cases reload with their head and return class scores.

Models are built with their full registry signal parameters so that CBraMod's head is not lazy; the Hub loader cannot load into uninitialized parameters.

On master 15561d1 the new tests give 47 failed and 6 passed: assert 2 == 5 for the stale config (19, plus assert 100 == 103 for MetaNeuromotorHand), RuntimeError: Error(s) in loading state_dict on reload (18), assert not True for the train-mode head (6), The keys of the mappings do not match for the dropped feature-extractor head (3). The 6 passes are the reload test for the six models that already called _set_n_outputs (BrainBERT, DANCE, EMG2QwertyNet, NeuroPose, SensingDynamics, VEMG2Pose), whose save and reload already worked.

Verification

  • pytest test/unit_tests/models/test_return_features.py -n 4 --run-network: 97 passed.
  • Full suite as in tests.yml, pytest -vv --durations=0 -n 16 --dist worksteal test/ (CPU only, Python 3.12, torch 2.14, Linux): 3533 passed, 248 skipped, 0 failed (the 2 extra skips are the new REVE cases, which need --run-network like the other REVE tests). Master 15561d1 with the same command: 3482 passed, 246 skipped, 0 failed.
  • pre-commit on the changed files: clean.
  • Not run: macOS, Windows, Python 3.13, the documentation build.

raghav-rathi and others added 6 commits September 26, 2026 07:15
reset_head rebuilt the head but kept the previous n_outputs in the saved
config in 18 models, so a model saved after changing its head could not
be loaded back. BENDR, CBraMod and EEGDINO built as feature extractors
also kept that setting in the config and reloaded without their trained
head. The new head was built in train mode, so a model in eval mode,
including the one from_pretrained(..., n_outputs=...) returns, kept
dropout active in the head.

The reset_head implementations now call _set_n_outputs, record other
changed constructor arguments with the new _update_init_kwargs helper,
and build the new head in the model's train/eval mode.

Fixes braindecode#1179
@codecov

codecov Bot commented Sep 29, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 91.66667% with 2 lines in your changes missing coverage. Please review.
✅ Project coverage is 87.07%. Comparing base (c7570c8) to head (e88864a).
⚠️ Report is 1 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1181      +/-   ##
==========================================
+ Coverage   86.99%   87.07%   +0.07%     
==========================================
  Files         144      144              
  Lines       16580    16602      +22     
==========================================
+ Hits        14424    14456      +32     
+ Misses       2156     2146      -10     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@bruAristimunha

bruAristimunha commented Sep 29, 2026 •

Copy link
Copy Markdown
Collaborator

the trick part is that the was that was propose, it would introduce some break in neuroai. I think the whole delegation where to put in eval and train, should be more responsability of the person or framework that use the model.

@bruAristimunha

Copy link
Copy Markdown
Collaborator

But this way should be better! Many thanks for the PR

@bruAristimunha
bruAristimunha merged commit 2c00ed8 into braindecode:master Sep 29, 2026
14 checks passed
bruAristimunha added a commit to Fashad-Ahmed/braindecode that referenced this pull request Sep 30, 2026
Master's braindecode#1181 requires reset_head implementations to record the new
n_outputs with _set_n_outputs, so get_config() and save_pretrained()
describe the head the model has. SleepFM and SleepFMStager still set
_n_outputs directly, which the merged test_reset_head_updates_config and
test_reset_head_model_reloads_after_saving cases catch.
@raghav-rathi

Copy link
Copy Markdown
Contributor Author

Thanks for merging and for fixing it up! Makes sense to leave train/eval to the caller, I hadn't thought about the neuroai side.

bruAristimunha added a commit that referenced this pull request Oct 7, 2026
* feat(models): add SleepFM model for multimodal polysomnography

Introduce the SleepFM model, which learns channel-agnostic representations from polysomnography data, including EEG, EOG, respiratory signals, ECG, and EMG. The model utilizes a convolutional tokenizer to process non-overlapping patches of input signals and incorporates attention mechanisms for pooling variable channel sets. Update summary.csv to include SleepFM and its parameters.

* feat(models): include SleepFM and SleepFMStager in model exports

Add SleepFM and SleepFMStager to the braindecode.models module, updating the __all__ export to ensure they are accessible for users. This change enhances the model offerings for multimodal polysomnography analysis.

* feat(models): add SleepFM and SleepFMStager parameters to model exports

Enhance the braindecode.models module by including mandatory parameters for the SleepFM and SleepFMStager models. This update ensures that both models are fully integrated into the existing framework, facilitating their use in multimodal polysomnography analysis.

* docs: update API and what's new for SleepFM and SleepFMStager

Enhance documentation by adding details about the new SleepFM and SleepFMStager models, including their functionalities and compatibility with existing checkpoints. Update the models categorization to reflect their inclusion and provide a comprehensive overview in the what's new section.

* feat(tests): add unit tests for SleepFM and SleepFMStager models

Introduce comprehensive unit tests for the SleepFM and SleepFMStager models, validating their functionality, including tokenizer behavior, attention pooling, and model output shapes. Ensure that the models correctly handle various input scenarios and maintain expected output dimensions, enhancing the test coverage for these new additions to the braindecode framework.

* docs: update NOTICE.txt to include SleepFM licensing information

Add licensing details for the SleepFM model, indicating its adaptation from `zou-group/sleepfm-clinical` under CC BY-NC 4.0. This update ensures proper attribution and compliance with licensing terms for the newly integrated model.

* refactor(models): enhance mask validation in _SleepFMAttentionPooling

Update the _validate_mask method to accept a mask_name parameter for improved error messaging. This change allows for more descriptive exceptions when validating key padding and channel masks, enhancing code clarity and maintainability.

* feat(tests): add additional unit tests for SleepFM attention pooling and channel mask validation

Introduce a new test for compiled attention pooling to ensure it correctly rejects all masked samples. Additionally, add parameterized tests to validate error handling for various channel mask shapes, enhancing the robustness of the SleepFM model's input validation.

* fix(models): enhance handling of all masked samples in _SleepFMAttentionPooling

Update the _SleepFMAttentionPooling class to safely handle cases where all input channels are masked. The changes ensure that outputs are correctly masked to zero when all channels are invalid, improving the robustness of the attention pooling mechanism. Additionally, modify the corresponding unit test to verify this behavior, ensuring that the model behaves as expected under these conditions.

* fix(tests): update SleepFM checkpoint URLs to Hugging Face

Modify unit tests for SleepFM to replace outdated GitHub URLs with updated links from Hugging Face. This change ensures that the tests reference the correct model checkpoints, improving the reliability of the test suite.

* docs(models): clarify SleepFMStager documentation with reference to downstream architecture

* FIX preserve SleepFM singleton parity

* Simplify SleepFM and load its weights through from_pretrained

Answers the review on braindecode/models/sleepfm.py:

- Public models first, private helper classes after, per the convention of
  the recently merged ports (mvpformer.py, steegformer.py).
- The two mask helpers become module-level functions: the stager was reaching
  into `SleepFM._prepare_channel_mask` and the pooling into
  `_SleepFMAttentionPooling._validate_mask`, neither of which needs a class.
  As a side effect the validation errors now quote `channel_mask`, the public
  argument, instead of the internal `key_padding_mask`.
- The tokenizer owns the patch arithmetic. Both constructors asked
  `n_times // patch_size` themselves and re-implemented the "at least one
  complete patch" check; they now ask `patch_embedding.n_patches()`.
- A wrong `sfreq` warns instead of raising. 128 Hz is what the released
  weights were trained at, not a constraint of the architecture, so the model
  stays usable on other data at the user's risk.
- einops for every channel/patch fold, which is where the shapes were hardest
  to follow, plus the math behind `encode` and the three stages of the
  staging head in the docstrings.
- `tokenize()` was a one-line delegation to `patch_embedding` and is gone.
- The bespoke checkpoint API (`load_pretrained_backbone`,
  `load_pretrained_staging_head` and the four key-remapping helpers) is
  replaced by `from_pretrained`, as for ZUNA and LUNA. The keys are rewritten
  once when mirroring the weights, so the library carries no remapping code,
  and the stager mirror merges the released tokenizer with the released
  staging head -- one call now returns a fully pretrained stager instead of
  two.
- `# nosec B105`: Bandit reads any dict key containing "token" with a
  constant value as a hardcoded credential, and `{"features": ...,
  "cls_token": None}` is braindecode's feature-return contract, shared with
  eight other models.

The refactor is numerically neutral: same state_dict and same input give
bit-exact outputs before and after, logits, features and masked forward alike.

* Fold the SleepFM tests into the shared model suites

`test_sleepfm.py` is deleted: shapes, the feature contract, `reset_head`,
compilation, export and registration are already covered for every model by
the parametrized suites, and duplicating them per model is what the shared
files exist to avoid.

What survives moves to test_foundation_models.py, next to the other foundation
models, and is only what is specific to SleepFM: the channel mask. A PSG
montage varies between subjects, so the tests worth keeping are that a masked
channel cannot reach any output whatever it contains, that a fully masked
sample is refused, and that the errors name `channel_mask`.

SleepFM and SleepFMStager also leave the `test_model_compiled` exclusion list.
The stated reason -- official-size configurations being too expensive to
compile -- does not hold: both compile and run together in 23 s on CPU at the
summary.csv size. They stay excluded from `test_model_torch_script`, with the
reason corrected to the measured one: TorchScript cannot type the
`tuple(x.shape)` shape-guard messages, and rejects the polymorphic
Dict/Tensor return exactly as it does for EEGDINO and MVPFormer.

The two network tests now go through `from_pretrained`, and additionally check
that the stager's tokenizer, recurrent head and output layer all come from the
checkpoint rather than from the random initialisation.

* Trim the SleepFM changelog entry

The entry described the loading API that no longer exists and the 128 Hz
check that is now a warning. It also spent five lines on the paper's
disease-prediction baselines; one sentence saying why they are out of scope
(they need age, sex, BMI and race/ethnicity, which are not
electrophysiological inputs) says the same thing.

* FIX: SleepFM reset_head keeps the saved config in sync

Master's #1181 requires reset_head implementations to record the new
n_outputs with _set_n_outputs, so get_config() and save_pretrained()
describe the head the model has. SleepFM and SleepFMStager still set
_n_outputs directly, which the merged test_reset_head_updates_config and
test_reset_head_model_reloads_after_saving cases catch.

* FIX: SleepFMStager reproduces the released staging pipeline

SleepFMStager.forward fed raw tokenizer outputs to the staging head,
whereas the released staging head was trained on SleepFM encoder
embeddings computed per modality on 5-minute chunks; its logits differed
from the official pipeline by up to 3.18 on released weights. It also had
no temporal mask, and in training mode masked channels entered the
tokenizer's BatchNorm statistics.

- The stager now holds the SleepFM encoder (same parameter names, minus
  the trial-level temporal pooling). Channels are grouped with the new
  channel_modalities argument; each modality is encoded chunk by chunk
  (encoder_chunk_patches=60), positions restarting per chunk, and the
  staging head pools at least four modality slots as released.
- forward accepts temporal_mask (batch, n_patches): padded patches reach
  the head as zero embeddings masked out of its Transformer, as the
  release collates them, and are masked in the encoder too.
- The tokenizer takes a padding mask. In training it only computes the
  valid patches, so masked channels and padded patches affect neither the
  BatchNorm batch statistics nor its running averages (the official
  pretraining normalises its zero padding with the signal; both agree
  when nothing is masked and always in eval mode). This also fixes
  SleepFM.encode.
- from_pretrained completes the tokenizer + head braindecode/SleepFMStager
  mirror with the encoder of the braindecode/SleepFM mirror, keys
  unchanged; checkpoints saved from a SleepFMStager load as they are.
- Shared encoder steps live in _SleepFMSequenceMixin and
  _temporal_transformer, used by SleepFM, the stager and the head.

With released weights, official two-stage pipeline vs public forward:
max |logit diff| 0.0 (float32), input gradients 0.0, BatchNorm running
stats 0.0, in eval and train mode, with masked channels and temporal
padding (train mode compared with a leak-free composition of the
official modules).

* FIX: SleepFM keeps the unmasked training path and tests the pipeline

Review follow-up to the SleepFMStager pipeline fix:
- Without channel_mask or temporal_mask the tokenizer takes its plain
  path again; the masked gather added host syncs and graph breaks to every
  training step (dynamo graph breaks back to 0 for SleepFM, 2 for the
  stager, as before the fix).
- In eval mode masked patches are zeroed before the tokenizer, so a
  non-finite padded channel no longer turns parameter gradients to NaN.
- from_pretrained warns when the released stager is loaded without
  channel_modalities, and accepts encoder_revision for the encoder repo.
- New tests compare the stager with the two-stage pipeline built from
  SleepFM.encode per modality and chunk plus the staging head, with the
  temporal mask aligned to a chunk and ending inside one; they fail when
  either the head's or the encoder's temporal mask is removed.
- Docstrings and changelog state precisely that only masked channels
  deviate from the official BatchNorm behaviour.

* FIX: SleepFMStager warns without channel_modalities for complete checkpoints

from_pretrained only warned about a missing channel_modalities when the
encoder weights were completed from braindecode/SleepFM, so a
self-contained stager checkpoint loaded silently as a single modality.
Warn whenever the loaded model has no channel_modalities, once per call,
and describe the head-only mirror layout as that of older revisions.

* DOC: keep the SleepFM NOTICE entry inside the CC BY-NC file list

* FIX register SleepFM and SleepFMStager under cc-by-nc-4.0

* STY ruff-format sleepfm.py class declarations

* MAINT SleepFM: model classes first, shared steps as functions, PatchTokenizer patching, paper details in docstrings

* MAINT SleepFM: cut redundant checks, mask helpers, loaders and tests; outputs unchanged

* SleepFM: docstring figure, pretraining note and weights box; fold one-use helpers

* SleepFM: note the SHHS-only replication in the docstrings; warn at the caller's line

---------

Co-authored-by: Bruno Aristimunha <[email protected]>
Co-authored-by: Bru <[email protected]>
Co-authored-by: Adam Mounir <[email protected]>
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.

reset_head() leaves the saved config at the old n_outputs, so a fine-tuned model can't be loaded back

2 participants