braindecode.models.SleepFM#
- class braindecode.models.SleepFM(n_outputs=None, n_chans=None, chs_info=None, n_times=None, input_window_seconds=None, sfreq=None, patch_size=640, embed_dim=128, num_heads=8, num_layers=6, pooling_heads=8, drop_prob=0.3, max_seq_length=128, activation=<class 'torch.nn.modules.activation.ELU'>, channel_strategy='native', channel_strategy_kwargs=None)[source]#
Sleep foundation model for multimodal polysomnography [sleepfm2026].
Foundation Model Attention/Transformer Convolution
This class is the pretrained encoder of panel b: a 1D CNN per channel, channel pooling, a temporal Transformer and temporal pooling. It was pretrained with leave-one-out contrastive learning, which aligns each modality (BAS, EKG, RESP, EMG) with the others, on over 585,000 h of PSG from SSC, BioSerenity, MESA and MrOS; SHHS was held out.
Every channel is cut into non-overlapping 5-second patches (
patch_sizesamples at 128 Hz; trailing samples are dropped) and embedded by a shared convolutional tokenizer. Attention pools the variable channel set of each patch, a Transformer models the patch sequence and a second attention layer pools it into one trial-level vector (seeencode()).- 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.
patch_size (
int) – Samples per patch; at least 64 and divisible by 64.embed_dim (
int) – Token and Transformer embedding dimension.num_heads (
int) – Heads of the temporal Transformer.num_layers (
int) – Layers of the temporal Transformer.pooling_heads (
int) – Heads of the channel and temporal attention pooling.drop_prob (
float) – Dropout probability of the attention and Transformer layers.max_seq_length (
int) – Maximum number of patches (length of the positional table).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
The encoder was checked through
SleepFMStageronly (SHHS sleep staging, see there); its pretraining was not reproduced.channel_mask((batch, n_chans),Truefor missing channels) keeps masked channels out of every output, and out of the tokenizer’s batch normalization in training. The official pretraining code normalizes zero-padded channels with the real ones; both agree when no channel is masked, and always in eval mode.Important
Pre-trained Weights Available
The released encoder is mirrored at braindecode/SleepFM (revision
acecf041, CC BY-NC 4.0).final_layeris not pretrained: the release is a contrastive encoder.from braindecode.models import SleepFM model = SleepFM.from_pretrained(n_chans=4, n_times=3840, n_outputs=5)
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 SleepFM # Train your model model = SleepFM(n_chans=22, n_outputs=4, n_times=1000) # ... training code ... # Push to the Hub model.push_to_hub( repo_id="username/my-sleepfm-model", commit_message="Initial model upload", )
Loading a model from the Hub:
from braindecode.models import SleepFM # Load pretrained model model = SleepFM.from_pretrained("username/my-sleepfm-model") # Load with a different number of outputs (head is rebuilt automatically) model = SleepFM.from_pretrained("username/my-sleepfm-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 = SleepFM.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
- encode(x, channel_mask=None)[source]#
Return the pooled and the per-patch representations.
With \(t_{b,c,p}\) the token of channel \(c\) and patch \(p\), and masked means over the unmasked items:
\[z_{b,p} = \operatorname{MaskedMean}_{c} \big(\operatorname{SelfAttn}(t_{b,:,p})\big), \quad h_{b} = \operatorname{Transformer} \big(\operatorname{LN}(z_{b} + \mathrm{PE})\big), \quad g_{b} = \operatorname{Mean}_{p} \big(\operatorname{SelfAttn}(h_{b})\big).\]- Parameters:
- Return type:
- Returns:
pooled (torch.Tensor) – \(g\), shape
(batch, embed_dim).contextual_tokens (torch.Tensor) – \(h\), shape
(batch, n_patches, embed_dim).
- forward(x, channel_mask=None, return_features=False)[source]#
Return trial logits, or the pooled features.
- classmethod from_pretrained(*args, **kwargs)[source]#
Load the encoder; the repo defaults to
"braindecode/SleepFM".- Parameters:
*args – The description is missing.
**kwargs – The description is missing.
- reset_head(n_outputs)[source]#
Replace the trial-level 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.