Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift
参数含义:
inputs: 输入
decay:衰减系数。合适的衰减系数值接近1.0,特别是含多个9的值:0.999,0.99,0.9。如果训练集表现很好而验证/测试集表现的不好,选择小的系数(推荐0.9)。如果想要提高稳定性,zero_debias_moving_mean设为True
center:如果为Ture,有beta偏移量;如果为False,无beta偏移量
scale:如果为True,则乘以gamma。如果为False,gamma则不用。当下一层是线性时(例如nn.relu),由于缩放可以由下一层完成,所以可以禁用该层。
epsilon:避免被0除。
activation_fn:用于激活,默认为线性激活函数。
param_initializers:beta, gamma, moving mean and moving variance的优化初始化。
param_regularizers:beta和gamma正则化优化。
updates_collections:Collections来收集计算的更新操作。updates_ops需要使用train_op来执行。如果为None,则会添加控件依赖项以确保更新已计算到位。
is_training:图层是否处于训练模式。在训练模式下,它将积累转入的统计量moving_mean并moving_variance使用给定的指数移动平均值delay。当它不是训练模式,那么它将使用的数值moving_mean和moving_variance。
scope:可选范围variable_scope。
注意:训练时,需要更新moving_mean和moving_variance。默认情况下,更新操作被放入tf.GraphKeys.UPDATE_OPS,所以需要添加它们作为依赖项train_op。
例如:
update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS) with tf.control_dependencies(update_ops):train_op = optimizer.minimize(loss)可以将updates_collections = None设置为强制更新,但可能会导致速度损失,尤其是在分布式设置中。
