┌──────────────────────────────────────────────────────────────┐
│ SleepStageNet │
├──────────────────────────────────────────────────────────────┤
│ │
│ 输入: (batch, T, 4) ← T个30秒epoch, 每个有4个特征 │
│ │
│ ┌──────────────────────────────────────────────────┐ │
│ │ 1. Feature Projection (参考SleepPPG-Net) │ │
│ │ MLP: 4 → d_model*2 → d_model │ │
│ └──────────────────┬───────────────────────────────┘ │
│ │ │
│ ┌──────────────────┴───────────────────────────────┐ │
│ │ 2. Cross-Feature Attention (参考wav2sleep) │ │
│ │ 每个特征独立投影 + CLS Token + Transformer │ │
│ │ 学习 HRV↔HR↔RR↔Movement 的交互关系 │ │
│ └──────────────────┬───────────────────────────────┘ │
│ │ (门控融合) │
│ ┌──────────────────┴───────────────────────────────┐ │
│ │ 3. Positional Encoding │ │
│ │ 正弦位置编码 (时间位置对睡眠结构很重要) │ │
│ └──────────────────┬───────────────────────────────┘ │
│ │ │
│ ┌──────────────────┴───────────────────────────────┐ │
│ │ 4. Dilated Temporal CNN (参考wav2sleep) │ │
│ │ 2 blocks × [d=1,2,4,8,16,32], k=7 │ │
│ │ 感受野 ≈ 6小时 → 捕获完整睡眠周期 │ │
│ └──────────────────┬───────────────────────────────┘ │
│ │ │
│ ┌──────────────────┴───────────────────────────────┐ │
│ │ 5. Classification Head │ │
│ │ Linear(d_model → d_model/2 → n_classes) │ │
│ └──────────────────────────────────────────────────┘ │
│ │
│ 输出: (batch, T, n_classes) ← 每个epoch的分类logits │
└──────────────────────────────────────────────────────────────┘
1import numpy as np
2import torch
3from sleep_staging_model import create_model
4
5# 创建模型
6model = create_model('base', n_features=4, n_classes=4)
7
8# 准备输入数据: T个30秒epoch的特征
9T = 1200 # 10小时 = 1200个epoch
10features = np.stack([
11 hrv_rmssd, # HRV (RMSSD) 序列
12 heart_rate, # 心率序列
13 respiratory_rate, # 呼吸频率序列
14 body_movement, # 体动序列
15], axis=-1) # shape: (T, 4)
16
17# Z-score标准化 (关键步骤!)
18features = (features - features.mean(axis=0)) / (features.std(axis=0) + 1e-8)
19features = np.clip(features, -5, 5)
20
21# 推理
22x = torch.tensor(features, dtype=torch.float32).unsqueeze(0) # (1, T, 4)
23model.eval()
24with torch.no_grad():
25 logits = model(x) # (1, T, n_classes)
26 predictions = torch.argmax(logits, dim=-1) # (1, T)
27
28# 标签: 0=Wake, 1=N1, 2=N2, 3=N3, 4=REM
29stage_names = {0: 'Wake', 1: 'N1', 2: 'N2', 3: 'N3', 4: 'REM'}