PTB-XL心电数据集分类实战:从环境搭建到PyTorch模型训练完整指南
2026/9/16 21:12:56 网站建设 项目流程

在这个领域折腾久了,你会发现一个规律:不管是刚入门的同学还是已经跑过不少CV模型的工程师,遇到心电信号分类任务时,第一反应都是去找PTB-XL数据集和对应的论文复现。原因很简单,PTB-XL是目前公开可用的、规模最大的12导联心电数据集,论文也相对成熟,作为学习深度学习的起点非常合适。但真到动手做的时候,很多人会卡在数据集下载、wfdb库的读取、标签处理、按患者划分数据集这些环节上,真正把整个pipeline从零到一跑通,比想象中要复杂不少。

这篇教程就是来解决这个问题的。我会基于Python和PyTorch,从环境准备开始,到PTB-XL数据集的下载与解析,再到模型结构设计、训练评估,最后给出可直接复现的完整代码思路。整个过程按我实际踩过坑的经验来讲,尽量让每个环节都能落地,而不是停留在概念层面。适合有三四个月Python基础、想入门医疗AI或时序信号分类的同学参考,也适合想快速拿PTB-XL做一个baseline结果的工程师。

1. 为什么选择PTB-XL做复现而不是自己造数据集

很多教程喜欢拿公开的Kaggle竞赛数据开始讲,但PTB-XL有它独特的位置。这个数据集包含了21837条12导联心电记录,每条都是10秒长度的原始信号,采样率有100Hz和500Hz两种版本可供选择。它比大多数竞赛数据集的规模都要大,而且带有结构化程度很高的标注信息,包括诊断类别、心律类别、形态类别等,这意味着你可以做多级多标签的分类任务,也可以只做superclass五分类,灵活性很强。

更关键的是,PTB-XL在论文中是按患者维度划分训练集、验证集和测试集的,比例是8:1:1。这个设计看似简单,但实际上做了充足的数据泄漏规避,因为同一个患者可能有多条记录,如果按记录而不是按患者划分,不同来源的数据可能混杂在一起,模型在测试集上的表现会虚高。用PTB-XL做复现时,这一点必须有意识地保持一致,否则你后续的对比实验、论文投稿都会出问题。

还有一个加分项:PTB-XL官方发布的paper提供了详细的baseline结果,包括使用不同模型结构在五分类superclass任务上的AUC、F1等指标。这就给我们复现时提供了直接对照的靶子,做完实验可以和官方指标对比,判断自己的实现是否合理。对学习者而言,有标准答案的训练远比自由发挥有效。

2. 环境搭建是复现的第一步:Anaconda、CUDA和PyTorch

很多同学会跳过环境这一步,认为装个PyTorch很简单,结果真到了复现代码的时候,不是缺少wfdb库就是CUDA版本不匹配,浪费大量时间。我建议把环境问题在最开始就彻底解决。

2.1 创建独立的conda环境

千万别在base环境里直接装PyTorch,时间久了依赖冲突会让你怀疑人生。用conda单独建一个环境,互不干扰。

conda create -n ecg python=3.9 conda activate ecg

Python版本选3.9或3.10都可以,PyTorch对这两个版本支持最稳定。如果机器上没有Anaconda,建议先安装Anaconda,Windows、Linux、macOS都有对应的安装包,安装完成后打开终端(Windows建议用Anaconda Prompt)执行上面的命令。

2.2 安装PyTorch:GPU版本和CPU版本的选择

PyTorch的安装方式推荐走官方源,但国内网络环境下经常遇到下载慢或者超时。一个比较稳妥的做法是在PyTorch官网选择对应CUDA版本后,把命令中的下载源替换为国内镜像源,比如清华源或阿里源。

pip install torch torchvision torchaudio -i https://mirrors.aliyun.com/pypi/simple/

如果机器没有NVIDIA GPU,装CPU版本即可:

pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu

安装完成后一定验证一下CUDA是否可用:

import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU mode")

这里有个容易踩的坑:很多人装完PyTorch发现torch.cuda.is_available()返回False,原因可能是安装的PyTorch是CPU版本,也可能是NVIDIA驱动太老导致CUDA版本不匹配。建议先通过nvidia-smi查看驱动支持的CUDA版本,再选择对应的PyTorch版本。实测下来,PyTorch 2.1以上版本对CUDA 11.8和12.1都兼容得比较好,优先在这两个版本里选。

2.3 项目依赖库清单

除了PyTorch,还需要安装以下库:

pip install wfdb numpy pandas scikit-learn matplotlib tqdm -i https://mirrors.aliyun.com/pypi/simple/
  • wfdb:读取PTB-XL数据集中WFDB格式的心电信号
  • scikit-learn:用于计算AUC、F1、混淆矩阵等评估指标
  • tqdm:训练时显示进度条
  • matplotlib:绘制ROC曲线和训练曲线

3. 拿到PTB-XL数据集之后,先别急着训练

我见过太多人的做法是数据集一解压就扔给模型,然后发现模型收敛慢、指标差,回头排查才发现是预处理环节出了问题。PTB-XL的预处理是整个复现流程里最容易出错、也最影响最终结果的一步。

3.1 数据集的下载与目录结构

PTB-XL需要通过PhysioNet网站下载,数据量大概在1GB左右(500Hz采样率版本更大些)。下载完成后解压到一个固定目录,比如data/ptbxl/。目录内核心文件如下:

  • ptbxl_database.csv:所有记录的基本信息,包括ecg_id、patient_id、采样率、心率、诊断标注等
  • scp_statements.csv:SCP编码对应的诊断类别和等级
  • records100/records500/:按ecg_id开头两位分目录存放的WFDB格式信号文件

ptbxl_database.csv是整个数据集的索引,里面的每一行对应一条心电记录,ecg_id是唯一标识,filename_lrfilename_hr分别对应当前采样率下的文件路径,label列是官方给出的superclass标签。

3.2 用wfdb读取心电信号

WFDB格式是心电领域通用的数据格式之一,用wfdb库读取非常简单。信号文件通常以.dat.hea为扩展名,前者是二进制波形数据,后者是文本头文件记录采样率、导联数和增益等信息。

import wfdb record = wfdb.rdsamp( record_path, # 不带扩展名的文件路径 sampfrom=0, # 从哪个采样点开始读取 sampto=None, # 读取到哪个采样点,None表示读完全部 channel_names=['II'] # 只读取指定导联,默认读全部12导联 ) signals, meta = record

signals是一个二维数组,形状是(采样点数, 导联数)meta包含采样率fs、导联名称、信号缩放因子等信息。

3.3 采样率选择:100Hz还是500Hz

PTB-XL官方同时发布了100Hz和500Hz两种采样率版本。复现时建议优先选择100Hz,原因有两个:

  • 每条10秒记录在100Hz下是1000个采样点,输入模型的计算量比500Hz少5倍,训练速度快很多
  • PTB-XL论文中大量baseline实验是在100Hz采样率上做的,便于对照结果

如果后续想挑战更高精度的复现,可以再做500Hz版本。

3.4 标签处理与类别映射

PTB-XL的标注体系分多个层级,常用的是superclass,共5类:

superclass含义
NORM正常心电图
MI心肌梗死
STTCST段和T波改变
CD传导阻滞
HYP心肌肥大

ptbxl_database.csvlabel列直接给出了每条的superclass标签,取值是NORMMISTTCCDHYP字符串。做分类任务时,需要将这5个类别映射为0-4的整型索引。

class_names = ['NORM', 'MI', 'STTC', 'CD', 'HYP'] label_dict = {name: idx for idx, name in enumerate(class_names)} df['label_idx'] = df['label'].map(label_dict)

3.5 数据泄漏预警

PTB-XL官方推荐按patient_id进行划分,而不是直接按行切分。由于同一个患者可能有多条记录,如果直接随机划分训练集和测试集,同一个患者的记录可能同时出现在两边,模型等于见过"考试答案",评估结果没有说服力。

实现时先按patient_id分组,把不重复的患者ID随机打乱,再按8:1:1划成训练、验证、测试三份,最后根据患者ID把对应记录归入相应集合。这是复现PTB-XL论文时最重要的一条规则,务必重视。

4. 预处理与DataLoader的完整实现

把WFDB文件读进来只是第一步,要把信号变成模型能直接吃的张量,中间还有几道工序。

4.1 信号标准化

不同记录的信号幅值范围不同,而深度学习模型对输入的量级很敏感,尤其是一开始就用带正则项训练时,量级差异过大会影响收敛。统一做法是Z-score标准化,按每条记录的全体采样点计算均值和标准差,然后做减法除法。

mean = signal.mean() std = signal.std() signal = (signal - mean) / (std + 1e-8)

也有论文选择只做幅度归一化到[-1, 1],或者按研究团队提供的官方预处理代码进行带通滤波。如果只是想复现baseline,Z-score标准化已经足够了,配合网络中的BatchNorm层,效果比较稳定。

4.2 降噪与滤波

原始心电信号中常混有工频干扰(50Hz/60Hz)、肌电干扰和基线漂移。为了屏蔽高频噪声和基线漂移,常见做法是做带通滤波,比如保留0.5Hz到40Hz或0.5Hz到100Hz频段。

很多人会直接用scipy.signal.butter设计数字滤波器,但在复现PTB-XL baseline时,官方并没有把滤波作为必要步骤,大多数论文也只做了轻预处理或直接卷积网络让模型自己过滤。考虑到这是保姆级教程,我会推荐一个简单有效的方案:如果网络结构比较深,滤波可以不做,让网络自动学习时域特征;如果网络结构比较简单,建议加一个带宽0.5Hz到45Hz的Butterworth带通滤波器。

from scipy.signal import butter, filtfilt def bandpass_filter(signal, lowcut=0.5, highcut=45.0, fs=100, order=4): nyquist = fs * 0.5 low = lowcut / nyquist high = highcut / nyquist b, a = butter(order, [low, high], btype='band') return filtfilt(b, a, signal, axis=0)

4.3 Dataset类实现

PyTorch里写Dataset类很简单,核心是实现__len____getitem__两个方法。PTB-XL的Dataset可以这样设计:

import torch from torch.utils.data import Dataset import pandas as pd import numpy as np import wfdb class PTBXLDataset(Dataset): def __init__(self, df, data_dir, fs=100, use_channels=None, transform=None): self.df = df.reset_index(drop=True) self.data_dir = data_dir self.fs = fs self.channels = use_channels or ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6'] self.transform = transform def __len__(self): return len(self.df) def __getitem__(self, idx): row = self.df.iloc[idx] record_path = f"{self.data_dir}/{row['filename_lr']}" record = wfdb.rdsamp(record_path, channel_names=self.channels, sampfrom=0, sampto=None) signals, meta = record signals = signals.astype(np.float32) # 标准化 for ch in range(signals.shape[1]): ch_mean = signals[:, ch].mean() ch_std = signals[:, ch].std() signals[:, ch] = (signals[:, ch] - ch_mean) / (ch_std + 1e-8) # 转成 (channels, time) 形状 signals = signals.T # (12, 1000) label = torch.tensor(row['label_idx'], dtype=torch.long) if self.transform: signals = self.transform(signals) return torch.tensor(signals), label

这里有几个关键细节需要注意:

  • 读取时通过channel_names显式指定导联顺序,保证每条记录读取出来的导联排列是相同的,否则DataLoader里会报形状不一致的错误
  • 标准化按每个导联独立计算,避免某个导联幅值过大把其他导联的信息掩盖掉
  • 输出形状统一为(12, 1000),通道维度在前,符合PyTorch对输入张量(batch, channels, length)的习惯

4.4 DataLoader参数设置

完成Dataset之后,用官方划分好的DataFrame创建三个Dataset,再套上DataLoader:

from torch.utils.data import DataLoader train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True)

num_workers在Linux下可以设置成4或8,Windows下建议从0开始调,否则多进程报错会让人抓狂。pin_memory在GPU训练时能减少数据传输时间,实测对训练速度有一定提升。

5. 模型结构:从一维CNN到多尺度特征融合

PTB-XL论文本身提供了多种baseline实现,包括简单的卷积网络、全连接网络等。复现的核心不是盲抄一个巨大网络,而是用和论文类似的结构跑出一个合理结果,再逐步优化。这里我给出两套方案:一套是快速baseline,适合验证pipeline通不通;另一套是多尺度融合结构,效果更好,也是很多后续论文的常用变体。

5.1 快速baseline:一维ResNet

如果之前跑过图像ResNet,改成1D非常容易,只需要把nn.Conv2d换成nn.Conv1d,把nn.BatchNorm2d换成nn.BatchNorm1d,池化层换成对应1D版本。这里给一个轻量级版本,参数量在百万级别,单卡训练很快。

import torch.nn as nn class BasicBlock1D(nn.Module): def __init__(self, in_channels, out_channels, stride=1): super().__init__() self.conv1 = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1) self.bn1 = nn.BatchNorm1d(out_channels) self.relu = nn.ReLU(inplace=True) self.conv2 = nn.Conv1d(out_channels, out_channels, kernel_size=3, padding=1) self.bn2 = nn.BatchNorm1d(out_channels) self.stride = stride self.shortcut = nn.Sequential() if stride != 1 or in_channels != out_channels: self.shortcut = nn.Sequential( nn.Conv1d(in_channels, out_channels, kernel_size=1, stride=stride), nn.BatchNorm1d(out_channels) ) def forward(self, x): residual = self.shortcut(x) out = self.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) out += residual out = self.relu(out) return out class ResNet1D(nn.Module): def __init__(self, in_channels=12, num_classes=5, base_channels=64): super().__init__() self.conv1 = nn.Conv1d(in_channels, base_channels, kernel_size=7, stride=2, padding=3) self.bn1 = nn.BatchNorm1d(base_channels) self.relu = nn.ReLU(inplace=True) self.maxpool = nn.MaxPool1d(kernel_size=3, stride=2, padding=1) self.layer1 = self._make_layer(base_channels, base_channels, 2, stride=1) self.layer2 = self._make_layer(base_channels, base_channels * 2, 2, stride=2) self.layer3 = self._make_layer(base_channels * 2, base_channels * 4, 2, stride=2) self.avgpool = nn.AdaptiveAvgPool1d(1) self.fc = nn.Linear(base_channels * 4, num_classes) def _make_layer(self, in_channels, out_channels, blocks, stride): layers = [] layers.append(BasicBlock1D(in_channels, out_channels, stride)) for _ in range(1, blocks): layers.append(BasicBlock1D(out_channels, out_channels)) return nn.Sequential(*layers) def forward(self, x): x = self.relu(self.bn1(self.conv1(x))) x = self.maxpool(x) x = self.layer1(x) x = self.layer2(x) x = self.layer3(x) x = self.avgpool(x) x = torch.flatten(x, 1) x = self.fc(x) return x

这个结构在PTB-XL五分类任务上,输入100Hz的12导联信号,测试集macro AUC通常能达到0.90左右,和PTB-XL论文里报道的baseline水平接近。

5.2 进阶方案:多尺度卷积 + 双向LSTM + 注意力

如果只用一个baseline网络,很难体会到深度学习做心电分类的精髓。心电信号有两个显著特点:一是局部形态特征(QRS波、ST段)对分类至关重要,二是时序上下文(心律节律)同样重要。单尺度卷积擅长捕捉局部模式,但对长程依赖表现一般。

我实际跑下来效果比较好的一条路线是:多尺度一维卷积提取局部特征,然后把特征序列输入双向LSTM建模时间依赖,最后用注意力池化汇总全局信息。结构大致如下:

  • 三个并行的Conv1d分支,kernel_size分别为5、15、31,分别捕捉不同宽度的波形特征
  • 将三个分支的输出在通道维度上拼接
  • 经过两个BiLSTM层,隐藏维度设置128
  • 用注意力池化代替简单的全局平均池化,让模型自动关注重要时间段
  • 接全连接分类层
class MultiScaleECGNet(nn.Module): def __init__(self, in_channels=12, num_classes=5): super().__init__() self.branch1 = nn.Sequential(nn.Conv1d(in_channels, 64, kernel_size=5, padding=2), nn.ReLU(inplace=True)) self.branch2 = nn.Sequential(nn.Conv1d(in_channels, 64, kernel_size=15, padding=7), nn.ReLU(inplace=True)) self.branch3 = nn.Sequential(nn.Conv1d(in_channels, 64, kernel_size=31, padding=15), nn.ReLU(inplace=True)) self.lstm = nn.LSTM(192, 128, bidirectional=True, batch_first=True) self.attention = nn.Sequential( nn.Linear(256, 64), nn.Tanh(), nn.Linear(64, 1) ) self.dropout = nn.Dropout(0.5) self.fc = nn.Linear(256, num_classes) def forward(self, x): # x: (batch, 12, 1000) b1 = self.branch1(x) b2 = self.branch2(x) b3 = self.branch3(x) feat = torch.cat([b1, b2, b3], dim=1) # (batch, 192, T) feat = feat.permute(0, 2, 1) # (batch, T, 192) lstm_out, _ = self.lstm(feat) # (batch, T, 256) attn_score = self.attention(lstm_out) # (batch, T, 1) attn_weight = torch.softmax(attn_score, dim=1) context = torch.sum(attn_weight * lstm_out, dim=1) # (batch, 256) out = self.fc(self.dropout(context)) return out

这个模型的参数量比单分支ResNet稍大,但训练时间还在可接受范围内。在PTB-XL五分类任务上,macro AUC能到0.92左右,F1分数也明显高于纯CNN结构。

5.3 为什么注意力池化比全局平均池化更有效

心电信号不是每个时间片段对分类都有同样的贡献,比如某段P波附近的信息对判断心肌梗死有多大作用,可能不如ST段附近的信息关键。如果做全局平均池化,所有时间点的特征都被等权压缩成一个向量,重要片段的贡献会被大量平常片段淹没。

注意力池化的思路恰恰是让网络自己学出一组权重,对每个时间点的特征向量打个分,加权求和得到融合表示。这个打分函数就是self.attention里的两线性层,输出一个标量再经过softmax变成和为1的权重。这个机制在长序列任务里几乎是免费的涨点手段。

6. 训练配置:损失函数、优化器、评估指标

基础设施搭好后,训练细节直接决定最终结果的好坏,尤其是类别不平衡问题。PTB-XL的superclass五分类中,NORM类别占了很大比例,如果直接拿交叉熵训练,模型会倾向于把所有样本都预测为NORM,整体准确率也许不低,但MI、CD这些类别的召回率会非常差。

6.1 类别权重处理

解决类别不平衡最常见的手段是在损失函数里给少数类别更高的权重。权重设置为各类样本数的倒数,再归一化即可。

from sklearn.utils.class_weight import compute_class_weight class_weights = compute_class_weight( class_weight='balanced', classes=np.array([0, 1, 2, 3, 4]), y=df['label_idx'].values ) class_weights = torch.tensor(class_weights, dtype=torch.float32).to(device) criterion = nn.CrossEntropyLoss(weight=class_weights)

compute_class_weight会自动根据各类样本数计算平衡权重,实测下来比手工设置更省心。训练时损失函数内部会根据当前batch中样本的类别自动按权重放大少数类别的梯度贡献。

6.2 优化器与学习率策略

优化器不用搞得太花哨,Adam就是很稳的选择,初始学习率设置在1e-3左右。训练过程中用学习率衰减配合早停,可以避免在末期振荡。

optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode='max', factor=0.5, patience=5, verbose=True )

这里的mode='max'是配合AUC等监控指标使用的,当指标连续5个epoch不再提升时,学习率减半。早停机制用验证集AUC作为监控指标,连续15个epoch不提升就停止训练,并保存最好的模型权重。

6.3 训练循环代码

一个完整的训练循环大概长这样:

from tqdm import tqdm import torch.nn.functional as F def train_one_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss = 0.0 for x, y in tqdm(dataloader, desc="Training"): x, y = x.to(device), y.to(device) optimizer.zero_grad() logits = model(x) loss = criterion(logits, y) loss.backward() optimizer.step() total_loss += loss.item() * x.size(0) return total_loss / len(dataloader.dataset)

验证时需要注意关闭梯度计算和BatchNorm的自动更新统计信息。

def evaluate(model, dataloader, criterion, device): model.eval() total_loss = 0.0 all_logits = [] all_labels = [] with torch.no_grad(): for x, y in tqdm(dataloader, desc="Evaluating"): x, y = x.to(device), y.to(device) logits = model(x) loss = criterion(logits, y) total_loss += loss.item() * x.size(0) all_logits.append(F.softmax(logits, dim=1).cpu().numpy()) all_labels.append(y.cpu().numpy()) probs = np.vstack(all_logits) labels = np.hstack(all_labels) return total_loss / len(dataloader.dataset), probs, labels

6.4 评估指标的选取与解读

PTB-XL论文和后续工作常用macro AUC和macro F1作为主要指标,而不是准确率。原因很简单:类别不平衡时准确率存在欺骗性。sklearn计算macro AUC需要先把多分类问题转换成OvR,具体做法是:

from sklearn.metrics import roc_auc_score, f1_score, accuracy_score auc = roc_auc_score(labels, probs, multi_class='ovr', average='macro') pred = np.argmax(probs, axis=1) f1 = f1_score(labels, pred, average='macro') acc = accuracy_score(labels, pred)

macro AUC的含义是对每个类别单独计算AUC后取平均,它不受类别数量分布的影响,能更客观地反映模型在每个类别上的区分能力。PTB-XL论文里的superclass五分类任务,baseline macro AUC在0.90左右,如果能跑到0.91、0.92,说明你的复现已经相当到位了。

7. 完整代码整合:从DataFrame到训练完成的一站式流程

前面每部分都是独立组件,这里把完整流程串起来,做成可直接运行的脚本流程,按顺序执行即可。

7.1 数据准备步骤

读取ptbxl_database.csv,添加标签索引,进行患者级别的8:1:1划分:

import pandas as pd import numpy as np df = pd.read_csv('data/ptbxl/ptbxl_database.csv', index_col='ecg_id') df['label_idx'] = df['label'].map(label_dict) # 患者级别划分 patients = df['patient_id'].unique() np.random.seed(42) np.random.shuffle(patients) n_train = int(len(patients) * 0.8) n_val = int(len(patients) * 0.1) train_patients = patients[:n_train] val_patients = patients[n_train:n_train + n_val] test_patients = patients[n_train + n_val:] train_df = df[df['patient_id'].isin(train_patients)] val_df = df[df['patient_id'].isin(val_patients)] test_df = df[df['patient_id'].isin(test_patients)] print(f"Train: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}")

如果是第一次复现,建议把划分好的DataFrame存成CSV方便复用,避免每次运行都重新读一遍。

train_df.to_csv('data/train_fold.csv') val_df.to_csv('data/val_fold.csv') test_df.to_csv('data/test_fold.csv')

7.2 训练主脚本

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') train_dataset = PTBXLDataset(train_df, data_dir='data/ptbxl', fs=100) val_dataset = PTBXLDataset(val_df, data_dir='data/ptbxl', fs=100) test_dataset = PTBXLDataset(test_df, data_dir='data/ptbxl', fs=100) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True) model = ResNet1D(in_channels=12, num_classes=5).to(device) criterion = nn.CrossEntropyLoss(weight=class_weights.to(device)) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=5) best_auc = 0.0 early_stop_counter = 0 num_epochs = 50 for epoch in range(num_epochs): train_loss = train_one_epoch(model, train_loader, optimizer, criterion, device) val_loss, val_probs, val_labels = evaluate(model, val_loader, criterion, device) val_auc = roc_auc_score(val_labels, val_probs, multi_class='ovr', average='macro') print(f"Epoch {epoch+1}/{num_epochs} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val AUC: {val_auc:.4f}") scheduler.step(val_auc) if val_auc > best_auc: best_auc = val_auc early_stop_counter = 0 torch.save(model.state_dict(), 'best_model.pt') print("Model saved.") else: early_stop_counter += 1 if early_stop_counter >= 15: print("Early stopping triggered.") break

7.3 测试集评估

训练结束后加载最优权重,在测试集上跑一遍:

model.load_state_dict(torch.load('best_model.pt')) test_loss, test_probs, test_labels = evaluate(model, test_loader, criterion, device) test_auc = roc_auc_score(test_labels, test_probs, multi_class='ovr', average='macro') test_pred = np.argmax(test_probs, axis=1) test_f1 = f1_score(test_labels, test_pred, average='macro') test_acc = accuracy_score(test_labels, test_pred) print(f"Test AUC: {test_auc:.4f}") print(f"Test F1: {test_f1:.4f}") print(f"Test Accuracy: {test_acc:.4f}")

7.4 混淆矩阵与ROC曲线可视化

在测试集上画混淆矩阵,能直观看出哪些类容易混淆:

from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay cm = confusion_matrix(test_labels, test_pred) disp = ConfusionMatrixDisplay(cm, display_labels=class_names) disp.plot(cmap='Blues') plt.title('Confusion Matrix') plt.show()

ROC曲线的绘制稍微麻烦些,需要为每个类别单独画:

from sklearn.preprocessing import label_binarize from sklearn.metrics import roc_curve, auc y_bin = label_binarize(test_labels, classes=[0, 1, 2, 3, 4]) plt.figure(figsize=(10, 8)) for i in range(5): fpr, tpr, _ = roc_curve(y_bin[:, i], test_probs[:, i]) roc_auc = auc(fpr, tpr) plt.plot(fpr, tpr, label=f'{class_names[i]} (AUC = {roc_auc:.3f})') plt.plot([0, 1], [0, 1], 'k--', label='Chance') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('ROC Curves for PTB-XL Superclass Classification') plt.legend(loc='lower right') plt.show()

8. 训练过程中的典型坑与排查思路

这个项目整体跑通不难,但中间有几个坑会让结果出现严重的劣化,我在复现时逐个踩过,列出排查思路。

8.1 第一个坑:wfdb读取时导联顺序错乱

wfdb.rdsamp在指定channel_names时,是按名称匹配导联,而不是按文件中的位置。不同记录的导联排列顺序大概率相同,但为了保险,必须在Dataset中显式传入channel_names参数,而不是靠默认顺序。

排查方法:随机抽取几条记录,打印record[1]['sig_name'],确认导联顺序和预期一致。如果顺序错乱,模型训练时输入通道语义混乱,结果会很差而且难以察觉。

8.2 第二个坑:病人划分不注意导致数据泄漏

如果直接df.sample(frac=0.8)划分训练集,同一个患者的多条心电记录可能同时出现在训练集和验证集中。这种泄漏会让验证集AUC虚高到0.95以上,测试时不升反降。判断方法很简单:检查训练集和验证集的重叠patient_id数量,如果为0说明划分正确。

8.3 第三个坑:NORM类别过多导致指标虚高

PTB-XL中NORM样本占比偏高,如果不做类别权重处理,模型预测全部朝NORM偏移,准确率可能超过70%但macro AUC只有0.8左右。在训练时打印每个batch的标签分布,发现某个类别占比过高,就要及时用compute_class_weight做平衡处理。

8.4 第四个坑:GPU显存不足

如果在训练到一半时遇到CUDA out of memory,优先把batch_size从32降到16或8,同时把pin_memory关掉。如果仍然不够,可以在训练循环里加上torch.cuda.empty_cache(),并且检查是否有历史变量占用了显存没有释放。

8.5 第五个坑:训练曲线看起来混乱,损失不下降

先检查学习率是否合适,1e-3对这个小网络一般是安全的。再检查DataLoader的num_workers,Windows下多进程可能引发死锁或性能下降,建议设置为0。如果验证曲线震荡很大,可以调小学习率并配合梯度裁剪。

9. 复现有哪些可以继续深挖的方向

基础五分类跑通之后,PTB-XL还有很多可以玩的方向,这部分对把教程当起点的同学会有启发。

第一是细粒度分类。PTB-XL不光有superclass标签,也有subclass和更细的诊断标签,共有71种诊断类别。把5分类改成多标签、多级别分类,模型需要更强的特征表达能力,也会遇到更严重的类别不平衡问题,是一个很好的进阶挑战。

第二是导联子集实验。12导联数据里,能否用更少的导联(比如只用II导联和V5导联)做分类,对便携式心电设备的算法设计很有实际意义。只需要在Dataset里设置use_channels参数就行,代码改动很小,实验价值很大。

第三是信号增强。对心电信号做随机裁剪、时间扭曲、加噪声、导联置零等增强,能在一定程度上提升模型的鲁棒性。需要自己写transform函数,思路和图像增强类似,但时序数据的变换要小心不要破坏心电波形的生理结构。

第四是模型轻量化。把训练好的ResNet1D做剪枝、量化或者知识蒸馏,压缩到可以在边缘设备上运行的规模,是一个工程价值很高的方向。

10. 个人实操后的几点体会

讲完所有代码和技术细节,最后分享几个在这个项目上验证过的经验。

深度学习做心电分类,结构复杂度不是第一位的,数据管线是否正确才是决定成败的关键。我见过很多同学拿PTB-XL跑不出官方指标,排查到最后发现是数据划分出了问题,或者标准化逻辑写错了。建议先按baseline代码把pipeline完全跑通,再考虑换更强的模型结构。

PTB-XL官方划分中,训练、验证、测试的比例是8:1:1,但实际操作时不同随机种子可能导致结果有零点几个百分点的波动。正确的做法是在固定随机种子后把划分结果落盘,所有实验共用同一份划分文件,这样模型对比才有意义。

学习率策略上,ReduceLROnPlateau比CosineAnnealing在这个任务上更稳定。心电分类的验证集曲线会有一定抖动,CosineAnnealing在周期末尾容易过早收敛到局部最优,而ReduceLROnPlateau能根据实际指标动态调整,实测最终指标更稳。

训练完成之后,别只看最终AUC,把混淆矩阵打出来看一眼哪些类别互相混淆非常有用。以PTB-XL的superclass为例,MI和STTC、CD和HYP在形态上本来就有一定重叠,混淆矩阵能直观告诉你模型的极限在哪儿,哪些混淆是数据本身的性质导致的,哪些是模型结构导致的,这对后续优化方向有直接指导意义。

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

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

立即咨询