ARTICLE DETAIL

资讯详情

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

基于12导联心电图与深度学习的ECG分类实战:从数据预处理到模型部署

基于12导联心电图与深度学习的ECG分类实战:从数据预处理到模型部署 简介这份资源是面向深度学习课程大作业场景的完整项目源码适合具备Python与神经网络基础、希望将算法落地到医疗信号分析的学习者。项目围绕12导联心电图数据展开目标是实现心脏疾病的自动诊断覆盖从数据加载、模型训练到性能测试的全流程。压缩包共7个文件包含5个py脚本、1个md说明文档和1个ipynb笔记本整体约74KB其中py文件分别承担训练、ResNet18网络定义、数据集处理、ViT模型与测试评估等职责md提供使用说明ipynb则整合了CNN、LSTM与注意力机制的综合实验。已有114人学习该资源。读者可借此获得一套可运行的课程设计参考方案理解卷积网络、循环网络、注意力机制与视觉转换器在心电信号上的实现差异并掌握数据预处理、数据增强、模型评估指标计算等关键环节便于快速搭建自己的实验基线或在此基础上做改进。1. 12导联心电图 深度学习一份课程大作业源码到底能跑出什么拿到「深度学习课程大作业-基于12导联心电图心脏疾病诊断python源码.zip」这个标题多数人的第一反应是找一份能直接跑的代码交差。但真正动手跑过心电图ECG分类的人都知道这类作业的坑不在模型结构而在数据本身——12导联信号怎么读、采样率怎么统一、标签怎么对齐、类别不平衡怎么处理任何一步翻车模型准确率都会掉到随机猜测的水平。这份源码对应的典型场景是用PTB-XL或CPSC2018这类公开12导联心电图数据集训练一个CNN或CNNLSTM的深度学习模型输出正常/房颤/心肌缺血等心脏疾病分类结果。适合正在做课程大作业的本科生、刚入门深度学习想找一个完整信号处理项目的工程师以及需要快速验证ECG分类baseline的研究者。下面从数据到推理把这条链路拆开讲清楚。2. 12导联心电图数据怎么读、怎么切、怎么对齐标签2.1 12导联ECG的信号结构与常见格式12导联心电图不是12条独立信号而是从不同电极组合导出的12个视角。标准布局是I、II、III、aVR、aVL、aVF、V1–V6采样率常见500Hz或100Hz单条记录长度从10秒到60秒不等。WFDB格式.dat .hea是PhysioNet系列数据集的标准存储方式PTB-XL用的是500Hz、10秒记录CPSC2018则是500Hz、10–60秒不等。读WFDB最稳的方式是用wfdb库不要自己解析二进制。import wfdb import numpy as np # 读取一条记录返回信号矩阵和元数据 record wfdb.rdrecord(data/ptbxl/00001_lr) # 不带扩展名 signal record.p_signal # shape: (5000, 12)500Hz × 10秒 fs record.fs # 采样率PTB-XL为500 lead_names record.sig_name # [I,II,III,aVR,aVL,aVF,V1...V6] # 检查缺失值和量纲 print(f信号范围: {signal.min():.3f} ~ {signal.max():.3f} mV) print(fNaN数量: {np.isnan(signal).sum()})这段代码做了三件事用wfdb.rdrecord读取原始记录拿到p_signal物理信号矩阵已转换为mV以及采样率和导联名称。关键参数是record.p_signal返回的是物理值而非ADU原始值省去了手动转换。如果数据集中有NaN常见做法是用前向填充或直接丢弃该记录不要用均值填充——ECG的基线漂移会被均值填充放大。2.2 采样率统一与带通滤波不同数据集采样率不同混用之前必须统一。常见做法是全部重采样到100Hz或250Hz降低计算量的同时保留ECG主要频段0.5–40Hz。重采样用scipy的resample_poly比resample更稳因为它是多相滤波实现不会引入相位失真。from scipy.signal import resample_poly, butter, filtfilt def preprocess_ecg(signal, fs_orig, fs_target100): # 重采样到目标采样率 if fs_orig ! fs_target: from math import gcd up fs_target // gcd(fs_orig, fs_target) down fs_orig // gcd(fs_orig, fs_target) signal resample_poly(signal, up, down, axis0) # 0.5–40Hz带通滤波去除基线漂移和工频干扰 nyq fs_target / 2 b, a butter(4, [0.5/nyq, 40.0/nyq], btypeband) signal filtfilt(b, a, signal, axis0) # 按导联做z-score归一化 signal (signal - signal.mean(axis0)) / (signal.std(axis0) 1e-8) return signal.astype(np.float32)resample_poly的up/down参数由新旧采样率的最大公约数决定比如500→100就是up1、down5。butter的阶数选4阶是经验值阶数太高会振铃太低则阻带衰减不够。filtfilt做零相位滤波避免T波位置偏移。归一化按导联独立做因为不同导联的幅值范围差异大aVR通常倒置且幅值小。2.3 滑窗切段与标签对齐10秒记录直接喂给模型太长通常切成2–5秒的片段。切的时候要注意两点一是切段之间要有重叠常见50%防止QRS波被切断二是标签要对齐到每个片段如果一条记录有多个标签多标签分类每个片段继承整条记录的标签。def segment_signal(signal, fs, window_sec5, overlap0.5): win_len int(window_sec * fs) step int(win_len * (1 - overlap)) segments [] for start in range(0, signal.shape[0] - win_len 1, step): segments.append(signal[start:start win_len]) return np.stack(segments) # shape: (n_segments, win_len, 12)window_sec5、overlap0.5是ECG分类的常用起点。如果数据集里有些记录不足5秒要么补零要么丢弃补零会让模型学到边界伪影建议丢弃。标签对齐时注意多标签场景PTB-XL的scp_codes字段是字典需要映射到超类NORM、MI、STTC、CD、HYP再做one-hot。3. 模型选型1D-CNN、CNNLSTM还是Transformer3.1 为什么1D-CNN是ECG分类的默认起点ECG是典型的一维时序信号1D-CNN在局部波形检测QRS、P波、T波上天然匹配。相比RNNCNN训练更快、更容易并行而且感受野可以通过堆叠层数控制。常见结构是4–6个卷积块每个块包含Conv1D BatchNorm ReLU MaxPool最后接全局平均池化和全连接分类头。import torch import torch.nn as nn class ECGNet(nn.Module): def __init__(self, n_leads12, n_classes5): super().__init__() self.features nn.Sequential( # block 1: 12 - 32, 感受野约50ms nn.Conv1d(n_leads, 32, kernel_size7, padding3), nn.BatchNorm1d(32), nn.ReLU(), nn.MaxPool1d(2), # block 2: 32 - 64 nn.Conv1d(32, 64, kernel_size5, padding2), nn.BatchNorm1d(64), nn.ReLU(), nn.MaxPool1d(2), # block 3: 64 - 128 nn.Conv1d(64, 128, kernel_size3, padding1), nn.BatchNorm1d(128), nn.ReLU(), nn.MaxPool1d(2), # block 4: 128 - 256 nn.Conv1d(128, 256, kernel_size3, padding1), nn.BatchNorm1d(256), nn.ReLU(), nn.AdaptiveAvgPool1d(1) # 全局平均池化 ) self.classifier nn.Linear(256, n_classes) def forward(self, x): # x: (batch, 12, seq_len) x self.features(x).squeeze(-1) return self.classifier(x)卷积核从7降到3是常见设计浅层大核抓QRS复合波深层小核抓细节形态。AdaptiveAvgPool1d(1)把时间维压成1避免全连接层参数爆炸。输入需要转置成(batch, leads, seq_len)因为Conv1d在通道维做卷积。如果显存不够把block 4的256改成128准确率通常只掉1–2个点。3.2 CNNLSTM什么时候值得上当分类依赖节律信息比如房颤的RR间期不规则而非单搏形态时在CNN后面接一层LSTM或GRU会有提升。做法是把CNN输出的特征序列未做全局池化送进LSTM取最后时间步的隐状态分类。class CNNLSTM(nn.Module): def __init__(self, n_leads12, n_classes5, hidden128): super().__init__() self.cnn nn.Sequential( nn.Conv1d(n_leads, 32, 7, padding3), nn.ReLU(), nn.MaxPool1d(2), nn.Conv1d(32, 64, 5, padding2), nn.ReLU(), nn.MaxPool1d(2), nn.Conv1d(64, 128, 3, padding1), nn.ReLU() ) self.lstm nn.LSTM(128, hidden, batch_firstTrue, bidirectionalTrue) self.fc nn.Linear(hidden * 2, n_classes) def forward(self, x): x self.cnn(x) # (B, 128, T) x x.permute(0, 2, 1) # (B, T, 128) out, _ self.lstm(x) # (B, T, 256) return self.fc(out[:, -1, :]) # 取最后时间步双向LSTM的hidden*2是因为前向和后向拼接。注意LSTM输入需要batch_firstTrue否则维度对不上。这个结构参数量比纯CNN大3–5倍训练时间翻倍如果数据集小于5000条记录提升可能不明显甚至过拟合。3.3 类别不平衡的处理策略ECG数据集中正常样本通常占60%以上少数类如心肌梗死可能只有5%。直接训练会让模型偏向多数类。常见做法有三种加权交叉熵、focal loss、过采样。加权交叉熵最省事权重设为类别频率的倒数。from collections import Counter labels [0, 0, 0, 1, 2, 2, 3, 4] # 示例标签 counts Counter(labels) total sum(counts.values()) weights torch.tensor([total / (len(counts) * counts[i]) for i in range(len(counts))]) criterion nn.CrossEntropyLoss(weightweights)weights的计算逻辑是类别i的权重 总样本数 / (类别数 × 类别i的样本数)。这样少数类权重更大梯度更新时被重视。如果效果还不够换focal lossgamma2是常用值但需要调alpha。过采样要小心重复采样少数类会导致过拟合建议用SMOTE的时序版本或加噪声增强。4. 训练、验证与推理从数据划分到模型导出4.1 按患者划分数据集别按记录随机分这是ECG分类最容易翻车的地方。同一条记录切出的多个片段如果被分到训练集和验证集模型会记住患者特征而非疾病特征验证准确率虚高。正确做法是按患者ID划分确保同一患者的所有记录只出现在一个集合中。import pandas as pd from sklearn.model_selection import GroupShuffleSplit # df包含record_id, patient_id, label gss GroupShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(gss.split(df, groupsdf[patient_id])) train_df df.iloc[train_idx] val_df df.iloc[val_idx]GroupShuffleSplit的groups参数指定患者ID保证同一患者不跨集。如果数据集没有患者ID用记录ID的前缀代替PTB-XL的patient_id在元数据里。验证集再切一半做测试集或者用交叉验证。4.2 训练循环与早停训练循环里要监控验证集的AUC或F1而不是准确率——类别不平衡时准确率没意义。早停耐心值设10–15个epoch学习率用余弦退火或ReduceLROnPlateau。from torch.optim import Adam from torch.optim.lr_scheduler import ReduceLROnPlateau from sklearn.metrics import f1_score model ECGNet().cuda() optimizer Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler ReduceLROnPlateau(optimizer, modemax, patience5, factor0.5) best_f1, patience_counter 0, 0 for epoch in range(100): model.train() for x, y in train_loader: x, y x.cuda(), y.cuda() optimizer.zero_grad() loss criterion(model(x), y) loss.backward() optimizer.step() model.eval() preds, targets [], [] with torch.no_grad(): for x, y in val_loader: preds.extend(model(x.cuda()).argmax(1).cpu().numpy()) targets.extend(y.numpy()) f1 f1_score(targets, preds, averagemacro) scheduler.step(f1) if f1 best_f1: best_f1 f1 torch.save(model.state_dict(), best_model.pth) patience_counter 0 else: patience_counter 1 if patience_counter 15: breakweight_decay1e-4是L2正则防止过拟合。ReduceLROnPlateau在验证F1不升时降学习率factor0.5每次减半。早停耐心15个epoch是经验值太小会错过后期提升太大浪费时间。保存最佳模型而不是最后一个epoch的模型。4.3 推理与模型导出推理时要注意预处理必须和训练时完全一致——同样的重采样、滤波、归一化参数。导出模型用TorchScript或ONNX方便部署到没有PyTorch环境的机器。# 导出为ONNX dummy_input torch.randn(1, 12, 500).cuda() # 5秒100Hz torch.onnx.export( model, dummy_input, ecg_model.onnx, input_names[ecg], output_names[logits], dynamic_axes{ecg: {0: batch}, logits: {0: batch}} )dynamic_axes让batch维可变推理时不用固定batch size。导出后可以用onnxruntime验证输出是否一致。注意如果模型里有LSTMONNX导出可能需要额外处理建议先用CNN版本。5. 避坑与排查ECG深度学习作业里最常见的5个翻车点5.1 验证集准确率95%测试集掉到60%现象训练时验证集F1很高换测试集或实际数据后性能暴跌。原因按记录随机划分导致同一患者跨集模型记住了患者特异性特征。解决用GroupShuffleSplit按患者ID划分确保患者不跨集。如果数据集没有患者ID检查元数据里是否有patient_id字段PTB-XL和CPSC2018都有。5.2 模型只预测多数类现象训练loss下降但F1不涨混淆矩阵显示所有样本被预测为正常。原因类别不平衡 未加权损失。解决用加权交叉熵权重按类别频率倒数计算。如果加权后仍不改善检查标签映射是否正确——有时候标签编码错了所有样本被映射到同一类。5.3 滤波后信号失真QRS波变形现象带通滤波后QRS波幅值被削平或出现振铃。原因滤波器阶数太高或截止频率太接近Nyquist。解决降低滤波器阶数到4阶截止频率上限设为40Hz100Hz采样率时Nyquist为50Hz40Hz留了10Hz过渡带。用filtfilt而非lfilter避免相位失真。5.4 训练loss为NaN现象几个epoch后loss变成NaN。原因学习率太大、梯度爆炸、或者输入数据有NaN/Inf。解决先检查输入数据是否有NaN用np.isnan(signal).sum()确认。然后降低学习率到1e-4加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)。如果还不行检查归一化是否除了零——std 1e-8就是防这个。5.5 推理时预处理不一致导致结果随机现象训练好的模型在推理时输出随机结果。原因推理时的预处理参数采样率、滤波截止频率、归一化方式和训练时不一致。解决把预处理参数写进配置文件训练和推理共用同一份代码。常见错误是推理时忘了重采样或者归一化用了全局均值而非按导联均值。6. 把模型推到能用的程度几个我踩过坑才总结出的技巧6.1 数据增强在ECG上怎么做才不破坏病理特征ECG的数据增强不能像图像那样随便旋转裁剪。可用的增强方式有限加高斯噪声信噪比不低于20dB、随机时间偏移±0.2秒、随机幅值缩放0.9–1.1倍、导联随机丢弃模拟电极脱落。不要做时间反转——ECG的时间方向有病理意义反转后P波和T波位置互换标签就错了。def augment_ecg(signal, noise_std0.01, shift_max20, scale_range(0.9, 1.1)): # 加高斯噪声 signal signal np.random.normal(0, noise_std, signal.shape) # 随机时间偏移 shift np.random.randint(-shift_max, shift_max) signal np.roll(signal, shift, axis0) # 随机幅值缩放 scale np.random.uniform(*scale_range) signal signal * scale return signal.astype(np.float32)noise_std0.01对应信噪比约40dB不会淹没波形。shift_max20在100Hz采样率下是0.2秒覆盖一个QRS间期。scale_range控制幅值变化模拟不同患者的电压差异。增强只在训练时做验证和测试不做。6.2 用Grad-CAM看模型到底学到了什么Grad-CAM可以可视化模型关注的时间区域判断它是看QRS波还是看噪声。如果热力图集中在基线区域说明模型没学到有效特征需要检查预处理或增加数据量。# 简化版Grad-CAM针对1D-CNN def grad_cam(model, x, target_layer): model.eval() features [] def hook(module, input, output): features.append(output) handle target_layer.register_forward_hook(hook) x.requires_grad True output model(x) pred_class output.argmax(1).item() output[0, pred_class].backward() handle.remove() grads x.grad # 输入梯度作为近似 weights grads.mean(dim2, keepdimTrue) cam (weights * features[0]).sum(dim1, keepdimTrue) return cam.detach().cpu().numpy()这个简化版用输入梯度代替了Grad-CAM的全局平均池化权重效果差不多但实现更简单。target_layer选最后一个卷积块。如果热力图集中在QRS波附近0.2–0.4秒区间说明模型学到了有效特征。6.3 模型集成与阈值调优单模型F1到0.8左右就上不去了可以试模型集成训练3–5个不同初始化的模型推理时取平均logits。另外分类阈值不要用默认的0.5按验证集调——对少数类降低阈值能提升召回。策略预期F1提升代价加权交叉熵3~5%无数据增强2~4%训练时间×1.5模型集成3模型2~3%推理时间×3阈值调优1~2%需要验证集CNNLSTM1~3%参数量×3这张表是我在PTB-XL上跑下来的经验值具体数字因数据集和任务而异。优先级建议先加权交叉熵再数据增强最后考虑集成。CNNLSTM如果数据集小于5000条提升可能为负。6.4 一个我坚持的习惯每次跑完实验把配置文件、随机种子、验证集F1、测试集F1记在一个CSV里。ECG分类的随机性很大同一个模型换个种子F1能差3–5个点。不记录的话两周后根本不知道哪个配置是最好的。这个习惯帮我省了至少几十次重复实验。希望帮到你。本文还有配套的精品资源点击获取
返回列表