You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 15:36:03