Keras实现的Adversarial Autoencoder无法正常训练学习问题
AAE训练异常问题修复方案
核心错误原因
- 对抗训练标签逻辑完全倒置:训练encoder的对抗损失时,你的目标是让encoder输出的隐向量被判别器判定为符合先验分布的真样本,应该把标签设为
y_real而不是y_fake。你当前的代码是要求encoder输出的向量被判为假,最终会出现判别器越准、encoder损失越高的情况,完全不符合AAE的训练目标。 - 硬编码batch_size存在隐患:你固定写死
batch_size=200,如果数据集总长度不能被200整除,最后一个不足长度的batch会直接报错,建议改为动态获取当前batch的尺寸。 - 重复计算冗余:多次调用
self.encoder(data)会增加不必要的计算开销,可以缓存第一次计算的隐向量结果复用。
修复后的核心代码(train_step部分)
def train_step(self, data): # 动态获取batch大小,兼容任意batch设置 batch_size = tf.shape(data)[0] # 只计算一次隐向量,复用结果 latent = self.encoder(data) # 生成符合先验的真样本 dists = tf.random.normal((batch_size,4,4,128)) y_real = tf.ones((batch_size, 1)) y_fake = tf.zeros((batch_size, 1)) real_dist_mix = tf.concat((dists, latent),axis=0) y_real_fake_mix = tf.concat((y_real, y_fake),axis=0) # 训练判别器 with tf.GradientTape() as tape: predictions = self.discriminator(real_dist_mix) d_loss = self.lossBCE(y_real_fake_mix, predictions) grads = tape.gradient(d_loss, self.discriminator.trainable_weights) self.dis_optimizer.apply_gradients(zip(grads, self.discriminator.trainable_weights)) # 训练encoder+decoder的重建任务 with tf.GradientTape() as Etape, tf.GradientTape() as Dtape: reconstruction = self.decoder(latent) reconstruction_loss = self.lossMAE(data, reconstruction) total_loss = reconstruction_loss Egrads = Etape.gradient(total_loss, self.encoder.trainable_weights) self.e_optimizer.apply_gradients(zip(Egrads, self.encoder.trainable_weights)) Dgrads = Dtape.gradient(total_loss, self.decoder.trainable_weights) self.de_optimizer.apply_gradients(zip(Dgrads, self.decoder.trainable_weights)) # 训练encoder的对抗任务,核心修改:标签改为y_real with tf.GradientTape() as tape: predictions = self.discriminator(latent) e_loss = self.lossBCE(y_real, predictions) # 可选:给对抗损失加权重,平衡重建和对抗任务,比如e_loss *= 0.2 grads = tape.gradient(e_loss, self.encoder.trainable_weights) self.e_optimizer.apply_gradients(zip(grads, self.encoder.trainable_weights)) self.total_loss_tracker.update_state(total_loss) self.reconstruction_loss_tracker.update_state(reconstruction_loss) self.disc_tracker.update_state(d_loss) self.discEnc_tracker.update_state(e_loss) return { "loss": self.total_loss_tracker.result(), "reconstruction_loss": self.reconstruction_loss_tracker.result(), "disc_loss": self.disc_tracker.result(), "discEnc_loss": self.discEnc_tracker.result(), }
额外调参建议
- 你当前判别器的学习率为1e-5,encoder为1e-4,如果还是出现判别器损失下降过快、encoder对抗损失波动大的情况,可以进一步把判别器的学习率下调到3e-6,或者给对抗损失增加权重系数,平衡两个任务的梯度强度。
- 可以尝试把重建损失和对抗损失合并后一次性更新encoder参数,避免两次更新带来的梯度冲突,训练稳定性会更好。
内容的提问来源于stack exchange,提问作者DmiSH
相关产品推荐
相关产品推荐

