PyTorch实现GAIN:用生成对抗网络填补缺失值

PyTorch实现GAIN:用生成对抗网络填补缺失值
简介缺失数据填补是数据预处理中的常见难题基于生成对抗网络的GAIN系列方法提供了生成式解决思路。这份PyTorch完整实现面向有一定Python与神经网络基础、希望探究生成式填补模型的研究者或开发者整合了GAIN、SGAIN、WSGAIN-CP、WSGAIN-GP四种算法并提供十个数据集用于实验对比。包内共27个文件包括5个Python源码、10个CSV数据文件以及XML工程配置与Markdown说明文档压缩后仅6.72MB便于本地快速运行。目前已有2191人学习下载既可作为教学示例也能嵌入实际预处理流程。代码结构清晰模型定义、数据加载与训练入口相互分离通过运行示例可直观理解生成器与判别器的对抗训练过程并在多个数据集上横向比较不同填补方案的效果便于读者复现和二次开发。 处理缺失值这件事几乎每个做数据项目的朋友都躲不掉。最开始我习惯用均值、中位数或者前向填充去糊弄缺失比例低的时候看起来没毛病一旦数据缺失超过10%或者变量之间并不是简单的线性关系这些传统方法会让下游模型的结果明显走形。后来接触到生成对抗网络GAN在数据填补方向上的变体GAINGenerative Adversarial Imputation Nets正好解决了“用生成器去拟合数据分布”这个关键问题。我这次用PyTorch把GAIN从零到一完整实现了一遍包含掩码生成、Hint机制、对抗训练、效果评估全套流程整个过程踩了不少坑也把教训整理成一篇可以照抄的实现笔记。1. 项目背景与整体设计思路1.1 缺失数据为什么不能用简单填充糊弄缺失数据在真实场景里太常见了传感器断线、用户跳题、后台日志漏记都会造成缺失。按照统计学的说法缺失机制可以分为三类完全随机缺失MCAR、随机缺失MAR和非随机缺失MNAR。MNAR是最棘手的缺失本身就和未观测值相关任何只依赖观测数据的填充方法都会有偏差。常见的处理方式一般有三种丢弃、单值填充和多重插补。丢弃数据简单粗暴但会损失样本量均值/中位数填充实现方便但会压缩方差导致各个变量之间的关系被扭曲多重插补比如MICE虽然能考虑变量之间的相关性但本质还是基于线性回归或树模型去迭代对复杂非线性分布的表达能力很有限。GAIN的思路完全不一样它不假设数据服从哪种分布而是用神经网络去逼近真实的数据分布再从学到的分布里采样缺失部分的合理取值。1.2 为什么选择GAIN而不是普通GAN刚开始我确实想过能不能直接拿WGAN或者DCGAN来填补缺失值但实际一跑就发现方向不对。普通GAN的输入是纯噪声生成器完全没有看到已有的观测数据判别器也分不清哪些位置是缺失的、哪些是真实观测到的训练出来的结果基本等于在随机生成数据。GAIN的核心改进在于增加了一个“提示矩阵”Hint Matrix。这个Hint有点像考试时老师给划的重点它不直接告诉判别器正确答案但会给出部分关于真实掩码的信息强制判别器不能只靠“缺失位置就是0”这个偷懒规律来区分真假。生成器在Hint的帮助下才能学会条件分布也就是在给定观测值的条件下生成缺失部分的合理填补。这是GAIN区别于普通GAN最核心的一点也是它能真正用在实际数据填补上的原因。1.3 完整实现流程整个项目的流程可以分成四步数据标准化把所有变量调整到0-1或均值为0方差为1的范围按指定缺失率随机生成掩码矩阵用掩码盖住部分真实值构建生成器和判别器两个全连接网络配合Hint矩阵计算对抗损失和重建损失交替迭代训练最后在测试集上比较填补值和真实值的误差。PyTorch的动态计算图在这个场景里非常有优势因为每个batch的掩码和Hint矩阵都是随机变化的动态图可以方便地处理这种输入结构变化写起来比静态图要自然很多。2. GAIN网络结构与关键原理拆解2.1 生成器和判别器怎么搭GAIN里的两个网络结构不需要太复杂我用的是多层全连接网络。以维度为d的输入数据为例生成器输入是“噪声z 观测值x·M 掩码M”其中M是0/1掩码1表示观测到0表示缺失。噪声z用来提供生成多样性生成器输出是一个和x同维度的向量表示对缺失位置的填补值判别器输入是“填补后的完整数据 掩码M Hint矩阵H”其中H是0/1矩阵判别器输出是每个位置属于“真实观测值”的概率维度同样是d。生成器和判别器内部我都加了BatchNorm激活函数用ReLU输出层用Sigmoid把数值压到[0,1]区间。需要注意的是输入数据标准化到[0,1]后生成器输出激活函数用Sigmoid比较适合如果后续要做复杂的连续值插补也可以改成带约束的线性激活我建议先用Sigmoid跑通流程再说。2.2 Hint Matrix到底做了什么Hint矩阵这一步可能是新手最容易误解的地方。原论文里的设计是对于每个样本以一定概率比如hint_rate0.9从真实掩码M中随机截取一部分信息剩下的位置设为0.5或者随机值构成H。简单说H不是完整的真实掩码而是带有噪声的部分掩码。为什么需要这个噪声如果HM判别器只要看H就知道哪些是缺失的生成器会变得非常容易骗过判别器但学习不到数据分布如果H全是0或者随机值判别器又得不到任何提示训练又会退回普通GAN那种混乱状态。所以我们要让H保持一个“模糊提示”的程度这样才能逼着生成器在有限信息下尽可能真实地填补。我在实现中发现hint_rate取0.8-0.9之间效果比较稳定太低和太高都会让训练不稳。2.3 损失函数设计细节GAIN的损失函数分成两个部分对抗损失和重建损失。对抗损失就是标准GAN那套判别器要最大化区分真实观测位置和生成填补位置的概率生成器要最小化判别器正确分类的概率也就是尽量让判别器认为填补出来的位置也是真实观测值。重建损失是GAIN另一个关键点。它的思想很朴素对于已经观测到的位置生成器应该把原始值尽量无损地还原出来对于缺失位置才去生成新的值。所以重建损失只在M1的位置上计算生成器输出和原始标准化数据之间的均方误差MSE。这个损失给生成器加了很强的约束让它在对抗训练的同时不会丢失输入信息。生成器总损失 对抗损失 alpha * 重建损失。alpha一般取1到100之间我最后选了alpha10既能保持对抗训练的强度又不会让重建损失压过生成多样性。3. 基于PyTorch的完整实现过程3.1 环境配置与依赖安装我这里用的是Python 3.10PyTorch 2.0以上版本实际操作中1.8以上应该都没问题。建议用Anaconda建一个干净环境避免和系统其他项目冲突。核心依赖只有torch、numpy、pandas、scikit-learn和matplotlib。conda create -n gain python3.10 conda activate gain # 按自己机器的CUDA版本选择合适的torchCPU版去掉cu后缀即可 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install numpy pandas scikit-learn matplotlib如果你是纯CPU跑数据量不大也完全够用。GAIN的网络本身比较小主要开销是迭代次数多CPU上跑UCI规模的数据也就几分钟到十几分钟新人不用一上来就纠结CUDA环境。3.2 数据预处理与掩码生成一定要先做数据标准化。因为网络输出层用了Sigmoid输入数据最好落到[0,1]区间这样生成器学习起来更稳定。我是用sklearn的MinMaxScaler把每个特征缩放到[0,1]。掩码生成的部分是核心要按缺失率生成0/1矩阵def generate_mask(x, missing_rate0.2): n, d x.shape mask np.random.rand(n, d) missing_rate mask mask.astype(np.float32) return mask def generate_hint(mask, hint_rate0.9): n, d mask.shape hint np.random.rand(n, d) hint_rate hint hint.astype(np.float32) # 随机翻转一半的位置让它成为“不完整提示” hint mask * hint 0.5 * (1 - hint) return hint这里有朋友可能会问为什么hint要把部分位置设置成0.5而不是0因为0在这个二值判别问题里语义很明确缺失用0.5作为中性值可以弱化判别器对提示信息的绝对信任让生成器有更多学习空间。这是我在反复对比后发现的一个小细节原论文的代码里也是这么处理的。3.3 模型定义与关键代码用PyTorch定义生成器和判别器核心结构如下import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, input_dim, hidden_dim128): super().__init__() self.fc nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.ReLU(), nn.BatchNorm1d(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.BatchNorm1d(hidden_dim), nn.Linear(hidden_dim, input_dim), nn.Sigmoid() ) def forward(self, x, m, z): # x: 原始标准化数据m: 掩码z: 噪声 inp torch.cat([x * m, m, z], dim1) return self.fc(inp) class Discriminator(nn.Module): def __init__(self, input_dim, hidden_dim128): super().__init__() self.fc nn.Sequential( nn.Linear(input_dim * 2 input_dim, hidden_dim), nn.ReLU(), nn.BatchNorm1d(hidden_dim), nn.Linear(hidden_dim, hidden_dim), nn.ReLU(), nn.BatchNorm1d(hidden_dim), nn.Linear(hidden_dim, input_dim), nn.Sigmoid() ) def forward(self, x_hat, m, h): inp torch.cat([x_hat, m, h], dim1) return self.fc(inp)注意判别器输入维度写的是input_dim*2 input_dim其实等价于3*input_dim拆开写是为了提示自己输入拼接的是“数据掩码提示”。BatchNorm在batch等于1时会报错训练时记得保证batch_size大于1。3.4 训练循环与超参配置训练的核心是交替更新判别器和生成器。我用的顺序是每个batch先更新判别器再更新生成器生成器更新两次这样能让对抗平衡更稳。核心训练逻辑如下def train_step(batch_x, missing_rate, hint_rate, G, D, opt_G, opt_D): n, d batch_x.shape # 生成mask和hint m torch.FloatTensor(generate_mask(batch_x.numpy(), missing_rate)).to(device) h torch.FloatTensor(generate_hint(m.cpu().numpy(), hint_rate)).to(device) # 构造带噪声的缺失输入 z torch.randn(n, d).to(device) g_loss 0 # 更新判别器 x_hat G(batch_x, m, z) D_output D(x_hat.detach(), m, h) # 真实数据位置 真缺失位置 假 d_target m # 真实观测位置为1 opt_D.zero_grad() d_loss nn.BCELoss()(D_output, d_target) d_loss.backward() opt_D.step() # 更新生成器两次 for _ in range(2): z torch.randn(n, d).to(device) x_hat G(batch_x, m, z) D_output D(x_hat, m, h) g_adv_loss nn.BCELoss()(D_output, m) # 希望判别器输出1 g_rec_loss nn.MSELoss()(x_hat * m, batch_x * m) # 只约束观测位置 g_loss g_adv_loss alpha * g_rec_loss opt_G.zero_grad() g_loss.backward() opt_G.step() return d_loss.item(), g_loss.item()训练主体循环加上早停和loss打印总共2000个epoch左右就能跑出比较稳定的结果。优化器用Adam初始学习率0.001batch_size设128。如果你发现损失震荡很厉害可以把学习率降到0.0005或者增加生成器的迭代次数。4. 实际实验效果与调参经验4.1 在公开数据集上的填补性能对比我拿UCI的“Heart Disease”数据集做验证一共13个数值特征随机删除20%和40%的数据用RMSE作为评估指标。对比了均值填充、MICE和GAIN三种方法GAIN在缺失率20%时RMSE比均值填充低了近30%比MICE低了近15%。缺失率40%时优势更明显说明数据缺失越多GAIN学习分布的优势越能体现。缺失率均值填充 RMSEMICE RMSEGAIN RMSE10%0.1380.1140.08920%0.2210.1870.15240%0.3620.2980.234这个结果其实符合预期因为GAIN不是简单拟合线性关系而是把每个样本当成一个整体去生成特征之间的非线性交互能被网络学出来。4.2 训练收敛与稳定性观察训练过程中我盯过两条loss曲线。判别器loss大概在前500个epoch内从0.7降到0.5左右然后缓慢波动生成器loss里对抗部分会慢慢上升重建部分持续下降。如果发现判别器loss迅速趋近0说明Hint矩阵给的信息太多判别器“开挂”了这时要适当降低hint_rate或者提高alpha让生成器更关注重建。如果发现生成器重建loss下降但对抗loss不下降说明生成器在偷懒把所有缺失值都填成均值来降低MSE。解决办法是把alpha从10降到1或者让生成器的学习率比判别器稍微大一点点比如生成器用0.002判别器用0.001形成优势平衡。4.3 参数调优的一点心得我试过几组参数比较稳定的组合是hidden_dim128batch_size128lr0.001alpha10hint_rate0.9epoch2000。hidden_dim不需要太大因为GAIN处理的数据维度往往不高太大的隐层只会增加过拟合风险。还有一个容易忽略的点是数据顺序。训练前一定要打乱样本顺序如果原始数据按某个变量排序会导致mini-batch之间分布不一致训练过程会非常飘。5. 常见问题与避坑记录5.1 生成器输出全变成均值了怎么办这种情况绝大多数是重建损失权重alpha太大。生成器发现与其费劲对抗不如把缺失位置都填成特征均值这样重建loss已经很低了整体loss看起来也很漂亮但数据方差被严重压缩。我的判断方法是把alpha降到1-5之间同时观察生成器输出矩阵的方差如果方差回升到正常水平说明平衡点找到了。另外要注意观测位置缺失率太低时重建loss能提供的信息很少此时应该适当增大alpha否则生成器过于自由。5.2 判别器器loss瞬间掉到0出现这个现象先怀疑Hint矩阵构造代码是不是写错了。我曾经把hint做成了完整掩码M也就是判别器每一行都直接看到缺失标记loss当然毫无悬念地崩到0。Hint应该是由原始掩码随机保留一部分信息而不是完整信息。如果代码没问题那就是hint_rate设得太高。试着把hint_rate调成0.7-0.8让判别器的信息来源更模糊一些迫使它去学习数据本身的分布特征而不是钻提示信息的空子。5.3 标准化与反标准化千万别搞反训练前用MinMaxScaler把数据缩放到[0,1]生成器输出也在这个范围填完之后我们拿到的还是标准化空间的值。要得到真实尺度的填补结果必须用同一个scaler做inverse_transform。踩坑点在于如果你的数据里有类别变量比如性别、等级不能直接做MinMaxScaler否则类别编码之间的顺序关系会被错误引入。我的处理方式是数值变量单独缩放类别变量做独热编码然后拼接。填补的时候也是分开处理数值部分用GAIN类别部分用最高概率的类别去还原。5.4 NaN和维度问题怎么排查训练时出现NaN第一反应是学习率太大。Adam默认学习率0.001在GAIN里偶尔也会不稳定尤其当生成器梯度更新太激进时可以试试把学习率降到0.0005或者0.0002。维度不匹配则几乎都出现在拼接环节。记住生成器输入维度是3d数据掩码噪声判别器输入维度是3d补全数据掩码提示。每次改网络结构前先打印一下tensor的shape养成这个习惯能省不少调试时间。我自己整套流程跑下来最大的感受是GAIN的训练稳定性比对错更重要。网上很多代码能跑通但真正要迁移到自己的数据集上还是要反复看loss曲线和生成值的分布是否符合逻辑。建议新手朋友先在一个小数据集上把流程跑顺再慢慢调参数不要一上来就追求大网络和高缺失率。另外模型保存可以用torch.save(G.state_dict(), generator.pt)后续只需要调用生成器就能做填补判别器训练完就可以丢掉了。如果后续还有时间可以在生成器里加入卷积层去处理时序数据或者把损失函数换成Wasserstein距离来进一步提升稳定性。本文还有配套的精品资源点击获取

最新新闻

日新闻

周新闻

月新闻