使用TensorRT加速TensorFlow模型推理时返回结果为空,求技术排查
解决TensorRT加速TensorFlow SSD模型时返回boxes和scores为None的问题
嘿,我看了你的代码,马上发现了几个明显的问题,这应该就是导致输出全为None的原因:
1. 漏掉了获取scores必需的输出节点
你想要拿到boxes和scores,但不管是创建TensorRT推理图还是导入图的时候,都只指定了detection_boxes和detection_classes,完全没把detection_scores这个关键节点加进去!模型本身是有这个输出的,你不把它加入输出列表,自然拿不到scores的值。
2. 输出节点名称拼写错误
在tf.import_graph_def的return_elements参数里,你写的是"detecction_classes"——注意这里多了一个c(正确的应该是detection_classes),拼写错误导致TensorFlow找不到对应的节点,返回的自然就是None了。
修正后的完整代码
我把所有问题都修正了,还加了注释标注修改点:
import tensorflow as tf import tensorflow.contrib.tensorrt as trt import os from PIL import Image import numpy as np with tf.Session() as sess: with tf.gfile.GFile('model/ssd_inceptionv2.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) # 👉 修正:添加detection_scores到输出列表,这样才能拿到置信度分数 trt_graph = trt.create_inference_graph( input_graph_def=graph_def, outputs=["detection_boxes", "detection_classes", "detection_scores"]) # 👉 修正1:把拼写错误的detecction_classes改成正确的detection_classes # 👉 修正2:添加detection_scores到返回节点列表 output_node = tf.import_graph_def( trt_graph, return_elements=["detection_boxes", "detection_classes", "detection_scores"], name='') input = tf.get_default_graph().get_tensor_by_name('image_tensor:0') # 👉 对应调整,把输出节点分别赋值给三个变量 boxes_tensor, classes_tensor, scores_tensor = output_node for root, dirs, files in os.walk('dataset/test'): for f in files: if not f.endswith('.jpg'): continue print(f) image = Image.open(os.path.join(root, f)) print(image) image = image.resize((300, 300), Image.ANTIALIAS) im_width, im_height = image.size image = np.array(image.getdata()).reshape((im_height, im_width, 3)).astype(np.uint8) image = np.expand_dims(image, axis=0) # 👉 修正:同时运行三个输出节点,对应我们指定的三个变量 boxes, classes, scores = sess.run([boxes_tensor, classes_tensor, scores_tensor], feed_dict={input: image}) boxes, scores = np.squeeze(boxes), np.squeeze(scores) print(boxes, scores) for i, s in enumerate(scores): if s > 0.5: print(boxes[i], s)
额外小提示
- 你可以用Netron工具打开你的pb模型文件,确认里面的输出节点名称是不是和代码里的一致,避免再次出现拼写错误。
- 看你的代码应该是用的TensorFlow 1.x,要是以后升级到TF2.x,记得
tensorflow.contrib.tensorrt已经移到tf.experimental.tensorrt了。
内容的提问来源于stack exchange,提问作者tidy
相关产品推荐
相关产品推荐

