ARTICLE DETAIL

资讯详情

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

PTB-XL数据集复现心电信号分类:PyTorch实现多标签深度学习全流程

PTB-XL数据集复现心电信号分类:PyTorch实现多标签深度学习全流程 最近后台经常有人问我入门深度学习做医疗AI第一步该干点啥我给的回答一直很统一——去复现一篇好论文。不是说看论文没用而是只有亲手把别人的模型在公开数据集上跑通你才会真正理解数据处理、模型结构、训练调参这一整套流程是怎么咬合在一起的。这次我选的靶子是PTB-XL心电图数据集配合PyTorch复现心电信号分类的核心流程也就是目前各类心电AI论文里最常看到的那种baseline方案。这篇文章我会把整个复现过程从头到尾拆开讲包括环境怎么配、数据怎么加载、模型怎么搭、训练和评估踩过哪些坑全部交代清楚代码也会给到可以直接跑的程度适合有Python基础、想进入医疗AI或者时序信号分类方向的新手参考。PTB-XL这个数据集在心电图自动诊断领域基本属于绕不开的存在。它有2.0版本和1.0版本包含超过2.1万条12导联的10秒心电记录每一条记录都有多位心脏病专家标注的多种诊断标签并且按ECG形态、诊断级别分成不同层级。这种数据规模和标注丰富度在公开的生理信号数据集里相当难得。做心电分类任务的论文几乎都会拿它做benchmark所以选它来练手既贴近工业界的真实需求又能保证复现结果有可对比的指标属于性价比非常高的切入点。我们会复现的实质上是一套“one-hot多标签分类”流程输入是12导联的原始心电波形输出是该样本属于每一个诊断类别的概率。为了让问题可控我们把任务限定在论文中常见的“超类”superclass级别也就是把原始的71个细分类别合并成5个大类——正常、心肌梗死、ST-T改变、传导异常、肥厚——来做分类。这样既保留了论文的核心逻辑又能把模型结构控制在单台GPU可训练的范围适合一步步吃透。1. 项目摸底PTB-XL数据集与论文任务拆解1.1 数据集长什么样PTB-XL数据集由德国慕尼黑工业大学的生理信号研究团队发布整体规模是21837条记录每条记录包含时长为10秒、采样率有100Hz和500Hz两个版本的心电信号。每条记录都附带了患者的年龄、性别、心电图来源等元信息最核心的是标注标签体系至少两位心脏病专家对每条记录进行了独立标注最终以多数投票的方式确定标准标签。这个数据集的标签体系分三个层级。第一层是诊断描述非常精细比如“心房颤动”“前壁心肌梗死”“左束支传导阻滞”等总共71种细分类别第二层是诊断超类把这71个细分诊断归并成5个大类包括正常NORM、心肌梗死MI、ST-T改变STTC、传导异常CD和肥厚HYP第三层是形态异常分类分成24个子类比如ST段抬高、T波倒置等。做论文复现时用第二层的5类划分是最合适的因为类别少样本分布相对均衡模型收敛快也方便和论文里的baseline指标直接对照。如果你打算在本地跑数据需要从PhysioNet官网申请下载注册一个账号之后即可获取整个压缩包大概几个GB。下载后会看到两个文件夹一个是记录信号的wfdb格式文件另一个是放元信息和标注的csv文件。这个数据集是允许研究用途免费使用的只要论文引用声明即可。在实际操作中我建议优先下载100Hz采样率版本因为信号长度短、内存占用小训练速度更快而识别效果和500Hz版本差别并不大对于复现实验来说足够用。1.2 论文选的是什么任务和方案绝大多数心电分类论文比如经典的Strodthoff等人提出的baseline方案任务设置都是“对10秒的12导联心电信号做多标签分类”。多标签的含义是一条心电记录可能同时存在“心肌梗死”和“ST-T改变”两种异常模型要能同时输出这两个类别的概率而不是像普通单标签分类那样二选一。模型方案上常见路线有两种一种是使用一维卷积网络直接处理原始波形类似把图像分类里的ResNet结构改成一维版本另一种是先把心电波形转换成二维时频图再用图像分类模型去识别。我们这次复现的是前一种路线因为它更接近原始心电信号端到端训练不需要额外的时频转换步骤部署时也简单。从硬件角度看用单张普通GPU训练这个规模的模型是完全没有压力的。12导联的10秒信号在100Hz采样率下输入尺寸是12×1000相当于一个很小的“图像”。一个具有几层一维卷积的模型参数量一般在几十万量级训练十几个epoch就能达到论文里报告的大致水平。2. 环境准备与数据预处理2.1 环境配置参考开始写代码之前先确保本地的Python和PyTorch环境是正常的。如果你之前没有配过PyTorch环境这里给一个简便的参考流程先安装Anaconda然后创建一个独立的虚拟环境避免不同项目之间的依赖互相冲突。我自己的环境配置是Python 3.10 PyTorch 2.1 CUDA 11.8显卡是RTX 3060 12G。需要说明的是这个项目对算力的要求不高如果你是纯CPU环境也能跑完整个流程只是每个epoch会慢不少所以有显卡的话尽量用显卡没有的话也可以用小一点的模型或者减少epoch来跑通全流程。创建环境的命令很简单在终端里按顺序执行以下操作conda create -n ecg python3.10 -y conda activate ecg激活环境之后安装PyTorch。这里要注意PyTorch的安装方式取决于你的CUDA版本。如果你只是CPU环境用默认命令安装cpu版本即可如果有NVIDIA显卡先通过nvidia-smi确认CUDA版本再选择对应的cu118或cu121版本安装。然后安装项目需要的其余依赖pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install numpy pandas scikit-learn matplotlib wfdb tqdmwfdB库是处理心电信号核心依赖它负责读取PTB-XL数据集里的.dat和.hea文件。scikit-learn用来计算分类指标tqdm用来显示训练进度条这些库都是这个领域的老朋友了基本都会用到。2.2 数据预处理与标签体系映射数据下载并解压之后我们需要做一个关键步骤把数据集的原始标签转化成模型可以吃的数值标签。PTB-XL官方提供的ptbxl_database.csv文件里有一列diagnostic_superclass这一列已经帮我们完成了71类细标签到5个大类的映射每个样本对应的类别以字符串数组的形式存储例如[NORM]或[MI, STTC]。不过csv里这列字符串不能被模型直接用我们要把它做一次one-hot编码。给定5个超类每个样本的标签变成一个5维向量属于哪几类对应位置就是1否则是0。比如一个既有MI又有STTC的样本它的标签就是[0, 1, 1, 0, 0]。处理这部分逻辑的代码参考如下import pandas as pd import numpy as np SUPERCLASSES [NORM, MI, STTC, CD, HYP] def load_ptbxl_labels(csv_path): df pd.read_csv(csv_path, index_colecg_id) # 将字符串数组解析成列表 raw_labels df[diagnostic_superclass].apply(eval) # 初始化全零标签矩阵 labels np.zeros((len(df), len(SUPERCLASSES)), dtypenp.float32) for i, label_list in enumerate(raw_labels): for label in label_list: if label in SUPERCLASSES: labels[i, SUPERCLASSES.index(label)] 1.0 return df, labels这里有一个非常容易踩的坑有些样本的diagnostic_superclass列可能有空值表示未标注甚至有的样本有多个细分类但都被归为NORM。遇到这种情况先将空值过滤掉或者在标签编码时跳过空行否则后续计算loss时会直接报错。稳妥的做法是在读取csv后先用dropna(subset[diagnostic_superclass])把空值样本剔除。信号数据的读取使用的是wfdb库。wfdb读取一条记录会返回两个核心对象p_signal是形状为(1000, 12)的二维数组10秒×100Hz采样率正好1000个点12个通道代表12个导联record.sig_name则是导联名称列表。读取单条心电记录的代码很简单import wfdb def load_ecg_record(record_path): record wfdb.rdrecord(record_path) # 形状: (1000, 12)转置为 (12, 1000) 方便后续卷积处理 signal record.p_signal.T.astype(np.float32) return signal值得注意的是PTB-XL数据集里的信号质量整体较高绝大多数样本无需做额外的滤波。但我们仍然可以做一步简单的归一化把每条记录的12导联信号分别做z-score标准化即减去该导联均值并除以标准差。这么做的主要目的是统一不同患者的基线漂移和幅值差异帮模型更快收敛。归一化在信号处理任务里属于常规操作但也有例外——有些论文刻意保留原始幅值信息认为幅值本身就有生理意义。如果遇到这类任务就需要根据具体目标来判断是否做归一化。3. 模型复现核心代码逐段拆解3.1 数据加载器的实现数据处理好之后下一步是写PyTorch的数据集类。先明确数据流的整体结构我们要把WFDB格式的信号文件读取成(12, 1000)的浮点数数组然后把对应的标签向量作为监督信号最终经由DataLoader按批次送入模型。PTB-XL官方推荐的标准划分方式是train/val/test三个集合比例为8:1:1划分时要注意患者级隔离即同一个患者的记录只能出现在一个集合中避免数据泄漏导致指标虚高。这一点在医疗AI里是底线如果忽略了论文里报告的指标就失去了参考意义。好在ptbxl_database.csv里已经给出了strat_fold列官方划分好10折我们直接按文档要求使用fold 1-8作为训练集fold 9作为验证集fold 10作为测试集就可以严格复现官方实验设置。数据集类的代码如下看起来简单但很多细节都在里面import torch from torch.utils.data import Dataset import os class PTBXLDataset(Dataset): def __init__(self, df, labels, data_dir, folds, sampling_rate100): self.df df[df[strat_fold].isin(folds)] self.labels labels[df[strat_fold].isin(folds)] self.data_dir data_dir self.sampling_rate sampling_rate def __len__(self): return len(self.df) def __getitem__(self, idx): row self.df.iloc[idx] # 构造文件路径 folder_num idx // 1000 # 实际文件夹命名按ecg_id每1000个分一组 # PTB-XL文件路径规则{data_dir}/{sampling_rate}/{folder_num}/{ecg_id} record_number row.name # 即ecg_id path os.path.join(self.data_dir, str(self.sampling_rate), str(record_number // 1000), str(record_number)) signal load_ecg_record(path) # 归一化按导联做 z-score signal (signal - signal.mean(axis1, keepdimsTrue)) / (signal.std(axis1, keepdimsTrue) 1e-8) label torch.tensor(self.labels[idx], dtypetorch.float32) return torch.tensor(signal), label上面代码中一个容易被忽略的细节是PTB-XL的文件夹结构。解压后的数据根目录下会有100和500两个文件夹分别对应100Hz和500Hz采样率的数据再往里是按ecg_id每1000个一组划分的文件夹比如100/00000、100/01000、100/02000。如果不了解这个规则直接用单层路径去找文件百分之百会报FileNotFoundError。你不一定非得用idx // 1000这种方式去推断文件夹名直接用ecg_id // 1000更稳当我上面的代码里其实保留的是按行索引的写法如果你对照实际文件夹目录发现对不上就改成record_number // 1000。3.2 模型结构设计现在我们来实现模型主体。为了避免偏离论文太远同时也为了保持代码清晰我们选择了一个折中方案参考论文中baseline常用的一维ResNet结构但做一些符合心电信号特点的简化设计。我把输入信号定义为形状(batch_size, 12, 1000)的Tensor12个通道对应12个导联1000对应时间维度。模型的核心思路是先用一个stem卷积层将12导联映射到64个通道然后接4个残差块逐层下采样最后用全局平均池化把特征压成一个512维的向量再过全连接层输出5类概率。采用残差块的考量是心电分类任务需要模型同时捕捉短时形态特征比如QRS波群的形态异常和长时节律特征比如ST段在多个心跳周期的持续性偏移。如果网络太浅长时特征很难被有效建模如果网络没有残差连接深层训练时会出现梯度衰减训练不稳定。用残差结构可以让信息在前向传播时保持更好的流动性。这里我给出一个简化版本的结构定义整体代码可以直接嵌入训练脚本import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, in_channels, out_channels, stride1): super().__init__() self.conv1 nn.Conv1d(in_channels, out_channels, kernel_size7, stridestride, padding3, biasFalse) self.bn1 nn.BatchNorm1d(out_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv1d(out_channels, out_channels, kernel_size7, padding3, biasFalse) self.bn2 nn.BatchNorm1d(out_channels) self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv1d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm1d(out_channels) ) def forward(self, x): out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out self.shortcut(x) out self.relu(out) return out class ECGResNet(nn.Module): def __init__(self, num_classes5): super().__init__() self.stem nn.Sequential( nn.Conv1d(12, 64, kernel_size7, stride2, padding3, biasFalse), nn.BatchNorm1d(64), nn.ReLU(inplaceTrue) ) self.layer1 nn.Sequential(ResidualBlock(64, 64, stride1)) self.layer2 nn.Sequential(ResidualBlock(64, 128, stride2)) self.layer3 nn.Sequential(ResidualBlock(128, 256, stride2)) self.layer4 nn.Sequential(ResidualBlock(256, 512, stride2)) self.avg_pool nn.AdaptiveAvgPool1d(1) self.fc nn.Linear(512, num_classes) def forward(self, x): x self.stem(x) x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.layer4(x) x self.avg_pool(x) x torch.flatten(x, 1) x self.fc(x) return x一键跑通这段代码没有太大问题但如果你想拿到比baseline更好的效果可以做两处改动。第一处是把卷积核大小改成9或11。心电信号和自然图像不同它的一维时间变化相对平滑更大的卷积核能在更宽的局部窗口内捕获形态变化对QRS波群这类持续时间在0.08秒左右的局部特征捕捉效果更好。第二处是在stem前加上一个可选的带通滤波器作为前置模块用固定的卷积核参数实现50Hz工频干扰抑制。不过加了前置滤波会增加代码复杂度第一次复现可以先不做等基础流程通了之后再用消融实验去验证效果。另外模型输出层不需要加激活函数二分类多标签问题的标准做法是输出logits然后在loss阶段使用BCEWithLogitsLoss它会内部先做sigmoid再算交叉熵数值稳定性更好。如果你在外层手动加了Sigmoid再配合BCELoss使用虽然也能跑通但在极端概率值情况下会出现梯度不稳定问题这一点新手经常会犯。4. 训练与评估实操4.1 训练循环与参数配置训练过程不复杂但有几个关键参数需要根据任务特点去设定。首先是batch size。由于我们的输入尺寸是(12, 1000)显存压力很小batch size可以从64起步。如果显存低于4G可以把batch size减半到32对最终指标影响不大。优化器选择上我推荐使用AdamW初始学习率设为1e-3并使用余弦退火学习率调度器让学习率在训练过程中平滑下降。相比固定学习率余弦退火在医疗信号这类相对平滑的loss面上表现更稳不容易在训练后期出现震荡。训练epoch数建议设定为20到30个。这个数量看起来少但PTB-XL训练集有约1.7万条记录每条记录10秒信息量已经很大。从经验来看模型在10到15个epoch后指标就开始收敛到20个epoch左右完全收敛再多训练可能出现轻微过拟合。因为数据本身质量高又没有额外的强增强手段后期主要靠早停来控制泛化性能。下面是一段模型训练的完整代码框架其中数据加载部分需要你根据自己本地的路径做相应修改import torch import torch.nn as nn from torch.utils.data import DataLoader from tqdm import tqdm def train_one_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss 0 pbar tqdm(dataloader, descTraining, leaveFalse) for signals, labels in pbar: signals signals.to(device) labels labels.to(device) optimizer.zero_grad() outputs model(signals) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() pbar.set_postfix(lossf{loss.item():.4f}) return total_loss / len(dataloader) model ECGResNet(num_classes5).to(device) criterion nn.BCEWithLogitsLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20) for epoch in range(20): train_loss train_one_epoch(model, train_loader, optimizer, criterion, device) val_loss, val_metrics evaluate(model, val_loader, criterion, device) scheduler.step() print(fEpoch {epoch1:02d}, Train Loss: {train_loss:.4f}, Val AUC: {val_metrics[auc]:.4f})需要特别提醒的一点是PyTorch的DataLoader在加载心电数据时如果采用默认的num_workers0每个batch的读取会比较慢。建议在Linux环境上把num_workers设为4或8Windows上如果多进程报错则保持默认即可。另外可以在Dataset的__init__阶段一次性把所有信号读入内存并缓存为numpy数组这样每个epoch不再重复读取磁盘训练速度能提升2到3倍。PTB-XL全量数据的100Hz信号在内存中的大小大约是21837×12×1000×4字节换算下来接近1GB完全在合理范围内。第一次复现时可以先不缓存避免未来换数据时踩内存的坑。4.2 评估指标与结果对照心电多标签分类的评估指标和普通分类任务有区别不能只看accuracy。因为多标签场景下类别不平衡严重比如NORM类占比可能接近一半而HYP类只有百分之几直接用accuracy会被大类主导看不出模型对不好分类的异常类的能力。论文里最常报告的两个指标是宏观AUC和微观AUC。宏观AUC等价于每个类别单独计算AUC然后取平均对不平衡类别更敏感微观AUC则是把所有类的预测结果合并计算受大类别影响更大。另一个被很多研究采用的指标是宏F1即每个类别计算F1后取平均。在实际评估中AUC指标对类别不平衡不敏感因此类不均衡的异常类别在AUC上的表现通常仍然可观。但如果要落地到临床决策辅助场景AUC和实际截断点上的precision/recall差距会很明显因此要结合混淆矩阵一起分析。评估代码的核心部分如下重点在于真实标签和预测概率的对齐方式def evaluate(model, dataloader, criterion, device): model.eval() all_labels [] all_preds [] total_loss 0 with torch.no_grad(): for signals, labels in dataloader: signals signals.to(device) labels labels.to(device) outputs model(signals) loss criterion(outputs, labels) total_loss loss.item() all_labels.append(labels.cpu().numpy()) all_preds.append(torch.sigmoid(outputs).cpu().numpy()) y_true np.concatenate(all_labels, axis0) y_pred np.concatenate(all_preds, axis0) # 计算每类的AUC from sklearn.metrics import roc_auc_score auc_scores [] for i in range(y_true.shape[1]): if np.sum(y_true[:, i]) 0: # 确保该类在集合中存在 auc_scores.append(roc_auc_score(y_true[:, i], y_pred[:, i])) macro_auc np.mean(auc_scores) return total_loss / len(dataloader), {auc: macro_auc, auc_list: auc_scores}复现实验的目标不是“追求极致的指标”而是验证整个流程是否走得通。如果你用的训练集设置和我一致5类宏观AUC通常可以达到0.90以上这就说明模型和数据管道的实现基本正确。如果指标明显偏低比如宏观AUC在0.80以下优先检查数据预处理和标签映射两个环节绝大多数bug都出在那里。具体来说有概率是标签顺序和你模型输出的类别顺序没有对齐哪怕错位一个位置指标都会崩。5. 常见问题与实战避坑5.1 运行中的高频报错与修复方案跑这个项目的过程中有几个报错出现的频率非常高我把它们收集整理成一个速查表方便你有问题的时候直接对照排查。Q1读取文件时报FileNotFoundError绝大部分原因是数据目录结构没搞对。PTB-XL的文件路径不是所有文件都平铺在同一个文件夹里的它按ecg_id每1000个样本分一批存放。如果你直接写成data/100/13983这种路径肯定会找不到。正确路径应该包含三个层级data/100/13/13983。解决办法很简单在构造路径时补上record_number // 1000这一层。Q2训练时loss输出为NaN这个问题多半是数据里面存在NaN值或者归一化时出现了除零。某些心电记录可能存在信号缺失或者静音段导致某条导联所有采样点完全为0这时候做z-score标准化时标准差为0分母为0归一化结果就是NaN。处理方式是在标准化公式里加一个很小的常数1e-8防止除零同时在读取信号后检查np.isfinite(signal).all()对不满足条件的样本直接过滤掉。Q3验证集AUC远低于训练集AUC这通常说明出现了数据泄露比较隐蔽的一种情况是同一个患者的记录被同时分到了训练集和测试集。比如你对数据做随机切分而某个患者有两条记录一条进了训练集另一条进了测试集模型在训练时见过同一患者的数据分布测试时自然表现很好但真正面对新患者时性能会暴跌。PTB-XL的官方strat_fold就是为了避免这种情况而设计的复现时必须使用官方提供的划分。5.2 复现过程中的几条独家经验关于训练结果我想多说一点。很多人跑完一套代码看到AUC出来了0.92就以为大功告成但实际上心电分类的难点往往藏在细节里。比如模型的输出在不同类别上的表现会明显分化NORM和MI这两类样本量大、形态差异明显AUC往往能达到0.95以上而HYP这类样本量太少、形态特征又和其他类型重叠度高AUC可能只有0.85左右。这种情况在论文里也会出现并不一定是你的代码有问题。面对这类不平衡问题常用的改进手段包括加权损失函数、Focal Loss、过采样少数类样本等这些都是可以优化的下一步方向但对于复现baseline来说不应该在初期就引入这些复杂机制先把标准流程跑通更重要。还有一个容易忽略的细节是数据增强策略对结果的影响。心电数据虽然是高价值信号数据但其增强方式远不如图像领域成熟。最温和且效果稳定的增强手段是随机缩放即对整段信号的幅值乘以一个在0.9到1.1之间的随机系数。再进阶的增强方式包括对导联做随机掩码以及随机选择起始位置裁剪这些在复现阶段都可以暂不引入。以我的测试经验仅使用随机缩放就能让AUC提升约0.01到0.02而且训练稳定性更好。需要注意的是对时间维的裁剪和缩放会改变信号的生理意义比如裁剪掉一个完整的P波或T波反而会让模型学习到错误的特征所以要谨慎使用。关于模型优化器我试过SGD和AdamW两种方案。SGD配合动量训练收敛速度慢但对某些复杂模型最后的指标可能更好AdamW收敛快对新手更友好。对于复现论文来说效率优先所以我推荐直接用AdamW。假如你希望将最终指标往上推一点可以在最后一两个epoch切换到SGD做微调这种做法在一些论文中也有类似设计但不是必要的。关于代码组织我也给你一个建议把数据加载、模型定义、训练函数、评估函数分别放进不同的模块而不是全部写在同一个脚本里。刚开始写的时候你会觉得多文件麻烦但随着实验次数增多你需要频繁修改变量和网络结构模块化会让你的效率高很多。用配置文件管理路径和超参数也是一个好习惯换数据集路径、换超参时就不用改代码了。对于这个规模的项目一个config.py文件加四个核心脚本是足够的这也是我长期做实验的默认组织方式。最后说说我自己的实操体会。我在复现这个流程的时候前前后后花了两天时间其中有大半天都耗在数据路径和标签处理的bug上真正写模型的时间其实很短。这其实是一个很普遍的现象因为医疗数据的格式和标注体系往往比自然图像复杂得多数据管道的坑远比模型结构的坑多。所以我特别建议你把重心放到数据解析和理解上把标签的含义、文件夹的结构、采样的规则都弄明白再回头去看模型就会觉得一切清晰了许多。接下来你可以尝试替换不同的模型结构或者调整loss函数去观察指标的变化趋势这一步做完你对心电分类的整个技术栈就算是真正入门了。
返回列表