
简介这份资源围绕LSTM-GAN生成逼真ECG信号展开面向具备Python与深度学习基础、关注医学信号处理与数据增强的研究者和开发者。项目以长短期记忆网络捕捉心电信号的周期性与波形模式再由生成器与判别器相互博弈产出足以以假乱真的合成心电数据可用于异常检测算法测试或扩充训练集。压缩包共13个文件约4.46MB包含5个py脚本、3张png结果图、2个h5权重文件以及1个ipynb交互式笔记、1个md说明文档和1个gitignore配置覆盖模型定义、训练、测试与信号清理等环节。已有313人学习下载。读者可借此理解LSTM-GAN在生物医学序列建模中的完整实现路径参考生成器与判别器的权重保存方式、噪声生成与ECG扩展脚本并对照可视化图像评估生成信号质量适合作为课程设计或科研入门的实践素材。1. 从一份 ECG 合成 Notebook 说起为什么“似是而非”比“以假乱真”更难拿到“用于生成似是而非的ECG信号的LSTM-GAN”这个题目时我第一反应不是模型结构而是“似是而非”这四个字。ECG 信号在临床上有一套硬约束P 波、QRS 复合波、T 波的时间关系RR 间期不能是负数幅值不能超出导联量程。一个 GAN 如果只追求判别器分不出真假很容易生成一段“看起来像波形、但医生一眼觉得不对劲”的东西——比如 QRS 波群宽到 300ms或者 T 波倒置出现在本该直立的位置。所以这个方向真正要解决的不是“生成得像”而是“生成得合理但又不完全重复”。这份 Jupyter Notebook 加 Python 的工程适合两类人一类是做生理信号数据增强的算法工程师手里只有几百条标注 ECG想扩样本又怕引入伪影另一类是刚学完 LSTM 和 GAN、想找一个比 MNIST 更有物理意义的练手项目的人。它不要求你懂心电诊断但要求你能把时序生成的基本功——序列对齐、梯度惩罚、模式崩溃排查——走一遍。下面我按自己复现这类项目的顺序把 LSTM-GAN 生成 ECG 的选型、代码骨架、参数和踩坑讲清楚。2. LSTM-GAN 生成 ECG 的选型逻辑与最小可跑骨架2.1 为什么是 LSTM 做生成器而不是纯全连接ECG 是典型的一维时序信号采样率常见 250Hz 或 360Hz一段 10 秒的片段就是 2500 到 3600 个点。如果用全连接网络直接输出这么长的向量参数量会爆炸而且模型学不到相邻采样点之间的局部相关性——表现出来就是生成的波形毛刺极多QRS 波群被拆成随机尖峰。LSTM 的循环结构天然适合处理这种依赖关系它的门控机制能在长序列里保留“上一个心跳的节律信息”让生成的下一个心跳和上一个在 RR 间期上保持连贯。但纯 LSTM 也有问题它倾向于生成过于平滑的均值波形因为 MSE 损失会惩罚任何偏离均值的输出。这就是为什么需要 GAN 的对抗损失来“逼”出高频细节。LSTM 做生成器、CNN 或 LSTM 做判别器是这类任务里比较稳的组合。判别器用一维卷积更常见因为卷积对局部形态QRS 的陡峭上升沿敏感而 LSTM 判别器训练慢、容易梯度消失。2.2 生成器和判别器的代码骨架下面这段是我一般会先跑通的最小结构不追求最优但能让你在 Jupyter Notebook 里快速看到 loss 有没有动。输入噪声维度设 100生成长度 1000 个采样点对应 4 秒左右的信号。import torch import torch.nn as nn class ECGGenerator(nn.Module): def __init__(self, noise_dim100, hidden_dim128, seq_len1000): super().__init__() self.seq_len seq_len self.hidden_dim hidden_dim # 把噪声映射成 LSTM 的初始状态 self.fc nn.Linear(noise_dim, hidden_dim * 2) self.lstm nn.LSTM( input_size1, hidden_sizehidden_dim, num_layers2, batch_firstTrue, dropout0.2 ) self.out nn.Linear(hidden_dim, 1) self.tanh nn.Tanh() # ECG 幅值归一化到 [-1, 1] def forward(self, z): # z: (batch, noise_dim) h0, c0 self.fc(z).chunk(2, dim1) h0 h0.unsqueeze(0).repeat(2, 1, 1) # num_layers2 c0 c0.unsqueeze(0).repeat(2, 1, 1) # 用零输入驱动 LSTM逐步生成 inp torch.zeros(z.size(0), self.seq_len, 1, devicez.device) out, _ self.lstm(inp, (h0, c0)) return self.tanh(self.out(out)) # (batch, seq_len, 1)这段代码的关键在self.fc(z).chunk(2, dim1)它把噪声向量拆成 LSTM 的隐状态 h 和细胞状态 c。这样噪声不是作为每一步的输入而是作为“初始记忆”让整个序列的生成受同一个噪声控制避免每一步输入不同噪声导致的节律断裂。num_layers2对应repeat(2, 1, 1)层数改了这里也要改这是新手最容易翻车的地方——报维度不匹配的错查半天发现是 repeat 次数写死成 2 了。判别器用一维卷积class ECGDiscriminator(nn.Module): def __init__(self, seq_len1000): super().__init__() self.net nn.Sequential( nn.Conv1d(1, 32, kernel_size5, stride2, padding2), nn.LeakyReLU(0.2), nn.Conv1d(32, 64, kernel_size5, stride2, padding2), nn.LeakyReLU(0.2), nn.Conv1d(64, 128, kernel_size5, stride2, padding2), nn.LeakyReLU(0.2), nn.AdaptiveAvgPool1d(1), nn.Flatten(), nn.Linear(128, 1) ) def forward(self, x): # x: (batch, seq_len, 1) - (batch, 1, seq_len) x x.permute(0, 2, 1) return self.net(x)判别器最后输出一个 logit不接 sigmoid因为训练时用BCEWithLogitsLoss更数值稳定。AdaptiveAvgPool1d(1)把任意长度的序列压成 1 个值这样判别器对输入长度不敏感方便你后续改 seq_len 做实验。2.3 训练循环里必须加的三个约束GAN 训练 ECG 最容易出现两种失败判别器太强导致生成器梯度消失或者生成器找到“捷径”只生成一种心跳形态。我在训练循环里固定加三样东西。第一梯度惩罚。WGAN-GP 的梯度惩罚比原始 GAN 的 BCE 稳得多尤其在小样本 ECG 上。第二判别器每步更新一次生成器每两步更新一次让判别器别跑太快。第三每轮记录生成样本的 RR 间期均值和标准差如果标准差趋近于 0说明模式崩溃已经开始了。def compute_gradient_penalty(D, real, fake, device): alpha torch.rand(real.size(0), 1, 1, devicedevice) interpolated (alpha * real (1 - alpha) * fake).requires_grad_(True) d_inter D(interpolated) grad torch.autograd.grad( outputsd_inter, inputsinterpolated, grad_outputstorch.ones_like(d_inter), create_graphTrue, retain_graphTrue )[0] grad grad.view(grad.size(0), -1) penalty ((grad.norm(2, dim1) - 1) ** 2).mean() return penaltyalpha的形状是(batch, 1, 1)因为信号是(batch, seq_len, 1)插值要在样本维和序列维同时广播。create_graphTrue不能省否则惩罚项没法反向传播。这个函数每步都调用计算开销不小如果显存吃紧可以把 batch 降到 16 或 32。3. 在 Jupyter Notebook 里把 ECG 数据喂进 LSTM-GAN3.1 数据预处理归一化和切窗的先后顺序ECG 原始数据常见两种格式WFDB 的.dat加.hea或者 CSV 里一列时间一列幅值。不管哪种第一步都是转成 numpy 数组然后做 z-score 归一化。注意归一化要按整条记录算均值和标准差不能按窗口算——按窗口算会把每个窗口的基线拉到 0破坏窗口之间的幅值关系生成器学到的就是“每个窗口都从零开始”的假模式。切窗用滑动窗口窗口长度 1000 点步长 500 点这样相邻窗口有 50% 重叠增加样本量。但重叠窗口在训练时要小心如果验证集也用重叠窗口评估指标会虚高因为相邻窗口高度相似。我一般训练集用重叠验证集用不重叠的独立片段。import numpy as np from scipy.signal import butter, filtfilt def bandpass_filter(signal, fs250, low0.5, high45): nyq 0.5 * fs b, a butter(4, [low / nyq, high / nyq], btypeband) return filtfilt(b, a, signal) def normalize(signal): return (signal - signal.mean()) / (signal.std() 1e-8) def make_windows(signal, window1000, step500): windows [] for start in range(0, len(signal) - window 1, step): windows.append(signal[start:start window]) return np.array(windows)butter(4, ...)的 4 是滤波器阶数阶数越高过渡带越陡但相位失真也越大。filtfilt做零相位滤波避免 QRS 波群位置偏移。1e-8是防止标准差为 0 的兜底实际 ECG 不会出现但代码里加上不亏。3.2 在 Notebook 里管理数据集和 DataLoaderJupyter Notebook 的交互性适合调参但不适合把数据加载逻辑散在各处。我习惯在第一个 cell 里定义 Dataset 类后面所有实验复用。这样改窗口长度或 batch size 时只动一个地方。from torch.utils.data import Dataset, DataLoader class ECGDataset(Dataset): def __init__(self, windows): self.data torch.FloatTensor(windows).unsqueeze(-1) # (N, 1000, 1) def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] dataset ECGDataset(train_windows) loader DataLoader(dataset, batch_size32, shuffleTrue, drop_lastTrue)drop_lastTrue在 GAN 训练里很重要如果最后一个 batch 只有 1 个样本梯度惩罚的alpha广播会出问题而且 BatchNorm 在 batch size 为 1 时直接报错。虽然上面的判别器没用 BatchNorm但养成这个习惯能省很多调试时间。3.3 训练循环的 Notebook 写法与实时监控在 Notebook 里训练我一般把训练循环写成一个函数每个 epoch 返回 loss 和几个监控指标然后用 matplotlib 画出来。不要用tqdm在 Notebook 里刷屏输出会乱。用clear_output加display做原地刷新更干净。from IPython.display import clear_output import matplotlib.pyplot as plt def train_epoch(G, D, loader, opt_g, opt_d, device, lambda_gp10): for i, real in enumerate(loader): real real.to(device) batch_size real.size(0) # 训练判别器 z torch.randn(batch_size, 100, devicedevice) fake G(z).detach() d_real D(real) d_fake D(fake) gp compute_gradient_penalty(D, real, fake, device) loss_d d_fake.mean() - d_real.mean() lambda_gp * gp opt_d.zero_grad() loss_d.backward() opt_d.step() # 每两步训练一次生成器 if i % 2 0: z torch.randn(batch_size, 100, devicedevice) fake G(z) loss_g -D(fake).mean() opt_g.zero_grad() loss_g.backward() opt_g.step() return loss_d.item(), loss_g.item()lambda_gp10是 WGAN-GP 原论文的推荐值我在 ECG 上试过 5 和 205 的时候判别器约束不够20 的时候生成器更新变慢10 是比较稳的中间值。loss_d的符号是d_fake.mean() - d_real.mean()这是 WGAN 的 Wasserstein 距离估计越小说明判别器越分不清真假但不要追求它降到 0那意味着判别器完全失效。4. 生成“似是而非”ECG 的避坑与排查清单4.1 生成波形全是直线或极小幅值震荡现象训练几十轮后生成器输出的信号幅值接近 0画出来是一条平线或者只有微小抖动。原因判别器太强生成器梯度消失。WGAN-GP 虽然比原始 GAN 稳但如果判别器学习率是生成器的 5 倍以上或者梯度惩罚系数太小判别器仍然会赢。解决把判别器学习率降到生成器的 1/2 到 1/4比如生成器 1e-4、判别器 2e-5。同时检查梯度惩罚项是否真的在反向传播——create_graphTrue漏掉的话惩罚项对判别器参数没有梯度等于没加。4.2 生成的心跳节律完全随机RR 间期忽长忽短现象生成的信号有 QRS 形态但两个 QRS 之间的距离从 200 点到 2000 点都有不像正常窦性节律。原因噪声只作为 LSTM 初始状态但 LSTM 在长序列生成时“忘记”了初始状态后面的心跳变成了由零输入和隐状态自行演化失去了全局节律控制。解决在生成器里加一个周期性的条件输入。常见做法是把噪声同时映射成一个“节律向量”在每个时间步拼接到 LSTM 输入上。或者更简单把 seq_len 缩短到 500 点约 2 秒只生成 2 到 3 个心跳这样 LSTM 的初始状态还能影响整段序列。4.3 判别器 loss 剧烈震荡生成样本质量时好时坏现象loss_d 在正负之间大幅跳变每隔几个 batch 生成的波形就变一个样。原因batch size 太小梯度估计方差大。ECG 窗口之间本身差异就大如果 batch 里恰好全是相似形态判别器会过拟合这个 batch。解决把 batch size 提到 64 或 128同时用shuffleTrue确保每个 batch 的形态多样。如果显存不够用梯度累积每 4 个 batch 才更新一次参数等效于大 batch。4.4 Notebook 重启后数据要重新处理浪费时间现象每次关掉 Jupyter Notebook 再打开都要重新跑滤波、切窗、归一化几分钟就没了。原因没有把预处理结果落盘。解决在第一个 cell 里加缓存逻辑处理完存成.npy下次直接np.load。注意存的时候把窗口长度和步长写进文件名比如ecg_win1000_step500.npy避免不同参数的结果混在一起。import os cache_file fecg_win{window}_step{step}.npy if os.path.exists(cache_file): windows np.load(cache_file) else: windows make_windows(normalize(bandpass_filter(raw_signal))) np.save(cache_file, windows)4.5 生成的信号在 QRS 波群处出现高频振铃现象QRS 的陡峭上升沿后面跟着一串衰减震荡像滤波器振铃。原因生成器的 tanh 输出加上 LSTM 的连续状态在快速变化处容易产生过冲。另外如果训练数据本身经过了截止频率很低的低通滤波生成器学到的 QRS 就是带振铃的。解决检查训练数据的滤波截止频率不要低于 40Hz否则 QRS 形态本身就不对。生成器最后一层可以改成nn.Hardtanh(-1, 1)硬限幅比 tanh 的软饱和更不容易过冲。如果振铃已经出现在生成后加一个 40Hz 的低通滤波但这是补救最好从数据源头解决。5. 用 RR 间期分布和形态模板做生成质量的量化验证训练完一个 LSTM-GAN光靠肉眼看波形图不够。我一般用两个指标做量化验证一个查节律一个查形态。节律用 RR 间期分布。对生成信号做 R 波检测可以用scipy.signal.find_peaks设置高度阈值为 0.5 倍最大幅值距离至少 200 个采样点然后算相邻 R 波位置的差值。真实 ECG 的 RR 间期标准差一般在 20 到 60ms 之间如果生成信号的 RR 间期标准差小于 10ms说明节律太死板大于 100ms说明节律失控。from scipy.signal import find_peaks def rr_intervals(signal, fs250): peaks, _ find_peaks(signal, height0.5 * signal.max(), distanceint(0.2 * fs)) rr np.diff(peaks) / fs * 1000 # 转成毫秒 return rr gen_signal G(torch.randn(1, 100, devicedevice)).detach().cpu().numpy().flatten() rr rr_intervals(gen_signal) print(fRR mean: {rr.mean():.1f} ms, RR std: {rr.std():.1f} ms)形态用模板匹配。从真实数据里取一条干净的窦性心跳作为模板对生成信号的每个心跳窗口算相关系数。相关系数在 0.7 到 0.9 之间是比较理想的“似是而非”——太像了说明生成器只是记住了训练样本太不像了说明形态不对。指标真实 ECG 参考范围生成质量判断RR 间期均值600-1000 ms超出范围说明节律异常RR 间期标准差20-60 ms10 太死板100 失控模板相关系数0.7-0.90.5 形态不对0.95 过拟合QRS 宽度80-120 ms超出说明波群形态失真最后说一个我自己的习惯每次改完生成器结构或损失函数先跑 200 个 batch用上面两个指标快速筛一遍不要等训练完 100 个 epoch 再看。ECG 生成这个方向调参的反馈周期越短你越容易找到那个“似是而非”的平衡点。希望帮到你。本文还有配套的精品资源点击获取