如何基于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文件,需要注意两个关键点:
- 加载PB后不需要执行
sess.run(tf.global_variables_initializer()),因为PB文件已经包含了训练好的权重,初始化会重置变量,反而可能导致结果异常。 - 对原始输出(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
相关产品推荐
相关产品推荐

