Keras实现WGAN-GP训练时生成图像异常排查
从DCGAN适配的WGAN-GP代码存在3个核心逻辑错误,和网络结构、学习率参数无关,属于训练流程和API调用的硬错误,直接导致训练不收敛,生成异常样例如下:
具体错误点及修正方案
错误1:梯度惩罚的插值系数采样不符合论文要求
当前使用tf.random.normal采样插值系数alpha,生成的值会超出[0,1]区间,导致插值点不在真实样本和生成样本的连线段上,梯度惩罚的Lipschitz约束完全失效。
修正方式:将alpha采样改为[0,1]区间的均匀分布:# 原错误写法 # alpha = tf.random.normal([batch_size, 1, 1, 1], 0.0, 1.0) # 修正后 alpha = tf.random.uniform([batch_size, 1, 1, 1], 0.0, 1.0)错误2:判别器前向计算放在了梯度带作用域外
代码在with tf.GradientTape() as tape:外部就执行了真实图、生成图的判别器前向推理,梯度带无法追踪这部分运算的计算图,算出来的判别器损失梯度缺失,参数更新完全错误。
修正方式:将判别器对真实图、生成图的前向计算全部移到梯度带作用域内,同时前向传播时指定training=True,保证BN、Dropout等层运行在训练模式:for i in range(self.d_steps): random_latent_vectors = tf.random.normal(shape=(batch_size, self.latent_dim)) generated_images = self.generator(random_latent_vectors, training=True) with tf.GradientTape() as tape: # 前向计算移到梯度带内部 generated_logits = self.discriminator(generated_images, training=True) real_logits = self.discriminator(real_images, training=True) d_cost = self.d_loss_fn(real_logits, generated_logits) gp = self.gradient_penalty(batch_size, real_images, generated_images) d_loss = d_cost + gp * self.gp_weight grads = tape.gradient(d_loss, self.discriminator.trainable_weights) self.d_optimizer.apply_gradients( zip(grads, self.discriminator.trainable_weights) )错误3:梯度惩罚函数参数传反
定义gradient_penalty函数时参数顺序为(batch_size, real_images, fake_images),但调用时传入的顺序是(batch_size, generated_images, real_images),会导致插值计算逻辑和预期不符,配合错误的alpha采样会进一步放大训练不稳定问题,调整调用时的参数顺序和函数定义一致即可。
额外适配注意点
从DCGAN迁移到WGAN-GP时,必须去掉判别器最后一层的Sigmoid激活函数:WGAN的判别器输出是无界的Wasserstein距离估计,不需要映射到[0,1]区间的分类概率,保留Sigmoid会直接导致损失计算错误。
内容的提问来源于stack exchange,提问作者pceccon

