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

TensorFlow中tf.GraphKeys.TRAINABLE_VARIABLES与UPDATE_OPS的区别解析

刚好对这两个GraphKeys的区别比较熟悉,结合你提到的batch norm场景给你拆解清楚:

tf.GraphKeys.TRAINABLE_VARIABLES vs tf.GraphKeys.UPDATE_OPS:核心差异与实际场景

1. TRAINABLE_VARIABLES:优化器的「直接操作对象」

  • 这个集合里装的是会被梯度下降等优化算法直接更新的可训练变量,比如神经网络的权重矩阵、偏置向量这些核心参数。
  • 当你调用优化器的minimize()方法时,它默认只对这个集合里的变量计算梯度并更新。你可以用tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES)获取所有这类变量,比如用来做「冻结部分层」的操作——把不想训练的变量从这个集合里移除,优化器就不会碰它们了。

2. UPDATE_OPS:维护性操作的「收纳箱」

  • 这个集合专门存放不需要梯度更新的辅助操作,这些操作是按照预设规则执行的,和梯度下降无关。最典型的例子就是批量归一化(batch norm)里的滑动平均均值、滑动方差的更新,还有指数移动平均(EMA)的更新步骤。
  • 这些操作不会被优化器自动执行,必须你手动把它们整合到训练流程里,否则相关的统计量就不会更新,模型的表现会出问题。

结合batch_norm()的实际例子

你提到tensorflow.contrib.layers.batch_norm()的默认参数updates_collections=tf.GraphKeys.UPDATE_OPS,意思是:

  • 这个batch norm层在计算过程中,会生成更新滑动均值和方差的操作,并把这些操作放到UPDATE_OPS集合里。
  • 如果你的训练代码只跑优化器的train_op,那这些滑动统计量永远不会更新,batch norm层相当于没起作用。正确的做法是把这些更新操作和训练操作绑定:
# 获取所有需要执行的维护更新操作
update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS)
# 确保先执行更新操作,再进行梯度下降
with tf.control_dependencies(update_ops):
    train_op = optimizer.minimize(loss)

这样每次运行train_op时,都会先完成UPDATE_OPS里的所有操作,再更新模型的可训练变量。

额外补充

你说的ops.py文件里定义了所有tf.GraphKeys的本质——它们就是TensorFlow计算图里用来给不同元素分类的「标签」,方便你批量获取和处理同类型的变量或操作,避免一个个手动管理的麻烦。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 08:14:50