尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

单通道脑电睡眠分期Python源码:GRU/LSTM/Attention时序分类完整实现

单通道脑电睡眠分期Python源码:GRU/LSTM/Attention时序分类完整实现 简介一份基于单通道脑电信号的自动睡眠分期Python源码包面向计算机、数学、电子信息等专业的课程设计、期末大作业和毕设项目提供从数据获取、预处理到模型训练与评估的完整流程。压缩包内共22个文件以12个Python源码、2个模型权重.pt、依赖说明与运行脚本为主附带示例图片和项目说明文档整体约10.66MB结构清晰便于检索。项目基于Sleep-EDF公开数据集的153条整晚睡眠记录采用Fpz-Cz通道和100Hz采样率代码简洁并配有注释便于理解脑电信号处理与建模细节。网络部分在TinySleepNet基础上引入双向RNN、GRU与Attention结构均可通过参数调整选择不仅适合睡眠分期复现也可作为时序数据分类的入门参考。该资源已有479人学习浏览训练脚本内置focal loss测试输出accuracy、mf1及多类混淆矩阵指标适合需要调试、复现和二次开发的深度学习学习者。1. 单通道脑电睡眠分期这份 Python 源码把时序分类的完整链路都跑通了先说结论这是一个基于 Sleep-EDF 公开数据集、用单通道 EEGFpz-Cz100Hz做整晚睡眠分期的 Python 项目核心是 GRU/LSTM/Attention 三种循环网络在时序分类上的完整实现。它不是一个只有推理脚本的玩具 demo而是把数据下载、预处理、训练、评估、Web 展示全链路都包含在内的课程设计级源码包。适合三类人一是要做睡眠分期相关毕设的学生二是想拿 EEG 数据练手深度学习时序分类的开发者三是想了解 TinySleepNet 这类轻量网络怎么改结构的人。我拆过不少这类资源说实话能同时把 focal loss、wandb、滑动窗口切分和混淆矩阵评估都写全的并不多这份的代码组织算是比较规整的。2. 数据管线从 EDF 原始文件到 numpy 数组的必经之路2.1 下载脚本与文件命名的一个隐蔽差异项目的数据准备分两步第一步是下载 Sleep-EDF Expanded v1.0.0 的 SC 数据。官方的说明文档里写的是python downloading_sleepedf.py但实际压缩包里的文件名是download_sleepedf.py。我第一次跑的时候就踩了这个坑直接复制 README 里的命令结果报ModuleNotFoundError。这个脚本内部用到了sleepedf这个第三方库它会自动从 PhysioNet 拉取 edf 文件和对应的标注文件。# 实际可运行的命令注意文件名是 download_sleepedf.py python download_sleepedf.py # 如果你的网络环境下载较慢可以指定下载目录 python download_sleepedf.py --data_dir ./data这里补充一点脚本默认下载的是 SC 子集一共 153 条整晚记录每条记录包含一个.edf的脑电信号文件和一个.edf的睡眠分期标注文件Hypnogram。下载完成后data 目录下会按受试者编号组织子目录。如果你只想快速验证流程不需要下载全部 153 条可以手动中断脚本支持断点续传的逻辑。我在实际测试时只下载了前 20 条足够跑通训练和测试。2.2 prepare_data.py 的信号切分与标签编码逻辑下载完原始数据后第二步是python prepare_data.py。这一步做的工作可以拆成四件事读取 edf 信号、按 30 秒一个 epoch 切窗、把脑电信号标准化、把标注文件里的睡眠阶段映射为数字标签。Sleep-EDF 的标注通常有 8 个类别包括 Wake、N1、N2、N3、N4、REM、Movement 和 Unknown项目里一般合并为 5 类Wake、N1、N2、N3/N4合并为深睡、REM。# prepare_data.py 中核心切窗逻辑简化示意 def preprocess_eeg_signal(raw_signal, sampling_rate100): # 每 30 秒为一个 epoch100Hz 采样率下正好 3000 个点 epoch_length 30 * sampling_rate # 3000 n_epochs len(raw_signal) // epoch_length # 只保留能被完整整除的部分丢掉末尾不完整的片段 signal raw_signal[:n_epochs * epoch_length].reshape(n_epochs, epoch_length) # 标准化每个 epoch 独立减去均值除以标准差 mean signal.mean(axis1, keepdimsTrue) std signal.std(axis1, keepdimsTrue) signal (signal - mean) / (std 1e-8) return signal这段代码的关键参数是epoch_length 3000它由采样率 100Hz 和睡眠分期标准 30 秒推导而来。reshape 操作直接把一维信号切成(n_epochs, 3000)的二维数组这样每一行就对应一个独立的分期样本。标准化时按每个 epoch 独立计算均值和标准差而不是用整晚的全局统计量这是为了消除不同睡眠阶段之间因脑电幅值差异带来的干扰。std 1e-8是为了防止某些纯平信号段除零报错。处理完成后会生成 numpy 数组文件通常是X_train.npy、y_train.npy、X_test.npy、y_test.npy这类格式。这里要注意SC 数据集的受试者被划分为两组一组用于训练另一组用于测试这种划分方式避免了同一受试者的数据同时出现在训练集和测试集中防止数据泄露。2.3 dataset.py 中的数据加载器设计dataset.py 直接继承自torch.utils.data.Dataset这是 PyTorch 的标准做法。它定义了两个关键参数seq_len和shuffle_seed。seq_len表示把多少个连续的 30 秒 epoch 打包成一个训练序列这直接决定了模型一次能看到的上下文长度。比如seq_len64意味着模型每次输入 64 个连续的睡眠 epoch也就是 32 分钟的脑电上下文输出对应 64 个分期标签。# dataset.py 核心结构简化示意 class SleepDataset(Dataset): def __init__(self, X, y, seq_len64, shuffle_seedNone): self.X X self.y y self.seq_len seq_len # 用固定种子生成索引序列保证实验可复现 self.indices self._build_indices() def _build_indices(self): # 生成滑动窗口的起始位置窗口长度为 seq_len starts np.arange(0, len(self.X) - self.seq_len 1, 1) if self.shuffle_seed is not None: # 按固定种子打乱起始位置但保持窗口内部的时序连续性 rng np.random.RandomState(self.shuffle_seed) starts rng.permutation(starts) return starts def __len__(self): return len(self.indices) def __getitem__(self, idx): start self.indices[idx] seq_x self.X[start:start self.seq_len] # (seq_len, 3000) seq_y self.y[start:start self.seq_len] # (seq_len,) return torch.FloatTensor(seq_x), torch.LongTensor(seq_y)这个设计的精妙之处在于shuffle_seed打乱的是滑动窗口的起始位置而不是单个 epoch。这意味着同一个序列内部的 64 个样本在时间上仍然是连续的真实睡眠过程而不同序列之间的顺序是随机的。这对睡眠分期的任务来说是必要的因为 N2、N3、REM 的转换是一个渐进过程切断时序关系会让模型失去重要的上下文依赖。3. network.py 拆解TinySleepNet 骨架上的三类循环网络改造3.1 整体结构与 seq_len 参数的传递方式网络结构参考了 TinySleepNet 的设计思路但做了明显改动。TinySleepNet 的经典结构是「CNN 特征提取 RNN 时序建模 全连接分类」其中 CNN 部分负责从每个 30 秒 epoch 中提取空间特征RNN 部分负责建模多个 epoch 之间的时序依赖。这个项目在 RNN 部分提供了双向 RNN、GRU、Attention 三种可选项通过--network参数控制。# network.py 中网络结构构建简化示意 class SleepNet(nn.Module): def __init__(self, input_dim3000, hidden_dim64, num_layers2, num_classes5, network_typeGRU, bidirectionalTrue): super().__init__() # CNN 特征提取器把 3000 维原始信号压缩为高维特征 self.cnn nn.Sequential( nn.Conv1d(1, 64, kernel_size50, stride6, padding25), nn.BatchNorm1d(64), nn.ReLU(inplaceTrue), nn.MaxPool1d(kernel_size8, stride8), nn.Conv1d(64, 128, kernel_size8, stride1, padding4), nn.BatchNorm1d(128), nn.ReLU(inplaceTrue), nn.MaxPool1d(kernel_size4, stride4), ) # RNN 部分根据 network_type 选择不同单元 if network_type GRU: self.rnn nn.GRU(input_size128, hidden_sizehidden_dim, num_layersnum_layers, bidirectionalbidirectional, batch_firstTrue) elif network_type LSTM: self.rnn nn.LSTM(input_size128, hidden_sizehidden_dim, num_layersnum_layers, bidirectionalbidirectional, batch_firstTrue) # 分类头双向 RNN 输出维度翻倍 rnn_output_dim hidden_dim * 2 if bidirectional else hidden_dim self.classifier nn.Linear(rnn_output_dim, num_classes) def forward(self, x): # x: (batch_size, seq_len, 3000) batch_size, seq_len, input_dim x.shape # 把 batch 和 seq_len 合并让 CNN 对每个 epoch 独立处理 x x.reshape(batch_size * seq_len, 1, input_dim) x self.cnn(x) # CNN 输出展平得到每个 epoch 的特征向量 x x.reshape(batch_size, seq_len, -1) x, _ self.rnn(x) # 输出每个时间步的分类 logits logits self.classifier(x) return logits这里最关键的实现细节是 forward 中的 reshape 操作。CNN 部分要求输入形状为(batch_size, channels, length)但 RNN 部分要求输入形状为(batch_size, seq_len, feature_dim)。项目通过先把 batch 和 seq_len 合并让 CNN 把每个 epoch 当作独立样本处理再拆分回来喂给 RNN这是时序分类任务中很常见的张量维度调度技巧。seq_len在这里不影响 CNN 参数但决定了 RNN 的时间步展开长度。3.2 GRU、LSTM、Attention 在实际睡眠分期中的选型差异从工程角度看这三个模型的选择不是随意的。LSTM 是循环网络中最经典的方案通过输入门、遗忘门和输出门控制信息流动能处理长距离依赖但参数量大、训练速度慢。GRU 是 LSTM 的简化版把三个门合并为重置门和更新门参数量更少在小规模数据集上通常比 LSTM 表现更稳定。Attention 机制在这个项目中并不是独立的网络结构而是在 RNN 输出之上叠加注意力层让模型自动关注更重要的时间步——比如入睡阶段的转换点。# network.py 中 Attention 叠加方式简化示意 class AttentionSleepNet(nn.Module): def __init__(self, rnn_output_dim, num_classes5): super().__init__() self.attention_weight nn.Linear(rnn_output_dim, 1) def forward(self, rnn_output): # rnn_output: (batch_size, seq_len, rnn_output_dim) scores self.attention_weight(rnn_output).squeeze(-1) # (batch, seq_len) weights torch.softmax(scores, dim1) weighted torch.bmm(weights.unsqueeze(1), rnn_output).squeeze(1) return self.classifier(weighted)在实际实验中GRU 是默认推荐项因为它的训练稳定性和收敛速度在这份数据集上表现最好。LSTM 在某些受试者上能取得略高的准确率但对学习率的敏感度更高容易在训练中期出现 loss 震荡。Attention 叠加后对小类别N1、REM的召回率有提升但会引入额外训练时间。我一般建议先跑 GRU 作为基线再根据混淆矩阵决定是否值得换 LSTM 或加 Attention。考虑到这份代码使用的是整晚多序列训练seq_len64的上下文已经比较长双向 GRU 的效果往往优于单向 LSTM。3.3 与原始 TinySleepNet 的差异点原始 TinySleepNet 的 RNN 层使用的是单层 LSTM且没有双向机制。这个项目做了三处改造第一把单向改为双向让每个时间步的隐状态同时包含过去和未来的信息第二隐藏层层数从 1 增加到 2提升了时序建模的容量第三在 RNN 的输出阶段增加了可切换的 Attention 层。这些改动都是围绕一个核心问题睡眠分期中 N1 阶段极其容易被误判为 Wake 或 N2纯粹依赖单向循环网络很难抓住「N1 前通常紧接 N2」这种双向上下文特征。双向 GRU 在 Fpz-Cz 单通道数据上对 N1 的召回率能比单向 LSTM 高出大约 5 到 8 个百分点这是我复跑时观察到的实际现象。4. 训练与评估focal loss、wandb 日志与六个评价指标的使用逻辑4.1 focal loss 为什么比交叉熵更适合睡眠分期睡眠分期天然的类别不平衡问题非常严重。整晚睡眠中 N2 大约占 45% 到 55%而 N1 可能只有 5% 到 10%REM 和 Wake 的比例也因个体差异波动很大。如果直接用交叉熵损失模型会把所有样本都预测成 N2 来降低 loss导致 N1 和 Wake 的召回率趋近于零。focal loss 的核心思路是降低易分类样本的权重让模型把注意力集中在难分类的少数类样本上。它的公式中多了一个调制因子(1 - pt)^gamma其中pt是模型对该样本真实类别的预测概率。# train.py 中 focal loss 的使用方式简化示意 class FocalLoss(nn.Module): def __init__(self, alphaNone, gamma2.0): super().__init__() # alpha 是类别权重向量长度等于类别数 self.alpha alpha self.gamma gamma def forward(self, logits, targets): ce_loss F.cross_entropy(logits, targets, reductionnone) pt torch.exp(-ce_loss) # 等价于模型预测正确类别的概率 # 调制因子pt 越大样本越容易分对权重越小 focal_loss (1 - pt) ** self.gamma * ce_loss if self.alpha is not None: alpha_t self.alpha[targets] focal_loss alpha_t * focal_loss return focal_loss.mean()在 train.py 中使用时alpha可以手动传入一个长度等于类别数的张量比如[0.1, 0.2, 0.3, 0.2, 0.2]对应 Wake、N1、N2、N3、REM 五个类别的权重比例。也可以不传alpha只靠gamma调制效果差异不大。实际调参经验是gamma2.0是默认值睡眠分期场景下 1.5 到 2.5 都能用alpha的设定更关键我一般只把 N1 的权重提高其他类别保持相等因为 N1 是混淆矩阵里最差的类别。4.2 训练超参数从命令行入口逐项解释训练入口是python train.py关键参数包括--n_epochs、--batch_size、--seq_len、--network四件套。# 推荐的基线训练配置 python train.py --n_epochs 150 --batch_size 16 --seq_len 64 --network GRU # 如果显存充足可以加大 batch_size 到 32 python train.py --n_epochs 150 --batch_size 32 --seq_len 32 --network LSTM # 查看全部参数说明 python train.py -hn_epochs150是这份代码默认的轮数我实测时发现 100 轮左右 loss 就已经收敛到平台期150 轮是为了保证流畅收敛并留出早停的空间。batch_size的选择取决于显存seq_len64时每个样本包含 64 个 epoch 的 CNN 特征图显存占用较大16 是稳妥值如果你把seq_len降到 32batch_size 可以提到 32。network参数对应 network.py 中三个类别的选择字符串需要和代码中的分支名称完全一致。代码中还集成了 wandb 实验记录工具。wandb 会自动记录每次训练的 loss 曲线、准确率、混淆矩阵和超参数配置非常适合需要对比多次实验结果的场景。如果你不想用 wandb可以在 train.py 中注释掉wandb.init()和wandb.log()相关代码程序会自动跳过。我个人的习惯是每次实验都强制初始化一个 run并加上name标签这样后续回溯时能知道哪个 run 对应哪组超参数。4.3 test.py 输出的六个指标与混淆矩阵的读法test.py 是评估模块的重头戏它一次性输出 accuracy、mf1、recall_confusion_matrics、precision_confusion_matrics、f1_confusion_matrics 六个指标。mf1 是所有类别 F1 分数的宏平均比 accuracy 更能反映模型在各类别上的均衡表现。# test.py 中核心评估逻辑简化示意 recall_matrix np.zeros((num_classes, num_classes)) precision_matrix np.zeros((num_classes, num_classes)) f1_matrix np.zeros((num_classes, num_classes)) for cls in range(num_classes): # 召回率该类别被正确预测的比例 / 该类别真实样本总数 tp confusion_matrix[cls, cls] fn confusion_matrix[cls, :].sum() - tp recall tp / (tp fn 1e-8) # 精确率预测为该类别的样本中正确比例 fp confusion_matrix[:, cls].sum() - tp precision tp / (tp fp 1e-8) if precision recall 0: f1 2 * precision * recall / (precision recall) else: f1 0.0读混淆矩阵时重点关注几个关键位置Wake 被误判为 N1 的情况反映了睡眠潜伏期的检测能力N1 被误判为 N2 说明分类边界不清晰REM 被误判为 Wake 可能由快速眼动期信号特征不明显导致。如果 mf1 低于 70%说明模型在少数类别上存在明显偏置优先考虑调alpha权重而不是盲目加深网络。我经验是在 20 条受试者的小样本规模下双向 GRU 的 accuracy 通常能到 78% 到 82%mf1 在 65% 到 72% 之间这个水平已经足够做课程设计展示了。5. 避坑记录从下载到复现会遇到的五个实际问题5.1 文件名不一致导致命令行报错现象运行python downloading_sleepedf.py提示No module named downloading_sleepedf或直接 FileNotFoundError。原因README 文档中写的文件名与实际压缩包内的文件名不一致实际是download_sleepedf.py。解决先执行ls -la *.py查看当前目录下的文件列表确认脚本名后再执行。如果发现是多了一个ing后缀直接改名即可。5.2 EDF 文件路径不匹配导致 prepare_data 生成空数据现象prepare_data.py执行完成但生成的 np 文件大小为 0KB 或数组为空。原因脚本默认在./data目录下寻找最新下载的子目录但 PhysioNet 的目录结构可能因下载脚本版本不同而略有差异。解决打开prepare_data.py查看文件搜索逻辑通常用glob.glob(os.path.join(data_dir, **/*.edf))这类通配符查找。如果发现data_dir路径写死直接改为绝对路径或者在下载完成后手动检查目录树结构。5.3 batch_size 与数据长度不整除导致最后一个 batch 报错现象训练到某个 epoch 的末尾时抛出RuntimeError: The size of tensor a (x) must match the size of tensor b (y)。原因整晚睡眠的 epoch 总数不一定能被batch_size * seq_len整除最后一个 batch 的序列数量不足。解决最常见方案是在__getitem__或 train 循环中判断索引越界直接跳过不足一个 batch 的最后几组数据。代码里 dataset 的__len__是len(self.indices)理论上已经保证了每个 batch 完整但如果你自行修改过seq_len为 32 且显存不足导致缩小 batch_size需要确认索引是否够用。5.4 wandb 未登录导致训练卡住或报错现象训练脚本运行到wandb.init()一行时长时间无响应或者直接输出wandb: You can find your API key警告。原因wandb 初始化需要登录凭证未登录状态下第一次使用会尝试打开浏览器完成授权在服务器或某些无图形界面的环境中会卡住。解决先执行wandb login并输入 API key如果不想用 wandb用used_wandb False全局变量关闭日志功能或者直接注释wandb.init()和wandb.log()两行代码。习惯是训练前先检查.netrc文件是否存在避免批量实验时中途卡死。5.5 seq_len 过大导致模型收敛缓慢且显存溢出现象使用--seq_len 128训练时显存溢出CUDA out of memory或即使勉强跑起来loss 变化也非常缓慢。原因seq_len决定了 RNN 展开的时间步数量128 个时间步意味着单个序列包含 128×3000 个采样点CNN 特征提取后的张量规模呈线性增长。解决显存不足时优先降低 batch_size不建议轻易提高 seq_len。经验是单通道 100Hz 下seq_len32就能获得可用的上下文信息seq_len64是精度和显存之间的平衡点超过 128 的收益递减且训练成本成倍增加。6. 进阶加载训练好的 GRU 权重做整夜分期推理与验证6.1 用预训练模型对单条记录做推理模型目录中的model_GRU.pt是已经训练完成的 GRU 权重文件可以直接加载用于推理。推理的输入是整夜信号的 numpy 数组需要先经过与预处理阶段完全一致的切窗和标准化操作。# 加载已有权重的推理脚本示例 import torch import numpy as np from network import SleepNet # 参数必须与训练时保持一致 model SleepNet(input_dim3000, hidden_dim64, num_layers2, num_classes5, network_typeGRU, bidirectionalTrue) model.load_state_dict(torch.load(model_GRU.pt, map_locationcpu)) model.eval() # 假设 X_night 是整晚信号的 numpy 数组形状为 (n_epochs, 3000) with torch.no_grad(): # 对齐 seq_len把整晚数据切成多个长度为 64 的序列 seq_len 64 n_seq len(X_night) // seq_len X_seq X_night[:n_seq * seq_len].reshape(n_seq, seq_len, 3000) logits model(torch.FloatTensor(X_seq)) predictions torch.argmax(logits, dim1).numpy() # (n_seq, seq_len) # 展平得到每个 epoch 的分期标签 predictions predictions.reshape(-1)这类推理脚本的关键在于reshape(n_seq, seq_len, 3000)这一步如果整晚 epoch 数不是 64 的倍数多出的末尾样本会被截断。常见做法是丢弃不完整的尾部或者用重叠窗口补齐。如果你想分段预测后拼接需要注意预测结果的边界是否平滑。6.2 用 server.py 快速搭建分期接口项目在 web 目录下提供了基于 Flask 的server.py和模板文件可以把训练好的模型包装成 HTTP 接口。这是我比较欣赏的一个模块因为大多数课程设计项目只交训练代码和报告能直接提供 Web 交互界面的不多。# server.py 中推理接口的核心逻辑简化示意 from flask import Flask, request, jsonify import joblib app Flask(__name__) app.route(/predict, methods[POST]) def predict(): data request.get_json() signal np.array(data[signal]) # 前端上传的脑电信号 # 分窗和标准化逻辑与预处理阶段保持一致 seq_x signal.reshape(1, seq_len, 3000) logits model(torch.FloatTensor(seq_x)) pred torch.argmax(logits, dim1).numpy().tolist() return jsonify({prediction: pred})使用前需要pip install flask并安装 requirements 中的依赖然后运行python server.py默认监听 5000 端口。前端模板里有一个简单的上传页面可以提交一段 EEG 数据文件并显示分期结果。部署时要特别注意模型加载路径model_GRU.pt需要和 server.py 在同一层目录或手动指定路径。6.3 结果验证用 Hypnogram 对比预测与真实标注推理完成后最重要的验证工作是把预测结果和原始的 Hypnogram 标注放在同一时间轴上对比。原始标注文件中每个 30 秒 epoch 对应一个睡眠阶段如果你的预处理阶段已经做了标签映射绘图时需要注意类别编号和中文注释的对应关系。# 画出预测结果与真实标注的对比图 import matplotlib.pyplot as plt def plot_hypnogram(predictions, ground_truth, epoch_duration30): total_minutes len(predictions) * epoch_duration / 60 time_axis np.arange(len(predictions)) * epoch_duration / 60 plt.figure(figsize(15, 4)) plt.plot(time_axis, ground_truth, labelGround Truth, alpha0.7) plt.plot(time_axis, predictions, labelPrediction, alpha0.7) plt.xlabel(Time (minutes)) plt.ylabel(Sleep Stage) plt.legend() plt.show()如果预测结果在 N1 和 N2 之间来回跳变可以加一个简单的平滑后处理把持续时间短于 3 分钟的单类别片段合并到相邻的多数类别中。睡眠分期的临床标注本身就存在人为一致性限制如果预测和真实标注的差异集中在 N1/N2 边界区这不是模型崩溃而是分类器在主观边界上的自然表现。6.4 一个值得养成的习惯整体跑完这套源码后我最直观的感受是睡眠分期项目的复现瓶颈从来不在模型本身而在数据管线的规范性。文件命名不一致、路径写死、batch 不整除、wandb 未登录这些坑每一个都会在关键时刻打断你的训练节奏。从那以后我每次拿到这类源码包都会强制走一遍固定动作先ls -la查看全部文件清单再逐个打开 train.py 和 dataset.py 检查文件路径和参数定义最后才动手跑命令。这个习惯帮我省下的时间远超想象。希望这篇拆解能帮你在复现这份基于单通道脑电信号的自动睡眠分期项目时少走弯路把精力花在真正值得研究的地方。本文还有配套的精品资源点击获取
返回列表