尧图精选

PTB-XL心电数据集分类实战:从环境搭建到PyTorch模型训练完整指南

🕒 发布时间:2026/9/16 21:13:16 📁 来源:尧图网络
在这个领域折腾久了你会发现一个规律不管是刚入门的同学还是已经跑过不少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 python3.9 conda activate ecgPython版本选3.9或3.10都可以PyTorch对这两个版本支持最稳定。如果机器上没有Anaconda建议先安装AnacondaWindows、Linux、macOS都有对应的安装包安装完成后打开终端Windows建议用Anaconda Prompt执行上面的命令。2.2 安装PyTorchGPU版本和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.csvSCP编码对应的诊断类别和等级records100/或records500/按ecg_id开头两位分目录存放的WFDB格式信号文件ptbxl_database.csv是整个数据集的索引里面的每一行对应一条心电记录ecg_id是唯一标识filename_lr和filename_hr分别对应当前采样率下的文件路径label列是官方给出的superclass标签。3.2 用wfdb读取心电信号WFDB格式是心电领域通用的数据格式之一用wfdb库读取非常简单。信号文件通常以.dat和.hea为扩展名前者是二进制波形数据后者是文本头文件记录采样率、导联数和增益等信息。import wfdb record wfdb.rdsamp( record_path, # 不带扩展名的文件路径 sampfrom0, # 从哪个采样点开始读取 samptoNone, # 读取到哪个采样点None表示读完全部 channel_names[II] # 只读取指定导联默认读全部12导联 ) signals, meta recordsignals是一个二维数组形状是(采样点数, 导联数)meta包含采样率fs、导联名称、信号缩放因子等信息。3.3 采样率选择100Hz还是500HzPTB-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.csv的label列直接给出了每条的superclass标签取值是NORM、MI、STTC、CD、HYP字符串。做分类任务时需要将这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]或者按研究团队提供的官方预处理代码进行带通滤波。如果只是想复现baselineZ-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, lowcut0.5, highcut45.0, fs100, order4): nyquist fs * 0.5 low lowcut / nyquist high highcut / nyquist b, a butter(order, [low, high], btypeband) return filtfilt(b, a, signal, axis0)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, fs100, use_channelsNone, transformNone): self.df df.reset_index(dropTrue) 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_namesself.channels, sampfrom0, samptoNone) 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], dtypetorch.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再套上DataLoaderfrom torch.utils.data import DataLoader train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) test_loader DataLoader(test_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)num_workers在Linux下可以设置成4或8Windows下建议从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, stride1): super().__init__() self.conv1 nn.Conv1d(in_channels, out_channels, kernel_size3, stridestride, padding1) self.bn1 nn.BatchNorm1d(out_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv1d(out_channels, out_channels, kernel_size3, padding1) 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_size1, stridestride), 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_channels12, num_classes5, base_channels64): super().__init__() self.conv1 nn.Conv1d(in_channels, base_channels, kernel_size7, stride2, padding3) self.bn1 nn.BatchNorm1d(base_channels) self.relu nn.ReLU(inplaceTrue) self.maxpool nn.MaxPool1d(kernel_size3, stride2, padding1) self.layer1 self._make_layer(base_channels, base_channels, 2, stride1) self.layer2 self._make_layer(base_channels, base_channels * 2, 2, stride2) self.layer3 self._make_layer(base_channels * 2, base_channels * 4, 2, stride2) 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_channels12, num_classes5): super().__init__() self.branch1 nn.Sequential(nn.Conv1d(in_channels, 64, kernel_size5, padding2), nn.ReLU(inplaceTrue)) self.branch2 nn.Sequential(nn.Conv1d(in_channels, 64, kernel_size15, padding7), nn.ReLU(inplaceTrue)) self.branch3 nn.Sequential(nn.Conv1d(in_channels, 64, kernel_size31, padding15), nn.ReLU(inplaceTrue)) self.lstm nn.LSTM(192, 128, bidirectionalTrue, batch_firstTrue) 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], dim1) # (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, dim1) context torch.sum(attn_weight * lstm_out, dim1) # (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_weightbalanced, classesnp.array([0, 1, 2, 3, 4]), ydf[label_idx].values ) class_weights torch.tensor(class_weights, dtypetorch.float32).to(device) criterion nn.CrossEntropyLoss(weightclass_weights)compute_class_weight会自动根据各类样本数计算平衡权重实测下来比手工设置更省心。训练时损失函数内部会根据当前batch中样本的类别自动按权重放大少数类别的梯度贡献。6.2 优化器与学习率策略优化器不用搞得太花哨Adam就是很稳的选择初始学习率设置在1e-3左右。训练过程中用学习率衰减配合早停可以避免在末期振荡。optimizer torch.optim.Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemax, factor0.5, patience5, verboseTrue )这里的modemax是配合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, descTraining): 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, descEvaluating): 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, dim1).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, labels6.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_classovr, averagemacro) pred np.argmax(probs, axis1) f1 f1_score(labels, pred, averagemacro) 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_colecg_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(fTrain: {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_dirdata/ptbxl, fs100) val_dataset PTBXLDataset(val_df, data_dirdata/ptbxl, fs100) test_dataset PTBXLDataset(test_df, data_dirdata/ptbxl, fs100) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) test_loader DataLoader(test_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) model ResNet1D(in_channels12, num_classes5).to(device) criterion nn.CrossEntropyLoss(weightclass_weights.to(device)) optimizer torch.optim.Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemax, factor0.5, patience5) 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_classovr, averagemacro) print(fEpoch {epoch1}/{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.) break7.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_classovr, averagemacro) test_pred np.argmax(test_probs, axis1) test_f1 f1_score(test_labels, test_pred, averagemacro) test_acc accuracy_score(test_labels, test_pred) print(fTest AUC: {test_auc:.4f}) print(fTest F1: {test_f1:.4f}) print(fTest 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_labelsclass_names) disp.plot(cmapBlues) 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, labelf{class_names[i]} (AUC {roc_auc:.3f})) plt.plot([0, 1], [0, 1], k--, labelChance) plt.xlabel(False Positive Rate) plt.ylabel(True Positive Rate) plt.title(ROC Curves for PTB-XL Superclass Classification) plt.legend(loclower right) plt.show()8. 训练过程中的典型坑与排查思路这个项目整体跑通不难但中间有几个坑会让结果出现严重的劣化我在复现时逐个踩过列出排查思路。8.1 第一个坑wfdb读取时导联顺序错乱wfdb.rdsamp在指定channel_names时是按名称匹配导联而不是按文件中的位置。不同记录的导联排列顺序大概率相同但为了保险必须在Dataset中显式传入channel_names参数而不是靠默认顺序。排查方法随机抽取几条记录打印record[1][sig_name]确认导联顺序和预期一致。如果顺序错乱模型训练时输入通道语义混乱结果会很差而且难以察觉。8.2 第二个坑病人划分不注意导致数据泄漏如果直接df.sample(frac0.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_workersWindows下多进程可能引发死锁或性能下降建议设置为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在形态上本来就有一定重叠混淆矩阵能直观告诉你模型的极限在哪儿哪些混淆是数据本身的性质导致的哪些是模型结构导致的这对后续优化方向有直接指导意义。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →