基于OpenCV、TensorFlow和Python的人体检测:仅识别80%置信度以上行人
解决仅检测置信度80%及以上行人的问题
嘿,我来帮你搞定这个问题!你当前的代码会把模型检测到的所有COCO类别物体都显示出来,而且没有做置信度过滤,所以才会出现球、球拍这些无关物体。咱们只需要做两个关键修改,就能精准只显示你要的行人:
关键修改逻辑
首先要明确两个核心信息:
- COCO标准数据集里,行人(person)的类别ID是1
- 你需要保留的置信度阈值是0.8(对应80%)
我们要从模型输出的所有检测结果里,只筛选出同时满足这两个条件的数据,再传给可视化函数。
修改后的完整代码
import numpy as np import os import six.moves.urllib as urllib import sys import tarfile import tensorflow as tf import zipfile from collections import defaultdict from io import StringIO from matplotlib import pyplot as plt from PIL import Image from utils import label_map_util from utils import visualization_utils as vis_util MODEL_NAME = 'ssd_mobilenet_v1_coco_11_06_2017' MODEL_FILE = MODEL_NAME + '.tar.gz' DOWNLOAD_BASE = 'http://download.tensorflow.org/models/object_detection/' PATH_TO_CKPT = MODEL_NAME + '/frozen_inference_graph.pb' PATH_TO_LABELS = os.path.join('data', 'mscoco_label_map.pbtxt') NUM_CLASSES = 90 if not os.path.exists(MODEL_NAME + '/frozen_inference_graph.pb'): print ('Downloading the model') opener = urllib.request.URLopener() opener.retrieve(DOWNLOAD_BASE + MODEL_FILE, MODEL_FILE) tar_file = tarfile.open(MODEL_FILE) for file in tar_file.getmembers(): file_name = os.path.basename(file.name) if 'frozen_inference_graph.pb' in file_name: tar_file.extract(file, os.getcwd()) print ('Download complete') else: print ('Model already exists') detection_graph = tf.Graph() with detection_graph.as_default(): od_graph_def = tf.GraphDef() with tf.gfile.GFile(PATH_TO_CKPT, 'rb') as fid: serialized_graph = fid.read() od_graph_def.ParseFromString(serialized_graph) tf.import_graph_def(od_graph_def, name='') label_map = label_map_util.load_labelmap(PATH_TO_LABELS) categories = label_map_util.convert_label_map_to_categories(label_map, max_num_classes=NUM_CLASSES, use_display_name=True) category_index = label_map_util.create_category_index(categories) import cv2 cap = cv2.VideoCapture(1) with detection_graph.as_default(): with tf.Session(graph=detection_graph) as sess: ret = True while (ret): ret,image_np = cap.read() image_np_expanded = np.expand_dims(image_np, axis=0) image_tensor = detection_graph.get_tensor_by_name('image_tensor:0') boxes = detection_graph.get_tensor_by_name('detection_boxes:0') scores = detection_graph.get_tensor_by_name('detection_scores:0') classes = detection_graph.get_tensor_by_name('detection_classes:0') num_detections = detection_graph.get_tensor_by_name('num_detections:0') # 获取模型原始检测结果 (boxes, scores, classes, num_detections) = sess.run( [boxes, scores, classes, num_detections], feed_dict={image_tensor: image_np_expanded}) # -------------------------- 新增过滤逻辑 -------------------------- # 把二维结果转成一维,方便筛选操作 boxes_squeeze = np.squeeze(boxes) scores_squeeze = np.squeeze(scores) classes_squeeze = np.squeeze(classes).astype(np.int32) # 生成筛选掩码:仅保留类别为行人(ID=1)且置信度≥0.8的结果 valid_mask = (classes_squeeze == 1) & (scores_squeeze >= 0.8) filtered_boxes = boxes_squeeze[valid_mask] filtered_scores = scores_squeeze[valid_mask] filtered_classes = classes_squeeze[valid_mask] # ----------------------------------------------------------------- # 传入过滤后的结果进行可视化,替换原来的全量结果 vis_util.visualize_boxes_and_labels_on_image_array( image_np, filtered_boxes, filtered_classes, filtered_scores, category_index, use_normalized_coordinates=True, line_thickness=8) cv2.imshow('image',cv2.resize(image_np,(1280,960))) if cv2.waitKey(27) & 0xFF == ord('q'): cv2.destroyAllWindows() cap.release() break
额外说明
- 如果你想调整置信度要求,直接修改
scores_squeeze >= 0.8里的数值即可,比如改成0.9就是只保留90%置信度以上的行人 - 确认你的
mscoco_label_map.pbtxt里行人的ID确实是1,COCO标准数据集里这个是固定的,一般不会有问题 - 这里的过滤是先筛选再可视化,比只在可视化函数里设置置信度阈值更精准——后者还是会显示其他类别里高置信度的物体,而我们的逻辑能确保只保留行人
内容的提问来源于stack exchange,提问作者Amrit Das
相关产品推荐
相关产品推荐

