心电图信号分类实战:基于Python的CNN+BiLSTM+Attention端到端模型

发布时间:2026/8/26 11:50:43
心电图信号分类实战:基于Python的CNN+BiLSTM+Attention端到端模型 简介深度学习在时间序列分类领域展现出强大的特征提取能力特别适合心电信号等生物医学信号的自动分析。基于Python与PyTorch构建端到端分类模型通过一维CNN、双向LSTM与注意力机制的组合可有效捕获心电波形中的形态特征与时序依赖实现对单导联心电图信号的五分类识别。数据预处理与按患者划分训练集是保证模型泛化能力的关键。结合AAMI标准对MIT-BIH数据集进行归组并利用类别加权损失处理不平衡问题最终在实际心律失常检测中达到稳定效果。该方案兼顾了算法原理与工程落地适用于可穿戴设备异常诊断、辅助决策等场景为ECG分类项目提供可复现的参考路径。 这套基于Python的心电图信号分类项目我前后完整跑通过好几遍最近一次是为了给团队的心律失常检测做技术验证。压缩包里装的是完整源码、训练好的模型权重和一份项目说明文档核心目标就一句话把单导联心电图信号通过端到端的深度学习模型自动划分成5个类别。这篇文章围绕整个项目的设计思路、模型结构、数据预处理、训练过程和踩过的坑做一次彻底复盘不管是刚接触心电信号处理的学生还是想快速落地ECG分类的工程师照着我这条线走能省下大量弯路。1. 项目概览心电图5分类到底在解决什么问题1.1 任务定义与5类标签含义先把分类目标确认清楚不然拿到代码很容易茫然。心电图信号分类有很多种玩法有病种级别的、有患者级别的、有节律级别的。这个项目做的是心律失常数据集上的搏动级beat-level5分类也就是对单个心搏波形进行判别。5类标签来自AAMI标准对MIT-BIH心律失常数据库的整合归组。MIT-BIH原始标注符有几十种比如N代表正常搏动、L代表左束支传导阻滞、R代表右束支传导阻滞、A代表房性早搏、V代表室性早搏、F代表融合搏动等。如果直接用原始标注符分类类别太多且样本量差异悬殊模型很难收敛。AAMI标准把语义相近的标注归并成5大类大类AAMI类别包含的原始标注常见含义N正常搏动N, L, R, e, j等正常心搏及束支传导阻滞S室上性异位搏动A, a, J, S房性早搏、交界性早搏V室性异位搏动V, E室性早搏、室性逸搏F融合搏动F正常与室性融合搏动Q未知搏动/起搏心搏、未分类心搏这个归组逻辑是整个项目里最容易被忽略、但最关键的设定。很多人拿到源码第一反应是“为什么我只看到5个数字标签”因为它们对应的不是原始ECG标注而是已经映射到AAMI大类后的结果。1.2 技术路线与整体流程整个项目的处理链路可以概括为原始心电图信号 → 信号去噪 → R峰检测 → 单搏动分割 → 归一化 → 训练/测试集划分 → 模型训练 → 模型评估 → 分类预测。这里面每一步都会直接影响最终准确率尤其是“信号分割”和“训练/测试划分”两个环节做不好模型精度再高也是纸上谈兵。模型结构采用了一维卷积神经网络1D CNN联合双向长短期记忆网络BiLSTM再加注意力机制Attention的混合结构。之所以不直接用传统特征工程加机器学习分类器是因为手动特征如RR间期、QRS宽度、ST段偏移等设计成本很高而且对信号质量敏感。端到端深度模型能在数据量足够的条件下直接从原始波形中学到判别性特征。这一点放到第3节详细讲。运行环境是Python 3.8深度学习框架用PyTorch信号处理主要依赖NumPy、SciPy数据读取用wfdb库。依赖清单在项目的requirements.txt里都列好了。1.3 适合谁来参考这套方案如果你属于下面任意一类这套源码和说明对你有实际价值生物医学工程方向的学生刚开始做深度学习分类任务需要一个能完整跑通的ECG基线项目做辅助诊断或可穿戴设备异常检测的工程师想快速对比不同模型结构的效果做时间序列分类的研究者想借鉴CNNBiLSTMAttention的混合结构ECG只是你的验证数据集。接下来我直接用实际运行过的代码和结果做讲解从数据到模型到评估每一步都可以复现。2. 数据层面心电信号的分割与预处理2.1 波形为什么要先分割再进模型原始心电信号是连续的长序列一段半小时的采样数据有几十万个采样点。直接把这些点全塞进模型即使能塞进去模型也几乎学不到每个搏动的细粒度特征。所以业内通用做法是先把心跳从长信号里切出来以每个心搏的R峰为对齐中心截取固定长度的窗口作为单个样本。做过语音识别或者事件检测的朋友应该很熟悉本质上就是“先把事件切出来再做小样本分类”。切得准不准直接决定模型看到的输入是不是对齐好的波形。切歪了即使同样的心律失常类型波形相位也对不上模型学不到稳定特征。项目里取的窗口长度是R峰前90个采样点、R峰后180个采样点一共270个采样点。这个窗口对应360Hz采样率下的0.75秒正常窦性心律的一个心动周期大约0.8秒所以窗口能覆盖完整的P-QRS-T波群又尽量避免引入相邻心跳的干扰。2.2 去噪与去除基线漂移心电信号在采集时常见干扰包括工频干扰50Hz、肌电干扰、电极移动造成的基线漂移。直接拿带噪声的信号去切搏动再丢进模型模型会把噪声当成特征去学一开始实验准确率可能还行但泛化能力会很差。项目里用带通滤波处理保留0.5Hz到45Hz的频率成分。低截至0.5Hz是为了去掉基线漂移高截至45Hz是为了去掉工频以及更高频的肌电干扰。滤波器用的是二阶巴特沃斯配合filtfilt做零相位滤波防止滤波过程引入相位偏移把波形弄变形。from scipy.signal import butter, filtfilt def bandpass_filter(signal, fs360, low0.5, high45): nyquist fs / 2 b, a butter(2, [low / nyquist, high / nyquist], btypeband) return filtfilt(b, a, signal)为什么不用中值滤波只做基线漂移中值滤波确实能去基线但对QRS波群的高频成分有一定切削效果处理不好会让R峰变钝影响后续分割准确性。带通滤波一步到位简单稳定。如果数据里存在严重工频干扰比如在普通病房环境下采集的可以在45Hz低通的基础上再加一个50Hz陷波器。MIT-BIH数据相对干净没有额外加。注意filtfilt要求信号不能有NaN否则直接报错。如果用的是自己采集的数据先检查信号里有没有缺失值。2.3 R峰检测与心跳分割R峰检测采用经典的Pan-Tompkins算法核心步骤是带通滤波、差分、平方、滑动窗口积分最终定位出每个QRS波群中的R波峰值位置。这个算法原理不难但实现时很考验细节窗口长度和阈值设置不对就容易多检或漏检。项目里直接调用wfdb库处理效率和准确率都不错。拿到R峰位置之后遍历所有R峰对每个R峰截取固定窗口内的波形同时把对应标注映射成5类标签。这是整个数据管线里最需要细心的一段逻辑。def extract_qualified_heartbeats(signal, r_peaks, labels, fs360, left90, right180): samples_list, label_list [], [] for i, peak in enumerate(r_peaks): start int(peak - left) end int(peak right) if start 0 or end len(signal): continue samples_list.append(signal[start:end]) label_list.append(labels[i]) return np.array(samples_list), np.array(label_list)2.4 数据划分打死都不能随机切网上很多分类项目喜欢把所有样本shuffle后按8:2划分训练测试集最后准确率很高。但这种做法在ECG项目里是致命的。同一个患者的心跳高度相似如果同一个人的心搏同时出现在训练集和测试集模型训练时已经“见过”这个人的波形形态测试时再见它拼的是记忆力而不是泛化能力。正确做法是把患者也就是原始record作为划分单位训练集出现的患者测试集不能出现。项目里用的是按record划分策略随机选取80%的record作为训练集剩余20%作为测试集。这样一来模型面对的是从未见过的心电形态评估结果才真正反映实际部署场景。records list(record_ids) random.seed(42) random.shuffle(records) split_idx int(len(records) * 0.8) train_records records[:split_idx] test_records records[split_idx:] def partition_by_records(samples, sample_records): train_idx [i for i, r in enumerate(sample_records) if r in train_records] test_idx [i for i, r in enumerate(sample_records) if r in test_records] return train_idx, test_idx经验之谈这个看似基础的划分方式直接影响最终结果的真实性。我见过太多项目把随机划分的准确率当成果汇报一上真实场景就露馅问题基本都出在这里。3. 模型结构设计从卷积到注意力3.1 为什么选择原始波形直接输入传统ECG分类方案先提取特征再交给SVM或随机森林。特征工程师会设计RR间期、QRS宽度、QT间期、T波形态等指标。这种方案在小样本、标注规范的数据集上很稳但有两方面局限一是特征设计需要很强的领域知识不同疾病的特征权重差异巨大二是特征提取过程容易漏掉波形中细微但关键的形态变化。把原始波形直接作为模型输入相当于让模型自己判断“哪些波形片段最有区分度”。CNN擅长提取局部形态特征比如QRS波群的尖锐程度、ST段的偏移趋势BiLSTM擅长捕捉时序依赖关系注意力机制负责找到“这个样本中最关键的几个时间点”。三种结构组合正好覆盖形态、时序、关键片段三个层面的信息。我在跑这个模型时对比过只保留CNN的版本去掉BiLSTM和Attention之后准确率掉了3到4个百分点。混合结构不是炫技确实有用。3.2 混合模型逐层拆解完整模型结构如下输入层270个采样点的一维序列shape为(batch, 270)。卷积块1Conv1d(1, 32, kernel_size15, padding7) → BatchNorm → ReLU → MaxPool1d(2)。卷积块2Conv1d(32, 64, kernel_size9, padding4) → BatchNorm → ReLU → MaxPool1d(2)。卷积块3Conv1d(64, 128, kernel_size5, padding2) → BatchNorm → ReLU → MaxPool1d(2)。BiLSTM128维输入64维隐藏状态双向输出为(batch, seq_len, 128)。注意力池化对BiLSTM每个时间步的输出计算权重加权求和得到128维上下文向量。全连接层128 → 5接Softmax输出类别概率。为什么卷积核尺寸从15逐渐减到5因为大卷积核感知野大能捕捉P波、T波这类持续时间较长的整体形态随着网络加深小卷积核能更精细地捕捉QRS波群内部的局部变化。这是多尺度特征提取的思路。池化步长取2三层池化后序列长度从270降到大约34个时间步这个长度对BiLSTM来说非常友好。批归一化放在卷积层之后、激活函数之前有两个作用加速收敛以及应对信号幅度差异较大的情况。即使做了标准化真实中心电信号的幅度仍有波动BN能把波动压到相对稳定的量级。import torch import torch.nn as nn import torch.nn.functional as F class Attention(nn.Module): def __init__(self, hidden_size): super().__init__() self.attn nn.Linear(hidden_size, 1) def forward(self, lstm_out): # lstm_out: (batch, seq_len, hidden_size) scores self.attn(lstm_out).squeeze(-1) # (batch, seq_len) weights F.softmax(scores, dim1) context torch.bmm(weights.unsqueeze(1), lstm_out).squeeze(1) return context class ECGMixedModel(nn.Module): def __init__(self, num_classes5, input_len270): super().__init__() self.conv_block1 nn.Sequential( nn.Conv1d(1, 32, kernel_size15, padding7), nn.BatchNorm1d(32), nn.ReLU(), nn.MaxPool1d(2) ) self.conv_block2 nn.Sequential( nn.Conv1d(32, 64, kernel_size9, padding4), nn.BatchNorm1d(64), nn.ReLU(), nn.MaxPool1d(2) ) self.conv_block3 nn.Sequential( nn.Conv1d(64, 128, kernel_size5, padding2), nn.BatchNorm1d(128), nn.ReLU(), nn.MaxPool1d(2) ) self.lstm nn.LSTM(128, 64, batch_firstTrue, bidirectionalTrue) self.attention Attention(128) self.dropout nn.Dropout(0.3) self.fc nn.Linear(128, num_classes) def forward(self, x): if x.dim() 2: x x.unsqueeze(1) x self.conv_block1(x) x self.conv_block2(x) x self.conv_block3(x) x x.permute(0, 2, 1) # (batch, seq_len, channels) lstm_out, _ self.lstm(x) context self.attention(lstm_out) out self.dropout(context) return self.fc(out)3.3 类别不平衡与损失函数设计心电数据集中正常搏动的数量远大于异常搏动。以MIT-BIH为例五大类中N类几乎占90%左右F类和Q类加起来可能不到2%。如果直接拿CrossEntropyLoss去训练模型会倾向于把所有样本预测成N类因为在训练集上这样就已经有接近90%的准确率但模型对异常搏动的识别能力几乎为零。项目里做了两件事。第一给少数类分配更高的类别权重CrossEntropyLoss的weight参数按类别样本数倒数归一化。第二在训练时对少数类做在线数据增强包括加高斯噪声、随机幅度缩放、轻微时间平移提高少数类样本的多样性。class_counts np.bincount(train_labels, minlength5) weights 1.0 / class_counts weights weights / weights.sum() * 5 weights torch.tensor(weights, dtypetorch.float32).to(device) criterion nn.CrossEntropyLoss(weightweights)这个改动在实测里很关键。不加类别权重时V类的F1分数大约在0.6左右加权重之后能到0.78左右。在严重不平衡的数据集上类别加权损失的作用非常直接。4. 训练实验与结果分析4.1 训练配置与超参数选择训练阶段用Adam优化器初始学习率1e-3weight_decay设为1e-4。batch size取256在RTX 3090上训练一个epoch大概几秒。整个训练跑60个epoch配置了ReduceLROnPlateau学习率调度器验证损失连续5个epoch不下降时学习率减半同时使用early stopping验证指标连续8个epoch不提升就停止训练。有一个容易被新手忽略的点卷积层和BiLSTM都对输入数据的尺度敏感所以训练前测试数据的预处理必须和训练数据完全一致。我在项目里把归一化参数均值和标准差在训练集上计算好并保存预测阶段直接使用训练集的统计量去归一化而不是在预测时重新计算。否则数据分布不一致性能会明显缩水。4.2 评估指标怎么看准确率只是其中一个维度。在类别不平衡数据集上必须辅以每个类别的精确率、召回率和F1分数来看。F类和Q类样本极少即使整体准确率高这两个类别也可能几乎不被正确识别。项目里每轮实验结束都会打印classification_report逐类检查。混淆矩阵是排查错误的利器。测试集上画了混淆矩阵热力图能很直观看到哪些类之间容易混淆。通常S类和V类、N类和F类之间最容易搞混因为波形形态上确实存在重叠。4.3 实测效果与指标对比在按患者划分的测试集上也就是测试集患者从未出现在训练集中模型整体分类效果大致如下类别精确率召回率F1分数测试样本数N0.980.970.97约18000S0.780.790.78约600V0.820.840.83约1400F0.650.580.61约150Q0.700.620.66约80整体准确率在95%左右但这个数字参考意义有限真正关键的是少数类的F1分数。F类样本太少提升空间还很大后续可以考虑扩充F类训练数据或者在损失函数中进一步加大F类的权重。这里有个细节表中样本数是测试集中该类别的搏动个数可以看到N类占据绝对主导Q类只有约80个。所以训练过程中我做了分层采样保证每个batch里少数类有一定占比而不是完全随机采样。5. 项目源码结构与使用说明5.1 目录结构与模块职责压缩包解开后的目录结构大致如下ecg_5class/ ├── README.md ├── requirements.txt ├── config.py ├── data/ │ ├── raw/ # 原始MIT-BIH记录 │ ├── processed/ # 预处理后的心跳样本npy │ └── record_split.json # 按record划分的训练/测试清单 ├── models/ │ ├── ecg_model.py # 模型结构定义 │ └── weights/ │ └── best_model.pth # 训练好的权重文件 ├── scripts/ │ ├── preprocess.py # 信号预处理心跳分割 │ ├── train.py # 训练入口 │ ├── test.py # 评估入口 │ └── predict.py # 对单条心电信号做预测 └── utils/ ├── dataset.py # Dataset与DataLoader封装 ├── metrics.py # 评估指标计算 └── visualization.py # 波形与混淆矩阵可视化config.py里集中管理所有超参数包括采样率、窗口长度、数据集路径、类别数量、训练epoch数、学习率等。改参数不用到处翻代码这一点对反复实验很有帮助。5.2 从零开始跑通的完整步骤如果你想从原始数据开始完整复现训练流程按下面的顺序执行# 1. 安装依赖 pip install -r requirements.txt # 2. 数据预处理生成processed目录下的心跳样本 python scripts/preprocess.py # 3. 启动训练 python scripts/train.py # 4. 评估训练好的模型 python scripts/test.py # 5. 对新的心电图信号做预测 python scripts/predict.py --signal_path data/raw/xxx.datpreprocess.py内部会自动调用wfdb库读取记录完成带通滤波、R峰检测、心搏分割、标签映射、归一化最后把样本保存成npy格式同时生成record_split.json记录哪些record用于训练、哪些用于测试。这一步跑完数据准备阶段就结束了。train.py里除了训练循环之外还会在每个epoch结束后做验证集评估记录最佳验证损失和最佳F1分数把效果最好的模型权重保存到weights/best_model.pth。日志打印包括loss、accuracy、每类F1便于实时掌握训练状态。predict.py可以直接加载训练好的权重对一段新的心电信号先做同样的预处理流程再调用模型输出5类概率分布。从这里可以很方便地扩展成接口服务为实时推理做准备。5.3 关键实现说明Dataset封装时需要注意预处理好的是全部样本但划分必须按照record来进行。项目里在utils/dataset.py做了一层包装dataset初始化时接收样本索引训练集和测试集各自持有不同的索引数组底层数据是同一份避免重复读盘。另外训练过程中做了动态数据增强。对每个样本以一定概率加高斯噪声或做小幅尺度缩放这有助于缓解过拟合尤其对少数类样本来说避免模型死记硬背那几个样本的原始波形。6. 常见问题与排查技巧实录6.1 训练不收敛或Loss震荡这个问题在调参期间遇到不少次。如果训练Loss不下降先检查输入数据是否存在NaN或异常大值。ECG数据里偶尔会出现电极脱落产生的巨大尖峰如果不处理梯度更新直接崩掉。项目里预处理阶段对幅值超过一定阈值的采样点做截断或剔除。Loss震荡的另一个原因是学习率太大尤其当模型里同时有卷积和LSTM时LSTM对学习率很敏感。用LRScheduler或者把Adam的学习率降到5e-4再试。6.2 准确率高但少数类几乎全错出现这种迹象基本可以断定是两类问题。第一没有做类别加权模型整体预测偏向多数类。第二数据划分时把同一个患者样本同时放进了训练集和测试集模型在测试集上表现亮眼但真实场景是陌生患者效果自然崩塌。排查方式看混淆矩阵。如果几乎所有样本都预测成N类说明类别不平衡没处理好。如果每个类都能预测但准确率虚高那多半是数据泄漏。6.3 预测阶段和训练阶段效果差异大训练和测试性能差异大首先检查预处理流程是否一致。比如训练时用了带通滤波预测时忘了加或者归一化的均值和标准差用了测试集的统计量。我在项目里把预处理逻辑统一封装在同一个函数里无论训练还是预测都走同一条代码路径从根本上避免这类不一致。另外注意模型权重的dropout模式。推理时必须调用model.eval()切换到评估模式否则dropout仍在工作输出结果会带有随机性多跑几次结果都不一样。6.4 环境依赖问题PyTorch版本和wfdb库版本对代码有影响。requirements.txt里固定了主要依赖版本底线建议在干净的虚拟环境里安装。如果你用的不是MIT-BIH而是自采数据注意采样率可能不是360Hz需要同步修改带通滤波的归一化频率和窗口截取的采样点数。最后再分享一点个人经验关于这套源码我最想强调的其实是数据处理那部分。模型结构可以改、可以抄但数据划分和预处理细节才决定你复现的结果和原作者差多少。我一开始图省事随机切数据结果准确率比按患者切高了好几个点还以为是模型写得好后来才发现是幻觉。做ECG分类数据诚实是第一原则。如果你准备基于这份源码做自己的实验建议先从跑通现有的数据管线开始再逐步替换成自己的数据集和模型结构这样排查问题时能有一条清晰的基准线。本文还有配套的精品资源点击获取