RGAN配置参数完全手册experiment.py 50选项逐项详解【免费下载链接】RGANRecurrent (conditional) generative adversarial networks for generating real-valued time series data.项目地址: https://gitcode.com/gh_mirrors/rg/RGANRGANRecurrent Generative Adversarial Network是一个专门用于生成实值时间序列数据的开源项目其所有训练行为都由experiment.py主脚本的配置参数控制。本文为你逐项详解 RGAN 配置参数从数据来源、模型结构到差分隐私手把手教你读懂并调好每一个选项快速跑通自己的时间序列生成实验。RGAN 由 ETH Zurich 生物医学信息学组开发论文《Real-valued (Medical) Time Series Generation with Recurrent Conditional GANs》提出了用 LSTM 构建生成器和判别器的 GAN 架构能够生成正弦波、高斯过程采样曲线、MNIST 数字乃至 ICU 患者生命体征等时间序列。对新手来说最大的门槛往往不是模型本身而是experiment.py那几十个配置参数——本文就是为你准备的参数地图。一、RGAN配置参数从哪里来RGAN 的配置参数统一由utils.py中的rgan_options_parser()函数定义约 40 个命令行参数加载顺序非常清晰先用命令行参数或默认值初始化settings字典如果指定了--settings_file则用experiments/settings/目录下的同名 JSON 文件整体覆盖命令行参数运行experiment.py后最终生效的配置会原样保存到./experiments/settings/{identifier}.txt。所以你有两种调参方式命令行直接传参或编辑设置文件。项目自带 4 个可参考的设置文件test.txt正弦波演示、mnistfull.txt全量 MNIST、cristobal_mnist.txt和cristobal_eicu.txt论文实验配置。最简运行方式python experiment.py --settings_file test所有参数都可以通过--参数名 值的方式在命令行覆盖例如python experiment.py --data sine --num_epochs 200 --latent_dim 10二、RGAN全部配置参数速查表下表按功能把配置参数分为六类先建立全局认知后面逐项详解。类别参数默认值作用一句话元配置settings_file空加载设置文件的名称不带 .txt元配置identifiertest所有输出文件的名字前缀数据datagp_rbf数据类型gp_rbf / sine / mnist / load 等数据num_samples14000训练样本数量数据seq_length30每条时间序列的长度数据num_signals1时间序列的通道/特征数数据num_generated_features1生成器输出的特征维度数据normalisefalse是否对数据做标准化数据cond_dim0条件向量的维度0 表示无条件 GAN数据max_val1条件编码的取值范围数据one_hotfalse条件信息是否用 one-hot 编码数据predict_labelsfalse让模型输出标签而非条件输入数据scale0.1gp_rbf 高斯过程核的尺度数据freq_low/freq_high1.0 / 5.0正弦波频率范围数据amplitude_low/amplitude_high0.1 / 0.9正弦波振幅范围数据multivariate_mnistfalseMNIST 是否按多变量序列处理数据full_mnistfalse是否使用完整 MNIST 数据集数据data_load_from空从已有 .data.npy 文件加载数据数据resample_rate_in_min15eICU 数据重采样间隔分钟模型hidden_units_g100生成器 LSTM 隐层单元数模型hidden_units_d100判别器 LSTM 隐层单元数模型latent_dim5隐空间/噪声维度模型kappa1判别器中间步损失权重模型batch_meanfalse判别器损失是否拼接 batch 均值模型learn_scalefalse生成器输出的 scale 参数是否可学习训练learning_rate0.1优化器学习率训练batch_size28批大小训练num_epochs100训练轮数训练D_rounds5每轮判别器训练次数训练G_rounds1每轮生成器训练次数训练use_timefalse隐空间第 0 维是否强制对应时间训练WGANfalse是否使用 WGAN 损失训练WGAN_clipfalse是否对判别器参数做裁剪训练shuffletrue每轮是否打乱训练数据训练wrong_labelsfalse是否用错误标签增强判别器隐私dpfalse是否启用差分隐私 SGD隐私l2norm_bound1e-5单样本梯度 L2 范数上界隐私batches_per_lot1每个 lot 包含的批数隐私dp_sigma1e-5差分隐私噪声标准差三、数据类RGAN配置参数详解1.data选择数据源这是最核心的配置参数决定整个实验喂什么数据gp_rbf默认用高斯过程 RBF 核采样光滑曲线sine生成带随机频率和振幅的正弦波mnist把 MNIST 数字当作逐行扫描的时间序列load从data_load_from指定的.data.npy文件读取resampled_eICU/eICU_task处理 ICU 患者生命体征数据。2. 序列形状三件套seq_length、num_signals、num_samplesseq_length每条序列的时间步数。MNIST 常用 2828x28 图像按行展开eICU 任务用 16num_signals每条序列的特征通道数例如 4 通道生命体征对应num_signals: 4num_samples生成的训练样本总数test.txt用 14000mnistfull.txt用 60000。3. 数据专属参数scale与正弦波四参数scalegp_rbf 数据核函数的长度尺度控制曲线的粗糙程度freq_low/freq_high正弦波频率的采样范围amplitude_low/amplitude_high正弦波振幅的采样范围。{ data: sine, freq_low: 1.0, freq_high: 5.0, amplitude_low: 0.1, amplitude_high: 0.9 }4. MNIST 专属参数multivariate_mnist、full_mnistmultivariate_mnist: true表示把 MNIST 当作多变量时间序列每行像素作为一个时间步的向量full_mnist: true使用完整 60000 样本训练集mnistfull.txt的配置否则使用子集。5. 数据加载与预处理data_load_from、normalisedata_load_from直接复用之前生成的.data.npy跳过重复造数据非常省时间normalise在划分训练/验证/测试集时对数据做标准化mnistfull.txt中为 false。6. 条件信息参数cond_dim、max_val、one_hot当cond_dim 0时RGAN 变为条件 GANRCGAN可以按类别生成数据cond_dim条件向量维度eICU 任务用 7 表示 7 种状态one_hot: true时会把类别转成 one-hot 编码此时max_val自动被调整为 1max_val非 one-hot 模式下条件编码的取值范围上界。四、模型结构类RGAN配置参数详解1.hidden_units_g与hidden_units_d生成器和判别器的 LSTM 隐层单元数默认为 100。增大单元数可提升模型容量但也会显著增加显存占用和训练时间小数据集建议从 50~100 起步。2.latent_dim隐空间维度生成器输入噪声的维度默认 5。cristobal_mnist.txt用 3cristobal_eicu.txt用 10。维度越高生成样本的多样性潜力越大但训练难度也随之上升。3.kappa判别器中间步损失权重判别器不仅看 LSTM 的最后输出还会考虑每个时间步的中间结果。kappa 1表示全部使用中间步损失这是对时序 GAN 很重要的一个稳定训练技巧。4.batch_mean与learn_scalebatch_mean为判别器损失拼接整个 batch 的均值特征帮助判别器捕捉分布整体信息learn_scale生成器输出层的 scale 参数默认固定为 1开启后变为可学习变量参见model.py中scale_out_G的定义。五、训练类RGAN配置参数详解1.learning_rate、batch_size、num_epochs训练三件套。默认学习率 0.1配合 SGD 类优化器batch_size默认 28num_epochs默认 100——注意cristobal_eicu.txt训练了 1005 轮医学数据的收敛通常需要更多轮次。2.D_rounds与G_rounds对抗节奏GAN 训练的关键旋钮D_rounds 5, G_rounds 1默认判别器每轮训练 5 次、生成器 1 次判别器更强cristobal_eicu.txt用D_rounds1, G_rounds3让生成器更激进mnistfull.txt用D_rounds1, G_rounds1平衡对抗。判别器和生成器太弱都会导致模式坍塌或训练崩溃建议从默认值开始出现 loss 不收敛时再调整。3.use_time、WGAN、WGAN_clip、shuffleuse_time强制隐空间第 0 维与时间步对应帮助模型学到时间结构WGAN/WGAN_clip切换为 Wasserstein GAN 训练方式若开启 WGAN 通常也要开启参数裁剪shuffle每轮训练后打乱数据顺序默认开启。4.wrong_labels给判别器额外喂标签配对错误的真实样本强迫判别器学会区分标签与数据的关系是条件 GAN 提升稳定性的辅助手段。六、差分隐私配置参数详解RGAN 的一大亮点是支持差分隐私训练为医学数据提供更强的隐私保护。相关实现在differential_privacy/dp_sgd/dp_optimizer/目录。dp总开关设为 true 后用差分隐私 SGD 训练判别器l2norm_bound单个样本梯度的 L2 范数上界裁剪阈值test.txt默认 1e-5cristobal_eicu.txt设为 4batches_per_lot每个 lot 包含的批数控制隐私预算的消耗粒度dp_sigma加入噪声的标准差越大隐私保护越强、但模型质量下降越明显。开启dp后experiment.py会在experiments/traces/{identifier}.dptrace.txt中按轮记录隐私消耗目标 eps 从 0.125 到 8 的 delta 曲线方便你监控隐私预算。七、评估与输出参数1.identifier输出命名所有输出文件都以identifier为前缀配置存档experiments/settings/{identifier}.txt训练数据experiments/data/{identifier}.data.npy训练轨迹experiments/traces/{identifier}.trace.txt可视化图片experiments/plots/{identifier}_*.png2. 自动评估机制无需配置但要知道experiment.py内置了 MMD最大均值差异评估用mmd.py中的median_pairwise_distance计算核带宽通过mix_rbf_mmd2_and_ratio在验证集上评估生成质量并在每轮记录D_loss、G_loss、mmd2、that等指标到 trace 文件。当mmd2优于历史最佳且epoch 10时模型参数会自动保存model.dump_parameters这就是你最终要用的最佳模型。八、实战从零配置一个RGAN实验步骤 1创建自己的设置文件在experiments/settings/下新建my_sine.txt配置一个正弦波生成实验{ data: sine, num_samples: 14000, seq_length: 30, num_signals: 1, cond_dim: 0, hidden_units_g: 100, hidden_units_d: 100, latent_dim: 5, learning_rate: 0.1, batch_size: 28, num_epochs: 100, D_rounds: 5, G_rounds: 1, identifier: my_sine }步骤 2运行实验python experiment.py --settings_file my_sine步骤 3查看结果训练结束后去experiments/plots/看生成的曲线图去experiments/traces/my_sine.trace.txt分析 loss 与 mmd2 曲线。官方参考配置在experiments/settings/目录下四个 txt 文件中test.txt是入门演示cristobal_eicu.txt是论文级别的完整配置。九、常见调参问题速查Q1生成的序列像噪声怎么办先加大D_rounds让判别器更强或降低learning_rate让训练更稳同时确认seq_length和num_samples是否匹配数据规模。Q2训练不收敛、loss 震荡检查batch_size是否过小尝试调整D_rounds/G_rounds的比例MNIST 场景可参考mnistfull.txt的 1:1 配置。Q3想按类别生成数据设置cond_dim为类别数并配合one_hot: true参考cristobal_eicu.txtcond_dim7。Q4数据隐私敏感开启dp: true并按隐私需求调l2norm_bound和dp_sigma参考cristobal_eicu.txt中的 l2norm_bound4、dp_sigma0.6。掌握以上 RGAN 配置参数你就能像调音师一样精准控制每一次时间序列生成实验。从experiment.py入手对照utils.py的参数定义逐项实验很快就能调出满意的生成效果。【免费下载链接】RGANRecurrent (conditional) generative adversarial networks for generating real-valued time series data.项目地址: https://gitcode.com/gh_mirrors/rg/RGAN创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考