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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 09:04:16