ARTICLE DETAIL

资讯详情

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

用PyTorch搭建ECG深度学习框架:从波形预处理到CPU推理

用PyTorch搭建ECG深度学习框架:从波形预处理到CPU推理 简介基于PyTorch实现的ECG深度学习框架面向深度学习和医疗AI方向的研究者与开发者聚焦心电图信号处理与识别。资源以PyTorch动态计算图为核心围绕CNN与RNN如LSTM构建混合模型系统覆盖数据清洗去噪、标准化分段、模型定义、损失函数与优化器选择、训练验证、评估以及模型保存与ONNX部署等关键环节。包内共749个文件以py源代码、mat数据文件、hea头文件、atr标注文件及dat原始记录为主配有pdf文档、png/svg架构图、md说明及ipynb示例便于对照学习。项目采用data、models、preprocess、train.py、evaluate.py、config.py等标准结构可直接复用或改造。压缩包仅22.76MB轻量易用。已有314人学习下载适合希望系统掌握PyTorch在ECG分类、节律检测等任务中落地方法的读者参考。1. 为什么用 PyTorch 搭建 ECG 深度学习框架长程动态心电图和 ICU 监护仪是 ECG 深度学习最高频的两类产出场景前者要处理 7×24 小时的单导联流式记录后者面对的是 12 导联、高采样率、逐秒判读的实时预警需求。这两种场景的数据形态、延迟要求和标签尺度完全不同但用 PyTorch 写起来只需要替换数据入口和分类头主干网络可以被两套任务共享。PyTorch 在这一领域比其他框架更顺手的点有两个一是序列信号预处理可以直接通过 Dataset 接进训练管道不用单独搭数据中转服务二是张量形状检查很方便对 1D-CNN 这种把波形长度逐层压缩的网络来说中间层尺寸出了问题一眼就能发现。标题里说的“ECG 深度学习框架”我理解为一条从原始波形到可部署模型的完整管线包括滤波切窗、导联对齐、模型主干、训练验证和模型导出。这篇文章按这条线展开带可复现的 PyTorch 代码和关键参数。新手可以按章节顺序搭建自己的版本有经验的工程师可以重点看患者独立划分、模型导入陷阱和单导联迁移这几节。2. 信号入口ECG 波形到 PyTorch DataLoader 的处理链路2.1 采样率、重采样与陷波先把原始波形整理成干净的一维张量原始心电数据的格式五花八门医院设备常见 EDF 或 XML 导出可穿戴设备通常是私有二进制格式采样率从 125Hz 到 1000Hz 都有。模型本身不关心原始采样率但卷积核的尺寸和池化步长都隐含了对特征频率的假设所以框架里需要把训练数据统一到一个采样率。我一般选 250Hz对 QRS 波的形态识别足够训练速度和内存占用又比 500Hz 少一半。如果要做精细的 ST 段分析才提高 500Hz。信号预处理按两步顺序执行。第一步带通滤波范围用 0.5Hz45Hz目的是去掉基线漂移和肌电干扰第二步陷波滤波针对电网工频国内是 50Hz。以下是一个可以直接放进框架的预处理函数import numpy as np from scipy import signal def preprocess_ecg(raw_wave: np.ndarray, fs_in: int, fs_out: int 250) - np.ndarray: # 1. 重采样线性插值先统一采样率再做滤波 if fs_in ! fs_out: n_out int(len(raw_wave) * fs_out / fs_in) raw_wave signal.resample(raw_wave, n_out) # 2. 带通滤波0.5Hz 高通去基线漂移45Hz 低通去高频肌电 sos signal.butter(4, [0.5, 45], btypebandpass, fsfs_out, outputsos) clean signal.sosfiltfilt(sos, raw_wave) # 3. 陷波工频 50HzQ30 控制带宽尽量窄 notch_b, notch_a signal.iirnotch(50, Q30, fsfs_out) clean signal.sosfiltfilt(signal.tf2sos(notch_b, notch_a), clean) return clean.astype(np.float32)注意处理顺序重采样必须在滤波之前否则滤波器的截止频率是相对原始采样率计算的重采样之后频率轴就错位了。sosfiltfilt 是零相位滤波输出信号没有群延迟这对后续需要精确对齐 QRS 波位置的任务很关键。普通分类任务用了带延迟的普通滤波可能无感但如果要做心拍分割延迟积累几毫秒就会错位。2.2 导联维度的统一单导联和 12 导联怎么喂进同一个模型PyTorch 里 1D-CNN 的标准输入形状是[batch, channels, seq_len]这里的 channels 就是导联数。单导联数据直接以[batch, 1, seq_len]输入12 导联则是[batch, 8, seq_len]标准 12 导联实际只有 8 个独立采集通道其余 4 个由数学推导得到。在框架里做多导联兼容时不要用补零把单导联强行扩展成 8 通道。补零通道会让第一层卷积的输出偏移因为卷积核会去响应全零输入相当于引入一个固定偏置。更合理的做法是让模型第一层的in_channels是一个可配置参数训练时按数据源传入。这样单导联和 12 导联各自训练各自的入口层但共享后续主干网络是一种很实用的权重复用方式。2.3 带防重叠窗口切分的 Dataset 实现动态心电图的判读单位通常是 10 秒或 30 秒片段30 秒窗口是房颤检测和心拍分类的常见选择。窗口切分有两个细节值得注意一是相邻窗口要留重叠这样可以缓解窗口边界切到异位心拍导致的误标问题二是在 Dataset 里预先算好窗口索引不要在__getitem__里反复计算。以下是框架里可以直接使用的 Dataset 代码import torch import numpy as np from torch.utils.data import Dataset class ECGWindowDataset(Dataset): def __init__(self, records, labels, window_len: int 7500, stride: int 3750): # window_len 7500 对应 250Hz 采样下 30 秒 # stride 取 window_len 的一半形成 50% 重叠窗口 self.records records self.labels labels self.window_len window_len self.stride stride self.index_map [] for rec_idx, rec in enumerate(records): n_windows max(1, (len(rec) - window_len) // stride 1) for w in range(n_windows): start w * stride self.index_map.append((rec_idx, start)) def __len__(self): return len(self.index_map) def __getitem__(self, idx): rec_idx, start self.index_map[idx] x self.records[rec_idx][start:start self.window_len] if len(x) self.window_len: x np.pad(x, (0, self.window_len - len(x))) x torch.from_numpy(x).float() x x.unsqueeze(0) # [1, window_len] y torch.tensor(self.labels[rec_idx], dtypetorch.long) return x, ystride window_len // 2意味着相邻窗口有一半数据重叠训练样本量几乎翻倍代价是每个样本之间有相关性会轻微增加过拟合风险。如果数据量本身足够大可以放宽到stride window_len完全无重叠。使用DataLoader时建议设置pin_memoryTrue和num_workers4以上因为心电数据量大数据加载常常比 GPU 计算更慢。2.4 针对心电的在线增强手段心电数据增强不能照搬图像那套。常用的有三种时间平移、幅度缩放、高斯噪声叠加。时间平移的幅度不应超过一个 RR 间期的三分之一否则会人为扭曲心搏的相对位置干扰模型对节律的判断。幅度缩放控制在 0.851.15 之间。噪声注入用信号标准差 5% 左右的高斯噪声用来模拟可穿戴设备佩戴晃动时的运动伪差。增强的强度要按任务区分心律失常分类对节律敏感时间平移要保守单心拍形态分类对幅度敏感噪声注入要保守。增强的目的不是让模型万无一失而是让训练分布覆盖测试时可能遇到的数据偏移过度增强反而会损害性能。3. 模型主干1D-CNN、注意力头与通道对齐的设计3.1 为什么不用 ResNet 套时频图把 ECG 转成时频图再用 2D ResNet 分类在一些公开竞赛里表现不错但工程落地时不推荐。原因有三个时频变换引入了 FFT 窗口长度和重叠比例两个新超参数增加调参成本2D 模型参数量普遍在 5M 以上在嵌入式设备上推理会触发内存瓶颈最深层的理由是 ECG 的强特征如 QRS 波、T 波是沿时间轴的形态学特征直接做 1D 卷积天然更高效。1D-CNN 在同等精度下参数量通常只有 2D 模型的四分之一。3.2 可复现的 ECGNet 主干下面这个模型采用四层一维卷积加多头注意力的结构在房颤和室性早搏分类任务上表现稳定。卷积通道逐步翻倍序列长度逐步压缩注意力层负责捕获长程 RR 间期的全局依赖import torch import torch.nn as nn class ECGNet(nn.Module): def __init__(self, in_channels: int 1, num_classes: int 4, dropout: float 0.3): super().__init__() # 特征金字塔4 个卷积阶段逐步压缩序列长度 self.features nn.Sequential( nn.Conv1d(in_channels, 32, kernel_size15, stride4, padding7), # 7500 - 1875 nn.BatchNorm1d(32), nn.ReLU(inplaceTrue), nn.Conv1d(32, 64, kernel_size9, stride4, padding4), # 1875 - 469 nn.BatchNorm1d(64), nn.ReLU(inplaceTrue), nn.Conv1d(64, 128, kernel_size5, stride4, padding2), # 469 - 118 nn.BatchNorm1d(128), nn.ReLU(inplaceTrue), nn.Conv1d(128, 256, kernel_size5, stride2, padding2), # 118 - 59 nn.BatchNorm1d(256), nn.ReLU(inplaceTrue), ) # 多头注意力把 CNN 输出的局部特征做全局建模 self.attn nn.MultiheadAttention(embed_dim256, num_heads4, batch_firstTrue) self.classifier nn.Sequential( nn.Dropout(dropout), nn.Flatten(), nn.Linear(256 * 59, 256), nn.ReLU(inplaceTrue), nn.Dropout(dropout), nn.Linear(256, num_classes), ) def forward(self, x): x self.features(x) # [B, 256, 59] x x.transpose(1, 2) # [B, 59, 256] x, _ self.attn(x, x, x) # 自注意力QKV x x.transpose(1, 2) # [B, 256, 59] return self.classifier(x)参数选择的逻辑第一层卷积 kernel_size15在 250Hz 采样下对应 60ms和 QRS 波群的典型宽度接近能有效提取心搏形态特征。stride4 是刻意选择让序列长度每层压缩约 4 倍30 秒波形到最后只剩 59 个时间步。分类前的全连接层把 256×59 展平成 15104 维再接一个宽隐藏层来整合信息这里是模型参数量最大的部分。注意力模块的 num_heads4 在序列长度仅 59 的情况下足够使用了增大到 8 不会带来明显收益反而增加显存占用。3.3 输出形状与感受野参考下表基于输入为[1, 1, 7500]也就是单导联、250Hz、30 秒波形的情况可以用 torchsummary 或直接打印每层张量形状来验证层名称输出尺寸参数量感受野约Conv1-1[1, 32, 1875]约 0.5K75msConv1-2[1, 64, 469]约 18K300msConv1-3[1, 128, 118]约 41K900msConv1-4[1, 256, 59]约 164K1.8s注意力全连接[1, 4]约 2.2M全局感受野从 75ms 增长到 1.8s 再靠注意力模块覆盖全局这个设计保证底层专注于 P-QRS-T 形态而高层能感知到节奏的周期性变化。参数量集中在最后的全连接层做嵌入式部署时这也是优先做剪枝的位置。4. 训练配置损失、优化器与患者独立验证4.1 类别不均衡为什么默认用 Focal Loss动态心电记录中正常窦性心律通常占比超过 90%心律失常的少数类别才是临床关注的重点。最直观的做法是加权交叉熵但权重设置稍大就会让模型倾向于把所有样本预测为多数类来降低整体损失。Focal Loss 更适合这种场景它通过调制因子降低对高置信度样本的关注把剩余的训练容量分配给难分类样本对房颤这类形态变异大的类别更友好class FocalLoss(nn.Module): def __init__(self, alpha: float 0.75, gamma: float 2.0, class_weightsNone): super().__init__() self.alpha alpha self.gamma gamma self.class_weights class_weights def forward(self, logits: torch.Tensor, targets: torch.Tensor) - torch.Tensor: ce nn.functional.cross_entropy( logits, targets, weightself.class_weights, reductionnone ) pt torch.exp(-ce) loss self.alpha * (1 - pt) ** self.gamma * ce return loss.mean()gamma 值调大难例的梯度占比越高但训练早期的收敛会变慢。默认 gamma2.0 一般够用alpha0.75 表示模型会把注意力向少数类倾斜。如果验证集里少数类的召回率始终上不去可以先把 alpha 提高到 0.85而不是动 gamma。需要注意数据增强和 Focal Loss 一起使用时容易被增强扭曲的样本会获得高梯度如果发现训练曲线震荡明显先减小时间平移幅度。提示类别权重不要直接按样本数反比来设对严重不均衡的数据这样会把多数类的梯度压得太低导致正常心拍也判断错误。建议从反比权重乘以 0.5 到 0.8 之间调起。4.2 优化器与学习率调度一维心电信号分类任务上AdamW 配合 OneCycleLR 的使用体验明显优于固定步长衰减。OneCycleLR 先用热身阶段让 BatchNorm 统计量跑稳再维持高学习率最后快速衰减整体训练时间能缩短约百分之三十。这个优势在处理 24 小时动态数据时很有价值数据量大导致单个 epoch 耗时较长收敛慢会拖慢整个开发周期import torch optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr3e-3, total_stepslen(train_loader) * epochs, pct_start0.3, ) for batch in train_loader: x, y batch optimizer.zero_grad() logits model(x) loss criterion(logits, y) loss.backward() optimizer.step() scheduler.step()pct_start0.3表示训练前 30% 的步数把学习率从初始值升到max_lr之后逐步衰减到接近零。max_lr不要超过lr的三倍否则前几个 batch 的 loss 容易跳到 Nan。weight_decay 固定 1e-4不需要随任务频繁调整。4.3 患者独立划分这是 ECG 框架泛化能力的分水岭这个问题决定了验证集的数据能不能真实反映模型上线效果。如果随机把所有样本混在一起洗牌划分同一个患者的不同时间段片段会同时出现在训练集和验证集里。模型会把患者 ID 的个体特征当作疾病特征学进去验证精度虚高 5 到 10 个百分点都是常见的。正确做法是使用 GroupShuffleSplit 按患者分组切分训练集和验证集import numpy as np from sklearn.model_selection import GroupShuffleSplit patient_ids np.array([rec[patient_id] for rec in records]) split GroupShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(split.split(records, labels, groupspatient_ids))关键点在于groups参数传的是患者 ID而不是记录片段 ID。验证集里的患者必须完全不出现在训练集中这样模型面对新患者时才有真实参考价值。如果数据集中每个患者只有一段记录那就需要用更严格的留一患者法来评估。5. 落地模型导出、CPU 推理与版本迁移注意事项5.1 用 ONNX 导出 ECGNet训练完成之后常见的落地路径是导出 ONNX再用 ONNX Runtime 在病房 CPU 服务器上做推理。导出前有两件事必须做模型切到eval()模式以及准备一个固定长度的 dummy 输入。下面的代码会把 ECGNet 导出为带动态 batch 维度的 ONNX 文件model.eval() dummy_input torch.randn(1, 1, 7500) torch.onnx.export( model, dummy_input, ecgnet.onnx, input_names[ecg_wave], output_names[logits], dynamic_axes{ ecg_wave: {0: batch}, logits: {0: batch}, }, opset_version17, )opset_version 至少要用 17因为这个版本开始 MultiheadAttention 的导出才比较稳定。PyTorch 2.x 环境下可以在 export 调用前加torch.backends.cudnn.enabled True正常情况下 ONNX Runtime 会返回一次成功的导出并用空白警告说明哪些节点是 fallback 的那些警告不影响推理结果。5.2 CPU 推理最小实现病房服务器上通常没有 GPU用 ONNX Runtime 的 CPU 推理就可以满足单路或几路并发的需求。预处理要放在 PyTorch 之外单独做原因是 scipy 的滤波函数在 Python 层面的开销比模型推理本身还高把预处理和推理分离后更容易定位性能瓶颈import onnxruntime as ort import numpy as np sess ort.InferenceSession(ecgnet.onnx, providers[CPUExecutionProvider]) wave preprocess_ecg(raw_wave, fs_in500, fs_out250) # 预处理模块 wave_tensor wave.reshape(1, 1, -1).astype(np.float32) logits sess.run([logits], {ecg_wave: wave_tensor})[0] pred int(np.argmax(logits, axis1)[0])实际工程里预处理往往比模型推理更耗时。如果每秒需要处理多段数据把原始波形批量读入内存一次性做 bandpass 和重采样避免逐段调用 scipy性能差距可能达到十倍。一段 30 秒的单导联波形CPU 单核上的模型推理通常在 20ms 以内预处理做到与推理同量级是合理的优化目标。5.3 PyTorch 版本迁移时模型导入的坑不同版本的 PyTorch 之间迁移pth文件最常见的麻烦是MultiheadAttention的权重键名和 BatchNorm 的缓冲区字段在不同版本间有过改名。直接用load_state_dict会抛出 key mismatch。稳妥的做法是允许strictFalse加载后手动检查缺了多少键然后决定是重新初始化还是接续训练import torch state torch.load(ecgnet_v1.pth, map_locationcpu, weights_onlyTrue) missing, unexpected model.load_state_dict(state, strictFalse) print(missing:, len(missing), unexpected:, len(unexpected))如果 missing 和 unexpected 的数量都在个位数大概率只是注意力层的 bias 或 BN 的num_batches_tracked字段对不上。加载成功后建议用一小段验证数据跑一遍推理对比旧版本输出的 argmax 结果是否一致不要直接认为能加载就等于行为一致。6. 实战把 ECGNet 压到单核 CPU 推理 10ms 的一个捷径6.1 先动全连接层而不是卷积层ECGNet 的参数大头在分类器里的第一个 Linear 层也就是 256×59 到 256 的那一层。它对推理内存的影响远大于卷积层但并不是每个神经元都有同样的重要性。一个低成本方法是计算每个输出神经元的平均激活值把激活长期接近零的神经元剪掉。这个操作理论上等价于移除全连接层的一整列权重配合torch.nn.utils.prune的 L1 结构化剪枝就能做到不需要额外的依赖import torch.nn.utils.prune as prune # 对分类器第一个全连接层做 L1 结构化剪枝剪掉 40% 输出通道 prune.l1_unstructured(model.classifier[2], nameweight, amount0.4) prune.remove(model.classifier[2], nameweight)prune.remove这一步很多人会漏掉它把剪枝后的权重固化回参数张量避免推理时仍然走 masked 计算的慢路径。剪枝后需要在验证集上跑一遍精度如果掉点超过 1 个百分点把 amount 回退到 0.3。这种方式对全连接层的压缩力度大而卷积层结构简单剪不好会直接破坏波形特征的局部性我不太建议动。6.2 量化时只量化卷积层保留注意力浮点ECGNet 走量化感知训练时一个容易踩的坑是把注意力模块也一起量化。注意力中 softmax 和矩阵乘的误差会被后续层放大量化后掉点通常比卷积层严重得多。常见做法是只对features里的卷积层做量化注意力模块保持 float32。PyTorch 2.x 下的写法model.qconfig torch.ao.quantization.QConfig( activationtorch.ao.quantization.FakeQuantize.with_args( observertorch.ao.quantization.MinMaxObserver, quant_min0, quant_max255, dtypetorch.quint8 ), weighttorch.ao.quantization.default_weight_observer, ) model torch.ao.quantization.prepare_qat(model, inplaceFalse) # 继续训练 1-2 个 epoch让模型适应量化噪声 model torch.ao.quantization.convert(model, inplaceFalse)量化感知训练不能用太少的数据否则统计出来的量化范围不准推理时会出现个别异常点把动态范围撑大的情况。转完后建议测三段不同患者的数据一段正常、一段房颤、一段室性早搏比整体准确率更能反映量化是否可靠。6.3 用输出一致性校验整套流程最后一步是把 PyTorch 原始模型和 ONNX 量化模型的输出做逐样本对齐。不要只比较 argmax而要比较 softmax 输出向量的余弦相似度。写一个脚本固定十段波形分别用原始模型和部署模型跑一遍记录每个类别的概率分布。如果相似度低于 0.98就要检查预处理是否在两端一致特别是重采样参数和滤波器的截止频率是否被改动了。经过剪枝和量化后ECGNet 在单核 CPU 上的单段推理通常能控制在 8ms 到 12ms这个量级已经足够在一个普通网关机上同时跑三四路实时心电监测。本文还有配套的精品资源点击获取
返回列表