SleepFMStager / README.md
bruAristimunha's picture
Self-contained stager export (braindecode PR #1106 @ b0ae1aa8)
95903b5 verified
|
Raw History Blame Contribute Delete
3.93 kB
metadata
license: cc-by-nc-4.0
library_name: braindecode
tags:
  - eeg
  - polysomnography
  - sleep-staging
  - foundation-model
  - braindecode

SleepFMStager — pretrained sleep stager

Mirror of the official SleepFM sleep-staging model, re-hosted for stable loading from Braindecode.

SleepFM is a multimodal polysomnography (PSG) foundation model introduced in:

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

Files

File Description
model.safetensors The complete stager (180 tensors: tokenizer, channel pooling, temporal Transformer and staging head), with the parameter names of braindecode.models.SleepFMStager
config.json Architecture of the checkpoint, read by from_pretrained()

Upstream ships the stager in two pieces: the encoder (channel-agnostic tokenizer, channel pooling and temporal Transformer) lives in model_base/best.pt and the staging head in model_sleep_staging/best.pth. This file merges both, so a single call returns a model that is pretrained end to end, its five-class output layer included. It holds the whole encoder except its trial-level temporal pooling, which sleep staging does not use. The tensors are those of the upstream artifacts; only the keys were rewritten to the library's parameter names. Loading this file or the two upstream ones gives bit-identical outputs.

The upstream artifacts themselves are kept, byte-for-byte, in braindecode/SleepFM.

Usage

from braindecode.models import SleepFMStager

# Defaults to this repository. The release encodes BAS, RESP, EKG and EMG
# channels as separate modalities; name the modality of each channel.
model = SleepFMStager.from_pretrained(
    n_chans=7,
    n_outputs=5,
    n_times=38400,
    sfreq=128,
    channel_modalities=["BAS"] * 3 + ["RESP"] * 2 + ["EKG", "EMG"],
)
model.eval()

config.json leaves channel_modalities unset (null), because it depends on the montage. Without it every channel is encoded as a single modality and from_pretrained warns.

The output has shape (batch, n_outputs, n_patches): one prediction per 5-second patch, not per 30-second scoring epoch, so six predictions cover one scored epoch. For this checkpoint the five classes are Wake, N1, N2, N3 and REM. Input must be sampled at 128 Hz. Pass n_outputs different from 5 to reinitialise the output layer for another label set.

Revisions

  • 8681fba: tokenizer and staging head only (93 tensors). SleepFMStager.from_pretrained reads the encoder's channel pooling and temporal Transformer from braindecode/SleepFM at load time.
  • Current revision: the complete stager (180 tensors), bit-identical to the upstream model_base/best.pt + model_sleep_staging/best.pth, re-exported with the fixed port of braindecode PR #1106. config.json now lists every constructor argument; the defaults (n_chans=4, n_times=3840, sfreq=128) are unchanged. Its outputs are bit-identical to those of 8681fba completed from braindecode/SleepFM, and both revisions load with SleepFMStager.from_pretrained (pass revision= to pin one).

License & attribution

These weights are not covered by Braindecode's BSD-3 license and inherit the upstream noncommercial terms. Re-hosted for reproducibility and stable availability only; attribution and the CC BY-NC 4.0 restriction are preserved.