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

为什么我的GAN训练无进展且仅生成灰色图像?

问题诊断与修复方案

核心问题:缺失训练参数更新逻辑

你当前的train函数循环里只完成了数据读取、生成假图的步骤,完全没有调用训练接口更新判别器和生成器的权重,这是模型完全没有训练进展的根本原因。

其他需要修复的问题

  • 训练流程顺序错误:GAN的标准训练顺序是先训练判别器,再固定判别器权重训练生成器,你当前代码顺序完全乱了,也没有把真假图喂给判别器完成训练步骤。
  • 优化器不匹配:判别器用了RMSprop,生成器用了Adam,DCGAN的标准实现是二者都用学习率0.0002、beta1=0.5的Adam优化器,优化器不匹配容易导致训练不稳定。
  • 输入数据归一化不匹配:生成器最后一层用的是tanh激活,输出范围是[-1,1],你需要确认你的训练集X_train是否也做了(x / 127.5) - 1的归一化,二者范围不匹配会导致判别器很容易区分真假样本,出现模式崩溃。
  • 图像保存逻辑缺少截断:生成器输出还原到[0,1]范围后可能存在超出0或1的数值,直接乘255会出现uint8溢出,导致颜色异常,需要加generated_images = np.clip(generated_images, 0, 1)做截断。
  • 生成器结构冗余:你在三次反卷积之后又加了两层512通道的卷积,参数过多容易出现梯度消失,建议删掉多余的卷积层,反卷积到128*128尺寸之后直接接输出卷积层即可。

修复后的核心训练逻辑示例

def train(epochs=1000, batchSize=128):
  batchCount = X_train.shape[0] // batchSize
  print(X_train.shape[0])
  print('Epochs:', epochs)
  print('Batch size:', batchSize)
  print('Batches per epoch:', batchCount)

  adam = get_optimizer()
  generator = get_generator()
  discriminator = get_discriminator()
  # 判别器统一更换为Adam优化器
  discriminator.compile(loss='binary_crossentropy', optimizer=adam)
  gan = get_gan_network(discriminator, random_dim, generator, adam)

  for epoch in range(epochs):
    for _ in tqdm(range(batchCount)):
      # 1. 训练判别器
      noise = np.random.normal(0, 1, size=[batchSize, random_dim])
      imageBatch = X_train[np.random.randint(0, X_train.shape[0], size=batchSize)]
      generatedImages = generator.predict(noise, verbose=0)
      
      # 单侧标签平滑
      y_real = np.ones(batchSize) * 0.9
      y_fake = np.zeros(batchSize)
      
      # 分别训练真实样本和假样本
      d_loss_real = discriminator.train_on_batch(imageBatch, y_real)
      d_loss_fake = discriminator.train_on_batch(generatedImages, y_fake)
      d_loss = 0.5 * np.add(d_loss_real, d_loss_fake)

      # 2. 固定判别器训练生成器
      noise = np.random.normal(0, 1, size=[batchSize, random_dim])
      y_gen = np.ones(batchSize)
      discriminator.trainable = False
      g_loss = gan.train_on_batch(noise, y_gen)
      discriminator.trainable = True
    
    # 每轮结束打印损失、保存图片
    print(f"Epoch {epoch+1} | 判别器损失: {d_loss:.4f} | 生成器损失: {g_loss:.4f}")
    save_images(epoch+1, fixed_noise, generator)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 04:18:04