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

树莓派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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 15:42:02