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

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

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_size samples 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 (see encode()).

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.

  • 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.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

The encoder was checked through SleepFMStager only (SHHS sleep staging, see there); its pretraining was not reproduced.

channel_mask ((batch, n_chans), True for 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_layer is 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.

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 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:
  • x (Tensor) – Input of shape (batch, n_chans, n_times).

  • channel_mask (Tensor | None) – (batch, n_chans) mask, True for missing channels.

Return type:

tuple[Tensor, Tensor]

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.

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

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

  • return_features (bool) – The description is missing.

Return type:

Tensor | dict[str, Tensor | None]

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.