简介:这份资源围绕LSTM-GAN生成逼真ECG信号展开,面向具备Python与深度学习基础、关注生物医学信号处理与数据增强的研究者和开发者。项目以生成对抗网络为核心,生成器负责合成心电波形,判别器负责辨别真伪,LSTM则用于捕捉ECG的周期性与时序模式,可用于异常检测算法测试或扩充训练数据集。压缩包共13个文件,约4.46MB,包含5个py脚本、3张png结果图、2个h5权重文件,以及1个ipynb交互式笔记、1个md说明文档和1个gitignore配置,覆盖模型定义、训练、测试与信号清理等环节。已有313人学习下载。读者可借此理解LSTM-GAN在序列数据建模中的完整实现路径,参考生成器与判别器的权重保存方式,并通过可视化图像对比生成信号与真实信号的差异,适合作为医学信号生成方向的入门实践素材。
1. 从一份 Jupyter Notebook 说起:LSTM-GAN 怎么造出"似是而非"的 ECG
拿到"用于生成似是而非的ECG信号的LSTM-GAN_Jupyter Notebook_Python_下载.zip"这个标题,很多人第一反应是:ECG 信号也能造假?能,而且这件事在临床上并不新鲜——动态心电 Holter 的算法验证、可穿戴设备的压力测试、教学演示里需要大量"看起来像但不对应任何真实病人"的心电数据,都绕不开合成 ECG。问题在于,普通 GAN 生成的波形要么形态崩坏、要么节律乱跳,一眼假;而 LSTM-GAN 的思路是让生成器带记忆地逐点吐出一段心拍序列,再让判别器去分辨"这段波形像不像真的心电"。所谓"似是而非",就是它具备 P 波、QRS 波群、T 波的基本形态和 RR 间期节律,但不对应任何真实个体,也不能用于诊断。这篇笔记面向的是手里有 Python 环境、想跑通一份 Jupyter Notebook 合成 ECG 的工程师和算法爱好者,从环境、数据、模型结构一路讲到参数和翻车点。
2. LSTM-GAN 合成 ECG 的原理与选型:为什么不是普通 GAN
2.1 心电信号的时序特性决定了生成器要带记忆
一段 10 秒的 ECG 在 250Hz 采样率下是 2500 个点,每个点都和前后的点强相关:QRS 波群的陡峭上升沿、T 波的缓慢回落,都是连续几十个采样点共同构成的形态。普通全连接 GAN 的生成器把噪声一次性映射成整段波形,它没有"上一个点是什么"的概念,于是生成出来的信号经常在 QRS 中间突然断裂,或者两个心拍之间出现不合理的平直线。LSTM 的隐藏状态天然适合这种场景——它逐时间步输出采样点,每一步都能"记得"前面已经画到心拍的哪个阶段,因此能维持波形的连续性。
从选型角度看,常见做法有三种:一是纯 LSTM 做自回归生成,简单但容易收敛到均值波形,生成结果千篇一律;二是普通 GAN 直接生成整段,快但形态差;三是 LSTM 当生成器、CNN 或 LSTM 当判别器的组合,兼顾时序建模和对抗训练。这份 Notebook 走的是第三条路,也是目前合成生理信号里比较稳的方案。判别器这边,如果也用 LSTM,它同样能捕捉时序依赖,判断"这段节律是否自然";如果用一维 CNN,则更擅长抓局部形态(比如 QRS 是否够窄够尖)。两种都能用,Notebook 里通常会给一个可切换的判别器实现。
2.2 对抗训练在生理信号上的特殊之处
普通图像 GAN 的判别器看的是像素,ECG 的判别器看的是"这段波形在生理上说不说得通"。这就带来一个麻烦:判别器很容易靠几个低级特征(比如整体幅值范围、是否有一条基线漂移)就把真假分开,导致生成器只学会"把幅值调到合理区间",形态依然一塌糊涂。缓解办法是在判别器输入里做归一化,让每条样本都独立标准化到零均值单位方差,逼判别器去看形态而不是绝对幅值。另一个办法是控制对抗训练的节奏,别让判别器太强——判别器 loss 掉到接近 0 的时候,生成器基本就学不动了。
提示:合成 ECG 只用于算法测试、教学和数据增强,不能当作真实临床数据使用,也不要用它去训练任何用于诊断的模型后直接上线。
2.3 数据从哪来:MIT-BIH 是绕不开的起点
做 ECG 合成,公开数据里 MIT-BIH 心律失常数据库是事实标准,PhysioNet 上可以拿到,格式是 WFDB。Notebook 里一般会先用wfdb库读一条记录,取其中一段单导联信号,重采样到统一频率(常见 250Hz 或 360Hz),再做分段。分段长度通常取 2 到 5 秒,太短学不到完整节律,太长 LSTM 训练会慢且梯度容易出问题。下面这段是读数据和预处理的典型写法:
import wfdb import numpy as np from scipy.signal import resample # 读取 MIT-BIH 一条记录,channel 0 通常是 MLII 导联 record = wfdb.rdrecord('100', pn_dir='mitdb') signal = record.p_signal[:, 0] # 去基线漂移:高通滤波,截止 0.5Hz from scipy.signal import butter, filtfilt b, a = butter(2, 0.5 / (record.fs / 2), btype='high') signal = filtfilt(b, a, signal) # 重采样到 250Hz,统一不同记录的采样率 target_fs = 250 if record.fs != target_fs: num_samples = int(len(signal) * target_fs / record.fs) signal = resample(signal, num_samples) # 按 3 秒一段切分,段间不重叠 seg_len = 3 * target_fs segments = np.array([signal[i:i+seg_len] for i in range(0, len(signal)-seg_len, seg_len)]) # 每条样本独立标准化,逼判别器看形态 segments = (segments - segments.mean(axis=1, keepdims=True)) / \ (segments.std(axis=1, keepdims=True) + 1e-8) np.save('ecg_segments.npy', segments)这段代码里几个参数值得说清楚。butter(2, 0.5/(fs/2), btype='high')里的 0.5Hz 是基线漂移的常见截止点,阶数 2 够用,阶数太高会引入相位失真,所以后面用filtfilt做零相位滤波。重采样到 250Hz 是为了让不同记录能拼进同一个 batch,如果你的数据源采样率统一,这步可以省。seg_len取 3 秒,在 250Hz 下是 750 个点,这个长度能覆盖两到三个完整心拍,LSTM 展开 750 步在显存上还能接受。最后一步的逐样本标准化是关键,不做的话判别器会偷懒只看幅值。
3. 在 Jupyter Notebook 里把 LSTM-GAN 跑起来:环境、结构与训练循环
3.1 环境准备:Python 版本、依赖和 Jupyter 的坑
这份 Notebook 是 Python 写的,跑之前先把环境理顺。Python 3.8 到 3.10 都比较稳,太新的版本有时会和旧版 TensorFlow 或 PyTorch 打架。依赖主要是深度学习框架(PyTorch 或 TensorFlow 二选一,Notebook 里通常写死一种)、numpy、scipy、wfdb、matplotlib。安装命令:
# 建独立环境,别污染系统 Python python -m venv ecg_gan_env source ecg_gan_env/bin/activate # Windows 用 ecg_gan_env\Scripts\activate # 装依赖,PyTorch 按官网对应 CUDA 版本选命令 pip install numpy scipy wfdb matplotlib jupyter pip install torch torchvision # 或 pip install tensorflow # 启动 Notebook jupyter notebookJupyter Notebook 默认保存路径是启动时所在的目录,很多人第一次用找不到文件存哪了,就是因为没注意这点。建议先cd到项目目录再启动,或者启动后用%pwd确认当前路径。如果 Notebook 里 import 报"要安装缺失的节点"之类的错,八成是环境没选对——Jupyter 的 kernel 可能还指向系统 Python,用python -m ipykernel install --user --name=ecg_gan_env把当前环境注册成 kernel,再在 Notebook 里切换。
3.2 生成器和判别器的结构:逐层拆开看
生成器的输入是一段噪声向量,输出是定长 ECG 序列。常见结构是:噪声先过一个全连接层升维,再 reshape 成(seq_len, feature_dim),然后进 LSTM,最后接一个全连接把每个时间步映射成单个采样值。判别器反过来,输入一段序列,过 LSTM 或一维卷积,最后输出一个标量判断真假。下面是一个能直接用的 PyTorch 版本:
import torch import torch.nn as nn class ECGGenerator(nn.Module): def __init__(self, noise_dim=100, hidden_dim=128, seq_len=750): super().__init__() self.seq_len = seq_len self.hidden_dim = hidden_dim # 噪声升维到 LSTM 输入维度 self.fc_in = nn.Linear(noise_dim, hidden_dim) # 两层 LSTM,batch_first 让输入形状是 (batch, seq, feature) self.lstm = nn.LSTM(hidden_dim, hidden_dim, num_layers=2, batch_first=True) # 每个时间步输出一个采样值 self.fc_out = nn.Linear(hidden_dim, 1) def forward(self, z): # z: (batch, noise_dim) -> (batch, seq_len, hidden_dim) h = self.fc_in(z).unsqueeze(1).repeat(1, self.seq_len, 1) out, _ = self.lstm(h) # (batch, seq_len, 1) -> (batch, seq_len) return self.fc_out(out).squeeze(-1) class ECGDiscriminator(nn.Module): def __init__(self, seq_len=750, hidden_dim=128): super().__init__() self.lstm = nn.LSTM(1, hidden_dim, num_layers=2, batch_first=True) self.fc = nn.Linear(hidden_dim, 1) def forward(self, x): # x: (batch, seq_len) -> (batch, seq_len, 1) out, (h_n, _) = self.lstm(x.unsqueeze(-1)) # 取最后一层最后时间步的隐藏状态 return self.fc(h_n[-1])生成器里fc_in把 100 维噪声升到 128 维,然后repeat成 750 步的序列喂给 LSTM——这里其实是用同一个噪声向量作为每个时间步的输入,LSTM 靠自己的隐藏状态产生时间上的变化。num_layers=2是经验值,一层太浅学不到复杂节律,三层以上训练慢且容易过拟合小数据集。判别器取h_n[-1],也就是最后一层 LSTM 在最后一个时间步的隐藏状态,它浓缩了整段序列的信息,再过一个全连接输出真假分数。如果你的数据量小,判别器可以换成一维 CNN,参数更少、更不容易过拟合。
3.3 训练循环:损失、优化器和几个必调参数
对抗训练的核心是交替更新判别器和生成器。判别器要最大化"真样本判真、假样本判假"的能力,生成器要骗过判别器。用 BCE 损失就行,但要注意标签平滑——把真样本的标签从 1 改成 0.9,能防止判别器过于自信。下面是训练循环的骨架:
import torch.optim as optim device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') G = ECGGenerator().to(device) D = ECGDiscriminator().to(device) # 生成器学习率通常比判别器低一点,防止它更新太猛 opt_G = optim.Adam(G.parameters(), lr=1e-4, betas=(0.5, 0.999)) opt_D = optim.Adam(D.parameters(), lr=2e-4, betas=(0.5, 0.999)) criterion = nn.BCEWithLogitsLoss() data = torch.tensor(np.load('ecg_segments.npy'), dtype=torch.float32) batch_size, epochs = 64, 200 for epoch in range(epochs): for i in range(0, len(data), batch_size): real = data[i:i+batch_size].to(device) bs = real.size(0) # ---- 更新判别器 ---- z = torch.randn(bs, 100, device=device) fake = G(z).detach() # 标签平滑:真样本用 0.9 而不是 1.0 loss_D = criterion(D(real), torch.full((bs, 1), 0.9, device=device)) + \ criterion(D(fake), torch.zeros(bs, 1, device=device)) opt_D.zero_grad(); loss_D.backward(); opt_D.step() # ---- 更新生成器 ---- z = torch.randn(bs, 100, device=device) fake = G(z) # 生成器希望判别器把假样本判成真 loss_G = criterion(D(fake), torch.ones(bs, 1, device=device)) opt_G.zero_grad(); loss_G.backward(); opt_G.step() if epoch % 20 == 0: print(f'epoch {epoch} loss_D {loss_D.item():.4f} loss_G {loss_G.item():.4f}')几个参数是血泪经验换来的。betas=(0.5, 0.999)是 GAN 训练的标配,动量项调低能减少震荡。生成器学习率 1e-4、判别器 2e-4,让判别器稍微强一点但别碾压,如果 loss_D 很快掉到 0.1 以下,就把判别器学习率再降或者给它加 dropout。batch_size=64在 750 点序列上对显存比较友好,显存够可以加到 128。训练轮数 200 是起步值,实际要看 loss 曲线和生成样本的形态,通常 100 到 300 轮之间能看出效果。判别器更新时对fake用了.detach(),这是必须的,否则梯度会回传到生成器,把它的更新搞乱。
4. 避坑与排查:LSTM-GAN 合成 ECG 最常见的 5 个翻车现场
4.1 生成波形全是平直线或单一正弦
现象:训练几十轮后,生成器输出的波形几乎是一条直线,或者是一个规整的正弦波,完全没有 P-QRS-T 形态。原因通常是模式崩溃,生成器发现"输出均值波形"就能骗过判别器,于是躺平了。解决办法是降低判别器强度——把判别器学习率降到和生成器一样,或者给判别器加 dropout、减少 LSTM 层数。另一个办法是在生成器 loss 里加一点多样性约束,比如让不同噪声生成的样本之间保持距离。还可以检查数据标准化是不是做过头了,如果所有训练样本被标准化得过于相似,生成器学不到变化。
4.2 判别器 loss 迅速归零,生成器学不动
现象:训练开始没多久,loss_D 就掉到 0.01 以下,loss_G 反而越来越大。原因是判别器太强,真假样本被它一眼看穿,生成器拿到的梯度几乎没有信息量。解决思路是给判别器加噪声(输入上加高斯噪声)、降低判别器学习率、或者把判别器的更新频率降到生成器的一半(每更新两次生成器才更新一次判别器)。还有一种情况是数据泄漏——训练集和验证集没分好,判别器见过生成器要学的样本,这种要从数据划分上查。
4.3 生成的波形幅值爆炸或全为零
现象:生成样本的数值范围远超训练数据,或者全部塌缩到零附近。前者通常是生成器最后一层没做约束,LSTM 输出经过全连接后数值无界。可以在生成器输出后加一个tanh再乘一个缩放系数,把输出限制在合理范围。后者往往是标准化的问题——如果训练时做了逐样本标准化,生成器学到的输出也是标准化后的,反标准化时如果标准差估计不对,就会塌缩。检查一下保存数据时的均值和方差有没有一起存下来。
4.4 Jupyter 里训练到一半 kernel 崩了
现象:训练到一半 Notebook 报 kernel died,或者显存溢出。最常见的原因是数据全部加载进内存后没释放,加上 LSTM 展开 750 步的中间状态很吃显存。解决办法是把数据做成DataLoader分批加载,别一次性torch.tensor整个数据集;训练循环里用with torch.no_grad()包住不需要梯度的部分;每轮结束调torch.cuda.empty_cache()。如果还是崩,把seq_len从 750 降到 500,或者把 LSTM 隐藏维度从 128 降到 64。
4.5 生成的 ECG 形态像但节律不对
现象:单看一个心拍,P 波、QRS、T 波都在,但连起来看 RR 间期忽长忽短,或者出现明显不合理的节律。这是因为 LSTM 学到了局部形态,但没学到全局节律。可以在判别器里加入对 RR 间期的约束,或者把生成器的输入从纯噪声改成"噪声 + 目标心率"的条件向量,让生成器知道该生成多快的心律。另一个办法是训练数据里如果心律失常样本太多,正常节律的样本会被淹没,需要做类别平衡。
5. 进阶技巧:怎么判断生成的 ECG 到底"像不像"
跑通训练只是第一步,真正难的是评估生成质量。肉眼看波形只能筛掉明显崩坏的,要量化"似是而非"的程度,得靠几个指标组合。下面这张表是我常用的评估维度:
| 评估维度 | 具体指标 | 判断标准 |
|---|---|---|
| 形态相似度 | 与真实样本的 DTW 距离 | 越小越像,但别追求过小,过小说明过拟合 |
| 分布距离 | MMD 或 Wasserstein 距离 | 衡量生成分布和真实分布的整体差距 |
| 节律合理性 | RR 间期标准差 | 应落在真实数据的合理区间内 |
| 频域特征 | 功率谱密度对比 | 主频和低频成分应与真实 ECG 接近 |
| 下游可用性 | 用合成数据训练分类器,在真实测试集上评估 | 掉点不超过 5% 说明合成数据有增强价值 |
实操上,我一般会先画一张对比图:上面一行真实 ECG,下面一行生成 ECG,各取 5 条,肉眼过一遍。然后算 MMD,用 RBF 核,带宽取中位数启发式。最后做一个下游任务验证——拿合成数据扩充训练集,训练一个简单的心拍分类器,看它在真实测试集上的 F1 有没有提升。如果合成数据让下游指标反而下降,说明生成质量还不够,得回去调模型。
还有一个容易被忽略的技巧:把生成器的噪声维度调大。很多人用 100 维噪声,但在 ECG 这种形态相对固定的信号上,噪声维度太大反而让生成器难以收敛到合理形态。我试过把噪声降到 32 维,生成波形的稳定性明显提升,代价是多样性略降。这个权衡要看你的用途——如果是为了数据增强,多样性重要,噪声维度别太小;如果是为了教学演示,稳定性优先,可以适当降维。
最后说个习惯:每次改完模型结构或超参,别急着跑 200 轮,先跑 20 轮看 loss 曲线和生成样本。LSTM-GAN 训练的不确定性很大,同样的代码换个随机种子结果可能差很多,所以我会固定种子、记录每次实验的配置,避免"上次明明能跑"这种玄学。合成 ECG 这件事,做到"似是而非"不难,难的是知道它哪里像、哪里不像、以及为什么。希望帮到你。
本文还有配套的精品资源,点击获取