TensorFlow目标检测推理加速求助:多进程报错及性能优化
解决TensorFlow目标检测推理慢+多进程报错问题
针对你遇到的单帧推理耗时久、资源利用率低,以及多进程运行报错的问题,我来一步步给你解决方案:
一、先解决多进程的报错问题
你的报错主要来自两个原因:
can't pickle _thread.RLock objects:TensorFlow的Graph和Session对象无法被序列化,主进程的模型资源传到子进程时会失败;CUDNN_STATUS_ALLOC_FAILED:多进程同时抢占GPU内存,导致内存分配失败。
修复后的多进程代码
核心思路是让每个子进程独立加载模型,避免序列化共享资源,同时配置GPU内存按需分配:
import numpy as np import os import tensorflow as tf import multiprocessing import time import cv2 # 用OpenCV替代PIL,提速预处理/后处理 sys.path.append("..") from object_detection.utils import ops as utils_ops from utils import label_map_util from utils import visualization_utils as vis_util # 全局变量:每个子进程会初始化自己的模型资源 local_graph = None local_sess = None category_index = None # 模型和路径配置 MODEL_NAME = 'ssd_mobilenet_v0_walrus_13_12_2019' PATH_TO_FROZEN_GRAPH = MODEL_NAME + '/frozen_inference_graph.pb' PATH_TO_LABELS = os.path.join('data', 'ssd_mobilenet_v0_walrus_13_12_2019_label_map.pbtxt') NUM_CLASSES = 5 PATH_TO_TEST_IMAGES_DIR = r'C:\temp\frames' OUTPUT_DIR = r'C:\temp\detected_frames' def init_worker(): """每个子进程启动时初始化模型和Session""" global local_graph, local_sess, category_index # 配置GPU内存按需分配,避免多进程抢占 config = tf.ConfigProto() config.gpu_options.allow_growth = True # 加载模型 local_graph = tf.Graph() with local_graph.as_default(): od_graph_def = tf.GraphDef() with tf.io.gfile.GFile(PATH_TO_FROZEN_GRAPH, 'rb') as fid: serialized_graph = fid.read() od_graph_def.ParseFromString(serialized_graph) tf.import_graph_def(od_graph_def, name='') local_sess = tf.Session(graph=local_graph, config=config) # 加载标签映射 category_index = label_map_util.create_category_index_from_labelmap(PATH_TO_LABELS, use_display_name=True) def run_inference_for_single_image(image): """使用子进程本地的模型进行推理""" global local_graph, local_sess with local_graph.as_default(): ops = local_graph.get_operations() all_tensor_names = {output.name for op in ops for output in op.outputs} tensor_dict = {} for key in ['num_detections', 'detection_boxes', 'detection_scores', 'detection_classes', 'detection_masks']: tensor_name = key + ':0' if tensor_name in all_tensor_names: tensor_dict[key] = local_graph.get_tensor_by_name(tensor_name) # 处理mask分支(如果存在) if 'detection_masks' in tensor_dict: detection_boxes = tf.squeeze(tensor_dict['detection_boxes'], [0]) detection_masks = tf.squeeze(tensor_dict['detection_masks'], [0]) real_num_detection = tf.cast(tensor_dict['num_detections'][0], tf.int32) detection_boxes = tf.slice(detection_boxes, [0, 0], [real_num_detection, -1]) detection_masks = tf.slice(detection_masks, [0, 0, 0], [real_num_detection, -1, -1]) detection_masks_reframed = utils_ops.reframe_box_masks_to_image_masks( detection_masks, detection_boxes, image.shape[1], image.shape[2]) detection_masks_reframed = tf.cast(tf.greater(detection_masks_reframed, 0.5), tf.uint8) tensor_dict['detection_masks'] = tf.expand_dims(detection_masks_reframed, 0) image_tensor = local_graph.get_tensor_by_name('image_tensor:0') output_dict = local_sess.run(tensor_dict, feed_dict={image_tensor: image}) # 整理输出格式 output_dict['num_detections'] = int(output_dict['num_detections'][0]) output_dict['detection_classes'] = output_dict['detection_classes'][0].astype(np.int64) output_dict['detection_boxes'] = output_dict['detection_boxes'][0] output_dict['detection_scores'] = output_dict['detection_scores'][0] if 'detection_masks' in output_dict: output_dict['detection_masks'] = output_dict['detection_masks'][0] return output_dict def detector(image_path): # 用OpenCV加载图片,比PIL快 image_np = cv2.imread(image_path) image_np = cv2.cvtColor(image_np, cv2.COLOR_BGR2RGB) # 转为RGB格式 image_np_expanded = np.expand_dims(image_np, axis=0) # 推理 output_dict = run_inference_for_single_image(image_np_expanded) # 可视化(如果不需要可以注释掉,节省时间) vis_util.visualize_boxes_and_labels_on_image_array( image_np, output_dict['detection_boxes'], output_dict['detection_classes'], output_dict['detection_scores'], category_index, instance_masks=output_dict.get('detection_masks'), use_normalized_coordinates=True, line_thickness=8) # 保存结果 img_name = os.path.basename(image_path) img_save_path = os.path.join(OUTPUT_DIR, img_name) cv2.imwrite(img_save_path, cv2.cvtColor(image_np, cv2.COLOR_RGB2BGR)) # 转回BGR保存 if __name__ == '__main__': os.makedirs(OUTPUT_DIR, exist_ok=True) TEST_IMAGE_PATHS = [ os.path.join(PATH_TO_TEST_IMAGES_DIR, 'image{}.jpg'.format(i)) for i in range(17, 20) ] # 不要用满所有CPU,给GPU留资源 n_cpu = multiprocessing.cpu_count() // 2 start_time = time.time() # 用initializer让每个子进程加载模型 pool = multiprocessing.Pool(processes=n_cpu, initializer=init_worker) pool.map(detector, TEST_IMAGE_PATHS, chunksize=4) # 调大chunksize减少进程通信开销 pool.close() pool.join() elapsed_time = time.time() - start_time print('Time with multiprocessing: ', elapsed_time)
关键修改点说明
- 每个子进程通过
init_worker独立加载模型,避免了主进程资源序列化的问题; - 配置
config.gpu_options.allow_growth = True,让TensorFlow按需分配GPU内存,解决内存抢占报错; - 用OpenCV替代PIL处理图片,提升预处理/后处理速度;
- 调大
chunksize参数,减少进程间任务传递的开销。
二、进一步提升推理速度的方案
解决多进程问题后,还可以通过以下方法进一步优化:
1. 模型量化优化
使用TensorFlow的量化工具将模型转为INT8量化模型,能在精度损失极小的情况下,大幅提升推理速度:
from tensorflow.contrib.quantize import create_eval_graph, quantize_graph # 在加载模型后添加量化节点 with local_graph.as_default(): create_eval_graph() quantized_graph = quantize_graph.create_eval_graph(input_graph=local_graph)
2. 使用TensorRT优化
TensorRT是NVIDIA的推理加速工具,能对模型进行层融合、量化、内核调优等优化,适合GPU推理:
from tensorflow.contrib.tensorrt.python import trt_convert as trt # 转换为TensorRT优化模型 converter = trt.TrtGraphConverter(input_graph_def=od_graph_def, nodes_blacklist=['num_detections', 'detection_boxes', 'detection_scores', 'detection_classes']) trt_graph = converter.convert() # 保存优化后的模型 with tf.io.gfile.GFile('trt_frozen_inference_graph.pb', 'wb') as f: f.write(trt_graph.SerializeToString())
之后加载这个优化后的模型进行推理即可。
3. 批量推理
如果GPU内存足够,可以一次性输入多张图片(比如8张/16张)进行批量推理,利用GPU的并行计算能力提升吞吐量:
def run_inference_for_batch(images_batch): # images_batch形状为[batch_size, height, width, 3] with local_graph.as_default(): image_tensor = local_graph.get_tensor_by_name('image_tensor:0') output_dict = local_sess.run(tensor_dict, feed_dict={image_tensor: images_batch}) # 拆分每个图片的检测结果并返回 # ...
4. 跳过不必要的后处理
如果不需要保存可视化后的图片,可以注释掉vis_util.visualize_boxes_and_labels_on_image_array和图片保存步骤,节省大量时间。
内容的提问来源于stack exchange,提问作者pkz
相关产品推荐
相关产品推荐

