Skip to content

About

[KDD 2026] StageGuard: Physiologically Constrained Sleep Staging

Topics

Resources

Stars

2 stars

Watchers

0 watching

Forks

Latest commit

Β 

History

6 Commits

Folders and files

Repository files navigation

[KDD 2026 AI4Science] StageGuard: Physiologically Constrained Sleep Staging

Backbone-agnostic physiological constraints for neural sleep staging.

KDD 2026 License: MIT Python PyTorch Zenodo

Official implementation of the KDD 2026 AI4Science paper "StageGuard: Physiologically Constrained Sleep Staging" by Juntang Wang, Yihan Wang, Hao Wu, Jiayu Gao, Shixin Xu, and Dongmian Zou (Zu Chongzhi Center, Duke Kunshan University). Equal contribution: Juntang Wang, Yihan Wang, Hao Wu, Jiayu Gao. Corresponding author: Dongmian Zou (dongmian.zou@duke.edu).

Automated sleep staging is increasingly used in large-scale studies and downstream scientific analyses to derive sleep-architecture endpoints, including total sleep time, REM latency, sleep efficiency, and bout-duration statistics. Deep learning models achieve epoch-level accuracy approaching inter-rater agreement, yet often produce hypnograms that violate physiological invariants, such as rare state transitions (e.g., direct Wake to REM) or excessively fragmented sequences. Such violations can bias downstream sleep metrics, regardless of overall accuracy. We propose StageGuard, a plug-and-play, backbone-agnostic structured-inference framework that wraps any neural sleep-staging backbone with physiology-informed priors. StageGuard combines (1) a differentiable soft transition penalty that discourages physiologically rare transitions during training, and (2) a semi-Markov constrained decoder with a duration-augmented state space that jointly enforces transition penalties and minimum bout durations at inference. Across six backbones and four datasets, StageGuard reduces the transition-violation rate to physiologically plausible levels and lowers the fragmentation index by 56 to 62%, while maintaining or slightly improving classification accuracy, and improves the accuracy of derived sleep-architecture statistics that are not directly optimized by the method.

πŸ“° News

  • 2026 Accepted to KDD 2026 (AI4Science track), Jeju Island, Republic of Korea (Aug 9 to 13, 2026).

πŸ“Œ Method

  1. Soft Transition Penalty - A differentiable loss term that discourages physiologically rare stage transitions (e.g., Wake to/from REM) during training.
  2. Semi-Markov Decoder - An augmented Viterbi decoder that enforces minimum bout durations, penalizes rare transitions, and suppresses flip-flop artifacts at inference time.
  3. Signal Quality Integration - Per-epoch signal quality scores modulate decoder confidence, gracefully degrading predictions for noisy segments.

πŸ“¦ Installation

# With conda
conda env create -f environment.yml
conda activate stageguard

# Or with pip
pip install -e .

# Optional: notebook/plotting extras
pip install -e ".[demo]"

πŸš€ Demo

Run the interactive demo on Colab:

Open In Colab

It trains the real AccuSleep backbone on real mouse EEG (downloaded at runtime from the open AccuSleep dataset), then runs the semi-Markov decoder on a held-out night and shows the transition-violation rate (TVR) and fragmentation index (FI) drop while accuracy is held or improved, with a before/after hypnogram plot. The backbone is trained EEG-only, so it confuses Wake and REM and fragments bouts - exactly the errors the decoder repairs. Expect a ~233 MB download and roughly 3-5 minutes on a Colab CPU. No raw data is shipped in this repository; the notebook downloads it at runtime. See notebooks/demo_stageguard.ipynb.

⚑ Quick Start

import torch
from stageguard import StageGuardWrapper, ModalityConfig
from stageguard.backbones import get_backbone

# Load modality config
config = ModalityConfig.from_yaml("configs/mouse_eeg.yaml")

# Create backbone + StageGuard wrapper
backbone = get_backbone("accusleep", num_classes=config.num_classes)
model = StageGuardWrapper(backbone, config)

# Training step
x = torch.randn(4, 50, 1, 128)  # (batch, epochs, channels, samples)
targets = torch.randint(0, 3, (4, 50))
loss, details = model.training_step(x, targets)
loss.backward()

# Inference with constrained decoding
predictions = model.predict(x)  # (4, 50) numpy array

🧩 Wrap Your Own Backbone

Any nn.Module that outputs (B, T, num_classes) logits works:

import torch.nn as nn

class MyBackbone(nn.Module):
    def forward(self, x):
        # Your architecture here
        return logits  # (B, T, C)

model = StageGuardWrapper(MyBackbone(), config)

See docs/backbones.md for the full interface contract and the ModalityConfig schema.

πŸ—‚οΈ Repository Structure

StageGuard/
β”œβ”€β”€ configs/                      # Modality-specific YAML configurations
β”‚   β”œβ”€β”€ mouse_eeg.yaml            # AccuSleep: 3-state, 4s epochs
β”‚   β”œβ”€β”€ actigraphy.yaml           # Sleep-Accel: 2-state, 30s epochs
β”‚   β”œβ”€β”€ cardiorespiratory.yaml    # SHHS: 3-state, 30s epochs
β”‚   β”œβ”€β”€ bioradar.yaml             # SLEEPBRL: 3-state, 30s epochs
β”‚   └── accusleep_demo.yaml       # 3-state mouse-EEG config used by the demo (2.5s epochs)
β”œβ”€β”€ stageguard/
β”‚   β”œβ”€β”€ config.py                 # ModalityConfig dataclass + YAML loader
β”‚   β”œβ”€β”€ losses.py                 # SoftTransitionPenalty + stageguard_loss
β”‚   β”œβ”€β”€ decoder.py                # SemiMarkovDecoder (augmented Viterbi)
β”‚   β”œβ”€β”€ wrapper.py                # StageGuardWrapper (backbone + loss + decoder)
β”‚   β”œβ”€β”€ metrics.py                # TVR, FI, accuracy, kappa, F1, sleep architecture
β”‚   β”œβ”€β”€ sqi.py                    # Signal quality index per modality
β”‚   β”œβ”€β”€ backbones/                # Backbone implementations + registry
β”‚   └── data/                     # Dataset loaders (expect pre-downloaded data)
β”œβ”€β”€ notebooks/
β”‚   └── demo_stageguard.ipynb     # Colab-runnable train + decode demo (real mouse EEG)
β”œβ”€β”€ scripts/
β”‚   β”œβ”€β”€ build_accusleep_demo.py   # reproduce the demo's data step locally (downloads, no data shipped)
β”‚   └── smoke_test.py             # standalone forward / train / predict sanity check
β”œβ”€β”€ data/
β”‚   └── README.md                 # demo-data provenance, license, stage mapping
β”œβ”€β”€ docs/
β”‚   β”œβ”€β”€ datasets.md               # dataset access + loader formats
β”‚   β”œβ”€β”€ backbones.md              # interface contract + config schema
β”‚   └── dependencies.md           # third-party provenance + attribution
β”œβ”€β”€ examples/demo.py              # end-to-end synthetic data demo
β”œβ”€β”€ tests/                        # test suite
β”œβ”€β”€ ROADMAP.md
β”œβ”€β”€ CITATION.cff
└── .zenodo.json

πŸ“Š Datasets

Dataset Modality Classes Epoch Access
AccuSleep Mouse EEG/EMG 3 (Wake, NREM, REM) 2.5s (native; canonical config 4s) Open (OSF)
Sleep-Accel Wrist actigraphy 2 (Wake, Sleep) 30s Open, ODC-By (PhysioNet)
SHHS Cardiorespiratory 3 (Wake, Light, Deep) 30s Data-use agreement (NSRR)
SLEEPBRL Bioradar 3 (Wake, Light, Deep) 30s Contact authors

Datasets must be downloaded separately; some require data use agreements. Each loader provides download instructions via Dataset.download_instructions(). See docs/datasets.md for the expected on-disk formats and data/README.md for the demo-data provenance.

πŸŽ›οΈ Hyperparameters

Parameter Symbol Code default Description
lambda_trans lambda 1.0 Transition penalty weight
epsilon epsilon 5.0 Rare transition penalty in decoder
gamma gamma 2.0 Anti-flip-flop penalty
k k 5 Flip-flop lookback window (epochs)
d_min d_min per-stage Minimum bout duration (epochs)
d_max d_max 30 Maximum tracked duration (epochs)

These are the ModalityConfig code defaults; per-modality configs/*.yaml override them (for example mouse_eeg.yaml sets d_max 50, and the demo config uses d_max 20, k 4, lambda_trans 0.1).

πŸ” Development

pytest                                   # run the test suite
pytest tests/test_decoder.py             # a single module
python examples/demo.py                  # end-to-end on synthetic data
python scripts/smoke_test.py             # quick forward / train / predict sanity check
python scripts/build_accusleep_demo.py   # reproduce the demo's mouse-EEG data locally (downloads; nothing committed)

build_accusleep_demo.py writes a local data/accusleep_demo.npz (git-ignored); no raw data is committed, since the AccuSleep recordings carry no redistribution license (see data/README.md).

πŸ”— Dependencies and Attribution

The bundled accusleep and usleep backbones are clean reimplementations of published architectures (Barger et al. 2019; Perslev et al. 2021), not ports; no third-party source is vendored here. The four datasets are obtained from their own sources under their own licenses, and the demo downloads AccuSleep mouse EEG from its public OSF project at runtime rather than redistributing it. See docs/dependencies.md for the full provenance table.

πŸ™ Acknowledgments

We thank the authors of the datasets used in this work: AccuSleep (Barger et al., 2019), Sleep-Accel (Walch et al., 2019), SHHS (Quan et al., 1997), and SLEEPBRL (Tataraidze et al., 2015), and the authors of the AccuSleep and U-Sleep architectures that the reference backbones reimplement. Provenance and licenses are detailed in docs/dependencies.md.

πŸ“š Citation

If you use StageGuard in your research, please cite:

@inproceedings{wang2026stageguard,
  title     = {StageGuard: Physiologically Constrained Sleep Staging},
  author    = {Wang, Juntang and Wang, Yihan and Wu, Hao and Gao, Jiayu and Xu, Shixin and Zou, Dongmian},
  booktitle = {Proceedings of the 32nd ACM SIGKDD Conference on Knowledge Discovery and Data Mining V.2 (KDD '26)},
  year      = {2026},
  doi       = {10.1145/3770855.3818916}
}

A machine-readable CITATION.cff is also provided, which GitHub renders into a "Cite this repository" button on the project page.

πŸ“ License

Released under the MIT License. See LICENSE for details.

About

[KDD 2026] StageGuard: Physiologically Constrained Sleep Staging

Topics

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages