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

TensorFlow中tf.Variable的trainable属性含义及相关疑问

关于TensorFlow中trainable参数与global_step的疑问解答

嘿,这个问题其实挺关键的,很多刚接触TensorFlow的同学都会对trainable参数和global_step的关系有疑惑,我来给你掰扯清楚~

先直接回答你的核心问题

1. 不可训练变量的值能不能在sess.run()期间被修改?反之呢?

完全可以!trainable=False这个参数根本不是用来限制你手动修改变量值的——它只是告诉TensorFlow:「这个变量不属于模型的可训练参数集合,别让优化器碰它」。

比如你的示例里,不管global_step_tensor的trainable是True还是False,你都可以通过tf.assign手动更新它的值,执行sess.run(assign_op)就能生效。反过来,trainable=True的变量,除了会被优化器自动更新(反向传播时调整值),你也一样可以手动修改它,两者不冲突。

你之前测试时觉得结果一致,是因为你的示例里没有用到优化器——只有当优化器介入时,trainable的差异才会显现出来。

2. 变量被设为不可训练的意义是什么?

结合global_step的场景,这个参数的意义主要有三点:

  • 避免被优化器误更新:这是最核心的原因。global_step的作用是计数训练步数,我们期望它每次执行训练操作后固定加1,而不是被梯度反向传播修改。如果把它设为trainable=True,优化器会把它当成模型参数,在计算梯度时试图调整它的值(比如你如果不小心把它放进损失函数里,它会被优化器“训练”得偏离计数逻辑),这完全违背了我们用它的初衷。
  • 减少不必要的计算开销:trainable=True的变量会被加入到GraphKeys.TRAINABLE_VARIABLES集合中,优化器会遍历这个集合计算每个变量的梯度。把不需要训练的变量设为False,能避免无意义的梯度计算,节省内存和计算资源。
  • 语义更清晰:明确告诉其他阅读代码的人(包括未来的你):这个变量是用来追踪训练状态的计数器,不是模型需要学习的参数,提升代码的可读性和可维护性。

举个直观的对比例子

错误示范:把global_step设为trainable=True

import tensorflow as tf

global_step_tensor = tf.Variable(10, trainable=True, name='global_step')
optimizer = tf.train.GradientDescentOptimizer(0.01)
# 随便定义一个损失函数,这里故意把global_step放进去
loss = tf.square(global_step_tensor - 20)
# 让优化器最小化损失,同时指定global_step
train_op = optimizer.minimize(loss, global_step=global_step_tensor)

with tf.Session() as sess:
    sess.run(tf.global_variables_initializer())
    print(f"初始global_step: {sess.run(global_step_tensor)}")  # 输出10
    sess.run(train_op)
    print(f"执行train_op后global_step: {sess.run(global_step_tensor)}")  # 输出10.1(被优化器梯度更新了,完全不是我们要的步数计数!)

正确示范:把global_step设为trainable=False

import tensorflow as tf

global_step_tensor = tf.Variable(10, trainable=False, name='global_step')
optimizer = tf.train.GradientDescentOptimizer(0.01)
# 损失函数用真正的模型参数,和global_step无关
model_var = tf.Variable(5.0)
loss = tf.square(model_var - 10.0)
# 优化器只更新模型参数,global_step负责计数
train_op = optimizer.minimize(loss, global_step=global_step_tensor)

with tf.Session() as sess:
    sess.run(tf.global_variables_initializer())
    print(f"初始global_step: {sess.run(global_step_tensor)}")  # 输出10
    sess.run(train_op)
    print(f"执行train_op后global_step: {sess.run(global_step_tensor)}")  # 输出11,符合步数计数的预期

总结一下

trainable参数的本质是标记变量是否属于模型的可训练参数,和你能不能手动修改它的值没有关系。global_step必须设为trainable=False,就是为了防止优化器误操作它,同时明确它的语义、减少不必要的计算开销。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:30:56