【tensorflow】库函数,tf.contrib.layers.batch

mac2026-08-16  3

tf.contrib.layers.batch_norm(     inputs,     decay=0.999,     center=True,     scale=False,     epsilon=0.001,     activation_fn=None,     param_initializers=None,     param_regularizers=None,     updates_collections=tf.GraphKeys.UPDATE_OPS,     is_training=True,     reuse=None,     variables_collections=None,     outputs_collections=None,     trainable=True,     batch_weights=None,     fused=None,     data_format=DATA_FORMAT_NHWC,     zero_debias_moving_mean=False,     scope=None,     renorm=False,     renorm_clipping=None,     renorm_decay=0.99,     adjustment=None )

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设置为强制更新,但可能会导致速度损失,尤其是在分布式设置中。

 

最新回复(0)