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

构建MNIST GAN遇阻:求解释损失函数工作原理及代码疑问

Understanding GAN Loss Functions for Your MNIST Implementation

Let’s break down what’s happening with your loss functions and why they might be holding back your GAN from working properly. I’ll start with the core goals of each component, then map that to your code.

Core GAN Objectives

The discriminator (D) exists to tell real data apart from fake data made by the generator (G). The generator’s job is to fool the discriminator into thinking its outputs are real. These two objectives are competing, which drives the GAN to learn.

Discriminator Loss: What’s Missing?

Your current discriminator loss only accounts for real data:

loss_d = -tf.reduce_mean(tf.log(discriminator(real_data)))

But the discriminator needs to learn two things:

  1. Assign high confidence (close to 1) to real data: maximize log(D(real_data))
  2. Assign low confidence (close to 0) to fake data: maximize log(1 - D(G(noise_input)))

Since optimizers minimize loss by default, we take the negative of this combined objective to turn it into a minimization problem. Your loss is missing the fake data term—this means the discriminator never learns to spot fakes, so the generator has no reason to improve its outputs.

The correct discriminator loss should look like this:

# Get predictions for real and fake data
d_real = discriminator(real_data)
d_fake = discriminator(generator(noise_input), trainable=False)

# Combine both objectives into a single loss to minimize
loss_d = -tf.reduce_mean(tf.log(d_real) + tf.log(1 - d_fake))

Generator Loss: You’re on the Right Track!

Your generator loss uses a common improvement over the original GAN formula:

loss_g = -tf.reduce_mean(tf.log(discriminator(generator(noise_input), trainable = False)))

Originally, generators were trained to maximize log(1 - D(G(noise))), but this leads to vanishing gradients early on (when the discriminator is good at spotting fakes, 1-D(G) is near 1, so the log is close to 0 and gradients are tiny). By maximizing log(D(G(noise))) instead (via minimizing its negative), you get stronger, more stable gradients early in training—great call here.

Fixing Your Full Training Setup

Here’s how to adjust your code to align with standard GAN training practices:

  1. Update Loss Calculations:

    # Get real/fake predictions
    d_real = discriminator(real_data)
    d_fake = discriminator(generator(noise_input), trainable=False)
    
    # Discriminator loss (minimize the negative of its desired maximization)
    loss_d = -tf.reduce_mean(tf.log(d_real + 1e-7) + tf.log(1 - d_fake + 1e-7))
    # Generator loss (minimize negative of log(D(fake)))
    loss_g = -tf.reduce_mean(tf.log(d_fake + 1e-7))
    

    The 1e-7 adds a small buffer to avoid log(0) errors, which can crash training.

  2. Separate Training Steps:
    You need to train the discriminator and generator separately (alternating steps) so they don’t interfere with each other’s learning:

    # Train discriminator (only update discriminator weights)
    train_d = tf.train.AdamOptimizer(learning_rate).minimize(loss_d, var_list=discriminator.trainable_variables)
    # Train generator (only update generator weights)
    train_g = tf.train.AdamOptimizer(learning_rate).minimize(loss_g, var_list=generator.trainable_variables)
    
  3. Alternate Training Cycles:
    During training, run 1-5 discriminator steps for every generator step to keep the discriminator from getting too weak or too strong. For example:

    # In your training loop
    for _ in range(5):
        sess.run(train_d, feed_dict={...})  # Train discriminator 5x
    sess.run(train_g, feed_dict={...})      # Train generator once
    

Quick Tips for Stability

  • Always freeze the opposite model’s weights when training one (you’re already doing this with trainable=False—keep that up!).
  • Start with a low learning rate (like 1e-4) to avoid unstable training.
  • Monitor both losses over time: if the discriminator loss drops to near 0, it’s too strong and the generator can’t learn. If both losses stay high, the model isn’t converging.

With these fixes, your GAN should start learning to generate plausible MNIST digits. Let me know if you run into specific issues like mode collapse or vanishing gradients after making these changes!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 08:13:32