You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 07:15:41