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
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.Encoder (
SleepFMwithout its trial pooling). Channels are grouped bychannel_modalities; each modality is encoded in independent chunks ofencoder_chunk_patchespatches (5 minutes in the release, positions restart at every chunk). A shorter trailing chunk is kept, whereas the official embedding script drops it.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 tomne.Infofor 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.Noneputs 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.ELUmatches 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 aValueError.
- 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),Truefor 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_maskworks as inSleepFM; 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.
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_hubpackage 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.
- 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") atencoder_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.