
Python实战 WGAN-LSTM时间序列预测从Wasserstein生成增强到LSTM多步预测与真实增益验证一句话看懂这篇文章WGAN并不是“造的数据越多越准”。在本文严格测试中低数据量LSTM的RMSE为 0.6117加入25%合成窗口后降到 0.5220改善约 14.7%但把合成比例继续提高到100%或200%误差反而显著恶化。Python 时间序列预测 WGAN Wasserstein GAN LSTM WGAN-GP 数据增强 生成对抗网络 PyTorch 深度学习生成对抗网络常被用于扩充时间序列训练样本但“生成得像”与“能提升预测”并不是同一件事。本文从可复现工程角度构建 WGAN-LSTM 时间序列预测流程先按时间严格切分训练、验证与测试区间再使用 WGAN-GP 仅学习训练窗口分布结合边际 Wasserstein 距离、自相关结构和真实/合成可区分性审计合成质量最后将不同数量的合成窗口加入 LSTM 多步预测器并与持续值、Ridge、原始 LSTM、简单噪声增强进行对照。真实实验表明适量 WGAN 增强可以缓解低数据量 LSTM 的误差但过多合成样本会将生成器偏差放大到下游预测中因此可信的 WGAN-LSTM 不应只展示 GAN 损失或漂亮曲线而必须在完全未参与生成与调参的真实测试集上验证最终收益。先看结果再学原理本实验采用 1958–2001 年 Mauna Loa 周度 CO₂ 序列52 周历史直接预测未来 4 周。低数据量LSTMRMSE0.6117WGAN-LSTM25%增强0.5220简单Jitter增强0.4904全量LSTM0.4786。这组结果说明 WGAN 有价值但不是“必胜组件”。图1WGAN-LSTM 的正确实验链路生成质量审计与预测增益验证缺一不可1. 为什么 WGAN-LSTM 值得研究也最容易被误用LSTM 擅长从历史序列中提取长期依赖但当真实训练样本数量有限、覆盖的波动形态不足时模型容易过拟合已有窗口。WGAN 的直觉是不直接修改标签而是学习训练窗口的整体分布从中采样更多“可能出现的历史—未来组合”让预测器看到更丰富的局部轨迹。问题在于GAN 生成的序列可以在视觉上非常自然却仍然缺少真实数据的联合分布、极值、相位关系或条件动态。如果把这些有偏合成样本大量灌入预测模型增强反而可能变成一种“分布污染”。因此本文把 WGAN-LSTM 拆成两个独立问题生成器是否可信可信到什么程度这些样本加入后真实测试集是否真正改善只比较 GAN 的生成损失不足以证明“生成质量”。只画几条真实/合成曲线不足以证明联合分布接近。只在随机切分样本上比较预测误差可能把时间泄漏当成模型提升。增强比例必须用验证集选测试集不能反过来决定“加多少合成数据”。2. 从 GAN 到 WGAN为什么要换一种“距离”经典 GAN 把生成器 G 与判别器 D 放在一个极小极大博弈中。当真实分布与生成分布的支撑集几乎不重合时Jensen–Shannon 类目标容易给生成器提供不稳定或饱和的梯度。Wasserstein GAN 改用 Earth Mover / Wasserstein-1 距离的对偶形式让 critic 输出任意实数评分并要求其满足 1-Lipschitz 约束。其核心目标可以写成WGAN 目标max_D E[D(x_real)] − E[D(G(z))]同时约束 D 为 1-Lipschitz 函数。生成器最小化 −E[D(G(z))]。这里的 critic 不是输出“真/假概率”而是在学习区分两个分布所需的连续评分函数。图2本文 WGAN 结构生成器输出完整时间窗critic 直接对真实/合成窗口打分2.1 为什么工程实现采用 WGAN-GP原始 WGAN 通过权重裁剪近似 Lipschitz 约束但裁剪范围难选过小会降低 critic 表达能力过大又可能破坏约束。WGAN-GP 在真实与合成样本之间随机插值并惩罚 critic 对这些插值样本的梯度范数偏离 1。实践中它通常比简单权重裁剪更稳定因此本文代码采用“WGAN核心目标 Gradient Penalty”的实现。图3梯度惩罚的直观含义限制 critic 在真实与生成样本之间出现异常陡峭的局部斜率def gradient_penalty(critic, real, fake):eps torch.rand(real.size(0), 1, devicereal.device)x_hat eps * real (1 - eps) * fakex_hat.requires_grad_(True)score critic(x_hat)grad torch.autograd.grad(outputsscore, inputsx_hat,grad_outputstorch.ones_like(score),create_graphTrue, retain_graphTrue)[0]return ((grad.norm(2, dim1) - 1.0) ** 2).mean()3. LSTM真正负责预测的序列模型WGAN 在本文中的角色是“生成训练窗口”最终的预测仍由 LSTM 完成。LSTM 通过遗忘门、输入门和输出门控制信息在单元状态中的保留与写入能够比普通 RNN 更稳定地学习长距离依赖。图4LSTM门控机制与本文的 52→4 直接多步预测任务本文不采用“预测一步—把预测值喂回去—继续预测”的递归策略而是让 LSTM 一次性输出未来 4 周。直接多步预测避免了测试阶段滚动输入自身误差造成的额外累积同时也使 WGAN 生成的每个 56 点窗口可以自然拆成 52 点输入与 4 点标签。class LSTMPredictor(nn.Module):def __init__(self, hidden24, horizon4):super().__init__()self.lstm nn.LSTM(1, hidden, batch_firstTrue)self.head nn.Linear(hidden, horizon)def forward(self, x):_, (h, _) self.lstm(x.unsqueeze(-1))return self.head(h[-1])4. 数据集与无泄漏实验协议实验使用 statsmodels 提供的 Mauna Loa 周度 CO₂ 公共序列覆盖 1958–2001 年。原始序列存在长期上升趋势和明显季节波动。为了让 WGAN 更聚焦于局部变化而不是直接记忆绝对浓度水平建模前先做一阶差分预测得到未来差分后再从最后一个已知 CO₂ 水平开始累加还原为 ppm。图5数据按时间顺序划分任何缩放、WGAN训练和模型选择都不能使用测试区间项目设置原因历史窗口52周覆盖约一年的季节周期预测长度4周直接多步输出避免递归误差训练/验证/测试65% / 15% / 20%严格按时间顺序WGAN训练集仅训练区间前40%的窗口模拟“样本有限”场景缩放MinMax[-1,1]参数只由训练段拟合增强比例候选0、25%、50%、100%、200%只用验证集选择最终测试测试集只评一次防止测试集反向调参为什么先差分CO₂绝对值存在明显趋势。如果直接让生成器学习绝对水平生成窗口容易把“时间位置”与“序列形态”混在一起。差分后WGAN主要学习局部增减、季节起伏和短期相关性最后再把预测差分累加回真实量纲。5. 先判断 WGAN 有没有学到“可用的序列”生成器训练完成后不能立即把所有合成窗口加入 LSTM。本文先从三个层次做质量审计边际分布是否接近、短期自相关是否接近、完整窗口是否仍然容易被一个简单分类器识别。第三项非常关键因为前两项都只是低阶统计。图6真实与合成窗口示例视觉相似不能替代定量审计图7真实/合成差分值的边际分布对比图8真实/合成窗口平均自相关函数对比图9质量审计结果低阶统计接近但完整窗口仍然容易被区分本次实验中缩放空间的一维 Wasserstein 距离为 0.0318Lag 1–12 的平均 ACF RMSE 为 0.0365两项都显示合成序列在“单点分布”和“短期相关性”上已经接近训练窗口。可是把完整 56 点窗口交给逻辑回归区分真实/合成时ROC-AUC 仍达到 0.966。如果生成分布真的高度逼真这个 AUC 应更靠近 0.5因此 0.966 明确说明生成器仍留下了可识别的结构性差异。这是全文最重要的生成质量结论“边际分布像 ACF像”仍然可能是假象。对于时间序列生成联合时序结构、极值、相位关系和高阶依赖都可能让合成数据被轻易识别。因此必须把下游任务表现作为最后一道验证。图10WGAN-GP训练曲线GAN的loss并不是监督学习意义上的“准确率”6. 核心实验WGAN增强到底有没有提升 LSTM为了让结论尽量接近 WGAN 的真实使用场景本文把训练区间的前 40% 窗口视为“低数据量训练集”。WGAN 只在这部分窗口上训练。增强比例通过验证集选择最终测试集完全不参与生成器训练和比例选择。模型RMSEMAER²相对低数据LSTMPersistence0.93100.75790.9635—Ridge-low0.52670.41290.9883—低数据LSTM0.61170.45870.98420.0%WGAN-LSTM 25%0.52200.39600.988514.7%Jitter-LSTM0.49040.37790.989919.8%Ridge-full0.48650.37960.9900全量训练LSTM-full0.47860.36500.9904全量训练图11测试集RMSE总览25% WGAN增强有效但并非所有方案中最优在低数据量设置下原始 LSTM 的测试 RMSE 为 0.6117。加入 25% WGAN 合成窗口后RMSE 降到 0.5220相对改善约 14.7%。这说明 WGAN 确实能向训练集补充一部分有用变化模式。但结果同样给出两个必要的限制条件。第一简单高斯 Jitter 增强的 RMSE 为 0.4904比本次 WGAN 增强更低第二使用全部真实训练窗口的 LSTM RMSE 为 0.4786仍明显优于低数据量下的 WGAN-LSTM。换言之WGAN 可以缓解缺数据但不能把“合成数据”当成真正新增的真实观测。7. 最容易被忽略的超参数合成数据比例图12验证集与测试集对增强比例的响应过量合成会迅速放大生成器偏差验证集在 25% 合成比例处取得最小 RMSE因此按预先规定的规则选择 25% 作为最终增强比例。继续把比例提高到 50%收益开始回退达到 100% 和 200% 时测试 RMSE 分别约为 0.928 和 0.935几乎退化到持续值基线。原因并不神秘当合成数据数量与真实样本相当甚至超过真实样本时预测器看到的主要“经验分布”已经不再是真实世界而是生成器近似出的分布。只要生成器存在一点系统偏差这个偏差就会被大量复制最终压过真实数据。经验法则不要默认“1:1增强”是合理设置。更稳妥的方法是把合成比例本身当作超参数在验证集上从小到大扫描并设置真实样本最低占比。8. 多步预测越远的步长越能暴露增强质量图131–4周预测误差远期误差持续增加说明长期动态比短期拟合更难一步预测很容易被局部平滑性“撑高成绩”因此本文同时看未来第 1、2、3、4 周的 RMSE。持续值模型在第 1 周尚能依靠序列惯性但第 4 周误差显著放大WGAN-LSTM 能部分降低这个增长速度却仍不如使用完整真实训练集的 LSTM。图14测试末段的一步预测WGAN-LSTM提升低数据模型但全量真实训练仍更稳9. 可信 WGAN-LSTM 的防泄漏清单图15从时间切分到最终测试每一步都可能把“未来”偷偷带进训练先按时间切原始序列再拟合 scaler不要用全序列均值、方差或最值。WGAN 只能训练在训练窗口上验证集和测试集不能被拿来“提升生成质量”。合成样本比例、GAN步数、LSTM结构等都只能根据训练/验证集确定。测试集只用于最终汇报如果看了测试结果再回去改GAN就已经产生测试集过拟合。若不同窗口来自同一长序列必须明确预测场景历史重叠本身不等于泄漏但预测标签绝不能越过时间边界。本文的测试窗口允许使用预测时点之前已经观测到的历史值这符合真实滚动预测但任何未来 4 周目标都不会进入训练集、GAN 或缩放器。10. 完整核心实现WGAN-GP LSTM下面给出可直接迁移的核心代码骨架。为了让文章重点保持清晰数据可视化与日志保存代码略去但生成器、critic、梯度惩罚、LSTM和增强训练逻辑完整保留。import numpy as npimport torchimport torch.nn as nnfrom torch.utils.data import DataLoader, TensorDatasetHISTORY, HORIZON, SEQ_LEN 52, 4, 56Z_DIM 64class Generator(nn.Module):def __init__(self):super().__init__()self.net nn.Sequential(nn.Linear(Z_DIM, 128), nn.LayerNorm(128), nn.LeakyReLU(0.2),nn.Linear(128, 256), nn.LayerNorm(256), nn.LeakyReLU(0.2),nn.Linear(256, SEQ_LEN), nn.Tanh())def forward(self, z):return self.net(z)class Critic(nn.Module):def __init__(self):super().__init__()self.net nn.Sequential(nn.Linear(SEQ_LEN, 256), nn.LeakyReLU(0.2),nn.Linear(256, 128), nn.LeakyReLU(0.2),nn.Linear(128, 1))def forward(self, x):return self.net(x).squeeze(-1)class LSTMPredictor(nn.Module):def __init__(self, hidden24):super().__init__()self.lstm nn.LSTM(1, hidden, batch_firstTrue)self.head nn.Linear(hidden, HORIZON)def forward(self, x):_, (h, _) self.lstm(x.unsqueeze(-1))return self.head(h[-1])def gradient_penalty(critic, real, fake):eps torch.rand(real.size(0), 1, devicereal.device)x_hat eps * real (1 - eps) * fakex_hat.requires_grad_(True)score critic(x_hat)grad torch.autograd.grad(score, x_hat, torch.ones_like(score),create_graphTrue, retain_graphTrue)[0]return ((grad.norm(2, dim1) - 1) ** 2).mean()def train_wgan(real_windows, steps2200, n_critic3, lambda_gp10.0):G, C Generator(), Critic()opt_g torch.optim.Adam(G.parameters(), 1e-4, betas(0.0, 0.9))opt_c torch.optim.Adam(C.parameters(), 1e-4, betas(0.0, 0.9))loader DataLoader(TensorDataset(torch.tensor(real_windows)),batch_size64, shuffleTrue, drop_lastTrue)it iter(loader)for step in range(steps):for _ in range(n_critic):try:real next(it)[0]except StopIteration:it iter(loader); real next(it)[0]z torch.randn(real.size(0), Z_DIM)fake G(z).detach()loss_c C(fake).mean() - C(real).mean()loss_c loss_c lambda_gp * gradient_penalty(C, real, fake)opt_c.zero_grad(); loss_c.backward(); opt_c.step()z torch.randn(real.size(0), Z_DIM)loss_g -C(G(z)).mean()opt_g.zero_grad(); loss_g.backward(); opt_g.step()return G, C# 生成增强窗口整段窗口同时包含 52 步历史与 4 步未来with torch.no_grad():synthetic G(torch.randn(n_synthetic, Z_DIM)).numpy()X_syn synthetic[:, :HISTORY]y_syn synthetic[:, HISTORY:]X_aug np.vstack([X_real, X_syn])y_aug np.vstack([y_real, y_syn])# 训练预测器验证集决定增强比例测试集不要参与选择model LSTMPredictor()optimizer torch.optim.Adam(model.parameters(), lr2e-3)loss_fn nn.HuberLoss(delta0.7)for epoch in range(max_epochs):model.train()pred model(torch.tensor(X_aug, dtypetorch.float32))loss loss_fn(pred, torch.tensor(y_aug, dtypetorch.float32))optimizer.zero_grad(); loss.backward(); optimizer.step()11. 不要跳过合成数据质量的三层检查from scipy.stats import wasserstein_distancefrom statsmodels.tsa.stattools import acffrom sklearn.linear_model import LogisticRegressionfrom sklearn.metrics import roc_auc_score# 1) 边际分布wd wasserstein_distance(real_windows.ravel(), synthetic.ravel())# 2) 时间依赖平均 ACFacf_real np.mean([acf(x, nlags12, fftFalse) for x in real_windows], axis0)acf_syn np.mean([acf(x, nlags12, fftFalse) for x in synthetic], axis0)acf_rmse np.sqrt(np.mean((acf_real[1:] - acf_syn[1:]) ** 2))# 3) 完整窗口可区分性AUC 越接近 0.5 越难区分X np.vstack([real_windows, synthetic[:len(real_windows)]])y np.r_[np.ones(len(real_windows)), np.zeros(len(real_windows))]clf LogisticRegression(max_iter2000).fit(X_train, y_train)auc roc_auc_score(y_test, clf.predict_proba(X_test)[:, 1])如果边际 Wasserstein 距离和 ACF 都很好但真实/合成分类 AUC 仍接近 1说明生成器只复现了部分低阶特征。此时更应该降低增强比例而不是继续无限生成。12. 换成自己的 CSV最少要改哪些地方实际项目通常不是单变量 CO₂。将本文迁移到电力负荷、设备传感器、金融指标或环境数据时建议先保持流程不变只替换数据读取、窗口定义和输出维度。场景输入窗口WGAN生成对象LSTM输出注意点单变量预测[B,T,1]TH整段序列H个未来值先处理趋势/季节性多变量单目标[B,T,F]可生成全部F维或仅目标相关特征H×1确保合成变量间相关性多变量多目标[B,T,F]TH的多变量块H×M生成器和critic需改成多通道稀有事件预测事件前历史窗含极端模式的条件窗口风险/数值普通WGAN可能抹平尾部建议条件WGANdf pd.read_csv(your_data.csv, parse_dates[time])df df.sort_values(time).set_index(time)feature_cols [load, temperature, humidity]target_col load# 关键先按时间切分再fit scalern len(df)train_end int(n * 0.65)val_end int(n * 0.80)train_df df.iloc[:train_end]val_df df.iloc[train_end:val_end]test_df df.iloc[val_end:]13. 生产级实践WGAN-LSTM真正难在“生成数据治理”为每一批合成数据记录生成器版本、训练区间、随机种子、增强比例和质量指标。不要把合成数据永久混入原始仓库应保留“真实/合成”来源标记便于回滚与消融。建立增益门槛只有当验证集和滚动回测同时改善时才启用生成增强。对极端值、结构突变和节假日等重要片段单独评估平均RMSE可能掩盖尾部风险。真实数据新增后优先重训/更新真实模型不要长期用旧生成器无限复制旧分布。在数据量足够时先尝试更强的基线、特征与正则化GAN并不是默认必选项。图16WGAN-LSTM的适用边界最适合样本不足且结构相对稳定的场景14. 为什么这次 WGAN 没有击败简单 Jitter这是一个很有价值的结果。Jitter 只是给真实窗口施加小幅噪声它不会改变窗口的大体结构因此在“已有样本附近做局部扩张”时非常保守。本文的 WGAN 虽然能生成更自由的新轨迹但真实/合成判别 AUC 高达 0.966说明它仍在某些高维结构上偏离真实窗口。25% 增强时这些新轨迹带来的多样性大于偏差因此总体有收益当比例继续升高偏差占主导模型性能迅速恶化。这也给出了一个非常实用的模型选择顺序先做可靠的无增强基线再尝试便宜的局部增强只有当数据稀缺或需要生成全新轨迹时再投入 WGAN 这类更昂贵的分布学习方法。15. 总结真正可靠的 WGAN-LSTM不是“GAN LSTM”四个字本文从严格实验协议出发完整实现了 WGAN-GP 生成增强与 LSTM 多步预测。实验没有刻意追求“组合模型一定第一”的结论而是得到更具工程价值的边界在低数据量场景中25% WGAN 合成窗口把 LSTM 测试 RMSE 从 0.6117 降到 0.5220证明适量生成增强可以有效但继续增加合成比例会快速恶化预测且本次简单 Jitter 增强与全量真实数据训练仍优于 WGAN-LSTM。因此WGAN-LSTM 的正确使用方式不是“先生成很多数据再训练”而是先严格切分时间 → 只用训练集训练 WGAN → 审计合成质量 → 在验证集上选择合成比例 → 最后只在未见真实测试集上评估一次。只要缺少其中任何一步漂亮的生成图和很低的训练损失都不足以证明模型真正有效。最终结论WGAN的价值是“补充训练分布”不是“创造新事实”。合成数据越多生成器偏差被复制得越多只有真实测试集上的稳定增益才是WGAN-LSTM值得上线的理由。参考资料1. Martin Arjovsky, Soumith Chintala, Léon Bottou. Wasserstein GAN. arXiv:1701.07875, 2017.2. Ishaan Gulrajani et al. Improved Training of Wasserstein GANs. arXiv:1704.00028, 2017.3. Sepp Hochreiter, Jürgen Schmidhuber. Long Short-Term Memory. Neural Computation, 1997.4. Jinsung Yoon, Daniel Jarrett, Mihaela van der Schaar. Time-series Generative Adversarial Networks. NeurIPS 2019.5. Data augmentation for time series regression: Applying transformations, autoencoders and adversarial networks to electricity price forecasting. Applied Energy, 2021.6. A LSTM-GAN Algorithm for Synthetic Data Generation of Time Series Data for Condition Monitoring. Procedia Computer Science, 2024.7. PyTorch Documentation: torch.nn.LSTM / autograd. Accessed 2026.