Skip to content

Add SleepFM foundation and sleep-staging models - #1106

Merged
bruAristimunha merged 35 commits into
braindecode:masterfrom
Fashad-Ahmed:sleepfm-model
Oct 7, 2026
Merged

bruAristimunha merged 35 commits into
braindecode:masterfrom
Fashad-Ahmed:sleepfm-model

Conversation

@Fashad-Ahmed

Copy link
Copy Markdown
Contributor

Summary

This PR adds support for SleepFM (https://github.com/zou-group/sleepfm-clinical), a multimodal polysomnography foundation model introduced in:

Thapa et al., “A multimodal sleep foundation model for disease prediction,” Nature Medicine (2026).
DOI: https://doi.org/10.1038/s41591-025-04133-4

The contribution introduces two public Braindecode models:

  • SleepFM: the pretrained channel-agnostic PSG encoder with a trial-level prediction head.
  • SleepFMStager: the released downstream architecture for patch-wise sleep-stage prediction.

Both models follow Braindecode’s EEGModuleMixin conventions and support the authors’ official base and sleep-staging checkpoints.

image

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:

  1. A shared six-block convolutional tokenizer for every channel and patch.
  2. Transformer-based attention pooling across available channels.
  3. Fixed sinusoidal positional encoding.
  4. A temporal Transformer encoder.
  5. Temporal attention pooling.
  6. A configurable trial-level prediction head.

The default parameters reproduce the released base configuration:


  SleepFM(
      patch_size=640,
      embed_dim=128,
      num_heads=8,
      num_layers=6,
      pooling_heads=8,
      drop_prob=0.3,
      max_seq_length=128,
  )

The model exposes:


  logits = model(x)
  result = model(x, return_features=True)
  tokens = model.tokenize(x)
  pooled, contextual_tokens = model.encode(x)

The returned shapes are:


  logits:             (batch, n_outputs)
  features:           (batch, embed_dim)
  tokens:             (batch, channels, patches, embed_dim)
  contextual_tokens:  (batch, patches, embed_dim)

SleepFMStager

SleepFMStager combines the released SleepFM tokenizer with the authors’ sleep-staging architecture:

  1. Shared raw-signal tokenizer.
  2. Attention pooling across channels.
  3. Temporal Transformer encoder.
  4. Bidirectional LSTM.
  5. Patch-wise output projection.

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:


  output = model(x, channel_mask=channel_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:

  • Input rank and expected (batch, channels, time) layout.
  • Sampling frequency of exactly 128 Hz.
  • Availability of at least one complete patch.
  • Compatibility of patch_size with the six convolutional blocks.
  • Divisibility of the embedding dimension by configured attention heads.
  • Positional-encoding capacity.
  • Channel-mask shape, type, and values.
  • Samples with every channel masked.
  • Compatibility of checkpoint keys and tensor shapes.

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:


  model.load_pretrained_backbone(base_checkpoint)

and, for sleep staging:

  stager.load_pretrained_backbone(base_checkpoint)
  stager.load_pretrained_staging_head(staging_checkpoint)

The loaders support:

  • Plain state dictionaries.
  • Dictionaries nested under state_dict.
  • Distributed module. prefixes.
  • The upstream positional-encoding key format.
  • Mapping the upstream staging fc layer to Braindecode’s final_layer.

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:

  
  sleepfm/checkpoints/model_base/best.pt
  sleepfm/checkpoints/model_sleep_staging/best.pth

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:


  age, sex, BMI, race/ethnicity
                │
             Linear(4, 32)
                │
               ReLU
                │
              Dropout
                │
         disease output layer

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:

  • The raw-signal convolutional tokenizer.
  • Attention pooling across PSG channels.
  • A bidirectional LSTM for temporal processing.
  • An embedding of age and sex.
  • A disease-prediction head.

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:

  • Public imports from braindecode.models.
  • Entries in all.
  • Model discovery and mandatory-parameter registration.
  • Model summary metadata and parameter counts.
  • Unified return_features behavior.
  • reset_head(n_outputs) support.
  • torch.export coverage.
  • Reduced-configuration torch.compile/Dynamo graph-capture coverage.
  • API autosummary entries.
  • Foundation-model categorization.
  • Release notes.
  • License and attribution notice.

SleepFMStager is registered as a non-trial-classification model because it returns token-wise logits rather than a single prediction per recording.

Usage example


  import torch

  from braindecode.models import SleepFM

  model = SleepFM(
      n_chans=4,
      n_outputs=2,
      n_times=3840,
      sfreq=128,
  )

  x = torch.randn(8, 4, 3840)

  # True marks a missing or padded channel.
  channel_mask = torch.tensor(
      [
          [False, False, False, False],
          [False, False, True, True],
          [False, False, False, True],
          [False, False, False, False],
          [False, True, True, True],
          [False, False, False, True],
          [False, False, True, True],
          [False, False, False, False],
      ]
  )

  logits = model(x, channel_mask=channel_mask)
  features = model(
      x,
      channel_mask=channel_mask,
      return_features=True,
  )["features"]

  print(logits.shape)    # (8, 2)
  print(features.shape)  # (8, 128)

Sleep-staging example:


  import torch

  from braindecode.models import SleepFMStager

  model = SleepFMStager(
      n_chans=4,
      n_outputs=5,
      n_times=3840,
      sfreq=128,
  )

  x = torch.randn(8, 4, 3840)
  logits = model(x)

  # 3840 samples / 640 samples per patch = 6 patches
  print(logits.shape)  # (8, 5, 6)

Official checkpoint loading:

  import torch

  from braindecode.models import SleepFMStager

  base_checkpoint = torch.load("best.pt", map_location="cpu")
  staging_checkpoint = torch.load("best.pth", map_location="cpu")

  model = SleepFMStager(
      n_chans=4,
      n_outputs=5,
      n_times=3840,
      sfreq=128,
  )

  model.load_pretrained_backbone(base_checkpoint)
  model.load_pretrained_staging_head(staging_checkpoint)
  model.eval()

Testing

The implementation was tested with:


  pytest -q \
    test/unit_tests/models/test_sleepfm.py \
    test/unit_tests/models/test_return_features.py


Result:

68 passed, 3 skipped

Shared SleepFM integration coverage:

  pytest -q test/unit_tests/models/test_integration.py -k "SleepFM"

Result:

18 passed, 6 skipped

Official checkpoint tests:

  pytest -q \
    test/unit_tests/models/test_foundation_models.py \
    -k "sleepfm_official" \
    --run-network

Result:

2 passed

The standard repository quality gate was also run:

  pre-commit run -a

All configured hooks passed, including:

  • Ruff formatting and linting.
  • MyPy.
  • isort.
  • codespell.
  • docstrfmt.
  • Sphinx RST lint.
  • File and whitespace checks.

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:

  • braindecode/models/sleepfm.py retains upstream attribution.
  • The file is explicitly identified as CC BY-NC 4.0.
  • NOTICE.txt records the source repository, copyright attribution, and inherited noncommercial restriction.
  • The implementation and official weights must not be treated as BSD-3-licensed components.

Upstream repository: https://github.com/zou-group/sleepfm-clinical
License: https://creativecommons.org/licenses/by-nc/4.0/

Files changed

  • braindecode/models/sleepfm.py

    • SleepFM models, internal components, validation, and checkpoint mapping.
  • braindecode/models/init.py

    • Public exports.
  • braindecode/models/util.py

    • Model discovery and mandatory signal parameters.
  • braindecode/models/summary.csv

    • Documentation metadata and parameter counts.
  • test/unit_tests/models/test_sleepfm.py

    • Focused architecture and behavior tests.
  • test/unit_tests/models/test_foundation_models.py

    • Official checkpoint tests.
  • test/unit_tests/models/test_return_features.py

    • Unified feature and head-reset coverage.
  • test/unit_tests/models/test_integration.py

    • Export, compile, and integration handling.
  • docs/api.rst

    • Public API entries.
  • docs/models/models_categorization.rst

    • Foundation-model categorization.
  • docs/whats_new.rst

    • Release note and baseline explanation.
  • NOTICE.txt

    • CC BY-NC attribution and licensing terms.

Checklist

  • Added public model implementations.
  • Followed EEGModuleMixin conventions.
  • Added return_features and reset_head.
  • Added validation and variable-channel masking.
  • Added official checkpoint loaders.
  • Tested official base and staging checkpoints.
  • Added focused and shared integration tests.
  • Added API and model-category documentation.
  • Documented the paper’s baselines accurately.
  • Added release notes.
  • Added upstream attribution and licensing notice.
  • Ran all configured pre-commit hooks.
  • No new runtime dependency was introduced.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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) and SleepFMStager (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.models API.
  • 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.

Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
@adammounir

Copy link
Copy Markdown
Collaborator

Thanks for this contribution @Fashad-Ahmed 👍 it's a thorough, well-documented port. A few things needed before we can move toward merge:

  1. Conflicts. The PR is currently conflicting with master. Could you please rebase on the latest master and re-resolve? CI won't run cleanly until then.

  2. Checkpoint hosting. The tests download the weights straight from the upstream repo. Our convention is to re-host on the braindecode HF org. I've just mirrored them for you (byte-identical & same layout) here: https://huggingface.co/braindecode/SleepFM. Please point the test URLs at the mirror:

https://huggingface.co/braindecode/SleepFM/resolve/main/model_base/best.pt
https://huggingface.co/braindecode/SleepFM/resolve/main/model_sleep_staging/best.pth

  1. Correctness. Copilot flagged a real one:
    In the torch.compile/torch.export paths, _validate_mask is skipped, so a sample with all channels masked divides by valid.sum()==0 → NaN/Inf (sleepfm.py:134).
    Minor: the mask validation error message says key_padding_mask rather than the public channel_mask (:362).

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.
@Fashad-Ahmed

Fashad-Ahmed commented Aug 17, 2026 •

Copy link
Copy Markdown
Contributor Author

@bruAristimunha Is there any thing left other than rebasing to the Codacy Static Code Analysis workflow, so that it can get merge

Comment thread test/unit_tests/models/test_sleepfm.py Outdated
Comment thread test/unit_tests/models/test_integration.py
Comment thread test/unit_tests/models/test_integration.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
Comment thread braindecode/models/sleepfm.py Outdated
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.
@bruAristimunha

bruAristimunha commented Aug 25, 2026 •

Copy link
Copy Markdown
Collaborator

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: 4cbb08c034f3be6f483cf336d1840a29cbf0dcac (master merged; outputs on the released weights identical to cd7645d7, max-abs difference 0).

SHHS test set (2,000 nights, released split): port 0.7925 macro-F1, authors' released code 0.7925, paper 0.78.

@codecov

codecov Bot commented Sep 21, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 95.85253% with 9 lines in your changes missing coverage. Please review.
✅ Project coverage is 88.61%. Comparing base (ea5bc7d) to head (9b6852e).
⚠️ Report is 2 commits behind head on master.

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:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

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.
@bruAristimunha bruAristimunha added model Adds a new model needs-replication Model PR: paper number must be replicated (NeuralBench) before merge labels Oct 5, 2026
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
@bruAristimunha
bruAristimunha dismissed their stale review October 7, 2026 13:06

Addressed: all review threads resolved; Bruno approved merging on 2026-10-07.

@bruAristimunha
bruAristimunha merged commit f869984 into braindecode:master Oct 7, 2026
15 checks passed
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Oct 7, 2026
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Oct 7, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

model Adds a new model needs-replication Model PR: paper number must be replicated (NeuralBench) before merge

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants