Repository navigation
Fix reset_head leaving the saved config on the previous head - #1181
Merged
bruAristimunha merged 7 commits intoSep 29, 2026
Merged
Conversation
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 Report❌ Patch coverage is 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:
|
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. |
Collaborator
|
But this way should be better! Many thanks for the PR |
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.
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]>
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
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
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.
Fixes #1179.
reset_head(n)rebuilt the head but leftn_outputsat its old value in the saved config, sosave_pretrainedfollowed byfrom_pretrainedfailed 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 inheritreset_headfrom their backbone.Two more problems on the same path:
final_layer=False, and CBraMod or EEGDINO built withreturn_encoder_output=True, gain a head inreset_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.reset_headimplementations. With dropout in the head (EEGPT, EEGDINO), two forward passes on the same input differed whilemodel.trainingwasFalse(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 callsreset_headafter huggingface_hub has put the model in eval.Changes
self._n_outputsdirectly now call_set_n_outputs, the helper added for this in Refactor shared EMG model support components #1145._update_init_kwargsrecords other constructor arguments thatreset_headchanges, and_set_n_outputsnow uses it. BENDR recordsfinal_layer=True; CBraMod and EEGDINO recordreturn_encoder_output=False.reset_headputs the new head inself.trainingmode. 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.reset_headdocstring now states both requirements for new models.reset_head(0)now raisesValueErrorin these models, as it already did in the six models using_set_n_outputs. Labram's timm-stylereset_classifier(0)is unchanged.Tests
In
test/unit_tests/models/test_return_features.py. The first two are parametrized over every registered model that overridesreset_head(24 cases, the interpolated wrappers included), so new models are picked up automatically:test_reset_head_updates_config_and_keeps_eval_mode: the newn_outputsis inget_config(), and no module is left in train mode on an eval model.test_reset_head_model_reloads_after_saving: afterreset_head,save_pretrainedthenfrom_pretrainedgives 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 == 5for the stale config (19, plusassert 100 == 103for MetaNeuromotorHand),RuntimeError: Error(s) in loading state_dicton reload (18),assert not Truefor the train-mode head (6),The keys of the mappings do not matchfor 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.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-networklike the other REVE tests). Master 15561d1 with the same command: 3482 passed, 246 skipped, 0 failed.