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

GAN判别器真假图像分批次训练失效?同批/分批结果差异原因解析

GAN分批次训练判别器失效问题分析

问题背景

我正在训练一款生成32×32小图像的静态GAN,为解决生成图像同质化问题,计划引入MinibatchStdev改进。原本将真假图像放在同一批次训练判别器时能得到尚可结果,但改为分两个批次分别训练真假图像后,生成图像始终是随机噪点,仿佛未进行训练,且loss_g与loss_d下降极快。调整训练顺序等方法均无改善,寻求该现象的原因。

原同批次训练代码

def train_gan(generator, discriminator, gan, dataset, batch_size, codings_size, noise8x8_dim=4, noise16x16_dim=0,
              noise32x32_dim=0, n_epochs=5, start_epoch=0, model_version=0, specific_model=''):
    # generator, discriminator = gan.layers
    for epoch in range(start_epoch, n_epochs):
        print(f"Epoch: {epoch + 1} / {n_epochs}")
        loss_d = []
        loss_g = []
        for X_batch in dataset:
            # 1. Train discriminator
            inputs = []
            noise = tf.random.normal(shape=[batch_size, codings_size])
            inputs.append(noise)
            if noise8x8_dim:
                inputs.append(tf.random.normal(shape=[batch_size, noise8x8_dim]))
            if noise16x16_dim:
                inputs.append(tf.random.normal(shape=[batch_size, noise16x16_dim]))
            if noise32x32_dim:
                inputs.append(tf.random.normal(shape=[batch_size, noise32x32_dim]))
            generated_images = generator(inputs)

            real_images = X_batch * 2 - 1  # Adapt real images between -1 and 1 because generator uses tanh as activation

# critical part of the code:
            X_fake_and_real = tf.concat([generated_images, real_images], axis=0)
            # y1 = tf.constant([[0.]] * batch_size + [[1.]] * batch_size)
            y1 = tf.constant([[0.]] * batch_size + [[1.]] * len(X_batch))

            discriminator.trainable = True
            l_d = discriminator.train_on_batch(X_fake_and_real, y1)
            loss_d.append(l_d)

            # 2. Train generator
            for i in range(2):
                inputs = []
                noise = tf.random.normal(shape=[batch_size, codings_size])
                inputs.append(noise)
                if noise8x8_dim:
                    inputs.append(tf.random.normal(shape=[batch_size, noise8x8_dim]))
                if noise16x16_dim:
                    inputs.append(tf.random.normal(shape=[batch_size, noise16x16_dim]))
                if noise32x32_dim:
                    inputs.append(tf.random.normal(shape=[batch_size, noise32x32_dim]))

                y2 = tf.constant([[1.]] * batch_size)
                discriminator.trainable = False
                l_g = gan.train_on_batch(inputs, y2)
                loss_g.append(l_g)


        l_g = np.mean(loss_g)
        l_d = np.mean(loss_d)
        print(f"loss_d: {l_d}, loss_g: {l_g}")

修改后的分批次训练关键代码

y_gen = tf.constant([[0.]] * batch_size)
y_real = tf.constant([[1.]] * len(X_batch))
l_d = discriminator.train_on_batch(real_images, y_real)
loss_d.append(l_d)
l_d = discriminator.train_on_batch(generated_images, y_gen)
loss_d.append(l_d)

核心原因分析

  • 判别器权重更新冲突:分批次训练时,判别器先在真实图像上更新权重(目标输出1),紧接着在生成图像上更新(目标输出0),两次更新的梯度方向完全相反,会相互抵消甚至导致权重震荡。而同批次训练时,真假样本的梯度混合计算,权重更新更稳定,能平衡对两类样本的判别能力。
  • MinibatchStdev功能失效:MinibatchStdev依赖批次内样本的统计差异来区分同质化生成图像,但分批次训练时每个批次只有单一类型样本,无法计算跨样本差异,直接失去约束生成器的作用,导致生成器退化为输出随机噪点。
  • 快速过拟合与无效反馈:判别器在单一类型样本上快速拟合,loss_d快速下降,但此时判别器对生成样本的判别逻辑已经偏移,生成器无法得到有效反馈——loss_g下降只是因为判别器失效,并非生成器学到了有效特征。

解决建议

  1. 若坚持分批次训练,可在两次判别器更新之间加入梯度累积,或降低学习率,缓解权重更新的冲突。
  2. 确保MinibatchStdev正常工作的前提是:输入批次包含真假混合样本,或分批次时每个批次保留一定比例的两类样本。
  3. 先在混合批次上训练判别器至稳定状态,再逐步切换到分批次训练,给模型适应过程。

内容的提问来源于stack exchange,提问作者Oniric

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 17:37:03