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

TensorFlow中如何收集GAN生成器与判别器的可训练变量?

解决GAN中无法正确收集生成器/判别器变量导致训练失效的问题

我帮你排查出了几个导致模型训练失效的核心问题,修正后就能正常训练出清晰的生成结果了,下面是具体分析和修正后的完整代码:

问题分析

  1. 网络结构缺失关键计算层
    你的生成器和判别器函数里,注释掉了G_h1和D_h1的定义,但后续代码却直接使用了这些变量,这导致网络完全没有有效的特征变换逻辑,生成器只能输出随机初始化后的噪声。

  2. 训练循环未执行优化操作
    你的训练循环里只实现了采样保存图片的逻辑,完全没有调用D_solver和G_solver来更新判别器和生成器的参数——相当于模型根本没在训练。

  3. 图片保存路径错误
    plt.savefig('out/')没有指定具体文件名,会触发IO错误,需要添加动态编号的文件名。

  4. 变量收集的小误区
    你最初用tf.name_scope收集变量失败,是因为name_scope仅给操作(op)添加命名前缀,不会影响变量的命名空间;而tf.variable_scope会同时标记操作和变量的归属,所以你后来用tf.trainable_variables()结合前缀筛选的方式是完全正确的。

修正后的完整代码

import tensorflow as tf
from tensorflow.examples.tutorials.mnist import input_data
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.gridspec as gridspec
import os

def xavier_init(size):
    in_dim = size[0]
    xavier_stddev = 1. / tf.sqrt(in_dim / 2.)
    return tf.random_normal(shape=size, stddev=xavier_stddev)

X = tf.placeholder(tf.float32, shape=[None, 784])
Z = tf.placeholder(tf.float32, shape=[None, 100])

def sample_Z(m, n):
    return np.random.uniform(-1., 1., size=[m, n])

def generator(z, reuse=False):
    with tf.variable_scope('generator', reuse=reuse):
        G_W1 = tf.Variable(xavier_init([100, 128]))
        G_b1 = tf.Variable(tf.zeros(shape=[128]))
        G_W2 = tf.Variable(xavier_init([128, 784]))
        G_b2 = tf.Variable(tf.zeros(shape=[784]))
        # 恢复关键计算层
        G_h1 = tf.nn.relu(tf.matmul(z, G_W1) + G_b1)
        G_log_prob = tf.matmul(G_h1, G_W2) + G_b2
        G_prob = tf.nn.sigmoid(G_log_prob)
        return G_prob

def discriminator(x, reuse=False):
    with tf.variable_scope('discriminator', reuse=reuse):
        D_W1 = tf.Variable(xavier_init([784, 128]))
        D_b1 = tf.Variable(tf.zeros(shape=[128]))
        D_W2 = tf.Variable(xavier_init([128, 1]))
        D_b2 = tf.Variable(tf.zeros(shape=[1]))
        # 恢复关键计算层
        D_h1 = tf.nn.relu(tf.matmul(x, D_W1) + D_b1)
        D_logit = tf.matmul(D_h1, D_W2) + D_b2
        D_prob = tf.nn.sigmoid(D_logit)
        return D_prob, D_logit

def plot(samples):
    fig = plt.figure(figsize=(4, 4))
    gs = gridspec.GridSpec(4, 4)
    gs.update(wspace=0.05, hspace=0.05)
    for i, sample in enumerate(samples):
        ax = plt.subplot(gs[i])
        plt.axis('off')
        ax.set_xticklabels([])
        ax.set_yticklabels([])
        ax.set_aspect('equal')
        plt.imshow(sample.reshape(28, 28), cmap='Greys_r')
    return fig

G_sample = generator(Z)
D_real, D_logit_real = discriminator(X)
# 调用判别器时传入reuse=True,共享判别器参数
D_fake, D_logit_fake = discriminator(G_sample, reuse=True)

# 损失函数定义
D_loss_real = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(logits=D_logit_real, labels=tf.ones_like(D_logit_real)))
D_loss_fake = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(logits=D_logit_fake, labels=tf.zeros_like(D_logit_fake)))
D_loss = D_loss_real + D_loss_fake
G_loss = tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(logits=D_logit_fake, labels=tf.ones_like(D_logit_fake)))

# 正确收集判别器和生成器的可训练变量
t_vars = tf.trainable_variables()
theta_D = [var for var in t_vars if var.name.startswith('discriminator')]
theta_G = [var for var in t_vars if var.name.startswith('generator')]

# 定义优化器
D_solver = tf.train.AdamOptimizer().minimize(D_loss, var_list=theta_D)
G_solver = tf.train.AdamOptimizer().minimize(G_loss, var_list=theta_G)

mb_size = 128
Z_dim = 100

mnist = input_data.read_data_sets('../../MNIST_data', one_hot=True)

sess = tf.Session()
sess.run(tf.global_variables_initializer())

if not os.path.exists('out/'):
    os.makedirs('out/')

i = 0

for it in range(1000000):
    # 加载批量数据
    X_mb, _ = mnist.train.next_batch(mb_size)
    
    # 执行判别器和生成器的训练步骤
    _, D_loss_curr = sess.run([D_solver, D_loss], feed_dict={X: X_mb, Z: sample_Z(mb_size, Z_dim)})
    _, G_loss_curr = sess.run([G_solver, G_loss], feed_dict={Z: sample_Z(mb_size, Z_dim)})
    
    # 每1000步保存生成结果并打印损失
    if it % 1000 == 0:
        print('Iter: {}, D loss: {:.4}, G loss: {:.4}'.format(it, D_loss_curr, G_loss_curr))
        samples = sess.run(G_sample, feed_dict={Z: sample_Z(16, Z_dim)})
        fig = plot(samples)
        plt.savefig('out/{}.png'.format(i))
        i += 1
        plt.close(fig)

额外说明

  • 调用判别器处理生成样本时,一定要传入reuse=True,这样才能共享判别器的参数,否则会重新创建一套判别器变量,导致训练逻辑完全错误。
  • 修正后的代码会在训练过程中打印损失值,方便你监控训练状态,同时每1000步保存一张包含16个生成样本的图片,随着训练推进,图片会从噪声逐渐变成清晰的MNIST手写数字。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 07:08:40