ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

单通道EEG睡眠分期实战:TinySleepNet+GRU+Attention轻量Pipeline

单通道EEG睡眠分期实战:TinySleepNet+GRU+Attention轻量Pipeline 简介本资源是一项基于Python实现的单通道脑电信号自动睡眠分期研究项目面向生物医学工程、人工智能交叉领域的初学者与进阶学习者适用于本科毕设、课程设计及工程实训等实践场景。项目复现并改进了TinySleepNet架构集成双向RNN、GRU与Attention机制支持灵活调整序列长度与网络组件并采用Focal Loss提升类别不平衡下的分类性能配套完整训练、测试与数据预处理流程且集成WandB实验追踪功能。压缩包共22个文件12个Python核心脚本、3个文本配置/说明、2个预训练模型.pt、1个Shell执行脚本及HTML/MD/PNG等辅助文件总大小10.68MB结构清晰、模块解耦便于理解模型构建、数据加载与评估全流程。目前已有274人学习下载提供从数据读取、网络定义、损失设计到指标可视化accuracy/mF1/混淆矩阵的端到端可运行代码附带详细README与requirements依赖说明开箱即用。1. 单通道EEG睡眠分期不是“降维偷懒”而是临床落地的关键突破口用Python跑通TinySleepNet变体3分钟加载预训练GRU模型复现mf10.78的轻量级分期 pipeline你是不是也试过把多导联EEG数据喂进ResNet或Transformer结果显存爆掉、训练卡在第2轮、测试时Kappa系数忽高忽低别急——这个项目专治“EEG模型太重、临床没法用”的焦虑。它不玩花哨的多模态融合也不堆参数就用单通道Fpz-Cz或C4-A1脑电信号靠一个修改过的TinySleepNet骨架双向GRUAttention机制在普通笔记本i5-8250U GTX 1050Ti上跑通完整训练→验证→预测闭环。核心价值不在SOTA指标而在可复现、可部署、可解释train.py里focal loss加权难分样本test.py输出recall_confusion_matrics直接定位N1期漏判率web/server.py甚至搭了个Flask轻量API供医生端调用。适合毕设学生快速出图出表也适合工程师拿去改造成嵌入式端推理模块——毕竟model_GRU.pt只有12.3MB比一张高清CT图还小。关键词“python 脑电信号 自动睡眠分期”不是泛泛而谈是真正在sleepedf数据集上跑出来的、带完整预处理链路preprocessing.py含陷波滤波重采样滑动窗分段、可一键下载download_sleepedf.py、可替换网络结构network.py里用config开关切换GRU/BiGRU/Attention的实操资源。2. 从原始EDF到PyTorch Dataset单通道EEG预处理的四个硬性约束与dataset.py实现逻辑2.1 为什么必须用单通道临床设备限制倒逼出的工程合理性市面上90%的家用/便携式睡眠监测设备如Emotiv EPOC、NextMind头环、国产NeuroScan Lite只提供1~2路有效EEG通道。强行插值补多导联不仅引入伪迹还会让模型学到噪声相关性。本项目默认使用Fpz-Cz通道对应SC4001E0.rec.edf中的ch[0]在dataset.py中通过raw.pick_channels([Fpz-Cz])硬性锁定——这不是妥协而是对真实部署场景的尊重。注意若你的数据是C4-A1请在prepare_data.py第47行手动修改channel_name C4-A1否则load_raw会报KeyError。2.2 preprocessing.py里的三道滤波工序陷波、带通、重采样缺一不可原始EDF采样率常为100Hz/200Hz但睡眠分期特征集中在0.5–30Hz工频干扰50Hz/60Hz又极顽固。preprocessing.py采用三级流水线mne.filter.notch_filter(raw, freqs50)→ 先削50Hz基频国内电网标准raw.filter(l_freq0.5, h_freq30)→ 再切带通砍掉肌电高频噪声和慢波漂移raw.resample(sfreq128)→ 统一重采样至128Hz确保后续滑动窗长度整除seq_len30秒×128Hz3840点提示若你的数据来自美国设备60Hz工频需将notch_filter的freqs改为60否则陷波失效——这是新手最常翻车的第一步。# preprocessing.py 关键代码段第89行起 def preprocess_eeg(raw_path, output_dir): raw mne.io.read_raw_edf(raw_path, preloadTrue) raw.pick_channels([Fpz-Cz]) # 强制单通道 raw.notch_filter(freqs50) # 国内工频 raw.filter(l_freq0.5, h_freq30) # 生理带宽 raw.resample(sfreq128) # 标准化采样率 # 后续写入.npy文件供dataset.py读取这段代码后接raw.save()生成.npy缓存避免每次训练都重复IO。参数说明sfreq128不是随意选的——它能被常见seq_len30s/20s/15s整除且满足Nyquist定理最高分析频率30Hz需≥60Hz采样128Hz留足余量。2.3 dataset.py如何用torch.utils.data.Dataset封装时序数据seq_len与shuffle_seed的实战意义dataset.py继承自torch.utils.data.Dataset但关键创新在__getitem__里对单通道信号做滑动窗切片。核心参数seq_len384030秒×128Hz不是固定死的若你改用20秒窗口需同步调整seq_len2560并在train.py第112行修改batch_size 32 // (seq_len // 3840)保持显存占用稳定shuffle_seed42保证实验可复现相同seed下train/val划分完全一致避免因随机打乱导致mf1波动超±0.03# dataset.py 第56行 __getitem__ 实现 def __getitem__(self, idx): # 加载预处理好的.npy文件已按30秒分段存好 data np.load(self.data_files[idx]) label self.labels[idx] # 滑动窗截取若data长度seq_len随机起点截取若seq_len零填充 if len(data) self.seq_len: start np.random.randint(0, len(data) - self.seq_len 1) segment data[start:startself.seq_len] else: segment np.pad(data, (0, self.seq_len - len(data)), constant) return torch.tensor(segment, dtypetorch.float32), torch.tensor(label, dtypetorch.long)逻辑说明这里没用torch.nn.utils.rnn.pad_sequence因为单通道信号pad效率更高np.random.randint配合shuffle_seed确保每个epoch内不同batch的起始点随机但跨epoch可复现。2.4 prepare_data.py把Sleep-EDF公开数据集转成项目专用格式的三步法download_sleepedf.py能自动下载SC-subset含197个夜间记录但原始EDF需转换才能喂给dataset.py。prepare_data.py完成解析EDF header获取采样率、通道名按30秒切片对应1个睡眠分期标签将每段信号存为{subject_id}_{segment_id}.npy标签存为labels.npy执行命令python prepare_data.py --edf_dir ./Sleep-EDF/sc/ --output_dir ./data/ --seq_len 3840参数说明--seq_len 3840必须与dataset.py中定义一致--edf_dir指向解压后的SC目录含*.rec.edf和*.hypnogram.edf输出目录./data/将生成eeg_signal.npy所有片段拼接和labels.npy对应标签数组供后续train.py直接加载。3. 网络结构拆解TinySleepNet变体怎么在GRU层加Attentionnetwork.py的5个可配置开关详解3.1 TinySleepNet原始结构 vs 本项目修改点轻量化不是删层而是重定向信息流TinySleepNet原设计含CNN提取局部特征 RNN建模时序依赖。本项目保留CNN backbone3层Conv1dkernel_size7/5/3但在RNN部分彻底重构原版用单向LSTM → 改为nn.GRU计算更快、显存更省增加bidirectionalTrue开关默认开启让前向/后向GRU分别捕获入睡/醒转过程在GRU输出后插入SelfAttention模块非Transformer full attention而是简化版Scaled Dot-Product最终分类头用nn.Linear(256, 5)5类W, N1, N2, N3, REM注意Attention模块不增加参数量仅用query/key/value投影各128维但显著提升N1/N2区分度——这是mf1从0.72升到0.78的关键。3.2 network.py里的config字典5个布尔开关决定模型行为所有网络结构由config字典控制避免改代码。关键字段config键默认值作用切换建议use_bidirectionalTrueGRU是否双向临床数据少时设False防过拟合use_attentionTrue是否启用Attention关闭后mf1↓0.02但推理快15%use_focal_lossTruetrain.py是否调用focal_loss.py数据不均衡时必开N2占比60%N1仅8%dropout_rate0.3全连接层Dropout显存紧张时调至0.5gru_hidden_size128GRU隐藏层维度笔记本GPU设128A100可提至256# network.py 第32行 config定义 config { use_bidirectional: True, use_attention: True, use_focal_loss: True, dropout_rate: 0.3, gru_hidden_size: 128, }3.3 SelfAttention模块实现为什么不用nn.MultiheadAttentionnetwork.py第142行的SelfAttention是自研轻量版query/key/value均用nn.Linear(128, 128)投影非矩阵乘省显存attention score用torch.softmax((q k.T) / sqrt(128), dim-1)计算输出为attn_weights v再经LayerNorm归一化不采用nn.MultiheadAttention是因为它强制要求embed_dim整除num_heads而GRU输出128维难以拆分多头机制在单通道EEG上易学噪声单头LayerNorm更鲁棒# network.py 第142行 SelfAttention.forward() def forward(self, x): # x: [batch, seq_len, hidden_size128] q self.w_q(x) # [b, s, 128] k self.w_k(x) v self.w_v(x) attn_scores torch.softmax(torch.bmm(q, k.transpose(1,2)) / (128**0.5), dim-1) out torch.bmm(attn_scores, v) # [b, s, 128] return self.layer_norm(out x) # 残差连接参数说明torch.bmm比torch.einsum快30%sqrt(128)是缩放因子防梯度爆炸——这些细节在论文里常被忽略但实测影响收敛速度。3.4 模型初始化策略为什么GRU权重要正交初始化train.py第78行调用init_gru_weights(model)def init_gru_weights(model): for name, param in model.named_parameters(): if gru.weight_ih_l0 in name: nn.init.orthogonal_(param.data) elif gru.weight_hh_l0 in name: nn.init.orthogonal_(param.data)原因GRU的weight_hh隐状态到隐状态若用默认xavier初始化易导致梯度消失。正交初始化使初始状态转移矩阵接近正交保障长序列梯度流动——在30秒3840帧输入下这能让loss下降曲线更平滑避免第10轮突然nan。4. 训练与评估focal loss怎么解决N1期漏判test.py输出的5个指标如何交叉验证模型缺陷4.1 focal_loss.pyγ2.0不是玄学是N1期召回率提升12%的实证参数Sleep-EDF中N1期占比仅7.8%传统CE loss会让模型忽略它。focal_loss.py实现class FocalLoss(nn.Module): def __init__(self, alpha1, gamma2, reductionmean): super().__init__() self.alpha alpha self.gamma gamma self.reduction reduction def forward(self, inputs, targets): ce_loss F.cross_entropy(inputs, targets, reductionnone) pt torch.exp(-ce_loss) # 预测概率 focal_weight (1-pt)**self.gamma loss self.alpha * focal_weight * ce_loss return loss.mean() if self.reductionmean else loss参数说明gamma2.0经网格搜索确定——γ1时N1 recall0.51γ2时升至0.63γ3时过拟合N2 precision↓5%。alpha设为类别频率倒数[1.0, 12.8, 1.5, 1.1, 1.8]对应W/N1/N2/N3/REM但代码中简化为标量1.0因focal_weight已足够调节。4.2 train.py里的wandb日志为什么必须记录per-class recalltrain.py第203行wandb.log()不仅记accuracy更关键的是# test.py第89行计算per-class recall recall_per_class recall_score(y_true, y_pred, averageNone) wandb.log({ val_recall_W: recall_per_class[0], val_recall_N1: recall_per_class[1], # 这才是医生关心的 val_recall_N2: recall_per_class[2], val_recall_N3: recall_per_class[3], val_recall_REM: recall_per_class[4], })原因整体accuracy达85%可能掩盖N1期漏判严重recall仅40%。医生需要知道“模型把多少N1误判为W”这直接影响临床决策。wandb面板里画出5条recall曲线一眼看出N1是否持续偏低——若连续3轮0.55立即检查focal_loss是否生效或数据增强是否过度。4.3 test.py输出的5个指标mf1、confusion matrix、precision/f1怎么组合诊断模型病灶运行python test.py --model_path models/model_GRU.pt后终端输出Accuracy: 0.842 Macro-F1: 0.783 Recall Confusion Matrix: [[0.89 0.03 0.05 0.01 0.02] # W类预测分布 [0.12 0.58 0.22 0.03 0.05] # N1类被错分为W/N2... [0.02 0.04 0.87 0.05 0.02] # N2类最准 [0.01 0.02 0.08 0.82 0.07] [0.03 0.04 0.05 0.12 0.76]] Precision Confusion Matrix: ...解读技巧看N1行[0.12, 0.58, 0.22, ...]表示N1期有58%被正确识别12%被当成W入睡初期易混淆22%当成N2分期边界模糊对比Recall vs Precision若N1 recall0.58但precision0.65说明模型保守宁可漏判也不乱判若precision0.45则说明它把太多W/N2标成N1激进mf10.783是5类recall的算术平均比accuracy更能反映均衡性——accuracy高但mf1低说明模型偏科严重。4.4 predict.py单样本推理怎么做到毫秒级输入格式与输出解码predict.py专为部署设计支持两种输入.npy文件预处理好的30秒单通道信号实时流式数据需自行接入硬件SDKpython predict.py --input eeg_signal.npy --model models/model_GRU.pt # 输出N2 (confidence: 0.92)关键优化模型设为model.eval()torch.no_grad()输入tensor用torch.float16半精度GPU支持时提速40%输出用torch.argmax(logits, dim1)而非softmax省去指数运算提示confidence不是softmax概率而是logits最大值——更稳定避免softmax数值溢出。5. 避坑指南单通道EEG睡眠分期的6个血泪经验从数据加载到模型崩溃全覆盖5.1 现象train.py报错RuntimeError: Expected 3-dimensional input for 3-dimensional weight原因preprocessing.py输出的.npy文件是2D[time_points, 1]但network.py的Conv1d要求3D输入[batch, channels, time]。dataset.py虽做了np.expand_dims(data, axis0)但若原始EDF通道名不匹配如存成EEG Fpz-Cz带空格raw.pick_channels([Fpz-Cz])失败data保持2D。解决检查preprocessing.py第33行raw.ch_names打印出来确认通道名精确匹配或强制data data.reshape(-1, 1)。5.2 现象test.py输出accuracy0.20mf10.0但loss曲线正常下降原因标签索引错位。Sleep-EDF的hypnogram标注中W0, N11, N22, N33, REM4但prepare_data.py若读取.hypnogram.edf时用错raw.annotations解析方式可能把N1映射成label5超出0-4范围导致cross_entropy返回nan但梯度仍传播。解决在prepare_data.py第120行添加断言assert np.all((labels 0) (labels 4))或用mne.events_from_annotations(raw)替代手动解析。5.3 现象wandb显示val_loss持续上升但train_loss下降原因validation set未shuffle。dataset.py中val_dataset默认shuffleFalse若val数据按时间顺序排列如前100段全是W期模型看到的val batch极度不均衡loss虚高。解决在train.py第155行创建val_loader时显式传参shuffleTrue或在dataset.py中val_dataset也设shuffle_seed。5.4 现象predict.py推理结果每次不同即使同一输入原因model.eval()未关闭dropout。network.py中nn.Dropout(p0.3)在eval模式下应自动关闭但若模型加载时未调用model.eval()或用了torch.compilePyTorch 2.0dropout可能未禁用。解决predict.py第45行强制model.eval()并在forward前加torch.set_grad_enabled(False)。5.5 现象web/server.py启动后浏览器访问500错误日志报CUDA out of memory原因Flask默认多进程每个worker加载一份model_GRU.pt到GPU10个worker吃光显存。解决server.py第22行改为单进程app.run(host0.0.0.0, port5000, threadedFalse, processes1)或改用CPU推理model.cpu()torch.set_default_device(cpu)。5.6 现象download_sleepedf.py下载中断重试后提示HTTP Error 403: Forbidden原因PhysioNet要求登录后下载但脚本未携带cookie。最新版PhysioNet API需OAuth token。解决手动下载Sleep-EDF SC subsethttps://physionet.org/content/sleep-edfx/1.0.0/解压到./Sleep-EDF/sc/跳过自动下载步骤——这是目前最稳方案比折腾token可靠。6. 进阶技巧用attention weights可视化“模型到底在看哪一秒”——三步提取时序注意力热力图6.1 修改network.py让SelfAttention模块返回attention weights原版SelfAttention只返回out需暴露attn_scores。修改network.py第158行def forward(self, x): q self.w_q(x) k self.w_k(x) v self.w_v(x) attn_scores torch.softmax(torch.bmm(q, k.transpose(1,2)) / (128**0.5), dim-1) out torch.bmm(attn_scores, v) out self.layer_norm(out x) return out, attn_scores # 返回attention weights同时在forward函数末尾第210行接收并传递# network.py 第210行 x, attn_weights self.attention(x) # 原来只接x return x, attn_weights # 向上传递6.2 predict_with_attn.py单样本推理并保存attention map新建predict_with_attn.py复用predict.py逻辑但捕获attn_weights# predict_with_attn.py model.eval() with torch.no_grad(): logits, attn_weights model(x.unsqueeze(0)) # x: [3840] pred torch.argmax(logits, dim1).item() # attn_weights: [1, 3840, 3840] → 取第一头单头平均 attn_map attn_weights[0].mean(dim0).cpu().numpy() # [3840] np.save(fattn_{pred}.npy, attn_map)执行python predict_with_attn.py --input data/sub01_001.npy --model models/model_GRU.pt6.3 用matplotlib绘制时序热力图定位模型决策依据用plot_attn.py可视化import numpy as np import matplotlib.pyplot as plt attn np.load(attn_N2.npy) # 形状[3840] time_axis np.arange(len(attn)) / 128 # 转秒 plt.figure(figsize(12,3)) plt.plot(time_axis, attn, linewidth1.2, colornavy) plt.fill_between(time_axis, 0, attn, alpha0.3, colorskyblue) plt.xlabel(Time (seconds)) plt.ylabel(Attention Weight) plt.title(Model Attention on 30-second EEG Segment (Predicted: N2)) plt.xlim(0, 30) plt.grid(True, alpha0.3) plt.savefig(attn_heatmap_N2.png, dpi300, bbox_inchestight) plt.show()效果图中峰值区域如8-12秒、22-26秒即模型认为最能代表N2期的纺锤波时段。对比原始EEG波形用mne.viz.plot_raw可验证模型是否真的聚焦在生理特征上——若峰值总在工频干扰段说明预处理没做好。6.4 临床验证技巧用混淆矩阵反推attention异常区当test.py显示N1→W误判率高混淆矩阵N1行第一列0.15提取所有N1样本的attn_map做均值# analysis_attn.py n1_samples [np.load(f) for f in glob(attn_N1_*.npy)] avg_attn_n1 np.stack(n1_samples).mean(axis0) # [3840] # 同理得avg_attn_w diff_map avg_attn_n1 - avg_attn_w # 正值区模型看N1比W多的地方若diff_map在0-5秒入睡初期持续为正说明模型正确捕捉了N1特征若在25-30秒快醒时段为正则可能把REM前期误当N1——这时该检查hypnogram标注一致性。从那以后我每次调参前都强制用predict_with_attn.py跑3个典型样本盯着热力图确认模型没学歪。哪怕mf1涨了0.01只要attention峰偏离生理区间我就回退版本。这习惯救了我两次毕设答辩——教授当场指着热力图问“为什么N3期关注点在0-2秒”我立刻答“因为delta波起始段”而不是背公式。希望帮到你。本文还有配套的精品资源点击获取
返回列表