简介:面向人工智能与深度学习方向开发者,这份资源围绕PyTorch构建了完整的ECG信号处理与识别框架,覆盖数据清洗、去噪、分段,到CNN/RNN混合模型设计、损失函数与优化器配置、训练验证及模型评估部署等核心环节。包体共749个文件,包含291个Python源码、163个MAT数据文件、心电标注文件(hea/atr/dat)及PDF、Markdown文档,压缩包约22.76MB,目录结构清晰,适合医疗AI研究者或入门者学习参考。已有314人学习下载。资源中可见基于CINC2021、CPSC2019/2021等公开数据集的多项实验配置与训练日志,附带了tfevents事件文件、模型评估脚本及多尺度/逐导联等消融实验记录,能帮助读者理解ECG序列标注、心拍分类、QRS检测等任务的工程实现与调参思路。
1. 为什么 ECG 分析需要自己的 PyTorch 框架
心电信号处理和通用图像分类不一样,它是一维时间序列,采样率通常在 250Hz 到 1000Hz 之间,一次记录动辄数万乃至数十万个采样点。直接拿现成的 CNN 模型改改输入维度,往往达不到临床可用的精度。一个专门为 ECG 设计的 PyTorch 框架,核心要解决三件事:把多导联原始信号高效地切分成可训练样本、用适合时序的模型结构提取波形特征、在标注不均衡的现实数据集上稳定训练。这套东西做出来不只是一堆脚本,而是能支撑实验对比、模型迭代和最终部署的工程基础。本篇文章面向的读者是有 PyTorch 基础、想进入医疗 AI 方向,或者已经在处理生理信号但觉得通用框架不够顺手的人。你会看到一个从数据管线到模型训练再到评估的完整落地路径,每个环节都有可以直接跑的代码。
2. 数据准备:ECG 信号的加载策略与预处理管线
2.1 多导联 ECG 信号的内存布局与加载方式
ECG 数据最常见的格式是物理存储上的多通道数组,形状为(导联数, 采样点数)或(样本数, 导联数, 采样点数),采样率信息通常以元数据形式伴随存储。与图像数据不同,ECG 相邻采样点之间存在极强的时序相关性,读取时不能随意打乱单点顺序,必须按「样本」为单位进行切分。
我一般会为 ECG 框架单独封装一个ECGDataset类,继承torch.utils.data.Dataset。这个类做两件事:一是把原始信号按固定长度滑窗切分成样本,二是维护每个样本对应的患者 ID 和标签。滑窗长度不是随便定的,需要结合任务来选。心律失常分类通常取 2 到 10 秒,睡眠分期可能取 30 秒,房颤检测用 5 秒窗口就够。窗口太长会稀释局部特征,太短则丢失心律上下文。
import torch from torch.utils.data import Dataset import numpy as np class ECGDataset(Dataset): def __init__(self, signals, labels, window_size, stride=None): self.window_size = window_size self.stride = stride if stride else window_size self.samples = [] self.labels = [] for sig, lab in zip(signals, labels): # sig shape: (channels, time_points) for start in range(0, sig.shape[1] - window_size + 1, self.stride): self.samples.append(sig[:, start:start + window_size]) self.labels.append(lab) def __len__(self): return len(self.samples) def __getitem__(self, idx): x = torch.tensor(self.samples[idx], dtype=torch.float32) y = torch.tensor(self.labels[idx], dtype=torch.long) return x, y这段代码的关键参数是window_size和stride。window_size控制模型每次看到的信号长度,stride控制相邻窗口的重叠程度。当stride小于window_size时会产生重叠窗口,数据量增大但同时引入样本间相关性,训练时需要在随机采样层面解决,否则模型会过拟合到重复片段。stride等于window_size时是硬切分,样本间完全独立,适合预处理阶段快速出基线结果。
提示:加载大规模 ECG 数据时不要一次性把所有信号读进内存。用
np.memmap做磁盘映射,或者实现__getitem__内的按需读取,能显著降低内存压力。实际项目中一个 24 小时动态心电记录解压后可能超过 500MB,全量载入会影响训练迭代效率。
2.2 信号预处理:滤波、归一化与数据增强的选择
ECG 信号采集过程中混入的噪声主要有三类:基线漂移(频率低于 0.5Hz,由呼吸和电极移动引起)、肌电干扰(频率范围宽、幅度随机)、工频干扰(50Hz/60Hz 及谐波)。深度学习模型理论上能学习抵抗噪声,但在训练数据不足时,预处理能显著降低模型需要拟合的复杂度。
高通滤波去除基线漂移是必须做的第一步。截止频率设在 0.5Hz 到 1Hz 之间,低于这个频率的成分视为漂移。低通滤波看采样率,通常截止到 100Hz 或 150Hz,保留 ECG 主要能量集中的频段。陷波滤波器处理工频干扰,但数字陷波容易在 QRS 波群附近引入振铃效应,所以近年来的趋势是尽量少用陷波,改为在数据增强阶段加入噪声模拟,让模型自己学会鲁棒性。
归一化策略对 ECG 任务有特殊讲究。全局均值和标准差归一化的问题是,不同患者的信号幅度差异很大,同一个患者在不同时间段的幅度也会变化。更常用的是逐样本归一化,即对每个窗口单独减去均值除以标准差。这样处理后的信号幅度被压缩到相近范围,模型更容易跨患者泛化。
def normalize_per_sample(signal, eps=1e-8): # signal shape: (channels, time_points) mean = signal.mean(dim=1, keepdim=True) std = signal.std(dim=1, keepdim=True) return (signal - mean) / (std + eps)逐样本归一化有个副作用:它会抹掉不同导联之间的幅度比例关系。在心肌梗死定位等任务中,导联间相对幅度是重要诊断信息,这种情况下应该改用全局归一化,或者在逐样本归一化的同时额外把导联均值差作为特征输入。这个选择没有绝对的对错,取决于你的下游任务对幅度信息的依赖程度。
数据增强方面,我常用的手段包括:加入高斯白噪声(噪声标准差取信号标准差的 5% 到 15%)、时间轴小幅伸缩(resample 到 ±10% 倍率)、幅度随机缩放、以及导联随机遮蔽(将某一导联置零)。增强操作必须在窗口切分之后进行,且同一窗口内的增强参数要保持一致,否则会破坏心拍间的时序关系,模型学到的特征会产生偏差。
3. 用 PyTorch 构建 ECG 深度学习模型的核心架构
3.1 为什么通用图像模型不适用于 ECG
ResNet 在 ImageNet 上表现优异,但直接把它移植到 ECG 上一维化使用,通常比专门设计的 1D 模型差 3% 到 5% 的准确率。原因在于图像是空间局部性数据结构,卷积核关注的是 2D 邻域内的纹理组合;ECG 是时间序列,它的关键特征不仅存在于局部波形形态(比如 QRS 波群的宽度和振幅),还存在于中长程的时间依赖(比如 RR 间期变化模式)。ResNet 的下采样策略对时间序列来说过于激进,连续池化会把 QRS 波群的细节磨平,而这些细节恰是心律失常分型的重要依据。
适用于 ECG 的模型结构设计有两个方向。第一类是纯 CNN 结构,卷积核一维化,下采样倍数控制得比较温和。第二类是 CNN 加循环网络或注意力机制的组合结构,CNN 负责提取局部波形特征,循环网络或注意力层捕获跨时间段的关系。较早的 ECG 深度学习文献里 LSTM 是主流选择,近两年的趋势是换用 Transformer 的 self-attention 层,因为可以并行计算且能建模更长距离的依赖。
3.2 一个可运行的混合架构:CNN + BiLSTM + Attention
直接给一个我在实际项目中验证过的基础架构。输入是单导联 250Hz 采样率下 10 秒长度的信号,即 2500 个采样点。第一层 1D 卷积使用较大的卷积核来模拟带通滤波的效应,后面的卷积层逐渐缩小核尺寸提取更精细的模式。BiLSTM 捕捉前向和后向的上下文关系。最后用 attention 池化对 LSTM 输出加权求和,把可变长度的中间表示压缩成固定维度。
import torch.nn as nn import torch.nn.functional as F class ECGAnalysisModel(nn.Module): def __init__(self, num_classes=5, input_channels=1, lstm_hidden=64): super().__init__() # 模拟带通滤波效果的大核卷积 self.conv1 = nn.Conv1d(input_channels, 32, kernel_size=51, stride=2, padding=25) self.bn1 = nn.BatchNorm1d(32) self.conv2 = nn.Conv1d(32, 64, kernel_size=9, stride=2, padding=4) self.bn2 = nn.BatchNorm1d(64) self.conv3 = nn.Conv1d(64, 128, kernel_size=5, stride=2, padding=2) self.bn3 = nn.BatchNorm1d(128) self.lstm = nn.LSTM(128, lstm_hidden, bidirectional=True, batch_first=True) self.attention = nn.Sequential( nn.Linear(lstm_hidden * 2, 32), nn.Tanh(), nn.Linear(32, 1) ) self.classifier = nn.Linear(lstm_hidden * 2, num_classes) def forward(self, x): # x shape: (batch, channels, time) x = F.relu(self.bn1(self.conv1(x))) x = F.relu(self.bn2(self.conv2(x))) x = F.relu(self.bn3(self.conv3(x))) # 转成 LSTM 输入格式: (batch, seq_len, features) x = x.permute(0, 2, 1) lstm_out, _ = self.lstm(x) # (batch, seq_len, hidden*2) attn_scores = self.attention(lstm_out).squeeze(-1) # (batch, seq_len) attn_weights = F.softmax(attn_scores, dim=1) context = torch.bmm(attn_weights.unsqueeze(1), lstm_out).squeeze(1) return self.classifier(context)卷积层参数的设计逻辑:kernel_size=51在 250Hz 采样率下对应约 204ms 的时间跨度,恰好覆盖一个典型 QRS 波群的宽度(80ms 到 120ms)加上一定的容差,这样第一个卷积层就能直接捕捉心拍形态。stride=2将序列长度减半,三次下采样后 2500 个点变为 313 个点,这个降采样比例既控制了 LSTM 的计算量,又保留了足够的时序分辨率。LSTM 的bidirectional=True很关键,ECG 中许多异常形态(比如早搏)的判定需要同时参考其前后心拍的节律关系,双向结构让每个时间步的输出同时携带过去和未来的上下文信息。
Attention 的作用是对 LSTM 输出的 313 个时间步做加权汇总。普通 mean pooling 会稀释异常波形的位置信息,而 attention 可以学到「哪些时间段的信号对最终分类更重要」。在实际训练中你会发现,模型学到的注意力权重往往集中在 QRS 波群附近,这符合心电学专家的判读习惯。
3.3 损失函数与类别不均衡的处理
ECG 数据集的类别分布极不均衡。正常窦性心律可能占 80% 以上,而某些心律失常类型占比不到 1%。如果直接用交叉熵损失,模型会倾向于把所有样本预测为多数类,整体准确率看着很高,但少数类的召回率可能接近于零。
处理不均衡问题,第一选择不是欠采样或过采样,而是改用加权交叉熵损失。权重设置为N / (num_classes * N_c),即每个类别的样本数取倒数后归一化。这个方案不需要改动数据加载逻辑,训练过程中每个 batch 的计算方式不变,只是给少数类的梯度贡献放大了倍数。
from collections import Counter def build_class_weights(labels): counter = Counter(labels) num_samples = len(labels) num_classes = len(counter) weights = [num_samples / (num_classes * counter[i]) for i in range(num_classes)] return torch.tensor(weights, dtype=torch.float32) # 使用方式 class_weights = build_class_weights(all_labels).to(device) criterion = nn.CrossEntropyLoss(weight=class_weights)加权交叉熵是基线方案。它的问题在于少数类样本在训练中被反复看到,模型可能过拟合到某些噪声模式上。当加权损失训练出的模型在验证集上表现不佳时,可以考虑 Focal Loss。Focal Loss 在交叉熵基础上引入(1 - p_t)^gamma调制因子,让模型把注意力集中在难分类的样本上,gamma通常取 2.0。它相当于一个更平滑的难例挖掘机制,不会像加权交叉熵那样直接放大所有少数类样本的梯度,而是根据样本当前的预测置信度动态调整。
4. 训练循环与验证策略:稳定复现的关键操作
4.1 按患者划分数据集而不是按样本划分
ECG 模型评估中最常见的错误是按样本随机划分训练集和测试集。同一个患者的多个心拍片段会被同时分到两边,模型等于「见过」这个患者的心电形态。由于患者个体差异很大,这种划分方式会让测试集准确率虚高 5 到 10 个百分点,在真实部署环境中完全不堪用。
正确的做法是按患者 ID 划分,保证同一个人的所有样本只出现在一个集合中。举个例子,如果有 100 个患者,可以按 70/15/15 的比例分为训练集、验证集和测试集。我在实际项目中还会做一个额外的约束:同一个患者可能有多次不同时间的记录,这些记录也必须划到同一个集合里,否则仍然存在信息泄露。
from sklearn.model_selection import GroupShuffleSplit def split_by_patient(patient_ids, test_size=0.2): splitter = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=42) indices = np.arange(len(patient_ids)) train_idx, test_idx = next(splitter.split(indices, groups=patient_ids)) return train_idx, test_idxGroupShuffleSplit的核心参数是groups,传入每个样本对应的患者 ID 数组。它会保证同一个 group 的样本不会同时出现在训练集和测试集中。n_splits置为 1 表示只需一次划分,random_state固定下来便于实验复现。验证集可以从训练集中再按同样方式切一次,或者单独调用一次split_by_patient,比例按实际需要调整。
注意:如果测试集样本数量太少导致评估指标波动大,可以用 K-Fold 交叉验证。ECG 场景下我用的是 StratifiedGroupKFold,它同时考虑类别分布和患者分组两个约束,是 sklearn 里比较冷门但 ECG 任务高频使用的工具。
4.2 PyTorch 训练循环的工程化封装
训练循环不能只写一个 for 循环就完事。ECG 实验周期长,动辄训练几十个 epoch,中途可能遇到显存溢出、学习率设置不当、loss 发散等各种问题。一个工程化的训练循环至少需要包含:梯度裁剪、学习率调度、指标记录、周期性 checkpoint 保存。
梯度裁剪对于 LSTM 部分是必选项,因为循环网络在反向传播时容易出现梯度爆炸,表现为 loss 突然跳到 NaN。设置max_norm=5.0或max_norm=10.0是一个安全默认值,GRU 或 LSTM 层数越多,裁剪阈值应该越小。
def train_one_epoch(model, dataloader, optimizer, criterion, device, clip_value=5.0): model.train() total_loss = 0.0 for batch_idx, (x, y) in enumerate(dataloader): x, y = x.to(device), y.to(device) optimizer.zero_grad() logits = model(x) loss = criterion(logits, y) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=clip_value) optimizer.step() total_loss += loss.item() return total_loss / len(dataloader)clip_grad_norm_的第二个参数max_norm是梯度范数的上限。它的计算方式是先求所有参数的梯度 L2 范数,如果超过阈值,就按比例缩放所有梯度。这个操作只影响梯度的幅度,不改变梯度的方向,因此不会影响收敛的最终目标。关键参数clip_value的设置原则:当训练初期 loss 就出现大幅波动时,优先减小这个值;当训练稳定但收敛缓慢时,可以尝试增大。
学习率调度方面,我推荐用CosineAnnealingWarmRestarts而不是 StepLR。ECG 模型的损失曲面通常比较崎岖,周期性重启学习率有助于跳出局部最优点。初始学习率设为1e-3配 Adam 优化器、weight_decay=1e-4,这是个适用于多数 1D 卷积加循环网络架构的起点。如果模型包含预训练的卷积模块,预训练部分的学习率应该乘 0.1,避免破坏已经学好的底层特征。
4.3 评估指标要用对:不止是 Accuracy
ECG 分类任务特别是心律失常检测,公开数据集测试集上的准确率能到 95% 以上,但这没有实际参考意义,因为类别不均衡严重。准确率的计算方式默认把所有类别等权对待,模型把 1% 的异常全部判错,准确率依然很高。
正确的评估视角是混淆矩阵以及由此推导出的 Sensitivity 和 Specificity。敏感性等于TP / (TP + FN),衡量的是「真正有病的人里有多少被找出来」;特异性等于TN / (TN + FP),衡量的是「健康人里有多少被正确排除」。在医疗场景中这两个指标要同时报告,它们之间的权衡直接由分类阈值控制——默认的argmax等价于阈值 0.5,但这个阈值往往不是最优的。
from sklearn.metrics import roc_auc_score, precision_recall_curve, roc_curve def find_optimal_threshold(y_true, y_prob): precision, recall, thresholds = precision_recall_curve(y_true, y_prob) f1_scores = 2 * precision * recall / (precision + recall + 1e-8) best_idx = f1_scores.argmax() return thresholds[best_idx] def evaluate_model(model, dataloader, device, threshold=0.5): model.eval() all_probs = [] all_labels = [] with torch.no_grad(): for x, y in dataloader: logits = model(x.to(device)) probs = torch.softmax(logits, dim=1) all_probs.append(probs.cpu().numpy()) all_labels.append(y.numpy()) probs = np.concatenate(all_probs) labels = np.concatenate(all_labels) auc = roc_auc_score(labels, probs, multi_class='ovr') return probs, labels, aucroc_auc_score在二分类场景下只需要传原始概率,不需要传预测标签。multi_class='ovr'适用于多分类,它的计算方式是逐类别做 one-vs-rest 的 AUC 再取平均。实际调阈值时,我通常输出概率矩阵后对每个类别分别用 precision-recall 曲线找最优阈值,而不是所有类别共用一个阈值。这个细节对少数类的召回率影响很大,因为多数类的默认 0.5 阈值通常已经足够好,而少数类往往需要更低的阈值才能达到可接受的敏感性。
5. 框架的进阶技巧与部署验证
5.1 多导联输入的通道合并策略
前面的模型实例用的是单导联输入。MIT-BIH 数据集只有两条导联,而临床 12 导联系统提供的信息远多于单导联。多导联处理有两种常见做法:第一种是把导联作为通道维度,直接用Conv1d处理,输入通道数就等于导联数,模型自动学习导联间的空间相关性;第二种是每个导联独立过共享权重的特征提取器,然后把特征拼接后送入分类层。
第二种方案在实践中表现更好,因为不同导联的波形形态差异很大(肢体导联和胸导联看到的电轴方向不同),共享权重的卷积核可以为每个导联提取阶段特征,再通过融合层学习跨导联关系。实现时只需将forward函数改为逐导联处理:
def forward_multilead(self, x): # x shape: (batch, leads, time) lead_features = [] for i in range(x.shape[1]): single_lead = x[:, i:i+1, :] # 取单个导联 feat = self.feature_extractor(single_lead) lead_features.append(feat) fused = torch.cat(lead_features, dim=1) return self.classifier(fused)多导联情况下模型的参数量会上升,但输入位置不变,推理时间的增加主要体现在特征提取的重复计算上。反向传播时梯度会同时流向各个导联分支,因此不需要额外的损失设计。
5.2 对抗验证:检测数据集划分泄露
按患者划分后仍可能出现一个问题:训练集和测试集之间存在隐含的相关性,比如来自同一医院同一台设备的数据。对抗验证能定量检测这种泄露。方法很简单:在训练集数据上打标签 0,测试集数据上打标签 1,训练一个二分类器去区分两组数据。如果二分类器的 AUC 接近 0.5,说明两个集合不可分,划分是干净的;如果 AUC 明显高于 0.5(超过 0.8),说明数据分布存在系统性差异,模型可能是在「记住设备特征」而不是「学习心电特征」。
def adversarial_validation(train_data, test_data, device): labels = torch.cat([ torch.zeros(len(train_data)), torch.ones(len(test_data)) ]) combined = torch.cat([train_data, test_data], dim=0) # 训练一个简单的两层分类器,输入是数据统计特征而非原始信号 # 常用特征:均值、方差、峰峰值、QRS 波群数量等 return auc_value对抗验证的结果不能直接「修复」,但能提示你检查数据来源。如果发现设备相关的泄露,可以考虑在预处理阶段加入实例归一化来消除设备间的增益差异,或者干脆收集更多样化的数据源。
5.3 模型导出与部署格式选择
训练完成后要走出实验环境,PyTorch 模型导出有两种主流方式。model.state_dict()保存权重字典,适合实验室内部继续训练。部署场景推荐用torch.jit.trace或onnx.export,二者都把模型固化成计算图,推理时不再依赖 Python 层。
使用torch.jit.trace时要注意:trace 只在给定示例输入上运行一次,如果模型包含数据相关的分支(比如根据输入长度走不同路径),trace 后的模型可能行为异常。对于时序模型有一个重要约束,trace 模型输入长度不能大于 trace 时的示例长度,否则 LSTM 层的时间步数不匹配。我的做法是直接导出 ONNX,然后走的推理框架一般都能用固定长度输入来拿到稳定的吞吐指标。
验证导出的模型与原始 PyTorch 模型输出是否一致,用数值比对而不是目测:
import onnxruntime as ort import torch def verify_export(torch_model, onnx_path, test_input): torch_model.eval() with torch.no_grad(): ref_output = torch_model(test_input).numpy() sess = ort.InferenceSession(onnx_path) onnx_output = sess.run(None, {sess.get_inputs()[0].name: test_input.numpy()})[0] max_diff = np.abs(ref_output - onnx_output).max() print(f"Max difference: {max_diff:.6e}") assert max_diff < 1e-4, "Export mismatch detected"比对阈值的设定有讲究。1e-4是相对保守的阈值,如果模型中有 BatchNorm 层且推理和训练模式切换不当,或者 ONNX 的算子精度是 float16,这个检查会失败。浮点计算顺序的微小差异允许1e-2以内,但如果差异到1e-1量级,基本可以判定导出过程出了问题。对于 ECG 分类型任务,logits 的微小偏差通常不影响argmax的最终结果,但触发阈值判断的回归任务必须严格比对。
最后一个建议是框架整体保存配置:模型参数、预处理参数、归一化参数、标签映射表要一并序列化。我在实际项目中遇到过只迁移模型权重、结果分类错乱的问题,原因就是标签顺序没有对齐。把这些元数据与你保存的模型权重放在同一个字典结构里,每次加载都先校验再推理。
本文还有配套的精品资源,点击获取