TensorFlow中GAN的Batch Normalization参数更新方法问询
关于TensorFlow中GAN集成Batch Normalization的参数更新问题解答
你好!针对你在使用TensorFlow 1.12实现带Batch Normalization(BN)的GAN时遇到的几个疑问,我结合你的代码详细解答如下:
一、TensorFlow中如何更新BN的移动均值、移动方差?
在TensorFlow 1.x的tf.layers.batch_normalization中,BN的移动均值(moving mean)和移动方差(moving variance)属于非可训练变量,它们不会被optimizer.minimize()自动更新。这些参数的更新逻辑是:
- 当
training=True时,层会计算当前batch的均值和方差,同时生成更新移动均值、方差的操作,这些操作会被自动添加到tf.GraphKeys.UPDATE_OPS集合中。 - 要让这些更新操作执行,你需要将训练操作(比如Adam的minimize返回的op)与对应的
UPDATE_OPS绑定,通过tf.control_dependencies()实现。
二、判别器和生成器用BN后,如何更新计算图?
不需要手动修改计算图结构,只需要在构建训练操作时,让训练op依赖各自的BN更新操作即可。因为生成器和判别器的BN参数是独立的(分别在Gen和Dis的variable_scope下),所以必须分开获取各自的UPDATE_OPS,避免训练时互相干扰。
三、BN参数会自动更新吗?
不会自动更新!因为移动均值和方差不属于tf.trainable_variables()(你代码里的t_vars),optimizer.minimize()只会更新可训练变量(比如生成器和判别器的权重w、偏置b)。BN的移动均值/方差更新操作在UPDATE_OPS里,必须手动绑定到训练流程中才会执行。
结合你的代码的修正方案
你的代码里尝试了三种方法,其中方法2的思路是对的,但没有区分生成器和判别器的UPDATE_OPS,这会导致训练生成器时更新判别器的BN参数,反过来也一样,是错误的。下面是修正后的代码:
步骤1:分别获取生成器、判别器的BN更新操作
在构建完生成器和判别器之后,添加以下代码:
# 按变量作用域分别获取BN更新操作 gen_update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS, scope="Gen") dis_update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS, scope="Dis") pre_update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS, scope="Pre") # 如果预训练判别器也用了BN
步骤2:将训练操作与对应更新操作绑定
替换你原来的优化器代码为:
# 预训练(如果预训练模块有BN的话) with tf.control_dependencies(pre_update_ops): p_train = tf.train.AdamOptimizer(learning_rate=learningrate_dis, beta1=0.5, beta2=0.999).minimize(p_loss, var_list=p_vars) # 生成器训练:仅依赖生成器的BN更新操作 with tf.control_dependencies(gen_update_ops): g_train = tf.train.AdamOptimizer(learning_rate=learningrate_gen, beta1=0.5, beta2=0.999).minimize(g_loss, var_list=g_vars) # 判别器训练:仅依赖判别器的BN更新操作 with tf.control_dependencies(dis_update_ops): d_train = tf.train.AdamOptimizer(learning_rate=learningrate_dis, beta1=0.5, beta2=0.999).minimize(d_loss, var_list=d_vars)
额外注意事项
- 训练与测试的training参数区分:训练时
feed_dict中trainingState要传True,生成样本/测试判别器时传False,这样测试时会使用已更新好的移动均值和方差做归一化。 - momentum参数调整:
tf.layers.batch_normalization的momentum参数默认是0.99,控制移动均值的更新速度——值越接近1,更新越慢,越依赖历史batch的统计信息,你可以根据训练情况调整。 - 参数共享的正确性:你的判别器第二次调用时设置了
reuse=True,这个是正确的,保证判别器在判别真实样本和生成样本时共享同一套参数。
内容的提问来源于stack exchange,提问作者Heewony
相关产品推荐
相关产品推荐

