PyTorch实现GAIN:用生成对抗网络填补缺失值
2026/9/8 23:10:28 网站建设 项目流程

简介:缺失数据填补是数据预处理中的常见难题,基于生成对抗网络的GAIN系列方法提供了生成式解决思路。这份PyTorch完整实现面向有一定Python与神经网络基础、希望探究生成式填补模型的研究者或开发者,整合了GAIN、SGAIN、WSGAIN-CP、WSGAIN-GP四种算法,并提供十个数据集用于实验对比。包内共27个文件,包括5个Python源码、10个CSV数据文件,以及XML工程配置与Markdown说明文档,压缩后仅6.72MB,便于本地快速运行。目前已有2191人学习下载,既可作为教学示例,也能嵌入实际预处理流程。代码结构清晰,模型定义、数据加载与训练入口相互分离,通过运行示例可直观理解生成器与判别器的对抗训练过程,并在多个数据集上横向比较不同填补方案的效果,便于读者复现和二次开发。 处理缺失值这件事,几乎每个做数据项目的朋友都躲不掉。最开始我习惯用均值、中位数或者前向填充去糊弄,缺失比例低的时候看起来没毛病,一旦数据缺失超过10%,或者变量之间并不是简单的线性关系,这些传统方法会让下游模型的结果明显走形。后来接触到生成对抗网络(GAN)在数据填补方向上的变体GAIN(Generative 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 完整实现流程

整个项目的流程可以分成四步:

  1. 数据标准化,把所有变量调整到0-1或均值为0方差为1的范围;
  2. 按指定缺失率随机生成掩码矩阵,用掩码盖住部分真实值;
  3. 构建生成器和判别器两个全连接网络,配合Hint矩阵计算对抗损失和重建损失;
  4. 交替迭代训练,最后在测试集上比较填补值和真实值的误差。

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_rate=0.9)从真实掩码M中随机截取一部分信息,剩下的位置设为0.5或者随机值,构成H。简单说,H不是完整的真实掩码,而是带有噪声的部分掩码。

为什么需要这个噪声?如果H=M,判别器只要看H就知道哪些是缺失的,生成器会变得非常容易骗过判别器,但学习不到数据分布;如果H全是0或者随机值,判别器又得不到任何提示,训练又会退回普通GAN那种混乱状态。所以我们要让H保持一个“模糊提示”的程度,这样才能逼着生成器在有限信息下尽可能真实地填补。我在实现中发现,hint_rate取0.8-0.9之间效果比较稳定,太低和太高都会让训练不稳。

2.3 损失函数设计细节

GAIN的损失函数分成两个部分:对抗损失和重建损失。

对抗损失就是标准GAN那套:判别器要最大化区分真实观测位置和生成填补位置的概率;生成器要最小化判别器正确分类的概率,也就是尽量让判别器认为填补出来的位置也是真实观测值。

重建损失是GAIN另一个关键点。它的思想很朴素:对于已经观测到的位置,生成器应该把原始值尽量无损地还原出来;对于缺失位置,才去生成新的值。所以重建损失只在M=1的位置上计算生成器输出和原始标准化数据之间的均方误差(MSE)。这个损失给生成器加了很强的约束,让它在对抗训练的同时不会丢失输入信息。

生成器总损失 = 对抗损失 + alpha * 重建损失。alpha一般取1到100之间,我最后选了alpha=10,既能保持对抗训练的强度,又不会让重建损失压过生成多样性。

3. 基于PyTorch的完整实现过程

3.1 环境配置与依赖安装

我这里用的是Python 3.10,PyTorch 2.0以上版本,实际操作中1.8以上应该都没问题。建议用Anaconda建一个干净环境,避免和系统其他项目冲突。核心依赖只有torch、numpy、pandas、scikit-learn和matplotlib。

conda create -n gain python=3.10 conda activate gain # 按自己机器的CUDA版本选择合适的torch,CPU版去掉+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_rate=0.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_rate=0.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_dim=128): 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], dim=1) return self.fc(inp) class Discriminator(nn.Module): def __init__(self, input_dim, hidden_dim=128): 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], dim=1) 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.001,batch_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 RMSE
10%0.1380.1140.089
20%0.2210.1870.152
40%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_dim=128,batch_size=128,lr=0.001,alpha=10,hint_rate=0.9,epoch=2000。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距离来进一步提升稳定性。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询