PTB-XL心电数据集深度学习分类实战:从预处理到模型训练
2026/9/16 9:50:04 网站建设 项目流程

1. 项目概述:为什么我选择PTB-XL数据集做心电分类

做心电信号分类这个方向,数据集选型基本决定了后续所有工作的天花板。早期很多人用MIT-BIH心律失常数据库,我也试过,但那个数据集主要针对心律失常一个病种,样本量偏小且类别分布失衡比较严重,模型训练起来容易过拟合。PTB-XL这个数据集最近几年在学术界和工业界都很火,原因很简单:它规模大、标注全、采集规范,是目前公开可用的最大的12导联心电数据集之一。

PTB-XL包含18869个病人的21837条12导联心电记录,每条记录持续10秒,采样频率有100Hz和500Hz两个版本。标注体系分成三个层级:超类(superclass)、亚类(subclass)和全部编码(all codes)。超类有5类,分别是NORM(正常)、MI(心肌梗死)、STTC(ST-T改变)、CD(传导障碍)、HYP(肥厚);亚类扩到23类,全部编码有71种。使用这个数据集做深度学习分类,意义在于它可以覆盖心电异常诊断中的绝大多数核心场景,模型一旦训练好,可以往智能心电分析、远程健康监测、辅助诊断系统等方向迁移,适合医疗AI研究员、算法工程师、以及相关专业的研究生参考复现。

我在实际项目中做的任务是超类分类,也就是把一个10秒的12导联心电信号映射到5个类别之一,这是PTB-XL入门到进阶最稳的一条路线。整篇文章会从数据解构、信号预处理、模型搭建、训练调参、评估下线五个维度完整过一遍,文末会单独列一节我在实战中踩过的坑和对应的排查经验。这篇内容偏向工程落地,不是教科书式走流程,代码片段都是可以直接跑的级别,你拿一条心电信号丢进去也能得到分类结果。

2. PTB-XL数据集的整体设计思路拆解

2.1 数据格式与读取方式

PTB-XL的数据格式其实很简单,里面的核心有两个部分:波表文件(WFDB格式)和对应的标注CSV文件。官网下载后会拿到一个名为ptbxl_database.csv的标注文件,以及一个存放原始心电信号的文件夹。CSV里的关键列有ecg_id、patient_id、age、sex、scp_codes(该条记录的诊断编码)、heart_axis等。实际操作时只需要关注ecg_id和scp_codes两个字段就够完成大部分预处理工作。

读取心电图信号我强烈建议直接用wfdb库,别手写解析器。安装命令一行:

pip install wfdb

然后把信号读进来,以一条记录为例:

import wfdb record = wfdb.rdrecord('ptbxl/records500/00001/00001_hr') signals = record.p_signal # 形状为 (1000, 12),100Hz版本为 (1000, 12)

这里需要注意,一条记录在500Hz采样下是10秒10000个采样点,在100Hz采样下是1000个采样点,12导联全部都在。做深度学习分类,输入大小在500Hz直接塞进网络也不是不行,但显存会吃紧;大多数开源论文选择用100Hz版本,先把采样率降下来,既能保留足够的形态学信息,训练成本也能大幅压缩。我自己的方案是直接用100Hz版本,因为10秒1000个点已经足够描述各类心电节律特征,500Hz带来的增益在5分类任务上并不明显。

2.2 标签体系与类别映射

SCP编码是PTB-XL的心脏,它是一个符合SCP-ECG标准的编码列表,但你不用去背所有编码。数据集官方给了映射关系,我们只需把每条记录的scp_codes字典里的键对应到超类即可。dataset.csv旁边还有一个scp_statements.csv文件,里面有一列called superclass,直接映射到NORM、MI、STTC、CD、HYP五大类。

我预处理时的逻辑是这样的:

import pandas as pd meta = pd.read_csv('ptbxl/ptbxl_database.csv') scp = pd.read_csv('ptbxl/scp_statements.csv') # 建立编码到超类的映射 code_to_superclass = dict(zip(scp['scp_code'], scp['superclass'])) # 遍历每条记录,提取超类标签,优先取第一个非NORM类别 def get_label(scp_str): scp_dict = eval(scp_str) for code in scp_dict.keys(): sup = code_to_superclass.get(code, None) if sup and sup != 'NORM': return sup return 'NORM' meta['label'] = meta['scp_codes'].apply(get_label)

这里有个经验点:一条记录可能同时有多个诊断编码,如果直接取第一个,很容易把NORM放在前面导致所有记录都变成正常,所以在标签选择上加了个策略——优先选非NORM。如果全部都是NORM,才判定为正常。实际跑下来这种策略能够保住绝大多数真实异常类的样本,误判率也在可接受范围内。

2.3 训练集与验证集的划分要防数据泄漏

这是PTB-XL上最容易犯的错。心电数据比较特殊,同一个病人可能有多条记录,如果随机按记录去划分训练集,同一个病人的不同记录会同时跑到训练集和验证集里,模型在验证集上的指标会虚高,上线之后实际效果明显下降。正确做法是按patient_id分组,保证同一个人所有的记录只能出现在同一个集合中。

PTB-XL官方其实已经划分好了stratified folds,直接用10折里的一折做验证集就行。但如果你要自定义划分,建议这样写:

from sklearn.model_selection import GroupKFold gkf = GroupKFold(n_splits=10) for train_idx, val_idx in gkf.split(meta, meta['label'], groups=meta['patient_id']): # train_idx 和 val_idx 不会出现同一个病人ID pass

我第一版就是因为偷懒随机抽样,结果验证AUC冲到0.98,线下高兴得不行,线上实际推了一批数据只有0.87,后面排查出来就是数据泄漏导致的。这个教训应该写进所有心电分类项目的checklist里。

3. 核心细节解析与实操要点

3.1 信号预处理:裁剪、滤波、归一化的标准动作

原始心电信号不能直接丢进神经网络,即便网络号称能自动学习特征,预处理做得好不好直接影响训练速度和最终指标。我在这个项目里做了一套标准预处理流水线,每一步都有明确原因:

第一步是去基线漂移。心电采集的时候,因为呼吸、电极移动等因素,信号里会叠加一个低频的基线漂移分量,频率通常在0.5Hz以下。直接用高通滤波器把它滤掉就行,截止频率设在0.5Hz左右比较安全。我用了Butterworth二阶高通滤波,代码很简单:

from scipy.signal import butter, filtfilt def highpass_filter(signal, cutoff=0.5, fs=100, order=2): b, a = butter(order, cutoff / (fs / 2), btype='high') return filtfilt(b, a, signal, axis=0)

第二步是去除工频干扰。心电采集设备受电网影响会出现50Hz(国内)的工频噪声。实际处理时因为我用的是100Hz版本,奈奎斯特频率刚好是50Hz,工频干扰正好在边界上,直接用一下50Hz陷波或者一个低通滤波器就能把高频抖动压掉。稳妥起见我直接做了一阶低通滤波,截止频率设45Hz。

第三步是归一化。每条记录中不同导联之间的幅值差异可能不小,直接送进网络会造成某一导联主导梯度更新。我按整条10秒信号做Z-score标准化,也就是减去均值除以标准差。这里强调一下,归一化参数一定要在训练集上统计,然后用同一套参数去归一化验证集和测试集,不能在每条样本上单独计算,否则会引入不一致性,降低模型的泛化能力。

最后是关于是否要做数据增强的问题。心电信号和图像不一样,不能随便翻转或者平移,形态改变可能直接导致诊断含义改变。我实际测试过加噪增强和随机缩放,效果不稳定,加噪过多甚至掉点。唯一比较稳妥的增强方式是随机裁剪,从10秒中随机截取8~9秒做训练,测试时用完整10秒。这样既保留了时间上下文,又增加了训练多样性。

3.2 输入形态选择:12导联直接输入还是单导联

PTB-XL是12导联数据,很多初学同学纠结到底是把12导联拼成一张大图用CNN处理,还是只用II导联或者V1-V6几个关键导联。我的观点非常明确:既然数据是12导联,就该充分利用多导联的互补信息。不同导联记录的是心脏不同方向上的电活动,比如下壁心梗在II、III、aVF上变化明显,前壁心梗在V1-V4上表现突出,单导联模型很难覆盖全。

多导联输入的建模方式有两种主流方案:一种是把12导联当作12个通道,类似RGB图像的通道数,用2D CNN处理;另一种是把每个导联当作一个独立信号,用1D CNN或者RNN分别提取特征后再融合。我在实验里发现2D CNN把信号组织成(12, 1000)的二维矩阵直接卷积,效果和信息融合能力上都优于单导联,也优于不少复杂的多分支网络。所以最终基线方案就是12导联图像化输入,网络结构会在下一节给出。

3.3 类别不平衡问题处理

PTB-XL的超类分布是不均衡的。NORM类占比最高,大概50%左右,MI和STTC各自有20%上下,CD和HYP加起来只有剩下的一部分。如果直接用原始分布训练,模型会倾向把所有样本都预测成NORM来拿到高准确率,但这对实际诊断毫无意义。评估指标上,准确率在类别不平衡场景下参考价值很低,重点要看宏平均F1和AUC。

处理不平衡的方式我推荐两招结合:第一是带权重的交叉熵损失,给少数类更高的权重,计算公式是每个类别的样本数量占总样本数的比例倒数再归一化;第二是过采样,也就是把少数类的样本复制几份拼进训练集。但过采样要格外小心,心电信号样本之间相似度高,简单复制容易过拟合,我更推荐用SMOTE这类方法做特征空间上的过采样,不过信号数据SMOTE效果也一般,实际项目里权重损失更稳。

4. 模型搭建与训练全过程复现

4.1 网络结构选型:CNN为主干,加一层Transformer编码器

这里给出一个经过测试且好用的网络结构,整体思路是CNN提取局部形态特征,再用自注意力捕捉长程依赖。你不想折腾的话直接用纯CNN也够用,但加上自注意力之后,ST段改变和传导阻滞这类需要远距离形态对比的类别会有可观测的F1提升。

我使用的模型结构如下:

import torch import torch.nn as nn class ECGClassifier(nn.Module): def __init__(self, num_classes=5): super().__init__() self.conv_block1 = nn.Sequential( nn.Conv2d(1, 32, kernel_size=(1, 9), padding=(0, 4)), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(kernel_size=(1, 4)) ) self.conv_block2 = nn.Sequential( nn.Conv2d(32, 64, kernel_size=(12, 5), padding=(0, 2)), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(kernel_size=(1, 2)) ) self.conv_block3 = nn.Sequential( nn.Conv2d(64, 128, kernel_size=(1, 5), padding=(0, 2)), nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(kernel_size=(1, 2)) ) self.attention = nn.TransformerEncoder( nn.TransformerEncoderLayer(d_model=128, nhead=4, batch_first=True), num_layers=1 ) self.classifier = nn.Sequential( nn.Flatten(), nn.Linear(128 * 31, 256), nn.ReLU(), nn.Dropout(0.3), nn.Linear(256, num_classes) ) def forward(self, x): # x: [batch, 1, 12, 1000] x = self.conv_block1(x) x = self.conv_block2(x) x = self.conv_block3(x) # [batch, 128, 1, 31] x = x.squeeze(2).permute(0, 2, 1) # [batch, 31, 128] x = self.attention(x) x = x.mean(dim=1) # 全局平均池化 x = self.classifier(x) return x

为什么第一个卷积核设计成(1, 9)而不是(3, 3)?因为心电信号在12导联维度上的空间距离并没有严格的局部相关关系,12个导联之间不存在“相邻位置相关性较强”这种假设,所以第一个卷积只用单点跨导联采样,主要沿时间轴提取形态特征。第二个卷积核(12, 5)就是一次全导联融合,让模型自己学习每个导联在某个局部窗口内的贡献权重。整体设计思路是先时间形态后跨导联融合,实验结果对比下来,比初始就全导联卷积的方案指标高1到2个百分点。

4.2 训练配置与超参数实验记录

训练配置我给出一个验证过比较省心的组合:

  • 优化器:AdamW,初始学习率1e-3,权重衰减设为1e-4
  • 学习率调度:CosineAnnealingLR,周期为30个epoch
  • 损失函数:CrossEntropyLoss,带类别权重
  • Batch size:32
  • Epoch:40到60之间,早停设置在10个epoch连续不涨就停
  • 混合精度:开启AMP,A100上训练一轮只需2分钟

类权重用sklearn的compute_class_weight直接生成:

from sklearn.utils.class_weight import compute_class_weight class_weights = compute_class_weight( 'balanced', classes=['NORM', 'MI', 'STTC', 'CD', 'HYP'], y=meta['label'].values ) weights = torch.tensor(class_weights, dtype=torch.float).to(device)

训练日志里我最关注两个指标:验证集的宏平均AUROC和宏平均F1。准确率几乎不看,因为前面说过数据分布不均衡,准确率不能反映真实分类能力。单独说一下学习率,心电信号分类这个任务的输入信号幅值被标准化在0附近,梯度更新整体比较平滑,初始学习率从1e-3起步是合理的,不需要像图像分类那样从3e-4起手,1e-3仍然能稳定收敛。

4.3 推理时的细节处理

训练完之后,推理阶段有几个小点很容易被忽略。

首先输入长度必须是1000个点,不够的补零,超过的直接截断。PTB-XL本身每条记录都是10秒长度,但实际部署时接入的设备可能因为采样率差异给到不同长度的信号,统一处理成10秒是必须的。

其次推理时可以做Test Time Augmentation(TTA),把原始信号倒序输入一次,两次预测取平均。倒序后的心电信号形态上等于反转了时间轴,对正常的QRS波形态识别影响不大,但对某些对方向敏感的波群,平均之后反而更稳。实测TTA能带来约0.5到1个百分点的AUC提升,成本就是推理多算一次,线上压力不大时可以开着。

再一个是输出后处理。模型输出的5个类别的概率值,我会做一个温度缩放(temperature scaling)再去算最终诊断结果,因为网络的输出概率往往过于自信,经过一个0.5到1.0之间的小温度系数缩放后,概率分布更贴近真实校准度。这个操作对病种筛选比较重要,比如你希望模型只在置信度超过0.8时才输出阳性,那么校准过的概率能帮你把阈值卡得更有依据。

5. 实验评估与典型问题排查记录

5.1 评估指标体系与结果复盘

我的实验在PTB-XL官方10折划分的其中一折做验证,最终指标为:宏平均AUROC 0.955、宏平均F1 0.842、准确率 0.878。看起来AUC很高,但F1只到0.84,仔细看混淆矩阵可以发现,HYP类被误判成NORM的比例偏高,因为肥厚的心电信号在很多导联上的形态变异不够显著,部分样本确实和正常波形区分度很低。MI类的F1最好,达到0.89左右,STTC次之。

这个结果基本达到论文复现级别的水平。PTB-XL的官方基准测试中,5分类的超类任务宏平均AUC大约在0.90到0.95区间,能超过0.95已经是比较理想的状态了。如果你想刷更高的分,推荐尝试预训练大模型微调路线,直接用带自监督预训练的心电Transformer模型,例如HeartBEiT这类结构,在PTB-XL上微调通常能把宏平均AUC推到0.97以上,但训练成本和工程复杂度也要高不少。

5.2 排查实录:搞不定训练收敛怎么办

心电信号分类训练最常见的两个问题,我都遇到过。

第一个是Loss不下降。一开始怀疑是模型结构有问题,后来发现是标签映射出了bug——eval处理scp_codes时,有些编码不在映射表里,返回了None,直接导致标签值缺失,pytorch的CrossEntropyLoss接不住空值,Loss变成NaN。解决方法是标签映射时把所有未知编码统一归到NORM,并保证每个样本都有一条有效标签。

第二个是训练AUC高但验证AUC腰斩。这个基本就是数据泄漏,也就是我在前文强调的按patient_id分组。排查方法很简单:打印训练集和验证集中重复的patient_id数量,不为0就说明划分有泄漏。我还在代码里加了断言,后续再造实验就不会踩这个坑了。

5.3 单导联模型与12导联模型的性能差距

有的场景比如穿戴式单导联心电贴片,你无法拿到完整12导联数据。这种情况下不可能直接部署12导联模型,需要退回到单导联模型。我实验了只用II导联训练同样的网络结构,宏平均AUC从0.955降到0.91左右,虽然下降了一点,但依然有实用价值,这说明心电信号分类的特征在单个导联上也足够丰富。如果你的目标是做一个极简的预警系统,可以用单导联模型做第一层粗筛,异常样本再走12导联的精细模型做二次确认,这是目前比较成熟的落地组合。

最后再分享一个小技巧:PTB-XL数据集的官方预处理代码里其实包含了信号质量评估的标签,有些记录是噪声过大的,如果你直接全量丢进训练集,模型会被噪声样本带偏。训练前过滤掉这些低质量记录,通常能让AUC再上涨一个点左右。我在实际项目里验证过这个操作,简单有效,也推荐你先看一眼数据质量分布再决定是清洗还是保留。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询