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

如何基于TensorFlow CNN的Checkpoint与Meta文件实现图像预测

问题分析

你现在打印出的数值是模型最后一层的原始输出(logits),不是最终的预测结果。如果是分类任务,这类输出是未经过归一化的得分,需要进一步处理才能得到类别概率或具体类别;如果是回归任务,可能需要根据训练时的归一化操作做反归一化。

解决方案

下面分两种场景给出修改方案,你可以根据自己的需求选择:


方案1:直接从Checkpoint恢复模型并预测(无需转PB)

转PB文件其实是多余的步骤,直接从meta和checkpoint恢复后就能直接做预测,流程更简洁:

import tensorflow as tf

# 假设test_images已经是预处理好的输入数据
with tf.Session() as sess:
    # 加载模型结构和权重
    saver = tf.train.import_meta_graph('./tmp1/my_model.meta', clear_devices=True)
    saver.restore(sess, "./tmp1/my_model")
    
    # 获取输入和原始输出tensor(注意名称要和你训练时的一致)
    input_tensor = sess.graph.get_tensor_by_name('input:0')
    logits_tensor = sess.graph.get_tensor_by_name('output:0')
    
    # 根据任务类型处理输出:
    # 情况A:分类任务 - 获取类别概率(softmax归一化)
    prob_tensor = tf.nn.softmax(logits_tensor)
    # 情况B:分类任务 - 获取直接的类别索引(取得分最高的类别)
    pred_class_tensor = tf.argmax(logits_tensor, axis=1)
    # 情况C:回归任务 - 如果训练时对标签做了归一化,这里需要反归一化(比如乘以标准差加均值)
    # pred_values = logits_tensor * std + mean  # 替换成你训练时的归一化参数
    
    # 运行得到预测结果
    pred_probs = sess.run(prob_tensor, feed_dict={input_tensor: test_images})
    pred_classes = sess.run(pred_class_tensor, feed_dict={input_tensor: test_images})
    
    print("每个样本的类别概率:", pred_probs)
    print("每个样本的预测类别:", pred_classes)

方案2:基于已生成的PB文件做预测

如果你坚持要用PB文件,需要注意两个关键点:

  1. 加载PB后不需要执行sess.run(tf.global_variables_initializer()),因为PB文件已经包含了训练好的权重,初始化会重置变量,反而可能导致结果异常。
  2. 对原始输出(logits)做后处理得到最终结果。

修改后的代码:

import tensorflow as tf

path="./my_model.pb"
def load_pb(path):
    with tf.gfile.GFile(path, "rb") as f:
        graph_def = tf.GraphDef()
        graph_def.ParseFromString(f.read())
    with tf.Graph().as_default() as graph:
        tf.import_graph_def(graph_def, name='')
    return graph

graph=load_pb(path)
with tf.Session(graph=graph) as sess:
    # 删掉多余的初始化步骤
    # sess.run(tf.global_variables_initializer())
    
    input_tensor = graph.get_tensor_by_name('input:0')
    logits_tensor = graph.get_tensor_by_name('output:0')
    
    # 同样根据任务处理输出
    pred_probs = sess.run(tf.nn.softmax(logits_tensor), feed_dict={input_tensor: test_images})
    pred_classes = sess.run(tf.argmax(logits_tensor, axis=1), feed_dict={input_tensor: test_images})
    
    print("预测类别概率:", pred_probs)
    print("预测类别索引:", pred_classes)

关键提示

  • 确保input:0和output:0的名称和你训练模型时定义的输入、输出tensor名称完全一致,如果不一致,可以用for op in graph.get_operations(): print(op.name)查看所有节点名称,替换成正确的标识。
  • 如果是回归任务,记得根据训练时对标签的预处理(比如归一化、标准化)做反向操作,才能得到真实的预测值。

内容的提问来源于stack exchange,提问作者Simran

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 05:21:42