TensorFlow 1.0:如何从恢复的RNN中获取隐藏状态
恢复RNN并获取隐藏状态的解决方案
要恢复你的LSTM模型并获取隐藏状态,关键在于准确定位图中的state张量,并在恢复权重后运行它得到具体值。下面分两种场景给出具体步骤:
一、最优方案:修改保存代码给state命名(推荐)
如果还能修改原始的保存代码,给state显式添加名字是最省心的方式,避免后续找张量名字的麻烦:
修改后的保存代码
loc="path/to/save/rnn" with tf.variable_scope("lstm") as scope: outputs, state = tf.nn.dynamic_rnn( cell=your_lstm_cell, # 替换成你的LSTM Cell实例 inputs=your_inputs, # 替换成你的输入张量 dtype=tf.float32 ) # 给隐藏状态加一个明确的名字,方便后续恢复 state = tf.identity(state, name="lstm_hidden_state") sess = tf.Session() sess.run(tf.global_variables_initializer()) saver = tf.train.Saver() save_path = saver.save(sess, loc)
对应的恢复代码
graph = tf.Graph() with graph.as_default(): sess = tf.Session(graph=graph) # 导入元图并恢复权重 saver = tf.train.import_meta_graph(loc + '.meta', clear_devices=True) saver.restore(sess, loc) # 这一步不能少,要加载保存的权重值 # 通过名字获取隐藏状态张量 # 如果是LSTM,state通常是LSTMStateTuple(c_state, h_state),对应两个张量 state_tensor = graph.get_tensor_by_name("lstm_hidden_state:0") # 同时需要获取输入的placeholder(确保原始代码给它加了名字) input_tensor = graph.get_tensor_by_name("your_input_placeholder:0") # 替换成你输入占位符的名字 # 准备符合模型要求的输入数据 input_data = ... # 比如形状为[batch_size, seq_len, input_dim]的numpy数组 # 运行得到隐藏状态的值 hidden_state = sess.run(state_tensor, feed_dict={input_tensor: input_data}) print("获取到的隐藏状态:", hidden_state)
二、如果无法修改原始保存代码:查找state的默认名字
如果已经不能修改保存代码,就得先找到state在图中的默认名称:
步骤1:遍历图中所有张量,找到state的名字
graph = tf.Graph() with graph.as_default(): sess = tf.Session(graph=graph) saver = tf.train.import_meta_graph(loc + '.meta') saver.restore(sess, loc) # 打印所有操作的名字,从中筛选出和state相关的 print("所有张量操作名称:") for op in graph.get_operations(): print(op.name)
运行后,你会看到类似lstm/rnn/while/Exit_2和lstm/rnn/while/Exit_3的名字(对应LSTM的细胞状态c和隐藏状态h),或者包含state关键词的名称。
步骤2:通过找到的名字获取并运行state
比如找到的名字是lstm/rnn/while/Exit_2:0(c_state)和lstm/rnn/while/Exit_3:0(h_state),就可以这样获取:
# 获取细胞状态和隐藏状态 c_state = graph.get_tensor_by_name("lstm/rnn/while/Exit_2:0") h_state = graph.get_tensor_by_name("lstm/rnn/while/Exit_3:0") # 同样需要喂入输入数据 input_tensor = graph.get_tensor_by_name("your_input_placeholder_name:0") input_data = ... c_val, h_val = sess.run([c_state, h_state], feed_dict={input_tensor: input_data}) print("细胞状态:", c_val) print("隐藏状态:", h_val)
关键注意点
- 必须执行
saver.restore(sess, loc)才能加载保存的权重参数,否则只是导入了图结构,权重还是初始化的值。 - 隐藏状态
state是计算张量,不是可训练变量,所以必须通过sess.run()并喂入输入数据才能得到具体值,不能直接通过get_variable()获取。 - 如果你的输入占位符没有命名,同样需要用上面的遍历方法找到它的名字。
内容的提问来源于stack exchange,提问作者mehh
相关产品推荐
相关产品推荐

