braindecode.models.SleepFMStager#

class braindecode.models.SleepFMStager(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, channel_modalities=None, patch_size=640, embed_dim=128, encoder_num_heads=8, encoder_num_layers=6, encoder_pooling_heads=8, encoder_drop_prob=0.0, encoder_chunk_patches=60, staging_num_heads=4, staging_num_layers=1, staging_pooling_heads=4, drop_prob=0.3, max_seq_length=8196, activation=<class 'torch.nn.modules.activation.ELU'>, channel_strategy='native', channel_strategy_kwargs=None)[source]#

SleepFM encoder with the released patch-wise sleep-staging head [sleepfm2026].

Foundation Model Attention/Transformer Convolution Recurrent

SleepFM overview (Thapa et al., 2026, Fig. 1).

This class is the fine-tuned head of panel c (modality pooling, LSTM, fully connected layer) on top of the panel b encoder of SleepFM. The released staging head was fine-tuned on frozen encoder embeddings of SSC, MESA, MrOS and SHHS.

  1. Encoder (SleepFM without its trial pooling). Channels are grouped by channel_modalities; each modality is encoded in independent chunks of encoder_chunk_patches patches (5 minutes in the release, positions restart at every chunk). A shorter trailing chunk is kept, whereas the official embedding script drops it.

  2. Staging head. Attention pools the modalities of every patch (always at least four slots, missing ones masked, as released), then a Transformer and a bidirectional LSTM run over the night.

The output has shape (batch, n_outputs, n_patches): one prediction per 5-second patch (Wake, N1, N2, N3, REM for the release), so a 30-second scored epoch spans six predictions.

Parameters:
  • n_outputs (int) – Number of outputs of the model. This is the number of classes in the case of classification.

  • n_chans (int) – Number of EEG channels.

  • chs_info (list of dict) – Information about each individual EEG channel. This should be filled with info["chs"]. Refer to mne.Info for more details.

  • n_times (int) – Number of time samples of the input window.

  • input_window_seconds (float) – Length of the input window in seconds.

  • sfreq (float) – Sampling frequency of the EEG recordings.

  • channel_modalities (Sequence[Hashable] | None) – Modality label of every channel, e.g. ["BAS", "RESP", "EKG", "EMG"] as in the release. None puts every channel in one modality.

  • patch_size (int) – Samples per patch.

  • embed_dim (int) – Token and recurrent feature dimension.

  • encoder_num_heads (int) – Heads of the encoder’s temporal Transformer.

  • encoder_num_layers (int) – Layers of the encoder’s temporal Transformer.

  • encoder_pooling_heads (int) – Heads of the encoder’s channel attention pooling.

  • encoder_drop_prob (float) – Encoder dropout; the release embeds with dropout disabled.

  • encoder_chunk_patches (int) – Patches the encoder sees at once. Above 128 the encoder’s positional table no longer loads from the released weights.

  • staging_num_heads (int) – Heads of the staging Transformer.

  • staging_num_layers (int) – Layers of the staging Transformer and of the bidirectional LSTM.

  • staging_pooling_heads (int) – Heads of the modality attention pooling.

  • drop_prob (float) – Dropout probability of the staging head.

  • max_seq_length (int) – Maximum number of patches.

  • activation (type[Module]) – Tokenizer activation; nn.ELU matches the released weights.

  • channel_strategy (str) – Only "native": the input is polysomnography grouped by modality (EEG, EOG, ECG, EMG, respiration), not an EEG montage, so no channel strategy applies. Any other value raises a ValueError.

  • channel_strategy_kwargs (dict | None) – Must be None.

Raises:

ValueError – If some input signal-related parameters are not specified: and can not be inferred.

Notes

Replication. With the released weights, this port and the authors’ released code both reach 0.7925 macro-F1 on the SHHS test set (2,000 nights of the released split; paper: 0.78). Only this SHHS sleep-staging result was checked; other datasets and tasks of the paper were not.

temporal_mask ((batch, n_patches), True for padding) marks padded patches at the end of shorter recordings. As in the release they enter the head as zero embeddings masked out of its Transformer, but the LSTM still runs over them, so valid predictions depend on the amount of padding (never on its content). channel_mask works as in SleepFM; a modality without a valid channel is masked.

Important

Pre-trained Weights Available

The released stager is mirrored at braindecode/SleepFMStager (revision 8681fbad, CC BY-NC 4.0).

from braindecode.models import SleepFMStager

model = SleepFMStager.from_pretrained(
    n_chans=7,
    n_times=38400,
    n_outputs=5,
    channel_modalities=["BAS"] * 3 + ["RESP"] * 2 + ["EKG", "EMG"],
)

Added in version 1.9.

License

References

[sleepfm2026]

Thapa, R., Kjaer, M. R., He, B., et al. (2026). A multimodal sleep foundation model for disease prediction. Nature Medicine, 32, 752–762. https://doi.org/10.1038/s41591-025-04133-4

Hugging Face Hub integration

When the optional huggingface_hub package is installed, all models automatically gain the ability to be pushed to and loaded from the Hugging Face Hub. Install with:

pip install braindecode[hub]

Pushing a model to the Hub:

from braindecode.models import SleepFMStager

# Train your model
model = SleepFMStager(n_chans=22, n_outputs=4, n_times=1000)
# ... training code ...

# Push to the Hub
model.push_to_hub(
    repo_id="username/my-sleepfmstager-model",
    commit_message="Initial model upload",
)

Loading a model from the Hub:

from braindecode.models import SleepFMStager

# Load pretrained model
model = SleepFMStager.from_pretrained("username/my-sleepfmstager-model")

# Load with a different number of outputs (head is rebuilt automatically)
model = SleepFMStager.from_pretrained("username/my-sleepfmstager-model", n_outputs=4)

Extracting features and replacing the head:

import torch

x = torch.randn(1, model.n_chans, model.n_times)
# Extract encoder features (consistent dict across all models)
out = model(x, return_features=True)
features = out["features"]

# Replace the classification head
model.reset_head(n_outputs=10)

Saving and restoring full configuration:

import json

config = model.get_config()            # all __init__ params
with open("config.json", "w") as f:
    json.dump(config, f)

model2 = SleepFMStager.from_config(config)    # reconstruct (no weights)

All model parameters (both EEG-specific and model-specific such as dropout rates, activation functions, number of filters) are automatically saved to the Hub and restored when loading.

See Loading and Adapting Pretrained Foundation Models for a complete tutorial.

Methods

forward(x, channel_mask=None, return_features=False, temporal_mask=None)[source]#

Return patch-wise logits, or the per-patch LSTM features.

Parameters:
  • x (Tensor) – The description is missing.

  • channel_mask (Tensor | None) – The description is missing.

  • return_features (bool) – The description is missing.

  • temporal_mask (Tensor | None) – The description is missing.

Return type:

Tensor | dict[str, Tensor | None]

classmethod from_pretrained(*args, encoder_model_name_or_path=None, encoder_revision=None, **kwargs)[source]#

Load the released stager; the repo defaults to "braindecode/SleepFMStager".

Older revisions of that mirror hold only the tokenizer and the head; the missing encoder weights are then read from encoder_model_name_or_path (repo id or local directory, default "braindecode/SleepFM") at encoder_revision.

Parameters:
  • *args – The description is missing.

  • encoder_model_name_or_path – The description is missing.

  • encoder_revision – The description is missing.

  • **kwargs – The description is missing.

reset_head(n_outputs)[source]#

Replace the patch-wise output layer.

Parameters:

n_outputs (int) – New number of output classes.

Examples

>>> from braindecode.models import BENDR
>>> model = BENDR(n_chans=22, n_times=1000, n_outputs=4)
>>> model.reset_head(10)
>>> model.n_outputs
10

Added in version 1.4.