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

TensorFlow如何确定需进行最小化更新的可训练变量?

关于TensorFlow GradientDescentOptimizer.minimize()的变量更新逻辑

是的,你观察到的现象完全正确——TensorFlow确实会自动识别并仅更新那些与损失函数计算相关的可训练变量,不会盲目更新所有TRAINABLE_VARIABLES集合里的变量。

背后的原理

minimize()方法本质上是compute_gradients()和apply_gradients()的组合:

  1. compute_gradients()的作用:它会沿着损失函数的计算图进行反向传播分析,只针对那些在损失值计算路径中出现的可训练变量计算梯度。如果某个可训练变量完全没有参与损失函数的计算(也就是损失值不会随该变量变化而变化),那么compute_gradients()不会为它生成梯度张量。
  2. 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 call compute_gradients() and apply_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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 06:48:38