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

TensorFlow:如何向导出的.pb模型输入图像并获取分类结果

嘿,我之前刚处理过类似的TensorFlow pb模型推理问题,给你一步步捋清楚怎么操作:

步骤1:先找到模型的输入/输出节点名称

加载pb模型后,你得先搞清楚模型里哪个是输入张量、哪个是输出张量的名称——这是最关键的一步,不然没法喂数据进去。

你可以用这段代码打印模型里所有的节点名称:

for op in graph.get_operations():
    print(op.name)

运行后,你会看到一堆节点名,通常输入节点会叫类似input_1、image_input这类,输出节点会是predictions/Softmax、dense_1/Softmax(对应分类概率的输出)。注意要取张量的完整名称,也就是节点名后面加:0,比如输入张量是input_1:0,输出张量是predictions/Softmax:0。

步骤2:图像预处理(必须和训练时对齐)

模型训练时怎么处理图像,推理时就得完全照搬,不然预测结果会乱掉。举个最常见的预处理流程示例(你要根据自己的训练代码调整):

import cv2
import numpy as np

# 1. 读取图像
img = cv2.imread("your_test_image.jpg")
# 2. 缩放成模型要求的输入尺寸(比如训练时用的224x224)
img = cv2.resize(img, (224, 224))
# 3. 通道转换:cv2默认读的是BGR,如果你训练时用的是RGB,就得转过来
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# 4. 归一化:比如训练时把像素值除以255缩到[0,1]区间
img = img / 255.0
# 5. 增加batch维度:模型一般接受(batch_size, height, width, channels)的输入,所以要加一个维度
img = np.expand_dims(img, axis=0)
步骤3:运行推理获取结果

现在就可以把预处理好的图像喂给模型,拿到预测结果了:

import tensorflow as tf

# 你已经写了的加载模型函数
def load_graph(frozen_graph_filename):
    with tf.gfile.GFile(frozen_graph_filename, "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_graph('model.pb')

# 替换成你刚才找到的输入/输出张量名称
input_tensor = graph.get_tensor_by_name("input_1:0")
output_tensor = graph.get_tensor_by_name("predictions/Softmax:0")

# 启动会话运行推理
with tf.Session(graph=graph) as sess:
    # 喂入图像数据,得到预测结果
    pred_probs = sess.run(output_tensor, feed_dict={input_tensor: img})

# 处理结果:pred_probs是一个二维数组,shape是(1, num_classes)
# 取概率最大的类别索引和对应的概率值
pred_class_idx = np.argmax(pred_probs[0])
pred_class_prob = pred_probs[0][pred_class_idx]

print(f"预测类别索引:{pred_class_idx},对应概率:{pred_class_prob:.4f}")
几个关键注意事项
  • 节点名称一定要加:0:因为每个节点可能对应多个张量,:0表示取该节点的第一个输出张量,这是TF的默认规则
  • 预处理必须和训练一致:比如训练时用的是tf.image.resize而不是cv2的resize,或者归一化到[-1,1](比如(img/127.5)-1),那推理时必须完全一样
  • 如果用的是TensorFlow 2.x,记得用兼容模式:把tf.Session换成tf.compat.v1.Session,并且在开头加tf.compat.v1.disable_eager_execution()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 03:32:25