TensorFlow如何确定需进行最小化更新的可训练变量?
是的,你观察到的现象完全正确——TensorFlow确实会自动识别并仅更新那些与损失函数计算相关的可训练变量,不会盲目更新所有TRAINABLE_VARIABLES集合里的变量。
背后的原理
minimize()方法本质上是compute_gradients()和apply_gradients()的组合:
compute_gradients()的作用:它会沿着损失函数的计算图进行反向传播分析,只针对那些在损失值计算路径中出现的可训练变量计算梯度。如果某个可训练变量完全没有参与损失函数的计算(也就是损失值不会随该变量变化而变化),那么compute_gradients()不会为它生成梯度张量。apply_gradients()的作用:它只会对有有效梯度的变量执行更新操作。没有梯度的变量自然不会被改动。
结合你的示例验证
你的损失函数是log_x_squared = tf.square(tf.log(x)),只有变量x参与了这个计算过程。哪怕你添加另一个可训练变量(比如y = tf.Variable(5.0)),它不在损失的计算链路里,反向传播时就不会产生关于y的梯度,minimize()也就不会对它做任何更新。
你提到的trainable=False参数是另一种场景:如果某个变量确实参与了损失计算,但你不希望它被优化,就可以在声明变量时设置这个参数,这样它不会被加入TRAINABLE_VARIABLES集合,compute_gradients()会直接忽略它。
补充API文档的佐证
正如你引用的API文档所说:
This method simply combines calls compute_gradients() and
apply_gradients(). If you want to process the gradient before applying them callcompute_gradients()andapply_gradients()explicitly instead of using this function.
如果你手动调用compute_gradients(log_x_squared),会发现返回的梯度列表里只有x对应的梯度值,其他无关可训练变量不会出现在这个列表中——这也直接证明了TensorFlow是基于计算依赖来筛选待更新变量的。
你的代码与输出回顾
你的极简示例清晰展示了这个逻辑:
import tensorflow as tf # Guess 2.5 as a starting point x = tf.Variable(2.5, name='x', dtype=tf.float32) log_x_squared = tf.square(tf.log(x)) optimizer = tf.train.GradientDescentOptimizer(0.5) train = optimizer.minimize(log_x_squared) init = tf.global_variables_initializer() with tf.Session() as session: session.run(init) print("start ", "x:", session.run(x), "log(x)^2:", session.run(log_x_squared)) for step in range(10): session.run(train) print("step", step, "x:", session.run(x), "log(x)^2:", session.run(log_x_squared))
输出中x逐步收敛到1.0(此时log(x)^2为0,达到损失最小值),完全符合预期:
start x: 2.5 log(x)^2: 0.83958876 step 0 x: 2.1334836 log(x)^2: 0.57419443 step 1 x: 1.7783105 log(x)^2: 0.33138883 step 2 x: 1.4545966 log(x)^2: 0.14042155 step 3 x: 1.1969798 log(x)^2: 0.032328587 step 4 x: 1.0467671 log(x)^2: 0.002089082 step 5 x: 1.0031027 log(x)^2: 9.596717e-06 step 6 x: 1.0000144 log(x)^2: 2.0805813e-10 step 7 x: 1.0 log(x)^2: 0.0 step 8 x: 1.0 log(x)^2: 0.0 step 9 x: 1.0 log(x)^2: 0.0
当你查看可训练变量集合时:
for v in tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES): print(v)
得到:
<tf.Variable 'x:0' shape=() dtype=float32_ref>
如果添加其他可训练变量,它会出现在这个集合里,但不会被minimize()更新,因为它和损失函数无关。
内容的提问来源于stack exchange,提问作者Chris Snow

