ARTICLE DETAIL

资讯详情

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

Python心电图5分类实战:CNN+BiLSTM模型构建与MIT-BIH预处理全流程

Python心电图5分类实战:CNN+BiLSTM模型构建与MIT-BIH预处理全流程 简介面向心血管疾病心电图ECG信号识别任务这份资源提供基于Python的完整5分类深度学习解决方案适合期末大作业、课程设计与毕业设计场景。内容覆盖从数据预处理、模型搭建到训练评估的整套流程包含CNN、LSTM、CNN-LSTM以及结合小波变换的多种网络结构实现并附有对应训练好的权重文件可直接用于验证与结果复现。资源包共36个文件总大小约50.19MB主要包含7个Python源码、10个pth模型权重、17张训练过程与结构示意图以及一份docx格式使用说明手册目录按模块划分清晰便于按需查阅。其中Python代码涵盖数据工具、不同模型定义与训练评估脚本训练好的模型覆盖多种结构组合与数据规模适合对比不同方法在有限数据条件下的表现。当前已有246人学习下载对于希望在心电图分类任务上快速上手、进行算法对比和撰写项目报告的开发者而言是一份即取即用的实操参考。1. 这个标题在讲什么一份心电图5分类工程的完整闭环收到一份“基于python的心电图信号设计模型结构完成5分类任务源码模型使用说明.zip”第一反应别急着解压。先问三件事5分类是哪5类模型结构长什么样训练数据从哪来这三件事没理顺模型跑通也只是个黑匣子。这个标题对应的是心电图ECG分类里最经典的AAMI五分类——正常搏动、室上性异位搏动、室性异位搏动、融合搏动、起搏/未知类落在工程上就是从MIT-BIH这类公开数据集出发完成信号滤波、滑动窗口切分、标签映射、模型结构设计、训练和评估的完整链路。适合刚入门生物信号处理的研究生也适合准备把分类模型落到真实心电数据上的算法工程师。2. 先定模型结构为什么CNN加BiLSTM是ECG 5分类的稳妥起点ECG信号和图像、文本都不太一样单条记录通常是两导联、360Hz采样率、长达几万到几十万采样点的连续波形。类别差异主要体现在心拍形态上——比如室性早搏的QRS波宽大畸形室上性早搏的P波位置异常——但单独的形态还不够某些类别需要结合前后几个心拍的节律上下文才能区分。这就是为什么模型结构不能拍脑袋选得同时兼顾局部形态和时序依赖。2.1 三类主流结构怎么选CNN、LSTM与Transformer先说结论在小数据集、低显存、要快速迭代出可解释结果的场景里CNN加BiLSTM的性价比最高。下面是我的实际对比结构局部形态建模时序依赖参数量小数据表现显存压力纯CNN强弱小中低BiLSTM弱强中中中CNNBiLSTM强强中好中Transformer中强大差易过拟合高TCN时序卷积强中中中低纯CNN的典型问题是感受野有限对“前一个心拍是正常、当前心拍是早搏”这类上下文关系建模不足。Transformer在长序列上确实强但心电图数据集往往只有几万到几十万条样本参数量稍大就过拟合而且要跑得动至少需要一块像样的显卡低显存环境很吃力。TCN是CNN的时序变体膨胀卷积拉长感受野表现不错但调参复杂度比LSTM高对新手不太友好。综合下来我一般会把CNNBiLSTM作为基线结构前几层一维卷积负责提取P波、QRS波群、ST段的局部形态特征BiLSTM在卷积输出的时序特征上再建模节律上下文最后用注意力机制把关键时间步聚合起来。这套组合在多个心电公开数据集上表现稳定参数量控制在几十万级别4G显存就能跑。2.2 核心结构代码三层卷积加双向LSTM怎么搭下面是一个可以直接抄的PyTorch模型结构输入是两导联、2.5秒窗口360Hz采样率下正好900个采样点输出是5个类别的logitsimport torch import torch.nn as nn class ECG5Classifier(nn.Module): def __init__(self, in_channels2, seq_len900, num_classes5): super().__init__() # 三段Conv1d逐层扩大感受野先抓QRS形态再抓整体波形趋势 self.conv nn.Sequential( nn.Conv1d(in_channels, 32, kernel_size15, padding7), nn.BatchNorm1d(32), nn.ReLU(inplaceTrue), nn.MaxPool1d(2), nn.Conv1d(32, 64, kernel_size9, padding4), nn.BatchNorm1d(64), nn.ReLU(inplaceTrue), nn.MaxPool1d(2), nn.Conv1d(64, 128, kernel_size5, padding2), nn.BatchNorm1d(128), nn.ReLU(inplaceTrue), nn.MaxPool1d(2), ) # 池化后序列长度900 - 450 - 225 - 112 self.lstm nn.LSTM(128, hidden_size64, bidirectionalTrue, num_layers1, batch_firstTrue) # 注意力让模型自己决定哪些时间步对分类更重要 self.attn nn.Sequential( nn.Linear(128, 64), nn.Tanh(), nn.Linear(64, 1), nn.Softmax(dim1) ) self.fc nn.Linear(128, num_classes) def forward(self, x): # x: (batch, channels, seq_len) x self.conv(x) # (batch, 128, 112) x x.transpose(1, 2) # (batch, 112, 128) out, _ self.lstm(x) # (batch, 112, 128) attn_w self.attn(out) # (batch, 112, 1) x (out * attn_w).sum(dim1) # 注意力加权聚合 return self.fc(x)逐层解释一下参数第一层卷积核取15个采样点大约对应42毫秒比单个QRS波群通常80~120毫秒短一些能捕捉波形内部的细微形态变化第二层卷积核9个点进一步组合局部特征第三层5个点做细粒度调整。每层卷积后面都跟BatchNorm和ReLU这是让训练稳定的关键千万别省。BiLSTM的hidden_size取64双向合并后输出维度是128与卷积层输出通道数一致这样注意力层的输入维度对齐不需要额外投影。注意力用两层MLP加Softmax对序列维度做归一化相当于给每个时间步学一个权重最后加权求和得到整条窗口的向量表达。2.3 输入长度与导联数900个采样点为什么够用窗口长度直接决定模型能看到多少心跳。360Hz采样率下正常心率约60~100次/分即每个心拍间隔0.6~1秒。取2.5秒窗口绝大多数情况下能覆盖2~3个完整心拍既保留了节律上下文又不至于让序列太长拖慢训练速度。900个采样点经过三层MaxPool后被压到112个时间步再进BiLSTM计算量很友好。导联数方面MIT-BIH原始数据是两导联通常MLII加V1或V2我建议保留两导联一起输入。有的实现只取MLII单导联输入变成单通道模型第一层的in_channels要改成1参数更少但会丢掉部分空间信息。如果只有一块4G显存的卡把in_channels改成1、LSTM的hidden_size降到32训练速度能快一半准确率通常只掉1~2个点。3. 把MIT-BIH改成可训练的5分类样本预处理与标签映射模型结构定下来之后最花时间的其实是数据。MIT-BIH心律失常数据库是心电图分类的事实标准48条半小时记录、360Hz采样率、每条带独立的心拍注释。但原始注释符号有几十种得先归并成AAMI规定的5大类再切窗、滤波、对齐标签。这一步做得干不干净直接决定模型上限。3.1 原始数据格式采样率、导联与AAMI五类符号用wfdb库读取MIT-BIH的常规做法是wfdb.rdsamp读波形、wfdb.rdann读注释文件。注释文件里每个心拍有一个符号比如N代表正常搏动、V代表室性早搏、A代表房性早搏常见符号映射关系如下AAMI类别含义原始符号N0正常搏动N, L, R, e, jS1室上性异位搏动A, a, J, SV2室性异位搏动V, EF3融合搏动FQ4起搏/未知/不可分类/, f, Q, U, P注意F类融合搏动在数据集里占比极低有的实现图省事把它并入V类结果所谓“5分类”实际只有4类。如果你拿到的是预训练好的模型第一件事就是确认它在F类上有没有真实的区分能力。3.2 滤波与滑动窗口切分把半小时信号变成一沓样本原始信号先要做带通滤波去掉基线漂移和工频干扰。心电图有效能量集中在0.5~45Hz我用四阶Butterworth带通滤波边缘效应比二阶更陡又不会像更高阶那样引入明显振铃。然后以固定窗口长度和步长做滑动窗口切分import numpy as np import wfdb from scipy.signal import butter, filtfilt def load_and_preprocess(record_path): record wfdb.rdsamp(record_path, channels[0, 1]) ann wfdb.rdann(record_path, atr) sig record[0].T # (2, n_samples)两导联 fs record[1][fs] # 通常360 b, a butter(4, [0.5, 45], btypebandpass, fsfs) sig np.stack([filtfilt(b, a, ch) for ch in sig]) return sig, ann.sample, ann.symbol, fs def make_windows(sig, ann_sample, ann_symbol, fs360, win_sec2.5, overlap0.5): win_len int(fs * win_sec) step_len int(win_len * (1 - overlap)) X, y [], [] for start in range(0, sig.shape[1] - win_len 1, step_len): end start win_len window sig[:, start:end] # 取窗口中心最近的R峰注释作为标签 center start win_len // 2 idx np.argmin(np.abs(ann_sample - center)) symbol ann_symbol[idx] label map_aami(symbol) # 映射到0~4 X.append(window) y.append(label) return np.stack(X), np.array(y)滑动窗口的重叠率我通常设0.5也就是窗口前进一半长度。这样做有两个作用第一变相做了数据增强同一个心拍可能出现在两个窗口里第二避免R峰刚好卡在窗口边缘导致形态被截断。窗口标签取的是“离窗口中心最近的注释”这是最常见做法因为中心位置的心拍受上下文影响最均衡。3.3 窗口标签不唯一中心R峰还是多数投票中心R峰取标签在多数场景够用但也有人用窗口内全部注释的多数投票。两种做法各有取舍中心R峰更简单实现起来就一行np.argmin适合R峰检测稳定、注释间隔均匀的记录多数投票对窗口内多个异常心拍的场景更稳比如窗口里有一个正常、一个室早、一个正常多数会投成正常但中心心拍明明是室早——这种就翻车。我的经验是先用中心R峰跑一版基线确认指标后再换多数投票做对比看F1是否显著提升。多数投票实现并不复杂遍历窗口内的注释计数即可但要注意窗口边缘的注释完整性我遇到过窗口恰好切掉半个心拍导致投票类别错乱的。3.4 按患者切分训练集这一步省得后面吃大亏预处理完的样本之间并不独立同一患者半小时记录里切出来的上千个窗口高度相关。如果随机打乱再切训练集和验证集同一个患者的窗口会同时出现在两边模型相当于“见过”验证数据指标虚高得离谱。正确做法是按患者记录ID切分保证训练集和验证集里没有任何一个患者的数据交叉def patient_split(record_ids, test_ratio0.2): # record_ids: 与每个样本对应的患者或记录编号 unique_patients np.unique(record_ids) n_test int(len(unique_patients) * test_ratio) rng np.random.default_rng(42) test_patients rng.choice(unique_patients, n_test, replaceFalse) test_mask np.isin(record_ids, test_patients) return ~test_mask, test_mask # train_mask, test_mask这段代码的作用是先拿到所有独立患者ID从中随机挑出20%作为测试患者再通过np.isin把属于这些患者的全部窗口划到测试集。这样训练集和测试集在患者维度上零交集评估结果才是真实的泛化能力。如果数据来自其他来源也要保证同样的切分逻辑。4. 训练与评估让模型不把一切预测成正常类心电图5分类最大的杀手是类别不均衡。正常搏动占比常常超过80%室上性和融合搏动加起来可能不到10%。如果不做任何处理模型学到的就是“全部预测成正常类也能有不错准确率”训练loss照降但每个异常类别的召回率几乎为零。解决这个问题要在训练配置和评估指标两个方向上同时下手。4.1 类别不均衡先处理类别权重与Focal Loss最简单的第一板斧是给损失函数加类别权重让少数类的梯度贡献更大。用sklearn计算权重后再传给CrossEntropyLossfrom sklearn.utils.class_weight import compute_class_weight import torch.nn as nn classes np.unique(y_train) weights compute_class_weight(balanced, classesclasses, yy_train) weights torch.tensor(weights, dtypetorch.float32) loss_fn nn.CrossEntropyLoss(weightweights)compute_class_weight的balanced模式按样本频率反比生成权重比如某类只占5%权重会接近占80%那类的16倍。这相当于告诉模型“分错少数类的代价更高”。要注意别把小批量内的权重直接传给loss必须用整个训练集的类别分布计算否则每个batch的权重都在变训练不稳定。如果加了类别权重后少数类召回率还是上不去第二板斧是换Focal Loss。Focal Loss在交叉熵基础上加了调节因子让模型聚焦于难分类样本。实现时gamma一般取2alpha用类别权重向量我的经验是在S类和F类上比纯类别权重多提升2~3个点的召回率。4.2 训练循环与早停保证不白跑训练循环里我固定用Adam优化器配ReduceLROnPlateau学习率初始1e-3验证loss连续3个epoch不降就减半。早停的耐心值设5个epoch防止最后几个epoch在小数据集上过拟合后把最佳模型覆盖掉opt torch.optim.Adam(model.parameters(), lr1e-3) sched torch.optim.lr_scheduler.ReduceLROnPlateau( opt, factor0.5, patience3) early_stopping EarlyStopping(patience5) for epoch in range(60): model.train() for xb, yb in train_loader: opt.zero_grad() out model(xb) # (batch, 5) loss loss_fn(out, yb) loss.backward() opt.step() model.eval() val_loss 0 with torch.no_grad(): for xb, yb in val_loader: out model(xb) val_loss loss_fn(out, yb).item() val_loss / len(val_loader) sched.step(val_loss) early_stopping(val_loss) if early_stopping.early_stop: print(fepoch {epoch} 触发早停) breakEarlyStopping类的核心逻辑是维护一个最佳验证loss连续patience次没有刷新就把early_stop置为True。用这段代码时要注意保存最优模型权重的时机我一般会在val_loss刷新时执行torch.save(model.state_dict(), best.pth)早停只负责终止不负责保存。Batch size在显存允许时取64太小梯度噪声大太大在类别不均衡下更容易陷入全预测正常类的局部最优。4.3 评估指标混淆矩阵和每类F1比accuracy重要训练完直接打印accuracy是远远不够的尤其是类别不均衡的数据。我用sklearn.metrics输出每个类别的precision、recall、F1同时画混淆矩阵热力图定位到底哪两类在互相纠缠from sklearn.metrics import confusion_matrix, classification_report import matplotlib.pyplot as plt preds predict_all(model, val_loader) cm confusion_matrix(y_val, preds) print(classification_report(y_val, preds, target_names[N, S, V, F, Q])) plt.figure(figsize(6, 5)) plt.imshow(cm, cmapBlues) for i in range(5): for j in range(5): plt.text(j, i, cm[i, j], hacenter, vacenter) plt.xticks(range(5), [N, S, V, F, Q]) plt.yticks(range(5), [N, S, V, F, Q]) plt.xlabel(Predicted) plt.ylabel(True) plt.show()看混淆矩阵有个实用技巧先看S类室上性和V类室性之间是否有大量互相误判这两类在形态上非常接近是区分度最差的组合再看F类是不是完全没人能分对F类样本太少时模型学不到有效特征输出会表现为F类recall接近0。这两种情况在后续调优时处理方式完全不同前者要加更长窗口或者多导联信息后者要针对F类做特殊数据增强。注意类别权重只影响训练时的梯度推理时模型输出概率直接用argmax取类别即可不要在执行推理时手动乘权重那样会扭曲概率分布。5. 最容易翻车的5个坑现象、原因与排查办法上面说的都是常规流程但真正让一个模型从“能跑”变成“能上线”靠的是排除一个个隐蔽的坑。我把自己踩过的、以及帮别人排查过的典型问题整理成下面五条每一条都按“现象→原因→解决”的口径展开。5.1 坑1训练和推理的窗口参数不一致现象训练时准确率到90%以上测试时跌到70%且预测结果里正常类占比异常高。原因训练用了重叠率0.5的滑动窗口推理时改了窗口长度或者没做重叠直接从头切到尾。模型看到的推理输入分布和训练分布完全不一致。最隐蔽的变体是训练做了滤波推理时忘了滤波波形形态全变。解决把滤波、切窗、归一化封装成一个preprocess_for_inference()函数训练和推理共用同一份代码路径不要手工复制粘贴参数。我习惯把窗口长度和重叠率做成配置常量放在单独配置类里训练脚本和推理脚本都引用它。5.2 坑25分类被悄悄做成3分类现象代码注释写的5分类模型输出层也确实有5个神经元但训练日志里F类样本数为0。原因标签映射时把F类符号映射错了或者F类样本在预处理阶段因为注释文件读取异常被全部丢弃。更常见的是网上下载的源码包把F类并入了V类美其名曰“临床不关心”但严格按标题要求就不是5分类。解决训练前打印每个类别的样本数量和比例确认5类都在再对测试集做一次预测看预测类别集合是否完整覆盖0~4。缺类别时优先查标签映射表不查模型结构。5.3 坑3同患者数据串进训练集和验证集现象验证集F1比预想高很多但换一批新患者数据立刻崩掉。原因切分样本时按行随机切同一个患者的窗口两边都有。心电图信号里同一患者的心拍形态高度自相关模型记住了患者特征而不是病理特征验证集等于开卷考试。解决严格按患者ID切分用3.4节的patient_split逻辑。如果数据集没有患者ID字段至少要用记录编号代替。判断有没有串集可以在训练后随机抽一个测试患者的原始信号看模型预测是否在相邻窗口间跳变频繁——跳变越少越可疑说明模型在记忆该患者整体特征。5.4 坑4滤波把QRS波群削平了现象训练loss下不去或者模型对V类室性早搏完全无感。原因带通滤波下截止频率设太高比如直接高通10HzQRS波群的低频分量被滤掉宽大畸形的QRS反而变得和正常波差不多。另一个常见原因是用了lfilter而不是filtfilt相位延迟导致波形整体偏移注释的R峰位置和实际波形尖峰错位。解决滤波后画一段波形对比原图和滤波后的图确认QRS波群形态和R峰位置没有明显变化。0.5~45Hz配四阶filtfilt是经过验证的安全组合尽量不要动参数如果新高通截止频率就要重新核对R峰检测准确率。5.5 坑5类别不均衡下验证损失不降现象加了类别权重后验证loss先降后升或者一直震荡但验证准确率却很高。原因损失函数加权后少数类的梯度占比变大训练早期模型用力学少数类导致正常类准确率下降而accuracy指标本身就偏向多数类掩盖了这种变化。解决训练过程同时记录loss和每类F1epoch结束时打印加权平均F1而不是accuracy。如果loss震荡先把学习率降到3e-4再把batch size提到64看两个指标是否同步稳定下来。我遇到过半数的“不降”其实都是lr偏大加上数据批次不均衡导致调参顺序是lr优先、batch次之、模型结构最后。6. 进阶验证患者级5折交叉验证与窗口投票推理把单次训练验证跑通之后想判断模型是真的能用还得做两层更硬的验证。第一层是用患者级5折交叉验证替代单次划分把“挑了一组好患者做验证集”的运气成分去掉第二层是在推理阶段用多数投票替代单窗口预测把窗口级别的抖动消掉。患者级5折交叉验证用GroupKFold最直接它会保证同一个患者的所有窗口在同一折内不会跨折泄漏from sklearn.model_selection import GroupKFold gkf GroupKFold(n_splits5) for fold, (train_idx, val_idx) in enumerate( gkf.split(X, y, groupsrecord_ids)): model ECG5Classifier(in_channels2, seq_len900, num_classes5) # 训练和评估这一折记录每类F1 ...5折跑完后看每类F1的均值和标准差。如果V类的均值高但标准差大说明模型对V类的识别依赖特定患者特征泛化仍不可靠如果每折指标都很稳这个模型才有讨论上线的基础。推理阶段的窗口投票是最后一道防线。对一段长心电图以步长等于窗口长度切成不重叠片段逐段预测后统计各类别出现次数取众数作为整段结论from collections import Counter def predict_sequence(model, sig, fs360, win_sec2.5): win_len int(fs * win_sec) preds [] model.eval() with torch.no_grad(): for start in range(0, sig.shape[1] - win_len 1, win_len): win torch.tensor( sig[:, start:start win_len] ).unsqueeze(0).float() # (1, 2, win_len) out model(win) preds.append(out.argmax(dim1).item()) return Counter(preds).most_common(1)[0][0]投票的意义不只是平滑噪声更是对“孤立的异常预测”做纠错——单窗口里一个室早可能因为窗口截断被错判但整段信号里多数窗口会指向正确类别。这个技巧在长时心电监测场景尤其有用代价只是计算量线性增加换来的是稳定性提升。我现在拿到一份心电图分类代码固定动作是先跑患者级5折、拉混淆矩阵、看S和V之间的纠缠情况再决定要不要调窗口长度而不是盯着accuracy欢呼。这两条习惯帮我在这个方向少走了很多弯路希望你也能少踩几个。希望帮到你。本文还有配套的精品资源点击获取
返回列表