TensorFlow随机森林分类器加载时操作丢失,如何正确保存复用?
解决TensorFlow随机森林模型加载后找不到
accuracy_op的问题 这个问题我之前处理过,核心原因是你训练阶段定义的accuracy_op只是一个临时计算的张量,既没有被赋予可识别的名称,也没有作为命名操作被保存到模型图中。当你用import_meta_graph加载模型时,图里根本不存在名叫"accuracy_op"的操作,所以才会抛出KeyError。
下面给你两种可行的解决方案,优先推荐第一种,更规范可靠:
方案一:修改训练代码,给关键张量显式命名
在训练时给需要后续复用的占位符、计算张量加上name参数,这样保存模型后就能通过名称精准检索到它们。
1. 调整训练代码中的关键部分
给输入占位符、accuracy_op、infer_op都加上名称:
# 给输入占位符命名,方便加载时获取 X = tf.placeholder(tf.float32, shape=[None, num_features], name="input_features") Y = tf.placeholder(tf.int32, shape=[None], name="input_labels") hparams = tensor_forest.ForestHParams(num_classes=num_classes, num_features=num_features, num_trees=num_trees).fill() forest_graph = tensor_forest.RandomForestGraphs(hparams) train_op = forest_graph.training_graph(X, Y) loss_op = forest_graph.training_loss(X, Y) # 给infer_op命名,后续预测会用到 infer_op, _, _ = forest_graph.inference_graph(X, name="predictions") # 给accuracy_op命名,明确标识这个计算节点 correct_prediction = tf.equal(tf.argmax(infer_op, 1), tf.cast(Y, tf.int64)) accuracy_op = tf.reduce_mean(tf.cast(correct_prediction, tf.float32), name="accuracy_op") # 剩下的初始化、训练、保存代码不变 init_vars = tf.group(tf.global_variables_initializer(), resources.initialize_resources(resources.shared_resources())) with tf.Session() as sess: sess.run(init_vars) saver = tf.train.Saver() for i in range(1, 100): for batch_x, batch_y in render_batch(batch_size): _, l = sess.run([train_op, loss_op], feed_dict={X: batch_x, Y: batch_y}) acc = sess.run(accuracy_op, feed_dict={X: batch_x, Y: batch_y}) print('Step %i, Loss: %f, Acc: %f' % (i, l, acc)) if acc >= 0.87: print("Stopping and saving") save_path = saver.save(sess, model_path) print("Model saved in file: %s" % save_path) break
2. 加载模型并使用命名节点
现在加载模型时,就可以通过get_tensor_by_name获取对应的张量(注意TensorFlow中张量的名称格式是操作名:张量索引,比如你命名的accuracy_op对应的张量名称是accuracy_op:0):
import tensorflow as tf from tensorflow.contrib.tensor_forest.python import tensor_forest model_path = "你的模型保存路径" checkpoint_file = tf.train.latest_checkpoint("./") with tf.Session() as sess: # 加载meta图和模型变量 saver = tf.train.import_meta_graph("{}.meta".format(model_path)) saver.restore(sess, checkpoint_file) # 获取命名的输入占位符和计算张量 X = tf.get_default_graph().get_tensor_by_name("input_features:0") Y = tf.get_default_graph().get_tensor_by_name("input_labels:0") accuracy_op = tf.get_default_graph().get_tensor_by_name("accuracy_op:0") infer_op = tf.get_default_graph().get_tensor_by_name("predictions:0") # 计算测试集准确率 test_acc = sess.run(accuracy_op, feed_dict={X: x_test, Y: y_test}) print(f"测试集准确率: {test_acc}") # 对未见过的数据做预测 unseen_predictions = sess.run(infer_op, feed_dict={X: x_unseen_data}) predicted_classes = tf.argmax(unseen_predictions, 1).eval(session=sess) print(f"预测结果: {predicted_classes}")
方案二:加载模型后重新构建准确率计算
如果不想修改训练代码,也可以在加载模型后,从图中找到必要的节点,重新构建accuracy_op:
with tf.Session() as sess: saver = tf.train.import_meta_graph("{}.meta".format(model_path)) saver.restore(sess, checkpoint_file) # 获取图中的输入占位符和预测张量(需要知道它们在图中的默认名称,或者训练时打印过名称) graph = tf.get_default_graph() # 比如训练时没给X命名,可以通过查找占位符的方式获取,或者打印graph.get_operations()看所有节点名称 X = graph.get_tensor_by_name("Placeholder:0") Y = graph.get_tensor_by_name("Placeholder_1:0") # 找到infer_op,训练时可以通过print(infer_op.name)获取它的名称 infer_op = graph.get_tensor_by_name("inference_graph/ArgMax:0") # 重新构建准确率计算逻辑 correct_prediction = tf.equal(tf.argmax(infer_op, 1), tf.cast(Y, tf.int64)) accuracy_op = tf.reduce_mean(tf.cast(correct_prediction, tf.float32)) test_acc = sess.run(accuracy_op, feed_dict={X: x_test, Y: y_test}) print(f"测试集准确率: {test_acc}")
关键提醒
- 训练时给所有后续需要复用的节点显式命名是最稳妥的方式,避免加载时因为节点名称不明确而踩坑。
- 你之前用
get_operation_by_name("accuracy_op")出错,是因为accuracy_op本质是一个张量(tf.reduce_mean的输出),它对应的操作名称默认是Mean,而不是accuracy_op。所以检索张量要用get_tensor_by_name,并且要加上:0后缀。
内容的提问来源于stack exchange,提问作者Steven
相关产品推荐
相关产品推荐

