Repository navigation
Add SleepFM foundation and sleep-staging models - #1106
Conversation
There was a problem hiding this comment.
Pull request overview
Adds Braindecode support for the SleepFM multimodal PSG foundation model and its downstream sleep-staging variant, including public exports, model registration/metadata, documentation, and a comprehensive unit/integration test suite (with optional network tests for official checkpoints).
Changes:
- Introduce
SleepFM(trial-level encoder + head) andSleepFMStager(patch-wise sleep staging) implementations with checkpoint-loading helpers and channel masking. - Register the new models for discovery/metadata and expose them via the public
braindecode.modelsAPI. - Add tests (including optional network download/load of official checkpoints) and update docs/release notes/licensing notices.
Reviewed changes
Copilot reviewed 12 out of 12 changed files in this pull request and generated 2 comments.
Show a summary per file
| File | Description |
|---|---|
| braindecode/models/sleepfm.py | Adds the SleepFM and SleepFMStager model implementations, masking/validation, and checkpoint key-mapping/load helpers. |
| braindecode/models/init.py | Exposes the new models as public imports and in __all__. |
| braindecode/models/util.py | Registers mandatory signal parameters and marks SleepFMStager as non-trial-classification output. |
| braindecode/models/summary.csv | Adds model summary metadata entries (domains, example configs, parameter counts, tags). |
| test/unit_tests/models/test_sleepfm.py | Adds focused unit tests for tokenizer/pooling behavior, masking invariance, loaders, and reduced-config compilation. |
| test/unit_tests/models/test_return_features.py | Adds SleepFM coverage to the shared return_features / reset_head behavior tests. |
| test/unit_tests/models/test_integration.py | Skips heavyweight compile/script coverage for SleepFM models in the generic integration suite (handled elsewhere). |
| test/unit_tests/models/test_foundation_models.py | Adds optional network tests that download and load official upstream SleepFM checkpoints. |
| docs/api.rst | Documents the new public API entries for SleepFM models. |
| docs/models/models_categorization.rst | Categorizes SleepFM/SleepFMStager as foundation-model examples. |
| docs/whats_new.rst | Adds release notes describing new models, baselines context, and licensing. |
| NOTICE.txt | Records CC BY-NC 4.0 attribution and noncommercial licensing constraints for the new derivative file. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
|
Thanks for this contribution @Fashad-Ahmed 👍 it's a thorough, well-documented port. A few things needed before we can move toward merge:
https://huggingface.co/braindecode/SleepFM/resolve/main/model_base/best.pt
Happy to help on any of these, thanks again :) ! |
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.
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.
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.
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.
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.
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.
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.
…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.
…ionPooling 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.
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.
87cf158 to
e15a34b
Compare
|
@bruAristimunha Is there any thing left other than rebasing to the Codacy Static Code Analysis workflow, so that it can get merge |
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.
`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.
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.
|
Integration gate (braindecode maintainers) SHHS test set (2,000 nights, released split): port 0.7925 macro-F1, authors' released code 0.7925, paper 0.78. Note: the replication check covers the SHHS sleep-staging cell only. Audited head: SHHS test set (2,000 nights, released split): port 0.7925 macro-F1, authors' released code 0.7925, paper 0.78. |
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1106 +/- ##
==========================================
+ Coverage 88.51% 88.61% +0.10%
==========================================
Files 158 159 +1
Lines 19306 19531 +225
==========================================
+ Hits 17089 17308 +219
- Misses 2217 2223 +6 🚀 New features to boost your workflow:
|
Resolve conflicts keeping both sides: - docs/whats_new.rst: keep the SleepFM/SleepFMStager entry and master's DIVER1, ZUNA, MSCFormer, BaRISTA and Brant entries. - test/unit_tests/models/test_foundation_models.py: union of the model imports (isort order) and both test blocks (SleepFM tests, then the DIVER1/STEEGFormer tests from master). - test/unit_tests/models/test_return_features.py: keep both the SleepFM and the MIRepNet parametrizations.
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.
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).
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.
…kpoints 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.
Resolve docs/whats_new.rst: keep master's braindecode#1159/braindecode#1155 entries and the SleepFM (braindecode#1106) entry. test_foundation_models.py auto-merged (master's LaBraM tests plus the SleepFM tests); _DIRECT_TORCHSCRIPT_MODELS stays 32 on base, PR and master.
Conflict resolution: - test/unit_tests/models/test_foundation_models.py: kept both the SleepFM/SleepFMStager imports (braindecode#1106) and master's PopulationTransformer import (braindecode#1105), alphabetized. - docs/whats_new.rst auto-merged cleanly, keeping both the SleepFM and PopulationTransformer entries.
# Conflicts: # test/unit_tests/models/test_foundation_models.py
…okenizer patching, paper details in docstrings
# Conflicts: # test/unit_tests/models/test_foundation_models.py
… outputs unchanged
Addressed: all review threads resolved; Bruno approved merging on 2026-10-07.
…nto feat/channel-layer
…M) into fix/model-dtypes
Summary
This PR adds support for SleepFM (https://github.com/zou-group/sleepfm-clinical), a multimodal polysomnography foundation model introduced in:
The contribution introduces two public Braindecode models:
Both models follow Braindecode’s EEGModuleMixin conventions and support the authors’ official base and sleep-staging checkpoints.
Links to #925.
Motivation
SleepFM was trained on multimodal PSG recordings and can process heterogeneous physiological channels, including EEG, EOG, respiratory signals, ECG, and EMG.
The upstream maintainers indicated that a contribution would be welcome and requested clarification about the baselines used in the paper. This PR therefore includes
both the model implementation and an explicit description of those baselines.
Model architecture
SleepFM
SleepFM processes raw signals with shape:
(batch, channels, time)
The reference configuration expects signals sampled at 128 Hz and uses non-overlapping 5-second patches, corresponding to 640 samples per patch.
The model contains:
The default parameters reproduce the released base configuration:
The model exposes:
The returned shapes are:
SleepFMStager
SleepFMStager combines the released SleepFM tokenizer with the authors’ sleep-staging architecture:
Its output shape is:
(batch, n_outputs, n_patches)
With the official checkpoint, n_outputs=5 represents the five sleep-stage classes used upstream.
No softmax is included, allowing the logits to be consumed directly by the selected Braindecode or PyTorch loss.
Variable-channel masking
Both models accept an optional channel padding mask:
The mask has shape:
(batch, channels)
A value of True marks a missing or padded channel.
Masked channels are excluded from attention pooling and the subsequent masked mean. Each sample must contain at least one valid channel.
This allows examples with different PSG montages or modality availability to be batched together without treating padded channel values as real observations.
Input validation
The implementation validates:
Trailing samples that do not form a complete patch are discarded, matching the upstream patching behavior.
Official checkpoint compatibility
The following loading methods are provided:
and, for sleep staging:
The loaders support:
Mappings are intentionally minimal. Unexpected or missing model keys are reported instead of silently discarding incompatible parameters.
Opt-in network tests download and load the official checkpoints:
Both official-checkpoint tests pass.
Baselines used in the paper
The paper compares SleepFM with two supervised baselines.
Demographics baseline
The demographics-only baseline is a shallow MLP:
This baseline is not added as a separate Braindecode model because it does not process electrophysiological signals.
End-to-End PSG baseline
The End-to-End PSG baseline is trained jointly from random initialization and contains:
This baseline does not use SleepFM self-supervised pretraining.
It is not equivalent to constructing SleepFM without pretrained weights: its downstream temporal architecture and demographic fusion are different. It is not included
as a separate public class in this PR because it requires nonelectrophysiological covariates and the paper’s disease-survival training pipeline.
The PR focuses on the reusable electrophysiological components and the released sleep-staging path.
Braindecode integration
The contribution includes:
SleepFMStager is registered as a non-trial-classification model because it returns token-wise logits rather than a single prediction per recording.
Usage example
Sleep-staging example:
Official checkpoint loading:
Testing
The implementation was tested with:
Result:
68 passed, 3 skipped
Shared SleepFM integration coverage:
Result:
18 passed, 6 skipped
Official checkpoint tests:
Result:
2 passed
The standard repository quality gate was also run:
All configured hooks passed, including:
git diff --check also passes.
Environmental verification notes
The complete model test directory was attempted, but the local Python 3.13/macOS environment encountered a native segmentation fault inside scikit-learn’s nearest-neighbor implementation. This occurred outside the SleepFM tests.
A complete Sphinx gallery build was also attempted. SleepFM autosummary pages were generated, but the full repository build was prevented by external intersphinx
network access, gallery dataset downloads, and sandbox multiprocessing restrictions. The configured RST, docstrfmt, and Sphinx-lint checks pass.
Licensing
The upstream SleepFM implementation and released checkpoints use the Creative Commons Attribution-NonCommercial 4.0 International license.
Consequently:
Upstream repository: https://github.com/zou-group/sleepfm-clinical
License: https://creativecommons.org/licenses/by-nc/4.0/
Files changed
braindecode/models/sleepfm.py
braindecode/models/init.py
braindecode/models/util.py
braindecode/models/summary.csv
test/unit_tests/models/test_sleepfm.py
test/unit_tests/models/test_foundation_models.py
test/unit_tests/models/test_return_features.py
test/unit_tests/models/test_integration.py
docs/api.rst
docs/models/models_categorization.rst
docs/whats_new.rst
NOTICE.txt
Checklist