ARTICLE DETAIL

资讯详情

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

GAIN缺失值填补实战:基于TensorFlow的生成对抗网络完整实现与调优

GAIN缺失值填补实战:基于TensorFlow的生成对抗网络完整实现与调优 简介基于生成对抗网络的缺失数据填补方法完整实现面向机器学习、数据挖掘等方向的开发者与研究者适用金融、医疗、传感等多领域数据预处理场景。以TensorFlow作为后端深度学习框架实现了GAIN、SGAIN、WSGAIN-CP、WSGAIN-GP四种主流对抗填补模型完整代码与十个标准评估数据集一同打包在6.72MB的zip压缩包内。压缩包共28个文件除7个Python源码文件分别承担模型构建、训练流程与数据装载外还包含10个CSV格式的数据集、XML工程配置文件、说明文档以及缓存文件目录划分清晰便于直接运行、修改与二次开发。目前已有1778人学习下载特别适合需要对比不同对抗填补机制、完成论文复现或课程设计的研究者使用。从说明文档、主程序到模型定义用户可快速定位入口基于自带数据进行实验也可替换自有数据集开展验证。 做数据分析和机器学习建模的人大概率都躲不过数据处理这个坎。尤其是拿到一批数据发现里面不是这里空一块、就是那里缺一截这时候最简单的办法是删行或者用均值填一下但稍微有点追求的人都会犹豫删掉太浪费均值填补又太粗暴万一把变量之间的关系搞坏了后面的模型效果直接崩。这几年我一直在折腾各种缺失数据填补方案从统计插补到多重插补都用过直到接触到基于生成对抗网络的缺失数据填补方法GAIN才发现原来这个老问题还能用生成模型来解而且效果确实能打。GAIN全称是Generative Adversarial Imputation Nets最早由Jinsung Yoon等人提出本质上是把生成对抗网络的思想迁移到缺失值补全场景中。我用TensorFlow复现了一遍完整版本跑通了训练、评估到落地应用的整套流程过程中踩了不少坑也积累了一些经验。这篇文章就把整个实现思路、关键代码、训练细节和排查技巧完整记录下来想折腾缺失数据填补的朋友可以直接参考这版实现去改。1. 内容整体设计与思路拆解1.1 为什么普通填补方法不够用偏偏要上生成对抗网络先聊一个最基础的问题为什么均值填补、中位数填补、前向填充这些方法不够用不是说它们完全不行而是它们有一个共同的致命弱点——没有充分挖掘特征之间的关联关系。比如一个用户表里年龄和收入有相关性如果只拿全局均值去填缺失的年龄等于把这种相关性抹掉了模型后续学到的模式就会变形。多重插补MICE好一些能利用其他特征去预测缺失值但它本质是回归思路对特征间复杂的高维非线性关系捕捉能力有限。而生成对抗网络的思路完全不一样。它不直接预测缺失值而是让一个生成器去学着造假数据另一个判别器去分辨哪些是观测数据、哪些是填补数据两边互相博弈、互相进步。放到缺失数据填补这个场景里生成器拿到的任务是针对缺失的位置生成尽可能真实的数值骗过判别器。判别器不仅要判断整个样本是否真实更关键的是要学会判断哪个位置的值是被填出来的。两边的博弈结果就是生成器被迫学习到了整个数据集的分布规律填出来的值自然比简单插补更符合特征间的关系。1.2 GAIN的核心结构拆解生成器、判别器和提示机制GAIN和普通GAN最大的区别在于它的数据形态和博弈方式都有调整。标准GAN处理的是图像那种完整数据而GAIN面对的是带缺失的表格数据。它的整体结构包含三个关键组成部分生成器Generator、判别器Discriminator还有一个很多人第一次看容易忽略的提示机制Hint Mechanism。生成器接收两部分输入一部分是带有缺失值的数据缺失位置先用0或者其他占位值填充另一部分是随机噪声向量用于给生成过程引入随机性。生成器的输出是一个完整的填补后数据但只对缺失位置的数值生效。判别器接收的是生成器补全后的完整数据以及提示向量输出的是一个掩码预测结果——也就是说它要判断数据中每个位置的值到底是真实观测值还是被生成的填补值。提示机制是GAIN的精髓所在。它按一定概率把真实的掩码信息泄露给判别器这样判别器才能在训练初期有足够的信息去区分真假。如果没有提示机制判别器几乎无法判断哪些值是真实观测值训练过程容易失控。实现中提示向量的生成逻辑是每个缺失位置按概率保留真实掩码值1表示缺失0表示观测其余位置用0.5填充作为模糊信号。这个设计初看不复杂实际训练中却直接影响收敛速度和最终填补质量后面我会详细说参数怎么调。1.3 为什么瞄准TensorFlow版本做完整实现现在说起深度学习框架PyTorch和TensorFlow的讨论度都很高各有各的生态。我这次选择TensorFlow版本首要原因是GAIN原版论文就是基于TensorFlow实现的官方代码也是TF架构直接还原TensorFlow版本更容易和原论文的参数设置对齐。另一个原因是TensorFlow在部署环节确实方便训练好之后导成SavedModel或者用TensorFlow Serving提供服务都不用额外做框架转换。我在网上搜了一圈发现GAIN的TensorFlow实现虽然能搜到但很多是搬运的论文附带的压缩包注释少、结构乱版本还停留在TF 1.x。用的时候一堆兼容性问题光是tf.contrib这块就劝退不少人。我当时想索性自己按TF 2.x的习惯写一版完整的动态图模式下调试直观得多控制流也不用绕来绕去。keras API配合自定义训练循环既能看清每一步的梯度流向也方便调整网络结构和损失函数。这版实现我后面用几个UCI标准数据集和一份业务部门的真实脱敏数据做了验证填补效果和稳定性都达到了可用的水平。2. 环境准备与工具选型解析2.1 Anaconda安装与TensorFlow环境配置要点真正动手写代码前环境是最容易消耗热情的环节。我见过很多人在这一步被各种版本兼容问题劝退尤其是TensorFlow这货版本之间的API变动非常大。第三个版本之前的代码迁移到2.x常有大量报错更别说那些年tf.contrib说没就没的情况了。关于Anaconda安装TensorFlow的实操我的建议是创建一个干净的独立环境不要直接装在base环境里。用conda管理环境的好处是可以锁定Python版本和依赖的版本范围不同项目互不干扰。我常用的创建命令是conda create -n tf_gain python3.8 conda activate tf_gain pip install tensorflow2.10这里要特别说一句版本选择的门道TensorFlow 2.10是最后一个原生支持Windows GPU的版本之后的版本在Windows下装GPU版得靠WSL2很多人第一次装不知道这个兼容坑。如果是Linux环境可以放心装更高版本2.10以上对Windows用户就不太顺了。CPU版本训练小数据集问题不大但如果数据量上去、隐藏层节点加多建议还是整个GPU版本。再补一个Anaconda安装时的经验conda默认的软件源在国内环境下往往下载速度感人建议先把conda源和pip源换成清华的镜像不然下载一个几百兆的包要等半天。这个操作属于老生常谈但确实能帮后面省出大量时间。2.2 版本兼容矩阵TensorFlow 2.x的API差异盘点这次实现里我主要用到了TensorFlow 2.x的几个核心API和网上流传的非官方1.x版本实现有不小的差别。具体包括用tf.keras.Sequential或者函数式API堆叠网络结构替代了1.x里繁琐的tf.variable_scope加手工初始化参数的操作用tf.GradientTape记录梯度并完成反向传播替代了tf.Session.run配合feed_dict的方式tf.compat.v1作为兼容层用来处理个别1.x风格的既有代码。写的时候要特别小心几个版本的坑tf.contrib在2.x里已经被彻底移除很多老代码里的tf.contrib.layers相关调用必须改写成tf.keras.layerstf.placeholder在动态图模式下也不需要了直接定义好输入张量就能前向传播tf.losses系列函数的位置也变了用tf.keras.losses更稳。我建议安装时直接定下版本不要装最新版因为最新版往往伴随新特性但也会引入潜在的不稳定因素。生产环境里稳定压倒一切选一个经过广泛验证的版本比追新更重要。TensorFlow 2.10配Python 3.8这组搭配我跑下来非常稳定整个训练过程没有遇到莫名其妙的底层报错。3. 核心细节解析与实操要点3.1 数据预处理标准化和掩码矩阵的构造逻辑GAIN对输入数据有一个硬性要求所有特征需要预先做归一化处理我习惯用最小最大归一化把数值压到0到1之间。这样做有多个好处生成器的激活函数通常选tanh或者sigmoid输出范围天然匹配判别器对输入尺度敏感范围一致可以避免某些特征值过大导致梯度震荡归一化后缺失位置用0填充也不会引入极端值干扰。数据预处理的第二个关键是构造掩码矩阵Mask Matrix。这个矩阵和原始数据形状一致观测位置值为1缺失位置值为0。这里有个容易搞反的点要先理清楚原论文掩码矩阵是1表示观测0表示缺失而后面提示机制生成的向量又是另一套逻辑。代码里这两处我特别加了注释防止自己和看代码的人绕晕。构造掩码矩阵的代码如下def create_mask(data): mask np.ones_like(data) mask[np.isnan(data)] 0 data[np.isnan(data)] 0 return data, mask实际操作时一定要在标准化之前把NaN标记单独提取出来不能先标准化再判断缺失。因为标准化过程用到的均值和方差本身就会受到缺失值影响处理顺序错了结果会有偏差。3.2 生成器与判别器结构设计的经验之谈网络结构这块我看过很多GAIN口径的实现隐藏层设置五花八门。有的把生成器和判别器都设计得特别庞大隐藏层节点数上千实际效果反而不如小而精的结构。原论文里使用的隐藏层节点数建议是数据维度的一到两倍左右我验证下来这个经验值比较靠谱。比如数据维度是22隐藏层设置为128节点表现就已经不错加到256收益有限但训练开销涨了一截。激活函数方面生成器中间隐藏层我用ReLU输出层用Sigmoid配合数据范围0-1判别器中间层同样用ReLU最后输出用Sigmoid因为每个位置输出的都是一类概率值。有个经验值得单独拿出来说生成器的输入除了带缺失的数据和随机噪声我建议再加上一个拼接的辅助特征比如逐样本的缺失比例或者缺失模式编号。这样生成器能感知到不同样本缺失严重程度的差异对重度缺失样本的填补效果有肉眼可见的提升。这个小改动原论文里没有直接提到但我在实际训练中对比过加上之后判别器的loss收敛更平稳。网络搭建的代码核心结构如下def build_generator(dim, hidden_units128): inputs tf.keras.Input(shape(dim * 2,)) x tf.keras.layers.Dense(hidden_units, activationrelu)(inputs) x tf.keras.layers.Dense(hidden_units, activationrelu)(x) x tf.keras.layers.Dense(dim, activationsigmoid)(x) return tf.keras.Model(inputs, x) def build_discriminator(dim, hidden_units128): inputs tf.keras.Input(shape(dim * 2,)) x tf.keras.layers.Dense(hidden_units, activationrelu)(inputs) x tf.keras.layers.Dense(hidden_units, activationrelu)(x) x tf.keras.layers.Dense(dim, activationsigmoid)(x) return tf.keras.Model(inputs, x)这里生成器输入维度是dim * 2分别是拼接后的带缺失数据和随机噪声。判别器输入维度同样是dim * 2是拼接后的填充数据和提示向量。网络结构不算复杂真正的技术含量在损失函数设计和训练循环的编写上。3.3 损失函数详解Hinge Loss和提示机制的配合GAIN的训练过程看似是两个网络在对抗实际损失函数的设计非常讲究。判别器输出的不是单个真假概率而是一个和特征维度相同的向量每个元素表示对应特征位置的值是不是被填补过的概率。这样设计的原因很明显GAIN的目标不只是判断样本整体真伪而是要精细到判断哪些位置是伪造的。损失函数这块我用的是Hinge Loss的变体形式。生成器的目标是最小化判别器对填补位置的预测概率等价于让判别器尽可能认为所有位置都是真实观测值。判别器的目标则是最小化它对真实观测位置和填补位置的分类误差即正确地区分哪些位置被生成器动了手脚。具体的损失计算逻辑如下def compute_loss(D_pred, G_pred, mask, hint_rate): # D_pred是判别器输出形状(B, dim) # G_pred是生成器输出形状(B, dim) # mask是掩码矩阵1表示观测0表示缺失 # 缺失位置权重设置为1观测位置权重降低 missing_w 1.0 observed_w 1.0 / hint_rate # 判别器对缺失位置的预测越接近0越好真实位置的可能性低 D_loss -tf.reduce_mean( missing_w * mask * tf.math.log(D_pred 1e-8) observed_w * (1 - mask) * tf.math.log(1 - D_pred 1e-8) ) # 生成器要让判别器对缺失位置也认为是真实位置 G_loss -tf.reduce_mean( (1 - mask) * tf.math.log(D_pred 1e-8) ) return D_loss, G_loss注意到这里的权重设置没有observed_w设成了1.0 / hint_rate目的就是调节判别器对真实观测位置和填补位置的关注度。hint_rate接近1时判别器看到的大多是真实的掩码信息它就可以非常明确地判断缺失的位置训练启动快hint_rate较小时判别器拿到的提示信息少任务难度增加生成器的压力反而小了。这样一个对抗过程的结果是生成器随着训练进行逐渐学会生成与真实观测值统计特征一致的填补值不仅单个特征分布接近特征间的联合分布也能在博弈中保留下来。3.4 提示机制的参数选择原则关于hint_rate这个参数我在实验里试过从0.5到1.0的多个值发现它对训练结果的影响不是线性的。hint_rate太高比如0.95以上判别器太容易判断出哪些是填的梯度对生成器的指导价值反而减弱生成器学到的只是表面规律hint_rate太低比如0.5以下判别器长期难以收敛生成器也得不到有效反馈整个对抗过程处于僵持状态。原论文里的默认值是0.9这是一个比较安全的起点但我实际用下来发现具体数据集的缺失模式也会影响最优值。如果数据缺失完全是随机缺失MCAR0.9效果很好如果缺失跟特征本身有关MNAR建议降到0.8左右给判别器一点模糊空间反而能让生成器探索到更深层的分布规律。实现提示向量生成的代码如下def generate_hint(mask, hint_rate): hint mask.copy() prob np.random.rand(*mask.shape) hint[prob hint_rate] 0.5 return hint4. 实操过程与核心环节实现4.1 完整训练循环的搭建与说明GAIN的训练流程可以用一句话概括交替更新生成器和判别器让两个网络在对抗中共同进步。但落实到代码要处理的细节其实挺多包括数据分批、梯度累积、loss记录等。我这份实现采用批量训练的方式每轮迭代从训练集中采样一个小批量样本依次完成正向传播、计算损失、反向传播、更新参数。生成器和判别器的更新频率可以1比1也可以2比1论文里没明确说我测试后觉得1比1最稳不然容易出现一方压倒另一方的情况。训练循环的核心代码如下# 训练参数 batch_size 128 epochs 2000 learning_rate 1e-3 # 初始化优化器 G_optimizer tf.keras.optimizers.Adam(learning_rate) D_optimizer tf.keras.optimizers.Adam(learning_rate) for epoch in range(epochs): for batch_start in range(0, len(data_normalized), batch_size): batch_end min(batch_start batch_size, len(data_normalized)) x_batch data_normalized[batch_start:batch_end] m_batch mask_matrix[batch_start:batch_end] # 生成提示向量 h_batch generate_hint(m_batch, hint_rate0.9) with tf.GradientTape() as tape_G, tf.GradientTape() as tape_D: # 拼接数据和噪声 z_batch np.random.uniform(0, 1, sizex_batch.shape) G_input np.concatenate([x_batch, z_batch], axis1) # 生成器前向传播 G_output generator(G_input, trainingTrue) # 只取缺失位置的生成值 filled_data x_batch * m_batch G_output * (1 - m_batch) # 判别器输入 D_input np.concatenate([filled_data, h_batch], axis1) D_output discriminator(D_input, trainingTrue) # 计算损失 D_loss, G_loss compute_loss(D_output, G_output, m_batch, hint_rate0.9) # 分别更新梯度 G_grads tape_G.gradient(G_loss, generator.trainable_variables) D_grads tape_D.gradient(D_loss, discriminator.trainable_variables) G_optimizer.apply_gradients(zip(G_grads, generator.trainable_variables)) D_optimizer.apply_gradients(zip(D_grads, discriminator.trainable_variables))一个容易被忽视的细节是填补值的还原操作训练时生成器输出的是整个样本的补全结果但实际填补只需要替换缺失位置。所以代码里专门用filled_data x_batch * m_batch G_output * (1 - m_batch)做了掩码合并确保观测位置的原始值不被覆盖。这一行看起来简单却是GAIN能保持数据原有信息的关键。4.2 观测值保留策略和收敛判断除了掩码合并我还在损失计算里做了一个小改进判别器更新时真实观测位置参与计算的权重只有缺失位置的1/hint_rate。这意味着训练初期判别器更关注缺失位置的判断而不会被大量真实观测样本带偏。这个细节原论文里的实现只是简单加权没解释为什么这么设计我在复现时分析下来认为这是为了平衡两类样本的数量差异。关于收敛判断GAIN有一个天然的优势判别器的loss可以当作训练进度的参考。训练初期判别器能轻松区分真实和生成的数据loss下降很快随着生成器能力的增强判别器的loss会缓慢上升这时说明生成器已经开始骗过判别器了。我一般观察生成器和判别器loss曲线的交叉点交叉之后再训练100到200个epoch效果基本就稳定了。如果两个loss都不动了但数值偏高很可能是学习率太大造成震荡建议调低一个数量级再训。4.3 填补效果评估不只是看RMSE训练完模型后怎么判断填补效果好不好很多人第一反应是算RMSE或者MAE。这确实是最常用的指标做法是把一批完整的数据手动挖掉一些值跑完填补后和真实值做对比。总结一下我自己实践中的做法# 手动制造缺失评估填补效果 def evaluate_imputation(model, data_complete, missing_rate0.2): data_missing data_complete.copy() mask np.ones_like(data_complete) for i in range(data_complete.shape[0]): n_missing int(data_complete.shape[1] * missing_rate) missing_idx np.random.choice(data_complete.shape[1], n_missing, replaceFalse) data_missing[i, missing_idx] 0 mask[i, missing_idx] 0 # 执行填补 filled model.impute(data_missing) # 计算RMSE只统计缺失位置 diff filled - data_complete mse np.sum(diff**2 * (1 - mask)) / np.sum(1 - mask) rmse np.sqrt(mse) return rmse但单纯看RMSE有盲区。我遇到过一种情况RMSE数值不算高但把所有样本的填补值拉出来看分布发现某些特征的填补值方差特别小几乎都在均值附近摆动。这说明生成器学会了预测均值但没有真正学会特征的波动范围。所以我还建议额外做一些分布检验比如对比填补后数据和原始数据在关键特征上的直方图、相关系数矩阵差异。判断标准很简单好的填补应该让缺失位置的数据分布和观测位置尽量一致特征间相关性也基本保持。原论文代码里实验场景主要是数据缺失率在20%到30%之间的情况这个范围内GAIN的填补效果非常明显。我额外测了40%和50%的高缺失率场景效果有下降但还是优于均值填补和MICE说明模型的稳健性还不错。5. 常见问题与排查技巧实录5.1 维度不匹配问题最频繁踩的坑GAIN实现过程碰到最多的报错就是维度不匹配。因为整个流程里要拼接带缺失数据、噪声、提示向量、掩码矩阵好几个参与方任何一个维度对不上都会炸。典型报错是ValueError: Dimensions must be equal类型。这类问题大部分情况出在数据经过标准化后形状发生变化或者掩码矩阵与数据矩阵的行列顺序不对。我的排查习惯是在网络输入构建之前用print(x_batch.shape, m_batch.shape, h_batch.shape, z_batch.shape)全部打出来逐一核对。四个张量的第一维batch size必须完全一致第二维通常是特征维度或者特征维度的两倍拼接后。把所有参与拼接的变量维度列出来对照检查比看报错信息省事得多。5.2 训练不收敛或损失函数震荡的排查方向训练不收敛是另一个高频问题。现象是loss值一直跳动没有规律地下降训练几百轮后生成器填出来的值还是一团糟。这种问题我遇到过太多次排查顺序可以按如下步骤来第一步是检查学习率是不是设置得过大GAIN的对抗训练本身就比较敏感学习率高很容易震荡。我在实践中发现学习率设置在1e-3到1e-4之间比较安全如果loss跳动明显优先调到1e-4。第二步是看hint_rate是否合适这个参数直接影响判别器的任务难度参数设置和缺失模式不匹配容易造成生成器和判别器loss收敛速度差距过大一个收敛了另一个还在原地打转。第三步是检查数据标准化操作。如果原始数据里有极端异常值还没处理掉即使做了归一化模型也会被个别样本牵着走。所以我在预处理阶段会先做异常值裁剪用3倍标准差截断会让训练过程踏实不少。5.3 填补结果全部趋同的处理经验还有一个比较隐蔽的问题训练成功了loss看起来正常但生成器对不同的缺失样本给出的填补值几乎一样。这说明生成器没有学会区分不同样本的特征模式退化成了类似均值填补的效果。我遇到这个问题的原因最初是随机噪声的维度设得太低。生成器输入里的噪声向量如果维度太小或者全部是常数生成器就没有随机性来源学来学去只会输出一个固定模式。我把噪声从标量扩展到和特征维度相同的向量并确保它每一轮训练都重新采样问题就解决了。另外生成器的能力如果太弱比如隐藏层节点太少它也学不到足够的样本差异。适当增加隐藏层宽度或者增加一层都可以缓解模式坍缩问题。5.4 不同缺失率下的表现和参数调整建议实测下来缺失率在10%到20%之间时GAIN填补效果显著优于均值填补和MICERMSE能降低20%到30%。缺失率超过40%时效果下降比较明显因为有效信息太少任何方法都很难精确恢复原始分布。这时建议先把缺失率高的特征标记出来看看有没有其他特征可以提供辅助信息。参数调节方面总结一下我常用的调参顺序先固定hint_rate为0.9把学习率调到稳定收敛的程度然后再动hint_rate观察生成器loss变化如果发现填补值的分布与实际不符再调整生成器和判别器的隐藏层宽度。不用一上来就疯狂调网络结构很多时候学习率和hint_rate的组合才是决定成败的关键。6. 经验总结与拓展思考GAIN这套方法我用到现在最深的感受是它把缺失值填补从插值思维提升到了分布拟合思维。传统方法试图用一个简单的函数去逼近缺失值而GAIN试图理解整个数据集背后的生成机制。当然这不意味着GAIN在所有场景都能碾压传统方法。小数据集、特征维度很低、数据模式简单的情况下多重插补可能已经够用堆一个对抗网络性价比不高。但数据维度高、特征关系复杂的场景GAIN真的很能打。目前GAIN还有一个可以改进的方向是结合自注意力机制。表格数据其实也存在特征的全局依赖关系比如某些特征之间隔了好几列但关系非常密切。自注意力可以捕捉这种跨特征依赖如果应用到GAIN的生成器里理论上效果会更好。另外一个方向是把时间序列的时序依赖引入生成器用LSTM或者Transformer结构替代普通全连接网络这条路我还没跑完但初步思路已经有了。如果你准备动手复现建议不要直接抄代码就跑先把数据预处理、掩码矩阵、提示机制这几块原理理解了再跑训练循环。否则遇到问题都不知道从哪里排查。有朋友问我要不要直接用官方GitHub的代码我的意见是可以先看但一定自己按2.x语法重写一遍——因为旧代码改动起来往往比自己从零写还费时间。最后再分享一个小技巧训练结束后保存好生成器模型对同类型的新数据做填补时直接加载模型推理即可不需要重新训练。只要新数据和训练数据分布相似填补质量基本不会差太多。这个特性在生产环境中特别实用相当于一次训练、多次使用省下的时间可不是一星半点。本文还有配套的精品资源点击获取
返回列表