【发布时间】:2017-01-31 04:27:08
【问题描述】:
我有不同的作用域,它们有相同名称但具有不同值的变量。我想在范围之间交换这些变量的值。 示例:
with tf.variable_scope('sc1'):
a1 = tf.Variable(0, name='test_var1')
b1 = tf.Variable(1, name='test_var2')
with tf.variable_scope('sc2'):
a2 = tf.Variable(2, name='test_var1')
b2 = tf.Variable(3, name='test_var2')
我想将a2 设置为 0,b2 设置为 1,a1 设置为 2,b1 设置为 3。
我正在考虑使用tf.get_collection_ref 获取所需的变量,但我看不到如何更改变量的范围,所以可能我需要更改变量的值。在这种情况下,我需要在临时变量中存储一个值,然后删除该临时变量。
我不确定它会起作用,这似乎太复杂了。
有简单的方法吗?
UPD1:我还需要从另一个集合中设置一个集合中的所有变量。我认为这是类似的问题。
例如,在上面的代码中,将a2 设置为 0,将b2 设置为 1。
UPD2:此代码不起作用:
with tf.variable_scope('sc1'):
a1 = tf.get_variable(name='test_var1', initializer=0.)
b1 = tf.Variable(0, name='test_var2')
with tf.variable_scope('sc2'):
a2 = tf.get_variable(name='test_var1', initializer=1.)
b2 = tf.Variable(1, name='test_var2')
def swap_tf_scopes(col1, col2):
col1_dict = {}
col2_dict = {}
for curr_var in col1:
curr_var_name = curr_var.name.split('/')[-1]
col1_dict[curr_var_name] = curr_var
for curr_var in col2:
curr_var_name = curr_var.name.split('/')[-1]
curr_col1_var = col1_dict[curr_var_name]
tmp_t = tf.identity(curr_col1_var)
assign1 = curr_col1_var.assign(curr_var)
assign2 = curr_var.assign(tmp_t)
return [assign1, assign2]
col1 = tf.get_collection(tf.GraphKeys.VARIABLES, scope='sc1')
col2 = tf.get_collection(tf.GraphKeys.VARIABLES, scope='sc2')
tf_ops_t = swap_tf_collections(col1, col2)
sess = tf.Session()
sess.run(tf.initialize_all_variables())
sess.run(tf_ops_t)
print sess.run(col1) #prints [0.0, 1] but I expect [1.0, 1]
print sess.run(col2) #prints [1.0, 1] but I expect [0.0, 0]
【问题讨论】:
-
一个猜测,基于快速阅读:写入一个变量正在覆盖另一个写入的输入。
tf.identity不足以强制复制张量数据。试试tmp_t = curr_col1_var + 0.0之类的东西。希望有帮助!
标签: python scope tensorflow