ARTICLE DETAIL

资讯详情

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

基于MIT-BIH数据库与2D CNN的心律不齐心拍分类实现

基于MIT-BIH数据库与2D CNN的心律不齐心拍分类实现 简介基于MIT-BIH数据库与二维卷积神经网络实现的心律不齐诊断算法源码项目采用二维CNN作为核心模型围绕八种心率不齐类型的自动识别问题完整覆盖从心电数据读取、预处理、二维特征构造、模型搭建到训练评估的整个流程。项目面向计算机、人工智能、大数据、数学、电子信息等专业学生适合用于课程设计、期末大作业或毕业设计也适合相关技术学习者作为参考要求具备Python和深度学习基础知识。压缩包共8个文件包含5个Python脚本分别承担数据加载、预处理、CNN模型定义、训练和主流程、一个心电数据压缩包、一份依赖清单及一份项目说明文档整体大小仅1.11MB轻巧易部署Python脚本按模块化思想组织便于按需阅读和二次修改。已有64人浏览学习可作为心电信号处理与卷积网络结合的入门实例。资源提供已调试可运行的完整源码配合随附的说明文档和依赖清单可快速搭建环境并复现八种心率不齐分类实验不仅便于理解算法实现细节也为后续模型优化或迁移学习提供了便捷的代码基础具有较好的可读性和可维护性。1. 心率不齐分类为什么落在2D CNN身上重症监护室的桌面心电图机每一秒都在推波形医生必须在几分钟里判断心拍性质。同样是早搏房性、室性、交界性的处理路径完全不一样。标题里的 mit-hib 数据库实际就是 MIT-BIH 心律失常数据库48 份带专家标注的心电记录是心拍分类最常用的基准2 维卷积神经网络则指明了实现路线不直接把一维波形喂给时序模型而是把单个心拍转成时频二维图再用 CNN 完成八种心率不齐的分类诊断。下面按数据读取、预处理、2D 表示、模型训练、工程部署的顺序把一套可复现源码的算法链路讲清楚。这篇内容面向医疗算法和可穿戴心电团队重点解决源码能跑但指标不严谨、模型一换设备就失效这两类问题。2. 从MIT-BIH数据库取数到构建训练集的关键步骤2.1 用 wfdb 读取记录先认清数据格式和标注结构MIT-BIH 心律失常数据库在 PhysioNet 上以记录为单位存放每条记录包含 .dat 信号文件、.hea 头文件和 .atr 标注文件总共 48 条每条时长约 30 分钟采样率 360 Hz。读取这类数据用 Python 生态里的 wfdb 库最省事一条rdrecord加一条rdann就能把波形和心拍标注同时拿到。import wfdb record wfdb.rdrecord(mitdb/100, sampfrom0, sampto360 * 60 * 30) annot wfdb.rdann(mitdb/100, atr, sampfrom0, sampto360 * 60 * 30) ecg record.p_signal[:, 0] # 取第一导联通常是 MLII fs record.fs # 360 Hz r_pos annot.sample # R 峰在信号中的采样点索引 sym annot.symbol # 对应心拍的类型符号sampfrom和sampto控制读取的起止采样点这里的360*60*30表示读取前 30 分钟正好覆盖完整记录。p_signal按列存放导联数据取第几列要先看头文件里的导联顺序annot.sample是标注的 R 峰位置也是后面切窗的中心annot.symbol是心拍类型符号。读取后先检查符号集合因为符号很多N代表正常L左束支传导阻滞R右束支传导阻滞V室性早搏A房性早搏/起搏心拍!心室颤动E交界性逸搏。标题要求的八种心率不齐一般就从这些符号里选。记录末尾偶尔会出现融合波、~信号质量差一类符号这类样本要么过滤要么单独处理直接扔进训练集只会污染边界样本。2.2 过滤基线漂移与工频干扰一次带通滤波就够了原始心电信号里有基线漂移和工频干扰。基线漂移频率通常在 0.5 Hz 以下工频干扰集中在 50/60 Hz。多数研究用的是一阶选择带通 0.545 Hz 的巴特沃斯滤波一次处理同时压掉低频和高频保留 P 波、QRS 波群和 T 波的主要形态。滤波必须用零相位的方式避免普通 IIR 滤波器带来的相位延迟让 R 峰位置发生几十毫秒的偏移。filtfilt先正向再反向各过一次滤波器相位延迟相互抵消波形上不会出现肉眼可见的位移。from scipy import signal b, a signal.butter(2, [0.5, 45], fsfs, btypebandpass, outputba) ecg_f signal.filtfilt(b, a, ecg)这里用 2 阶巴特沃斯[0.5, 45]是通带范围。阶数越高带外衰减越快但阶数到 4 以上在边缘端嵌入式平台实现的成本明显上升算法验证先用 2 阶足够。如果发现滤波后波形细节异常比如 QRS 顶部变钝问题多半出在 45 Hz 截止频率过低而不是滤波本身可以改成[0.5, 60]再对比一次结果。2.3 以R峰为中心切窗得到训练样本清单拿到干净波形后开始切窗。切窗要回答两个问题窗子多长以及以谁为中心。常见做法是以标注的 R 峰为中心R 峰前取 120 个采样点约 0.33 秒覆盖 P 波和 PR 段R 峰后取 240 个采样点约 0.67 秒覆盖 QRS 波群和 T 波。总窗长 360 个采样点正好 1 秒。win_before, win_after 120, 240 CLASS_MAP {N: 0, L: 1, R: 2, A: 3, V: 4, /: 5, !: 6, E: 7} windows [] labels [] for idx, s in zip(r_pos, sym): if s not in CLASS_MAP: continue start, end idx - win_before, idx win_after if start 0 or end len(ecg_f): continue windows.append(ecg_f[start:end]) labels.append(CLASS_MAP[s])这段逻辑把标注符号映射到 07 的整数标签只保留窗口完整的心拍最后得到平行数组windows和labels。这里的/对应起搏心拍后文可读标签记作 P!对应室颤 VFE对应交界性逸搏方便在评估报告里显示。照这个流程处理完整数据库会得到一个类别严重不均衡的样本集。正常心拍数量占绝对多数心室颤动和交界性逸搏可能只有几百份下表是分类时的主要注意点。类别标注符号分布特点分类难点正常N绝对多数P 波存在但形态多变容易与房早混淆左束支阻滞L约 8%QRS 明显增宽需要抓住左束支特征右束支阻滞R约 7%与左束支同属宽 QRS容易互相干扰房性早搏A约 2.5%P 波提前出现且形态异常差异微小室性早搏V约 7%QRS 宽大畸形前面没有 P 波起搏心拍/约 7%起搏脉冲尾迹干扰特征提取心室颤动!极少波形无规律样本本身难以定义交界性逸搏E极少样本最少最容易被模型忽略切窗阶段先不急着做扩增保持原始分布跑一次基线能直观看到模型天然偏向哪几类。不均衡的修正手段放到训练阶段处理切窗阶段只负责把样本质量做扎实。3. 把一维心电信号变成2D输入喂给卷积神经网络3.1 为什么用时频图而不是直接把波形矩阵折叠拿到 360 点的一维波形后一个直觉做法是直接reshape(30, 12)或reshape(18, 20)塞给二维卷积层。但这样有逻辑问题二维卷积核此时感受到的相邻像素在原始时间轴上可能并不是相邻采样点。把 360 点排成 30×12原本时间上连续的第 20、21 个采样点会被折到第一行末尾和第二行开头卷积核同时看到它们却不知道它们在时间上只是前后脚而真正时间邻近的信息被矩阵结构打散了。换成时频图就顺理成章横轴是时间纵轴是频率每个像素强度表示该时刻该频率成分的能量。QRS 波群的能量集中在 1040 Hz 频段T 波能量集中在 15 Hz 频段二者在时频图上天然分居不同区域卷积核可以在时间和频率两个维度上分别学习局部模式既保留时间顺序又不丢失形态细节。这是 2 维卷积神经网络做心拍分类最常见的实现路线先做时频变换再做图像分类。3.2 用短时傅里叶变换生成单通道频谱图具体实现用scipy.signal.stft最直接。以一段 360 点的滤波后心拍为例汉宁窗窗口长度 128 点重叠 96 点FFT 点数 256得到 129×8 的复数矩阵取模并归一化后就是单通道频谱图。这个原始尺寸直接进网络太扁在送入 CNN 前用双线性插值缩放到 64×64。from scipy import signal from scipy.ndimage import zoom import numpy as np def beat_to_spec(beat, fs360): f, t, Zxx signal.stft(beat, fsfs, windowhann, nperseg128, noverlap96, nfft256) spec np.abs(Zxx) spec spec / (spec.max() 1e-8) spec zoom(spec, (64 / spec.shape[0], 64 / spec.shape[1]), order1) return spec.astype(np.float32)函数输入beat是长度 360 的滤波后心拍输出是 64×64 的浮点矩阵。nperseg128让每帧包含约 0.36 秒的信号对应频率分辨率约 2.8 Hznoverlap96表示 75% 的重叠率让时间轴帧与帧之间平滑过渡nfft256把频率分辨率细化到 129 个频点。zoom的双线性插值负责把 129×8 规整到 64×64这个尺寸是工程折中再大训练慢容易过拟合再小丢失纹理细节。参数值效果与注意nperseg128每帧约 0.36 秒越小时间分辨率越高频率分辨率越低noverlap9675% 重叠过小会导致时间轴帧数不足帧间不连续nfft256零填充到 256得到 129 个频点不影响时间分辨率缩放尺寸64×64兼顾训练速度与纹理信息可调到 96×96 换取精度另一个常用备选是连续小波变换它对低频段频率分辨率更好但计算成本高且没有stft这么开箱即用。如果目标是边缘端部署STFT 更容易在 DSP 或 FPGA 上落地CWT 更适合离线研究场景。3.3 直接把360点折叠成30×12矩阵的兜底方案如果暂时不引入时频变换直接折叠波形矩阵也能做快速验证但有两个前提先把每个窗口单独做 z-score 归一化避免心电振幅差异主导分类再保证训练和推理时使用相同的折叠方向。def beat_to_image(beat): seg beat[120:480] seg (seg - seg.mean()) / (seg.std() 1e-8) return seg.reshape(30, 12).astype(np.float32)beat[120:480]取以 R 峰为中心的 360 点seg.std() 1e-8防止全零段除零。这个方案的时间邻接关系被矩阵折行破坏了前面已经解释过。它存在的意义是当目标推理环境缺少scipy.fftpack等依赖时少一个依赖就少一个部署坑。卷积层照样能学到分布统计特征但分类性能通常低于时频图方案只适合做基线对比不建议当最终交付物。4. 构建8分类2D CNN模型与训练管线4.1 网络结构三层卷积加批归一化的设计考量2D CNN 的输入是 64×64 的单通道图输出 8 个心拍类别的概率。网络结构不需要太深心电时频图是纹理类图片没有 ImageNet 那种层级语义。常见可靠的结构是三层卷积第一层 32 个 3×3 卷积核提取边缘和波形成分第二层 64 个卷积核组合局部纹理第三层 128 个卷积核捕捉类别特定的高频细节。每层后面都带 BatchNorm不光是加速收敛更重要的是缓解心拍样本间振幅差异导致的内部协变量偏移。import torch import torch.nn as nn class BeatCNN(nn.Module): def __init__(self, num_classes8): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, 3, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(64, 128, 3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.fc nn.Sequential( nn.Flatten(), nn.Dropout(0.5), nn.Linear(128 * 8 * 8, 256), nn.ReLU(inplaceTrue), nn.Linear(256, num_classes), ) def forward(self, x): return self.fc(self.features(x))三次池化后特征图从 64×64 降到 8×8所以全连接层输入维度是128*8*8。padding1让卷积不改变空间尺寸尺寸变化只由 MaxPool 决定。池化窗口 2×2 对心拍在窗内位置不完全一致的情况有平移容忍度。Dropout 放在全连接层前训练时随机丢掉一半神经元是八分类任务里对抗过拟合的直接手段。注意不要在模块里手动接 softmax最后一层输出的 logits 直接喂给损失函数即可CrossEntropyLoss内部会做 softmax。4.2 PyTorch训练脚本与超参数设定训练超参数直接决定模型能不能落地。优化器选 AdamW理由是它把权重大小与梯度解耦实测比 Adam 对学习率的敏感度更低。由于样本类别不均衡交叉熵必须带类别权重学习率从 1e-3 起步用余弦退火在计划轮数内衰减。from torch.utils.data import DataLoader, TensorDataset import numpy as np X np.stack(all_specs) # shape: [N, 64, 64] X X[:, None, :, :] # 扩展通道维 - [N, 1, 64, 64] y np.array(all_labels) dataset TensorDataset(torch.from_numpy(X).float(), torch.from_numpy(y).long()) loader DataLoader(dataset, batch_size256, shuffleTrue, num_workers4, pin_memoryTrue) counts np.bincount(y, minlength8).astype(np.float32) class_weights torch.tensor(counts.sum() / (8 * counts), dtypetorch.float32) criterion nn.CrossEntropyLoss(weightclass_weights) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max60, eta_min1e-5)batch_size256在单块 12GB 显存的 GPU 上够跑显存小就降到 128反向传播的梯度稳定性不会明显变差。class_weights按每个类别的样本数取倒数再归一化极小类计数接近 0 时权重会被放大这是预期的目的是让梯度在少数类上产生更大更新。T_max60和计划训练轮数保持一致让学习率在完整周期内从 1e-3 平滑降到 1e-5。准备数据时最容易报错的地方是张量维度ceil后的数组是[N, 64, 64]必须显式加一维成[N, 1, 64, 64]Conv2d 才能正确处理。4.3 类别不均衡的两种处理损失加权和过采样并用训练时建议同时对次要类过采样其一是少数类在全部数据中的占比心室颤动和交界性逸搏可能不足 2%单靠损失加权会让模型在交叉熵的细节上过度拟合其二是过采样要放在训练集内部依据记录分组进行否则同一患者相邻心拍会跨越训练验证边界。counts np.bincount(y_train, minlength8) rare_classes np.where(counts counts.sum() * 0.03)[0] for idx in rare_classes: idx_rare np.where(y_train idx)[0] repeat int(counts.max() / counts[idx]) idx_rare_expanded np.tile(idx_rare, repeat)过采样时对少数类心拍顺序做一次性 shuffle 后再复制避免同一份波形在一个 batch 里大量紧邻出现。配合损失权重后模型在少数类上的召回率通常能提高 1020 个百分点代价是训练时间变长。验证阶段不要再看总准确率宏平均 F1 才是八分类的体检指标。5. 训练中必须盯紧的验证指标与常见坑5.1 按记录划分数据集心拍样本不能乱拆做心拍级分类最容易犯的错误是直接把所有心拍塞进train_test_split(X, y, test_size0.2)。同一个患者相邻几个心拍的波形几乎一样随机拆分会让训练集里某个心拍的“近亲”出现在验证集验证指标虚高。正确做法是按记录分组从 48 条记录里先选训练记录再选验证和测试记录每条记录内的窗口不会跨越数据集边界。train_recs [100, 101, 103, 105, 106, 111, ...] valid_recs [208, 209, 215, 220, ...] test_recs [217, 219, 222, ...] X_train, y_train build_from_records(train_recs) X_val, y_val build_from_records(valid_recs) X_test, y_test build_from_records(test_recs)build_from_records内部做的就是第 2.3 节切窗的逻辑只是换成遍历一个记录列表。关键是先分记录再切心拍这条顺序一旦颠倒后续所有调参都建立在不可靠的验证指标上。测试集至少保留 7 条记录以获得相对稳定的统计结果如果验证集恰好落在与训练分布差异明显的记录上F1 会突然掉一截那不是模型坏了而是泛化边界被真实暴露出来了。5.2 早停要盯着验证损失的移动平均而不是每个epoch的瞬时值八分类的验证损失曲线会因为个别类别样本分布波动出现锯齿一个 epoch 涨不代表模型在变坏。保留最近三轮的验证损失均值作为早停依据连续 10 轮均值不下降才停同时保存验证宏 F1 最高的权重。best_f1, wait 0.0, 0 for epoch in range(60): train_one_epoch(model, loader, criterion, optimizer) v_loss, v_f1 validate(model, val_loader) if v_f1 best_f1: best_f1 v_f1 torch.save(model.state_dict(), best_model.pth) wait 0 else: wait 1 if wait 10: print(fearly stop at epoch {epoch}, best F1 {best_f1:.4f}) break scheduler.step()如果早停触发得太早第一件事不是调大忍耐轮数而是检查学习率是否衰减过快。学习率降到 1e-5 以下后损失变化很小早停会把这种平台期误判为不收敛。注意scheduler.step()放在每个 epoch 最后中断时刻的权重未必是最优的最终取的是best_model.pth。5.3 用混淆矩阵排查 L/R 与 A/N 这两类经典混淆训练结束后只报一个总准确率没有意义必须看每一类的查准率、召回率和 F1。sklearn.metrics里现成的工具一次算完。from sklearn.metrics import classification_report, confusion_matrix def collect_pred_labels(model, loader): model.eval() preds, truths [], [] with torch.no_grad(): for x, y in loader: logits model(x) preds.extend(logits.argmax(dim1).cpu().numpy()) truths.extend(y.cpu().numpy()) return np.array(truths), np.array(preds) y_true, y_pred collect_pred_labels(model, test_loader) print(classification_report(y_true, y_pred, target_names[N, L, R, A, V, P, VF, JE])) cm confusion_matrix(y_true, y_pred)看纵坐标是真实类别、横坐标是预测类别的混淆矩阵最常出现两处集中错误左束支阻滞和右束支阻滞互相串房性早搏和正常心拍分不清。左右束支都体现为 QRS 宽大畸形只有时频图的高频细节和束支特征能区分房早和正常心拍的差异在一个提前出现的 P 波上P 波幅度小容易被频谱能量归一化抹平。针对性处理是先调整 STFT 参数让低频分辨率更细再把 P 波片段单独作为辅助输入分支这两招都比无脑加深网络有效。6. 模型导出与批量推理把八分类模型落到真实设备6.1 用 TorchScript 固化模型保持预处理与训练完全一致训练完成的BeatCNN在 PyTorch 里只是动态图对象。落地前先把它固化成不依赖 Python 的 TorchScript 格式再把滤波系数、窗口长度、STFT 参数全部写进同一个推理类里。推理环境最容易出问题的地方不是网络本身而是预处理与训练时不一致比如训练时做了 filtfilt推理脚本却用 lfilter训练时窗口取 360 点推理时丢尾取 359 点训练时做了 z-score推理时忘了做。任何一处不一致都会直接改变输入分布。class InferencePipeline(torch.nn.Module): def __init__(self, cnn): super().__init__() self.cnn cnn self.b, self.a signal.butter(2, [0.5, 45], fs360, btypebandpass, outputba) self.win_before, self.win_after 120, 240 def forward(self, raw_ecg): ecg raw_ecg.cpu().numpy() ecg_f signal.filtfilt(self.b, self.a, ecg) beats [] for rpos in detect_r_peaks(ecg_f): seg ecg_f[rpos - self.win_before: rpos self.win_after] if len(seg) 360: beats.append(beat_to_spec(seg)) if len(beats) 0: return torch.zeros(0, 8, dtypetorch.float32) x torch.from_numpy(np.stack(beats))[:, None, :, :] with torch.no_grad(): logits self.cnn(x) return logitsdetect_r_peaks是训练阶段绕开的环节训练用的是数据库标注位置实际部署必须自己检测。一个能跑通的方案是对滤波后信号做scipy.signal.find_peaks限制最小峰高为信号标准差的 0.6 倍最小峰间距 200 个采样点。R 峰检测的误差超过 20 个采样点时切出的窗口会偏离形态中心后续分类精度再高也救不回来。R 峰检测模块最好单独做一次评估用正确检测率而不是肉眼观察来验收。固化并导出pipeline InferencePipeline(model.eval()) scripted torch.jit.script(pipeline) torch.jit.save(scripted, beat_classifier.pt) torch.jit.save(pipeline, beat_classifier_full.zip)导出后把模型权重、预处理脚本、推理样例、版本说明和八个类别的标签定义一起整理进压缩包这正好对应标题里的交付形态。最后一步是用一条完全没有参与训练的新记录跑完整推理打印出每个心拍的预测类别与该记录原始标注的对比确认八类的分布比例和混淆情况与论文阶段一致整套八种心率不齐诊断算法才算真正交付完成而不是止步于训练集上的漂亮曲线。本文还有配套的精品资源点击获取
返回列表