树莓派4项目:TensorFlow转TensorFlow Lite并适配Coral加速器
迁移TensorFlow目标检测代码到TensorFlow Lite并适配Coral USB Accelerator
前置准备
- 将原
ssd_inception_v2_coco_2017_11_17模型转换为TFLite格式,或直接下载官方预编译的INT8量化TFLite模型;若使用Coral USB Accelerator,需通过edgetpu_compiler工具将模型编译为EdgeTPU兼容版本 - 安装依赖:
pip install tflite_runtime opencv-python # 若用Coral,额外安装: pip install edgetpu
代码改造步骤
1. 替换导入模块
移除TensorFlow V1相关代码,改用TFLite运行时库:
import os import cv2 import numpy as np import argparse import sys # 替换原TF导入为TFLite运行时 import tflite_runtime.interpreter as tflite from object_detection.utils import label_map_util from object_detection.utils import visualization_utils as vis_util
2. 重构模型加载逻辑
删除原TensorFlow Graph和Session初始化代码,改为加载TFLite模型(适配Coral加速器):
# 设置相机参数 IM_WIDTH = 640 IM_HEIGHT = 480 # 解析命令行参数 parser = argparse.ArgumentParser() parser.add_argument('--usbcam', help='使用USB摄像头替代PiCamera', action='store_true') args = parser.parse_args() camera_type = 'usb' if args.usbcam else 'picamera' # 模型路径配置 MODEL_NAME = 'ssd_inception_v2_coco_2017_11_17' CWD_PATH = os.getcwd() # 替换为你的TFLite模型路径,Coral用户使用编译后的_edgetpu.tflite文件 PATH_TO_TFLITE = os.path.join(CWD_PATH, MODEL_NAME, 'detect.tflite') PATH_TO_LABELS = os.path.join(CWD_PATH, 'data', 'mscoco_label_map.pbtxt') NUM_CLASSES = 90 # 加载标签映射(原逻辑复用) 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) # 加载TFLite解释器及Coral加速器Delegate delegate = tflite.load_delegate('libedgetpu.so.1') interpreter = tflite.Interpreter(model_path=PATH_TO_TFLITE, experimental_delegates=[delegate]) # 不使用Coral时,注释上面两行,启用下面一行: # interpreter = tflite.Interpreter(model_path=PATH_TO_TFLITE) interpreter.allocate_tensors() # 获取输入输出张量信息 input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() input_shape = input_details[0]['shape']
3. 修改推理逻辑
将原sess.run替换为TFLite的推理流程,适配模型输入要求:
# 初始化计数器与控制变量(原逻辑复用) frame_rate_calc = 1 freq = cv2.getTickFrequency() font = cv2.FONT_HERSHEY_SIMPLEX TL_inside = (int(IM_WIDTH*0.016),int(IM_HEIGHT*0.021)) BR_inside = (int(IM_WIDTH*0.323),int(IM_HEIGHT*0.979)) TL_outside = (int(IM_WIDTH*0.333),int(IM_HEIGHT*0.021)) BR_outside = (int(IM_WIDTH*0.673),int(IM_HEIGHT*0.979)) TL_right = (int(IM_WIDTH*0.683),int(IM_HEIGHT*0.021)) BR_right = (int(IM_WIDTH*0.986),int(IM_HEIGHT*0.979)) detected_inside = False detected_outside = False detected_right = False inside_counter = 0 outside_counter = 0 right_counter = 0 pause = 0 pause_counter = 0 def pet_detector(frame): global detected_inside, detected_outside, detected_right global inside_counter, outside_counter, right_counter global pause, pause_counter # 图像预处理:转换RGB+调整尺寸+适配模型数据类型 frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) frame_resized = cv2.resize(frame_rgb, (input_shape[2], input_shape[1])) input_data = np.expand_dims(frame_resized, axis=0) # 处理量化模型的输入转换 if input_details[0]['dtype'] == np.uint8: input_scale, input_zero_point = input_details[0]['quantization'] input_data = (input_data / input_scale + input_zero_point).astype(np.uint8) # 设置输入张量并执行推理 interpreter.set_tensor(input_details[0]['index'], input_data) interpreter.invoke() # 获取推理结果 boxes = interpreter.get_tensor(output_details[0]['index']) classes = interpreter.get_tensor(output_details[1]['index']) scores = interpreter.get_tensor(output_details[2]['index']) num_detections = interpreter.get_tensor(output_details[3]['index']) # 可视化检测结果(原逻辑复用) vis_util.visualize_boxes_and_labels_on_image_array( frame, np.squeeze(boxes), np.squeeze(classes).astype(np.int32), np.squeeze(scores), category_index, use_normalized_coordinates=True, line_thickness=8, min_score_thresh=0.40) # 绘制区域框(原逻辑复用) cv2.rectangle(frame,TL_outside,BR_outside,(255,20,20),3) cv2.putText(frame,"Outside box",(TL_outside[0]+10,TL_outside[1]-10),font,1,(255,20,255),3,cv2.LINE_AA) cv2.rectangle(frame,TL_inside,BR_inside,(20,20,255),3) cv2.putText(frame,"Inside box",(TL_inside[0]+10,TL_inside[1]-10),font,1,(20,255,255),3,cv2.LINE_AA) cv2.rectangle(frame,TL_right,BR_right,(20,255,25),3) cv2.putText(frame,"right box",(TL_right[0]+10,TL_right[1]-10),font,1,(20,255,255),3,cv2.LINE_AA) # 目标位置判断与计数器逻辑(原逻辑复用,注意classes索引是否为1-based,和原模型保持一致) if (((int(classes[0][0]) == 1) or (int(classes[0][0]) == 18) or (int(classes[0][0]) == 88)) and (pause == 0)): x = int(((boxes[0][0][1]+boxes[0][0][3])/2)*IM_WIDTH) y = int(((boxes[0][0][0]+boxes[0][0][2])/2)*IM_HEIGHT) cv2.circle(frame,(x,y), 5, (75,13,180), -1) if ((x > TL_inside[0]) and (x < BR_inside[0]) and (y > TL_inside[1]) and (y < BR_inside[1])): inside_counter += 1 if ((x > TL_outside[0]) and (x < BR_outside[0]) and (y > TL_outside[1]) and (y < BR_outside[1])): outside_counter += 1 if ((x > TL_right[0]) and (x < BR_right[0]) and (y > TL_right[1]) and (y < BR_right[1])): right_counter += 1 # 触发检测后的逻辑(原逻辑复用) if inside_counter == 1: detected_inside = True inside_counter = 0 outside_counter = 0 right_counter = 0 pause = 1 if outside_counter == 1: detected_outside = True inside_counter = 0 outside_counter = 0 right_counter = 0 pause = 1 if right_counter == 1: detected_right = True inside_counter = 0 outside_counter = 0 right_counter = 0 pause = 1 # 暂停状态处理(原逻辑复用) if pause == 1: if detected_inside == True: cv2.putText(frame,'Left detected!',(int(IM_WIDTH*0.027),int(IM_HEIGHT-60)),font,3,(0,0,0),7,cv2.LINE_AA) cv2.putText(frame,'Left detected!',(int(IM_WIDTH*0.967),int(IM_HEIGHT-60)),font,3,(95,176,23),5,cv2.LINE_AA) if detected_outside == True: cv2.putText(frame,'Mid detected!',(int(IM_WIDTH*0.027),int(IM_HEIGHT-60)),font,3,(0,0,0),7,cv2.LINE_AA) cv2.putText(frame,'Mid detected!',(int(IM_WIDTH*0.967),int(IM_HEIGHT-60)),font,3,(95,176,23),5,cv2.LINE_AA) if detected_right == True: cv2.putText(frame,'Right detected!',(int(IM_WIDTH*0.027),int(IM_HEIGHT-60)),font,3,(0,0,0),7,cv2.LINE_AA) cv2.putText(frame,'Right detected!',(int(IM_WIDTH*0.967),int(IM_HEIGHT-60)),font,3,(95,176,23),5,cv2.LINE_AA) pause_counter += 1 if pause_counter > 3: pause = 0 pause_counter = 0 detected_inside = False detected_outside = False detected_right = False # 绘制计数器信息(原逻辑复用) cv2.putText(frame,'Detection counter: ' + str(max(inside_counter,outside_counter, right_counter)),(10,100),font,0.5,(255,255,0),1,cv2.LINE_AA) cv2.putText(frame,'Pause counter: ' + str(pause_counter),(10,150),font,0.5,(255,255,0),1,cv2.LINE_AA) return frame
4. 摄像头与主循环(原逻辑复用)
# USB摄像头初始化(原逻辑复用) camera = cv2.VideoCapture(0) ret = camera.set(3,IM_WIDTH) ret = camera.set(4,IM_HEIGHT) while(True): t1 = cv2.getTickCount() ret, frame = camera.read() frame = pet_detector(frame) # FPS计算与显示(原逻辑复用) cv2.putText(frame,"FPS: {0:.2f}".format(frame_rate_calc),(30,50),font,1,(255,255,0),2,cv2.LINE_AA) cv2.imshow('Object detector', frame) t2 = cv2.getTickCount() time1 = (t2-t1)/freq frame_rate_calc = 1/time1 if cv2.waitKey(1) == ord('q'): break camera.release() cv2.destroyAllWindows()
关键优化点
- Coral加速器适配:必须使用EdgeTPU编译的量化模型,640x480分辨率下帧率可提升至15-20fps左右
- 输入数据类型:量化模型使用
uint8输入,避免浮点运算开销,进一步提升性能 - 减少冗余操作:图像尺寸直接匹配模型输入要求,避免额外缩放
内容的提问来源于stack exchange,提问作者Michael Dixon
相关产品推荐
相关产品推荐

