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

为何我的GAN生成噪声而非清晰图像?附代码求解决方案

GAN生成噪声的调整方案

以下是针对代码问题的具体修复与优化建议:

1. 降低优化器学习率并调整Adam参数

WGAN对学习率极度敏感,当前设置的1e-3过高,会导致梯度不稳定、模型无法收敛。建议将学习率降至1e-4,同时调整Adam的beta1参数为0.5(WGAN通用稳定参数),并暂时移除混合精度优化(简化训练流程,避免额外精度问题):

optimizer = Adam(lr=1e-4, beta_1=0.5, beta_2=0.9)
# 注释掉混合精度相关代码
# optimizer = mixed_precision.LossScaleOptimizer(optimizer)

2. 修复联合模型的训练逻辑

你的gan函数中错误地在添加判别器后重新开启了其可训练性,导致训练生成器时判别器权重也被更新,违反了WGAN“固定判别器训练生成器”的核心逻辑。正确实现如下:

def gan(generator, discriminator):
    discriminator.trainable = False  # 固定判别器权重,仅训练生成器
    model = Sequential()
    model.add(generator)
    model.add(discriminator)
    return model

训练判别器时,判别器默认处于可训练状态(初始化时已设置),无需额外修改。

3. 增加判别器训练次数

WGAN需要让判别器充分训练至接近最优,再更新生成器。当前每轮仅训练判别器1次,建议改为训练判别器5次,再训练生成器1次:
修改train函数:

def train(x_train, epochs, batch_size):
    with tqdm(total=epochs,unit="iters") as pbar:
        for epoch in range(epochs):
            # 重复训练判别器5次
            d_loss_total = [0.0, 0.0]
            d_valid_total = [0.0, 0.0]
            d_fake_total = [0.0, 0.0]
            for _ in range(5):
                d_loss, d_loss_valid, d_loss_fake = train_discriminator(x_train, batch_size)
                d_loss_total = np.add(d_loss_total, d_loss)
                d_valid_total = np.add(d_valid_total, d_loss_valid)
                d_fake_total = np.add(d_fake_total, d_loss_fake)
            
            # 计算平均损失
            d_loss = np.divide(d_loss_total, 5)
            d_loss_valid = np.divide(d_valid_total, 5)
            d_loss_fake = np.divide(d_fake_total, 5)
            
            # 训练生成器1次
            g_loss = train_generator(batch_size)
            
            pbar.set_description('Epoch {0}'.format(epoch+1))
            pbar.set_postfix(
                d_loss='{0:.4f}'.format(d_loss[0]),
                g_loss='{0:.4f}'.format(g_loss[0]),
                d_loss_valid='{0:.4f}'.format(d_loss_valid[0]),
                d_loss_fake='{0:.4f}'.format(d_loss_fake[0]),
            )
            pbar.update(1)

4. 确保训练数据归一化匹配生成器输出

生成器最后一层使用sigmoid激活,输出范围为[0,1],因此训练数据必须归一化到同一区间:

# 假设x_train是原始MNIST像素数据(0-255)
x_train = x_train.astype('float32') / 255.0
# 增加通道维度(适配判别器输入)
x_train = np.expand_dims(x_train, axis=-1)

5. 移除无用的生成器单独编译

生成器的训练完全依赖联合模型(gan),单独编译生成器没有实际作用,可直接删除:

# 注释或删除该行
# g.compile(loss=wasserstein_loss, optimizer=optimizer, metrics=['accuracy'])

6. 可选:调整生成器输出激活函数

若想进一步提升稳定性,可将生成器最后一层的sigmoid替换为tanh,同时将训练数据归一化到[-1,1]区间:
修改生成器:

model.add(Conv2D(1, (7,7), activation='tanh', padding='same'))

修改数据预处理:

x_train = (x_train.astype('float32') - 127.5) / 127.5  # 归一化到[-1,1]
x_train = np.expand_dims(x_train, axis=-1)

7. 可选:增加训练轮数

1000轮对于WGAN来说可能不足以收敛,建议增加至5000-10000轮,观察生成图像的逐步变化。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 21:21:59