ARTICLE DETAIL

资讯详情

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

基于CNN与PyTorch的EEG睡眠分期实践:从数据处理到模型部署

基于CNN与PyTorch的EEG睡眠分期实践:从数据处理到模型部署 简介面向毕业设计、课程设计与期末大作业场景的深度学习睡眠状态检测项目代码包适合人工智能、生物信号处理方向的初学者或需要在EEG数据上完成分类任务的学生。资源围绕脑电EEG信号处理与CNN模型构建展开包含2个Python脚本和1个Markdown说明文档分别负责数据加载与模型训练并提供使用说明共3个文件、整个压缩包仅5KB便于快速阅读和二次修改。已有35人学习下载。借助该代码包读者可以掌握从EEG信号预处理、数据集划分到卷积神经网络分类的完整流程理解卷积层、池化层、全连接层等结构在睡眠分期中的具体用法并可参考说明文档快速复现实验、调整模型参数。对于需要提交代码或完成报告的学生而言是一份轻量、可直接运行的参考范例有助于快速完成课程项目或毕业设计。1. 为什么睡眠分期任务把卷积层放在循环层前面睡眠状态检测本质上是把连续的EEG信号切成片段再给每个片段打上Wake、N1、N2、N3、REM这样的标签。早期做法是人工提取频域特征再喂给SVM或随机森林但这套流程非常依赖先验知识不同睡眠阶段在δ波、θ波、α波、纺锤波上的分布差异很大一旦特征工程做得不够细模型上限就被卡住了。深度学习方法之所以能替代传统流程是因为它可以在原始波形的时频表示上直接学习特征而CNN在这里的优势尤其明显。一个反直觉的结论是对于单导联EEG睡眠分期一维卷积网络往往比LSTM更好训。EEG信号虽然有时序依赖但睡眠阶段的判别更依赖局部形态特征比如K复合波、纺锤波这类持续一到两秒的瞬态结构正好对应卷积核在时间轴上的局部感受野。而且CNN参数少、训练稳定在数据量不够大的课程设计场景里更容易收敛。如果你想更深入地理解为什么CNN能提取脑电特征可以对照经典的SleepEEGNet或DeepSleepNet的结构再看这份压缩包里的cnn-eeg-classification.py会发现设计思路是一致的。这个项目适合深度学习入门者、做生物信号分析的课程设计或期末大作业的人帮助你把数据加载、模型构建、训练评估整条链路跑通。2. EEG数据集的加载与预处理从原始信号到CNN输入张量2.1 load-dataset.py 的整体设计逻辑拿到一份睡眠状态检测的项目第一个要弄清楚的事情不是网络结构多先进而是数据怎么进模型。EEG数据集常见格式包括EDF、EDF、MAT和CSV不同来源的数据通道数、采样率、标注粒度都不一样。load-dataset.py 的作用就是把这些异构数据标准化成CNN能直接消费的三维张量。通常的流程是读取原始记录文件和对应的睡眠分期标签文件按固定窗口长度切分信号例如30秒为一个epoch因为睡眠分期的国际标准就是每30秒标注一次对每个epoch做滤波、去伪影、归一化返回数据集对象支持按批次迭代常见的做法是用PyTorch的Dataset和DataLoader封装这样后续训练代码只需要关注模型本身。下面的示例展示了load-dataset.py中核心的数据集类实现import numpy as np import torch from torch.utils.data import Dataset import scipy.io as sio class SleepEEGDataset(Dataset): def __init__(self, mat_path, window_sec30, fs100): mat_path: 预处理后的.mat文件, 包含eeg_data和labels字段 window_sec: 每个样本的时长, 睡眠分期标准为30秒 fs: 采样率, 常见为100Hz或128Hz data sio.loadmat(mat_path) # 原始信号形状: (n_epochs, n_channels, n_times) self.eeg data[eeg_data].astype(np.float32) self.labels data[labels].flatten() # 关键把多通道的通道维放到后面方便后续使用一维卷积 # 新形状: (n_epochs, n_times, n_channels) self.eeg self.eeg.transpose(0, 2, 1) self.fs fs def __len__(self): return len(self.labels) def __getitem__(self, idx): x self.eeg[idx] # (n_times, n_channels) y self.labels[idx] # 归一化到[0,1]避免不同受试者幅值差异过大 x_min x.min() x_max x.max() x (x - x_min) / (x_max - x_min 1e-8) # 增加通道维CNN卷积核期望输入为 (channels, length) x torch.tensor(x, dtypetorch.float32).permute(1, 0) y torch.tensor(y, dtypetorch.long) return x, y这段代码把数据加载和归一化放在了一起。关键点在于transpose之后将通道维放到了最后而permute又把它调回(channels, length)这是为了匹配一维卷积的输入约定。对EEG数据做min-max归一化而不是z-score是因为不同受试者的脑电幅值差异很大min-max可以把所有样本拉回同一尺度避免模型在训练初期被方差过大的样本主导。2.2 滤波和伪影去除的细节load-dataset.py 里如果包含信号预处理通常会用带通滤波保留EEG的有效频段。睡眠EEG的主要能量集中在0.5Hz到30Hz之间超过30Hz的多为肌电伪影。常见做法是使用Butterworth滤波器from scipy.signal import butter, filtfilt def bandpass_filter(eeg, fs100, low0.5, high30, order4): 零相位带通滤波, 使用filtfilt避免相位偏移 nyquist 0.5 * fs low_norm low / nyquist high_norm high / nyquist b, a butter(order, [low_norm, high_norm], btypeband) # axis-1 表示对最后一个维度时间轴滤波 filtered filtfilt(b, a, eeg, axis-1) return filtered这里用filtfilt而不是lfilter因为EEG分析对相位敏感零点相位滤波能保持波形特征。滤波后还需要检查是否存在明显的基线漂移如果信号基线上下大幅抖动可以先做一次0.5Hz高通滤波。很多开源预处理工具还会进一步做独立成分分析去除眼电伪影但在课程设计级别带通滤波加简单的幅值阈值裁剪通常就够用了。2.3 标签映射与类别不平衡处理睡眠分期有五种标准阶段但有些数据集会把N3和N4合并有的还会标注Movement或Unknown类别。load-dataset.py里需要把原始标签映射到统一的编码。一种常见方式是把类别映射为整数Wake0, N11, N22, N33, REM4。def map_labels(raw_labels): 标签归一化: 将标准AASM规则转换为统一的0-4编码 mapping {W: 0, N1: 1, N2: 2, N3: 3, REM: 4, SLEEP-S1: 1, SLEEP-S2: 2, SLEEP-S3: 3, SLEEP-S4: 3} mapped np.array([mapping.get(str(l).strip(), 0) for l in raw_labels]) return mapped注意N3和N4合并是常见的做法因为深度学习模型往往难以区分这两者而且临床中也常将它们视为深睡眠。标签映射后强烈建议用np.bincount检查一下每类样本数量睡眠数据集中N1通常很少有时只占5%这会导致模型严重偏向多数类。处理办法是给损失函数加类别权重或者在训练时用加权采样器。这个内容会在第4章详细讨论。3. CNN模型架构与训练循环的构建3.1 cnn-eeg-classification.py 中的网络结构模型文件是整个项目的核心。针对单导联或双导联的EEG信号一维卷积网络是最直接的选择。网络设计通常包含几个特征提取块每个块由一维卷积、批归一化、ReLU激活和最大池化组成。在睡眠分期中输入长度为3000个采样点30秒×100Hz感受野需要覆盖至少1秒的波形才能识别纺锤波和K复合波所以第一层卷积核不宜太小。下面给出一个典型的网络定义import torch.nn as nn import torch.nn.functional as F class SleepCNN(nn.Module): def __init__(self, n_channels2, n_classes5): super(SleepCNN, self).__init__() # 第一层: 大卷积核捕捉慢波 self.conv1 nn.Conv1d(n_channels, 32, kernel_size101, stride2, padding50) self.bn1 nn.BatchNorm1d(32) self.pool1 nn.MaxPool1d(kernel_size2, stride2) # 第二层: 中等卷积核捕捉纺锤波 self.conv2 nn.Conv1d(32, 64, kernel_size31, stride2, padding15) self.bn2 nn.BatchNorm1d(64) self.pool2 nn.MaxPool1d(kernel_size2, stride2) # 第三层: 小卷积核组合高阶特征 self.conv3 nn.Conv1d(64, 128, kernel_size11, stride1, padding5) self.bn3 nn.BatchNorm1d(128) self.pool3 nn.MaxPool1d(kernel_size2, stride2) # 全局平均池化替代全连接层, 减少参数 self.gap nn.AdaptiveAvgPool1d(1) self.fc nn.Linear(128, n_classes) def forward(self, x): x self.pool1(F.relu(self.bn1(self.conv1(x)))) x self.pool2(F.relu(self.bn2(self.conv2(x)))) x self.pool3(F.relu(self.bn3(self.conv3(x)))) x self.gap(x) x x.view(x.size(0), -1) return self.fc(x)使用全局平均池化替代全连接层能显著减少参数量降低过拟合风险。第一层卷积核设为101在100Hz采样率下正好对应1秒的时域窗口能捕获delta波和theta波这类慢波。而后续31和11的核负责提取局部瞬态特征。3.2 训练循环与超参数设置训练代码通常需要自己写循环因为PyTorch没有内置的高级训练器。关键超参数如下参数推荐值说明batch_size64太小训练波动大太大显存压力高learning_rate1e-3使用Adam时常用若用SGD需要调整到0.01~0.1epoch30~50配合早停观察验证集lossoptimizerAdambeta10.9, beta20.999loss functionCrossEntropyLoss多分类标准选择dropout0.3全连接层之前使用训练循环的标准写法如下import torch.optim as optim from sklearn.metrics import accuracy_score, f1_score def train_one_epoch(model, dataloader, criterion, optimizer, device): model.train() total_loss 0 all_preds [] all_labels [] for x, y in dataloader: x, y x.to(device), y.to(device) optimizer.zero_grad() out model(x) loss criterion(out, y) loss.backward() optimizer.step() total_loss loss.item() * x.size(0) preds torch.argmax(out, dim1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(y.cpu().numpy()) avg_loss total_loss / len(dataloader.dataset) acc accuracy_score(all_labels, all_preds) f1 f1_score(all_labels, all_preds, averagemacro) return avg_loss, acc, f1这里在训练阶段就同步计算准确率和宏平均F1原因是睡眠分期的类别不平衡严重只看准确率会被N2类的高占比带偏。宏平均F1能给每一类相同的权重如果N1识别效果差F1会立刻掉下来比准确率敏感得多。3.3 验证与测试的划分策略睡眠数据不能简单地随机打乱分割因为同一个人的相邻epoch具有强相关性随机分割会让模型在测试时“偷看”到同一受试者的睡眠片段导致性能虚高。正确做法是按受试者划分训练集包含若干人的全部数据测试集只用未见过的个体数据。load-dataset.py如果支持受试者ID应当优先采用这种划分方式。def split_by_subject(subjects, eeg_data, labels, test_subjects): 按受试者划分数据 subjects: 每个epoch对应的受试者编号 test_subjects: 用作测试的受试者编号列表 test_mask np.isin(subjects, test_subjects) train_mask ~test_mask train_data eeg_data[train_mask] train_labels labels[train_mask] test_data eeg_data[test_mask] test_labels labels[test_mask] return (train_data, train_labels), (test_data, test_labels)使用这种划分后模型的泛化能力会明显下降但这才更接近真实应用场景。很多期末大作业为了追求报告上的漂亮数字随意随机切分训练集和测试集这是最常踩的坑。4. 训练结果判读准确率之外还要看什么4.1 混淆矩阵里的睡眠阶段偏移模型训练结束后cnn-eeg-classification.py通常会输出分类报告或混淆矩阵。睡眠分期任务中N1是一个典型的难分类别因为它和Wake、N2的边界非常模糊。如果只看准确率可能训练集上达到90%但混淆矩阵里N1的召回率只有30%。这时候需要调整损失函数的权重。import torch.nn as nn # 统计每个类别的样本数量 class_counts np.bincount(all_train_labels) total class_counts.sum() # 权重与样本数成反比, 让少数类获得更高权重 weights torch.tensor([total / (len(class_counts) * c) for c in class_counts], dtypetorch.float32).to(device) criterion nn.CrossEntropyLoss(weightweights)加了类别权重之后N1的分类能力会显著提升但Wake的精确率可能会轻微下降。这是常见的trade-off。如果你发现N1和N2之间的混淆特别严重还可以把N1与N2合并为浅睡眠简化任务。4.2 早停与模型保存训练过程中需要监控验证集loss当连续多个epoch验证loss不再下降时就应该停止训练并恢复最佳模型。常见做法如下best_f1 0 patience 10 no_improve 0 for epoch in range(max_epochs): train_loss, train_acc, train_f1 train_one_epoch(model, train_loader, criterion, optimizer, device) val_loss, val_acc, val_f1 evaluate(model, val_loader, criterion, device) if val_f1 best_f1: best_f1 val_f1 torch.save(model.state_dict(), best_sleep_cnn.pth) no_improve 0 else: no_improve 1 if no_improve patience: print(fEarly stopping at epoch {epoch}) break这里使用验证集F1作为早停指标而不是loss是因为F1对类别不平衡更稳健。只保存验证集表现最好的模型避免最后几个epoch过拟合后的参数覆盖掉好的结果。4.3 过拟合的其他防线睡眠EEG数据量通常不大一个受试者的整夜记录大约只有900个epoch按30秒计算8小时睡眠。如果不做任何正则化CNN很容易在训练集上达到接近100%准确率而验证集上只有70%。建议的组合拳是在卷积层之间插入Dropout比例从0.1到0.3不要太高否则信号特征难以传递对训练集做数据增强比如以95%到105%的比例随机缩放时间轴模拟采样率差异使用AdamW替代Adam并设置weight_decay1e-4时间轴缩放是一种有效的EEG增强方式因为不同采集设备的采样率标定存在误差模型对小幅时间拉伸应该保持鲁棒。实现时可以用torch.nn.functional.interpolate或np.interp重采样。5. 把CNN睡眠分期模型接进真实数据流预测新记录与特征可视化5.1 对新的EEG原始记录做预测训练好的模型最终要拿整夜的EEG记录来跑预测。这时需要先把原始信号切成30秒窗口然后逐窗口预测再把预测结果按时间顺序拼接成完整的睡眠图谱。这个流程要在python脚本里串起来。def predict_sleep_stages(model, raw_eeg, fs100, window_sec30): raw_eeg: 一维或二维原始EEG信号 (n_channels, n_samples) 返回: 每个epoch的预测标签序列 model.eval() window_len fs * window_sec # 计算能切出多少个完整epoch n_windows raw_eeg.shape[1] // window_len # 丢弃末尾不足一个窗口的数据 cut_len n_windows * window_len raw_eeg raw_eeg[:, :cut_len] # 按窗口拆分, 返回形状 (n_windows, n_channels, window_len) windows raw_eeg.reshape(raw_eeg.shape[0], n_windows, window_len).transpose(1, 0) preds [] with torch.no_grad(): for w in windows: # 归一化方式必须与训练时完全一致 w_min w.min() w_max w.max() w_norm (w - w_min) / (w_max - w_min 1e-8) x torch.tensor(w_norm, dtypetorch.float32).unsqueeze(0) out model(x) pred torch.argmax(out, dim1).item() preds.append(pred) return np.array(preds)注意归一化参数必须在每个窗口上独立计算还是在整个记录上统一计算这个细节非常关键。训练时是在每个epoch内部做min-max归一化那么预测时也要对每个窗口单独做否则训练和推理的数据分布不一致。很多同学的模型测试集准确率不错但部署到新数据上效果暴跌原因就在这里。5.2 用特征图验证模型学到了什么为了确认CNN确实学到了睡眠阶段相关的波形特征可以把中间卷积层的输出可视化。第一层卷积核的输出可以看成一堆经过滤波后的信号频率选择特性一目了然。import matplotlib.pyplot as plt def visualize_first_layer(model, sample_x, layer_index0, top_k8): 可视化第一层卷积后的特征图 sample_x: 形状为 (1, n_channels, n_times) 的输入张量 activations {} def hook_fn(name): def hook(model, input, output): activations[name] output.detach().cpu().numpy() return hook # 注册hook到第一层卷积 handle model.conv1.register_forward_hook(hook_fn(conv1)) model.eval() with torch.no_grad(): model(sample_x) handle.remove() feat_maps activations[conv1][0] # (n_filters, n_times_downsampled) fig, axes plt.subplots(top_k, 1, figsize(12, 8)) for i in range(top_k): axes[i].plot(feat_maps[i][:500]) axes[i].set_ylabel(fK{i}) plt.tight_layout() plt.savefig(conv1_features.png)如果某些卷积核的输出有明显的波形节律性波动说明它们学到了慢波或纺锤波的滤波模式如果输出全是噪声那可能是训练不充分或卷积核尺寸不合理。这个可视化步骤写进报告里比单纯的准确率图表更有说服力。5.3 部署时的常见坑位检查把模型从训练环境搬到推理环境时有一个容易被忽略的点PyTorch模型默认处于训练模式。如果你直接用model(raw_eeg)做预测BatchNorm会使用当前batch的统计量导致输出不稳定。正确做法是调用model.eval()并包裹在torch.no_grad()里。另一个坑是输入通道数不匹配。有些睡眠数据集包含两个通道比如F4-M1和C4-M1但有些设备只有单导联。如果你的模型定义了2个输入通道就必须保证预测时也喂2个通道。兼容做法是在模型开头加一个通道适配层class FlexibleInputCNN(nn.Module): def __init__(self, expected_channels2, **kwargs): super().__init__() # 让第一层卷积接受任意通道数, 通过padding或者重复填充 self.channel_adapter nn.Conv1d(expected_channels, expected_channels, kernel_size1)但最稳妥的方式还是在预处理阶段就统一通道数和采样率。你可以在README.md里明确写清楚模型期望的输入格式比如“输入形状为(batch, 2, 3000)采样率100Hz数值范围[0,1]”。这样任何拿到代码的人都能快速跑通而不用去猜数据管道到底做了什么。本文还有配套的精品资源点击获取
返回列表