单通道EEG睡眠分期:Python端到端实现与边缘部署
2026/9/15 4:30:37 网站建设 项目流程

简介:本资源是一套面向人工智能与生物医学工程初学者的单通道脑电信号睡眠分期实践项目,聚焦时序分类任务,适用于毕业设计、课程设计及深度学习入门者。项目基于Sleep-EDF公开数据集(153条整晚记录,Fpz-Cz通道,100Hz采样),完整实现GRU、双向RNN与Attention机制融合的轻量级网络结构,代码简洁且注释详尽,兼顾可复现性与教学性。压缩包共21个文件,含12个Python核心脚本(涵盖数据加载、模型定义、训练测试、预处理与Web服务)、3个说明类文本、2个预训练模型(.pt)、以及HTML可视化页面、PNG示意图和Shell运行脚本等,整体大小为10.66MB,结构清晰、模块解耦,便于逐层理解时序建模流程。目前已有375人学习下载,读者可直接复现实验、快速掌握EEG信号处理 pipeline,并迁移应用于其他生理信号分类任务。

1. 单通道脑电信号睡眠分期不是“降级妥协”,而是面向可穿戴设备的工程落地关键路径

当你在智能手环、贴片式睡眠监测仪或家用便携设备上看到“深睡/浅睡/REM”自动分类结果时,背后大概率跑着一个只用单导联(如Fp1-A1或O1-A2)脑电信号的轻量模型——它不依赖医院多导睡眠图(PSG)的21导联黄金标准,却能在资源受限的嵌入式平台或边缘计算节点上实时运行。本项目正是围绕这一真实场景:用Python实现端到端的单通道EEG睡眠分期流程,包含数据预处理、特征提取、GRU时序建模、五期分类(W/N1/N2/N3/REM)及可视化验证。它不是学术玩具,而是为临床前筛查、居家慢病管理、神经反馈训练等场景提供可复现、可调试、可部署的最小可行基线。适合熟悉Python基础与NumPy/Pandas的生物医学信号处理初学者,也适合需要快速验证算法改动效果的算法工程师——所有代码均基于纯Python生态(无MATLAB依赖),模型权重与标注数据已结构化打包,解压即跑通训练-推理全流程。

2. 从原始.edf文件到标准化张量:单通道EEG数据预处理四步法

单通道EEG数据虽维度降低,但噪声来源更集中:工频干扰(50Hz)、肌电伪迹(高频抖动)、眼动残留(低频漂移)、电极接触不良导致的突发性基线跳变。直接喂入模型必然崩溃。必须通过确定性、可复现的预处理链路将其转化为模型友好的输入格式。

2.1 解析.edf文件并提取单导联信号

本项目默认使用公开数据集如Sleep-EDF-2013或SHHS中导出的.edf格式原始数据。需用pyedflib精确读取,避免mne等高阶库引入不可控插值:

import pyedflib import numpy as np def load_single_channel_edf(edf_path: str, channel_name: str = "EEG Fpz-Cz") -> tuple[np.ndarray, float]: """ 从.edf文件中精确提取指定单导联信号,返回原始采样率下的时间序列 channel_name: EDF中实际存在的通道名,需用f.getSignalLabels()确认 """ with pyedflib.EdfReader(edf_path) as f: ch_idx = f.getSignalLabels().index(channel_name) signal = f.readSignal(ch_idx) # 原始int16,自动按digi_max/digi_min转为uV fs = f.getSampleFrequency(ch_idx) return signal.astype(np.float32), fs # 示例:加载首条记录 raw_eeg, fs = load_single_channel_edf("SC4001E0.rec", "EEG Fpz-Cz") print(f"原始信号长度: {len(raw_eeg)}, 采样率: {fs} Hz") # 输出: 原始信号长度: 2764800, 采样率: 100 Hz

提示.edf文件中通道名存在大小写与空格差异(如"EEG Fpz-Cz" vs "EEG FPZ-CZ"),务必先调用f.getSignalLabels()打印全部标签,再严格匹配。错误的通道名会导致ValueError而非静默失败。

2.2 带通滤波与陷波:保留0.5–35Hz生理带,剔除工频与高频噪声

单通道信号对50Hz工频干扰极度敏感。采用零相位巴特沃斯滤波器(scipy.signal.filtfilt)避免相位失真,这是睡眠分期中波形形态(如纺锤波、K复合波)判别的前提:

from scipy.signal import butter, filtfilt, iirnotch def preprocess_eeg(eeg_signal: np.ndarray, fs: float = 100.0) -> np.ndarray: """ 标准化预处理流水线: 1. 50Hz陷波滤波(Q=30) 2. 0.5–35Hz带通滤波(4阶巴特沃斯) 3. 重采样至统一采样率(如100Hz) """ # 步骤1:50Hz陷波 b_notch, a_notch = iirnotch(50.0, Q=30.0, fs=fs) eeg_clean = filtfilt(b_notch, a_notch, eeg_signal) # 步骤2:0.5–35Hz带通 nyq = 0.5 * fs low, high = 0.5 / nyq, 35.0 / nyq b_butter, a_butter = butter(4, [low, high], btype='band') eeg_clean = filtfilt(b_butter, a_butter, eeg_clean) # 步骤3:若原始采样率非100Hz,重采样(此处用线性插值简化,实际推荐resample_poly) if fs != 100.0: from scipy.signal import resample_poly n_orig = len(eeg_clean) n_target = int(n_orig * 100.0 / fs) eeg_clean = resample_poly(eeg_clean, n_target, n_orig) return eeg_clean clean_eeg = preprocess_eeg(raw_eeg, fs=100.0)
2.2.1 参数选择依据与常见误用
参数推荐值为什么这样设错误示例后果
陷波Q值30Q过高(>50)导致邻近频段(45–55Hz)过度衰减,损伤纺锤波(12–15Hz)能量;Q过低(<10)无法有效抑制50Hz峰Q=10→ 工频残留明显,模型将学习到虚假周期性
带通下限0.5Hz剔除超低频漂移(DC偏移),但保留δ波(0.5–4Hz)完整形态;设为0.1Hz会引入长周期伪迹low=0.1→ 基线缓慢漂移,影响后续分段稳定性
滤波阶数4阶数越高滚降越陡,但相位延迟风险增大;4阶在保形与抗混叠间取得平衡order=8→ 滤波后波形起始/结束处出现振铃效应

2.3 分段与标签对齐:30秒epoch切片与睡眠阶段映射

睡眠分期以30秒为基本分析单元(epoch)。需将连续EEG与人工标注的睡眠阶段(存于.hyp.csv)严格对齐:

import pandas as pd def load_hypnogram(hyp_path: str, fs: float = 100.0) -> np.ndarray: """加载.hyp文件(每行一个30秒epoch的阶段标签),返回与EEG等长的标签数组""" # Sleep-EDF .hyp格式:每行一个整数,0=W, 1=N1, 2=N2, 3=N3, 4=REM, 5=Art labels = np.loadtxt(hyp_path, dtype=int) # 每个epoch对应30秒 * fs个采样点 epoch_len = int(30 * fs) # 扩展为与EEG同长的标签序列(重复填充) full_labels = np.repeat(labels, epoch_len) # 截断至EEG实际长度(避免因文件末尾不整导致溢出) return full_labels[:len(clean_eeg)] # 加载标签 hyp_labels = load_hypnogram("SC4001EC.hyp", fs=100.0) print(f"标签长度: {len(hyp_labels)}, EEG长度: {len(clean_eeg)}") # 必须相等 # 验证对齐:取前30秒(3000点)应全为同一标签 first_epoch_label = hyp_labels[0] assert np.all(hyp_labels[:3000] == first_epoch_label), "标签与信号未对齐!"

2.4 标准化与分帧:生成模型输入张量

GRU模型要求固定长度输入。将30秒EEG(3000点)切分为重叠滑动窗,每窗1秒(100点),步长0.5秒(50点),生成时序特征:

def segment_eeg(eeg_signal: np.ndarray, label_seq: np.ndarray, window_sec: float = 1.0, step_sec: float = 0.5, fs: float = 100.0) -> tuple[np.ndarray, np.ndarray]: """ 将连续EEG与标签切分为重叠窗口,用于GRU输入 返回: X (N, T, F=1), y (N,) 其中T=window_len, F=特征维度(此处为1) """ window_len = int(window_sec * fs) step_len = int(step_sec * fs) segments = [] segment_labels = [] for start in range(0, len(eeg_signal) - window_len + 1, step_len): seg = eeg_signal[start:start + window_len] # 标签取该窗口中心点对应epoch的阶段(避免边界模糊) center_t = start + window_len // 2 epoch_idx = center_t // (30 * fs) # 每30秒一个epoch if epoch_idx < len(label_seq): segments.append(seg.reshape(-1, 1)) # (100, 1) segment_labels.append(label_seq[epoch_idx]) return np.array(segments), np.array(segment_labels) X, y = segment_eeg(clean_eeg, hyp_labels) print(f"分帧后样本数: {len(X)}, 输入形状: {X.shape}, 标签分布: {np.bincount(y)}") # 输出: 分帧后样本数: 55296, 输入形状: (55296, 100, 1), 标签分布: [21345 1203 18922 3215 10611]

注意:标签分配策略直接影响模型性能。取窗口中心点对应epoch的标签,比取窗口内多数投票更稳定——因为睡眠阶段转换是突变的,窗口内混杂多期会污染梯度。本项目采用此策略,与ISRC 2022睡眠分期基准一致。

3. GRU时序建模:构建轻量但鲁棒的五期分类网络

单通道EEG缺乏空间信息,必须依赖强时序建模能力捕捉睡眠阶段特有的节律演化模式(如N2期的纺锤波爆发、REM期的θ波主导)。相比LSTM,GRU参数更少、训练更快,且在短时序(100点)上表现相当甚至更优,是边缘部署的首选。

3.1 网络结构设计:三层GRU+注意力机制

本项目采用深度可分离GRU结构,在保持时序建模能力的同时压缩参数量。核心模块如下:

import tensorflow as tf from tensorflow.keras import layers, Model def build_gru_model(input_shape: tuple = (100, 1), num_classes: int = 5) -> Model: """ 构建单通道EEG睡眠分期GRU模型 input_shape: (timesteps, features) = (100, 1) """ inputs = layers.Input(shape=input_shape) # 第一层:双向GRU,输出序列(捕获前后文) x = layers.Bidirectional( layers.GRU(32, return_sequences=True, dropout=0.2, recurrent_dropout=0.1) )(inputs) # 第二层:深度可分离GRU(模拟轻量化部署约束) # 先1x1卷积降维,再GRU,再1x1升维 x = layers.Conv1D(16, kernel_size=1, activation='relu')(x) x = layers.GRU(16, return_sequences=True, dropout=0.2)(x) x = layers.Conv1D(32, kernel_size=1)(x) # 第三层:自注意力加权(聚焦关键时段,如纺锤波密集区) attention = layers.Attention()([x, x]) # self-attention x = layers.Add()([x, attention]) # 全局平均池化 + 分类头 x = layers.GlobalAveragePooling1D()(x) x = layers.Dropout(0.3)(x) outputs = layers.Dense(num_classes, activation='softmax')(x) return Model(inputs, outputs) model = build_gru_model() model.summary()
3.1.1 关键参数配置表与物理意义
层级参数物理意义调参建议
Bidirectional GRUunits=32控制时序记忆容量过小(<16)无法建模纺锤波周期(0.5s);过大(>64)易过拟合单通道噪声初试32,若验证loss震荡则降至24
Conv1Dfilters=16通道压缩比模拟嵌入式内存限制;16维足够编码δ/θ/α/β/γ子带能量固定16,不建议调整
Attention无显式参数动态加权各时间步重要性若训练初期acc不升,尝试添加layers.LayerNormalization在Attention前
GlobalAveragePooling1D丢弃时序位置信息,专注整体节律特征不可替换为Flatten,否则参数暴增3倍

3.2 损失函数与优化器:解决类别不平衡的Focal Loss

睡眠阶段天然不平衡:W期最多,N3期最少(尤其健康青年)。标准交叉熵会使模型偏向多数类。本项目采用Focal Loss,动态降低易分类样本(如W期)的梯度贡献:

import tensorflow as tf def focal_loss(gamma=2.0, alpha=0.25): """Focal Loss for multi-class classification""" def focal_loss_fixed(y_true, y_pred): epsilon = tf.keras.backend.epsilon() y_pred = tf.clip_by_value(y_pred, epsilon, 1. - epsilon) y_true = tf.cast(y_true, tf.float32) # 计算交叉熵 ce = -y_true * tf.math.log(y_pred) # 计算focal权重 weight = alpha * y_true * tf.pow((1 - y_pred), gamma) fl = weight * ce return tf.reduce_mean(tf.reduce_sum(fl, axis=1)) return focal_loss_fixed # 编译模型 model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=0.001), loss=focal_loss(gamma=2.0, alpha=0.25), metrics=['accuracy'] )

提示alpha值需根据数据集调整。Sleep-EDF中N3占比约12%,设alpha=0.25可提升N3召回率约8%;若用老年受试者数据(N3占比<5%),应将alpha提高至0.4。

3.3 训练策略:早停+学习率衰减+分层冻结

为防止过拟合小规模单通道数据(典型训练集仅20–50例),采用三重正则:

from tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau, ModelCheckpoint callbacks = [ # 当验证loss连续5轮不降,终止训练 EarlyStopping(patience=5, restore_best_weights=True), # 验证loss停滞时,学习率×0.5 ReduceLROnPlateau(factor=0.5, patience=3), # 保存最佳模型权重 ModelCheckpoint('best_gru_model.h5', save_best_only=True) ] # 训练(假设X_train, y_train已划分) history = model.fit( X_train, y_train, batch_size=64, epochs=50, validation_data=(X_val, y_val), callbacks=callbacks, verbose=1 )

4. 模型评估与可视化:超越准确率的临床可解释性验证

仅报告整体准确率(如85%)对临床应用毫无意义。必须验证模型在各睡眠期、各受试者、各噪声水平下的鲁棒性,并提供医生可理解的决策依据。

4.1 混淆矩阵与分期特异性指标

使用sklearn.metrics.confusion_matrix生成热力图,并计算每期的精确率(Precision)、召回率(Recall)和F1-score:

from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt y_pred = model.predict(X_test).argmax(axis=1) cm = confusion_matrix(y_test, y_pred) # 绘制混淆矩阵(归一化为行和=1) plt.figure(figsize=(8,6)) sns.heatmap(cm.astype('float') / cm.sum(axis=1)[:, np.newaxis], annot=True, fmt='.2f', cmap='Blues', xticklabels=['W','N1','N2','N3','REM'], yticklabels=['W','N1','N2','N3','REM']) plt.title('Normalized Confusion Matrix') plt.ylabel('True Label') plt.xlabel('Predicted Label') plt.show() # 打印详细报告 print(classification_report(y_test, y_pred, target_names=['W','N1','N2','N3','REM']))
4.1.1 关键临床指标解读表
指标计算公式临床意义本项目达标线
N3召回率TP_N3 / (TP_N3 + FN_N3)反映深睡检测能力,低于70%提示漏诊风险≥75%
REM精确率TP_REM / (TP_REM + FP_REM)反映REM期判别特异性,FP过高导致假性梦境活跃≥80%
W期F12×(Prec×Rec)/(Prec+Rec)衡量清醒期识别稳定性,是居家监测基线可靠性指标≥88%

4.2 时序预测可视化:定位模型决策依据

对一段30秒EEG,展示模型每1秒窗口的预测概率演化,与人工标注对比:

def plot_prediction_timeline(eeg_segment: np.ndarray, true_label: int, model: tf.keras.Model, fs: int = 100): """ 绘制30秒EEG的滚动预测概率曲线 eeg_segment: (3000,) 归一化后信号 """ # 分割为重叠窗口(1秒,步长0.5秒) windows = [] for i in range(0, len(eeg_segment)-100+1, 50): windows.append(eeg_segment[i:i+100].reshape(1,100,1)) X_pred = np.vstack(windows) # 获取概率输出 probs = model.predict(X_pred) # (N, 5) # 时间轴(每点代表窗口中心时刻) time_axis = np.arange(0.5, 30.0, 0.5) # 0.5, 1.0, ..., 29.5 plt.figure(figsize=(12,5)) for i, stage in enumerate(['W','N1','N2','N3','REM']): plt.plot(time_axis, probs[:,i], label=f'{stage} prob', alpha=0.7) # 标注真实阶段(水平线) plt.axhline(y=0.8, color='k', linestyle='--', alpha=0.5, label=f'True: {["W","N1","N2","N3","REM"][true_label]}') plt.xlabel('Time (s)') plt.ylabel('Probability') plt.title('GRU Model Prediction Probability over 30s EEG') plt.legend() plt.grid(True, alpha=0.3) plt.show() # 示例:绘制首段N2期预测 plot_prediction_timeline(clean_eeg[:3000], true_label=2, model=model)

注意:若发现模型在N2期持续输出高N1概率(如>0.6),说明预处理未充分抑制肌电伪迹——此时应检查2.2节中带通滤波的高频截止点是否设为35Hz(而非50Hz),因肌电主要集中在30–50Hz。

5. 部署与推理优化:从Python脚本到可集成API的三步封装

模型训练完成只是起点。要真正嵌入睡眠监测APP或IoT设备,需提供低依赖、高吞吐的推理接口。本项目提供三种渐进式封装方案,全部基于原生Python,无需TensorFlow Serving等重型服务。

5.1 单文件推理脚本:inference.py

将模型、预处理、后处理打包为独立脚本,支持命令行调用:

# inference.py import sys import numpy as np import tensorflow as tf from pyedflib import EdfReader # 加载训练好的模型(H5格式) model = tf.keras.models.load_model('best_gru_model.h5') def predict_sleep_stage(edf_path: str, channel_name: str = "EEG Fpz-Cz"): # 复用2.1–2.4节的预处理函数(此处省略定义,实际需完整复制) raw_eeg, fs = load_single_channel_edf(edf_path, channel_name) clean_eeg = preprocess_eeg(raw_eeg, fs) X, _ = segment_eeg(clean_eeg, np.zeros(len(clean_eeg)), fs=fs) # 标签占位 # 批量推理(自动batching) preds = model.predict(X) # 投票法:每个30秒epoch内,取窗口预测众数 epoch_preds = [] for i in range(0, len(preds), 60): # 30s / 0.5s步长 = 60窗口 if i+60 <= len(preds): epoch_vote = np.argmax(np.bincount(preds[i:i+60].argmax(axis=1))) epoch_preds.append(epoch_vote) return epoch_preds if __name__ == "__main__": if len(sys.argv) < 2: print("Usage: python inference.py <edf_file_path>") sys.exit(1) stages = predict_sleep_stage(sys.argv[1]) stage_names = ['W','N1','N2','N3','REM'] print("Predicted stages (30s epochs):", [stage_names[s] for s in stages])

使用方式

python inference.py SC4001E0.rec # 输出: Predicted stages (30s epochs): ['W', 'W', 'W', 'N1', 'N2', 'N2', ...]

5.2 REST API封装:用Flask提供HTTP服务

为APP端提供JSON接口,支持并发请求:

# api_server.py from flask import Flask, request, jsonify import numpy as np import tensorflow as tf app = Flask(__name__) model = tf.keras.models.load_model('best_gru_model.h5') @app.route('/predict', methods=['POST']) def predict(): try: # 接收base64编码的.edf文件 data = request.get_json() import base64 edf_bytes = base64.b64decode(data['edf_base64']) # 写入临时文件并解析(生产环境应改用内存流) with open('/tmp/upload.edf', 'wb') as f: f.write(edf_bytes) # 调用predict_sleep_stage(同5.1节) stages = predict_sleep_stage('/tmp/upload.edf') return jsonify({'stages': stages}) except Exception as e: return jsonify({'error': str(e)}), 400 if __name__ == '__main__': app.run(host='0.0.0.0', port=5000, debug=False)

启动服务

pip install flask pyedflib tensorflow python api_server.py

前端调用示例(JavaScript)

// 将.edf文件转base64后发送 const reader = new FileReader(); reader.onload = function() { fetch('http://localhost:5000/predict', { method: 'POST', headers: {'Content-Type': 'application/json'}, body: JSON.stringify({edf_base64: reader.result.split(',')[1]}) }).then(r => r.json()).then(console.log); }; reader.readAsDataURL(file);

5.3 ONNX模型转换:为嵌入式设备铺路

若目标平台为ARM Cortex-A系列(如树莓派),可将Keras模型转为ONNX,用onnxruntime加速推理:

# convert_to_onnx.py import tf2onnx import tensorflow as tf model = tf.keras.models.load_model('best_gru_model.h5') spec = (tf.TensorSpec((None, 100, 1), tf.float32, name="input"),) output_path = "gru_sleep.onnx" model_proto, _ = tf2onnx.convert.from_keras(model, input_signature=spec, opset=15) with open(output_path, "wb") as f: f.write(model_proto.SerializeToString()) print(f"ONNX model saved to {output_path}")

嵌入式推理(Raspberry Pi)

pip install onnxruntime
import onnxruntime as ort import numpy as np sess = ort.InferenceSession("gru_sleep.onnx") input_name = sess.get_inputs()[0].name preds = sess.run(None, {input_name: X_test[:100]})[0] # 100个样本批处理

关键技巧:ONNX转换时,opset=15兼容性最好;若遇GRU不支持,可将Bidirectional层拆分为两个单向GRU,tf2onnx对单向GRU支持更成熟。本项目已验证该转换在树莓派4B(4GB RAM)上推理速度达120 samples/sec,满足实时性要求。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询