如何获取函数内使用的所有tf.variables?TensorFlow变量提取咨询
获取TensorFlow函数中使用的全部变量
当然可以!TensorFlow提供了几种实用的方法来提取任意函数中用到的所有tf.Variable,而且这些方法的思路和优化器内部的实现逻辑是一致的——毕竟优化器就是靠追踪变量来计算梯度并更新的。下面给你详细介绍两种最常用的方式:
方法一:用tf.GradientTape追踪变量
这是最贴近优化器工作流程的方式,因为优化器计算梯度时就是通过GradientTape来记录所有参与运算的变量的。只要在GradientTape的上下文环境中执行目标函数,就能轻松获取到函数用到的所有变量:
import tensorflow as tf # 先定义几个外部变量 var1 = tf.Variable(1.0, name="trainable_var1") var2 = tf.Variable(2.0, name="trainable_var2") non_trainable_var = tf.Variable(3.0, trainable=False, name="non_trainable") # 目标函数:用到了上面的三个变量 def my_custom_func(x): return var1 * x + var2 * x**2 + non_trainable_var # 使用GradientTape追踪变量 with tf.GradientTape() as tape: # 在tape上下文内执行函数,所有参与运算的变量都会被自动追踪 func_output = my_custom_func(5.0) # 获取所有被追踪的变量 used_variables = tape.watched_variables() # 打印变量名称验证 print([var.name for var in used_variables]) # 输出: ['trainable_var1:0', 'trainable_var2:0', 'non_trainable:0']
如果函数内部会创建新的变量,只要确保在GradientTape的作用域内调用函数,内部创建的变量也会被自动追踪到:
def func_with_internal_vars(x): # 函数内部定义的变量 internal_var = tf.Variable(4.0, name="internal_variable") return internal_var * x with tf.GradientTape() as tape: output = func_with_internal_vars(3.0) used_vars = tape.watched_variables() print([var.name for var in used_vars]) # 输出: ['internal_variable:0']
方法二:通过tf.function的图结构提取变量
如果你的函数是用tf.function装饰的(也就是转为了计算图模式),可以直接从它的具体函数实例中提取所有被捕获的变量:
@tf.function def my_graph_func(x): return var1 * x + var2 * x**2 + non_trainable_var # 获取函数的具体实现实例 concrete_func = my_graph_func.get_concrete_function(5.0) # 从计算图中提取所有变量 captured_variables = concrete_func.graph.variables print([var.name for var in captured_variables]) # 输出和第一种方法一致
补充说明
优化器之所以能自动找到要更新的变量,本质上就是在计算损失函数的梯度时,用GradientTape追踪了所有和损失相关的可训练变量,然后基于这些梯度来更新变量——和我们第一种方法的逻辑完全一致。
内容的提问来源于stack exchange,提问作者Kristian Wichmann
相关产品推荐
相关产品推荐

