【问题标题】:How to initialize tensorflow variable that wasn't saved other than with tf.global_variables_initializer()如何初始化除 tf.global_variables_initializer() 之外未保存的张量流变量
【发布时间】:2017-10-30 07:18:24
【问题描述】:

我正在研究如何在 tensorflow 中保存/加载特定变量。

我可以毫无问题地加载和保存特定变量,但是,我不知道如何在不使用的情况下初始化剩余的未保存变量

sess.run(tf.global_variables_initializer())   

然后用以下代码覆盖保存的变量:

new_saver.restore(sess,'my_test_model2')

这可以正常工作并初始化未保存的变量 (w2) 并恢复已保存的变量 (w1),但看起来非常笨拙和不自然。

我想知道如何摆脱

tf.global_variables_initializer()

,在我将 w1 变量恢复为 pythonic 的最后。

我尝试了sess.run(tf.variables_initializer([w2])) 并得到了输入:“^w2/Assign”不是该图表的元素。)

我也试过sess.run(tf.variables_initializer(["w2:0"])) 并得到 AttributeError: 'str' object has no attribute 'initializer' 将张量流导入为 tf

print(tf.__version__)
w1 = tf.Variable(tf.linspace(0.0, 0.5, 6), name="w1")
w2 = tf.Variable(tf.linspace(1.0, 5.0, 6), name="w2")
saver = tf.train.Saver({'w1':w1})
sess = tf.Session()
sess.run(tf.global_variables_initializer())
for v in tf.global_variables():
      print (v.name)

print(sess.run(["w1:0"]))
print(sess.run(["w2:0"]))

saver.save(sess, 'my_test_model')

tf.reset_default_graph()

print ('-'*80 )

w1 = tf.Variable(tf.linspace(10.0, 50.0, 6), name="w1")
w2 = tf.Variable(tf.linspace(100.0, 500.0, 6), name="w2")
saver = tf.train.Saver({'w1':w1})
sess = tf.Session()
sess.run(tf.global_variables_initializer())
for v in tf.global_variables():
      print (v.name)

print(sess.run(["w1:0"]))
print(sess.run(["w2:0"]))

saver.save(sess, 'my_test_model2')  

tf.reset_default_graph()

print ('-'*80 )
print("Let's load w1 \n")  

with tf.Session() as sess:
  # Loading the model structure from 'my_test_model.meta'
  new_saver = tf.train.import_meta_graph('my_test_model.meta')
  # I do this to make sure w1:0 and w2:0 are variables
  for v in tf.global_variables():
        print (v.name)  
  sess.run(tf.global_variables_initializer())  #<----- line I want to make more pythonic
#   sess.run(tf.variables_initializer([w2]))  # input: "^w2/Assign" is not an element of this graph.)
#   sess.run(tf.variables_initializer(["w2:0"])) #AttributeError: 'str' object has no attribute 'initializer'

# Loading the saved "w1" Variable
  new_saver.restore(sess,'my_test_model2')

  print(sess.run(["w1:0"]))
  print(sess.run(["w2:0"]))    

【问题讨论】:

    标签: python tensorflow save


    【解决方案1】:

    终于看完了:

    In TensorFlow is there any way to just initialize uninitialised variables?

    我喜欢https://stackoverflow.com/users/1090562/salvador-dali 的答案并将其修改为使用itertools.compress,如果变量不只少数,这会更快。

    def initialize_uninitialized_vars(sess):
        from itertools import compress
        global_vars = tf.global_variables()
        is_not_initialized = sess.run([~(tf.is_variable_initialized(var)) \
                                       for var in global_vars])
        not_initialized_vars = list(compress(global_vars, is_not_initialized))
    
        if len(not_initialized_vars):
            sess.run(tf.variables_initializer(not_initialized_vars))
    

    然后我的代码变成:

    with tf.Session() as sess:
      # Loading the model structure from 'my_test_model.meta'
      new_saver = tf.train.import_meta_graph('my_test_model.meta')  
    
      # Loading the saved "w1" Variable
      new_saver.restore(sess,'my_test_model2')
    
      # initialize the unitialized variables
      initialize_uninitialized_vars(sess)
    
      print(sess.run(["w1:0"]))
      print(sess.run(["w2:0"]))  
    

    【讨论】:

      猜你喜欢
      • 2017-12-10
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2022-06-15
      • 1970-01-01
      相关资源
      最近更新 更多