TensorFlow+OpenCV推理多线程的可行性与实现方案咨询
推理阶段多线程的意义与优化方案
绝对有意义!尤其是在视频实时检测这种IO和计算交织的场景里,单线程串行处理会浪费大量等待时间——读帧时推理闲置,推理时读帧等待,完全没把硬件资源用满。多线程能把「读帧」(IO密集)和「推理+可视化」(计算密集)拆成两个并行任务,直接提升整体FPS。我来给你拆解下怎么优化你的代码:
核心实现思路
用线程安全的队列做帧的中转,拆分两个独立线程:
- 线程1:专门负责从视频流读取帧,存入队列(专注IO任务)
- 线程2:从队列取出帧,执行推理和可视化,最后显示(专注计算任务)
这样两个任务互不阻塞,彻底解决单线程下的等待浪费问题。
优化后的代码
我基于你的原代码修改,关键改动都加了注释:
#!/usr/bin/env python2 # -*- coding: utf-8 -*- """ Optimized with Multi-Threading for Real-Time Object Detection @author: GustavZ (modified by Stack Overflow contributor) """ import numpy as np import os import six.moves.urllib as urllib import tarfile import tensorflow as tf import cv2 import queue import threading from threading import Event # Protobuf Compilation (once necessary) os.system('protoc object_detection/protos/*.proto --python_out=.') from object_detection.utils import label_map_util from object_detection.utils import visualization_utils as vis_util from stuff.helper import FPS2 # -------------------------- 新增:线程与队列配置 -------------------------- FRAME_QUEUE_SIZE = 10 # 限制队列大小,防止内存溢出 stop_event = Event() # 用于通知线程停止的信号 frame_queue = queue.Queue(maxsize=FRAME_QUEUE_SIZE) # Define Video Input video_input = 0 width = 640 height = 480 fps_interval = 3 # Model preparation MODEL_NAME = 'ssd_mobilenet_v1_coco_2017_11_17' MODEL_FILE = MODEL_NAME + '.tar.gz' DOWNLOAD_BASE = 'http://download.tensorflow.org/models/object_detection/' PATH_TO_CKPT = 'models/' + MODEL_NAME + '/frozen_inference_graph.pb' LABEL_MAP = 'mscoco_label_map.pbtxt' PATH_TO_LABELS = 'object_detection/data/' + LABEL_MAP NUM_CLASSES = 90 # Download Model if not os.path.isfile(PATH_TO_CKPT): print('Model not found. Downloading it now.') 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()) os.remove(MODEL_FILE) # 修正原代码路径问题 else: print('Model found. Proceed.') # Load a (frozen) Tensorflow model into memory. 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='') # Loading label map 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) # -------------------------- 新增:读帧线程函数 -------------------------- def frame_reader(): video_stream = cv2.VideoCapture(video_input) video_stream.set(cv2.CAP_PROP_FRAME_WIDTH, width) video_stream.set(cv2.CAP_PROP_FRAME_HEIGHT, height) while not stop_event.is_set() and video_stream.isOpened(): ret_val, image_np = video_stream.read() if not ret_val: break # 队列满时自动阻塞,避免内存占用过高 if not frame_queue.full(): frame_queue.put(image_np) video_stream.release() print("[INFO] Frame reader thread stopped.") # 启动读帧线程 reader_thread = threading.Thread(target=frame_reader) reader_thread.start() # Detection print ("Press 'q' to Exit") with detection_graph.as_default(): with tf.Session(graph=detection_graph) as sess: # 定义模型输入输出张量 image_tensor = detection_graph.get_tensor_by_name('image_tensor:0') detection_boxes = detection_graph.get_tensor_by_name('detection_boxes:0') detection_scores = detection_graph.get_tensor_by_name('detection_scores:0') detection_classes = detection_graph.get_tensor_by_name('detection_classes:0') num_detections = detection_graph.get_tensor_by_name('num_detections:0') # FPS计算 fps = FPS2(fps_interval).start() while not stop_event.is_set(): try: # 从队列取帧,超时等待避免无限阻塞 image_np = frame_queue.get(timeout=1) except queue.Empty: continue # 扩展维度适配模型输入要求 image_np_expanded = np.expand_dims(image_np, axis=0) # 执行推理 (boxes, scores, classes, num) = sess.run( [detection_boxes, detection_scores, detection_classes, num_detections], feed_dict={image_tensor: image_np_expanded}) # 可视化结果(降低线宽减少开销) vis_util.visualize_boxes_and_labels_on_image_array( image_np, np.squeeze(boxes), np.squeeze(classes).astype(np.int32), np.squeeze(scores), category_index, use_normalized_coordinates=True, line_thickness=4) cv2.imshow('object_detection', image_np) # 退出逻辑 if cv2.waitKey(1) & 0xFF == ord('q'): stop_event.set() break fps.update() frame_queue.task_done() # 标记帧处理完成 # 清理资源 stop_event.set() reader_thread.join() # 等待读帧线程安全退出 cv2.destroyAllWindows() fps.stop() print('[INFO] elapsed time (total): {:.2f}'.format(fps.elapsed())) print('[INFO] approx. FPS: {:.2f}'.format(fps.fps()))
额外FPS提升建议
除了多线程,这些方法也能进一步优化性能:
- 换轻量模型:比如
ssd_mobilenet_v2_coco或efficientdet-lite0,推理速度更快且精度损失极小 - TensorRT优化:如果用NVIDIA GPU,把冻结图转成TensorRT格式,能大幅加速推理
- 可视化精简:只显示置信度≥0.5的框,或者简化可视化样式
- Session多线程配置:给Session添加
config=tf.ConfigProto(intra_op_parallelism_threads=4, inter_op_parallelism_threads=4),让TensorFlow内部启用多线程计算
内容的提问来源于stack exchange,提问作者gustavz
相关产品推荐
相关产品推荐

