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

资讯详情

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

PyTorch实现单通道EEG睡眠分期:从预处理到LSTM模型全流程解析

PyTorch实现单通道EEG睡眠分期:从预处理到LSTM模型全流程解析 简介基于PyTorch框架的单通道脑电睡眠分期项目完整源码面向计算机相关专业学生适用于毕业设计、课程设计或期末大作业场景也可作为深度学习初学者的实战练习素材。项目整合了模型定义、数据预处理、训练与评估、PyTorch Lightning封装等核心环节并附有依赖说明与使用文档代码结构清晰有助于掌握睡眠分期的完整实现流程。压缩包共包含19个文件以Python源码为主辅以项目配置、依赖清单和说明文档整体大小仅20KB下载运行十分轻便。这一项目目前已有174人学习下载是经导师指导并获评审99分的高分项目参考价值较高。对正在准备毕设或需要项目实战的学生而言这份代码既能提供可运行的基础实现又展示了模型搜索与工程化封装等进阶思路是一条值得深入研读的完整项目路径。1. 单通道EEG睡眠分期为什么PyTorch比手工特征更适合课程设计睡眠分期是睡眠医学与可穿戴监测共同依赖的基础任务。传统流程由技师按 AASM 规则把整夜 PSG 切成 30 秒 epoch逐个标注 W/N1/N2/N3/REM一夜数据耗时两小时以上不同标注者的一致性通常只有 70% 到 82%。改用深度学习后输入单通道脑电片段即可输出分期标签整个流程端到端完成模型还能部署到只有一导 EEG 采集能力的便携设备上。这套 PyTorch 实现的单通道睡眠分期项目preprocess.py、dataset.py、model.py、lightning_wrapper.py、train.py、benchmark.py 等文件覆盖了从原始 EDF 到评估指标的全链路源码直接可运行且评审达到 99 分。对正在做课程设计和期末大作业的同学它可以作为逐文件精读的完整样本跑通之后再换自己的数据做迁移实验比从零搭一个训练管线省去大量调错时间。2. 预处理链路preprocess.py 与 dataset.py 如何把原始脑电变成训练样本睡眠分期模型的输入是固定长度的脑电片段也就是 30 秒一个 epoch。以 Sleep-EDF 这类公开数据集为例原始记录是 EDF 格式的多导信号文件里面除了 EEG 还有 EOG、EMG不同记录的采样率也不统一。preprocess.py 的首要任务就是从 EDF 里抽出目标单通道重采样到统一频率再完成滤波、分段最后交给 dataset.py 按窗口组织成 PyTorch Dataset。2.1 单通道抽取与重采样pyedflib 的读取方式常见做法是用 pyedflib 直接读取 EDF 文件按通道名定位目标导联再用线性插值把信号重采样到 100Hz。这样后面所有 epoch 的长度都是固定的 3000 个采样点模型输入维度不会因为记录设备的采样率差异而变化训练和推理阶段的数据形状保持一致。import numpy as np import pyedflib def load_single_channel(edf_path, channel_nameFpz-Cz, target_fs100): reader pyedflib.EdfReader(edf_path) ch_index reader.getSignalLabels().index(channel_name) raw reader.readSignal(ch_index) fs reader.getSampleFrequency(ch_index) reader.close() if fs ! target_fs: n_target int(len(raw) * target_fs / fs) raw np.interp(np.linspace(0, len(raw) - 1, n_target), np.arange(len(raw)), raw).astype(np.float32) return raw, target_fs这里np.interp做的是分段线性插值对 EEG 这种连续信号足够用不需要引入更重的重采样库。readSignal返回的是原始微伏量纲的浮点数组后续归一化会把它拉回标准范围所以读取阶段不需要做幅值换算。通道名用getSignalLabels().index()定位避免硬编码通道序号导致换数据集时报错。2.2 带通滤波与 30 秒分段scipy.signal 的实用组合EEG 的有效频率成分集中在 0.5Hz 到 30Hz 之间30Hz 以上主要是肌电伪迹和工频干扰。preprocess.py 里用四阶 Butterworth 零相位滤波sosfiltfilt相比直接filtfilt在高阶滤波器下数值稳定性更好这是实际跑数据时对比之后换过来的写法低阶butter配合sosfiltfilt能避免滤波结果出现边缘振荡。from scipy.signal import butter, sosfiltfilt def bandpass(signal, low0.5, high30, fs100): sos butter(4, [low, high], btypebandpass, fsfs, outputsos) return sosfiltfilt(sos, signal) def to_epochs(signal, fs100, epoch_sec30): win epoch_sec * fs n_epochs len(signal) // win return signal[: n_epochs * win].reshape(n_epochs, win)butter的阶数定为 4 的原因很直接阶数太低过渡带太宽会保留 30Hz 附近的肌肉噪声阶数太高会引入相位畸变虽然sosfiltfilt做了零相位补偿但计算量明显上升。reshape之前先把信号截断到 epoch 长度的整数倍避免最后一个不完整片段混进训练集这也是睡眠分期预处理里最容易被忽略的细节。提示先重采样再滤波。如果顺序反了重采样会在滤波后信号的高频截止边界附近引入新的混叠噪声影响 N3 期慢波的波形判断。2.3 Dataset 封装归一化与上下文窗口的最终形态预处理最后一步是标签映射和归一化。AASM 标准下 W/N1/N2/N3/REM 映射为 0 到 4 的整数如果项目里沿用 RK 旧标准N3 和 N4 需要合并。归一化按照每个 epoch 独立做 z-score而不是按整夜记录做因为不同受试者脑电幅值差异很大按记录整体归一化会把个体差异当作特征学进去跨受试者测试时泛化会变差。import numpy as np import torch from torch.utils.data import Dataset class SleepDataset(Dataset): def __init__(self, epochs, labels, context3): self.epochs epochs # (N, 3000) self.labels labels # (N,) self.context context # 上下文窗口内 epoch 数量 def __getitem__(self, idx): half self.context // 2 start, end idx - half, idx half 1 win self.epochs[max(0, start):end] if len(win) self.context: # 边界处补零 padded np.zeros((self.context, win.shape[1]), dtypenp.float32) pad_front max(0, -start) padded[pad_front:pad_front len(win)] win win padded x (win - win.mean(axis1, keepdimsTrue)) / \ (win.std(axis1, keepdimsTrue) 1e-8) return torch.from_numpy(x[:, None, :]), torch.tensor(self.labels[idx], dtypetorch.long) def __len__(self): return len(self.labels)context3表示每个样本包含当前 epoch 及其前后各一个 epoch标签取中间位置的类别。增加的通道维度[1, 3000]对应 Conv1d 的输入布局[B, C, L]。1e-8防止某个 epoch 信号过于平坦导致除零边界补零比边缘复制更省事LSTM 也能通过零向量感知序列边界。表预处理关键参数参数取值影响重采样频率100Hz决定每个 epoch 采样点数影响卷积核感受野带通范围0.5–30Hz滤除基线漂移与高频肌电保留主要脑电频段epoch 长度30 秒AASM 标准评分窗口不可随意修改归一化方式按 epoch z-score消除个体幅值差异提升跨受试者泛化3. EmbedSleepNet 网络拆解model.py 里嵌入层、双向LSTM与分类头的三段式设计model.py 实现的核心网络叫 EmbedSleepNet从命名可以看出设计思路先用可学习的卷积嵌入层把原始波形压缩成高维特征再用双向 LSTM 捕捉相邻 epoch 的时间依赖最后用分类头输出五类概率。这个三段式结构与 DeepSleepNet、SleepEEGNet 等经典工作一脉相承但嵌入层用了更小的卷积核和全局池化训练参数量控制得更好。3.1 一维卷积嵌入层把 3000 点波形压缩成 128 维特征输入是[B, T, 1, 3000]的上下文窗口嵌入层对每个 epoch 独立做一维卷积。第一层用较大的卷积核覆盖约 90ms 的波形相当于先提取一次局部的波形形态第二层 stride 为 2 扩大感受野。两层卷积后接全局平均池化把每个 epoch 压缩成一个 128 维向量LSTM 再在这个向量序列上建模跨 epoch 依赖计算量比直接在原始信号上接循环网络少一个数量级。import torch.nn as nn import torch.nn.functional as F class EEGEmbedding(nn.Module): def __init__(self, in_ch1, feat_dim128, kernel9): super().__init__() self.conv1 nn.Conv1d(in_ch, 64, kernel, stride2, padding4) self.bn1 nn.BatchNorm1d(64) self.conv2 nn.Conv1d(64, feat_dim, kernel, stride2, padding4) self.bn2 nn.BatchNorm1d(feat_dim) def forward(self, x): x F.relu(self.bn1(self.conv1(x))) x F.relu(self.bn2(self.conv2(x))) return F.adaptive_avg_pool1d(x, 1).squeeze(-1) # [B, feat_dim]padding 设为kernel // 2并在 stride2 的情况下保证时间维度能整除3000 点经过两层卷积后变成 750 点再池化成向量。BatchNorm 放在卷积之后激活函数之前这是 ResNet 时代验证过更稳定的顺序。全局平均池化会损失单 epoch 内部的时序细节但对睡眠分期影响不大因为纺锤波、K 复合波这类瞬态结构在卷积层的局部感受野里已经捕获池化只损失它们在时间轴上的精确位置。3.2 双向 LSTM 捕捉跨 epoch 上下文睡眠是一个连续过程N2 的纺锤波和 N3 的慢波往往跨越多个 epoch 边界。如果只对单个 epoch 做分类模型无法利用前后文信息N1 和 REM 这类波形相似的阶段特别容易混淆。SeqContextEncoder 把上下文窗口内 T 个 epoch 的嵌入向量拼接成序列送进双向 LSTM再取中间位置对应的输出作为当前 epoch 的上下文特征。class SeqContextEncoder(nn.Module): def __init__(self, feat_dim128, hidden_dim64, num_layers1): super().__init__() self.lstm nn.LSTM(feat_dim, hidden_dim, num_layersnum_layers, bidirectionalTrue, batch_firstTrue) def forward(self, x): # x: [B, T, feat_dim]T 是上下文窗口内的 epoch 数量 out, _ self.lstm(x) mid x.size(1) // 2 return out[:, mid, :] # 取当前 epoch 位置[B, 2*hidden_dim]双向 LSTM 的输出维度是hidden_dim * 2前向和后向隐状态拼接起来当前 epoch 的分类同时利用前文和后文信息。窗口 T 取 3 到 5太小上下文不足太大边界效应明显且训练时需要更多补零。num_layers默认设为 1因为这个任务的特征序列本身不长加深 LSTM 层数带来的收益远不如加大嵌入层宽度这一点是消融实验里反复验证过的。3.3 分类头与类别权重N1 占比低不能靠网络自己学五个睡眠阶段里 N1 通常只占 5% 左右而 N2 往往超过 45%。如果不做处理网络把所有样本预测成 N2 也能获得很高的整体准确率但这样的模型没有医学价值。解决方法是给 CrossEntropyLoss 传入按类别频率反比的权重让少数类在反向传播中获得更大的梯度贡献。import numpy as np def make_class_weight(labels, num_classes5): counts np.bincount(labels, minlengthnum_classes).astype(np.float32) weights counts.sum() / (num_classes * counts 1e-6) return torch.tensor(weights, dtypetorch.float32) class SleepClassifier(nn.Module): def __init__(self, in_dim128, num_classes5): super().__init__() self.head nn.Sequential( nn.Linear(in_dim, in_dim // 2), nn.ReLU(), nn.Dropout(0.5), nn.Linear(in_dim // 2, num_classes), ) def forward(self, x): return self.head(x)counts.sum() / (num_classes * counts)比直接取倒数更稳权重之和保持稳定不会因为某类占比极低导致权重爆炸。分类头的 Dropout 设为 0.5睡眠分期数据集规模通常在几万到几十万 epoch 量级过拟合风险比一般图像任务更高丢弃率太低时验证 loss 与训练 loss 的差距会明显拉大。完整前向过程如下每个 epoch 波形经卷积嵌入得到 128 维向量T 个向量拼成[B, T, 128]双向 LSTM 输出中间位置上下文特征最后经两层 MLP 得到 5 类 logits。这个结构参数量远小于直接堆全连接层的方案单卡 GPU 几小时就能完成训练。class EmbedSleepNet(nn.Module): def __init__(self, feat_dim128, hidden_dim64, num_classes5): super().__init__() self.embedding EEGEmbedding(feat_dim) self.seq_encoder SeqContextEncoder(feat_dim, hidden_dim) self.classifier SleepClassifier(hidden_dim * 2, num_classes) def forward(self, x): # x: [B, T, 1, 3000] B, T x.size(0), x.size(1) embedded self.embedding(x.view(B * T, 1, -1)) embedded embedded.view(B, T, -1) context self.seq_encoder(embedded) return self.classifier(context)4. PyTorch Lightning 训练管线train.py 与 lightning_wrapper.py 的工程化封装原生 PyTorch 的训练循环本身不复杂但睡眠分期实验需要反复调整学习率、加早停、做断点续训和超参数搜索这些逻辑堆在同一个文件里很快失控。项目用 PyTorch Lightning 做了一层封装把网络、损失函数、优化器和评估逻辑全部装进 LightningModule 子类训练循环交给框架调度实验代码与数据加载彻底解耦。4.1 LightningModule 子类训练步与验证步的定义方式lightning_wrapper.py 的核心是SleepStageLightning它把 loss 计算、准确率计算、优化器配置写进框架约定的钩子方法。training_step只需返回 loss框架负责反向传播和参数更新validation_step在每轮验证时自动调用返回的张量会累积到 epoch 结束时统一计算。import pytorch_lightning as pl class SleepStageLightning(pl.LightningModule): def __init__(self, model, class_weight, lr1e-3): super().__init__() self.model model self.criterion nn.CrossEntropyLoss(weightclass_weight) self.lr lr def training_step(self, batch, batch_idx): x, y batch logits self.model(x) loss self.criterion(logits, y) self.log(train_loss, loss, prog_barTrue) return loss def validation_step(self, batch, batch_idx): x, y batch logits self.model(x) loss self.criterion(logits, y) acc (logits.argmax(dim1) y).float().mean() self.log(val_loss, loss, prog_barTrue) self.log(val_acc, acc, prog_barTrue)self.log的第一个参数是指标名prog_barTrue会把指标显示在进度条上。验证阶段框架自动启用torch.no_grad()并清理梯度不需要手动管理。val_acc 只是训练过程的粗粒度监控最终评估交给 benchmark.py 计算 macro-F1因为整体准确率会被 N2 主导看不出少数类的真实表现。4.2 优化器与学习率调度ReduceLROnPlateau 的实践配置对这个任务我不用余弦退火而是用ReduceLROnPlateau。睡眠分期验证 loss 曲线经常出现平台期按验证指标衰减学习率比按 step 线性衰减更容易跳出不显著的局部极小点尤其当数据集类别分布不均衡时。def configure_optimizers(self): optimizer torch.optim.Adam(self.parameters(), lrself.lr, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.5, patience5, min_lr1e-6 ) return { optimizer: optimizer, lr_scheduler: {scheduler: scheduler, monitor: val_loss}, }patience5表示验证 loss 连续 5 个 epoch 不下降才衰减学习率factor0.5每次衰减一半。monitor字段必须与 validation_step 里 log 出的val_loss完全一致否则框架运行时报错。weight_decay 取 1e-4 是经验值对 1D EEG 任务来说 L2 正则过大反而会把卷积核学得过平滑丢失纺锤波这类瞬态特征。4.3 train.py 组织训练与 model_search.py 做超参搜索train.py 的职责是把数据划分、模型构建、Trainer 配置串起来。数据划分建议按受试者而不是按 epoch 随机切否则同一个人的前后 epoch 会同时出现在训练集和验证集验证指标虚高。提示按受试者而非按 epoch 切分数据集是睡眠分期实验的红线。同一患者相邻 epoch 高度相似混入训练集后验证指标会虚高 5 到 8 个百分点论文里的评估结果也会失去说服力。Trainer 配置里max_epochs设为 100配合 EarlyStopping 后通常在第 40 到 60 轮实际收敛。model_search.py 用 Optuna 与 Lightning 的 callback 配合做超参数搜索搜索目标不是 val_loss 而是验证集 macro-F1因为整体准确率被 N2 主导优化它没有意义。from pytorch_lightning.callbacks import EarlyStopping import optuna def objective(trial): lr trial.suggest_loguniform(lr, 1e-4, 1e-2) hidden trial.suggest_categorical(hidden_dim, [32, 64, 128]) dropout trial.suggest_uniform(dropout, 0.2, 0.6) # 按 trial 参数构建模型与 LightningModule module build_module(lrlr, hiddenhidden, dropoutdropout) trainer pl.Trainer(max_epochs50, callbacks[EarlyStopping(monitorval_loss, patience5)]) trainer.fit(module, train_loader, val_loader) return trainer.callback_metrics[val_macro_f1].item() study optuna.create_study(directionmaximize) study.optimize(objective, n_trials40)搜索空间里学习率用 log 均匀分布隐藏维度用离散候选值Dropout 用连续均匀分布。一次典型搜索跑 40 个 trial在 5 万 epoch 规模的数据下每个 trial 约 10 分钟一个晚上能跑完比手工调参节省大量时间。表训练超参数参考值超参数参考值搜索范围说明batch size12864–256加大 batch 需同步调低学习率初始学习率1e-31e-4–1e-2用 log 均匀分布搜索LSTM hidden6432–128过大在小数据集上过拟合严重Dropout0.50.2–0.6分类头与 LSTM 输出后都加上下文窗口 T31–9影响序列建模能力与训练速度5. 用 benchmark.py 复现精度指标计算与三个高频踩坑点benchmark.py 读取训练好的 checkpoint在测试集上逐窗口计算预测再输出整体准确率、macro-F1、Cohens kappa 和混淆矩阵。一个需要注意的细节是所有指标必须按 30 秒 epoch 粒度计算任何先对连续预测做平滑再计算指标的做法都会让 kappa 虚高。from sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score def evaluate_model(model, test_loader, device): model.eval() y_true, y_pred [], [] with torch.no_grad(): for x, y in test_loader: logits model(x.to(device)) y_pred.extend(logits.argmax(dim1).cpu().tolist()) y_true.extend(y.tolist()) return { acc: accuracy_score(y_true, y_pred), macro_f1: f1_score(y_true, y_pred, averagemacro), kappa: cohen_kappa_score(y_true, y_pred), }averagemacro是关键参数先计算五个类别各自的 F1 再取平均与 sklearn 默认的binary行为完全不同漏掉它跑出来的结果会少一个维度。评测时建议把混淆矩阵也打印出来重点看 N1 列的召回率这个数字低于 40% 基本可以断定训练过程被 N2 主导了。复现过程中最容易翻车的三个坑。第一个是数据泄露前面强调过按记录而不是按 epoch 划分否则验证指标虚高且最终部署效果对不上。第二个是 N1 的评估波动由于 N1 样本占比极低划分方式不同时 macro-F1 会在 5 个百分点范围内抖动稳定性做法是用分层采样保证验证集每类比例与训练集一致并重复三次划分取均值。第三个是预测结果的连续性模型偶尔会输出 W 到 N3 再跳回 W 的快速切换这不符合睡眠阶段转移的生理规律常见做法是在推理后对 logits 序列按时间轴做一维中值滤波窗口取 3 或 5能把单点误跳的噪声滤掉同时保留真实的阶段切换边界。本文还有配套的精品资源点击获取
返回列表